Skip to content

Python API

This page gives the public Python interface, in groups by task. All other parts of ace_jax are internal. They can change between releases.

Enable float64 in Python

Fitting and radial learning need float64. The aj command enables float64 itself. In Python, enable float64 before any other code imports JAX:

import jax
jax.config.update("jax_enable_x64", True)

Alternatively, set JAX_ENABLE_X64=1 in the environment. ace-jax never changes the JAX precision setting itself. The JAX default is float32.

import ace_jax as aj
model, meta, arrays = aj.load("model.npz")      # an ACE model, its metadata, the raw arrays

Loading and evaluating

ace_jax.load

load(
    path,
    dtype=jnp.float64,
    a2b_sparse="auto",
    edge_a_kind="gather",
    fold=True,
)

Load a model. Caller controls dtype; nothing here touches jax.config, so f64 requires the caller to have enabled x64 first.

a2b_sparse selects how A2B is held and contracted. "auto" (default): sparse triplets only, no dense A2B, when at most A2B_SPARSE_MAX_DENSITY of it is nonzero (every ACE coupling: 0.025-1.9% occupied), with B and dB/dA by sorted segment_sums (with_a2b_sparse). A file that stores a dense A2B is converted. True / False force the sparse / dense form. The dense form at 5456 basis functions is 2.6 GB, and its Jacobian costs 2 n_B n_AA n_A flops per node.

edge_a_kind selects how the A-basis product is formed: "gather" (default, today's behaviour) or "matmul", an algebraically identical one-hot form whose adjoint is a matmul rather than a scatter. Which is faster depends on the backend, the dtype and the edge-buffer length -- see calibrate_edge_a, and do not guess from the platform.

fold (default True) folds the linear readout through A2B (see fold_readout); pass False to keep the B-materialising path, e.g. to measure the fold by difference. Descriptors are unaffected either way.

ace_jax.calc.point.ACECalculator

ACECalculator(
    model,
    meta=None,
    cutoff=None,
    dtype=None,
    edge_a_kind="auto",
    layout="auto",
    skin=1.0,
    lean=True,
    spline_tol=AUTO,
    spline_intervals=None,
    radial_table=None,
    posterior=None,
    forces_std_every_call=False,
    energy_reference="absolute",
    shape_path="rows",
    shape_tau=1.0,
    shape_rank=None,
    **kw,
)

Bases: Calculator

ACECalculator("si_fitted.npz") is the intended form: cutoff, species and dtype all come from the file. A pre-loaded (model, meta) pair is still accepted, which is what the validation tests use.

edge_a_kind picks the A-basis form (see EdgeSiteModel.edge_a): "gather", "matmul", or "auto" (default), which calibrates both on the actual neighbour list once per power-of-two edge-count bucket and keeps the faster, using the gather below AUTO_MIN_EDGES edges. It does not apply to a pool-first model (PACE, uses_edge_a False), which is used as given and reports last_edge_a_kind None.

layout picks the neighbour layout: "sparse" (edge list), "dense" (padded (n, K) per-node blocks; A by a batched outer product), or "auto" (default): dense when the padding is efficient (see MIN_DENSE_FILL), else sparse. edge_a_kind applies to the sparse layout.

skin (A, default 1.0) is the Verlet skin of the dense layout (see calc.skin): the neighbour list is built for cutoff + skin and reused, one compiled step per call, until an atom has moved skin / 2 or the cell, pbc or species change. skin=0 rebuilds the list every call. last_timing["rebuilds"] counts the calls that built a neighbour list, and last_timing["nlist_s"] is this call's build time (0 on reuse).

lean (default True) evaluates energies, forces and stress with ace_jax.eval.model.lean(model): exact to roundoff, with the per-edge work the energy never reads removed (docs/dev/ace-vs-pace-gap.md). spline_tol="auto" (default) first splines a learned analytic tensor radial (radial_learned, set by the radial learner) at 1e-10 (to_spline). That is not roundoff: energies agree with lean=False to up to ~1e-9 relative and forces to up to ~2.3e-8 of max|F| on the benchmark models (docs/dev/learned-radial-splining.md). Other analytic models -- ACEpotentials ace_model exports, built bases -- stay exact; a float spline_tol (e.g. 1e-10) opts them in, None never splines. spline_intervals pins the grid. calc.splined (and last_timing["spline_tol"]) says what was splined, None when nothing was. The spline is cached on the radial's content, and its grid is bucketed, so radial swaps usually keep the compiled step. It is eval_model; model stays the model as given, and descriptors use it. A PACE or unfolded model is evaluated as given either way. Setting calc.model recomputes the lean form on the host (a device-to-host copy of the model's arrays): negligible for MD, but a per-step cost if the model is swapped every step.

