diff --git a/sevenn/_const.py b/sevenn/_const.py index 6dc45589..6d44b768 100644 --- a/sevenn/_const.py +++ b/sevenn/_const.py @@ -132,6 +132,9 @@ def error_record_condition(x): KEY.USE_FLASH_TP: False, KEY.CUEQUIVARIANCE_CONFIG: {}, KEY.USE_OEQ: False, + KEY.USE_LES: False, + # les_config keys: les_args (dict), compute_bec (bool), bec_output_index (int|None) + KEY.LES_CONFIG: {}, } @@ -180,6 +183,8 @@ def error_record_condition(x): KEY.USE_FLASH_TP: bool, KEY.CUEQUIVARIANCE_CONFIG: dict, KEY.USE_OEQ: bool, + KEY.USE_LES: bool, + KEY.LES_CONFIG: dict, } diff --git a/sevenn/_keys.py b/sevenn/_keys.py index 4b8d63d7..75a7ce22 100644 --- a/sevenn/_keys.py +++ b/sevenn/_keys.py @@ -58,6 +58,12 @@ ATOMIC_ENERGY: Final[str] = 'atomic_energy' PRED_TOTAL_ENERGY: Final[str] = 'inferred_total_energy' +# LES (Latent Ewald Summation) outputs +LES_Q: Final[str] = 'les_latent_charge' # (N_atoms, n_charges) per-atom latent charges +SR_ENERGY: Final[str] = 'les_sr_energy' # (n_graphs,) short-range energy sum +LR_ENERGY: Final[str] = 'les_lr_energy' # (n_graphs,) long-range Ewald energy +LES_BEC: Final[str] = 'les_born_eff_charge' # (N_atoms, 3, 3) Born effective charges + PRED_PER_ATOM_ENERGY: Final[str] = 'inferred_per_atom_energy' PER_ATOM_ENERGY: Final[str] = 'per_atom_energy' @@ -225,6 +231,10 @@ CUEQUIVARIANCE_CONFIG = 'cuequivariance_config' USE_OEQ = 'use_oeq' +# LES model configuration keys +USE_LES = 'use_les' +LES_CONFIG = 'les_config' + _NORMALIZE_SPH = '_normalize_sph' OPTIMIZE_BY_REDUCE = 'optimize_by_reduce' diff --git a/sevenn/checkpoint.py b/sevenn/checkpoint.py index e0422ec2..e8b42345 100644 --- a/sevenn/checkpoint.py +++ b/sevenn/checkpoint.py @@ -369,6 +369,140 @@ def build_model( return model + def build_model_with_les( + self, + les_config: Optional[Dict[str, Any]] = None, + freeze_sr: bool = True, + *, + enable_cueq: Optional[bool] = None, + enable_flash: Optional[bool] = None, + enable_oeq: Optional[bool] = None, + ) -> AtomGraphSequential: + """ + Build a LES-equipped model initialised from this (non-LES) checkpoint. + + New LES parameters not present in the pretrained state dict use their + default construction-time initialisation: + - ``les_charge_readout.linear.weight``: e3nn default (non-zero) + so that the Ewald gradient is non-zero and fine-tuning starts + immediately. Pass ``les_config={'zero_init': True}`` to force + zero weights, which preserves the SR baseline exactly at epoch 0 + but sets all charge gradients to zero — useful for verification + tests, not for real training. + - ``les_lr_energy.les.*``: Les() default init + + All other parameters are loaded from the checkpoint unchanged. + + Args: + les_config: LES configuration dict forwarded to the model builder as + ``LES_CONFIG``. Supported keys: + les_args (dict) — kwargs for ``Les()``, + default ``{'use_atomwise': False}`` + n_charges (int) — number of latent charge channels + per atom; the Ewald energy is the + sum of n_charges independent + Coulomb interactions, default 1 + hidden_channels (list) — hidden layer widths for the charge + readout MLP, e.g. ``[128]`` for a + two-layer 128→128→1 network, + default ``[]`` (single layer) + zero_init (bool) — zero-initialise the charge readout + weights, default ``False`` + compute_bec (bool) — compute Born effective charges, + default ``False`` + bec_output_index (int) — 0/1/2 for BEC component, + default ``None`` + freeze_sr: if ``True`` (default), set ``requires_grad=False`` on + all parameters whose top-level module is not + ``les_charge_readout`` or ``les_lr_energy``. The Trainer + already filters parameters by ``requires_grad``, so no + Trainer changes are needed. + enable_cueq/enable_flash/enable_oeq: backend overrides (same + semantics as :meth:`build_model`). + + Returns: + :class:`~sevenn.nn.sequential.AtomGraphSequential` with LES layers + added and SR weights loaded from the checkpoint. + + Raises: + RuntimeError: if any *non*-LES keys are missing after loading, + indicating an unexpected model mismatch. + """ + from .model_build import build_E3_equivariant_model + + if les_config is None: + les_config = {} + + # Resolve backend flags using the same logic as build_model() + try: + cp_using_cueq = self.config[KEY.CUEQUIVARIANCE_CONFIG]['use'] + except KeyError: + cp_using_cueq = False + final_cueq = cp_using_cueq if enable_cueq is None else enable_cueq + + cp_using_flash = self.config.get(KEY.USE_FLASH_TP, False) + final_flash = cp_using_flash if enable_flash is None else enable_flash + + cp_using_oeq = self.config.get(KEY.USE_OEQ, False) + final_oeq = cp_using_oeq if enable_oeq is None else enable_oeq + + if sum([final_cueq, final_flash, final_oeq]) > 1: + raise ValueError('Only one TP accelerator can be enabled.') + + # Step 1: use build_model() to get a correctly backend-converted non-LES + # state dict. This handles FlashTP↔e3nn and cueq↔e3nn in one place + # without duplicating the conversion logic here. + sr_state_dict = self.build_model( + enable_cueq=enable_cueq, + enable_flash=enable_flash, + enable_oeq=enable_oeq, + ).state_dict() + + # Step 2: build the LES model with the identical resolved backend. + cfg_new = self.config + cfg_new[KEY.USE_LES] = True + cfg_new[KEY.LES_CONFIG] = les_config + cfg_new[KEY.USE_OEQ] = final_oeq + cfg_new[KEY.CUEQUIVARIANCE_CONFIG] = {'use': final_cueq} + cfg_new[KEY.USE_FLASH_TP] = final_flash + + model = build_E3_equivariant_model(cfg_new) + + # Step 3: load the already-converted SR state dict into the LES model. + # strict=False because LES-specific keys are absent from sr_state_dict. + missing, not_used = model.load_state_dict(sr_state_dict, strict=False) + + # Only LES-specific keys should be absent from the pretrained checkpoint. + # Any other missing key means an unexpected model mismatch. + les_prefixes = ('les_charge_readout.', 'les_lr_energy.') + unexpected_missing = [ + k for k in missing if not k.startswith(les_prefixes) + ] + if unexpected_missing: + raise RuntimeError( + 'Unexpected missing keys when loading SR checkpoint into LES ' + f'model: {unexpected_missing}' + ) + if missing: + warnings.warn( + 'LES parameters initialised from scratch (not in checkpoint): ' + f'{missing}', + UserWarning, + ) + if not_used: + warnings.warn( + f'Checkpoint keys not used in LES model: {not_used}', + UserWarning, + ) + + if freeze_sr: + les_module_names = {'les_charge_readout', 'les_lr_energy'} + for name, param in model.named_parameters(): + if name.split('.')[0] not in les_module_names: + param.requires_grad_(False) + + return model + def yaml_dict(self, mode: str) -> Dict[str, Any]: """ Return dict for input.yaml from checkpoint config diff --git a/sevenn/model_build.py b/sevenn/model_build.py index c548c34e..795ec2e4 100644 --- a/sevenn/model_build.py +++ b/sevenn/model_build.py @@ -14,11 +14,12 @@ from .nn.edge_embedding import ( BesselBasis, EdgeEmbedding, + EdgePreprocess, PolynomialCutoff, SphericalEncoding, XPLORCutoff, ) -from .nn.force_output import ForceStressOutputFromEdge +from .nn.force_output import ForceStressOutput, ForceStressOutputFromEdge from .nn.interaction_blocks import NequIP_interaction_block from .nn.linear import AtomReduce, FCN_e3nn, IrrepsLinear from .nn.node_embedding import OnehotEmbedding @@ -28,6 +29,11 @@ SelfConnectionLinearIntro, SelfConnectionOutro, ) +from .nn.les import ( + AddLREnergy, + LatentChargeReadout, + LatentEwaldSum, +) from .nn.sequential import AtomGraphSequential # warning from PyTorch, about e3nn type annotations @@ -458,8 +464,22 @@ def build_E3_equivariant_model( for data w/o cell volume, pred_stress has garbage values """ + if parallel and config.get(KEY.USE_LES, False): + raise NotImplementedError( + 'LES parallel mode is not supported on LES-legacy branch. ' + 'LES-legacy uses EdgePreprocess (serial) for unified pos/strain ' + 'gradients. Use the LES branch for training / parallel deployment.' + ) + layers = OrderedDict() + # Legacy serial LES: insert EdgePreprocess before edge_embedding so that + # pos → EDGE_VEC is live in the autograd graph. The strain leaf created + # here is shared by LatentEwaldSum (cell) and ForceStressOutput (forces + + # complete stress), replacing the separate Path-2 / LES_STRAIN approach. + if config.get(KEY.USE_LES, False): + layers['edge_preprocess'] = EdgePreprocess(is_stress=True) + cutoff = config[KEY.CUTOFF] num_species = config[KEY.NUM_SPECIES] feature_multiplicity = config[KEY.NODE_FEATURE_MULTIPLICITY] @@ -599,19 +619,65 @@ def build_E3_equivariant_model( layers.update(interaction_builder(**param_interaction_block)) irreps_x = irreps_out + if config.get(KEY.USE_LES, False): + les_cfg = config.get(KEY.LES_CONFIG, {}) + # Latent charge readout must sit BEFORE init_feature_reduce because + # reduce_input_to_hidden overwrites KEY.NODE_FEATURE in-place. + # Reading here guarantees we use the full (pre-projection) node features. + layers['les_charge_readout'] = LatentChargeReadout( + irreps_in=irreps_x, # type: ignore + data_key_in=KEY.NODE_FEATURE, + data_key_out=KEY.LES_Q, + n_charges=les_cfg.get('n_charges', 1), + hidden_channels=les_cfg.get('hidden_channels', None), + zero_init=les_cfg.get('zero_init', False), + ) + layers.update(init_feature_reduce(config, irreps_x)) # type: ignore - layers.update( - { - 'rescale_atomic_energy': init_shift_scale(config), - 'reduce_total_enegy': AtomReduce( - data_key_in=KEY.ATOMIC_ENERGY, - data_key_out=KEY.PRED_TOTAL_ENERGY, - ), - } - ) + if config.get(KEY.USE_LES, False): + layers.update( + { + 'rescale_atomic_energy': init_shift_scale(config), + # SR energy: sum of per-atom (local) energies + 'reduce_sr_energy': AtomReduce( + data_key_in=KEY.ATOMIC_ENERGY, + data_key_out=KEY.SR_ENERGY, + ), + # LR energy: Ewald summation on latent charges + 'les_lr_energy': LatentEwaldSum( + les_args=les_cfg.get('les_args', {'use_atomwise': False}), + data_key_in=KEY.LES_Q, + data_key_out=KEY.LR_ENERGY, + compute_bec=les_cfg.get('compute_bec', False), + bec_output_index=les_cfg.get('bec_output_index', None), + ), + # Total = SR + LR + 'add_lr_to_total': AddLREnergy( + key_sr=KEY.SR_ENERGY, + key_lr=KEY.LR_ENERGY, + data_key_out=KEY.PRED_TOTAL_ENERGY, + ), + } + ) + else: + layers.update( + { + 'rescale_atomic_energy': init_shift_scale(config), + 'reduce_total_enegy': AtomReduce( + data_key_in=KEY.ATOMIC_ENERGY, + data_key_out=KEY.PRED_TOTAL_ENERGY, + ), + } + ) - gradient_module = ForceStressOutputFromEdge() + if config.get(KEY.USE_LES, False): + # Legacy serial mode: unified pos + _strain gradients via EdgePreprocess. + # ForceStressOutput differentiates w.r.t. strained pos (force) and + # _strain (complete stress — SR virial + LR positional + LR cell). + gradient_module = ForceStressOutput() + else: + gradient_module = ForceStressOutputFromEdge() grad_key = gradient_module.get_grad_key() layers.update({'force_output': gradient_module}) diff --git a/sevenn/nn/edge_embedding.py b/sevenn/nn/edge_embedding.py index 86ce5d94..ff4394d9 100644 --- a/sevenn/nn/edge_embedding.py +++ b/sevenn/nn/edge_embedding.py @@ -12,8 +12,13 @@ @compile_mode('script') class EdgePreprocess(nn.Module): """ - preprocessing pos to edge vectors and edge lengths - currently used in sevenn/scripts/deploy for lammps serial model + Computes edge vectors and lengths from atomic positions, with optional + affine strain for stress computation. + + Used exclusively by LES models (LES branch) as the first model + layer, restoring the pos → EDGE_VEC autograd connection so that + ForceStressOutput can recover forces and stress via positional and strain + gradients rather than edge-virial decomposition. """ def __init__(self, is_stress: bool) -> None: @@ -30,9 +35,9 @@ def forward(self, data: AtomGraphDataType) -> AtomGraphDataType: cell_shift = data[KEY.CELL_SHIFT] pos = data[KEY.POS] - batch = data[KEY.BATCH] # for deploy, must be defined first if self.is_stress: if self._is_batch_data: + batch = data[KEY.BATCH] num_batch = int(batch.max().cpu().item()) + 1 strain = torch.zeros( (num_batch, 3, 3), @@ -60,12 +65,25 @@ def forward(self, data: AtomGraphDataType) -> AtomGraphDataType: pos = pos + torch.mm(pos, sym_strain) cell = cell + torch.mm(cell, sym_strain) + # Write strained pos and cell back so that downstream modules + # (LatentEwaldSum, ForceStressOutput) receive tensors that are + # connected to _strain in the autograd graph. This enables a + # single d(E)/d(_strain) call to capture the complete stress + # (SR virial + LR positional + LR cell/k-space) without a + # separate LES_STRAIN leaf inside LatentEwaldSum. + data[KEY.POS] = pos + if self._is_batch_data: + data[KEY.CELL] = cell.reshape(-1, 3) + else: + data[KEY.CELL] = cell + idx_src = data[KEY.EDGE_IDX][0] idx_dst = data[KEY.EDGE_IDX][1] edge_vec = pos[idx_dst] - pos[idx_src] if self._is_batch_data: + batch = data[KEY.BATCH] edge_vec = edge_vec + torch.einsum( 'ni,nij->nj', cell_shift, cell[batch[idx_src]] ) diff --git a/sevenn/nn/les.py b/sevenn/nn/les.py new file mode 100644 index 00000000..c1ab3e20 --- /dev/null +++ b/sevenn/nn/les.py @@ -0,0 +1,233 @@ +""" +LES (Latent Ewald Summation) modules for SevenNet. + +Architecture: + NODE_FEATURE (last conv. layer, all-scalar) + │ + ├─→ [LatentChargeReadout] → LES_Q (N_atoms, n_charges) + │ + └─→ [init_feature_reduce] → ATOMIC_ENERGY → [AtomReduce] → SR_ENERGY + │ + LES_Q ──→ [LatentEwaldSum] ──→ LR_ENERGY ──→ [AddLREnergy] ─────────┘ + │ + PRED_TOTAL_ENERGY + │ + [ForceStressOutput] + +EdgePreprocess (first layer) applies an affine strain to pos and cell and +computes EDGE_VEC from the strained pos, connecting all three to the _strain +leaf. ForceStressOutput then recovers: + Forces: -d(E_total)/d(strained_pos) SR + q-path LR + direct Ewald + Stress: -d(E_total)/d(_strain) SR virial + LR positional + LR cell + +References: + - LES library: https://github.com/ChengUCB/les + - NequIP-LES: https://github.com/ChengUCB/nequip-les +""" +from typing import Optional + +import torch +import torch.nn as nn +from e3nn.o3 import Irreps + +import sevenn._keys as KEY +from sevenn._const import AtomGraphDataType + +from .linear import IrrepsLinear + + +class LatentChargeReadout(nn.Module): + """ + Projects node features to per-atom latent charges. + + Architecture (controlled by ``hidden_channels``): + hidden_channels=[] (default): + irreps_in ──[IrrepsLinear]──► (N, n_charges) + hidden_channels=[H, ...]: + irreps_in ──[IrrepsLinear]──► (N, H) ──[SiLU + nn.Linear]──► (N, n_charges) + + The first layer is SevenNet's IrrepsLinear (a thin wrapper around + e3nn.o3.Linear that operates on AtomGraphData dicts). Modality dependence + flows in through the upstream conv stack's modality-aware features; this + layer does not concatenate the modality one-hot itself. + + Args: + irreps_in: e3nn irreps of the input node features. + n_charges: number of latent charge channels per atom (default 1). + With n_charges > 1 the Ewald energy is the sum of + n_charges independent Coulomb interactions, one per + channel: E_LR = Σ_α E_Coulomb(q^α). + hidden_channels: hidden layer widths, e.g. [128] for one hidden layer. + zero_init: zero-initialise all weights so E_LR = 0 at init. + Useful for transparent-wrapper tests; not for training. + """ + + def __init__( + self, + irreps_in: Irreps, + data_key_in: str = KEY.NODE_FEATURE, + data_key_out: str = KEY.LES_Q, + n_charges: int = 1, + hidden_channels: Optional[list] = None, + zero_init: bool = False, + ): + super().__init__() + self.key_input = data_key_in + self.key_output = data_key_out + self.n_charges = n_charges + + if hidden_channels is None: + hidden_channels = [] + self._hidden_channels = list(hidden_channels) + + first_out = hidden_channels[0] if hidden_channels else n_charges + # Intermediate key only needed when a scalar MLP follows. + self._intermediate_key = ( + f'{data_key_out}_intermediate' if hidden_channels else data_key_out + ) + self.first_linear = IrrepsLinear( + irreps_in=irreps_in, + irreps_out=Irreps(f'{first_out}x0e'), + data_key_in=data_key_in, + data_key_out=self._intermediate_key, + biases=False, + ) + + scalar_layers: list[nn.Module] = [] + if hidden_channels: + dims = hidden_channels + [n_charges] + for i in range(len(dims) - 1): + scalar_layers.append(nn.SiLU()) + scalar_layers.append(nn.Linear(dims[i], dims[i + 1], bias=False)) + self.scalar_mlp = nn.Sequential(*scalar_layers) + + self._zero_init = zero_init + if zero_init: + for m in self.scalar_mlp.modules(): + if isinstance(m, nn.Linear): + nn.init.zeros_(m.weight) + + @property + def layer_instantiated(self) -> bool: + # AtomGraphSequential._instantiate_modules only walks top-level modules, + # so we expose the inner IrrepsLinear's lazy-instantiation status here. + return self.first_linear.layer_instantiated + + def instantiate(self) -> None: + self.first_linear.instantiate() + if self._zero_init: + nn.init.zeros_(self.first_linear.linear.weight) + + def forward(self, data: AtomGraphDataType) -> AtomGraphDataType: + data = self.first_linear(data) + if self._hidden_channels: + data[self.key_output] = self.scalar_mlp(data[self._intermediate_key]) + return data + + +class LatentEwaldSum(nn.Module): + """ + Computes long-range energy via Ewald summation on latent charges. + + Expects EdgePreprocess to have already run, which: + - creates the _strain leaf and connects pos and cell to it + - writes strained pos to data[KEY.POS] + - writes strained cell to data[KEY.CELL] + + ForceStressOutput then differentiates w.r.t. strained pos (forces) and + _strain (complete stress: SR virial + LR positional + LR cell/k-space). + + Args: + les_args: kwargs forwarded to Les(). + data_key_in: per-atom latent charges (N_atoms, n_charges). + data_key_out: per-graph LR energy output. + compute_bec: if True, compute Born effective charges. + bec_output_index: 0/1/2 for x/y/z component of BEC. + """ + + def __init__( + self, + les_args: Optional[dict] = None, + data_key_in: str = KEY.LES_Q, + data_key_out: str = KEY.LR_ENERGY, + compute_bec: bool = False, + bec_output_index: Optional[int] = None, + ): + super().__init__() + try: + from les import Les # https://github.com/ChengUCB/les + except ImportError as e: + raise ImportError( + "The 'les' package is required for LES support. " + "Install it with: pip install git+https://github.com/ChengUCB/les.git" + ) from e + + if les_args is None: + les_args = {'use_atomwise': False} + self.key_input = data_key_in + self.key_output = data_key_out + self.compute_bec = compute_bec + self.bec_output_index = bec_output_index + self.les = Les(les_args) + self._is_batch_data = True # set by AtomGraphSequential.set_is_batch_data() + + def forward(self, data: AtomGraphDataType) -> AtomGraphDataType: + q = data[self.key_input] # (N_atoms, n_charges) + pos = data[KEY.POS] # strained pos from EdgePreprocess + + if self._is_batch_data: + batch = data[KEY.BATCH].long() + n_graphs = int(batch.max().item()) + 1 + else: + batch = torch.zeros(pos.shape[0], dtype=torch.long, device=pos.device) + n_graphs = 1 + + # Batched cell: SevenNet stores (3,3) per graph; PyG stacks to (3*n,3). + # EdgePreprocess wrote the strained cell here, so les() receives a + # tensor connected to _strain for correct stress computation. + if KEY.CELL in data: + cell = data[KEY.CELL].view(-1, 3, 3) # (n_graphs, 3, 3) + else: + cell = torch.zeros((n_graphs, 3, 3), device=pos.device, dtype=pos.dtype) + + les_result = self.les( + latent_charges=q, + positions=pos, + batch=batch, + cell=cell, + compute_energy=True, + compute_bec=self.compute_bec, + bec_output_index=self.bec_output_index, + ) + + e_lr = les_result['E_lr'] # (n_graphs,) + assert e_lr is not None + + # Non-batch mode: squeeze to scalar to match SR_ENERGY from AtomReduce. + data[self.key_output] = e_lr if self._is_batch_data else e_lr.squeeze() + + if self.compute_bec: + bec = les_result.get('BEC') + if bec is not None: + data[KEY.LES_BEC] = bec + + return data + + +class AddLREnergy(nn.Module): + """Adds LR energy to SR energy: PRED_TOTAL_ENERGY = SR_ENERGY + LR_ENERGY.""" + + def __init__( + self, + key_sr: str = KEY.SR_ENERGY, + key_lr: str = KEY.LR_ENERGY, + data_key_out: str = KEY.PRED_TOTAL_ENERGY, + ): + super().__init__() + self.key_sr = key_sr + self.key_lr = key_lr + self.key_output = data_key_out + + def forward(self, data: AtomGraphDataType) -> AtomGraphDataType: + data[self.key_output] = data[self.key_sr] + data[self.key_lr] + return data diff --git a/sevenn/scripts/train.py b/sevenn/scripts/train.py index 28371857..1266e3ba 100644 --- a/sevenn/scripts/train.py +++ b/sevenn/scripts/train.py @@ -78,6 +78,18 @@ def train_v2(config: Dict[str, Any], working_dir: str) -> None: model = build_E3_equivariant_model(config) log.print_model_info(model, config) + # LES fine-tuning: optionally freeze all non-LES parameters so the + # Trainer's requires_grad filter only optimises les_charge_readout and + # les_lr_energy. Must happen before Trainer.from_config captures params. + if config.get(KEY.USE_LES, False) and config.get( + KEY.LES_CONFIG, {} + ).get('freeze_sr', False): + les_module_names = {'les_charge_readout', 'les_lr_energy'} + for name, param in model.named_parameters(): + if name.split('.')[0] not in les_module_names: + param.requires_grad_(False) + log.writeline('LES fine-tuning: SR parameters frozen (freeze_sr=True).') + trainer = Trainer.from_config(model, config) if state_dicts: trainer.load_state_dicts(*state_dicts, strict=False) diff --git a/tests/unit_tests/test_les_finetune.py b/tests/unit_tests/test_les_finetune.py new file mode 100644 index 00000000..5ee9c700 --- /dev/null +++ b/tests/unit_tests/test_les_finetune.py @@ -0,0 +1,524 @@ +""" +Integration test for LES fine-tuning of SevenNet-omni. + +Covers: + 1. build_model_with_les() — loads SR weights, leaves LES params at init + 2. Parameter freezing (freeze_sr=True) — only LES params require grad + 3. Inference pass — no crash, no NaN/Inf in energy/force/stress + 4. Training pass — backward through E_total, only LES params accumulate grad + 5. Optimizer step — LES params update, SR params unchanged +""" + +import warnings + +import pytest +import torch +from ase.build import bulk + +import sevenn._keys as KEY +from sevenn.atom_graph_data import AtomGraphData +from sevenn.checkpoint import SevenNetCheckpoint +from sevenn.train.dataload import unlabeled_atoms_to_graph +from sevenn.util import pretrained_name_to_path + + +# ── fixtures ───────────────────────────────────────────────────────────────── + +@pytest.fixture(scope='module') +def omni_checkpoint(): + path = pretrained_name_to_path('7net-omni') + return SevenNetCheckpoint(path) + + +pytestmark = pytest.mark.skipif( + not torch.cuda.is_available(), reason='FlashTP requires CUDA' +) + +DEVICE = 'cuda' + + +@pytest.fixture(scope='module') +def sr_model(omni_checkpoint): + """Original non-LES SevenNet-omni on CUDA (reference for comparison).""" + model = omni_checkpoint.build_model() + model.set_is_batch_data(False) + return model.to(DEVICE) + + +@pytest.fixture(scope='module') +def les_model(omni_checkpoint): + """LES model built from SevenNet-omni with SR params frozen, on CUDA.""" + with warnings.catch_warnings(record=True): + warnings.simplefilter('always') + model = omni_checkpoint.build_model_with_les( + les_config={ + 'les_args': {'use_atomwise': False}, + 'zero_init': True, # needed for transparent-wrapper tests + }, + freeze_sr=True, + # No backend overrides: inherits FlashTP from the checkpoint + ) + return model.to(DEVICE) + + +@pytest.fixture # function-scoped: each test gets a fresh, un-mutated graph +def nacl_graph(omni_checkpoint): + """NaCl rocksalt structure as an AtomGraphData (with cell for LES).""" + atoms = bulk('NaCl', 'rocksalt', a=5.63) + atoms.rattle(stdev=0.02, seed=0) + cutoff = omni_checkpoint.config['cutoff'] + # with_shift=True → KEY.CELL included in data dict (needed for Ewald) + graph = AtomGraphData.from_numpy_dict( + unlabeled_atoms_to_graph(atoms, cutoff, with_shift=True) + ) + # SevenNet-omni is multimodal; pick one modality for testing + graph[KEY.DATA_MODALITY] = 'omat24' + return graph.to(DEVICE) + + +@pytest.fixture # function-scoped: each test gets a fresh batch +def nacl_batch(nacl_graph): + """Single NaCl graph wrapped as a PyG Batch (is_batch_data=True).""" + from torch_geometric.data import Batch + # Batch.from_data_list adds KEY.BATCH and turns DATA_MODALITY into a list + return Batch.from_data_list([nacl_graph]) + + +# ── tests ───────────────────────────────────────────────────────────────────── + +class TestBuildModelWithLES: + def test_model_type(self, les_model): + from sevenn.nn.sequential import AtomGraphSequential + assert isinstance(les_model, AtomGraphSequential) + + def test_les_modules_present(self, les_model): + module_names = dict(les_model.named_modules()).keys() + assert 'les_charge_readout' in module_names + assert 'les_lr_energy' in module_names + assert 'add_lr_to_total' in module_names + + def test_les_charge_readout_zero_init(self, les_model): + """les_charge_readout should start at zero — E_LR = 0 initially.""" + w = les_model._modules['les_charge_readout'].first_linear.linear.weight + assert torch.allclose(w, torch.zeros_like(w)), \ + 'les_charge_readout.first_linear.linear.weight should be zero-initialised' + + def test_sr_structure_preserved(self, les_model, omni_checkpoint): + """All SR param values must match the original checkpoint exactly.""" + orig_sd = omni_checkpoint.model_state_dict # CPU tensors (raw checkpoint) + les_sd = les_model.state_dict() # CUDA tensors (model on GPU) + for key, orig_val in orig_sd.items(): + assert key in les_sd, f'SR key {key!r} missing from LES model' + assert torch.allclose(les_sd[key].cpu().float(), orig_val.float()), \ + f'SR param {key!r} changed during build_model_with_les' + + def test_is_batch_data_propagates(self, les_model): + """set_is_batch_data must reach top-level LES modules. + + Regression guard for the IrrepsLinear-based readout: removing the + _is_batch_data property from LatentChargeReadout must not break + propagation to other top-level layers. + """ + les_model.set_is_batch_data(False) + for name in ('edge_preprocess', 'les_lr_energy', 'force_output'): + assert les_model._modules[name]._is_batch_data is False, ( + f'{name}._is_batch_data not set to False' + ) + + les_model.set_is_batch_data(True) + for name in ('edge_preprocess', 'les_lr_energy', 'force_output'): + assert les_model._modules[name]._is_batch_data is True, ( + f'{name}._is_batch_data not restored to True' + ) + + def test_state_dict_round_trip(self, les_model, omni_checkpoint, nacl_graph): + """state_dict() → load_state_dict() must reproduce identical output.""" + sd = les_model.state_dict() + model2 = omni_checkpoint.build_model_with_les( + les_config={ + 'les_args': {'use_atomwise': False}, + 'zero_init': True, + }, + freeze_sr=True, + ).to(DEVICE) + missing, unexpected = model2.load_state_dict(sd, strict=True) + assert not missing, f'Unexpected missing keys: {missing}' + assert not unexpected, f'Unexpected extra keys: {unexpected}' + + for m in (les_model, model2): + m.eval() + m.set_is_batch_data(False) + out1 = les_model(nacl_graph.clone()) + out2 = model2(nacl_graph.clone()) + assert torch.allclose( + out1[KEY.PRED_TOTAL_ENERGY], out2[KEY.PRED_TOTAL_ENERGY], atol=1e-6 + ) + assert torch.allclose( + out1[KEY.PRED_FORCE], out2[KEY.PRED_FORCE], atol=1e-6 + ) + les_model.set_is_batch_data(True) # restore + + +class TestParameterFreezing: + def test_les_params_require_grad(self, les_model): + les_prefixes = ('les_charge_readout.', 'les_lr_energy.') + les_trainable = [ + name for name, p in les_model.named_parameters() + if name.startswith(les_prefixes) and p.requires_grad + ] + assert len(les_trainable) > 0, 'No trainable LES parameters found' + + def test_sr_params_frozen(self, les_model): + les_prefixes = ('les_charge_readout.', 'les_lr_energy.') + sr_trainable = [ + name for name, p in les_model.named_parameters() + if not name.startswith(les_prefixes) and p.requires_grad + ] + assert len(sr_trainable) == 0, \ + f'SR params should be frozen; found trainable: {sr_trainable}' + + def test_optimizer_only_has_les_params(self, les_model): + trainable = [p for p in les_model.parameters() if p.requires_grad] + assert len(trainable) > 0 + # Verify these are exactly the LES params + les_prefixes = ('les_charge_readout.', 'les_lr_energy.') + for name, p in les_model.named_parameters(): + if p.requires_grad: + assert name.startswith(les_prefixes), \ + f'Unexpected trainable param: {name}' + + +class TestInference: + """Model in eval mode, single-graph (is_batch_data=False, like the calculator).""" + + @pytest.fixture(autouse=True) + def set_single_graph_mode(self, les_model): + les_model.set_is_batch_data(False) + yield + les_model.set_is_batch_data(True) # restore for training tests + + def test_forward_no_crash(self, les_model, nacl_graph): + # Note: no torch.no_grad() — SevenNet calls torch.autograd.grad + # internally to compute forces, so grad computation must stay enabled. + les_model.eval() + out = les_model(nacl_graph) + assert KEY.PRED_TOTAL_ENERGY in out + assert KEY.PRED_FORCE in out + + def test_energy_finite(self, les_model, nacl_graph): + les_model.eval() + out = les_model(nacl_graph) + e = out[KEY.PRED_TOTAL_ENERGY] + assert torch.isfinite(e).all(), f'Energy contains NaN/Inf: {e}' + + def test_force_finite(self, les_model, nacl_graph): + les_model.eval() + out = les_model(nacl_graph) + f = out[KEY.PRED_FORCE] + assert torch.isfinite(f).all(), f'Force contains NaN/Inf: {f}' + + def test_force_shape(self, les_model, nacl_graph): + les_model.eval() + out = les_model(nacl_graph) + n_atoms = int(nacl_graph[KEY.NUM_ATOMS].item()) + assert out[KEY.PRED_FORCE].shape == (n_atoms, 3) + + def test_stress_finite(self, les_model, nacl_graph): + les_model.eval() + out = les_model(nacl_graph) + if KEY.PRED_STRESS in out: + s = out[KEY.PRED_STRESS] + assert torch.isfinite(s).all(), f'Stress contains NaN/Inf: {s}' + + def test_lr_energy_zero_at_init(self, les_model, nacl_graph): + """With zero-init les_charge_readout, LR charges = 0 → E_LR ≈ 0. + + Also verifies that LR_ENERGY is a scalar in non-batch mode so that + AddLREnergy(SR_ENERGY + LR_ENERGY) does not silently broadcast. + """ + les_model.eval() + out = les_model(nacl_graph) + e_lr = out.get(KEY.LR_ENERGY) + if e_lr is not None: + assert e_lr.shape == (), ( + f'LR_ENERGY should be scalar in non-batch mode, ' + f'got shape {tuple(e_lr.shape)}' + ) + assert torch.allclose(e_lr, torch.zeros_like(e_lr), atol=1e-6), \ + f'E_LR should be ~0 with zero-init charges, got {e_lr}' + + def test_total_energy_equals_sr_at_init(self, les_model, nacl_graph): + """With zero charges, PRED_TOTAL_ENERGY == SR_ENERGY.""" + les_model.eval() + out = les_model(nacl_graph) + e_total = out[KEY.PRED_TOTAL_ENERGY] + e_sr = out.get(KEY.SR_ENERGY) + if e_sr is not None: + assert torch.allclose(e_total, e_sr, atol=1e-5), \ + 'Total energy should equal SR energy when LR charges are zero' + + def test_multiple_forward_passes_stable(self, les_model, nacl_graph, + omni_checkpoint): + """retain_graph bug would crash on second inference pass. + Use a second fresh graph (function-scoped fixture gives one per call, + but we need two here so we build the second manually).""" + atoms = bulk('NaCl', 'rocksalt', a=5.63) + atoms.rattle(stdev=0.02, seed=0) + cutoff = omni_checkpoint.config['cutoff'] + graph2 = AtomGraphData.from_numpy_dict( + unlabeled_atoms_to_graph(atoms, cutoff, with_shift=True) + ) + graph2[KEY.DATA_MODALITY] = 'omat24' + graph2 = graph2.to(DEVICE) + + les_model.eval() + out1 = les_model(nacl_graph) + out2 = les_model(graph2) + assert torch.allclose( + out1[KEY.PRED_TOTAL_ENERGY], out2[KEY.PRED_TOTAL_ENERGY] + ), 'Two identical forward passes give different energies' + + +class TestTraining: + """Model in train mode with batched data (is_batch_data=True, like training).""" + + def test_backward_no_crash(self, les_model, nacl_batch): + les_model.train() + out = les_model(nacl_batch) + loss = out[KEY.PRED_TOTAL_ENERGY].sum() + loss.backward() # must not raise RuntimeError + + def test_only_les_params_get_grad(self, les_model, nacl_batch): + """After backward, only LES params should have .grad set.""" + les_model.train() + # Zero out any stale grads from previous tests + les_model.zero_grad() + out = les_model(nacl_batch) + loss = out[KEY.PRED_TOTAL_ENERGY].sum() + loss.backward() + + les_prefixes = ('les_charge_readout.', 'les_lr_energy.') + for name, param in les_model.named_parameters(): + if name.startswith(les_prefixes): + # LES params must have a gradient + assert param.grad is not None, \ + f'LES param {name!r} has no grad after backward' + else: + # SR params should have no gradient (frozen) + assert param.grad is None, \ + f'SR param {name!r} has grad despite being frozen: {param.grad}' + + def test_optimizer_step_updates_les_only(self, les_model, nacl_batch): + """One Adam step should move LES params and leave SR params unchanged.""" + les_model.train() + les_model.zero_grad() + + # Snapshot SR params before step + sr_snapshot = { + name: p.data.clone() + for name, p in les_model.named_parameters() + if not p.requires_grad + } + + # Perturb les_charge_readout weights so grad is non-zero + # (They start at zero, so the *gradient* of loss w.r.t. them may be zero + # at q=0, but the Les() params should receive non-zero grads.) + with torch.no_grad(): + les_model._modules['les_charge_readout'].first_linear.linear.weight.fill_(0.01) + + out = les_model(nacl_batch) + loss = out[KEY.PRED_TOTAL_ENERGY].sum() + loss.backward() + + trainable = [p for p in les_model.parameters() if p.requires_grad] + opt = torch.optim.Adam(trainable, lr=1e-3) + opt.step() + + # SR params must be unchanged + for name, orig in sr_snapshot.items(): + cur = dict(les_model.named_parameters())[name].data + assert torch.allclose(cur, orig), \ + f'SR param {name!r} changed after optimizer step' + + # Restore les_charge_readout weights to zero for subsequent tests + with torch.no_grad(): + les_model._modules['les_charge_readout'].first_linear.linear.weight.zero_() + + def test_force_computed_in_train_mode(self, les_model, nacl_batch): + les_model.train() + les_model.zero_grad() + out = les_model(nacl_batch) + assert KEY.PRED_FORCE in out + assert torch.isfinite(out[KEY.PRED_FORCE]).all() + + def test_stress_computed_in_train_mode(self, les_model, nacl_batch): + les_model.train() + les_model.zero_grad() + out = les_model(nacl_batch) + if KEY.PRED_STRESS in out: + assert torch.isfinite(out[KEY.PRED_STRESS]).all() + + +class TestZeroInitEquivalence: + """ + Core correctness test: with zero-init les_charge_readout (E_LR = 0), the + LES model must produce energy, forces, and stress consistent with the + original non-LES checkpoint. + + Note on tolerance: the LES-legacy model uses ForceStressOutput (positional + gradient via EdgePreprocess), while the SR model uses ForceStressOutputFromEdge + (edge virial). These are mathematically equivalent but not bit-identical in + float32, so atol is relaxed to 1e-4 for forces/stress. + """ + + @pytest.fixture(autouse=True) + def single_graph_mode(self, les_model, sr_model): + les_model.eval() + les_model.set_is_batch_data(False) + sr_model.eval() + yield + les_model.set_is_batch_data(True) + + def test_energy_matches_original(self, les_model, sr_model, nacl_graph): + out_sr = sr_model(nacl_graph) + # nacl_graph is function-scoped so this is a fresh copy + out_les = les_model(nacl_graph) + assert torch.allclose( + out_les[KEY.PRED_TOTAL_ENERGY], + out_sr[KEY.PRED_TOTAL_ENERGY], + atol=1e-5, + ), ( + f'Energy mismatch: LES={out_les[KEY.PRED_TOTAL_ENERGY].item():.8f}, ' + f'SR={out_sr[KEY.PRED_TOTAL_ENERGY].item():.8f}' + ) + + def test_forces_match_original(self, les_model, sr_model, nacl_graph): + out_sr = sr_model(nacl_graph) + out_les = les_model(nacl_graph) + assert torch.allclose( + out_les[KEY.PRED_FORCE], + out_sr[KEY.PRED_FORCE], + atol=1e-4, + ), ( + f'Force mismatch (max abs diff): ' + f'{(out_les[KEY.PRED_FORCE] - out_sr[KEY.PRED_FORCE]).abs().max().item():.2e}' + ) + + def test_stress_matches_original(self, les_model, sr_model, nacl_graph): + out_sr = sr_model(nacl_graph) + out_les = les_model(nacl_graph) + if KEY.PRED_STRESS not in out_sr or KEY.PRED_STRESS not in out_les: + pytest.skip('Stress not computed for this structure') + assert torch.allclose( + out_les[KEY.PRED_STRESS], + out_sr[KEY.PRED_STRESS], + atol=1e-4, + ), ( + f'Stress mismatch (max abs diff): ' + f'{(out_les[KEY.PRED_STRESS] - out_sr[KEY.PRED_STRESS]).abs().max().item():.2e}' + ) + + +N_CHARGES = 4 # number of latent charge channels used in multi-q tests + + +class TestMultiDimensionalQ: + """ + Tests for multi-channel latent charges (n_charges > 1). + + With n_charges=N_CHARGES the readout maps node features to + (N_atoms, N_CHARGES) and the Ewald energy is the sum of N_CHARGES + independent Coulomb interactions, one per channel. + """ + + @pytest.fixture(scope='class') + def les_model_mq(self, omni_checkpoint): + """LES model with n_charges=N_CHARGES, SR params frozen.""" + with warnings.catch_warnings(record=True): + warnings.simplefilter('always') + model = omni_checkpoint.build_model_with_les( + les_config={ + 'les_args': {'use_atomwise': False}, + 'n_charges': N_CHARGES, + # non-zero init so that charges and LR energy are nonzero + }, + freeze_sr=True, + ) + return model.to(DEVICE) + + # ── shape ──────────────────────────────────────────────────────────────── + + def test_charge_shape(self, les_model_mq, nacl_graph): + """LES_Q must have shape (N_atoms, N_CHARGES).""" + les_model_mq.eval() + les_model_mq.set_is_batch_data(False) + out = les_model_mq(nacl_graph) + q = out[KEY.LES_Q] + n_atoms = int(nacl_graph[KEY.NUM_ATOMS].item()) + assert q.shape == (n_atoms, N_CHARGES), ( + f'Expected LES_Q shape ({n_atoms}, {N_CHARGES}), got {tuple(q.shape)}' + ) + + def test_force_shape(self, les_model_mq, nacl_graph): + """Force shape must be (N_atoms, 3) regardless of n_charges.""" + les_model_mq.eval() + les_model_mq.set_is_batch_data(False) + out = les_model_mq(nacl_graph) + n_atoms = int(nacl_graph[KEY.NUM_ATOMS].item()) + assert out[KEY.PRED_FORCE].shape == (n_atoms, 3) + + # ── inference ──────────────────────────────────────────────────────────── + + def test_energy_finite(self, les_model_mq, nacl_graph): + les_model_mq.eval() + les_model_mq.set_is_batch_data(False) + out = les_model_mq(nacl_graph) + assert torch.isfinite(out[KEY.PRED_TOTAL_ENERGY]).all() + + def test_force_finite(self, les_model_mq, nacl_graph): + les_model_mq.eval() + les_model_mq.set_is_batch_data(False) + out = les_model_mq(nacl_graph) + assert torch.isfinite(out[KEY.PRED_FORCE]).all() + + def test_stress_finite(self, les_model_mq, nacl_graph): + les_model_mq.eval() + les_model_mq.set_is_batch_data(False) + out = les_model_mq(nacl_graph) + if KEY.PRED_STRESS in out: + assert torch.isfinite(out[KEY.PRED_STRESS]).all() + + def test_lr_energy_nonzero(self, les_model_mq, nacl_graph): + """With default (non-zero) init, LR energy should be nonzero.""" + les_model_mq.eval() + les_model_mq.set_is_batch_data(False) + out = les_model_mq(nacl_graph) + e_lr = out.get(KEY.LR_ENERGY) + if e_lr is not None: + assert not torch.allclose(e_lr, torch.zeros_like(e_lr), atol=1e-6), \ + 'E_LR is unexpectedly zero with non-zero-init multi-channel readout' + + # ── training ───────────────────────────────────────────────────────────── + + def test_backward_no_crash(self, les_model_mq, nacl_batch): + les_model_mq.train() + les_model_mq.set_is_batch_data(True) + les_model_mq.zero_grad() + out = les_model_mq(nacl_batch) + out[KEY.PRED_TOTAL_ENERGY].sum().backward() + + def test_only_les_params_get_grad(self, les_model_mq, nacl_batch): + les_model_mq.train() + les_model_mq.set_is_batch_data(True) + les_model_mq.zero_grad() + out = les_model_mq(nacl_batch) + out[KEY.PRED_TOTAL_ENERGY].sum().backward() + + les_prefixes = ('les_charge_readout.', 'les_lr_energy.') + for name, param in les_model_mq.named_parameters(): + if name.startswith(les_prefixes): + assert param.grad is not None, \ + f'LES param {name!r} has no grad after backward' + else: + assert param.grad is None, \ + f'SR param {name!r} has grad despite being frozen: {param.grad}' diff --git a/tests/unit_tests/test_les_legacy.py b/tests/unit_tests/test_les_legacy.py new file mode 100644 index 00000000..620f8e83 --- /dev/null +++ b/tests/unit_tests/test_les_legacy.py @@ -0,0 +1,602 @@ +""" +Unit tests for the LES-legacy architecture. + +LES-legacy specific design under test: + - EdgePreprocess (first model layer, is_stress=True) restores the + pos -> EDGE_VEC autograd connection so that ForceStressOutput + captures all gradient paths through a single _strain leaf. + - ForceStressOutput (positional + strain gradient, not edge virial) + - LatentEwaldSum reads strained pos/cell written back by EdgePreprocess + +All tests run on CPU with a small model built from scratch. +The 'les' package (https://github.com/ChengUCB/les) must be installed. +""" + +import pytest +import torch +from ase.build import bulk +from torch_geometric.data import Batch + +import sevenn._keys as KEY +import sevenn.train.dataload as dl +from sevenn.atom_graph_data import AtomGraphData +from sevenn.model_build import build_E3_equivariant_model +from sevenn.nn.edge_embedding import EdgePreprocess +from sevenn.nn.force_output import ForceStressOutput +from sevenn.util import chemical_species_preprocess + +# ── skip if les not installed ───────────────────────────────────────────────── + +try: + import les as _les_pkg # noqa: F401 + HAS_LES = True +except ImportError: + HAS_LES = False + +pytestmark = pytest.mark.skipif(not HAS_LES, reason='les package not installed') + +# ── constants ───────────────────────────────────────────────────────────────── + +CUTOFF = 4.0 +DELTA = 5e-4 # Angstrom, perturbation for numerical gradient checks +ATOL_FD = 1e-2 # tolerance for FD vs autograd comparison + + +# ── config helpers ──────────────────────────────────────────────────────────── + +def _base_config(): + """Minimal SevenNet config for fast CPU testing.""" + config = { + 'cutoff': CUTOFF, + 'channel': 4, + 'radial_basis': {'radial_basis_name': 'bessel'}, + 'cutoff_function': {'cutoff_function_name': 'poly_cut'}, + 'interaction_type': 'nequip', + 'lmax': 1, + 'is_parity': True, + 'num_convolution_layer': 2, + 'weight_nn_hidden_neurons': [16], + 'act_radial': 'silu', + 'act_scalar': {'e': 'silu', 'o': 'tanh'}, + 'act_gate': {'e': 'silu', 'o': 'tanh'}, + 'conv_denominator': 10.0, + 'train_denominator': False, + 'self_connection_type': 'nequip', + 'shift': 0.0, + 'scale': 1.0, + 'train_shift_scale': False, + 'irreps_manual': False, + 'lmax_edge': -1, + 'lmax_node': -1, + 'readout_as_fcn': False, + 'use_bias_in_linear': False, + '_normalize_sph': True, + } + config.update(**chemical_species_preprocess(['Na', 'Cl'])) + return config + + +def _les_config(zero_init=True, n_charges=1): + cfg = _base_config() + cfg['use_les'] = True + cfg['les_config'] = { + 'les_args': {'use_atomwise': False}, + 'n_charges': n_charges, + 'zero_init': zero_init, + } + return cfg + + +# ── fixtures ────────────────────────────────────────────────────────────────── + +@pytest.fixture(scope='module') +def nacl_atoms(): + atoms = bulk('NaCl', 'rocksalt', a=5.63) + atoms.rattle(stdev=0.01, seed=42) + return atoms + + +@pytest.fixture(scope='module') +def nacl_graph(nacl_atoms): + """NaCl AtomGraphData with cell info for Ewald summation.""" + return AtomGraphData.from_numpy_dict( + dl.unlabeled_atoms_to_graph(nacl_atoms, CUTOFF, with_shift=True) + ) + + +@pytest.fixture(scope='module') +def les_model_zero(): + """LES model, zero-init charges -> E_LR = 0 at construction.""" + return build_E3_equivariant_model(_les_config(zero_init=True), parallel=False) + + +@pytest.fixture(scope='module') +def les_model(): + """LES model with non-zero charges -> E_LR != 0.""" + return build_E3_equivariant_model(_les_config(zero_init=False), parallel=False) + + +# ── graph helpers ────────────────────────────────────────────────────────────── + +def _fresh(graph): + """Clone graph for a single forward pass (EdgePreprocess mutates data).""" + return graph.clone() + + +def _run(model, graph, batch=False): + model.eval() + model.set_is_batch_data(batch) + return model(_fresh(graph)) + + +def _energy(model, graph): + return _run(model, graph)[KEY.PRED_TOTAL_ENERGY].item() + + +def _perturbed(graph, atom_idx, direction, delta): + """Graph with one Cartesian coordinate shifted by delta.""" + g = _fresh(graph) + pos = g[KEY.POS].clone() + pos[atom_idx, direction] += delta + g[KEY.POS] = pos + return g + + +def _strained(graph, alpha, beta, delta): + """ + Graph with symmetric strain delta applied to the (alpha, beta) component. + + Applies new_pos = pos + pos @ eps_sym, new_cell = cell + cell @ eps_sym. + CELL_SHIFT (integer PBC shifts) is unchanged for small delta. + CELL_VOLUME is updated for correct stress normalisation. + """ + g = _fresh(graph) + pos = g[KEY.POS].clone().float() + cell = g[KEY.CELL].view(3, 3).clone().float() + + eps = torch.zeros(3, 3) + eps[alpha, beta] += 0.5 * delta + eps[beta, alpha] += 0.5 * delta # symmetric + + g[KEY.POS] = pos + pos @ eps + new_cell = cell + cell @ eps + g[KEY.CELL] = new_cell + g[KEY.CELL_VOLUME] = torch.det(new_cell).abs() + return g + + +# ── architecture tests ───────────────────────────────────────────────────────── + +class TestLESLegacyArchitecture: + """Verify that build_E3_equivariant_model produces the correct layer structure.""" + + def test_edge_preprocess_is_first_layer(self, les_model): + first_name, first_mod = next(iter(les_model._modules.items())) + assert first_name == 'edge_preprocess' + assert isinstance(first_mod, EdgePreprocess) + + def test_edge_preprocess_is_stress_true(self, les_model): + assert les_model._modules['edge_preprocess'].is_stress is True + + def test_force_output_is_ForceStressOutput(self, les_model): + fo = les_model._modules['force_output'] + assert isinstance(fo, ForceStressOutput) + + def test_les_modules_present(self, les_model): + names = set(les_model._modules.keys()) + for expected in ('les_charge_readout', 'les_lr_energy', 'add_lr_to_total'): + assert expected in names, f'Missing module: {expected}' + + def test_sr_energy_reduce_present(self, les_model): + # LES uses reduce_sr_energy, not the non-LES reduce_total_enegy + assert 'reduce_sr_energy' in les_model._modules + assert 'reduce_total_enegy' not in les_model._modules + + def test_parallel_raises(self): + with pytest.raises(NotImplementedError): + build_E3_equivariant_model(_les_config(), parallel=True) + + def test_strain_leaf_created_during_forward(self, les_model, nacl_graph): + """EdgePreprocess must write _strain to data on every forward pass.""" + les_model.eval() + les_model.set_is_batch_data(False) + g = _fresh(nacl_graph) + les_model(g) + assert '_strain' in g, '_strain leaf not found in data after forward' + # Must be a true autograd leaf so that d(E)/d(_strain) is well-defined. + assert g['_strain'].requires_grad + assert g['_strain'].is_leaf + + def test_is_batch_data_propagates(self, les_model): + """set_is_batch_data must reach top-level LES modules and EdgePreprocess.""" + les_model.set_is_batch_data(False) + for name in ('edge_preprocess', 'les_lr_energy', 'force_output'): + assert les_model._modules[name]._is_batch_data is False, ( + f'{name}._is_batch_data not set to False' + ) + + les_model.set_is_batch_data(True) + for name in ('edge_preprocess', 'les_lr_energy', 'force_output'): + assert les_model._modules[name]._is_batch_data is True, ( + f'{name}._is_batch_data not restored to True' + ) + + def test_state_dict_round_trip(self): + """save → load → forward must reproduce original output bit-for-bit.""" + model1 = build_E3_equivariant_model( + _les_config(zero_init=False), parallel=False + ) + model2 = build_E3_equivariant_model( + _les_config(zero_init=False), parallel=False + ) + # Different random init → outputs would normally differ. + sd = model1.state_dict() + missing, unexpected = model2.load_state_dict(sd, strict=True) + assert not missing, f'Unexpected missing keys: {missing}' + assert not unexpected, f'Unexpected extra keys: {unexpected}' + + atoms = bulk('NaCl', 'rocksalt', a=5.63) + graph = AtomGraphData.from_numpy_dict( + dl.unlabeled_atoms_to_graph(atoms, CUTOFF, with_shift=True) + ) + for m in (model1, model2): + m.eval() + m.set_is_batch_data(False) + out1 = model1(_fresh(graph)) + out2 = model2(_fresh(graph)) + assert torch.allclose( + out1[KEY.PRED_TOTAL_ENERGY], out2[KEY.PRED_TOTAL_ENERGY], atol=1e-6 + ) + assert torch.allclose(out1[KEY.PRED_FORCE], out2[KEY.PRED_FORCE], atol=1e-6) + assert torch.allclose( + out1[KEY.PRED_STRESS], out2[KEY.PRED_STRESS], atol=1e-6 + ) + + +# ── non-batch inference ──────────────────────────────────────────────────────── + +class TestNonBatchInference: + + @pytest.fixture(autouse=True) + def setup(self, les_model): + les_model.eval() + les_model.set_is_batch_data(False) + + def test_energy_finite(self, les_model, nacl_graph): + out = _run(les_model, nacl_graph) + assert torch.isfinite(out[KEY.PRED_TOTAL_ENERGY]) + + def test_force_finite(self, les_model, nacl_graph): + out = _run(les_model, nacl_graph) + assert torch.isfinite(out[KEY.PRED_FORCE]).all() + + def test_stress_finite(self, les_model, nacl_graph): + out = _run(les_model, nacl_graph) + assert torch.isfinite(out[KEY.PRED_STRESS]).all() + + def test_energy_shape(self, les_model, nacl_graph): + out = _run(les_model, nacl_graph) + assert out[KEY.PRED_TOTAL_ENERGY].shape == () + + def test_force_shape(self, les_model, nacl_graph): + out = _run(les_model, nacl_graph) + n = int(nacl_graph[KEY.NUM_ATOMS].item()) + assert out[KEY.PRED_FORCE].shape == (n, 3) + + def test_stress_shape(self, les_model, nacl_graph): + out = _run(les_model, nacl_graph) + assert out[KEY.PRED_STRESS].shape == (6,) + + def test_total_equals_sr_plus_lr(self, les_model, nacl_graph): + out = _run(les_model, nacl_graph) + assert torch.allclose( + out[KEY.PRED_TOTAL_ENERGY], + out[KEY.SR_ENERGY] + out[KEY.LR_ENERGY], + atol=1e-6, + ) + + def test_lr_energy_is_scalar(self, les_model, nacl_graph): + """Non-batch LR_ENERGY must be a scalar tensor so AddLREnergy doesn't + accidentally broadcast a (1,) tensor onto the scalar SR_ENERGY.""" + out = _run(les_model, nacl_graph) + assert out[KEY.LR_ENERGY].shape == (), ( + f'LR_ENERGY should be a scalar, got shape {tuple(out[KEY.LR_ENERGY].shape)}' + ) + assert out[KEY.SR_ENERGY].shape == () + + +# ── batch inference ──────────────────────────────────────────────────────────── + +class TestBatchInference: + + @pytest.fixture + def batch(self, nacl_graph): + return Batch.from_data_list([_fresh(nacl_graph), _fresh(nacl_graph)]) + + @pytest.fixture(autouse=True) + def setup(self, les_model): + les_model.eval() + les_model.set_is_batch_data(True) + + def test_energy_shape(self, les_model, batch): + out = les_model(batch) + assert out[KEY.PRED_TOTAL_ENERGY].shape == (2,) + + def test_force_shape(self, les_model, batch, nacl_graph): + out = les_model(batch) + n = int(nacl_graph[KEY.NUM_ATOMS].item()) + assert out[KEY.PRED_FORCE].shape == (2 * n, 3) + + def test_stress_shape(self, les_model, batch): + out = les_model(batch) + assert out[KEY.PRED_STRESS].shape == (2, 6) + + def test_all_finite(self, les_model, batch): + out = les_model(batch) + for key in (KEY.PRED_TOTAL_ENERGY, KEY.PRED_FORCE, KEY.PRED_STRESS): + assert torch.isfinite(out[key]).all(), f'{key} contains NaN/Inf' + + +# ── batch == sequential consistency ─────────────────────────────────────────── + +class TestBatchConsistency: + """Batch output must match running the same graph twice in non-batch mode.""" + + def _seq_outputs(self, model, graph): + model.eval() + model.set_is_batch_data(False) + o1 = _run(model, graph) + o2 = _run(model, graph) + return o1, o2 + + def _batch_output(self, model, graph): + model.eval() + model.set_is_batch_data(True) + batch = Batch.from_data_list([_fresh(graph), _fresh(graph)]) + return model(batch) + + def test_energy_consistent(self, les_model, nacl_graph): + o1, o2 = self._seq_outputs(les_model, nacl_graph) + ob = self._batch_output(les_model, nacl_graph) + e_seq = torch.stack([o1[KEY.PRED_TOTAL_ENERGY], o2[KEY.PRED_TOTAL_ENERGY]]) + assert torch.allclose(e_seq, ob[KEY.PRED_TOTAL_ENERGY], atol=1e-5) + + def test_force_consistent(self, les_model, nacl_graph): + o1, o2 = self._seq_outputs(les_model, nacl_graph) + ob = self._batch_output(les_model, nacl_graph) + f_seq = torch.cat([o1[KEY.PRED_FORCE], o2[KEY.PRED_FORCE]]) + assert torch.allclose(f_seq, ob[KEY.PRED_FORCE], atol=1e-5) + + def test_stress_consistent(self, les_model, nacl_graph): + o1, o2 = self._seq_outputs(les_model, nacl_graph) + ob = self._batch_output(les_model, nacl_graph) + s_seq = torch.stack([o1[KEY.PRED_STRESS], o2[KEY.PRED_STRESS]]) + assert torch.allclose(s_seq, ob[KEY.PRED_STRESS], atol=1e-5) + + +# ── training backward ────────────────────────────────────────────────────────── + +class TestTraining: + + @pytest.fixture(autouse=True) + def setup(self, les_model): + les_model.train() + les_model.set_is_batch_data(False) + les_model.zero_grad() + + def test_backward_energy(self, les_model, nacl_graph): + out = les_model(_fresh(nacl_graph)) + out[KEY.PRED_TOTAL_ENERGY].sum().backward() + + def test_backward_force(self, les_model, nacl_graph): + """Force loss requires create_graph=True inside ForceStressOutput.""" + out = les_model(_fresh(nacl_graph)) + loss = out[KEY.PRED_TOTAL_ENERGY].sum() + out[KEY.PRED_FORCE].sum() + loss.backward() + + def test_backward_stress(self, les_model, nacl_graph): + out = les_model(_fresh(nacl_graph)) + loss = (out[KEY.PRED_TOTAL_ENERGY].sum() + + out[KEY.PRED_FORCE].sum() + + out[KEY.PRED_STRESS].sum()) + loss.backward() + + def test_params_receive_grad(self, les_model, nacl_graph): + out = les_model(_fresh(nacl_graph)) + (out[KEY.PRED_TOTAL_ENERGY].sum() + out[KEY.PRED_FORCE].sum()).backward() + params_with_grad = [n for n, p in les_model.named_parameters() + if p.grad is not None and p.grad.abs().max() > 0] + assert len(params_with_grad) > 0, 'No parameters received a non-zero gradient' + + +# ── numerical gradient: forces ───────────────────────────────────────────────── + +class TestNumericalForce: + """ + PRED_FORCE must equal -dE/d(pos) via central finite differences. + + This validates that EdgePreprocess correctly restores the pos -> EDGE_VEC + autograd connection so ForceStressOutput captures all gradient paths + (SR through EDGE_VEC chain + LR direct from les()) in one call. + """ + + @pytest.fixture(autouse=True) + def setup(self, les_model): + les_model.eval() + les_model.set_is_batch_data(False) + + def _fd_force(self, model, graph, atom_idx, direction): + """Central-difference: -(E(pos+δ) - E(pos-δ)) / 2δ.""" + ep = model(_perturbed(graph, atom_idx, direction, +DELTA))[KEY.PRED_TOTAL_ENERGY].item() + em = model(_perturbed(graph, atom_idx, direction, -DELTA))[KEY.PRED_TOTAL_ENERGY].item() + return -(ep - em) / (2 * DELTA) + + @pytest.mark.parametrize('atom_idx,direction', [(0, 0), (0, 1), (1, 2)]) + def test_force_vs_fd(self, les_model, nacl_graph, atom_idx, direction): + fd = self._fd_force(les_model, nacl_graph, atom_idx, direction) + f_model = _run(les_model, nacl_graph)[KEY.PRED_FORCE][atom_idx, direction].item() + assert abs(fd - f_model) < ATOL_FD, ( + f'Force mismatch at atom {atom_idx} dir {direction}: ' + f'FD={fd:.6f} model={f_model:.6f}' + ) + + +# ── numerical gradient: stress ───────────────────────────────────────────────── + +class TestNumericalStress: + """ + PRED_STRESS must equal -(1/V) dE/dε via central finite differences. + + The _strain leaf created by EdgePreprocess must capture SR virial + + LR positional + LR cell contributions so that the total stress is + reproduced by a single d(E)/d(_strain) call. + """ + + # Voigt index -> (alpha, beta) + VOIGT = [(0, 0), (1, 1), (2, 2), (0, 1), (1, 2), (0, 2)] + + @pytest.fixture(autouse=True) + def setup(self, les_model): + les_model.eval() + les_model.set_is_batch_data(False) + + def _fd_stress(self, model, graph, voigt_idx): + """Central-difference stress: -(E(+δ) - E(-δ)) / (2δ V_0).""" + alpha, beta = self.VOIGT[voigt_idx] + ep = model(_strained(graph, alpha, beta, +DELTA))[KEY.PRED_TOTAL_ENERGY].item() + em = model(_strained(graph, alpha, beta, -DELTA))[KEY.PRED_TOTAL_ENERGY].item() + vol = graph[KEY.CELL_VOLUME].item() + return -(ep - em) / (2 * DELTA * vol) + + @pytest.mark.parametrize('voigt_idx', [0, 1, 2]) # diagonal (xx, yy, zz) + def test_stress_vs_fd(self, les_model, nacl_graph, voigt_idx): + fd = self._fd_stress(les_model, nacl_graph, voigt_idx) + s_model = _run(les_model, nacl_graph)[KEY.PRED_STRESS][voigt_idx].item() + assert abs(fd - s_model) < ATOL_FD, ( + f'Stress mismatch at Voigt index {voigt_idx}: ' + f'FD={fd:.6f} model={s_model:.6f}' + ) + + +# ── SR isolation: zero-init ──────────────────────────────────────────────────── + +class TestSRIsolation: + """ + With zero-init charges (LES_Q = 0 -> E_LR = 0), PRED_TOTAL_ENERGY + must equal SR_ENERGY. Validates that LES adds no spurious energy at init. + """ + + @pytest.fixture(autouse=True) + def setup(self, les_model_zero): + les_model_zero.eval() + les_model_zero.set_is_batch_data(False) + + def test_lr_energy_zero(self, les_model_zero, nacl_graph): + out = _run(les_model_zero, nacl_graph) + e_lr = out[KEY.LR_ENERGY] + assert torch.allclose(e_lr, torch.zeros_like(e_lr), atol=1e-6), \ + f'E_LR should be 0 with zero-init charges, got {e_lr.item()}' + + def test_total_equals_sr(self, les_model_zero, nacl_graph): + out = _run(les_model_zero, nacl_graph) + assert torch.allclose(out[KEY.PRED_TOTAL_ENERGY], out[KEY.SR_ENERGY], atol=1e-6) + + +# ── LR contribution to force/stress ──────────────────────────────────────────── + +class TestLRContribution: + """ + With non-zero charges, the LR Ewald term must produce non-trivial + contributions to forces and stress. Compares against a zero-init + (SR-only) model that uses the same SR weights. + + Catches bugs where the LR gradient path is silently broken (e.g. + LatentEwaldSum not connected to the autograd graph, or pos write-back + from EdgePreprocess missing). + """ + + @pytest.fixture(scope='class') + def shared_config(self): + return _les_config(zero_init=False, n_charges=1) + + @pytest.fixture(scope='class') + def les_model_lr(self, shared_config): + return build_E3_equivariant_model(shared_config, parallel=False) + + @pytest.fixture(scope='class') + def les_model_sr_only(self, shared_config, les_model_lr): + """Same SR weights as les_model_lr, but charges forced to zero.""" + cfg = dict(shared_config) + cfg['les_config'] = {**cfg['les_config'], 'zero_init': True} + m = build_E3_equivariant_model(cfg, parallel=False) + # Copy SR weights so only the LR term differs. + src_sd = les_model_lr.state_dict() + dst_sd = m.state_dict() + for k in dst_sd: + if not k.startswith(('les_charge_readout.', 'les_lr_energy.')): + dst_sd[k] = src_sd[k] + m.load_state_dict(dst_sd, strict=True) + return m + + def test_lr_changes_force(self, les_model_lr, les_model_sr_only, nacl_graph): + for m in (les_model_lr, les_model_sr_only): + m.eval() + m.set_is_batch_data(False) + f_lr = les_model_lr(_fresh(nacl_graph))[KEY.PRED_FORCE] + f_sr = les_model_sr_only(_fresh(nacl_graph))[KEY.PRED_FORCE] + diff = (f_lr - f_sr).abs().max().item() + assert diff > 1e-5, ( + f'LR term should contribute non-trivially to forces; max diff = {diff:.2e}' + ) + + def test_lr_changes_stress(self, les_model_lr, les_model_sr_only, nacl_graph): + for m in (les_model_lr, les_model_sr_only): + m.eval() + m.set_is_batch_data(False) + s_lr = les_model_lr(_fresh(nacl_graph))[KEY.PRED_STRESS] + s_sr = les_model_sr_only(_fresh(nacl_graph))[KEY.PRED_STRESS] + diff = (s_lr - s_sr).abs().max().item() + assert diff > 1e-5, ( + f'LR term should contribute non-trivially to stress; max diff = {diff:.2e}' + ) + + +# ── multi-charge channels ────────────────────────────────────────────────────── + +class TestMultiCharge: + """n_charges > 1: LES_Q shape and basic sanity.""" + + N_CHARGES = 3 + + @pytest.fixture(scope='class') + def les_model_mq(self): + return build_E3_equivariant_model( + _les_config(zero_init=False, n_charges=self.N_CHARGES), parallel=False + ) + + def test_charge_shape(self, les_model_mq, nacl_graph): + les_model_mq.eval() + les_model_mq.set_is_batch_data(False) + out = les_model_mq(_fresh(nacl_graph)) + n = int(nacl_graph[KEY.NUM_ATOMS].item()) + assert out[KEY.LES_Q].shape == (n, self.N_CHARGES) + + def test_energy_finite(self, les_model_mq, nacl_graph): + les_model_mq.eval() + les_model_mq.set_is_batch_data(False) + out = les_model_mq(_fresh(nacl_graph)) + assert torch.isfinite(out[KEY.PRED_TOTAL_ENERGY]) + + def test_force_shape(self, les_model_mq, nacl_graph): + les_model_mq.eval() + les_model_mq.set_is_batch_data(False) + out = les_model_mq(_fresh(nacl_graph)) + n = int(nacl_graph[KEY.NUM_ATOMS].item()) + assert out[KEY.PRED_FORCE].shape == (n, 3) + + def test_backward_no_crash(self, les_model_mq, nacl_graph): + les_model_mq.train() + les_model_mq.set_is_batch_data(False) + les_model_mq.zero_grad() + out = les_model_mq(_fresh(nacl_graph)) + (out[KEY.PRED_TOTAL_ENERGY].sum() + out[KEY.PRED_FORCE].sum()).backward()