From 787617c67391592466491efe276cc2ce267f0019 Mon Sep 17 00:00:00 2001 From: Lester Hedges Date: Wed, 23 Sep 2026 13:04:28 +0100 Subject: [PATCH] Backport fix from PR #93. [ci skip] --- emle/models/_mace.py | 47 ++++++++++++++++++++++++++++++++++--------- emle/models/_utils.py | 40 ++++++++++++++++++++++++++++++++++++ tests/test_models.py | 28 ++++++++++++++++++++++++++ 3 files changed, 106 insertions(+), 9 deletions(-) diff --git a/emle/models/_mace.py b/emle/models/_mace.py index 8972a72..8bce103 100644 --- a/emle/models/_mace.py +++ b/emle/models/_mace.py @@ -27,6 +27,7 @@ __all__ = ["MACEEMLE", "MACEEMLEJoint"] +import io as _io import os as _os import torch as _torch import numpy as _np @@ -35,7 +36,6 @@ from ._emle import EMLE as _EMLE from ._utils import _get_neighbor_pairs -from ._utils import _has_neighbor_pairs from torch import Tensor @@ -81,6 +81,31 @@ def _is_energy_emle_mace(model) -> bool: return name == "EnergyEMLEMACE" +def _get_mace_state(model) -> dict: + # Compiled MACE models are ScriptModules, which can't be pickled, so + # serialise them to bytes with torch.jit.save. + state = model.__dict__.copy() + modules = model._modules.copy() + del modules["_mace"] + mace_models = [] + for mace in modules.pop("_mace_models"): + buffer = _io.BytesIO() + _torch.jit.save(mace, buffer) + mace_models.append(buffer.getvalue()) + state["_modules"] = modules + state["_mace_model_bytes"] = mace_models + return state + + +def _set_mace_state(model, state: dict) -> None: + mace_models = state.pop("_mace_model_bytes") + _torch.nn.Module.__setstate__(model, state) + model._mace_models = _torch.nn.ModuleList( + [_torch.jit.load(_io.BytesIO(b)) for b in mace_models] + ) + model._mace = model._mace_models[0] + + class MACEEMLE(_torch.nn.Module): """ Combined MACE and EMLE model. Predicts the in vacuo MACE energy along with @@ -172,10 +197,6 @@ def __init__( ) if not _has_e3nn: raise ImportError("e3nn is required to compile the MACEmodel.") - if not _has_neighbor_pairs: - raise ImportError( - "NNPOps.neighbors.getNeighborPairs is required to use the MACEEMLE model." - ) if device is not None: if not isinstance(device, _torch.device): @@ -455,6 +476,12 @@ def _get_node_attrs(self, atomic_numbers: _torch.Tensor) -> _torch.Tensor: ids = self._atomic_numbers_to_indices(atomic_numbers, z_table=self._z_table) return self._to_one_hot(ids, num_classes=len(self._z_table)) + def __getstate__(self): + return _get_mace_state(self) + + def __setstate__(self, state): + _set_mace_state(self, state) + def to(self, *args, **kwargs): """ Performs Tensor dtype and/or device conversion on the model. @@ -786,10 +813,6 @@ def __init__( ) if not _has_e3nn: raise ImportError("e3nn is required to compile the MACEmodel.") - if not _has_neighbor_pairs: - raise ImportError( - "NNPOps.neighbors.getNeighborPairs is required to use the MACEEMLE model." - ) if device is not None: if not isinstance(device, _torch.device): @@ -1096,6 +1119,12 @@ def _get_node_attrs(self, atomic_numbers: _torch.Tensor) -> _torch.Tensor: ids = self._atomic_numbers_to_indices(atomic_numbers, z_table=self._z_table) return self._to_one_hot(ids, num_classes=len(self._z_table)) + def __getstate__(self): + return _get_mace_state(self) + + def __setstate__(self, state): + _set_mace_state(self, state) + def to(self, *args, **kwargs): """ Performs Tensor dtype and/or device conversion on the model. diff --git a/emle/models/_utils.py b/emle/models/_utils.py index 9c7e76f..69eee3f 100644 --- a/emle/models/_utils.py +++ b/emle/models/_utils.py @@ -102,6 +102,46 @@ def _get_neighbor_pairs( return edge_index, shifts +def _get_neighbor_pairs_torch( + positions: _torch.Tensor, + cell: Optional[_torch.Tensor], + cutoff: float, + dtype: _torch.dtype, + device: _torch.device, +) -> Tuple[_torch.Tensor, _torch.Tensor]: + """ + Pure PyTorch fallback for _get_neighbor_pairs, used when NNPOps is not + available. Has the same signature and return values. + """ + num_atoms = positions.shape[0] + pairs = _torch.triu_indices(num_atoms, num_atoms, 1, device=positions.device) + i = pairs[0] + j = pairs[1] + deltas = positions[i] - positions[j] + if cell is not None: + wrapped_deltas = _minimum_image(deltas, cell) + else: + wrapped_deltas = deltas + mask = _torch.linalg.norm(wrapped_deltas, dim=1) < cutoff + i = i[mask] + j = j[mask] + + edge_index = _torch.stack((_torch.cat((i, j)), _torch.cat((j, i)))).to( + _torch.int64 + ) + if cell is not None: + shifts = deltas[mask] - wrapped_deltas[mask] + shifts = _torch.vstack((shifts, -shifts)) + else: + shifts = _torch.zeros((edge_index.shape[1], 3), dtype=dtype, device=device) + + return edge_index, shifts + + +if not _has_neighbor_pairs: + _get_neighbor_pairs = _get_neighbor_pairs_torch + + def _minimum_image(delta: _torch.Tensor, cell: _torch.Tensor) -> _torch.Tensor: """ Apply the minimum image convention to a batch of displacement vectors. diff --git a/tests/test_models.py b/tests/test_models.py index 8e9ef5a..a2e0ef2 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -75,6 +75,13 @@ def xyz_mm(): except: has_sire = False +try: + import emle_mace # noqa: F401 + + has_emle_mace = True +except: + has_emle_mace = False + MACE_EMLE_MODEL = "tests/input/mace-emle.model" has_emle_mace_model = os.path.exists(MACE_EMLE_MODEL) @@ -202,6 +209,27 @@ def test_mace(alpha_mode, mace_model, atomic_numbers, charges_mm, xyz_qm, xyz_mm @pytest.mark.skipif(not has_mace, reason="mace-torch not installed") @pytest.mark.skipif(not has_e3nn, reason="e3nn not installed") +def test_mace_pickle(atomic_numbers, charges_mm, xyz_qm, xyz_mm): + """ + Check that a MACEEMLE model can be pickled and still gives the same energy. + """ + import pickle + + try: + model = MACEEMLE() + except RuntimeError as e: + pytest.skip(f"MACE model unavailable: {e}") + unpickled = pickle.loads(pickle.dumps(model)) + + energy = model(atomic_numbers, charges_mm, xyz_qm, xyz_mm) + assert torch.allclose( + energy, unpickled(atomic_numbers, charges_mm, xyz_qm, xyz_mm) + ) + + +@pytest.mark.skipif(not has_mace, reason="mace-torch not installed") +@pytest.mark.skipif(not has_e3nn, reason="e3nn not installed") +@pytest.mark.skipif(not has_emle_mace, reason="emle-mace not installed") @pytest.mark.skipif(not has_emle_mace_model, reason="Test emle-mace model not found") def test_emle_mace(atomic_numbers, charges_mm, xyz_qm, xyz_mm): """