radial_table (None, the default: off; True: 4000 intervals; an int: that many) tabulates the evaluation model's radial stage in r (eval.model.with_radial_table, after lean; ACE: R_nl with its transform and envelope, and the pair radial; PACE: g_k), one cubic B-spline per species pair on [0.5 A, rcut], exactly zero beyond each pair's cutoff. A CPU speed-up (1.1-1.4x on one core), and an approximation: at 4000 intervals energies agree to ~1e-11 relative and forces to ~1e-8 of max|F| (tests/test_radial_table.py). Applies with lean=False too (to the model as given). calc.radial_table is the table's info (n_intervals, r_min, r_max, max_rel_err, max_rel_deriv_err), None when off; last_timing["radial_table"] its n_intervals. Cached on the radial's content, so a readout-only calc.model swap does not rebuild it.

posterior (a posterior.npz from fit --uq ard, with model the matching model.npz FILE) adds the per-atom force uncertainty. For a revision-2 (schema-3) posterior, each atom has a 3x3 shape V -- the centred delete-one-cluster (PRESS) jackknife covariance of its force (--ard-variance kappa: phi A^-1 phi^T) -- and a Mondrian group g that selects the calibrated scales: forces_std = lam_rms[g] sqrt(tr V), (N,); forces_cov = lam_rms[g]^2 V, (N, 3, 3); forces_q, the conformal radius at the fit's coverage (q[g] sqrt(tr V / 3) iso, the ellipsoid's largest semi-axis aniso); forces_q_mahal = q[g] (aniso); forces_group = g; and forces_support (on request). Schema-1/2 posteriors serve only the scalar forces_std (kappa times the epistemic std, or lambda times an uncentred configuration-clustered sandwich std). Energy and stress carry no uncertainty. By default the uncertainty is computed only when requested (calc.get_property("forces_std", atoms) or any of the above, which reuses the cached E/F/stress): a design-row rebuild plus an L^2 solve per step is not a silent MD cost. forces_std_every_call=True adds it to every calculation. posterior= requires jax.config.update("jax_enable_x64", True) (RuntimeError otherwise).

shape_path picks how a schema-3 jackknife shape V is evaluated: "rows" (default) builds the whole cell's force design rows (N, 3, L) and contracts them with R; "committee" evaluates the forces of the r-output linear ACE with coefficients D^-1 R (fit.jackknife.committee_shape), equal to roundoff, in memory O(N r) instead of O(N L) and without the edge Jacobian. shape_tau < 1 truncates R to the smallest rank holding that fraction of its sum sigma^2, and shape_rank (an int) caps the rank, on either path -- an approximation whose ranking and coverage loss against rank is measured in bench/defect_uq; the defaults keep R exact.

energy_reference is "absolute" (default: the model's energy, isolated-atom energies E0 included) or "E0": the energy relative to the isolated atoms, sum_i (E_i - E0[z_i]). With "E0" the per-atom constant (~-160 eV for Si) is never added to the site energies, so the total of a large cell stays small and keeps resolution: an energy-based line search (ASE PreconLBFGS's Armijo test) otherwise stops resolving energy decreases once they fall below the ulp of |E| ~ 160 eV x N, at fmax ~ 1e-4..1e-5 eV/A on 10^4-10^5 atoms (docs/dev/energy-sum-results.md). Forces and stress are unchanged. results["e0_offset"] is sum_i E0[z_i] (math.fsum; 0.0 for "absolute"), so the absolute energy is energy + e0_offset.

Memory: forces_std builds the force design rows of the WHOLE cell, about N3L*8 bytes (N atoms padded, L = (n_B + n_pair) * NZ columns) -- 7 GB for N = 100k at L = 3k -- on the JAX device, besides the posterior's L^2 factor. The edge Jacobian is node-chunked (linear_rows_chunked), the rows themselves are not: size cells to fit them. The rows come from the FULL model (never the energy-only eval_model), whatever lean is.

splined property

splined

What the lean form splined, None when nothing: {"spline_tol", "radials", "n_intervals"} (eval.model.splining).

eval_model property

eval_model

The model energies, forces and stress are evaluated with: lean(model) (or model itself with lean=False). Energy only: no descriptors.

model property writable

model

skin property writable

skin

ace_jax.calc.gp.GPCalculator

GPCalculator(fitted, meta, deriv_dtc=True, **kw)

Bases: Calculator

deriv_dtc=False drops the derivative-DTC term from forces_std (SoR-only force variance): cheaper, and needed when a big cell's (n, K, d, 3) arrays do not fit.

from_file classmethod

from_file(path, **kw)

A calculator from the gp_model.npz that ace-jax fit writes.

ace_jax.site_descriptors

site_descriptors(
    model,
    positions,
    numbers,
    cell=None,
    pbc=False,
    meta=None,
    cutoff=None,
    dtype=None,
    domain=None,
)

