diff --git a/README.md b/README.md
index 249b2a1..0a27207 100644
--- a/README.md
+++ b/README.md
@@ -178,6 +178,8 @@ The preprocessing command accepts `.xyz`, `.lmdb`/`.aselmdb`, and `.h5` inputs;
For MACE-POLAR/PolarMACE inputs, Equitrain preserves system-level `charge`, `spin`, and `external_field` metadata and maps them to the MACE graph keys `total_charge`, `total_spin`, and `external_field`; the XYZ key names can be changed with `--total-charge-key`, `--total-spin-key`, and `--external-field-key`.
+For reaction-relative training data, Equitrain also preserves integer `source_id`, `reaction_id`, and `state_id` metadata. Ordinary frames default to `source_id=0`, `reaction_id=-1`, and `state_id=-1`; reactive triplets can use `state_id=0` for reactant, `1` for transition state, and `2` for product. The XYZ key names can be changed with `--source-id-key`, `--reaction-id-key`, and `--state-id-key`.
+
Under the hood, each processed file is organised as:
- `/structures`: per-configuration metadata (cell, energy, stress, charge, spin, external field, weights, etc.) and pointers into the per-atom arrays.
@@ -268,6 +270,12 @@ HDF5 inputs can be a directory, a glob (e.g. `data/train_*.h5`), or a comma-sepa
list of files; all shards are concatenated in order. This applies to
`--train-file`, `--valid-file`, and `--test-file` when training with either backend.
+Torch training can add reaction-relative energy targets with `--barrier-weight` for
+`E_TS - E_reactant` and `--reaction-energy-weight` for
+`E_product - E_reactant`. These losses require complete reaction groups in the
+Torch batch, are averaged once per reaction rather than per frame, and are not
+currently available for the JAX backend.
+
#### Python Script:
@@ -701,9 +709,25 @@ Delta fine-tuning is the simplest adapter method in the repository:
- the forward pass uses `base_parameter + delta`
- the base model stays frozen throughout optimisation
-This is effectively LoRA without any rank compression. It is useful when you
-want the simplest possible residual fine-tuning scheme and do not need to limit
-adapter size aggressively.
+This is the Equitrain residual-parameter implementation of L2-SP ("Starting
+Point") regularization from [Li, Grandvalet, and Davoine, 2018, *Explicit
+Inductive Bias for Transfer Learning with Convolutional
+Networks*](https://proceedings.mlr.press/v80/li18a.html). L2-SP regularizes
+fine-tuned parameters toward their pre-trained starting values instead of toward
+zero:
+
+```text
+Omega(theta) = lambda / 2 * ||theta - theta_0||_2^2
+```
+
+Equitrain parameterizes this as `theta = theta_0 + delta`, with `theta_0`
+frozen and `delta` initialized at zero. Optimizer weight decay on decayed delta
+tensors therefore regularizes `||delta||_2^2`, i.e. the distance between the
+effective fine-tuned parameters and the pre-trained parameters.
+
+Compared with LoRA, delta fine-tuning uses full-size residuals rather than
+low-rank residuals, so it is useful when you want the simplest residual
+fine-tuning scheme and do not need to limit adapter size aggressively.
Implementation details:
@@ -740,6 +764,12 @@ zero-based layer indices or ranges. For MACE models, Equitrain groups deltas as
`TorchDeltaFineTuneWrapper(base_model, freeze_layers="2-")` keeps only the
node embedding and first interaction block trainable.
+When delta fine-tuning is combined with `freeze_layers`, Equitrain calls this
+targeted L2-SP (L2-TSP): the L2-SP
+penalty is applied only to the selected trainable delta layers, while frozen
+layers keep `delta = 0` and remain exactly at their pre-trained starting
+values.
+
#### Freeze Fine-Tuning
For Torch models, `TorchFreezeFineTuneWrapper` provides the same semantic layer
diff --git a/equitrain/argparser.py b/equitrain/argparser.py
index 63f357a..5cd1433 100644
--- a/equitrain/argparser.py
+++ b/equitrain/argparser.py
@@ -274,6 +274,18 @@ def add_loss_weights_args(parser: argparse.ArgumentParser) -> argparse.ArgumentP
parser.add_argument(
'--stress-weight', help='Weight for stress loss', type=float, default=1.0
)
+ parser.add_argument(
+ '--barrier-weight',
+ help='Weight for relative barrier loss E_TS - E_reactant',
+ type=float,
+ default=0.0,
+ )
+ parser.add_argument(
+ '--reaction-energy-weight',
+ help='Weight for reaction energy loss E_product - E_reactant',
+ type=float,
+ default=0.0,
+ )
return parser
@@ -607,6 +619,24 @@ def get_args_parser(script_type: str) -> argparse.ArgumentParser:
type=str,
default='external_field',
)
+ parser.add_argument(
+ '--source-id-key',
+ help='Key of integer source id in training xyz',
+ type=str,
+ default='source_id',
+ )
+ parser.add_argument(
+ '--reaction-id-key',
+ help='Key of integer reaction group id in training xyz',
+ type=str,
+ default='reaction_id',
+ )
+ parser.add_argument(
+ '--state-id-key',
+ help='Key of integer reaction state id in training xyz',
+ type=str,
+ default='state_id',
+ )
parser.add_argument(
'--output-dir', help='Output directory', type=str, default=''
)
@@ -885,19 +915,57 @@ def _ensure_losses_defined(args, backend_name: str) -> None:
energy = getattr(args, 'energy_weight', 0.0) or 0.0
forces = getattr(args, 'forces_weight', 0.0) or 0.0
stress = getattr(args, 'stress_weight', 0.0) or 0.0
- if energy == 0.0 and forces == 0.0 and stress == 0.0:
+ barrier = getattr(args, 'barrier_weight', 0.0) or 0.0
+ reaction_energy = getattr(args, 'reaction_energy_weight', 0.0) or 0.0
+
+ if backend_name == 'jax' and (barrier != 0.0 or reaction_energy != 0.0):
+ raise ArgumentError(
+ 'The JAX backend does not support relative reaction losses yet; '
+ 'set --barrier-weight 0 and --reaction-energy-weight 0.'
+ )
+
+ if getattr(args, 'weighted_sampler', False) and (
+ barrier != 0.0 or reaction_energy != 0.0
+ ):
+ raise ArgumentError(
+ 'The weighted sampler does not support relative reaction losses yet; '
+ 'disable --weighted-sampler or set relative loss weights to zero.'
+ )
+
+ if (
+ energy == 0.0
+ and forces == 0.0
+ and stress == 0.0
+ and barrier == 0.0
+ and reaction_energy == 0.0
+ ):
raise ArgumentError(
f'{backend_name} backend requires at least one non-zero loss weight.'
)
+def fine_tune_export_config(model):
+ config_fn = getattr(model, 'get_fine_tune_export_config', None)
+ if callable(config_fn):
+ return config_fn()
+ return None
+
+
+def args_dict_with_runtime_metadata(args):
+ args_dict = dict(vars(args))
+ fine_tune_config = fine_tune_export_config(args_dict.get('model'))
+ if fine_tune_config is not None:
+ args_dict['fine_tune_export'] = fine_tune_config
+ return args_dict
+
+
class ArgsFormatter:
def __init__(self, args):
"""
Initialize the ArgsFormatter with parsed arguments.
:param args: argparse.Namespace object
"""
- self.args = vars(args) # Convert Namespace to dictionary
+ self.args = args_dict_with_runtime_metadata(args)
def format(self):
"""
@@ -932,6 +1000,10 @@ def is_simple(self, value):
def filter(self, args):
"""Filter the list of arguments to include only allowed types."""
- return {
- key: value for key, value in vars(args).items() if self.is_simple(value)
+ args_dict = args_dict_with_runtime_metadata(args)
+ filtered = {
+ key: value for key, value in args_dict.items() if self.is_simple(value)
}
+ if 'fine_tune_export' in args_dict:
+ filtered['fine_tune_export'] = args_dict['fine_tune_export']
+ return filtered
diff --git a/equitrain/backends/jax_backend.py b/equitrain/backends/jax_backend.py
index 5d9d941..eb703eb 100644
--- a/equitrain/backends/jax_backend.py
+++ b/equitrain/backends/jax_backend.py
@@ -20,6 +20,7 @@
from jax import tree_util as jtu
from equitrain.argparser import (
+ ArgsFilterSimple,
ArgsFormatter,
check_args_consistency,
validate_training_args,
@@ -921,6 +922,90 @@ def _run_eval_loop(
return mean_loss, loss_collection
+def _jax_runtime_config(
+ args,
+ *,
+ requested_batch_size,
+ requested_batch_max_nodes,
+ multi_device: bool,
+ device_count: int,
+ effective_workers: int,
+ prefetch_batches: int,
+ process_count: int | None = None,
+ process_index: int | None = None,
+) -> dict[str, object]:
+ graph_multiple = device_count if multi_device else 1
+ config: dict[str, object] = {
+ 'backend': 'jax',
+ 'jax_runtime_batching': 'graph-packing',
+ 'jax_requested_batch_size': requested_batch_size,
+ 'jax_runtime_batch_size': getattr(args, 'batch_size', None),
+ 'jax_requested_batch_max_nodes': requested_batch_max_nodes,
+ 'jax_runtime_batch_max_nodes': getattr(args, 'batch_max_nodes', None),
+ 'jax_runtime_batch_max_edges': getattr(args, 'batch_max_edges', None),
+ 'jax_runtime_graph_multiple': graph_multiple,
+ 'jax_runtime_multi_device': multi_device,
+ 'jax_runtime_device_count': device_count,
+ 'jax_runtime_num_workers': effective_workers,
+ 'jax_runtime_prefetch_batches': prefetch_batches,
+ }
+ if process_count is not None:
+ config['jax_runtime_process_count'] = process_count
+ if process_index is not None:
+ config['jax_runtime_process_index'] = process_index
+ return config
+
+
+def _log_jax_runtime_summary(logger, runtime_config: dict[str, object]) -> None:
+ if logger is None:
+ return
+
+ logger.log(
+ 1,
+ 'JAX runtime batching : '
+ f'{runtime_config["jax_runtime_batching"]} '
+ f'(requested batch_size={runtime_config["jax_requested_batch_size"]}, '
+ 'runtime batch_size='
+ f'{runtime_config["jax_runtime_batch_size"]})',
+ )
+ logger.log(
+ 1,
+ 'JAX runtime node limit : '
+ f'requested={runtime_config["jax_requested_batch_max_nodes"]}, '
+ f'runtime={runtime_config["jax_runtime_batch_max_nodes"]}',
+ )
+ logger.log(
+ 1,
+ f'JAX runtime edge limit : {runtime_config["jax_runtime_batch_max_edges"]}',
+ )
+ logger.log(
+ 1,
+ f'JAX runtime graph multiple: {runtime_config["jax_runtime_graph_multiple"]}',
+ )
+ logger.log(
+ 1,
+ 'JAX runtime devices : '
+ f'{runtime_config["jax_runtime_device_count"]} '
+ f'(multi_device={runtime_config["jax_runtime_multi_device"]})',
+ )
+ logger.log(
+ 1,
+ 'JAX runtime workers : '
+ f'{runtime_config["jax_runtime_num_workers"]} '
+ f'(prefetch={runtime_config["jax_runtime_prefetch_batches"]})',
+ )
+ if (
+ 'jax_runtime_process_index' in runtime_config
+ and 'jax_runtime_process_count' in runtime_config
+ ):
+ logger.log(
+ 1,
+ 'JAX runtime process : '
+ f'{runtime_config["jax_runtime_process_index"]}/'
+ f'{runtime_config["jax_runtime_process_count"]}',
+ )
+
+
def train(args):
exit_code = _launch_local_processes(args)
if exit_code is not None:
@@ -952,20 +1037,6 @@ def train(args):
logger.log(1, ArgsFormatter(args))
wandb_run = None
- if is_primary and getattr(args, 'wandb_project', None):
- try:
- import wandb
- except ModuleNotFoundError as exc: # pragma: no cover - optional dependency
- raise RuntimeError(
- 'wandb is required for the JAX backend when wandb_project is set.'
- ) from exc
-
- init_kwargs = {'project': args.wandb_project}
- if getattr(args, 'wandb_name', None):
- init_kwargs['name'] = args.wandb_name
- if getattr(args, 'wandb_group', None):
- init_kwargs['group'] = args.wandb_group
- wandb_run = wandb.init(**init_kwargs, config={'backend': 'jax'})
bundle = load_model_bundle(
args.model,
@@ -988,6 +1059,8 @@ def train(args):
local_devices = jax.local_devices()
device_count = len(local_devices) if multi_device else 1
+ requested_batch_size = getattr(args, 'batch_size', None)
+ requested_batch_max_nodes = getattr(args, 'batch_max_nodes', None)
args.batch_size = None
if getattr(args, 'batch_max_edges', None) is None:
raise ValueError(
@@ -1008,6 +1081,36 @@ def train(args):
else:
prefetch_batches = max(int(prefetch_requested or 0), 0)
+ runtime_config = _jax_runtime_config(
+ args,
+ requested_batch_size=requested_batch_size,
+ requested_batch_max_nodes=requested_batch_max_nodes,
+ multi_device=multi_device,
+ device_count=device_count,
+ effective_workers=effective_workers,
+ prefetch_batches=prefetch_batches,
+ process_count=process_count,
+ process_index=process_index,
+ )
+ _log_jax_runtime_summary(logger, runtime_config)
+
+ if is_primary and getattr(args, 'wandb_project', None):
+ try:
+ import wandb
+ except ModuleNotFoundError as exc: # pragma: no cover - optional dependency
+ raise RuntimeError(
+ 'wandb is required for the JAX backend when wandb_project is set.'
+ ) from exc
+
+ init_kwargs = {'project': args.wandb_project}
+ if getattr(args, 'wandb_name', None):
+ init_kwargs['name'] = args.wandb_name
+ if getattr(args, 'wandb_group', None):
+ init_kwargs['group'] = args.wandb_group
+ wandb_config = ArgsFilterSimple().filter(args)
+ wandb_config.update(runtime_config)
+ wandb_run = wandb.init(**init_kwargs, config=wandb_config)
+
def _build_streaming_loader(path: str | None, shuffle: bool):
if path in (None, 'None'):
return None
diff --git a/equitrain/backends/jax_evaluate.py b/equitrain/backends/jax_evaluate.py
index 2fe894e..2ba2ea4 100644
--- a/equitrain/backends/jax_evaluate.py
+++ b/equitrain/backends/jax_evaluate.py
@@ -10,7 +10,9 @@
from equitrain.backends.jax_backend import (
_build_eval_step,
_initialize_distributed,
+ _jax_runtime_config,
_launch_local_processes,
+ _log_jax_runtime_summary,
_run_eval_loop,
_shutdown_distributed,
)
@@ -119,6 +121,8 @@ def _evaluate_initialized(args):
)
multi_device = _is_multi_device()
device_count = jax.local_device_count() if multi_device else 1
+ requested_batch_size = getattr(args, 'batch_size', None)
+ requested_batch_max_nodes = getattr(args, 'batch_max_nodes', None)
args.batch_size = None
if getattr(args, 'batch_max_edges', None) is None:
raise ValueError(
@@ -139,6 +143,19 @@ def _evaluate_initialized(args):
else:
prefetch_batches = max(int(prefetch_requested or 0), 0)
+ runtime_config = _jax_runtime_config(
+ args,
+ requested_batch_size=requested_batch_size,
+ requested_batch_max_nodes=requested_batch_max_nodes,
+ multi_device=multi_device,
+ device_count=device_count,
+ effective_workers=effective_workers,
+ prefetch_batches=prefetch_batches,
+ process_count=getattr(jax, 'process_count', lambda: 1)(),
+ process_index=process_index,
+ )
+ _log_jax_runtime_summary(logger, runtime_config)
+
test_loader = get_dataloader(
data_file=test_file,
atomic_numbers=z_table,
diff --git a/equitrain/backends/torch_backend.py b/equitrain/backends/torch_backend.py
index 108a534..c9765e6 100644
--- a/equitrain/backends/torch_backend.py
+++ b/equitrain/backends/torch_backend.py
@@ -11,6 +11,7 @@
ArgsFilterSimple,
ArgsFormatter,
check_args_consistency,
+ fine_tune_export_config,
get_loss_monitor,
validate_training_args,
)
@@ -34,6 +35,25 @@
warnings.filterwarnings('ignore', message=r'.*TorchScript type system.*')
+def _log_fine_tune_summary(
+ model: torch.nn.Module,
+ accelerator: Accelerator,
+ logger: FileLogger,
+) -> None:
+ unwrapped_model = accelerator.unwrap_model(model)
+ config = fine_tune_export_config(unwrapped_model)
+ if not config:
+ return
+
+ wrapper = config.get('wrapper')
+ if wrapper is not None:
+ logger.log(1, f'Fine-tune wrapper : {wrapper}')
+ for key, value in config.items():
+ if key == 'wrapper':
+ continue
+ logger.log(1, f'Fine-tune {key:<15} : {value}')
+
+
def fix_gradients(args, model: torch.nn.Module, accelerator: Accelerator):
# Remove NaN and Inf from gradients
for param in model.parameters():
@@ -271,6 +291,7 @@ def _train_with_accelerator(args, accelerator: Accelerator):
n_parameters = sum(p.numel() for p in model.parameters() if p.requires_grad)
logger.log(1, f'Number of params : {n_parameters}')
+ _log_fine_tune_summary(model, accelerator, logger)
logger.log(1, f'Number of training points : {len(train_loader.dataset)}')
logger.log(
1,
diff --git a/equitrain/backends/torch_loss.py b/equitrain/backends/torch_loss.py
index 3b8683e..b7c912a 100644
--- a/equitrain/backends/torch_loss.py
+++ b/equitrain/backends/torch_loss.py
@@ -9,6 +9,12 @@ def __init__(self, value: torch.Tensor = None, n: torch.Tensor = None, device=No
self.n = n if n is not None else torch.tensor(0.0, device=device)
def __iadd__(self, component: LossComponent):
+ if torch.as_tensor(component.n).detach().sum().item() <= 0:
+ return self
+ if torch.as_tensor(self.n).detach().sum().item() <= 0:
+ self.value = component.value
+ self.n = component.n
+ return self
self.value = (self.value * self.n + component.value * component.n) / (
self.n + component.n
)
@@ -42,6 +48,8 @@ def __init__(self, device=None):
self['energy'] = LossComponent(device=device)
self['forces'] = LossComponent(device=device)
self['stress'] = LossComponent(device=device)
+ self['barrier'] = LossComponent(device=device)
+ self['reaction_energy'] = LossComponent(device=device)
def __iadd__(self, loss: Loss):
for key, component in loss.items():
diff --git a/equitrain/backends/torch_loss_fn.py b/equitrain/backends/torch_loss_fn.py
index fa217ce..1c87efc 100644
--- a/equitrain/backends/torch_loss_fn.py
+++ b/equitrain/backends/torch_loss_fn.py
@@ -118,12 +118,58 @@ def forward(self, input, target):
return loss, error.detach().mean(dim=(1, 2))
+class LossFnRelativeEnergy(torch.nn.Module):
+ def __init__(self, target_state: int, **args):
+ super().__init__()
+ args = dict(args)
+ if args.get('loss_type_energy') is not None:
+ args['loss_type'] = args['loss_type_energy']
+ args['loss_weight_type'] = None
+ self.error_fn = ErrorFn(**args)
+ self.target_state = int(target_state)
+
+ def forward(self, input, target, reaction_group_id, state_id):
+ input = input.reshape(-1)
+ target = target.reshape(-1).to(device=input.device, dtype=input.dtype)
+ reaction_group_id = reaction_group_id.reshape(-1).to(device=input.device)
+ state_id = state_id.reshape(-1).to(device=input.device)
+
+ pred_diffs = []
+ target_diffs = []
+ for group_id in torch.unique(reaction_group_id[reaction_group_id >= 0]):
+ group_mask = reaction_group_id == group_id
+ reference_mask = group_mask & (state_id == 0)
+ target_mask = group_mask & (state_id == self.target_state)
+ if not torch.any(reference_mask) or not torch.any(target_mask):
+ continue
+
+ pred_diffs.append(input[target_mask].mean() - input[reference_mask].mean())
+ target_diffs.append(
+ target[target_mask].mean() - target[reference_mask].mean()
+ )
+
+ if not pred_diffs:
+ zero = input.sum() * 0.0
+ count = torch.tensor(0.0, device=input.device, dtype=input.dtype)
+ return zero, zero.detach(), count
+
+ pred = torch.stack(pred_diffs)
+ true = torch.stack(target_diffs)
+ error = self.error_fn(pred, true)
+ count = torch.tensor(
+ float(error.numel()), device=input.device, dtype=input.dtype
+ )
+ return error.mean(), error.detach(), count
+
+
class LossFn(torch.nn.Module):
def __init__(
self,
energy_weight: float = 1.0,
forces_weight: float = 1.0,
stress_weight: float = 0.0,
+ barrier_weight: float = 0.0,
+ reaction_energy_weight: float = 0.0,
loss_energy_per_atom: bool = True,
**args,
):
@@ -131,10 +177,14 @@ def __init__(
self.loss_energy = LossFnEnergy(**args)
self.loss_forces = LossFnForces(**args)
self.loss_stress = LossFnStress(**args)
+ self.loss_barrier = LossFnRelativeEnergy(target_state=1, **args)
+ self.loss_reaction_energy = LossFnRelativeEnergy(target_state=2, **args)
self.energy_weight = energy_weight
self.forces_weight = forces_weight
self.stress_weight = stress_weight
+ self.barrier_weight = barrier_weight
+ self.reaction_energy_weight = reaction_energy_weight
self.loss_energy_per_atom = loss_energy_per_atom
def compute_weighted(self, energy_value, forces_value, stress_value):
@@ -153,6 +203,18 @@ def compute_weighted(self, energy_value, forces_value, stress_value):
result += self.stress_weight * stress_value
return result
+ @staticmethod
+ def _graph_attr(y_true, name: str, default: int, *, length: int, device):
+ value = getattr(y_true, name, None)
+ if value is None:
+ return torch.full((length,), default, device=device, dtype=torch.long)
+ value = value.to(device=device, dtype=torch.long).reshape(-1)
+ if value.numel() < length:
+ padded = torch.full((length,), default, device=device, dtype=torch.long)
+ padded[: value.numel()] = value
+ return padded
+ return value[:length]
+
def forward(self, y_pred, y_true):
loss = Loss(device=y_true.batch.device)
energy_weights = None
@@ -172,7 +234,9 @@ def forward(self, y_pred, y_true):
s_pred = y_pred['stress']
loss_e = loss_f = loss_s = None
+ loss_b = loss_r = None
error_e = error_f = error_s = None
+ count_b = count_r = None
if self.energy_weight > 0.0:
loss_e, error_e = self.loss_energy(e_pred, e_true, energy_weights)
@@ -181,8 +245,47 @@ def forward(self, y_pred, y_true):
if self.stress_weight > 0.0:
loss_s, error_s = self.loss_stress(s_pred, s_true)
- loss['total'].value += self.compute_weighted(loss_e, loss_f, loss_s)
- loss['total'].n += y_true.batch.max() + 1
+ graph_count = e_pred.reshape(-1).numel()
+ reaction_group_id = self._graph_attr(
+ y_true,
+ 'reaction_group_id',
+ -1,
+ length=graph_count,
+ device=e_pred.device,
+ )
+ if torch.all(reaction_group_id < 0):
+ reaction_group_id = self._graph_attr(
+ y_true,
+ 'reaction_id',
+ -1,
+ length=graph_count,
+ device=e_pred.device,
+ )
+ state_id = self._graph_attr(
+ y_true,
+ 'state_id',
+ -1,
+ length=graph_count,
+ device=e_pred.device,
+ )
+
+ total = self.compute_weighted(loss_e, loss_f, loss_s)
+ if not isinstance(total, torch.Tensor):
+ total = torch.tensor(total, device=e_pred.device, dtype=e_pred.dtype)
+
+ if self.barrier_weight > 0.0:
+ loss_b, _, count_b = self.loss_barrier(
+ e_pred, e_true, reaction_group_id, state_id
+ )
+ total = total + self.barrier_weight * loss_b
+ if self.reaction_energy_weight > 0.0:
+ loss_r, _, count_r = self.loss_reaction_energy(
+ e_pred, e_true, reaction_group_id, state_id
+ )
+ total = total + self.reaction_energy_weight * loss_r
+
+ loss['total'].value += total
+ loss['total'].n += graph_count
if self.energy_weight > 0.0:
loss['energy'].value = loss_e
@@ -193,6 +296,12 @@ def forward(self, y_pred, y_true):
if self.stress_weight > 0.0:
loss['stress'].value += loss_s
loss['stress'].n += s_true.numel()
+ if self.barrier_weight > 0.0:
+ loss['barrier'].value += loss_b
+ loss['barrier'].n += count_b
+ if self.reaction_energy_weight > 0.0:
+ loss['reaction_energy'].value += loss_r
+ loss['reaction_energy'].n += count_r
error = self.compute_weighted(error_e, error_f, error_s)
if not isinstance(error, torch.Tensor):
@@ -234,6 +343,7 @@ def forward(self, y_pred, y_true):
'LossFnEnergy',
'LossFnForces',
'LossFnStress',
+ 'LossFnRelativeEnergy',
'LossFn',
'LossFnCollection',
]
diff --git a/equitrain/backends/torch_loss_metrics.py b/equitrain/backends/torch_loss_metrics.py
index feb94d4..232831a 100644
--- a/equitrain/backends/torch_loss_metrics.py
+++ b/equitrain/backends/torch_loss_metrics.py
@@ -11,6 +11,24 @@ def __init__(self, args):
self['energy'] = AverageMeter() if args.energy_weight > 0.0 else None
self['forces'] = AverageMeter() if args.forces_weight > 0.0 else None
self['stress'] = AverageMeter() if args.stress_weight > 0.0 else None
+ self['barrier'] = (
+ AverageMeter() if getattr(args, 'barrier_weight', 0.0) > 0.0 else None
+ )
+ self['reaction_energy'] = (
+ AverageMeter()
+ if getattr(args, 'reaction_energy_weight', 0.0) > 0.0
+ else None
+ )
+ self._weights = {
+ 'energy': float(getattr(args, 'energy_weight', 0.0)),
+ 'forces': float(getattr(args, 'forces_weight', 0.0)),
+ 'stress': float(getattr(args, 'stress_weight', 0.0)),
+ 'barrier': float(getattr(args, 'barrier_weight', 0.0)),
+ 'reaction_energy': float(getattr(args, 'reaction_energy_weight', 0.0)),
+ }
+ self._use_composite_total = (
+ self['barrier'] is not None or self['reaction_energy'] is not None
+ )
def update(self, loss):
self['total'].update(
@@ -28,6 +46,26 @@ def update(self, loss):
self['stress'].update(
loss['stress'].value.detach().item(), n=loss['stress'].n.detach().item()
)
+ if self['barrier'] is not None:
+ self['barrier'].update(
+ loss['barrier'].value.detach().item(),
+ n=loss['barrier'].n.detach().item(),
+ )
+ if self['reaction_energy'] is not None:
+ self['reaction_energy'].update(
+ loss['reaction_energy'].value.detach().item(),
+ n=loss['reaction_energy'].n.detach().item(),
+ )
+ if self._use_composite_total:
+ self['total'].avg = self._composite_total()
+
+ def _composite_total(self) -> float:
+ total = 0.0
+ for key, weight in self._weights.items():
+ meter = self.get(key)
+ if meter is not None and meter.count > 0:
+ total += weight * meter.avg
+ return total
def log(
self, logger, mode: str, epoch=None, step=None, time=None, lr=None, force=False
@@ -50,6 +88,10 @@ def log(
message += f', forces: {self["forces"].avg:.5f}'
if self['stress'] is not None:
message += f', stress: {self["stress"].avg:.5f}'
+ if self['barrier'] is not None:
+ message += f', barrier: {self["barrier"].avg:.5f}'
+ if self['reaction_energy'] is not None:
+ message += f', reaction_energy: {self["reaction_energy"].avg:.5f}'
message += suffix
logger.log(1, message, force=force)
@@ -71,6 +113,10 @@ def log_step(
message += f', forces: {self["forces"].avg:.6f}'
if self['stress'] is not None:
message += f', stress: {self["stress"].avg:.6f}'
+ if self['barrier'] is not None:
+ message += f', barrier: {self["barrier"].avg:.6f}'
+ if self['reaction_energy'] is not None:
+ message += f', reaction_energy: {self["reaction_energy"].avg:.6f}'
message += suffix
logger.log(1, message, force=force)
@@ -82,6 +128,12 @@ def __init__(self, args):
self['energy'] = float('inf') if args.energy_weight > 0.0 else None
self['forces'] = float('inf') if args.forces_weight > 0.0 else None
self['stress'] = float('inf') if args.stress_weight > 0.0 else None
+ self['barrier'] = (
+ float('inf') if getattr(args, 'barrier_weight', 0.0) > 0.0 else None
+ )
+ self['reaction_energy'] = (
+ float('inf') if getattr(args, 'reaction_energy_weight', 0.0) > 0.0 else None
+ )
self['epoch'] = None
def update(self, loss, epoch):
@@ -94,6 +146,10 @@ def update(self, loss, epoch):
self['forces'] = loss['forces'].avg
if self['stress'] is not None:
self['stress'] = loss['stress'].avg
+ if self['barrier'] is not None:
+ self['barrier'] = loss['barrier'].avg
+ if self['reaction_energy'] is not None:
+ self['reaction_energy'] = loss['reaction_energy'].avg
self['epoch'] = epoch
update = True
return update
diff --git a/equitrain/data/backend_jax/atoms_to_graphs_impl.py b/equitrain/data/backend_jax/atoms_to_graphs_impl.py
index 91d44cf..ec32a62 100644
--- a/equitrain/data/backend_jax/atoms_to_graphs_impl.py
+++ b/equitrain/data/backend_jax/atoms_to_graphs_impl.py
@@ -89,6 +89,9 @@ def graph_from_configuration(
if config.external_field is None
else config.external_field
),
+ source_id=_int_array(config.source_id),
+ reaction_id=_int_array(config.reaction_id),
+ state_id=_int_array(config.state_id),
)
nodes = _AttrDict(
@@ -174,6 +177,12 @@ def _scalar_array(value: float | int | None) -> np.ndarray:
return np.asarray([float(value)], dtype=np.float32)
+def _int_array(value: int | None) -> np.ndarray:
+ if value is None:
+ value = -1
+ return np.asarray([int(value)], dtype=np.int32)
+
+
def _matrix_array(value: np.ndarray | None) -> np.ndarray:
if value is None:
base = np.zeros((3, 3), dtype=np.float32)
@@ -288,6 +297,11 @@ def graph_to_data(graph: jraph.GraphsTuple, num_species: int) -> dict[str, jnp.n
if value is not None:
data_dict[field] = jnp.asarray(value, dtype=positions.dtype)
+ for field in ('source_id', 'reaction_id', 'state_id'):
+ value = getattr(graph.globals, field, None)
+ if value is not None:
+ data_dict[field] = jnp.asarray(value, dtype=jnp.int32)
+
if hasattr(graph.nodes, 'head'):
data_dict['head'] = graph.nodes.head
diff --git a/equitrain/data/backend_torch/atoms_to_graphs.py b/equitrain/data/backend_torch/atoms_to_graphs.py
index 8ab1ce4..fbdbf00 100644
--- a/equitrain/data/backend_torch/atoms_to_graphs.py
+++ b/equitrain/data/backend_torch/atoms_to_graphs.py
@@ -89,6 +89,15 @@ def convert(self, atoms):
_external_field(atoms),
dtype=dtype,
).view(1, 3),
+ source_id=torch.tensor(
+ _info_int(atoms, 'source_id', default=0), dtype=torch.long
+ ),
+ reaction_id=torch.tensor(
+ _info_int(atoms, 'reaction_id', default=-1), dtype=torch.long
+ ),
+ state_id=torch.tensor(
+ _info_int(atoms, 'state_id', default=-1), dtype=torch.long
+ ),
fermi_level=torch.tensor(0.0, dtype=dtype),
volume=torch.tensor(np.linalg.det(cell_array), dtype=dtype),
rcell=torch.tensor(_reciprocal_cell(cell_array), dtype=dtype),
@@ -137,6 +146,10 @@ def _info_float(atoms, *keys, default: float) -> float:
return default
+def _info_int(atoms, key: str, default: int) -> int:
+ return int(np.asarray(atoms.info.get(key, default)))
+
+
def _external_field(atoms) -> np.ndarray:
return np.asarray(
atoms.info.get('external_field', np.zeros(3, dtype=float)),
diff --git a/equitrain/data/backend_torch/loaders.py b/equitrain/data/backend_torch/loaders.py
index 8631307..7cbac8f 100644
--- a/equitrain/data/backend_torch/loaders.py
+++ b/equitrain/data/backend_torch/loaders.py
@@ -10,6 +10,11 @@
from equitrain.logger import FileLogger
from .loaders_impl import DynamicGraphLoader
+from .loaders_reaction import (
+ get_reaction_loader,
+ prepare_reaction_dataset,
+ relative_reaction_losses_enabled,
+)
def _should_pin_memory(requested: bool, accelerator: Accelerator | None) -> bool:
@@ -127,12 +132,28 @@ def get_dataloader(
else:
data_set = torch.utils.data.ConcatDataset(datasets)
+ reaction_grouping = relative_reaction_losses_enabled(args)
+ reaction_group_ids = None
+ if reaction_grouping:
+ data_set, reaction_group_ids = prepare_reaction_dataset(
+ args, data_set, label=str(data_file)
+ )
+
pin_memory = _should_pin_memory(args.pin_memory, accelerator)
num_workers = _resolve_num_workers(args.num_workers, accelerator)
+ if reaction_grouping:
+ return get_reaction_loader(
+ args,
+ data_set,
+ reaction_group_ids,
+ pin_memory=pin_memory,
+ num_workers=num_workers,
+ accelerator=accelerator,
+ )
+
data_loader = DynamicGraphLoader(
dataset=data_set,
- errors=None,
batch_size=args.batch_size,
shuffle=args.shuffle,
drop_last=False,
diff --git a/equitrain/data/backend_torch/loaders_reaction.py b/equitrain/data/backend_torch/loaders_reaction.py
new file mode 100644
index 0000000..bac521b
--- /dev/null
+++ b/equitrain/data/backend_torch/loaders_reaction.py
@@ -0,0 +1,439 @@
+from __future__ import annotations
+
+import warnings
+from collections import OrderedDict
+from collections.abc import Iterator, Sequence
+
+import torch
+import torch_geometric
+from accelerate import Accelerator
+from accelerate.data_loader import prepare_data_loader as prepare_accelerate_data_loader
+from torch.utils.data import Sampler
+
+
+class ReactionMetadataDataset(torch.utils.data.Dataset):
+ """Attach loader-local reaction group ids to graph objects."""
+
+ def __init__(self, dataset, reaction_group_ids: Sequence[int]):
+ self.dataset = dataset
+ self.reaction_group_ids = [int(group_id) for group_id in reaction_group_ids]
+ if len(self.reaction_group_ids) != len(dataset):
+ raise ValueError('Reaction group id count must match dataset length.')
+
+ def __len__(self):
+ return len(self.dataset)
+
+ def __getitem__(self, index):
+ index = int(index)
+ data = self.dataset[index]
+ data.idx = index
+ data.reaction_group_id = torch.tensor(
+ self.reaction_group_ids[index], dtype=torch.long
+ )
+ return data
+
+
+class ReactionGroupBatchSampler(Sampler[list[int]]):
+ """Yield index batches that never split a reaction group."""
+
+ def __init__(
+ self,
+ reaction_group_ids: Sequence[int],
+ *,
+ batch_size: int,
+ shuffle: bool = False,
+ generator: torch.Generator | None = None,
+ num_replicas: int = 1,
+ rank: int = 0,
+ drop_last: bool = False,
+ seed: int = 0,
+ ):
+ if batch_size is None or int(batch_size) <= 0:
+ raise ValueError('A positive batch size is required for reaction grouping.')
+ if int(num_replicas) <= 0:
+ raise ValueError('num_replicas must be positive.')
+ if int(rank) < 0 or int(rank) >= int(num_replicas):
+ raise ValueError('rank must satisfy 0 <= rank < num_replicas.')
+ self.max_batch_size = int(batch_size)
+ self.shuffle = bool(shuffle)
+ self.generator = generator
+ self.num_replicas = int(num_replicas)
+ self.rank = int(rank)
+ self.drop_last = bool(drop_last)
+ self.seed = int(seed)
+ self.epoch = 0
+ self.units = _reaction_units(reaction_group_ids)
+ self.batches = _pack_reaction_units(self.units, self.max_batch_size)
+
+ def __iter__(self) -> Iterator[list[int]]:
+ batch_indices = list(range(len(self.batches)))
+ if self.shuffle:
+ generator = self.generator
+ if generator is None:
+ generator = torch.Generator()
+ generator.manual_seed(self.seed + self.epoch)
+ permutation = torch.randperm(
+ len(batch_indices), generator=generator
+ ).tolist()
+ batch_indices = [batch_indices[index] for index in permutation]
+
+ if self.num_replicas > 1:
+ batch_indices = self._make_even_batch_indices(batch_indices)
+ batch_indices = batch_indices[self.rank :: self.num_replicas]
+
+ for batch_index in batch_indices:
+ yield list(self.batches[batch_index])
+
+ def set_epoch(self, epoch: int) -> None:
+ self.epoch = int(epoch)
+
+ def __len__(self) -> int:
+ batch_count = len(self.batches)
+ if self.num_replicas <= 1:
+ return batch_count
+ if self.drop_last:
+ return batch_count // self.num_replicas
+ return (batch_count + self.num_replicas - 1) // self.num_replicas
+
+ def _make_even_batch_indices(self, batch_indices: list[int]) -> list[int]:
+ if not batch_indices:
+ return []
+ remainder = len(batch_indices) % self.num_replicas
+ if remainder == 0:
+ return batch_indices
+ if self.drop_last:
+ return batch_indices[: len(batch_indices) - remainder]
+
+ needed = self.num_replicas - remainder
+ padding = [batch_indices[index % len(batch_indices)] for index in range(needed)]
+ return [*batch_indices, *padding]
+
+
+class ReactionGraphCollater:
+ def __init__(self, collate_fn, max_nodes=None, max_edges=None, drop=False):
+ self.max_nodes = max_nodes
+ self.max_edges = max_edges
+ self.drop = drop
+ self.collate_fn = collate_fn
+
+ def __call__(self, batch):
+ dynamic_batches = []
+ current_batch = []
+ current_node_sum = 0
+ current_edge_sum = 0
+
+ for unit in _atomic_units(batch):
+ unit_node_sum = sum(item.num_nodes for item in unit)
+ unit_edge_sum = sum(item.num_edges for item in unit)
+ group_id = _reaction_group_id(unit[0])
+ grouped_reaction = group_id is not None and group_id >= 0
+
+ if grouped_reaction:
+ _raise_if_group_exceeds_limits(
+ group_id,
+ unit_node_sum,
+ unit_edge_sum,
+ self.max_nodes,
+ self.max_edges,
+ )
+ elif self._drop_oversized(unit_node_sum, unit_edge_sum):
+ continue
+
+ if current_batch:
+ if (
+ self.max_nodes is not None
+ and current_node_sum + unit_node_sum > self.max_nodes
+ ):
+ dynamic_batches.append(self.collate_fn(current_batch))
+ current_batch = []
+ current_node_sum = 0
+ current_edge_sum = 0
+
+ if (
+ self.max_edges is not None
+ and current_edge_sum + unit_edge_sum > self.max_edges
+ ):
+ dynamic_batches.append(self.collate_fn(current_batch))
+ current_batch = []
+ current_node_sum = 0
+ current_edge_sum = 0
+
+ current_batch.extend(unit)
+ current_node_sum += unit_node_sum
+ current_edge_sum += unit_edge_sum
+
+ if current_batch:
+ dynamic_batches.append(self.collate_fn(current_batch))
+
+ return dynamic_batches
+
+ def _drop_oversized(self, node_count: int, edge_count: int) -> bool:
+ if self.max_edges is not None and self.drop and edge_count > self.max_edges:
+ return True
+ if self.max_nodes is not None and self.drop and node_count > self.max_nodes:
+ return True
+ return False
+
+
+class ReactionGraphLoader(torch_geometric.loader.DataLoader):
+ def __init__(
+ self,
+ *args,
+ max_nodes=None,
+ max_edges=None,
+ drop=False,
+ **kwargs,
+ ):
+ super().__init__(*args, **kwargs)
+
+ self.collate_fn = ReactionGraphCollater(
+ self.collate_fn, max_nodes=max_nodes, max_edges=max_edges, drop=drop
+ )
+
+
+def relative_reaction_losses_enabled(args) -> bool:
+ return (
+ float(getattr(args, 'barrier_weight', 0.0) or 0.0) != 0.0
+ or float(getattr(args, 'reaction_energy_weight', 0.0) or 0.0) != 0.0
+ )
+
+
+def prepare_reaction_dataset(args, dataset, *, label: str):
+ metadata = _reaction_metadata(dataset)
+ _validate_relative_reaction_metadata(args, metadata, label=label)
+ reaction_group_ids = _reaction_group_ids_from_metadata(metadata)
+ return ReactionMetadataDataset(dataset, reaction_group_ids), reaction_group_ids
+
+
+def get_reaction_loader(
+ args,
+ dataset,
+ reaction_group_ids: Sequence[int],
+ *,
+ pin_memory: bool,
+ num_workers: int,
+ accelerator: Accelerator | None,
+):
+ data_loader = ReactionGraphLoader(
+ dataset=dataset,
+ batch_sampler=ReactionGroupBatchSampler(
+ reaction_group_ids,
+ batch_size=args.batch_size,
+ shuffle=args.shuffle,
+ num_replicas=int(getattr(accelerator, 'num_processes', 1) or 1),
+ rank=int(getattr(accelerator, 'process_index', 0) or 0),
+ seed=int(getattr(args, 'seed', 0) or 0),
+ ),
+ pin_memory=pin_memory,
+ num_workers=num_workers,
+ max_nodes=args.batch_max_nodes,
+ max_edges=args.batch_max_edges,
+ drop=args.batch_drop,
+ )
+ if accelerator is None:
+ return data_loader
+ return _prepare_rank_sharded_reaction_loader(data_loader, accelerator)
+
+
+def _reaction_metadata(dataset) -> list[tuple[int, int, int]]:
+ if isinstance(dataset, torch.utils.data.ConcatDataset):
+ result: list[tuple[int, int, int]] = []
+ for child in dataset.datasets:
+ result.extend(_reaction_metadata(child))
+ return result
+
+ reaction_metadata = getattr(dataset, 'reaction_metadata', None)
+ if not callable(reaction_metadata):
+ return [(0, -1, -1) for _ in range(len(dataset))]
+
+ return [
+ (int(source_id), int(reaction_id), int(state_id))
+ for source_id, reaction_id, state_id in reaction_metadata()
+ ]
+
+
+def _reaction_group_ids(dataset) -> list[int]:
+ return _reaction_group_ids_from_metadata(_reaction_metadata(dataset))
+
+
+def _reaction_group_ids_from_metadata(
+ reaction_metadata: list[tuple[int, int, int]],
+) -> list[int]:
+ key_to_group_id: dict[tuple[int, int], int] = {}
+ group_ids = []
+ next_group_id = 0
+ for source_id, reaction_id, _state_id in reaction_metadata:
+ reaction_id = int(reaction_id)
+ if reaction_id < 0:
+ group_ids.append(-1)
+ continue
+ key = (int(source_id), reaction_id)
+ if key not in key_to_group_id:
+ key_to_group_id[key] = next_group_id
+ next_group_id += 1
+ group_ids.append(key_to_group_id[key])
+ return group_ids
+
+
+def _validate_relative_reaction_metadata(
+ args,
+ reaction_metadata: list[tuple[int, int, int]],
+ *,
+ label: str,
+) -> None:
+ roles_by_reaction: dict[tuple[int, int], set[int]] = {}
+ for source_id, reaction_id, state_id in reaction_metadata:
+ if reaction_id < 0:
+ continue
+ roles_by_reaction.setdefault((source_id, reaction_id), set()).add(state_id)
+
+ if getattr(args, 'barrier_weight', 0.0) > 0.0:
+ _validate_relative_role_coverage(
+ roles_by_reaction,
+ required_roles={0, 1},
+ option='--barrier-weight',
+ label=label,
+ )
+ if getattr(args, 'reaction_energy_weight', 0.0) > 0.0:
+ _validate_relative_role_coverage(
+ roles_by_reaction,
+ required_roles={0, 2},
+ option='--reaction-energy-weight',
+ label=label,
+ )
+
+
+def _validate_relative_role_coverage(
+ roles_by_reaction: dict[tuple[int, int], set[int]],
+ *,
+ required_roles: set[int],
+ option: str,
+ label: str,
+) -> None:
+ complete = sum(
+ 1 for roles in roles_by_reaction.values() if required_roles.issubset(roles)
+ )
+ if complete == 0:
+ roles = ', '.join(str(role) for role in sorted(required_roles))
+ raise ValueError(
+ f'No complete reaction groups found for {option} in {label}; '
+ f'expected reaction_id >= 0 frames with state_id roles {roles}.'
+ )
+
+ incomplete = len(roles_by_reaction) - complete
+ if incomplete > 0:
+ warnings.warn(
+ f'{incomplete} reaction groups in {label} are missing roles required by '
+ f'{option} and will be skipped for that relative loss.',
+ RuntimeWarning,
+ stacklevel=3,
+ )
+
+
+def _prepare_rank_sharded_reaction_loader(data_loader, accelerator: Accelerator):
+ rng_types = getattr(accelerator, 'rng_types', None)
+ if rng_types is not None:
+ rng_types = list(rng_types)
+ return prepare_accelerate_data_loader(
+ data_loader,
+ device=accelerator.device,
+ num_processes=1,
+ process_index=0,
+ split_batches=False,
+ put_on_device=True,
+ rng_types=rng_types,
+ dispatch_batches=False,
+ even_batches=False,
+ non_blocking=getattr(accelerator, 'non_blocking', False),
+ use_stateful_dataloader=getattr(accelerator, 'use_stateful_dataloader', False),
+ )
+
+
+def _pack_reaction_units(
+ units: Sequence[Sequence[int]], max_batch_size: int
+) -> list[list[int]]:
+ batches: list[list[int]] = []
+ batch: list[int] = []
+ batch_count = 0
+ for unit in units:
+ unit = list(unit)
+ unit_count = len(unit)
+ if batch and batch_count + unit_count > max_batch_size:
+ batches.append(batch)
+ batch = []
+ batch_count = 0
+ batch.extend(unit)
+ batch_count += unit_count
+ if batch:
+ batches.append(batch)
+ return batches
+
+
+def _reaction_units(reaction_group_ids: Sequence[int]) -> list[list[int]]:
+ grouped: OrderedDict[int, list[int]] = OrderedDict()
+ units: list[list[int]] = []
+ for index, group_id in enumerate(reaction_group_ids):
+ group_id = int(group_id)
+ if group_id >= 0:
+ if group_id not in grouped:
+ grouped[group_id] = []
+ units.append(grouped[group_id])
+ grouped[group_id].append(index)
+ else:
+ units.append([index])
+ return units
+
+
+def _atomic_units(batch) -> list[list]:
+ units_by_key: OrderedDict[tuple[str, int], list] = OrderedDict()
+ ordinary_count = 0
+ for item in batch:
+ group_id = _reaction_group_id(item)
+ if group_id is not None and group_id >= 0:
+ key = ('reaction', group_id)
+ else:
+ key = ('ordinary', ordinary_count)
+ ordinary_count += 1
+ units_by_key.setdefault(key, []).append(item)
+ return list(units_by_key.values())
+
+
+def _reaction_group_id(item) -> int | None:
+ value = getattr(item, 'reaction_group_id', None)
+ if value is None:
+ value = getattr(item, 'reaction_id', None)
+ if value is None:
+ return None
+ if hasattr(value, 'detach'):
+ value = value.detach().reshape(-1)[0].item()
+ return int(value)
+
+
+def _raise_if_group_exceeds_limits(
+ group_id: int,
+ node_count: int,
+ edge_count: int,
+ max_nodes: int | None,
+ max_edges: int | None,
+) -> None:
+ if max_nodes is not None and node_count > max_nodes:
+ raise ValueError(
+ f'Reaction group {group_id} has {node_count} nodes, exceeding '
+ f'--batch-max-nodes={max_nodes}; reaction groups cannot be split.'
+ )
+ if max_edges is not None and edge_count > max_edges:
+ raise ValueError(
+ f'Reaction group {group_id} has {edge_count} edges, exceeding '
+ f'--batch-max-edges={max_edges}; reaction groups cannot be split.'
+ )
+
+
+__all__ = [
+ 'ReactionGraphCollater',
+ 'ReactionGraphLoader',
+ 'ReactionGroupBatchSampler',
+ 'ReactionMetadataDataset',
+ 'get_reaction_loader',
+ 'prepare_reaction_dataset',
+ 'relative_reaction_losses_enabled',
+]
diff --git a/equitrain/data/configuration.py b/equitrain/data/configuration.py
index 48cff44..1e59584 100644
--- a/equitrain/data/configuration.py
+++ b/equitrain/data/configuration.py
@@ -28,6 +28,9 @@ class Configuration:
total_charge: float | None = None
total_spin: float | None = None
external_field: Vector | None = None
+ source_id: int = 0
+ reaction_id: int = -1
+ state_id: int = -1
cell: Cell | None = None
pbc: Pbc | None = None
@@ -51,6 +54,9 @@ def from_atoms(
total_charge_key: str = 'charge',
total_spin_key: str = 'spin',
external_field_key: str = 'external_field',
+ source_id_key: str = 'source_id',
+ reaction_id_key: str = 'reaction_id',
+ state_id_key: str = 'state_id',
) -> 'Configuration':
"""Convert ase.Atoms to Configuration"""
@@ -100,6 +106,11 @@ def from_atoms(
aliases=('external_field',),
default=np.zeros(3),
)
+ source_id = _info_int(atoms, source_id_key, aliases=('source_id',), default=0)
+ reaction_id = _info_int(
+ atoms, reaction_id_key, aliases=('reaction_id',), default=-1
+ )
+ state_id = _info_int(atoms, state_id_key, aliases=('state_id',), default=-1)
# Charges default to 0 instead of None if not found
charges = atoms.arrays.get(charges_key, np.zeros(len(atoms)))
@@ -145,6 +156,9 @@ def from_atoms(
total_charge=float(np.asarray(total_charge)),
total_spin=float(np.asarray(total_spin)),
external_field=external_field,
+ source_id=source_id,
+ reaction_id=reaction_id,
+ state_id=state_id,
pbc=pbc,
cell=cell,
energy_weight=energy_weight,
@@ -179,6 +193,9 @@ def to_atoms(self):
atoms.info['external_field'] = (
np.zeros(3) if self.external_field is None else self.external_field
)
+ atoms.info['source_id'] = int(self.source_id)
+ atoms.info['reaction_id'] = int(self.reaction_id)
+ atoms.info['state_id'] = int(self.state_id)
atoms.info['energy_weight'] = self.energy_weight
atoms.info['forces_weight'] = self.forces_weight
@@ -201,6 +218,10 @@ def _info_value(atoms, key: str | None, *, aliases=(), default=None):
return default
+def _info_int(atoms, key: str | None, *, aliases=(), default: int) -> int:
+ return int(np.asarray(_info_value(atoms, key, aliases=aliases, default=default)))
+
+
# Replacement class for ase SinglePointCalculator, which is not stable across releases
class CachedCalc:
def __init__(self, energy, forces, stress):
diff --git a/equitrain/data/format_hdf5/dataset.py b/equitrain/data/format_hdf5/dataset.py
index 0f4d164..c4ebdea 100644
--- a/equitrain/data/format_hdf5/dataset.py
+++ b/equitrain/data/format_hdf5/dataset.py
@@ -72,6 +72,9 @@ def create_dataset(self):
('total_charge', np.float64),
('total_spin', np.float64),
('external_field', np.float64, (3,)),
+ ('source_id', np.int32),
+ ('reaction_id', np.int64),
+ ('state_id', np.int32),
('energy_weight', np.float32),
('forces_weight', np.float32),
('stress_weight', np.float32),
@@ -192,6 +195,11 @@ def __getitem__(self, i: int) -> Atoms:
atoms.info['total_charge'] = atoms.info['charge'] = total_charge
atoms.info['total_spin'] = atoms.info['spin'] = total_spin
atoms.info['external_field'] = external_field
+ atoms.info['source_id'] = int(_entry_value(entry, field_names, 'source_id', 0))
+ atoms.info['reaction_id'] = int(
+ _entry_value(entry, field_names, 'reaction_id', -1)
+ )
+ atoms.info['state_id'] = int(_entry_value(entry, field_names, 'state_id', -1))
atoms.info['energy_weight'] = entry['energy_weight']
atoms.info['forces_weight'] = entry['forces_weight']
atoms.info['stress_weight'] = entry['stress_weight']
@@ -235,6 +243,9 @@ def __setitem__(self, i: int, atoms: Atoms) -> None:
atoms.info.get('external_field', np.zeros(3, dtype=np.float64)),
dtype=np.float64,
).reshape(3)
+ source_id = np.int32(atoms.info.get('source_id', 0))
+ reaction_id = np.int64(atoms.info.get('reaction_id', -1))
+ state_id = np.int32(atoms.info.get('state_id', -1))
energy_weight = np.float32(atoms.info.get('energy_weight', 1.0))
forces_weight = np.float32(atoms.info.get('forces_weight', 1.0))
stress_weight = np.float32(atoms.info.get('stress_weight', 1.0))
@@ -266,6 +277,9 @@ def __setitem__(self, i: int, atoms: Atoms) -> None:
total_charge,
total_spin,
external_field,
+ source_id,
+ reaction_id,
+ state_id,
energy_weight,
forces_weight,
stress_weight,
@@ -306,6 +320,9 @@ def __setitem__(self, i: int, atoms: Atoms) -> None:
total_charge,
total_spin,
external_field,
+ source_id,
+ reaction_id,
+ state_id,
energy_weight,
forces_weight,
stress_weight,
@@ -313,6 +330,17 @@ def __setitem__(self, i: int, atoms: Atoms) -> None:
dipole_weight,
)
+ def reaction_metadata(self):
+ structures = self.file[self.STRUCTURES_DATASET]
+ return list(
+ zip(
+ _entry_array(structures, 'source_id', 0, np.int32).tolist(),
+ _entry_array(structures, 'reaction_id', -1, np.int64).tolist(),
+ _entry_array(structures, 'state_id', -1, np.int32).tolist(),
+ strict=True,
+ )
+ )
+
def check_magic(self):
try:
grp = self.file['MAGIC']
@@ -340,12 +368,9 @@ def _entry_value(entry, field_names, key, default):
return entry[key] if key in field_names else default
-def _has_polar_fields(structures) -> bool:
+def _has_fields(structures, fields: tuple[str, ...]) -> bool:
field_names = structures.dtype.names or ()
- return all(
- field in field_names
- for field in ('total_charge', 'total_spin', 'external_field')
- )
+ return all(field in field_names for field in fields)
def _structure_values(
@@ -361,6 +386,9 @@ def _structure_values(
total_charge,
total_spin,
external_field,
+ source_id,
+ reaction_id,
+ state_id,
energy_weight,
forces_weight,
stress_weight,
@@ -368,14 +396,23 @@ def _structure_values(
dipole_weight,
):
values = [offset, length, cell, pbc, energy, stress, virials, dipole]
- if _has_polar_fields(structures):
+ if _has_fields(structures, ('total_charge', 'total_spin', 'external_field')):
values.extend([total_charge, total_spin, external_field])
+ if _has_fields(structures, ('source_id', 'reaction_id', 'state_id')):
+ values.extend([source_id, reaction_id, state_id])
values.extend(
[energy_weight, forces_weight, stress_weight, virials_weight, dipole_weight]
)
return tuple(values)
+def _entry_array(structures, field_name: str, default, dtype):
+ field_names = structures.dtype.names or ()
+ if field_name in field_names:
+ return np.asarray(structures[field_name], dtype=dtype)
+ return np.full((structures.shape[0],), default, dtype=dtype)
+
+
class HDF5GraphDataset(HDF5Dataset):
def __init__(
self,
diff --git a/equitrain/data/format_lmdb/lmdb.py b/equitrain/data/format_lmdb/lmdb.py
index 51a85fa..47ec00b 100644
--- a/equitrain/data/format_lmdb/lmdb.py
+++ b/equitrain/data/format_lmdb/lmdb.py
@@ -115,6 +115,9 @@ def lmdb_entry_to_atoms(entry: Mapping) -> Atoms:
atoms.info.setdefault('total_spin', total_spin)
atoms.info.setdefault('spin', total_spin)
atoms.info.setdefault('external_field', external_field)
+ atoms.info.setdefault('source_id', int(np.asarray(entry.get('source_id', 0))))
+ atoms.info.setdefault('reaction_id', int(np.asarray(entry.get('reaction_id', -1))))
+ atoms.info.setdefault('state_id', int(np.asarray(entry.get('state_id', -1))))
atoms.info.setdefault('energy_weight', 1.0)
atoms.info.setdefault('forces_weight', 1.0)
atoms.info.setdefault('stress_weight', 1.0)
diff --git a/equitrain/data/format_xyz/reader.py b/equitrain/data/format_xyz/reader.py
index 31f72b2..ee89939 100644
--- a/equitrain/data/format_xyz/reader.py
+++ b/equitrain/data/format_xyz/reader.py
@@ -19,6 +19,9 @@ def __init__(
total_charge_key: str = 'charge',
total_spin_key: str = 'spin',
external_field_key: str = 'external_field',
+ source_id_key: str = 'source_id',
+ reaction_id_key: str = 'reaction_id',
+ state_id_key: str = 'state_id',
extract_atomic_numbers: bool = False,
extract_atomic_energies: bool = False,
):
@@ -32,6 +35,9 @@ def __init__(
self.total_charge_key = total_charge_key
self.total_spin_key = total_spin_key
self.external_field_key = external_field_key
+ self.source_id_key = source_id_key
+ self.reaction_id_key = reaction_id_key
+ self.state_id_key = state_id_key
self.z_set = set()
self.atomic_energies = {}
self.extract_atomic_numbers = extract_atomic_numbers
@@ -63,6 +69,9 @@ def __iter__(self):
total_charge_key=self.total_charge_key,
total_spin_key=self.total_spin_key,
external_field_key=self.external_field_key,
+ source_id_key=self.source_id_key,
+ reaction_id_key=self.reaction_id_key,
+ state_id_key=self.state_id_key,
).to_atoms()
yield atoms
diff --git a/equitrain/finetune/delta_jax.py b/equitrain/finetune/delta_jax.py
index 1b7a535..53b90b5 100644
--- a/equitrain/finetune/delta_jax.py
+++ b/equitrain/finetune/delta_jax.py
@@ -47,7 +47,7 @@ def merge_delta_params(base_params, delta_params) -> dict:
def ensure_delta_params(variables, delta_template) -> flax_core.FrozenDict:
"""
- Wrap a full MACE-JAX NNX state into Equitrain's delta fine-tuning layout.
+ Wrap a full MACE-JAX NNX state into Equitrain's delta/L^2-SP layout.
"""
unfrozen = _as_mutable_tree(variables)
@@ -67,8 +67,12 @@ def ensure_delta_params(variables, delta_template) -> flax_core.FrozenDict:
class DeltaFineTuneModule:
"""
- Wrap an NNX module so Equitrain can fine-tune additive deltas on top of the
- frozen imported MACE-JAX state.
+ Wrap an NNX module with additive residual parameters for L^2-SP fine-tuning.
+
+ The imported state is frozen under ``base_params`` and trainable deltas are
+ stored under ``params.delta``. Applying the module evaluates
+ ``theta = theta_0 + delta``, so optimizer weight decay on deltas corresponds
+ to the L^2-SP penalty on distance from the imported starting weights.
"""
def __init__(self, inner_module):
diff --git a/equitrain/finetune/delta_torch.py b/equitrain/finetune/delta_torch.py
index 6321169..7d1465e 100644
--- a/equitrain/finetune/delta_torch.py
+++ b/equitrain/finetune/delta_torch.py
@@ -30,13 +30,20 @@ def _sanitize(name: str) -> str:
class DeltaFineTuneWrapper(AbstractWrapper):
"""
- Wrap a :class:`~equitrain.backends.torch_wrappers.AbstractWrapper` instance with
- additive (delta) parameters that are trained while the original parameters remain
- frozen.
-
- The wrapper keeps a reference to the underlying base wrapper and proxies all
- attribute access to it. During the forward pass, the deltas are temporarily added
- to the base parameters.
+ Wrap a :class:`~equitrain.backends.torch_wrappers.AbstractWrapper` instance
+ with additive residual parameters for L^2-SP fine-tuning.
+
+ The wrapped base model is frozen at its pre-trained starting point. Each
+ base parameter is mirrored by a zero-initialized ``delta`` parameter, and
+ the forward pass evaluates the effective parameter
+ ``theta = theta_0 + delta``. Optimizer weight decay on trainable deltas
+ therefore corresponds to the L^2-SP penalty on distance from the starting
+ weights.
+
+ When ``freeze_layers`` freezes selected semantic delta layers, Equitrain
+ calls the configuration targeted L^2-SP (L^2-TSP): L^2-SP is applied
+ only to the remaining trainable delta layers, while frozen layers keep
+ ``delta = 0``.
"""
def __init__(self, base_wrapper: AbstractWrapper, *, freeze_layers=None):
diff --git a/equitrain/preprocess.py b/equitrain/preprocess.py
index 579449a..823398c 100644
--- a/equitrain/preprocess.py
+++ b/equitrain/preprocess.py
@@ -155,6 +155,9 @@ def _convert_to_hdf5(
total_charge_key=getattr(args, 'total_charge_key', 'charge'),
total_spin_key=getattr(args, 'total_spin_key', 'spin'),
external_field_key=getattr(args, 'external_field_key', 'external_field'),
+ source_id_key=getattr(args, 'source_id_key', 'source_id'),
+ reaction_id_key=getattr(args, 'reaction_id_key', 'reaction_id'),
+ state_id_key=getattr(args, 'state_id_key', 'state_id'),
extract_atomic_numbers=extract_atomic_numbers,
extract_atomic_energies=extract_atomic_energies,
)
diff --git a/tests/test_finetune_delta_torch.py b/tests/test_finetune_delta_torch.py
index b002620..6258ec0 100644
--- a/tests/test_finetune_delta_torch.py
+++ b/tests/test_finetune_delta_torch.py
@@ -3,6 +3,7 @@
import pytest
import torch
+from equitrain.argparser import ArgsFilterSimple, ArgsFormatter
from equitrain.backends.torch_optimizer import create_optimizer_impl
from equitrain.backends.torch_wrappers import AbstractWrapper
from equitrain.finetune._layer_selection import infer_semantic_layer_names
@@ -189,6 +190,32 @@ def test_delta_wrapper_freezes_semantic_layer_range():
}
+def test_args_formatter_includes_delta_freeze_layers():
+ args = type('Args', (), {})()
+ args.model = DeltaFineTuneWrapper(_ToyMaceLikeWrapper(), freeze_layers='2-')
+
+ formatted = ArgsFormatter(args).format()
+
+ assert 'fine_tune_export' in formatted
+ assert 'wrapper' in formatted
+ assert 'delta' in formatted
+ assert 'freeze_layers' in formatted
+ assert '2-' in formatted
+
+
+def test_args_filter_simple_includes_delta_freeze_layers():
+ args = type('Args', (), {})()
+ args.model = DeltaFineTuneWrapper(_ToyMaceLikeWrapper(), freeze_layers='2-')
+ args.lr = 1e-3
+
+ filtered = ArgsFilterSimple().filter(args)
+
+ assert filtered['fine_tune_export'] == {
+ 'wrapper': 'delta',
+ 'freeze_layers': '2-',
+ }
+
+
def test_delta_wrapper_freezes_from_forward_order_index():
wrapper = DeltaFineTuneWrapper(_ToyMaceLikeWrapper(), freeze_layers='3-')
diff --git a/tests/test_finetune_freeze_torch.py b/tests/test_finetune_freeze_torch.py
index 81bb09a..4283be6 100644
--- a/tests/test_finetune_freeze_torch.py
+++ b/tests/test_finetune_freeze_torch.py
@@ -3,6 +3,7 @@
import pytest
import torch
+from equitrain.argparser import ArgsFilterSimple, ArgsFormatter
from equitrain.backends.torch_optimizer import create_optimizer_impl
from equitrain.backends.torch_wrappers import AbstractWrapper
from equitrain.finetune.freeze_torch import FreezeFineTuneWrapper
@@ -115,6 +116,32 @@ def test_freeze_wrapper_freezes_semantic_layer_range():
}
+def test_args_formatter_includes_freeze_freeze_layers():
+ args = type('Args', (), {})()
+ args.model = FreezeFineTuneWrapper(_ToyMaceLikeWrapper(), freeze_layers='2-')
+
+ formatted = ArgsFormatter(args).format()
+
+ assert 'fine_tune_export' in formatted
+ assert 'wrapper' in formatted
+ assert 'freeze' in formatted
+ assert 'freeze_layers' in formatted
+ assert '2-' in formatted
+
+
+def test_args_filter_simple_includes_freeze_freeze_layers():
+ args = type('Args', (), {})()
+ args.model = FreezeFineTuneWrapper(_ToyMaceLikeWrapper(), freeze_layers='2-')
+ args.lr = 1e-3
+
+ filtered = ArgsFilterSimple().filter(args)
+
+ assert filtered['fine_tune_export'] == {
+ 'wrapper': 'freeze',
+ 'freeze_layers': '2-',
+ }
+
+
def test_freeze_wrapper_freezes_from_forward_order_index():
wrapper = FreezeFineTuneWrapper(_ToyMaceLikeWrapper(), freeze_layers='3-')
diff --git a/tests/test_jax_evaluate.py b/tests/test_jax_evaluate.py
index 7b70b88..1c3cfba 100644
--- a/tests/test_jax_evaluate.py
+++ b/tests/test_jax_evaluate.py
@@ -123,9 +123,54 @@ def fake_run_eval_loop(
assert captured['run_eval_loop_multi_device'] is True
assert captured['run_eval_loop_loader'] == ['g0', 'g1', 'g2', 'g3']
assert captured['loader_kwargs']['niggli_reduce'] is False
+ assert captured['loader_kwargs']['graph_multiple'] == 2
assert captured['run_eval_loop_params'] == {'weights': 1.0}
assert captured['wrapper_kwargs']['compute_force'] is False
assert captured['wrapper_kwargs']['compute_stress'] is False
+ messages = [
+ str(args[1]) for args, _kwargs in captured['log_calls'] if len(args) > 1
+ ]
+ assert any('JAX runtime batching' in message for message in messages)
+ assert any('runtime batch_size=None' in message for message in messages)
+ assert args.batch_size is None
+ assert args.batch_max_nodes is None
+
+
+def test_jax_runtime_config_records_requested_and_effective_batching():
+ args = SimpleNamespace(
+ batch_size=None,
+ batch_max_nodes=None,
+ batch_max_edges=4096,
+ )
+
+ config = jax_evaluate._jax_runtime_config(
+ args,
+ requested_batch_size=8,
+ requested_batch_max_nodes=1024,
+ multi_device=True,
+ device_count=2,
+ effective_workers=4,
+ prefetch_batches=3,
+ process_count=2,
+ process_index=1,
+ )
+
+ assert config == {
+ 'backend': 'jax',
+ 'jax_runtime_batching': 'graph-packing',
+ 'jax_requested_batch_size': 8,
+ 'jax_runtime_batch_size': None,
+ 'jax_requested_batch_max_nodes': 1024,
+ 'jax_runtime_batch_max_nodes': None,
+ 'jax_runtime_batch_max_edges': 4096,
+ 'jax_runtime_graph_multiple': 2,
+ 'jax_runtime_multi_device': True,
+ 'jax_runtime_device_count': 2,
+ 'jax_runtime_num_workers': 4,
+ 'jax_runtime_prefetch_batches': 3,
+ 'jax_runtime_process_count': 2,
+ 'jax_runtime_process_index': 1,
+ }
def test_jax_evaluate_requires_pack_limits(monkeypatch):
diff --git a/tests/test_polar_mace_data.py b/tests/test_polar_mace_data.py
index 2489614..4ba28e8 100644
--- a/tests/test_polar_mace_data.py
+++ b/tests/test_polar_mace_data.py
@@ -12,7 +12,15 @@
from equitrain.data.format_hdf5 import HDF5Dataset
-def _atoms(*, charge=-1.0, spin=2.0, external_field=(0.1, -0.2, 0.3)):
+def _atoms(
+ *,
+ charge=-1.0,
+ spin=2.0,
+ external_field=(0.1, -0.2, 0.3),
+ source_id=1,
+ reaction_id=7,
+ state_id=1,
+):
atoms = Atoms(
symbols='OH',
positions=np.array([[0.0, 0.0, 0.0], [0.0, 0.0, 0.96]], dtype=float),
@@ -30,6 +38,9 @@ def _atoms(*, charge=-1.0, spin=2.0, external_field=(0.1, -0.2, 0.3)):
atoms.info['charge'] = charge
atoms.info['spin'] = spin
atoms.info['external_field'] = np.asarray(external_field, dtype=float)
+ atoms.info['source_id'] = source_id
+ atoms.info['reaction_id'] = reaction_id
+ atoms.info['state_id'] = state_id
atoms.info['energy_weight'] = 1.0
atoms.info['forces_weight'] = 1.0
atoms.info['stress_weight'] = 1.0
@@ -54,6 +65,9 @@ def test_configuration_preserves_polar_mace_metadata():
assert config.total_charge == -2.0
assert config.total_spin == 3.0
np.testing.assert_allclose(config.external_field, [0.4, 0.5, 0.6])
+ assert config.source_id == 1
+ assert config.reaction_id == 7
+ assert config.state_id == 1
roundtrip = config.to_atoms()
assert roundtrip.info['charge'] == -2.0
@@ -61,6 +75,9 @@ def test_configuration_preserves_polar_mace_metadata():
assert roundtrip.info['spin'] == 3.0
assert roundtrip.info['total_spin'] == 3.0
np.testing.assert_allclose(roundtrip.info['external_field'], [0.4, 0.5, 0.6])
+ assert roundtrip.info['source_id'] == 1
+ assert roundtrip.info['reaction_id'] == 7
+ assert roundtrip.info['state_id'] == 1
def test_hdf5_roundtrip_stores_polar_mace_metadata(tmp_path):
@@ -73,6 +90,9 @@ def test_hdf5_roundtrip_stores_polar_mace_metadata(tmp_path):
assert 'total_charge' in names
assert 'total_spin' in names
assert 'external_field' in names
+ assert 'source_id' in names
+ assert 'reaction_id' in names
+ assert 'state_id' in names
atoms = dataset[0]
assert atoms.info['charge'] == -1.5
@@ -80,6 +100,10 @@ def test_hdf5_roundtrip_stores_polar_mace_metadata(tmp_path):
assert atoms.info['spin'] == 4.0
assert atoms.info['total_spin'] == 4.0
np.testing.assert_allclose(atoms.info['external_field'], [0.2, 0.0, -0.1])
+ assert atoms.info['source_id'] == 1
+ assert atoms.info['reaction_id'] == 7
+ assert atoms.info['state_id'] == 1
+ assert dataset.reaction_metadata() == [(1, 7, 1)]
def test_hdf5_old_schema_defaults_to_neutral_singlet_zero_field(tmp_path):
@@ -161,6 +185,9 @@ def test_hdf5_old_schema_defaults_to_neutral_singlet_zero_field(tmp_path):
assert atoms.info['spin'] == 1.0
assert atoms.info['total_spin'] == 1.0
np.testing.assert_allclose(atoms.info['external_field'], np.zeros(3))
+ assert atoms.info['source_id'] == 0
+ assert atoms.info['reaction_id'] == -1
+ assert atoms.info['state_id'] == -1
def test_torch_graph_contains_polar_mace_inputs():
@@ -183,11 +210,17 @@ def test_torch_graph_contains_polar_mace_inputs():
assert graph.total_charge.shape == torch.Size([])
assert graph.total_spin.shape == torch.Size([])
assert graph.external_field.shape == (1, 3)
+ assert graph.source_id.shape == torch.Size([])
+ assert graph.reaction_id.shape == torch.Size([])
+ assert graph.state_id.shape == torch.Size([])
assert graph.fermi_level.shape == torch.Size([])
assert graph.volume.shape == torch.Size([])
assert graph.rcell.shape == (3, 3)
assert graph.total_charge.item() == -1.0
assert graph.total_spin.item() == 2.0
+ assert graph.source_id.item() == 1
+ assert graph.reaction_id.item() == 7
+ assert graph.state_id.item() == 1
np.testing.assert_allclose(graph.external_field.numpy(), [[0.1, -0.2, 0.3]])
cell = np.eye(3) * 5.0
@@ -203,9 +236,15 @@ def test_torch_graph_contains_polar_mace_inputs():
assert batch.total_charge.shape == (2,)
assert batch.total_spin.shape == (2,)
assert batch.external_field.shape == (2, 3)
+ assert batch.source_id.shape == (2,)
+ assert batch.reaction_id.shape == (2,)
+ assert batch.state_id.shape == (2,)
assert batch.fermi_level.shape == (2,)
assert batch.volume.shape == (2,)
assert batch.rcell.shape == (6, 3)
np.testing.assert_allclose(batch.total_charge.numpy(), [-1.0, 0.5])
np.testing.assert_allclose(batch.total_spin.numpy(), [2.0, 1.5])
+ np.testing.assert_array_equal(batch.source_id.numpy(), [1, 1])
+ np.testing.assert_array_equal(batch.reaction_id.numpy(), [7, 7])
+ np.testing.assert_array_equal(batch.state_id.numpy(), [1, 1])
np.testing.assert_allclose(batch.external_field.numpy()[1], [0.0, 0.0, 1.0])
diff --git a/tests/test_reaction_relative_loss.py b/tests/test_reaction_relative_loss.py
new file mode 100644
index 0000000..e2641ab
--- /dev/null
+++ b/tests/test_reaction_relative_loss.py
@@ -0,0 +1,358 @@
+from __future__ import annotations
+
+from types import SimpleNamespace
+
+import pytest
+import torch
+
+from equitrain.argparser import (
+ ArgumentError,
+ get_args_parser_train,
+ validate_training_args,
+)
+from equitrain.backends.torch_loss_fn import LossFnCollection
+from equitrain.backends.torch_loss_metrics import LossMetrics
+from equitrain.data.backend_torch.loaders_impl import DynamicGraphCollater
+from equitrain.data.backend_torch.loaders_reaction import (
+ ReactionGraphCollater,
+ ReactionGroupBatchSampler,
+ _reaction_group_ids,
+ _validate_relative_reaction_metadata,
+)
+
+
+def _loss_args(**overrides):
+ args = dict(
+ energy_weight=0.0,
+ forces_weight=0.0,
+ stress_weight=0.0,
+ barrier_weight=1.0,
+ reaction_energy_weight=1.0,
+ loss_energy_per_atom=True,
+ loss_type='mae',
+ loss_type_energy='mae',
+ loss_type_forces='mae',
+ loss_type_stress='mae',
+ loss_weight_type=None,
+ loss_weight_type_energy=None,
+ loss_weight_type_forces=None,
+ loss_weight_type_stress=None,
+ smooth_l1_beta=1.0,
+ huber_delta=1.0,
+ loss_clipping=None,
+ loss_monitor=['mse'],
+ )
+ args.update(overrides)
+ return args
+
+
+class _MetadataDataset(torch.utils.data.Dataset):
+ def __init__(self, metadata):
+ self.metadata = metadata
+
+ def __len__(self):
+ return len(self.metadata)
+
+ def __getitem__(self, index):
+ raise IndexError(index)
+
+ def reaction_metadata(self):
+ return self.metadata
+
+
+def _reaction_graph(*, energy, reaction_group_id, state_id, num_nodes=1):
+ pytest.importorskip('torch_geometric')
+ from torch_geometric.data import Data
+
+ return Data(
+ pos=torch.zeros(num_nodes, 3),
+ positions=torch.zeros(num_nodes, 3),
+ y=torch.tensor(float(energy)),
+ force=torch.zeros(num_nodes, 3),
+ stress=torch.zeros(1, 3, 3),
+ edge_index=torch.zeros(2, 0, dtype=torch.long),
+ reaction_group_id=torch.tensor(reaction_group_id, dtype=torch.long),
+ reaction_id=torch.tensor(reaction_group_id, dtype=torch.long),
+ state_id=torch.tensor(state_id, dtype=torch.long),
+ )
+
+
+def test_reaction_group_ids_are_global_across_concat_datasets():
+ left = _MetadataDataset([(1, 42, 0)])
+ right = _MetadataDataset([(1, 42, 1), (1, 42, 2), (2, 42, 0)])
+ dataset = torch.utils.data.ConcatDataset([left, right])
+
+ assert _reaction_group_ids(dataset) == [0, 0, 0, 1]
+
+
+def test_relative_reaction_metadata_requires_complete_requested_roles():
+ args = SimpleNamespace(barrier_weight=1.0, reaction_energy_weight=0.0)
+
+ with pytest.raises(
+ ValueError, match='No complete reaction groups.*--barrier-weight'
+ ):
+ _validate_relative_reaction_metadata(
+ args,
+ [(1, 7, 0), (1, 7, 2)],
+ label='train.h5',
+ )
+
+
+def test_relative_reaction_metadata_warns_for_incomplete_requested_roles():
+ args = SimpleNamespace(barrier_weight=1.0, reaction_energy_weight=0.0)
+
+ with pytest.warns(RuntimeWarning, match='missing roles required'):
+ _validate_relative_reaction_metadata(
+ args,
+ [(1, 7, 0), (1, 7, 1), (1, 8, 0)],
+ label='train.h5',
+ )
+
+
+def test_relative_reaction_losses_are_averaged_per_reaction_not_frame():
+ pytest.importorskip('torch_geometric')
+ from torch_geometric.data import Batch
+
+ target = Batch.from_data_list(
+ [
+ _reaction_graph(energy=0.0, reaction_group_id=0, state_id=0, num_nodes=3),
+ _reaction_graph(energy=5.0, reaction_group_id=0, state_id=1, num_nodes=1),
+ _reaction_graph(energy=-1.0, reaction_group_id=0, state_id=2, num_nodes=2),
+ _reaction_graph(energy=100.0, reaction_group_id=-1, state_id=-1),
+ ]
+ )
+ pred = {
+ 'energy': torch.tensor([0.0, 7.0, -0.5, 50.0]),
+ 'forces': torch.zeros(target.num_nodes, 3),
+ 'stress': torch.zeros(target.num_graphs, 3, 3),
+ }
+
+ loss, _ = LossFnCollection(
+ **_loss_args(barrier_weight=2.0, reaction_energy_weight=3.0)
+ )(pred, target)
+
+ assert loss.main['barrier'].value.item() == pytest.approx(2.0)
+ assert loss.main['barrier'].n.item() == pytest.approx(1.0)
+ assert loss.main['reaction_energy'].value.item() == pytest.approx(0.5)
+ assert loss.main['reaction_energy'].n.item() == pytest.approx(1.0)
+ assert loss.main['total'].value.item() == pytest.approx(5.5)
+
+ assert loss['mse']['barrier'].value.item() == pytest.approx(4.0)
+ assert loss['mse']['reaction_energy'].value.item() == pytest.approx(0.25)
+ assert loss['mse']['total'].value.item() == pytest.approx(8.75)
+
+
+def test_relative_loss_metric_total_uses_reaction_averages():
+ pytest.importorskip('torch_geometric')
+ from torch_geometric.data import Batch
+
+ args = SimpleNamespace(
+ energy_weight=0.0,
+ forces_weight=0.0,
+ stress_weight=0.0,
+ barrier_weight=1.0,
+ reaction_energy_weight=0.0,
+ loss_type='mae',
+ loss_monitor=[],
+ )
+ loss_fn = LossFnCollection(
+ **_loss_args(reaction_energy_weight=0.0, loss_monitor=[])
+ )
+ metrics = LossMetrics(args)
+
+ target_1 = Batch.from_data_list(
+ [
+ _reaction_graph(energy=0.0, reaction_group_id=0, state_id=0),
+ _reaction_graph(energy=0.0, reaction_group_id=0, state_id=1),
+ *[
+ _reaction_graph(energy=0.0, reaction_group_id=-1, state_id=-1)
+ for _ in range(8)
+ ],
+ ]
+ )
+ loss_1, _ = loss_fn(
+ {
+ 'energy': torch.tensor([0.0, 10.0] + [0.0] * 8),
+ 'forces': torch.zeros(target_1.num_nodes, 3),
+ 'stress': torch.zeros(target_1.num_graphs, 3, 3),
+ },
+ target_1,
+ )
+ metrics.update(loss_1)
+
+ target_2 = Batch.from_data_list(
+ [
+ _reaction_graph(energy=0.0, reaction_group_id=1, state_id=0),
+ _reaction_graph(energy=0.0, reaction_group_id=1, state_id=1),
+ ]
+ )
+ loss_2, _ = loss_fn(
+ {
+ 'energy': torch.tensor([0.0, 2.0]),
+ 'forces': torch.zeros(target_2.num_nodes, 3),
+ 'stress': torch.zeros(target_2.num_graphs, 3, 3),
+ },
+ target_2,
+ )
+ metrics.update(loss_2)
+
+ assert metrics.main['barrier'].avg == pytest.approx(6.0)
+ assert metrics.main['total'].avg == pytest.approx(6.0)
+
+
+def test_relative_reaction_losses_ignore_batches_without_complete_roles():
+ pytest.importorskip('torch_geometric')
+ from torch_geometric.data import Batch
+
+ target = Batch.from_data_list(
+ [
+ _reaction_graph(energy=0.0, reaction_group_id=0, state_id=0),
+ _reaction_graph(energy=5.0, reaction_group_id=0, state_id=1),
+ ]
+ )
+ pred = {
+ 'energy': torch.tensor([1.0, 5.0]),
+ 'forces': torch.zeros(target.num_nodes, 3),
+ 'stress': torch.zeros(target.num_graphs, 3, 3),
+ }
+
+ loss, _ = LossFnCollection(
+ **_loss_args(barrier_weight=1.0, reaction_energy_weight=1.0)
+ )(pred, target)
+
+ assert loss.main['barrier'].n.item() == pytest.approx(1.0)
+ assert loss.main['reaction_energy'].n.item() == pytest.approx(0.0)
+ assert loss.main['total'].value.item() == pytest.approx(1.0)
+
+
+def test_reaction_group_batch_sampler_keeps_groups_atomic():
+ sampler = ReactionGroupBatchSampler(
+ [0, -1, 0, 1, 1],
+ batch_size=2,
+ shuffle=False,
+ )
+
+ assert list(sampler) == [[0, 2], [1], [3, 4]]
+
+
+def test_reaction_graph_collater_does_not_split_reaction_groups_across_sub_batches():
+ items = [
+ _reaction_graph(energy=0.0, reaction_group_id=0, state_id=0),
+ _reaction_graph(energy=1.0, reaction_group_id=-1, state_id=-1, num_nodes=2),
+ _reaction_graph(energy=2.0, reaction_group_id=0, state_id=1),
+ ]
+ collater = ReactionGraphCollater(
+ lambda graphs: [int(graph.y.item()) for graph in graphs],
+ max_nodes=2,
+ max_edges=None,
+ drop=False,
+ )
+
+ assert collater(items) == [[0, 2], [1]]
+
+
+def test_dynamic_collater_keeps_standard_per_graph_batching():
+ items = [
+ _reaction_graph(energy=0.0, reaction_group_id=0, state_id=0),
+ _reaction_graph(energy=1.0, reaction_group_id=-1, state_id=-1, num_nodes=2),
+ _reaction_graph(energy=2.0, reaction_group_id=0, state_id=1),
+ ]
+ collater = DynamicGraphCollater(
+ lambda graphs: [int(graph.y.item()) for graph in graphs],
+ max_nodes=2,
+ max_edges=None,
+ drop=False,
+ )
+
+ assert collater(items) == [[0], [1], [2]]
+
+
+def test_reaction_graph_collater_rejects_oversized_reaction_group():
+ items = [
+ _reaction_graph(energy=0.0, reaction_group_id=0, state_id=0, num_nodes=2),
+ _reaction_graph(energy=2.0, reaction_group_id=0, state_id=1, num_nodes=2),
+ ]
+ collater = ReactionGraphCollater(
+ lambda graphs: graphs,
+ max_nodes=3,
+ max_edges=None,
+ drop=True,
+ )
+
+ with pytest.raises(ValueError, match='reaction groups cannot be split'):
+ collater(items)
+
+
+def test_reaction_group_batch_sampler_shards_evenly_by_whole_batches():
+ rank_0 = ReactionGroupBatchSampler(
+ [0, -1, 0, 1, 1],
+ batch_size=2,
+ shuffle=False,
+ num_replicas=2,
+ rank=0,
+ )
+ rank_1 = ReactionGroupBatchSampler(
+ [0, -1, 0, 1, 1],
+ batch_size=2,
+ shuffle=False,
+ num_replicas=2,
+ rank=1,
+ )
+
+ assert not hasattr(rank_0, 'batch_size')
+ assert len(rank_0) == len(rank_1) == 2
+ assert list(rank_0) == [[0, 2], [3, 4]]
+ assert list(rank_1) == [[1], [0, 2]]
+
+
+def test_reaction_group_batch_sampler_shuffle_is_epoch_seeded_across_ranks():
+ group_ids = [0, -1, 0, 1, 1, -1]
+ global_sampler = ReactionGroupBatchSampler(
+ group_ids,
+ batch_size=2,
+ shuffle=True,
+ seed=17,
+ )
+ rank_0 = ReactionGroupBatchSampler(
+ group_ids,
+ batch_size=2,
+ shuffle=True,
+ seed=17,
+ num_replicas=2,
+ rank=0,
+ )
+ rank_1 = ReactionGroupBatchSampler(
+ group_ids,
+ batch_size=2,
+ shuffle=True,
+ seed=17,
+ num_replicas=2,
+ rank=1,
+ )
+ for sampler in (global_sampler, rank_0, rank_1):
+ sampler.set_epoch(4)
+
+ global_batches = list(global_sampler)
+ if len(global_batches) % 2:
+ global_batches.append(global_batches[0])
+ interleaved = [
+ batch for pair in zip(list(rank_0), list(rank_1), strict=True) for batch in pair
+ ]
+
+ assert interleaved == global_batches
+
+
+def test_jax_rejects_relative_reaction_losses():
+ args = get_args_parser_train().parse_args([])
+ args.train_file = 'train.h5'
+ args.valid_file = 'valid.h5'
+ args.output_dir = 'out'
+ args.model = 'model'
+ args.energy_weight = 0.0
+ args.forces_weight = 0.0
+ args.stress_weight = 0.0
+ args.barrier_weight = 1.0
+ args.reaction_energy_weight = 0.0
+
+ with pytest.raises(ArgumentError, match='JAX backend does not support'):
+ validate_training_args(args, 'jax')