Equitrain is a Python toolkit for preprocessing atomistic datasets, training machine-learning interatomic potentials (MLIPs), fine-tuning existing checkpoints, and running evaluation or prediction through one CLI/API.
- Unified Torch and JAX training entry points.
- Model wrappers for MACE, SevenNet, ORB, ANI, and M3GNet.
- Native HDF5 preprocessing for large atomistic datasets.
- Torch reaction-relative losses for barrier and reaction energies.
- Fine-tuning adapters for Delta/L2-SP, Freeze, and LoRA workflows.
- ASE calculator helpers for batched prediction and relaxation.
| Wrapper | Backends | Upstream / Companion Project | Notes |
|---|---|---|---|
mace |
Torch, JAX | mace-model |
Companion repository for MACE model definitions, conversion, and foundation-model export. |
sevennet |
Torch | MDIL-SNU/SevenNet |
Torch SevenNet checkpoints and models. |
orb |
Torch | orbital-materials/orb-models |
Torch ORB force-field models. |
ani |
Torch, JAX | aiqm/torchani |
Torch uses TorchANI; JAX uses a JAX-native bundle. |
m3gnet |
Torch, JAX | materialsvirtuallab/matgl |
Torch uses MatGL; JAX uses a JAX-native bundle. |
For MACE, use mace-model for
model construction/conversion and equitrain for preprocessing, training,
fine-tuning, checkpointing, evaluation, and prediction.
Full documentation is published at https://bamescience.github.io/equitrain/:
- Installation
- Quickstart
- Data and Preprocessing
- CLI
- Training Options
- Python API
- Model Wrappers
- JAX Bundles
- Fine-Tuning
- Calculators
- Reaction-Relative Losses
- Resources
The documentation source is in docs/. Build or serve it locally with:
pip install -e '.[docu]'
mkdocs servepip install equitrainUntil the package is fully available on PyPI, install from a local clone:
git clone https://github.com/BAMeScience/equitrain.git
cd equitrain
python3.10 -m venv .venv
source .venv/bin/activate
pip install --upgrade pip
pip install uv
uv pip install -e '.[dev,docu]'Install model/runtime extras as needed:
pip install 'equitrain[torch,mace]'
pip install 'equitrain[jax,mace-jax]'
pip install 'equitrain[torch,ani]'Preprocess data:
equitrain-preprocess \
--train-file data-train.xyz \
--valid-file data-valid.xyz \
--compute-statistics \
--atomic-energies average \
--output-dir data \
--r-max 4.5Train a Torch/MACE model:
equitrain -v \
--train-file data/train.h5 \
--valid-file data/valid.h5 \
--output-dir runs/mace \
--model path/to/mace.model \
--model-wrapper mace \
--epochs 10 \
--tqdmEvaluate and predict:
equitrain-evaluate -v \
--test-file data/test.h5 \
--model path/to/mace.model \
--model-wrapper mace \
--output-dir evaluation_mace
equitrain-predict \
--predict-file data/valid.h5 \
--model path/to/mace.model \
--model-wrapper mace \
--output-dir predictions_maceSee the Quickstart for the full workflow, including JAX bundles and fine-tuned checkpoint export.
Equitrain's Delta adapter is a residual-parameter implementation of
L2-SP ("Starting Point") regularization from Li, Grandvalet, and
Davoine, 2018,
Explicit Inductive Bias for Transfer Learning with Convolutional Networks.
It parameterizes fine-tuning as theta = theta_0 + delta, so weight decay on
trainable deltas regularizes ||delta||_2^2.
Delta combined with freeze_layers is targeted L2-SP
(L2-TSP): the L2-SP penalty applies only to selected
trainable delta layers while frozen layers remain exactly at their pre-trained
starting values. See Fine-Tuning.
Example data-preparation scripts are in resources/data, training scripts are
in resources/training, and initial model examples are in resources/models.