Site descriptors, (n_atoms, (n_B + n_pair) * n_species).

Parity target is ACEpotentials.site_descriptors, which is marked in the ACEpotentials source as "RETIRING THIS FOR NOW BECAUSE IT IS HIGHLY INEFFICIENT" because it recomputes per site. This takes the whole batch from one forward pass -- the same pass the energy uses -- so the port is genuinely faster here, not merely equivalent.

domain restricts the returned rows (as ACEpotentials does); the forward pass still covers the whole structure, since a site's descriptor needs its neighbours regardless.

ace_jax.highest_precision

highest_precision()

Force true f32/f64 matmuls. On Ampere+ the TF32 default silently costs ~400x accuracy (1.17e-3 vs 2.93e-6 on this descriptor).

Building a basis

ace_jax.basis.model.BasisSpec dataclass

BasisSpec(
    order,
    max_degree,
    elements=None,
    wL=1.5,
    rcut=None,
    rin=0.0,
    maxl=None,
    d_max=None,
    reduction="pca",
    radial_mode="onehot",
    pair_mode="onehot",
    embedding=None,
    no_gamma=False,
    no_coupling_cache=False,
    coupling_cache_dir=None,
)

A basis definition: the aj basis flags (and the basis: block of a fit.yaml). elements=None means "the species of the training data".

ace_jax.basis.model.build_basis

build_basis(spec, *, seed=0)

The one basis builder (aj basis, aj fit, Python): a Basis from a BasisSpec. Categorical (build_model) unless spec.embedding is set (build_embedding_model).

ace_jax.basis.model.Basis

Bases: NamedTuple

An authored model plus everything needed to reproduce or package it.

eval_pair

eval_pair()

(model, meta) ready for direct evaluation -- the in-memory hand-off.

The four branch-selector meta keys are derived from the tree being handed over, exactly as export.save_npz does for the file path: the loader-side consumers (eval.io.load, and anything that reads meta to route branches) must never see authoring defaults that a patched tree has outgrown -- JAX clamps out-of-range branch indices silently, so a mismatch fails numerically, not loudly. The meta is deep-copied; the caller may mutate it freely. Structural checks guard the derivable keys (_resolve consumers read rcut/elements/n_* from meta).

ACECalculator(*auth.eval_pair()) and site_descriptors(auth.model, ..., meta=meta) evaluate the in-memory tree directly -- no npz round-trip; save_npz remains for ACEfit-fitted interchange and shell hand-off only.

ace_jax.basis.model.basis_r0

basis_r0(meta)

Mean radial length scale the basis was built with (per-pair table or a scalar): the default GP hyperprior centre for a fit of this basis.

ace_jax.basis.export.save_npz

save_npz(path, auth)

Write auth (a Basis) to path in the eval-io schema.

The meta written here is derived from the model being saved, not from the authoring defaults: any tree patched onto a different branch (e.g. the fixture-injection round trip) must save as what it now is, or the loader silently routes branch arrays to placeholders. The derivation and the structural checks live in Basis.eval_pair -- the file path and the in-memory hand-off share one source of truth.

ace_jax.basis.coupling.BasisUnavailable

Bases: RuntimeError

Building a new basis needs the compiled coupling library, which is not installed (unsupported platform) or has no compiled payload (dev build).

Fitting

The pipeline that aj fit uses.

FitConfig defaults are different from the command line

The FitConfig defaults are those of the research driver. Some are different from the command-line defaults (for example arm="gp", e0="lsq", map_steps=150). Set each field that you need explicitly. To get exactly the command-line result, also set predict_stats="recompute".

ace_jax.fit.pipeline.FitConfig dataclass

