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:
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 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.
ace_jax.calc.gp.GPCalculator ¶
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
¶
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 ¶
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 ¶
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 ¶
(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 ¶
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 ¶
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 ¶
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 ¶
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 ¶
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 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.radial_model.to_analytic ¶
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=" suffix appended
when spec_grid has more than one value, and a "_gap=" 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 ¶
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 ¶
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 ¶
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 ¶
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.write_yace ¶
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 ¶
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 ¶
Whether the installed lammps-jax exports the neighbour-matrix input layout (export_model(max_neighbors=...); lammps-jax 4a7f4fb and later).