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')