FitConfig(
    model,
    arm="gp",
    energy_key="energy",
    force_key="forces",
    virial_key="virial",
    stress_key=None,
    ntrain=800,
    ntest=200,
    test_start=None,
    seed=0,
    batch=4,
    batch_pack="auto",
    weights=None,
    factors=None,
    sigma_type=False,
    route=None,
    baseline=None,
    base_npz=None,
    e0="lsq",
    m_per_species=100,
    kernel="cosine",
    bump=True,
    density="none",
    pca_d=128,
    warp="none",
    embedding=None,
    delta_s_floor_q=None,
    fix_rho=None,
    r0=2.5,
    objective="lml",
    lml="device",
    lml_chunk=64,
    lml_solver="qr",
    devices=1,
    opt="lbfgs",
    map_steps=150,
    map_lr=0.02,
    map_restarts=1,
    map_polish="auto",
    strict=False,
    init=None,
    noise="per-quantity",
    rungs=("map",),
    solver="evidence",
    laplace="fd",
    n_draws=64,
    vi_steps=1000,
    nuts_warmup=100,
    nuts_samples=100,
    nuts_chains=1,
    pf_samples=4,
    pf_maxiter=10,
    uq="blr",
    ard_mode="joint",
    ard_variance="sandwich",
    ard_val_frac=0.2,
    ard_cond_max=100000000000000.0,
    ard_laplace=False,
    ard_force_shape="aniso",
    ard_shape_eps=0.001,
    ard_coverage=0.9,
    ard_groups="distortion",
    ard_cluster_size=3.0,
    ard_press="exact",
    ard_shape_tau=1.0,
    ard_transfer="exponent",
    ard_n_min=20,
    ard_support=True,
    ard_support_features="raw",
    ard_support_max_atoms=50000,
    _shape_variant="press",
    _score_source="fit",
    deriv_dtc=True,
    predict_stats="cached",
    predict_train=True,
    pops_posterior="hypercube",
    pops_leverage_pct=0.0,
    pops_ridge="auto",
    pops_ridge_grid=(
        0.01,
        0.001,
        0.0001,
        1e-05,
        1e-06,
        1e-07,
        1e-08,
        1e-09,
        1e-10,
        1e-11,
        1e-12,
        1e-13,
        1e-14,
    ),
    pops_val_frac=0.2,
    pops_env_nf=2000,
    pops_rows="auto",
    learn_radial=False,
    radial_n_q=12,
    radial_steps=40,
    radial_lam_grid=(0.0, 0.01),
    radial_val_frac=0.2,
)

Every option of the fitting pipeline. Defaults are run.py's; the CLI overrides the ones where it differs (see cli.py).

ace_jax.fit.pipeline.load_fit_data

load_fit_data(
    cfg,
    *,
    data=None,
    train=None,
    test=None,
    ood=None,
    log=print,
)

Configs, baselines, E0 and datasets. Either data (one file, split with split_configs) or train (+ optional test; defaults to train) files. Each may be a path or a list of ase.Atoms (see fit.data.load_configs).

ace_jax.fit.pipeline.fit

fit(cfg, data, log=print, on_stage=None)

Run the pipeline. on_stage(name, payload), if given, is called as each expensive stage finishes ("data" -> FitData, "radial" -> RadialResult when cfg.learn_radial, "map" -> MapFit, "rungs" -> Rungs, "ard" -> ARDResult and "model" -> the ARD-mean model.npz arrays when uq == "ard"), so a driver can write those results before a later stage (e.g. POPS or prediction running out of memory) can lose them.

ace_jax.fit.pipeline.save_model

save_model(res, out, n_draws=1, log=print)

Write the fitted model into directory out: model.npz (linear arm) or gp_model.npz (GP arm). Only the model: the fitted hyperparameters are res.theta (and res.map.log_evidence), and write_outputs(res, out) also writes them (theta_map.json) with the metrics. Returns the path, or None when the fit cannot be represented as a model file (a dimer baseline is added back outside the model; a .yace input has no npz schema to write into).

ace_jax.fit.pipeline.write_outputs

write_outputs(
    res,
    out,
    layout=("run",),
    argv=None,
    save_model=True,
    model_draws=1,
    log=print,
)

Write the run artefacts; with save_model, also the fitted model file (model.npz for the linear arm, gp_model.npz for the GP arm; see export.py).

ace_jax.fit.data.load_configs

load_configs(
    source,
    energy_key="energy",
    force_key="forces",
    virial_key="virial",
    stress_key=None,
    weights=None,
    weight_key="config_type",
    factors=None,
)

