Learn the radial basis¶
In a built basis, the radial functions are frozen: aj fit fits only the
linear coefficients. The radial learner also optimises the radial mixing
weights, by variable projection (VarPro). A validation gate then keeps the
learned radials only if they predict better than the initial radials.
The output is a usual model.npz file. Thus all later steps
(aj fit --model, ACECalculator, LAMMPS export) work without changes.
Concepts explains the method.
Tutorial 2 shows all the steps.
Radial learning needs float64. The aj command enables float64 itself.
Radial learning costs much more than a linear fit. Each L-BFGS step does one pass over the fit split, in batches. For example:
- 26 two-atom silicon cells: approximately 2 minutes on a CPU;
- production datasets of a few hundred cells: tens of minutes on a GPU.
On the command line¶
aj fit --order 3 --max-degree 10 \
--train train.xyz --test test.xyz \
--e0 lsq --m-per-species 0 --opt lbfgs \
--learn-radial \
--out fit
--learn-radial does these steps:
- It keeps a seeded validation set of the training configurations
(
--radial-val-frac, default 0.2). - It learns the radials on the remaining configurations, and the validation set gates the result.
- It fits the coefficients on the full training set, as usual.
--learn-radial also works with --model, and with all final arms and UQ
options. Its options are:
--radial-n-q 12: the number of polynomials for each radial;--radial-steps 40: the number of L-BFGS steps for each roughness weight;--radial-lam-grid 0,1e-2: the roughness weights. The gate selects from these weights and the initial radials.
The fit writes these files:
fit/radial_info.json: the selection of the gate, and the validation score of each candidate.fit/model.npz: the learned radials. The file has a "learned" mark, so ace-jax changes the radials to splines at deployment (see below). If the gate keeps the initial radials, the file has no mark.fit/fit.yaml: all settings of the run, including the radial options.
--learn-radial does not support species-embedded bases
(--basis-embedding) at this time.
In Python¶
import jax
jax.config.update("jax_enable_x64", True)
from ace_jax.basis.export import save_npz
from ace_jax.basis.model import BasisSpec, build_basis
from ace_jax.fit.pipeline import FitConfig, load_fit_data
from ace_jax.fit.pipeline.problem import build_problem
from ace_jax.fit.radial_learn import fit_radial, save_result
from ace_jax.fit.radial_model import rnl_degrees, to_analytic
save_npz("basis.npz", build_basis(BasisSpec(order=3, max_degree=10, elements=("Si",))))
# the linear problem on the fit split, with the validation split as its test set
cfg = FitConfig(model="basis.npz", arm="linear", m_per_species=0, e0="lsq", opt="lbfgs",
r0=2.35, rungs=("map",)).validate()
d = load_fit_data(cfg, train="fit.xyz", test="val.xyz")
model, _ = to_analytic(d.model, 12) # widen to n_q = 12 polynomials
prob = build_problem(cfg, d).prob._replace(model=model)
W, info = fit_radial(prob, d.ds_train, d.ds_test, model.rnl_Wnlq,
lam_grid=(0.0, 1e-2), # relative roughness weights to try
rough_weights=1.0 / (1.0 + rnl_degrees(d.meta)) ** 2,
steps=40, reprofile_every=20, log=print)
print(info["selected"], info["scores"]) # the gate's choice and scores
save_result("learned", W, info, src_npz="basis.npz", model=model) # learned/model.npz
Then fit the coefficients again on the full training set, and use the model as usual:
aj fit --model learned/model.npz --r0 2.35 \
--train train.xyz --test test.xyz \
--e0 lsq --m-per-species 0 --opt lbfgs \
--out fit_learned
Options¶
| Argument | Meaning |
|---|---|
to_analytic(model, n_q) |
the polynomial span of each radial. 12 is a small increase that optimises well. 30 is ill-conditioned and learns only small high-frequency changes |
steps, reprofile_every |
the maximum number of L-BFGS steps, and the interval at which the evidence hyperparameters are fitted again at the current radials |
lam_grid |
relative weights of a roughness penalty on the radials. Each value is a gate candidate (learned_lam=<value>) |
spec_grid, gap_grid |
optional relative weights of a spectral prior and a data-gap prior on the change of the radials. The gate also selects from these |
If no learned candidate is better than the initial radials on the
validation split, info["selected"] is "init". The saved model then has
the initial radials and coefficients fitted for them, and it has no
"learned" mark.
With the research driver¶
bench/learn_radial/run.py in the repository does the same steps on one
extxyz file, with a seeded fit/validation split. It is not part of the
installed package.
python bench/learn_radial/run.py --model basis.npz --data train.xyz --out learned \
--ntrain 26 --nval 13 --r0 2.35 --n-q 12 --steps 40 --lam-grid 0,1e-2
It writes learned/model.npz, the radial weights, the coefficients and a
summary.json file with the gate scores.
Note
The research driver uses the E0 in the basis file. It does not fit E0.
Deployment: spline speed¶
A learned radial is a polynomial mixture. Its evaluation is slower than the cubic splines of a stock model. Thus:
- The learner marks its output as learned (
radial_learnedin the metadata of the file). - With the default
spline_tol="auto",ACECalculatorandexport_lammpschange the learned radials to cubic splines before evaluation, at a relative tolerance of 1e-10.
from ace_jax import ACECalculator
calc = ACECalculator("fit_learned/model.npz")
calc.splined # {'spline_tol': 1e-10, 'radials': ['rnl'], 'n_intervals': {'rnl': ...}}
- The splined model agrees with the exact model to approximately the tolerance, not to round-off. On the benchmark models, energies agree to approximately 1e-9 (relative), and forces to approximately 2e-8 of the largest force.
spline_tol=Nonenever makes splines. A float (for example1e-10) makes splines of all analytic radials, learned or not.- Bases built by ace-jax are analytic but not learned. Thus, by default, they stay exact.
- The calculator keeps the splines in a cache, keyed on the content of the
radials. Thus, if you replace only the coefficients (
calc.model = ...), it does not make the splines again. - A learned-radial file written before the "learned" mark existed loads as
not learned. To correct this, mark it with
ace_jax.basis.export.mark_radial_learned("model.npz"), or give a floatspline_tol.