factors: an optional list of weights.WeightFactors (see ace_jax.fit.weights), composed via compose() into a single fn(meta, quantity) -> float that produces w_E/w_F/w_V. factors=None (the default) reconstructs the CLASSIC weighting -- structural 1/sqrt(n) on E,V (1 on F), times the resolved per-config-type {E,F,V} dict from weights/weight_key -- as [Structural(), ConfigType(...)], so every existing caller (none of which pass factors) gets bit-identical w_E/w_F/w_V to before this was added. weights/weight_key ALSO still drive type_idx (Task 6's per-type sigma routing) independently of factors -- that wiring is untouched.

Read with libAtoms extxyz (fit.xyz.read_extxyz): every label comes back under the name it was written with, including energy / forces, which ase.io hides in a calculator. source may instead be a list of ase.Atoms (labels in info / arrays, or a SinglePointCalculator's results). stress_key: a periodic config with no virial label takes virial = -stress * volume from it.

Learned radials

See Learn the radial basis for how these fit together.

ace_jax.fit.pipeline.problem.build_problem

build_problem(cfg, d)

ace_jax.fit.radial_model.to_analytic

to_analytic(model, n_q, n_x=2001)

Convert a splined tensor radial to the analytic branch. Each pair's spline S(x) (the envelope-free part) is sampled on a uniform x-grid, envS is projected onto envP_q(x) (radial_init.from_table), and the branch is swapped. An analytic model is only widened. Returns (model, relres) with relres (NZ, NZ, n_rnl) the per-radial relative projection residual.

ace_jax.fit.radial_learn.fit_radial

fit_radial(
    prob,
    ds_fit,
    ds_val,
    W0,
    *,
    lam_grid=(0.0, 0.001, 0.01, 0.1),
    spec_grid=(0.0,),
    gap_grid=(0.0,),
    theta0=None,
    map_steps=300,
    log=None,
    checkpoint=None,
    noise="per-quantity",
    **learn_kw,
)

learn_radial on ds_fit once per (roughness, spectral, data-gap weight) triple in lam_grid x spec_grid x gap_grid, then keep the best of {init, learned per triple} on the disjoint ds_val (ties -> init).

Candidate labels are "learned_lam=" with a "spec=" suffix appended when spec_grid has more than one value, and a "_gap=" suffix appended when gap_grid has more than one value (each independently; backward compatible with plain "learned_lam=" when both grids are single-valued, an implicit spec=0/gap=0); checkpoints use the same scheme without the "learned" prefix (lam_label = f"{l:g}", plus "_spec=..."/"_gap=..." under the same conditions).

Every candidate, init included, is scored by ONE procedure: a_fit = theta_map_linear(ds_fit, W, map_steps, init=a0), readout = posterior mean on ds_fit at a_fit, score on ds_val with sigma from a0 (a0 = theta0, or the theta-MAP at the normalised init). So with steps=0 (learned == init) the scores are identical and the tie goes to init. Also recorded: info["scores_at_a0"] (every candidate's readout at the common a0, no theta optimisation involved), info["map_diag"] (per-candidate MAP convergence diagnostic, see theta_map_linear), info["theta_fit"] and info["readout"] (the SELECTED candidate's M = 0 readout on ds_fit at its a_fit, length len_basis, species-blocked as fit.rows._place). Each candidate costs one ds_fit pass plus one ds_val pass. Returns (W_sel, info).

log: optional one-arg callable given one-line progress strings (forwarded to learn_radial, plus fit_radial's own per-candidate and gate-score lines); None (default) is silent and leaves all other behaviour unchanged.

checkpoint: optional callable checkpoint(lam_label, W, run_info), called as soon as each (lam, spec, gap) triple's learn_radial finishes, so an interrupted grid keeps its finished runs (e.g. save_result without src_npz). Q, D2, U and the relative-lambda/spec/gap reference are computed once here and shared by every candidate (U only when some gap_grid value is > 0, via one extra streaming pass for data_r_range); M > 0 raises.

ace_jax.fit.radial_learn.learn_radial

learn_radial(
    prob,
    ds,
    W0,
    *,
    theta0=None,
    profile=True,
    lam_rough=0.0,
    rough_weights=None,
    lam_spec=0.0,
    spec_p=4.0,
    lam_gap=0.0,
    steps=200,
    reprofile_every=10,
    tol=1e-06,
    patience=3,
    map_steps=300,
    n_prior=None,
    seed=0,
    log=None,
    Q=None,
    D2=None,
    r0=None,
    U=None,
    learn_sigma_e_mult=1.0,
    noise="per-quantity",
)

VarPro-learn the tensor radials of prob.model (analytic branch, M = 0).

learn_sigma_e_mult scales sigma_E inside the radial objective only (after every theta re-MAP): > 1 learns the radials under a force-heavier weighting than the evidence picks. Everything downstream (gate, final linear fit) keeps the plain MAP theta. With a multiplier != 1, r0 (the reference scale of the relative priors) is taken at the scaled theta.

Minimises r(W; theta) + lam * roughness(W) + lam_spec_abs * spectral_penalty(W, W_ref) + lam_gap_abs * gap_penalty(W, W_ref, U) over V with W = normalise(V) (unit empirical norm per radial; rows that are zero in W0 stay zero), by L-BFGS in rounds of reprofile_every steps. With profile=True theta is re-MAP'd on the M = 0 LML after every round (and at the start unless theta0 is given), and L-BFGS restarts because the objective changed. lam_rough is RELATIVE: lam = lam_rough * r(W0) / roughness(W0). lam_spec is also RELATIVE: lam_spec_abs = lam_spec * r0 / n_active, n_active = the number of active (row_active) radials; the spectral prior penalises the departure of the normalised W from W_ref = normalise(W0) (the starting V), weighted per Legendre degree q by spectral_weights(n_q, spec_p) = (1 + q)^spec_p, so it grows with degree -- the analogue of the Gamma smoothness prior's degree weighting, but on the CHANGE rather than the absolute radial. lam_gap is also RELATIVE (relative_lambda_gap): lam_gap_abs = lam_gap * r0 / n_active; the data-gap prior penalises the same departure from W_ref, but measured under U = uniform_gram(prob.model, 0.8 * r_min, rcut) (r_min from data_r_range(ds), rcut = prob.cfg.rcut) -- a uniform-in-r measure, so it costs change equally everywhere on that range, including gaps between coordination shells where the data-weighted gauge Q costs nothing. Stops at steps total or at the first round that ends early (converged, line search, non-finite). Returns (W, info); steps=0 returns normalise(W0).

log: optional one-arg callable given one-line progress strings (start hyperparameters, then a line per round); None (default) is silent and leaves all other behaviour unchanged.

Precomputed pieces (fit_radial passes them so a lambda grid shares them; each is computed here when None): Q = radial_gram(prob.model, ds, n_prior) (n_prior None -> radial_gram's default), D2 = roughness_matrix(prob.model), r0 = r(normalise(W0); theta0), the projected residual at the start -- only valid together with that theta0 -- and U = uniform_gram(...) as above (only computed, an extra streaming pass over ds via data_r_range, when lam_gap > 0; a zeros array of the right shape otherwise, so the default path adds no streaming passes and the traced arg structure of _objective stays fixed). M > 0 (inducing points) is unsupported and raises.

ace_jax.fit.radial_learn.save_result

save_result(
    out_dir,
    W,
    info,
    *,
    src_npz=None,
    model=None,
    readout=None,
)

Write rnl_Wnlq.npy and radial_info.json to out_dir. With src_npz and model (the analytic model W belongs to, e.g. after widen_radial / to_analytic) also write model.npz = src_npz with the learned radial AND the readout fitted for it (readout, default info["readout"] as returned by fit_radial; also saved as readout.npy). The source npz's WB/Wpair belong to the old radials, so writing model.npz without a readout is refused rather than silently stale. info["readout"] is kept out of the JSON (it is len_basis long). model.npz is marked meta "radial_learned" (so lean splines it by default) unless info["selected"] == "init".

ace_jax.fit.radial_model.with_radial

with_radial(model, W, learned=True)

model with its tensor-radial weights replaced by W (same shape), marked radial_learned (the learner's weights; lean's default splines those). learned=False for weights that are not learned.

ace_jax.basis.export.mark_radial_learned

mark_radial_learned(src, dst=None, learned=True)

Set meta_json "radial_learned" in the npz at src (in place, or into dst), leaving every array as is. For learned-radial files written before the flag existed, which load as not learned, so lean's default keeps their radial analytic: mark_radial_learned("model.npz"). (Or pass spline_tol=1e-10 to ACECalculator / export_lammps / lean instead.)

Deployment forms

ace_jax.eval.lean

lean(
    model,
    spline_tol=AUTO,
    spline_intervals=None,
    radial_table=None,
)

The evaluation form of a folded ACEModel: prune_columns, fold_pair and the l-blocked dense A (block_dense). Exact to roundoff in E, F and the virial for a splined model; 1.1-3.3x faster forces on the benchmark models (docs/dev/ace-vs-pace-gap.md section 8).

Splining (splinify.spline_plan, the one decision point): spline_tol="auto" (default) splines an analytic tensor radial only when it was learned (radial_learned), at DEFAULT_SPLINE_TOL = 1e-10; the pair radial is never learned and stays as is. ACEpotentials ace_model exports and built bases are analytic but not learned, so they stay exact. A float spline_tol opts every analytic radial in (e.g. 1e-10 for an old learned-radial file without the flag); None never splines. A splined radial gets the spline gather and, when each R_nl column belongs to one neighbour species (ACE1's pattern, which learned radials keep), the species-compact blocks. That step is an approximation, not roundoff: at 1e-10 the lean energies agree with the full model to up to ~1e-9 relative and forces to up to ~2.3e-8 of max|F| on the benchmark models (docs/dev/learned-radial-splining.md). spline_intervals pins the grid (default: the smallest splinify.BUCKETS bucket meeting tol, so radial swaps keep the compiled step). The conversion is cached on the radial's content (splinify), so a re-lean after a readout-only change does not re-spline. For evaluation and export only: fitting keeps the analytic model, and a UQ variance should come from the full model. splining(model, lean(model), tol) reports what was splined.

A wrapper model (.base and with_base, e.g. FSModel(base, ...)) returns model.with_base(lean_keep_basis(model.base, spline_tol, spline_intervals)): the wrapper reads the basis, so only the basis-preserving transforms apply.

radial_table (None, the default: off; True: splinify.DEFAULT_RADIAL_TABLE = 4000 intervals; an int: that many) applies splinify.radial_table last, via with_radial_table: the whole radial stage -- R_nl with its transform and envelope, and the folded pair radial -- as one cubic B-spline table per species pair in r on [0.5 A, rcut], masked to exact zeros beyond each pair's cutoff (the end cubic extrapolates below 0.5 A). A CPU speed-up (1.1-1.4x on one core), and an approximation: at 4000 intervals energies agree to ~1e-11 relative and forces to ~1e-8 of max|F|. A PACEModel gets its g_k tabulated (and is otherwise returned as given); a wrapper model refuses it.

Energy only (see fold_pair): keep the original for descriptors and fitting. Anything that is not a folded ACEModel (a PACEModel, an unfolded model) is returned as given, as is a model that is already lean. Edit the full model, never this one: a lean model holds the radial twice and its pair channel is the readout, so replacing rnl_coefs, Wnlq, Wpair, WB or ctilde on it changes one layout and not the other (require_full guards the radial helpers). Change the full model and re-apply lean.

ace_jax.eval.to_spline

to_spline(
    model,
    n_intervals=None,
    tol=DEFAULT_SPLINE_TOL,
    deriv_tol=None,
    return_info=False,
    radials=("rnl", "pair"),
)

Convert an ACEModel's analytic radials to the spline branch.

Converts every analytic radial named in radials ("rnl", "pair"), learned or not (an ACEpotentials ace_model export or a built basis too): this is the conversion itself. Whether lean applies it is spline_plan's call (by default only for learned radials).

The tensor radial (rnl_Wnlq, radial_kind "analytic") and the pair radial (pair_Wnlq, pair_radial_kind "analytic") are each tabulated as S(x) = sum_q W P_q(x) on a uniform grid of n intervals over x in [-1, 1], (x0, h, n + 1) = (-1, 2 / n, n + 1), and interpolated with the cubic B-spline spline_eval evaluates, clamped to the polynomial's own second derivative at the ends. The envelope, the transform and every other array are unchanged.

Grid: n_intervals pins it. None (default) takes the smallest bucket of BUCKETS (quarter-octave steps, <= 19% finer than needed) whose error is <= tol, separately for the tensor and the pair radial; bucketing keeps the table shape, hence the compiled step, across radial swaps.

Error: the largest, over species pairs and the radial columns the A basis reads (aspec_r; all pair columns), of max|spline - exact| / max|exact| on 10 points per interval. For R_nl it includes the envelope (a function of x), for the pair radial it is envelope-free (its envelope is in r and multiplies both alike). tol bounds VALUES. The derivative error, which forces see, is O(h^3) where values are O(h^4): it is measured the same way on d/dx and reported (return_info), and gated too when deriv_tol is given. At 1e-10 the lean energies agree with the full model to up to ~1e-9 relative and forces to up to ~2.3e-8 of max|F| on the benchmark models (docs/dev/learned-radial-splining.md).

tol must be in [TOL_FLOOR, inf); it is floored at 10 eps of the model's dtype (a float32 table cannot do better). Returns (model, max_rel_err), or with return_info=True (model, info): max_rel_err, max_rel_deriv_err, n_intervals {"rnl", "pair"} and the effective tol. A spline or spline_factorised R_nl, and a spline pair radial, are kept as they are. The result is an ordinary full model (float64 tables cast to the model's dtype), trainable as a spline; it agrees with model to tol, not to roundoff. Needs the full model (require_full).

Cached (an LRU of at most CACHE_BYTES, clear_cache()) on the content of each radial's Wnlq, recursion, (R_nl) envelope and measured columns plus n_intervals and the tolerances, so converting a model whose radials are unchanged -- e.g. after a readout-only swap -- costs a hash, not a fit.

PACE .yace models

ace_jax.eval.load_yace

load_yace(path, dtype=jnp.float64, edge_a_kind='gather')

ace_jax.eval.write_yace

write_yace(model, spec, path)

Serialise a PACEModel: numeric fields from its leaves, everything else verbatim from spec.tree (function layout is never regenerated).

LAMMPS export

ace_jax.export.lammps.export_lammps

export_lammps(
    model,
    meta,
    path,
    *,
    max_atoms,
    max_edges=None,
    k_dense=None,
    dtype="float64",
    layout="auto",
    type_elements=None,
    max_owned=None,
    lean=True,
    spline_tol=AUTO,
    spline_intervals=None,
    max_neighbors=None,
    radial_table=None,
)

Write a lammps-jax JSON bundle for model; returns the bundle dict.

Capacities: max_atoms (owned + ghost positions); max_edges (sparse / dense: the packed edge buffer, pairs within rcut); k_dense (dense / matrix: model slots per atom, for the matrix default max_neighbors); max_neighbors (matrix only, required: list slots per row -- LAMMPS copies its rcut + skin list whole, so a k_dense sized for pairs within rcut is too small, and a wider row aborts the run). neighbour_capacity sizes all of them for a structure; its slots="cutoff" gives tight model slots (k_dense < max_neighbors).

type_elements: atomic numbers in LAMMPS type order (type 1 first). LAMMPS hands the model species = type - 1, and a model's own element order (a .yace lists its elements as fitted) need not match; default: the model's.

max_owned: dense / matrix row capacity (owned atoms only, LAMMPS numbers them first). When given, the dense energy function evaluates only rows < it (see make_energy_fn's n_rows) and the matrix has that many rows (lammps-jax's max_owned); always recorded in the bundle, including for the sparse layout, which has no row concept and ignores it otherwise.

layout="auto" picks a dense-family layout when k_dense is given and estimate_a_bytes for one dense block (the row capacity -- max_owned if given, else max_atoms -- capped at BUNDLE_BLOCK_ROWS) fits ace_jax.calc.point.dense_budget_bytes(), else sparse. The dense-family layout is "matrix" when max_neighbors is given, the installed lammps-jax supports it (matrix_supported) and the block plus the matrix's unblocked pre-processing (matrix_prep_bytes) fits; else "dense". A LAMMPS plugin older than the Python package rejects a matrix bundle: rebuild the plugin, or export layout="dense". The bundle records the exporting lammps-jax as ace_jax.lammps_jax.

lean (default True): export ace_jax.eval.model.lean(model, spline_tol, spline_intervals), the evaluation form with the dead per-edge work removed (docs/dev/ace-vs-pace-gap.md). spline_tol="auto" (default) first splines a learned analytic tensor radial (radial_learned) at 1e-10, which is not roundoff: energies agree with lean=False to up to ~1e-9 relative and forces to up to ~2.3e-8 of max|F| on the benchmark models (docs/dev/learned-radial-splining.md). Other analytic models (ACEpotentials ace_model exports, built bases) stay exact; a float spline_tol opts them in, None never splines. Recorded from what lean actually did (looking through a wrapper's .base): ace_jax.lean (False when lean returned the model as given, e.g. PACE or an unfolded model), ace_jax.spline_tol and ace_jax.spline_intervals ({radial: n}), both None when nothing was splined.

radial_table (None, the default: off; True: 4000 intervals; an int: that many): tabulate the exported model's radial stage in r after lean (eval.model.with_radial_table; ACE R_nl and the pair radial, PACE g_k), for every layout. A CPU speed-up and an approximation (at 4000 intervals energies to ~1e-11 relative, forces to ~1e-8 of max|F|); exactly zero beyond each pair's cutoff, so the matrix layout's skin pairs still drop out. Recorded as ace_jax.radial_table (n_intervals, r_min, r_max, max_rel_err, max_rel_deriv_err; None when off), next to ace_jax.spline_tol.

ace_jax.export.lammps.neighbour_capacity

neighbour_capacity(
    atoms,
    rcut,
    skin=1.0,
    slots="skin",
    margin=8,
    owned=1.1,
    list_headroom=0.5,
)

lammps-jax buffer sizes for atoms (a periodic ase.Atoms, the structure the run starts from). Returns max_atoms (owned + ghost shell rcut + skin deep perpendicular to each face -- the face spacings, 1/|reciprocal row|, so sheared cells are covered -- x1.1), max_owned (owned x owned), k_list / k_cut (largest coordination within rcut + skin / rcut), max_neighbors (neighbour-matrix list slots: LAMMPS copies its whole rcut + skin list at every rebuild, and a wider row aborts the run, so max(k_list + margin, ceil((1 + list_headroom) * k_list))), k_dense (model slots; the matrix compacts the in-cutoff pairs into them) and max_edges (max_owned * k_dense).

The list needs more headroom than the model slots: when a structure compresses, the rcut + skin count grows as the rcut count does, from a larger base. list_headroom = 0.5 comes from ONE observed overflow: on the benchmark deck SiGe medium's widest list row grew from 34 to 45 by the first rebuild (k_list + 8 = 42 aborted), while its rcut slots never overflowed (docs/dev/perf-lammps-large-n.md). It is a guess, not a bound; for stable MD list_headroom=0 (with margin >= 8) is cheaper.

slots="skin" (default, always safe between list rebuilds): k_dense = k_list + margin. No atom can gain more neighbours within rcut than its rcut + skin list holds, so this covers structures that compress during the run (the benchmark's random-weight models push Cantor's coordination within rcut from 42 to 50 in 250 steps; docs/dev/perf-lammps-large-n.md).

slots="cutoff": k_dense = k_cut + margin, 1.2-1.4x faster on Cantor (fewer model slots). Safe for stable MD with a fitted model -- a thermalised crystal or liquid whose coordination within rcut stays within margin of the starting structure's. If it does not, an atom overflows its slots and the step's energy and forces are NaN (the run fails loudly; never a truncation).

ace_jax.export.lammps.matrix_supported

matrix_supported()

Whether the installed lammps-jax exports the neighbour-matrix input layout (export_model(max_neighbors=...); lammps-jax 4a7f4fb and later).