From aff4054b563c1e753614a181a81702a9317be122 Mon Sep 17 00:00:00 2001 From: Justin Chu Date: Fri, 4 Sep 2026 10:19:42 -0700 Subject: [PATCH 1/2] test: add MatterGen source parity coverage Validate the official mp_20_base score core against pinned MatterGen source and a source-derived golden fixture. Cover PBC invariance and the dft_band_gap adapter's conditional and unconditional score paths. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> Copilot-Session: 332c98c9-cf56-4bab-8552-ac1b6d730556 Signed-off-by: Justin Chu --- .../diffusion/mattergen-mp20-score.json | 291 ++++++ tests/integration/mattergen_parity_test.py | 949 ++++++++++++++++++ 2 files changed, 1240 insertions(+) create mode 100644 testdata/golden/diffusion/mattergen-mp20-score.json create mode 100644 tests/integration/mattergen_parity_test.py diff --git a/testdata/golden/diffusion/mattergen-mp20-score.json b/testdata/golden/diffusion/mattergen-mp20-score.json new file mode 100644 index 000000000..1b6584eed --- /dev/null +++ b/testdata/golden/diffusion/mattergen-mp20-score.json @@ -0,0 +1,291 @@ +{ + "format": "mobius.mattergen-score-golden.v1", + "source_repository": "https://github.com/microsoft/mattergen", + "source_commit": "842ffe735f7d06cec89d56aa23d9f001e1124b30", + "checkpoint": "mp_20_base-last.ckpt", + "checkpoint_sha256": "ffb80e4425a6f99f479a67b8cd111885d45117234e8947ff77eb3a55df420b9a", + "timestep": 0.25, + "crystal": { + "atomic_numbers": [ + 3, + 8 + ], + "fractional_coordinates": [ + [ + 0.10000000149011612, + 0.15000000596046448, + 0.20000000298023224 + ], + [ + 0.4000000059604645, + 0.44999998807907104, + 0.3499999940395355 + ] + ], + "cell": [ + [ + [ + 6.0, + 0.0, + 0.0 + ], + [ + 0.0, + 6.0, + 0.0 + ], + [ + 0.0, + 0.0, + 6.0 + ] + ] + ] + }, + "outputs": { + "atom_logits": [ + [ + -3.1085011959075928, + -12.377018928527832, + 10.0474853515625, + 1.7480764389038086, + -4.710264205932617, + -4.163971900939941, + -4.576262950897217, + -0.9128561615943909, + -2.197728395462036, + -3.7759506702423096, + -0.43625929951667786, + 1.1311371326446533, + 1.2050434350967407, + -2.747300624847412, + -3.6841583251953125, + -2.899646520614624, + -3.412635564804077, + -11.81222152709961, + 1.7368035316467285, + 0.17785242199897766, + -1.1944823265075684, + 0.4001528024673462, + 4.191551685333252, + 3.63641357421875, + 3.3493564128875732, + 1.6105612516403198, + 1.819764494895935, + -0.5741809606552124, + -0.30905216932296753, + 0.39869967103004456, + -1.0135079622268677, + -2.204998254776001, + -6.7173051834106445, + -5.047140121459961, + -3.843327760696411, + 0.362662672996521, + 1.5078366994857788, + 3.4961049556732178, + -1.110459327697754, + -2.49055552482605, + 0.45650729537010193, + -1.3272736072540283, + 0.6715148091316223, + -2.648557424545288, + -2.5672764778137207, + -3.8022468090057373, + -0.5727097392082214, + -2.3542284965515137, + -0.5075894594192505, + -1.4091328382492065, + -3.9549498558044434, + -0.95721834897995, + -5.114755630493164, + -0.15306557714939117, + 0.9815680980682373, + -1.3402780294418335, + -1.9143892526626587, + 1.0523079633712769, + -4.194021701812744, + -1.9072030782699585, + -2.5036749839782715, + -3.66318416595459, + -2.3778951168060303, + 0.8506289720535278, + -4.43316125869751, + -2.1699888706207275, + -3.8878345489501953, + -1.6545499563217163, + -2.274446725845337, + 0.11025649309158325, + 2.128303289413452, + 0.34579765796661377, + -4.287592887878418, + -0.2032458633184433, + -3.151463031768799, + -3.792062997817993, + -0.6335768699645996, + -4.393718242645264, + -3.284034252166748, + -6.507627487182617, + -2.2801928520202637, + -5.363531112670898, + -4.640564441680908, + -10.857362747192383, + -10.537946701049805, + -11.074407577514648, + -12.097562789916992, + -11.082813262939453, + -3.7654311656951904, + -0.5181636214256287, + -1.1919262409210205, + -2.995697498321533, + -5.51726770401001, + -7.6755828857421875, + -10.672794342041016, + -10.97872257232666, + -11.936291694641113, + -10.407349586486816, + -11.379374504089355, + -11.006378173828125, + -10.8950777053833 + ], + [ + 1.420569658279419, + -3.286576271057129, + 2.380751371383667, + -0.34605100750923157, + -1.7767198085784912, + -1.102645993232727, + 0.23695425689220428, + 6.552952766418457, + 1.3139729499816895, + -2.7440903186798096, + 1.0794779062271118, + 0.09439793974161148, + -1.6759899854660034, + -1.5101711750030518, + -0.8124418258666992, + -0.26274511218070984, + 0.8627647161483765, + -8.79258918762207, + 1.4639447927474976, + -0.10246194154024124, + -1.4410068988800049, + -0.8441038727760315, + -0.10325182229280472, + -0.06950928270816803, + 0.9870234727859497, + 0.16743877530097961, + 0.7330787777900696, + -0.13685491681098938, + 2.1348018646240234, + -0.6191890835762024, + -2.0433735847473145, + -1.2162950038909912, + -1.4412394762039185, + -1.2107605934143066, + 0.9676874279975891, + -2.5670368671417236, + 1.176652193069458, + 0.6676841974258423, + -0.8970199227333069, + -2.188666343688965, + -1.5398725271224976, + -1.1777150630950928, + -3.0990047454833984, + -1.9012398719787598, + -2.4490959644317627, + -1.501819133758545, + -0.6260356307029724, + -0.3859741985797882, + -0.34553948044776917, + -1.2225031852722168, + -1.469612717628479, + -0.969379723072052, + 0.4760497510433197, + -0.8974651098251343, + 1.4012274742126465, + 0.2332736998796463, + -1.491938591003418, + -1.0161073207855225, + -2.2546138763427734, + -1.0184361934661865, + -3.914781332015991, + -0.9634543061256409, + -0.7626060247421265, + -1.3901968002319336, + -3.145031452178955, + -2.681318998336792, + -1.7764226198196411, + -0.7355911731719971, + -1.4520633220672607, + -2.1841349601745605, + -1.1496822834014893, + -1.8152446746826172, + -2.5183801651000977, + -0.717405378818512, + -1.6223053932189941, + -2.2340903282165527, + -1.917813777923584, + -1.336418867111206, + 0.6232991218566895, + -1.2227025032043457, + -1.2338305711746216, + -0.4629520773887634, + 0.7116777300834656, + -8.087118148803711, + -8.19482421875, + -8.282251358032227, + -9.403785705566406, + -9.18935489654541, + -2.9567177295684814, + -2.0076496601104736, + -3.680656671524048, + -1.7048240900039673, + -3.480633497238159, + -4.144234657287598, + -7.990667819976807, + -8.333154678344727, + -9.872822761535645, + -8.682591438293457, + -9.104074478149414, + -8.641095161437988, + -8.7645263671875 + ] + ], + "coordinate_score": [ + [ + -1.6257960796356201, + -1.320042610168457, + 3.2187538146972656 + ], + [ + 1.2291163206100464, + 1.4477763175964355, + -1.9523234367370605 + ] + ], + "lattice_score": [ + [ + [ + -6.4367146492004395, + 3.065068244934082, + 0.72104811668396 + ], + [ + 3.065068244934082, + -6.799010753631592, + 0.7477848529815674 + ], + [ + 0.72104811668396, + 0.7477848529815674, + -1.791745662689209 + ] + ] + ], + "energy": [ + [ + -0.434081494808197 + ] + ] + } +} diff --git a/tests/integration/mattergen_parity_test.py b/tests/integration/mattergen_parity_test.py new file mode 100644 index 000000000..8a61945e0 --- /dev/null +++ b/tests/integration/mattergen_parity_test.py @@ -0,0 +1,949 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +"""L3/L4 CPU parity for the pinned MatterGen GemNet-T score core. + +These tests intentionally require local, immutable MatterGen artifacts rather +than downloading them. Set ``MOBIUS_MATTERGEN_SOURCE_DIR`` to source commit +``842ffe735f7d06cec89d56aa23d9f001e1124b30`` and +``MOBIUS_MATTERGEN_MP20_CHECKPOINT`` to the official ``mp_20_base`` +``last.ckpt``. Optionally set ``MOBIUS_MATTERGEN_DFT_BAND_GAP_CHECKPOINT`` to +the released ``dft_band_gap`` adapter checkpoint. The source repository's +optional PyG extensions do not publish Apple-Silicon wheels for the Torch +version used by this project. The narrowly scoped native-Torch compatibility +definitions below implement only the sum scatter, CSR segment sum, sparse-row +lookup, and Gaussian basis operations that the pinned source invokes on this +fixture. The network layers, periodic graph construction, triplet generation, +checkpoint loading, and score computation are all executed from the pinned +source tree. +""" + +from __future__ import annotations + +import hashlib +import importlib +import importlib.util +import json +import math +import os +import subprocess +import sys +import types +from collections.abc import Iterator, Mapping +from contextlib import contextmanager +from dataclasses import dataclass +from itertools import pairwise +from pathlib import Path +from typing import Any + +import numpy as np +import pytest +import torch + +from mobius import build_from_module +from mobius._testing.ort_inference import OnnxModelSession +from mobius.integrations.mattergen import MatterGenConfig, MatterGenModel +from mobius.integrations.mattergen._configs import MATTERGEN_SOURCE_COMMIT +from mobius.integrations.mattergen._weights import apply_mattergen_checkpoint +from mobius.tasks import MatterGenScoreTask + +pytestmark = pytest.mark.integration + +_SOURCE_DIR_ENV = "MOBIUS_MATTERGEN_SOURCE_DIR" +_MP20_CHECKPOINT_ENV = "MOBIUS_MATTERGEN_MP20_CHECKPOINT" +_DFT_BAND_GAP_CHECKPOINT_ENV = "MOBIUS_MATTERGEN_DFT_BAND_GAP_CHECKPOINT" +_GOLDEN_PATH = ( + Path(__file__).parents[2] + / "testdata" + / "golden" + / "diffusion" + / "mattergen-mp20-score.json" +) +_RTOL = 1e-3 +_ATOL = 1e-3 +# The fine-tuned adapter checkpoint amplifies sub-ULP differences between +# PyTorch and ORT float32 kernels; this remains an output-level parity bound. +_ADAPTER_RTOL = 1e-2 +_ADAPTER_ATOL = 5e-3 + + +@dataclass(frozen=True) +class _SourceModules: + """Pinned source modules used to evaluate the independent PyTorch reference.""" + + gemnet: Any + gemnet_ctrl: Any + data_utils: Any + atom_embedding: Any + model_utils: Any + property_embeddings: Any + + +@dataclass(frozen=True) +class _Crystal: + """A small periodic two-atom crystal in MatterGen's row-vector cell convention.""" + + atomic_numbers: np.ndarray + fractional_coordinates: np.ndarray + cell: np.ndarray + + +@dataclass +class _ReferenceModel: + """Executable pinned-source GemNet and its timestep encoder/classification head.""" + + gemnet: Any + noise_level_encoding: Any + fc_atom: torch.nn.Linear + adapter_property: Any | None = None + adapter_name: str | None = None + + +@dataclass(frozen=True) +class _ScoreCase: + """Host ABI feeds and the exact pinned-source outputs for one score evaluation.""" + + feeds: dict[str, np.ndarray] + outputs: dict[str, np.ndarray] + + +@dataclass(frozen=True) +class _SourceGraph: + """Initial PBC graph for source evaluation plus the normalized ONNX host ABI.""" + + initial_edges: torch.Tensor + to_jimages: torch.Tensor + num_bonds: torch.Tensor + feeds: dict[str, np.ndarray] + + +@dataclass +class _Mp20Runtime: + """Shared real-weight source references and a loaded ONNX Runtime score graph.""" + + session: OnnxModelSession + original: dict[float, _ScoreCase] + translated: _ScoreCase + permuted: _ScoreCase + checkpoint_sha256: str + + def close(self) -> None: + self.session.close() + + +def _required_artifact(environment_variable: str) -> Path: + """Resolve an explicit local integration artifact or skip without network access.""" + raw_path = os.environ.get(environment_variable) + if raw_path is None: + pytest.skip(f"set {environment_variable} to run MatterGen source parity") + path = Path(raw_path) + if not path.is_file() and environment_variable == _SOURCE_DIR_ENV: + if not path.is_dir(): + pytest.skip(f"{environment_variable} is not a readable source directory: {path}") + elif not path.is_file(): + pytest.skip(f"{environment_variable} is not a readable checkpoint: {path}") + return path + + +def _source_revision(source_dir: Path) -> str: + """Return the checked-out source revision, refusing an unpinned reference tree.""" + completed = subprocess.run( + ["git", "-C", str(source_dir), "rev-parse", "HEAD"], + check=True, + capture_output=True, + text=True, + ) + return completed.stdout.strip() + + +def _module_is_available(name: str) -> bool: + """Handle persistent in-process compatibility modules without re-resolving specs.""" + return name in sys.modules or importlib.util.find_spec(name) is not None + + +def _install_source_dependency_compatibility() -> None: + """Provide native-Torch equivalents only for unavailable source extension APIs. + + MatterGen's pinned source calls all of these helpers only with leading-axis + sum reductions. Keeping the compatibility surface this small avoids a + second implementation of any GemNet layer in this test. + """ + # MatterGen v1.0.3 declares NumPy <2. NumPy 2 removed this historical + # alias; restoring the identical stdlib module lets its source basis + # generator run without changing any generated polynomial. + if not hasattr(np, "math"): + np.math = math # type: ignore[attr-defined] + + if not _module_is_available("omegaconf"): + omegaconf = types.ModuleType("omegaconf") + + class OmegaConf: + """Minimal import-time resolver registration interface.""" + + @staticmethod + def register_new_resolver(*_args: object, **_kwargs: object) -> None: + return None + + omegaconf.OmegaConf = OmegaConf + sys.modules["omegaconf"] = omegaconf + + if not _module_is_available("torch_scatter"): + torch_scatter = types.ModuleType("torch_scatter") + + def scatter( + source: torch.Tensor, + index: torch.Tensor, + dim: int = 0, + dim_size: int | torch.Tensor | None = None, + reduce: str = "sum", + ) -> torch.Tensor: + if dim != 0 or reduce not in {"sum", "add"}: + raise NotImplementedError( + "MatterGen fixture needs only leading-axis sum scatter" + ) + if index.ndim != 1 or index.shape[0] != source.shape[0]: + raise ValueError("source and index must share their leading dimension") + rows = ( + int(index.max().item()) + 1 + if dim_size is None + else int(torch.as_tensor(dim_size).item()) + ) + result = source.new_zeros((rows, *source.shape[1:])) + return result.index_add_(0, index, source) + + def segment_coo( + source: torch.Tensor, + index: torch.Tensor, + dim_size: int | torch.Tensor | None = None, + reduce: str = "sum", + ) -> torch.Tensor: + return scatter(source, index, dim_size=dim_size, reduce=reduce) + + def segment_csr( + source: torch.Tensor, indptr: torch.Tensor, reduce: str = "sum" + ) -> torch.Tensor: + if reduce not in {"sum", "add"}: + raise NotImplementedError("MatterGen fixture needs only CSR segment sums") + return torch.stack( + [ + source[int(start.item()) : int(end.item())].sum(dim=0) + for start, end in pairwise(indptr) + ] + ) + + torch_scatter.scatter = scatter + torch_scatter.scatter_add = scatter + torch_scatter.segment_coo = segment_coo + torch_scatter.segment_csr = segment_csr + sys.modules["torch_scatter"] = torch_scatter + + if not _module_is_available("torch_sparse"): + torch_sparse = types.ModuleType("torch_sparse") + + class _SparseStorage: + """Storage view exposing the two accessors GemNetT.get_triplets uses.""" + + def __init__(self, row: torch.Tensor, value: torch.Tensor): + self._row = row + self._value = value + + def row(self) -> torch.Tensor: + return self._row + + def value(self) -> torch.Tensor: + return self._value + + class SparseTensor: + """Row-indexed multigraph lookup preserving MatterGen's edge ordering.""" + + def __init__( + self, + *, + row: torch.Tensor, + col: torch.Tensor, + value: torch.Tensor, + sparse_sizes: tuple[torch.Tensor, torch.Tensor] | tuple[int, int], + ): + del col, sparse_sizes + self._row = row + self._value = value + self.storage = _SparseStorage(row.new_empty(0), value.new_empty(0)) + + def __getitem__(self, queried_rows: torch.Tensor) -> SparseTensor: + result_rows: list[torch.Tensor] = [] + result_values: list[torch.Tensor] = [] + for output_row, source_row in enumerate(queried_rows): + match = torch.nonzero(self._row == source_row, as_tuple=False).squeeze(1) + result_rows.append( + torch.full( + (len(match),), + output_row, + dtype=queried_rows.dtype, + device=queried_rows.device, + ) + ) + result_values.append(self._value[match]) + result = object.__new__(SparseTensor) + result._row = self._row + result._value = self._value + result.storage = _SparseStorage( + torch.cat(result_rows), + torch.cat(result_values), + ) + return result + + def to(self, _device: torch.device) -> SparseTensor: + return self + + torch_sparse.SparseTensor = SparseTensor + sys.modules["torch_sparse"] = torch_sparse + + if not _module_is_available("torch_geometric"): + torch_geometric = types.ModuleType("torch_geometric") + torch_geometric.__path__ = [] + pyg_data = types.ModuleType("torch_geometric.data") + pyg_typing = types.ModuleType("torch_geometric.typing") + pyg_utils = types.ModuleType("torch_geometric.utils") + pyg_nn = types.ModuleType("torch_geometric.nn") + pyg_nn.__path__ = [] + pyg_models = types.ModuleType("torch_geometric.nn.models") + pyg_models.__path__ = [] + pyg_schnet = types.ModuleType("torch_geometric.nn.models.schnet") + + class Data: + """Import-only Data base sufficient for source type declarations.""" + + def __init__(self, **kwargs: object): + self.__dict__.update(kwargs) + + class Batch: + """Import-only Batch factory used while defining ChemGraphBatch.""" + + def __new__(cls, _base_cls: type | None = None, **_kwargs: object): + if _base_cls is None: + return super().__new__(cls) + return type(f"{_base_cls.__name__}Batch", (_base_cls, cls), {})() + + class GaussianSmearing(torch.nn.Module): + """Pinned PyG GaussianSmearing formula used by MatterGen's radial basis.""" + + def __init__( + self, start: float, stop: float, num_gaussians: int, **_kwargs: object + ): + super().__init__() + offset = torch.linspace(start, stop, num_gaussians) + self.register_buffer("offset", offset) + self.coeff = -0.5 / float((offset[1] - offset[0]).square()) + + def forward(self, distance: torch.Tensor) -> torch.Tensor: + return torch.exp( + self.coeff * (distance.view(-1, 1) - self.offset.view(1, -1)).square() + ) + + pyg_data.Data = Data + pyg_data.Batch = Batch + pyg_typing.OptTensor = object + pyg_schnet.GaussianSmearing = GaussianSmearing + torch_geometric.data = pyg_data + torch_geometric.typing = pyg_typing + torch_geometric.utils = pyg_utils + torch_geometric.nn = pyg_nn + pyg_nn.models = pyg_models + pyg_models.schnet = pyg_schnet + sys.modules.update( + { + "torch_geometric": torch_geometric, + "torch_geometric.data": pyg_data, + "torch_geometric.typing": pyg_typing, + "torch_geometric.utils": pyg_utils, + "torch_geometric.nn": pyg_nn, + "torch_geometric.nn.models": pyg_models, + "torch_geometric.nn.models.schnet": pyg_schnet, + } + ) + + if not _module_is_available("pymatgen"): + pymatgen = types.ModuleType("pymatgen") + pymatgen.__path__ = [] + pymatgen_core = types.ModuleType("pymatgen.core") + + class Element: + """Import-only placeholder; this fixture supplies atomic numbers directly.""" + + def __init__(self, *_args: object, **_kwargs: object): + raise RuntimeError( + "MatterGen parity fixture does not construct pymatgen Elements" + ) + + pymatgen_core.Element = Element + pymatgen.core = pymatgen_core + sys.modules["pymatgen"] = pymatgen + sys.modules["pymatgen.core"] = pymatgen_core + + if not _module_is_available("emmet"): + emmet = types.ModuleType("emmet") + emmet.__path__ = [] + emmet_core = types.ModuleType("emmet.core") + emmet_core.__path__ = [] + emmet_material = types.ModuleType("emmet.core.material") + + class PropertyOrigin: + """Import-only provenance type used only by source annotations.""" + + emmet_material.PropertyOrigin = PropertyOrigin + emmet.core = emmet_core + emmet_core.material = emmet_material + sys.modules.update( + { + "emmet": emmet, + "emmet.core": emmet_core, + "emmet.core.material": emmet_material, + } + ) + + +@contextmanager +def _pinned_source_modules(source_dir: Path) -> Iterator[_SourceModules]: + """Import all reference layers from the requested immutable source checkout.""" + if _source_revision(source_dir) != MATTERGEN_SOURCE_COMMIT: + pytest.fail( + f"MatterGen source must be {MATTERGEN_SOURCE_COMMIT}, got {_source_revision(source_dir)}" + ) + _install_source_dependency_compatibility() + stale = [ + name for name in sys.modules if name == "mattergen" or name.startswith("mattergen.") + ] + for name in stale: + del sys.modules[name] + sys.path.insert(0, str(source_dir)) + try: + yield _SourceModules( + gemnet=importlib.import_module("mattergen.common.gemnet.gemnet"), + gemnet_ctrl=importlib.import_module("mattergen.common.gemnet.gemnet_ctrl"), + data_utils=importlib.import_module("mattergen.common.utils.data_utils"), + atom_embedding=importlib.import_module( + "mattergen.common.gemnet.layers.embedding_block" + ), + model_utils=importlib.import_module("mattergen.diffusion.model_utils"), + property_embeddings=importlib.import_module("mattergen.property_embeddings"), + ) + finally: + sys.path.remove(str(source_dir)) + for name in [ + name + for name in sys.modules + if name == "mattergen" or name.startswith("mattergen.") + ]: + del sys.modules[name] + + +def _load_checkpoint(checkpoint: Path) -> tuple[dict[str, torch.Tensor], Mapping[str, object]]: + """Safely read the official Lightning checkpoint and return its score-core state.""" + payload = torch.load(checkpoint, map_location="cpu", weights_only=True) + if not isinstance(payload, Mapping): + raise TypeError("MatterGen checkpoint must deserialize to a mapping") + state_dict = payload.get("state_dict") + config = payload.get("config") + if not isinstance(state_dict, Mapping) or not isinstance(config, Mapping): + raise TypeError( + "MatterGen checkpoint must provide mapping state_dict and config values" + ) + prefix = "diffusion_module.model." + state = { + name.removeprefix(prefix): value.detach().clone() + for name, value in state_dict.items() + if isinstance(name, str) + and name.startswith(prefix) + and isinstance(value, torch.Tensor) + } + if not state: + raise ValueError("MatterGen checkpoint contains no score-core tensors") + return state, config + + +def _gemnet_kwargs(config: MatterGenConfig, atom_embedding: Any) -> dict[str, object]: + """Build the exact GemNet constructor argument set reflected in the Hydra config.""" + return { + "num_targets": config.num_targets, + "latent_dim": config.latent_dim, + "atom_embedding": atom_embedding, + "num_spherical": config.num_spherical, + "num_radial": config.num_radial, + "num_blocks": config.num_blocks, + "emb_size_atom": config.emb_size_atom, + "emb_size_edge": config.emb_size_edge, + "emb_size_trip": config.emb_size_trip, + "emb_size_rbf": config.emb_size_rbf, + "emb_size_cbf": config.emb_size_cbf, + "emb_size_bil_trip": config.emb_size_bil_trip, + "num_before_skip": config.num_before_skip, + "num_after_skip": config.num_after_skip, + "num_concat": config.num_concat, + "num_atom": config.num_atom, + "regress_stress": config.regress_stress, + "cutoff": config.cutoff, + "max_neighbors": config.max_neighbors, + "max_cell_images_per_dim": config.max_cell_images_per_dim, + "otf_graph": False, + } + + +def _make_reference_model( + source: _SourceModules, + config: MatterGenConfig, + state: Mapping[str, torch.Tensor], +) -> _ReferenceModel: + """Load official tensors into the independently imported source neural modules.""" + atom_embedding = source.atom_embedding.AtomEmbedding( + config.hidden_size, + with_mask_type=True, + ) + gemnet = source.gemnet.GemNetT(**_gemnet_kwargs(config, atom_embedding)) + gemnet.load_state_dict( + { + name.removeprefix("gemnet."): value + for name, value in state.items() + if name.startswith("gemnet.") + }, + strict=True, + ) + noise_level_encoding = source.model_utils.NoiseLevelEncoding(config.hidden_size) + noise_level_encoding.load_state_dict( + { + name.removeprefix("noise_level_encoding."): value + for name, value in state.items() + if name.startswith("noise_level_encoding.") + }, + strict=True, + ) + fc_atom = torch.nn.Linear(config.hidden_size, config.num_atom_types) + fc_atom.load_state_dict( + { + name.removeprefix("fc_atom."): value + for name, value in state.items() + if name.startswith("fc_atom.") + }, + strict=True, + ) + gemnet.eval() + noise_level_encoding.eval() + fc_atom.eval() + return _ReferenceModel(gemnet, noise_level_encoding, fc_atom) + + +class _SourcePropertyBatch(dict[str, object]): + """Minimal mapping interface consumed by source PropertyEmbedding.forward.""" + + def __init__( + self, + *, + value: torch.Tensor, + use_unconditional: torch.Tensor, + position: torch.Tensor, + ): + super().__init__( + dft_band_gap=value, + num_atoms=torch.tensor([len(position)], dtype=torch.long), + _USE_UNCONDITIONAL_EMBEDDING={"dft_band_gap": use_unconditional}, + ) + self.pos = position + + +def _make_dft_band_gap_reference_model( + source: _SourceModules, + config: MatterGenConfig, + state: Mapping[str, torch.Tensor], +) -> _ReferenceModel: + """Load the source GemNet-T control adapter and its property encoder.""" + assert config.condition_on_adapt == ("dft_band_gap",) + atom_embedding = source.atom_embedding.AtomEmbedding( + config.hidden_size, + with_mask_type=True, + ) + gemnet = source.gemnet_ctrl.GemNetTCtrl( + list(config.condition_on_adapt), + **_gemnet_kwargs(config, atom_embedding), + ) + gemnet.load_state_dict( + { + name.removeprefix("gemnet."): value + for name, value in state.items() + if name.startswith("gemnet.") + }, + strict=True, + ) + noise_level_encoding = source.model_utils.NoiseLevelEncoding(config.hidden_size) + noise_level_encoding.load_state_dict( + { + name.removeprefix("noise_level_encoding."): value + for name, value in state.items() + if name.startswith("noise_level_encoding.") + }, + strict=True, + ) + property_embedding = source.property_embeddings.PropertyEmbedding( + name="dft_band_gap", + conditional_embedding_module=source.model_utils.NoiseLevelEncoding(config.hidden_size), + unconditional_embedding_module=source.property_embeddings.ZerosEmbedding( + config.hidden_size + ), + scaler=source.data_utils.StandardScalerTorch(), + ) + property_embedding.load_state_dict( + { + name.removeprefix("property_embeddings_adapt.dft_band_gap."): value + for name, value in state.items() + if name.startswith("property_embeddings_adapt.dft_band_gap.") + }, + strict=True, + ) + fc_atom = torch.nn.Linear(config.hidden_size, config.num_atom_types) + fc_atom.load_state_dict( + { + name.removeprefix("fc_atom."): value + for name, value in state.items() + if name.startswith("fc_atom.") + }, + strict=True, + ) + gemnet.eval() + noise_level_encoding.eval() + property_embedding.eval() + fc_atom.eval() + return _ReferenceModel( + gemnet, + noise_level_encoding, + fc_atom, + adapter_property=property_embedding, + adapter_name="dft_band_gap", + ) + + +def _crystal( + *, + translated: bool = False, + permutation: np.ndarray | None = None, +) -> _Crystal: + """Return a non-boundary periodic fixture with an optional rigid translation/permutation.""" + fractional_coordinates = np.array( + [[0.10, 0.15, 0.20], [0.40, 0.45, 0.35]], + dtype=np.float32, + ) + if translated: + # No coordinate wraps, so this is a true rigid translation in the source cell convention. + fractional_coordinates = fractional_coordinates + np.array( + [0.10, 0.20, 0.10], dtype=np.float32 + ) + atomic_numbers = np.array([3, 8], dtype=np.int64) + if permutation is not None: + atomic_numbers = atomic_numbers[permutation] + fractional_coordinates = fractional_coordinates[permutation] + return _Crystal( + atomic_numbers=atomic_numbers, + fractional_coordinates=fractional_coordinates, + cell=np.diag(np.array([6.0, 6.0, 6.0], dtype=np.float32))[None, ...], + ) + + +def _source_host_feeds( + reference: _ReferenceModel, + source: _SourceModules, + crystal: _Crystal, + timestep: float, +) -> _SourceGraph: + """Build host ABI tensors through source PBC graph, symmetrization, and triplet code.""" + atomic_numbers = torch.from_numpy(crystal.atomic_numbers) + fractional_coordinates = torch.from_numpy(crystal.fractional_coordinates) + lattice = torch.from_numpy(crystal.cell) + num_atoms = torch.tensor([len(atomic_numbers)], dtype=torch.long) + batch = torch.zeros(len(atomic_numbers), dtype=torch.long) + cartesian_coordinates = source.data_utils.frac_to_cart_coords_with_lattice( + fractional_coordinates, num_atoms, lattice + ) + initial_edges, to_jimages, num_bonds = source.data_utils.radius_graph_pbc( + cart_coords=cartesian_coordinates, + lattice=lattice, + num_atoms=num_atoms, + radius=7.0, + max_num_neighbors_threshold=50, + max_cell_images_per_dim=5, + ) + ( + edge_index, + _neighbors, + edge_distance, + edge_direction, + id_swap, + id3_ba, + id3_ca, + id3_ragged_idx, + _cell_offsets, + ) = reference.gemnet.generate_interaction_graph( + cartesian_coordinates, + lattice, + num_atoms, + initial_edges, + to_jimages, + num_bonds, + ) + edge_batch = batch[edge_index[0]] + edge_lattice_cosines = torch.cosine_similarity( + edge_direction[:, None], + lattice[edge_batch], + dim=-1, + ) + feeds = { + "atomic_numbers": atomic_numbers.numpy(), + "batch": batch.numpy(), + "timestep": np.array([timestep], dtype=np.float32), + "edge_index": edge_index.numpy(), + "edge_distance": edge_distance.numpy(), + "edge_direction": edge_direction.numpy(), + "edge_lattice_cosines": edge_lattice_cosines.numpy(), + "id_swap": id_swap.numpy(), + "id3_ba": id3_ba.numpy(), + "id3_ca": id3_ca.numpy(), + "id3_ragged_idx": id3_ragged_idx.numpy(), + } + return _SourceGraph(initial_edges, to_jimages, num_bonds, feeds) + + +def _source_score( + reference: _ReferenceModel, + source: _SourceModules, + crystal: _Crystal, + timestep: float, + *, + adapter_value: float | None = None, + use_unconditional: bool = False, +) -> _ScoreCase: + """Evaluate all score heads in the pinned source with the same host graph passed to ONNX.""" + graph = _source_host_feeds(reference, source, crystal, timestep) + adapter_inputs: dict[str, object] = {} + if reference.adapter_property is not None: + assert reference.adapter_name is not None + if adapter_value is None: + raise ValueError("adapter score requires a concrete source property value") + adapter_mask = torch.tensor([[use_unconditional]], dtype=torch.bool) + property_value = torch.tensor([adapter_value], dtype=torch.float32) + adapter_embedding = reference.adapter_property( + _SourcePropertyBatch( + value=property_value, + use_unconditional=adapter_mask, + position=torch.from_numpy(crystal.fractional_coordinates), + ) + ) + adapter_inputs = { + "cond_adapt": {reference.adapter_name: adapter_embedding}, + "cond_adapt_mask": {reference.adapter_name: adapter_mask}, + } + graph.feeds[f"condition.{reference.adapter_name}"] = property_value.numpy() + graph.feeds[f"condition.{reference.adapter_name}.use_unconditional"] = ( + adapter_mask.squeeze(1).numpy() + ) + with torch.inference_mode(): + output = reference.gemnet( + z=reference.noise_level_encoding(torch.from_numpy(graph.feeds["timestep"])), + frac_coords=torch.from_numpy(crystal.fractional_coordinates), + atom_types=torch.from_numpy(graph.feeds["atomic_numbers"]), + num_atoms=torch.tensor([len(crystal.atomic_numbers)], dtype=torch.long), + batch=torch.from_numpy(graph.feeds["batch"]), + edge_index=graph.initial_edges, + to_jimages=graph.to_jimages, + num_bonds=graph.num_bonds, + lattice=torch.from_numpy(crystal.cell), + **adapter_inputs, + ) + outputs = { + "atom_logits": reference.fc_atom(output.node_embeddings).numpy(), + "coordinate_score": output.forces.numpy(), + "lattice_score": output.stress.numpy(), + "energy": output.energy.numpy(), + } + return _ScoreCase(graph.feeds, outputs) + + +def _sha256(path: Path) -> str: + """Hash an immutable checkpoint for golden-fixture provenance.""" + digest = hashlib.sha256() + with path.open("rb") as file: + for block in iter(lambda: file.read(1024 * 1024), b""): + digest.update(block) + return digest.hexdigest() + + +@pytest.fixture(scope="module") +def mp20_runtime() -> Iterator[_Mp20Runtime]: + """Run pinned-source references first, then load the same checkpoint into standard ONNX.""" + source_dir = _required_artifact(_SOURCE_DIR_ENV) + checkpoint = _required_artifact(_MP20_CHECKPOINT_ENV) + with _pinned_source_modules(source_dir) as source: + state, hydra_config = _load_checkpoint(checkpoint) + config = MatterGenConfig.from_hydra_config(hydra_config, variant="mp_20_base") + reference = _make_reference_model(source, config, state) + original = { + timestep: _source_score(reference, source, _crystal(), timestep) + for timestep in (0.25, 0.75) + } + translated = _source_score(reference, source, _crystal(translated=True), 0.25) + permuted = _source_score( + reference, + source, + _crystal(permutation=np.array([1, 0], dtype=np.int64)), + 0.25, + ) + del reference + del state + + module = MatterGenModel(config) + package = build_from_module(module, config, task=MatterGenScoreTask()) + assert package.export_report is not None + score_core_report = package.export_report.component("score_core") + assert score_core_report.runtime_validation_status == "validated" + assert score_core_report.evidence_id == "mattergen-score-core-ort" + apply_mattergen_checkpoint(package, module, checkpoint) + # This materializes the reported standard-ONNX score core in CPU ORT; the + # tests below execute all exported score heads against pinned-source values. + session = OnnxModelSession(package["model"], device="cpu") + del module + del package + runtime = _Mp20Runtime( + session=session, + original=original, + translated=translated, + permuted=permuted, + checkpoint_sha256=_sha256(checkpoint), + ) + try: + yield runtime + finally: + runtime.close() + + +def _assert_outputs_close( + actual: Mapping[str, np.ndarray], + expected: Mapping[str, np.ndarray], + *, + rtol: float = _RTOL, + atol: float = _ATOL, +) -> None: + """Compare every source score-core output, retaining array-specific failure context.""" + assert set(actual) == {"atom_logits", "coordinate_score", "lattice_score", "energy"} + for name in actual: + np.testing.assert_allclose( + actual[name], + expected[name], + rtol=rtol, + atol=atol, + err_msg=f"{name} differs between ONNX Runtime and pinned MatterGen source", + ) + + +def _unpermute_atom_outputs( + outputs: Mapping[str, np.ndarray], permutation: np.ndarray +) -> dict[str, np.ndarray]: + """Restore atom-indexed outputs after evaluating a source-compatible input atom permutation.""" + inverse = np.argsort(permutation) + return { + name: value[inverse] if name in {"atom_logits", "coordinate_score"} else value + for name, value in outputs.items() + } + + +def test_mp20_source_score_core_matches_onnx_at_two_timesteps( + mp20_runtime: _Mp20Runtime, +) -> None: + """L3: exact-source synthetic periodic graph parity for all requested score heads.""" + for expected in mp20_runtime.original.values(): + _assert_outputs_close(mp20_runtime.session.run(expected.feeds), expected.outputs) + + +def test_mp20_periodic_translation_and_permutation_invariance( + mp20_runtime: _Mp20Runtime, +) -> None: + """L3: establish PBC translation/permutation invariance in source before asserting it in ONNX.""" + baseline = mp20_runtime.original[0.25].outputs + translation = mp20_runtime.translated + permutation = np.array([1, 0], dtype=np.int64) + permuted = _unpermute_atom_outputs(mp20_runtime.permuted.outputs, permutation) + + _assert_outputs_close(translation.outputs, baseline) + _assert_outputs_close(permuted, baseline) + _assert_outputs_close(mp20_runtime.session.run(translation.feeds), translation.outputs) + _assert_outputs_close( + _unpermute_atom_outputs( + mp20_runtime.session.run(mp20_runtime.permuted.feeds), permutation + ), + baseline, + ) + + +@pytest.mark.golden +def test_mp20_one_step_source_golden(mp20_runtime: _Mp20Runtime) -> None: + """L4: compare a one-step real ``mp_20_base`` source output with committed provenance.""" + golden = json.loads(_GOLDEN_PATH.read_text(encoding="utf-8")) + assert golden["source_commit"] == MATTERGEN_SOURCE_COMMIT + assert golden["checkpoint_sha256"] == mp20_runtime.checkpoint_sha256 + assert np.isclose(golden["timestep"], 0.25) + expected = { + name: np.asarray(value, dtype=np.float32) for name, value in golden["outputs"].items() + } + source_case = mp20_runtime.original[0.25] + _assert_outputs_close(source_case.outputs, expected) + _assert_outputs_close(mp20_runtime.session.run(source_case.feeds), expected) + + +def test_dft_band_gap_adapter_matches_source_for_conditional_and_unconditional_scores() -> ( + None +): + """L3: exercise the released scalar adapter through both source embedding modes.""" + source_dir = _required_artifact(_SOURCE_DIR_ENV) + checkpoint = _required_artifact(_DFT_BAND_GAP_CHECKPOINT_ENV) + with _pinned_source_modules(source_dir) as source: + state, hydra_config = _load_checkpoint(checkpoint) + config = MatterGenConfig.from_hydra_config(hydra_config, variant="dft_band_gap") + reference = _make_dft_band_gap_reference_model(source, config, state) + conditional = _source_score( + reference, + source, + _crystal(), + 0.0, + adapter_value=1.5, + ) + unconditional = _source_score( + reference, + source, + _crystal(), + 0.0, + adapter_value=1.5, + use_unconditional=True, + ) + + # The property value must reach GemNet-T control blocks; otherwise the two + # source evaluations would be indistinguishable despite different adapter masks. + assert not np.allclose( + conditional.outputs["coordinate_score"], + unconditional.outputs["coordinate_score"], + rtol=_RTOL, + atol=_ATOL, + ) + + module = MatterGenModel(config) + package = build_from_module(module, config, task=MatterGenScoreTask()) + apply_mattergen_checkpoint(package, module, checkpoint) + session = OnnxModelSession(package["model"], device="cpu") + try: + _assert_outputs_close( + session.run(conditional.feeds), + conditional.outputs, + rtol=_ADAPTER_RTOL, + atol=_ADAPTER_ATOL, + ) + _assert_outputs_close( + session.run(unconditional.feeds), + unconditional.outputs, + rtol=_ADAPTER_RTOL, + atol=_ADAPTER_ATOL, + ) + finally: + session.close() From 92771e388af5279892e1b2b973c6e8e31f46deda Mon Sep 17 00:00:00 2001 From: Justin Chu Date: Fri, 4 Sep 2026 11:06:26 -0700 Subject: [PATCH 2/2] feat: add MatterGen score-core export Export the pinned official MatterGen GemNet-T score core with Hydra configuration parsing, exhaustive checkpoint routing, a documented dynamic PBC graph ABI, and CLI integration. Add a source-faithful CPU host sampler with the released D3PM/VE/VP schedule and a real ONNX sampling golden while explicitly retaining host orchestration in the partial export report. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> Copilot-Session: 332c98c9-cf56-4bab-8552-ac1b6d730556 Signed-off-by: Justin Chu --- docs/index.md | 1 + docs/mattergen.md | 179 +++ docs/model-catalog.md | 12 + pyproject.toml | 1 + src/mobius/__init__.py | 2 + src/mobius/__main__.py | 71 +- src/mobius/integrations/mattergen/__init__.py | 100 ++ src/mobius/integrations/mattergen/_builder.py | 183 +++ .../integrations/mattergen/_builder_test.py | 187 +++ src/mobius/integrations/mattergen/_configs.py | 454 +++++++ .../integrations/mattergen/_contract.py | 323 +++++ .../integrations/mattergen/_contract_test.py | 100 ++ src/mobius/integrations/mattergen/_runtime.py | 1199 ++++++++++++++++ .../integrations/mattergen/_runtime_test.py | 325 +++++ src/mobius/integrations/mattergen/_weights.py | 148 ++ .../integrations/mattergen/_weights_test.py | 115 ++ src/mobius/models/__init__.py | 3 + src/mobius/models/mattergen.py | 1206 +++++++++++++++++ src/mobius/tasks/__init__.py | 3 + src/mobius/tasks/_mattergen.py | 208 +++ .../diffusion/mattergen-mp20-host-sample.json | 42 + tests/cli_test.py | 67 + tests/integration/mattergen_parity_test.py | 222 ++- 23 files changed, 5147 insertions(+), 4 deletions(-) create mode 100644 docs/mattergen.md create mode 100644 src/mobius/integrations/mattergen/__init__.py create mode 100644 src/mobius/integrations/mattergen/_builder.py create mode 100644 src/mobius/integrations/mattergen/_builder_test.py create mode 100644 src/mobius/integrations/mattergen/_configs.py create mode 100644 src/mobius/integrations/mattergen/_contract.py create mode 100644 src/mobius/integrations/mattergen/_contract_test.py create mode 100644 src/mobius/integrations/mattergen/_runtime.py create mode 100644 src/mobius/integrations/mattergen/_runtime_test.py create mode 100644 src/mobius/integrations/mattergen/_weights.py create mode 100644 src/mobius/integrations/mattergen/_weights_test.py create mode 100644 src/mobius/models/mattergen.py create mode 100644 src/mobius/tasks/_mattergen.py create mode 100644 testdata/golden/diffusion/mattergen-mp20-host-sample.json diff --git a/docs/index.md b/docs/index.md index aae56e1ff..5ac25a2c7 100644 --- a/docs/index.md +++ b/docs/index.md @@ -14,6 +14,7 @@ getting-started cli_reference module-architecture model-catalog +mattergen models/index ``` diff --git a/docs/mattergen.md b/docs/mattergen.md new file mode 100644 index 000000000..6800cbb43 --- /dev/null +++ b/docs/mattergen.md @@ -0,0 +1,179 @@ +# MatterGen crystal diffusion score core + +Mobius supports the deterministic neural score component from the official +[`microsoft/mattergen`](https://huggingface.co/microsoft/mattergen) release. +It is a periodic-crystal graph diffusion system, not a Transformers or +Diffusers model. The integration pins the Hub revision +`5244495dd9a979ff71abc7548a0b14b9deb0069a` and replicates the matching +MatterGen v1.0.3 source at +`842ffe735f7d06cec89d56aa23d9f001e1124b30`. + +```mermaid +flowchart LR + S[Host: noisy atomic numbers, fractional coordinates, row-vector cell] + G[Host: periodic radius graph, symmetric ordering and sparse triplets] + C[Host: raw condition values and unconditional selectors] + D[ONNX: time/property embeddings and GemNet-T score core] + O[ONNX: atom logits, Cartesian coordinate score, lattice score, energy diagnostic] + P[Host: D3PM/SDE scheduler, CFG, wrapping and lattice projection] + V[Host: dependency-free crystal validation] + S --> G --> D --> O --> P --> V + C --> D +``` + +## ONNX score contract + +The exported `model.onnx` consumes a host-normalized, dynamic periodic graph: + +| Tensor | Type and shape | Meaning | +|---|---|---| +| `atomic_numbers` | `int64[N]` | MatterGen D3PM species IDs: `1..100`, with `101` as the absorbing mask. | +| `batch` | `int64[N]` | Crystal index for each atom. | +| `timestep` | `float32[B]` | Diffusion time for each crystal. | +| `edge_index` | `int64[2,E]` | Source-ordered periodic edges after MatterGen symmetric reordering. | +| `edge_distance` | `float32[E]` | Periodic Cartesian edge lengths. | +| `edge_direction` | `float32[E,3]` | MatterGen `V_st`: the **negative** normalized periodic distance vector. | +| `edge_lattice_cosines` | `float32[E,3]` | Host-computed `cosine_similarity(V_st, cell[batch[edge_index[0]]])`. | +| `id_swap`, `id3_ba`, `id3_ca`, `id3_ragged_idx` | `int64[...]` | Symmetric-edge and sparse-triplet indexes generated by the source ordering. | +| condition input(s) | family-specific | Raw chemical-system multihot, space-group, or scalar values plus explicit boolean unconditional selectors. | + +It returns `atom_logits: float32[N,101]`, a Cartesian `coordinate_score: +float32[N,3]`, `lattice_score: float32[B,3,3]`, and an `energy: +float32[B,1]` diagnostic. The energy output keeps every trained GemNet +OutputBlock path observable; MatterGen's denoiser does not use it in its +sampling state update. The core intentionally does not mask atom logits, +convert Cartesian scores to fractional scores, wrap coordinates, or sample +atom types. + +## Host orchestration and limits + +MatterGen reconstructs its periodic radius graph on every score evaluation, +including data-dependent periodic image enumeration, nearest-neighbor +selection, symmetric edge reordering, and ragged triplets. Those operations +are not portable as a faithful dynamic ONNX contract. Mobius provides the +source-faithful CPU host implementation in +`mobius.integrations.mattergen.MatterGenHostSampler`; the application still +supplies the ONNX score callback. It has no MatterGen, PyTorch Geometric, +`torch_scatter`, or `torch_sparse` runtime dependency. In particular, +triplets exclude matching **edge IDs**, not matching atom IDs; valid periodic +self-image triplets remain possible. + +The adapter owns all stochastic semantics: the fixed 1,000-step +absorbing-mask D3PM, wrapped VE coordinates, VP lattice updates, +predictor-corrector scheduling, classifier-free guidance, modulo-one +coordinate wrapping, lattice projection, and final structural validation. +It intentionally rejects a shortened/re-scheduled path rather than claiming +it is MatterGen. ONNX Runtime GenAI does not provide this runtime. + +The `ModelPackage` remains a **partial score-core export**: its +`export_report.json` continues to mark periodic graph construction, sampling, +and crystal validation as deferred host stages with +`end_to_end_runnable: false`. `MatterGenHostSampler` is a separate, +application-composed CPU adapter around an explicit score callback; its +source-semantics and L5 test execute the real `mp_20_base` score artifact but +do not upgrade that partial package report. + +```python +import onnxruntime as ort +import torch +from mobius.integrations.mattergen import ( + MatterGenHostSampler, + create_onnxruntime_score_callback, +) + +session = ort.InferenceSession("model.onnx", providers=["CPUExecutionProvider"]) +sampler = MatterGenHostSampler( + create_onnxruntime_score_callback(session), + condition_names=(), # Use the exported config's condition-input names when present. +) +samples = sampler.sample(torch.tensor([4, 8], dtype=torch.long), seed=1234) +crystals = samples.crystals() # Dependency-free structural validation has run. +``` + +For a conditioned checkpoint, pass the exact raw ports declared by the +exported score graph, such as +`condition_values={"chemical_system": torch.from_numpy(...).reshape(B, 101)}`. +`chemical_system_multihot()` creates the one-based `[101]` value. The host +sets every `condition..use_unconditional` selector itself and performs +source classifier-free guidance in conditional-then-unconditional order. To +draw an unconditional sample from an adapter checkpoint, omit a condition +value; Mobius supplies a shape-valid internal placeholder and marks that +condition unconditional, matching the source's missing-property behavior. +The score callback receives `MatterGenScoreInputs`, including every graph +tensor, so a non-ORT inference host is equally supported. + +`MatterGenCrystal` is a validated array artifact, not a replacement for +Pymatgen's optional `Structure`/CIF APIs. The default gate fails closed on a +non-finite or non-positive-volume cell, unwrapped coordinates, unsupported +species, or a count outside 1–20. Applications that require a Pymatgen +`Structure` or CIF must perform that optional serialization after validation; +Mobius does not add Pymatgen as a production dependency. + +Official count priors support one through 20 atoms. Sampling must apply the +pinned 78-element allowlist (ending at Bi); it must not infer support from +the broader 101-class vocabulary. A chemical-system condition is a +101-element, one-based atomic-number multihot vector, where index zero is +unused. Scalar conditions use checkpoint-loaded standardization; `ml_bulk_modulus` +applies `log10` before standardization and must be strictly positive. + +## Pinned-source evidence + +The host adapter is a direct Torch port of these paths at source commit +`842ffe735f7d06cec89d56aa23d9f001e1124b30`: + +- `mattergen/common/utils/ocp_graph_utils.py::radius_graph_pbc` (periodic + candidate enumeration and nearest-neighbor truncation); +- `mattergen/common/utils/data_utils.py::get_pbc_distances` and + `mattergen/common/gemnet/gemnet.py::{reorder_symmetric_edges,get_triplets, + generate_interaction_graph}` (row-vector cell offsets, `V_st`, symmetry, + and sparse triplets); +- `mattergen/diffusion/sampling/pc_sampler.py::_denoise`, + `mattergen/diffusion/sampling/classifier_free_guidance.py::_score_fn`, and + `mattergen/diffusion/d3pm/d3pm_predictors_correctors.py:: + D3PMAncestralSamplingPredictor.update_given_score` (PC, CFG, and D3PM + ordering); +- `mattergen/common/diffusion/corruption.py` and + `mattergen/common/diffusion/predictors_correctors.py` (wrapped VE and + lattice VP processes). + +The committed L5 fixture runs all 1,000 released timesteps through a real +CPU ONNX `mp_20_base` score callback with seed `814`, then validates its +one-atom `MatterGenCrystal` artifact. Its checkpoint SHA-256 and final +species, fractional coordinates, cell, and volume are recorded in +`testdata/golden/diffusion/mattergen-mp20-host-sample.json`. Pymatgen is not a +production dependency, so that fixture validates the portable structural +artifact rather than serializing a CIF. + +## Official checkpoint families + +`mp_20_base` is the smallest official release and the default evidence +checkpoint. The config reader recognizes the following pinned families and +their proven condition inputs: + +| Checkpoint | Condition inputs | +|---|---| +| `mattergen_base`, `mp_20_base` | none | +| `chemical_system` | `chemical_system` | +| `chemical_system_energy_above_hull` | `chemical_system`, `energy_above_hull` | +| `space_group` | `space_group` | +| `dft_band_gap` | `dft_band_gap` | +| `dft_mag_density` | `dft_mag_density` | +| `dft_mag_density_hhi_score` | `dft_mag_density`, `hhi_score` | +| `ml_bulk_modulus` | `ml_bulk_modulus` | + +Each adapter family is only loadable when its Hydra configuration and +Lightning checkpoint tensor layout agree. The loader rejects a source tensor +that cannot be routed to the exported inference graph rather than silently +dropping it. + +## Building + +```bash +mobius build --model microsoft/mattergen --mattergen-checkpoint mp_20_base \ + --no-weights --output mattergen-score +``` + +The command resolves the immutable revision above by default. Only float32 and +the portable/default or CPU execution-provider paths are assessed; f16, bf16, +CUDA, and ONNX Runtime GenAI are rejected rather than advertised as equivalent +to the official float32 host pipeline. diff --git a/docs/model-catalog.md b/docs/model-catalog.md index 94c983b54..d235a18be 100644 --- a/docs/model-catalog.md +++ b/docs/model-catalog.md @@ -263,6 +263,18 @@ from mobius import build pkg = build("stabilityai/stable-diffusion-xl-base-1.0") ``` +## Periodic crystal diffusion + +| Model | Exported component | Task | Example HuggingFace Model | +|---|---|---|---| +| MatterGen | GemNet-T crystal score core | `mattergen-score` | `microsoft/mattergen` | + +MatterGen is a native periodic-crystal diffusion integration rather than a +Diffusers pipeline. Mobius exports its deterministic neural score core; the +periodic neighbor graph, diffusion scheduler, sampling, and crystal validation +remain source-compatible host responsibilities. See [MatterGen crystal +diffusion score core](mattergen.md) for the staged contract and runtime limits. + ## Quantization Support All decoder-only LLMs and MoE models support quantized weight loading: diff --git a/pyproject.toml b/pyproject.toml index 0c262c3dc..00faa5e5e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -19,6 +19,7 @@ dependencies = [ "onnx_ir>=1.0.0", "onnx-shape-inference>=0.3.1", "onnxscript>=0.7.1", + "PyYAML", "rfc8785", "safetensors", "torch>=2.10.0", diff --git a/src/mobius/__init__.py b/src/mobius/__init__.py index bffa8a4bc..921698b3a 100644 --- a/src/mobius/__init__.py +++ b/src/mobius/__init__.py @@ -59,6 +59,7 @@ "build_context", "build_diffusers_pipeline", "build_from_gguf", + "build_mattergen", "build_from_module", "build_from_nemo", "compose_adapter_deltas", @@ -145,6 +146,7 @@ from mobius.integrations._weight_loading import apply_weights, stream_safetensors_to_model from mobius.integrations.diffusers import build_diffusers_pipeline from mobius.integrations.gguf import build_from_gguf +from mobius.integrations.mattergen import build_mattergen from mobius.integrations.nemo import build_from_nemo from mobius.integrations.transformers import build from mobius.models import MLPWorldModel diff --git a/src/mobius/__main__.py b/src/mobius/__main__.py index 8a5a86a78..cac479147 100644 --- a/src/mobius/__main__.py +++ b/src/mobius/__main__.py @@ -329,12 +329,72 @@ def _resolve_static_cache_task(model_type: str) -> ModelTask: revision = REUSE_REVISION output_dir = args.output_dir - os.makedirs(output_dir, exist_ok=True) dtype_override = resolve_dtype(args.dtype) optimize = args.optimize component_filter = args.component execution_provider = args.execution_provider + from mobius.integrations.mattergen._builder import ( + build_mattergen, + is_mattergen_checkpoint, + ) + + mattergen_source = args.model or args.config + is_mattergen = is_mattergen_checkpoint(mattergen_source) + if args.mattergen_checkpoint is not None and not is_mattergen: + raise SystemExit( + "Error: --mattergen-checkpoint requires --model microsoft/mattergen or a " + "local MatterGen checkpoint root." + ) + if is_mattergen: + if task is not None: + raise SystemExit( + "Error: MatterGen uses its fixed mattergen-score task; do not pass --task." + ) + if args.runtime is not None: + raise SystemExit( + "Error: MatterGen cannot produce an ONNX Runtime GenAI package. Its " + "periodic graph and stochastic crystal scheduler remain host-owned." + ) + if component_filter is not None: + raise SystemExit( + "Error: MatterGen exports exactly one score-core component; --component is unsupported." + ) + if optimize is not None: + raise SystemExit( + "Error: Transformer rewrite rules are unsupported for MatterGen score-core exports." + ) + if ( + args.text_only + or args.static_cache + or fp8_kv_cache + or prune_prefill_prefix + or args.glm_full_attention + or export_paged_attention + or args.trust_remote_code + or args.dequantize + ): + raise SystemExit( + "Error: Transformer/compressed-weight build options are unsupported for MatterGen." + ) + if input_sampling_rate is not None or bwe_sampling_rate is not None: + raise SystemExit( + "Error: --input-sample-rate and --bwe-sample-rate are unsupported for MatterGen." + ) + try: + pkg = build_mattergen( + mattergen_source, + checkpoint=args.mattergen_checkpoint or "mp_20_base", + revision=revision, + dtype=dtype_override, + load_weights=load_weights, + execution_provider=execution_provider, + ) + except ValueError as error: + raise SystemExit(f"Error: {error}") from error + _save_package(pkg, output_dir, args, optimize, component_filter) + return + # Auto-detect diffusers pipelines. Skipped when the text-only feature is set: # that flag only applies to transformers decoder exports, so we let the # central build() validation reject a diffusers/unsupported repo rather @@ -1416,6 +1476,15 @@ def build_parser() -> argparse.ArgumentParser: default=None, help="Model task (auto-detected if not specified). Use 'mobius list tasks' to see available tasks.", ) + build_parser.add_argument( + "--mattergen-checkpoint", + default=None, + metavar="FAMILY", + help=( + "Official MatterGen checkpoint family (default: mp_20_base). Only valid " + "with --model microsoft/mattergen or a local MatterGen checkpoint root." + ), + ) build_parser.add_argument( "--no-weights", action="store_true", diff --git a/src/mobius/integrations/mattergen/__init__.py b/src/mobius/integrations/mattergen/__init__.py new file mode 100644 index 000000000..c7646b5f6 --- /dev/null +++ b/src/mobius/integrations/mattergen/__init__.py @@ -0,0 +1,100 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +"""Pinned MatterGen checkpoint configuration integration.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING + +__all__ = [ + "MATTERGEN_CONDITION_FAMILY", + "MATTERGEN_CONDITION_SPECS", + "MATTERGEN_HUB_REVISION", + "MATTERGEN_MODEL_ID", + "MATTERGEN_SOURCE_COMMIT", + "MatterGenConditionSpec", + "MatterGenConfig", + "MatterGenGemNetTModel", + "MatterGenGraph", + "MatterGenHostSampler", + "MatterGenSampleBatch", + "MatterGenScoreCallback", + "MatterGenScoreInputs", + "MatterGenScoreOutputs", + "MatterGenModel", + "MatterGenCrystal", + "build_mattergen", + "build_periodic_graph", + "create_onnxruntime_score_callback", + "is_mattergen_checkpoint", +] + +if TYPE_CHECKING: + from mobius.integrations.mattergen._builder import ( + build_mattergen, + is_mattergen_checkpoint, + ) + from mobius.integrations.mattergen._configs import ( + MATTERGEN_CONDITION_FAMILY, + MATTERGEN_CONDITION_SPECS, + MATTERGEN_HUB_REVISION, + MATTERGEN_MODEL_ID, + MATTERGEN_SOURCE_COMMIT, + MatterGenConditionSpec, + MatterGenConfig, + ) + from mobius.integrations.mattergen._runtime import ( + MatterGenCrystal, + MatterGenGraph, + MatterGenHostSampler, + MatterGenSampleBatch, + MatterGenScoreCallback, + MatterGenScoreInputs, + MatterGenScoreOutputs, + build_periodic_graph, + create_onnxruntime_score_callback, + ) + from mobius.models.mattergen import MatterGenGemNetTModel, MatterGenModel + + +def __getattr__(name: str): + """Lazily expose configuration without creating model import cycles.""" + if name in { + "MATTERGEN_CONDITION_FAMILY", + "MATTERGEN_CONDITION_SPECS", + "MATTERGEN_HUB_REVISION", + "MATTERGEN_MODEL_ID", + "MATTERGEN_SOURCE_COMMIT", + "MatterGenConditionSpec", + "MatterGenConfig", + }: + from mobius.integrations.mattergen import _configs + + return getattr(_configs, name) + if name in {"MatterGenGemNetTModel", "MatterGenModel"}: + from mobius.models.mattergen import MatterGenGemNetTModel, MatterGenModel + + return { + "MatterGenGemNetTModel": MatterGenGemNetTModel, + "MatterGenModel": MatterGenModel, + }[name] + if name in {"build_mattergen", "is_mattergen_checkpoint"}: + from mobius.integrations.mattergen import _builder + + return getattr(_builder, name) + if name in { + "MatterGenCrystal", + "MatterGenGraph", + "MatterGenHostSampler", + "MatterGenSampleBatch", + "MatterGenScoreCallback", + "MatterGenScoreInputs", + "MatterGenScoreOutputs", + "build_periodic_graph", + "create_onnxruntime_score_callback", + }: + from mobius.integrations.mattergen import _runtime + + return getattr(_runtime, name) + raise AttributeError(name) diff --git a/src/mobius/integrations/mattergen/_builder.py b/src/mobius/integrations/mattergen/_builder.py new file mode 100644 index 000000000..a5481a72a --- /dev/null +++ b/src/mobius/integrations/mattergen/_builder.py @@ -0,0 +1,183 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +"""Pinned configuration and checkpoint builder for MatterGen score-core exports.""" + +from __future__ import annotations + +__all__ = ["build_mattergen", "is_mattergen_checkpoint"] + +import dataclasses +from pathlib import Path + +import onnx_ir as ir +import yaml +from huggingface_hub import hf_hub_download + +from mobius._builder import build_from_module, resolve_dtype +from mobius._model_package import ModelPackage +from mobius.integrations.mattergen._configs import MatterGenConfig +from mobius.integrations.mattergen._contract import ( + MATTERGEN_HUB_ID, + MATTERGEN_HUB_REVISION, + OFFICIAL_CHECKPOINT_CONDITIONS, +) +from mobius.integrations.mattergen._weights import apply_mattergen_checkpoint +from mobius.models.mattergen import MatterGenModel +from mobius.tasks import MatterGenScoreTask + + +def _validate_checkpoint_family(family: str) -> str: + """Return a declared official family or reject a path-like/checkpoint typo.""" + if family not in OFFICIAL_CHECKPOINT_CONDITIONS: + options = ", ".join(sorted(OFFICIAL_CHECKPOINT_CONDITIONS)) + raise ValueError(f"Unknown MatterGen checkpoint family {family!r}. Available: {options}.") + return family + + +def _local_checkpoint_file(root: Path, family: str, name: str) -> Path: + """Resolve one local checkpoint artifact without following a link outside *root*.""" + candidate = root / "checkpoints" / family / name + if candidate.is_symlink() or not candidate.is_file(): + raise FileNotFoundError(f"MatterGen local artifact must be a regular file: {candidate}") + resolved = candidate.resolve(strict=True) + if root not in resolved.parents: + raise ValueError(f"MatterGen local artifact escapes its checkpoint root: {candidate}") + return resolved + + +def is_mattergen_checkpoint(source: str | Path) -> bool: + """Return whether *source* is the canonical Hub ID or a MatterGen directory.""" + if str(source) == MATTERGEN_HUB_ID: + return True + root = Path(source).expanduser() + if not root.is_dir() or root.is_symlink(): + return False + return any( + (root / "checkpoints" / family / "config.yaml").is_file() + for family in OFFICIAL_CHECKPOINT_CONDITIONS + ) + + +def _load_config( + source: str | Path, + family: str, + revision: str | None, + *, + load_weights: bool, +) -> tuple[dict[object, object], Path | None, str]: + """Load the family Hydra YAML from a validated local root or immutable Hub revision.""" + source_path = Path(source).expanduser() + if source_path.is_dir(): + if revision is not None: + raise ValueError("MatterGen local checkpoint directories cannot use --revision.") + root = source_path.resolve(strict=True) + if root.is_symlink(): + raise ValueError("MatterGen local checkpoint root must not be a symlink.") + config_path = _local_checkpoint_file(root, family, "config.yaml") + checkpoint_path = ( + _local_checkpoint_file(root, family, "checkpoints/last.ckpt") + if load_weights + else None + ) + effective_revision = "local" + else: + if str(source) != MATTERGEN_HUB_ID: + raise ValueError( + f"MatterGen exports only support {MATTERGEN_HUB_ID!r} or a local checkpoint root." + ) + effective_revision = MATTERGEN_HUB_REVISION if revision is None else revision + if effective_revision != MATTERGEN_HUB_REVISION: + raise ValueError( + "MatterGen requires the pinned Hub revision " + f"{MATTERGEN_HUB_REVISION}; got {effective_revision!r}." + ) + config_path = Path( + hf_hub_download( + repo_id=MATTERGEN_HUB_ID, + filename=f"checkpoints/{family}/config.yaml", + revision=effective_revision, + ) + ) + checkpoint_path = None + + with config_path.open(encoding="utf-8") as file: + parsed = yaml.safe_load(file) + if not isinstance(parsed, dict): + raise TypeError(f"MatterGen Hydra config must be a mapping: {config_path}") + return parsed, checkpoint_path, effective_revision + + +def build_mattergen( + source: str | Path = MATTERGEN_HUB_ID, + *, + checkpoint: str = "mp_20_base", + revision: str | None = None, + dtype: str | ir.DataType | None = None, + load_weights: bool = True, + execution_provider: str = "default", +) -> ModelPackage: + """Build one official MatterGen GemNet-T score core from its pinned Hydra config. + + ``source`` is either the immutable official Hub repository or a local + checkpoint root containing ``checkpoints//config.yaml`` and, when + weights are requested, ``checkpoints//checkpoints/last.ckpt``. + The model intentionally exports neither MatterGen's dynamic periodic graph + construction nor its stochastic crystal sampling host loop. + """ + family = _validate_checkpoint_family(checkpoint) + resolved_dtype = resolve_dtype(dtype) + if resolved_dtype not in {None, ir.DataType.FLOAT}: + raise ValueError( + "MatterGen score-core export is assessed only for float32; f16 and bf16 are " + "refused rather than changing the source float32 numerical contract." + ) + if execution_provider not in {"default", "cpu"}: + raise ValueError( + "MatterGen score-core export is currently assessed only for default/CPU ONNX; " + f"execution provider {execution_provider!r} is not supported." + ) + + hydra_config, local_checkpoint, effective_revision = _load_config( + source, + family, + revision, + load_weights=load_weights, + ) + config = MatterGenConfig.from_hydra_config( + hydra_config, + model_id=str(source), + revision=effective_revision, + variant=family, + ) + actual_conditions = tuple(spec.name for spec in config.condition_input_specs) + expected_conditions = OFFICIAL_CHECKPOINT_CONDITIONS[family] + if actual_conditions != expected_conditions: + raise ValueError( + f"MatterGen checkpoint family {family!r} declares condition inputs " + f"{actual_conditions!r}; expected the pinned contract {expected_conditions!r}." + ) + if resolved_dtype is not None: + config = dataclasses.replace(config, dtype=resolved_dtype) + config.validate() + + module = MatterGenModel(config) + package = build_from_module( + module, + config, + task=MatterGenScoreTask(), + execution_provider=execution_provider, + ) + package["model"].graph.name = f"{source}/{family}/model" + if load_weights: + checkpoint_path = local_checkpoint + if checkpoint_path is None: + checkpoint_path = Path( + hf_hub_download( + repo_id=MATTERGEN_HUB_ID, + filename=f"checkpoints/{family}/checkpoints/last.ckpt", + revision=effective_revision, + ) + ) + apply_mattergen_checkpoint(package, module, checkpoint_path) + return package diff --git a/src/mobius/integrations/mattergen/_builder_test.py b/src/mobius/integrations/mattergen/_builder_test.py new file mode 100644 index 000000000..8edaf4cf8 --- /dev/null +++ b/src/mobius/integrations/mattergen/_builder_test.py @@ -0,0 +1,187 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +from __future__ import annotations + +from unittest import mock + +import pytest +import yaml + +from mobius import build_from_module +from mobius.integrations.mattergen import MatterGenConfig, MatterGenModel +from mobius.integrations.mattergen import _builder +from mobius.integrations.mattergen._contract import ( + MATTERGEN_HUB_ID, + MATTERGEN_HUB_REVISION, + OFFICIAL_CHECKPOINT_CONDITIONS, +) + + +def _tiny_hydra_config(*, adapter: bool = False) -> dict[str, object]: + """Return a source-shaped configuration small enough for graph unit tests.""" + model: dict[str, object] = { + "hidden_dim": 8, + "denoise_atom_types": True, + "atom_type_diffusion": "mask", + "gemnet": { + "num_targets": 1, + "num_spherical": 2, + "num_radial": 4, + "num_blocks": 1, + "emb_size_atom": 8, + "emb_size_edge": 8, + "emb_size_trip": 4, + "emb_size_rbf": 4, + "emb_size_cbf": 4, + "emb_size_bil_trip": 4, + "num_before_skip": 1, + "num_after_skip": 1, + "num_concat": 1, + "num_atom": 1, + "max_neighbors": 4, + "max_cell_images_per_dim": 1, + "regress_stress": True, + "atom_embedding": {"with_mask_type": True}, + }, + } + if adapter: + model["property_embeddings_adapt"] = { + "dft_band_gap": { + "conditional_embedding_module": { + "_target_": "mattergen.property_embeddings.NoiseLevelEncoding" + }, + "scaler": {"_target_": "mattergen.common.utils.data_utils.StandardScalerTorch"}, + } + } + gemnet = model["gemnet"] + assert isinstance(gemnet, dict) + gemnet["condition_on_adapt"] = ["dft_band_gap"] + return {"lightning_module": {"diffusion_module": {"model": model}}} + + +class TestMatterGenGraphTask: + @pytest.mark.parametrize("adapter", [False, True]) + def test_builds_tiny_score_graph_with_exact_condition_ports(self, adapter: bool) -> None: + config = MatterGenConfig.from_hydra_config(_tiny_hydra_config(adapter=adapter)) + package = build_from_module(MatterGenModel(config), config, task="mattergen-score") + model = package["model"] + + expected_inputs = { + "atomic_numbers", + "batch", + "timestep", + "edge_index", + "edge_distance", + "edge_direction", + "edge_lattice_cosines", + "id_swap", + "id3_ba", + "id3_ca", + "id3_ragged_idx", + } + if adapter: + expected_inputs |= { + "condition.dft_band_gap", + "condition.dft_band_gap.use_unconditional", + } + assert {value.name for value in model.graph.inputs} == expected_inputs + assert [value.name for value in model.graph.outputs] == [ + "atom_logits", + "coordinate_score", + "lattice_score", + "energy", + ] + assert model.metadata_props["mobius.source_revision"] == MATTERGEN_HUB_REVISION + assert model.metadata_props["mobius.max_atoms"] == "20" + assert model.metadata_props["mobius.checkpoint_family"] == "mattergen_base" + assert package.export_report is not None + assert package.export_report.status == "partial" + assert package.export_report.component("score_core").runtime_validation_status == "validated" + + +class TestMatterGenBuilder: + def test_builds_no_weights_from_a_safe_local_checkpoint_root(self, tmp_path) -> None: + root = tmp_path / "mattergen" + config_path = root / "checkpoints" / "mp_20_base" / "config.yaml" + config_path.parent.mkdir(parents=True) + config_path.write_text(yaml.safe_dump(_tiny_hydra_config()), encoding="utf-8") + + package = _builder.build_mattergen(root, load_weights=False) + + assert package.config.variant == "mp_20_base" + assert package["model"].graph.name == f"{root}/mp_20_base/model" + assert _builder.is_mattergen_checkpoint(root) + + def test_no_weights_package_saves_an_atomic_partial_export_report(self, tmp_path) -> None: + root = tmp_path / "mattergen" + config_path = root / "checkpoints" / "mp_20_base" / "config.yaml" + config_path.parent.mkdir(parents=True) + config_path.write_text(yaml.safe_dump(_tiny_hydra_config()), encoding="utf-8") + package = _builder.build_mattergen(root, load_weights=False) + output = tmp_path / "score-core" + + package.save(output, check_weights=False, progress_bar=False) + + assert (output / "model.onnx").is_file() + assert (output / "export_report.json").is_file() + + def test_resolves_the_immutable_hub_revision_for_config_download(self, tmp_path, monkeypatch) -> None: + config_path = tmp_path / "config.yaml" + config_path.write_text(yaml.safe_dump(_tiny_hydra_config()), encoding="utf-8") + download = mock.Mock(return_value=str(config_path)) + monkeypatch.setattr(_builder, "hf_hub_download", download) + + config, checkpoint, revision = _builder._load_config( + MATTERGEN_HUB_ID, + "mp_20_base", + None, + load_weights=False, + ) + + assert config == _tiny_hydra_config() + assert checkpoint is None + assert revision == MATTERGEN_HUB_REVISION + assert download.call_args.kwargs == { + "repo_id": MATTERGEN_HUB_ID, + "filename": "checkpoints/mp_20_base/config.yaml", + "revision": MATTERGEN_HUB_REVISION, + } + + def test_rejects_mutable_or_incompatible_build_options(self) -> None: + with pytest.raises(ValueError, match="pinned Hub revision"): + _builder.build_mattergen(revision="main", load_weights=False) + with pytest.raises(ValueError, match="float32"): + _builder.build_mattergen(dtype="f16", load_weights=False) + with pytest.raises(ValueError, match="default/CPU"): + _builder.build_mattergen(execution_provider="cuda", load_weights=False) + + def test_rejects_a_local_config_with_the_wrong_declared_family_conditions(self, tmp_path) -> None: + root = tmp_path / "mattergen" + config_path = root / "checkpoints" / "mp_20_base" / "config.yaml" + config_path.parent.mkdir(parents=True) + config_path.write_text(yaml.safe_dump(_tiny_hydra_config(adapter=True)), encoding="utf-8") + + with pytest.raises(ValueError, match="expected the pinned contract"): + _builder.build_mattergen(root, load_weights=False) + + @pytest.mark.arch_validation + @pytest.mark.parametrize( + ("family", "conditions"), + sorted(OFFICIAL_CHECKPOINT_CONDITIONS.items()), + ) + def test_all_official_pinned_hydra_configs_build_score_graph( + self, + family: str, + conditions: tuple[str, ...], + ) -> None: + package = _builder.build_mattergen( + MATTERGEN_HUB_ID, + checkpoint=family, + revision=MATTERGEN_HUB_REVISION, + load_weights=False, + ) + + assert package.config.variant == family + assert tuple(spec.name for spec in package.config.condition_input_specs) == conditions + assert package["model"].graph.outputs[0].name == "atom_logits" diff --git a/src/mobius/integrations/mattergen/_configs.py b/src/mobius/integrations/mattergen/_configs.py new file mode 100644 index 000000000..f2eeea4fb --- /dev/null +++ b/src/mobius/integrations/mattergen/_configs.py @@ -0,0 +1,454 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +"""Typed configuration for the pinned MatterGen GemNet-T score core. + +The MatterGen Hub checkpoints store an expanded Hydra configuration rather +than a Transformers ``config.json``. This module deliberately accepts a +plain nested mapping so reading a checkpoint configuration does not require +Hydra, OmegaConf, or PyYAML at import time. +""" + +from __future__ import annotations + +import dataclasses +from collections.abc import Mapping +from typing import Any, ClassVar + +import onnx_ir as ir + +from mobius._configs import BaseModelConfig +from mobius.integrations.mattergen._contract import ( + MATTERGEN_HUB_ID, + MATTERGEN_HUB_REVISION, + MATTERGEN_SOURCE_COMMIT, +) + +MATTERGEN_MODEL_ID = MATTERGEN_HUB_ID + +# ``PROPERTY_SOURCE_IDS`` in MatterGen v1.0.3. Keep this full tuple even +# though only the listed entries occur in the published adapter checkpoints: +# it is the authoritative identifier family for expanded Hydra configs. +MATTERGEN_CONDITION_FAMILY: tuple[str, ...] = ( + "dft_mag_density", + "dft_bulk_modulus", + "dft_shear_modulus", + "energy_above_hull", + "formation_energy_per_atom", + "space_group", + "hhi_score", + "ml_bulk_modulus", + "chemical_system", + "dft_band_gap", +) + + +@dataclasses.dataclass(frozen=True) +class MatterGenConditionSpec: + """One declared MatterGen property-embedding input. + + ``kind`` names the source conditional encoder. Scalar properties use the + source ``NoiseLevelEncoding`` after their optional fitted standard scaler; + chemical systems are host-provided 101-wide multi-hot vectors; space-group + values are one-based indices and are decremented before their lookup. + """ + + name: str + kind: str + scaler: str = "identity" + log10_transform: bool = False + unconditional: str = "embedding_vector" + is_adapter: bool = False + + @property + def input_shape_suffix(self) -> tuple[int, ...]: + """Required host value shape excluding the batch dimension.""" + return (101,) if self.kind == "chemical_system_multihot" else () + + +# Defaults inferred from the official config groups. The two condition names +# without public group files are still represented as scalar sources because +# they are legal identifiers in MatterGen v1.0.3's property family. +MATTERGEN_CONDITION_SPECS: tuple[MatterGenConditionSpec, ...] = ( + MatterGenConditionSpec("dft_mag_density", "scalar_sinusoidal", "standard"), + MatterGenConditionSpec( + "dft_bulk_modulus", "scalar_sinusoidal", "standard", log10_transform=True + ), + MatterGenConditionSpec("dft_shear_modulus", "scalar_sinusoidal", "standard"), + MatterGenConditionSpec("energy_above_hull", "scalar_sinusoidal", "standard"), + MatterGenConditionSpec("formation_energy_per_atom", "scalar_sinusoidal", "standard"), + MatterGenConditionSpec("space_group", "space_group_index"), + MatterGenConditionSpec("hhi_score", "scalar_sinusoidal", "standard"), + MatterGenConditionSpec( + "ml_bulk_modulus", "scalar_sinusoidal", "standard", log10_transform=True + ), + MatterGenConditionSpec("chemical_system", "chemical_system_multihot"), + MatterGenConditionSpec("dft_band_gap", "scalar_sinusoidal", "standard"), +) + +_SPEC_BY_NAME: dict[str, MatterGenConditionSpec] = { + spec.name: spec for spec in MATTERGEN_CONDITION_SPECS +} + + +def _mapping(value: object | None) -> Mapping[str, Any] | None: + """Return a string-keyed mapping, including OmegaConf-like mapping values.""" + if isinstance(value, Mapping): + return value + items = getattr(value, "items", None) + if callable(items): + raw_items = items() + return {str(key): item for key, item in raw_items} + return None + + +def _get(value: object | None, key: str, default: object | None = None) -> object | None: + """Read *key* from a mapping or attribute-style Hydra object.""" + mapping = _mapping(value) + if mapping is not None: + return mapping.get(key, default) + return getattr(value, key, default) + + +def _nested(value: object | None, *keys: str) -> object | None: + """Look up nested mapping/object fields without imposing a config library.""" + current = value + for key in keys: + current = _get(current, key) + if current is None: + return None + return current + + +def _as_tuple(value: object | None) -> tuple[str, ...]: + """Normalize a Hydra list-like value to a tuple of strings.""" + if value is None: + return () + if isinstance(value, str): + return (value,) + if isinstance(value, (list, tuple)): + return tuple(str(item) for item in value) + raise TypeError(f"Expected a condition sequence, got {type(value).__name__}") + + +def _target_name(value: object | None) -> str: + """Read a Hydra ``_target_`` value, accepting absent identity declarations.""" + target = _get(value, "_target_", "") + return target if isinstance(target, str) else "" + + +def _condition_spec( + name: str, + source: object | None, + *, + is_adapter: bool, +) -> MatterGenConditionSpec: + """Parse one expanded official ``PropertyEmbedding`` mapping.""" + if name not in _SPEC_BY_NAME: + raise ValueError(f"Unsupported MatterGen condition {name!r}") + + fallback = _SPEC_BY_NAME[name] + conditional = _get(source, "conditional_embedding_module") + conditional_target = _target_name(conditional) + if not conditional_target: + kind = fallback.kind + elif conditional_target.endswith("NoiseLevelEncoding"): + kind = "scalar_sinusoidal" + elif conditional_target.endswith("ChemicalSystemMultiHotEmbedding"): + kind = "chemical_system_multihot" + elif conditional_target.endswith("SpaceGroupEmbeddingVector"): + kind = "space_group_index" + else: + raise ValueError( + f"Unsupported MatterGen condition encoder for {name!r}: {conditional_target!r}" + ) + + scaler_config = _get(source, "scaler") + scaler_target = _target_name(scaler_config) + if not scaler_target or scaler_target.endswith("Identity"): + scaler = "identity" + elif scaler_target.endswith("StandardScalerTorch"): + scaler = "standard" + else: + raise ValueError(f"Unsupported MatterGen scaler for {name!r}: {scaler_target!r}") + + log10_transform = bool(_get(scaler_config, "log10_transform", False)) + if kind != "scalar_sinusoidal" and scaler != "identity": + raise ValueError(f"Only scalar MatterGen conditions may use a scaler: {name!r}") + return MatterGenConditionSpec( + name=name, + kind=kind, + scaler=scaler, + log10_transform=log10_transform, + # GemNetTAdapter replaces this source module with ZerosEmbedding after + # construction, so its checkpoint contains no unconditional vector. + unconditional="zeros" if is_adapter else "embedding_vector", + is_adapter=is_adapter, + ) + + +def _condition_specs( + value: object | None, *, is_adapter: bool +) -> tuple[MatterGenConditionSpec, ...]: + """Parse a Hydra ModuleDict mapping in source's sorted concatenation order.""" + entries = _mapping(value) + if not entries: + return () + return tuple( + _condition_spec(str(name), source, is_adapter=is_adapter) + for name, source in sorted(entries.items()) + ) + + +@dataclasses.dataclass +class MatterGenConfig(BaseModelConfig): + """Configuration of the v1.0.3 GemNet-T MatterGen neural score core.""" + + model_id: str = MATTERGEN_MODEL_ID + revision: str = MATTERGEN_HUB_REVISION + source_commit: str = MATTERGEN_SOURCE_COMMIT + variant: str = "mattergen_base" + + # BaseModelConfig compatibility fields. + vocab_size: int = 101 + hidden_size: int = 512 + num_hidden_layers: int = 4 + hidden_act: str | None = "silu" + dtype: ir.DataType = ir.DataType.FLOAT + + # GemNet-T dimensions and graph construction bounds from mattergen.yaml. + num_targets: int = 1 + num_spherical: int = 7 + num_radial: int = 128 + num_blocks: int = 4 + emb_size_atom: int = 512 + emb_size_edge: int = 512 + emb_size_trip: int = 64 + emb_size_rbf: int = 16 + emb_size_cbf: int = 16 + emb_size_bil_trip: int = 64 + num_before_skip: int = 1 + num_after_skip: int = 2 + num_concat: int = 1 + num_atom: int = 3 + cutoff: float = 7.0 + max_neighbors: int = 50 + max_cell_images_per_dim: int = 5 + num_atom_types: int = 101 + denoise_atom_types: bool = True + atom_type_diffusion: str = "mask" + regress_stress: bool = True + + # ``condition_family`` documents all legal source identifiers; the two + # embedding tuples declare only paths actually instantiated in this graph. + condition_family: tuple[str, ...] = MATTERGEN_CONDITION_FAMILY + condition_catalog: tuple[MatterGenConditionSpec, ...] = MATTERGEN_CONDITION_SPECS + property_embeddings: tuple[MatterGenConditionSpec, ...] = () + property_embeddings_adapt: tuple[MatterGenConditionSpec, ...] = () + condition_on_adapt: tuple[str, ...] = () + + # A later task consumes this stable, config-derived list to declare one + # tensor value and one bool ``use_unconditional`` mask per condition. + condition_input_prefix: str = "condition" + + _MODEL_PATH: ClassVar[tuple[str, ...]] = ( + "lightning_module", + "diffusion_module", + "model", + ) + + @property + def latent_dim(self) -> int: + """GemNet atom latent width: noise encoding plus base property encodings.""" + return self.hidden_size * (1 + len(self.property_embeddings)) + + @property + def condition_input_specs(self) -> tuple[MatterGenConditionSpec, ...]: + """Config-ordered inputs: source sorts each ModuleDict lexicographically.""" + return self.property_embeddings + self.property_embeddings_adapt + + def validate(self) -> None: + """Validate the source-compatible dimensions and property declarations.""" + if self.dtype != ir.DataType.FLOAT: + raise ValueError( + "MatterGen score-core exports preserve the official float32 numerical contract; " + "f16 and bf16 are not assessed." + ) + positive_fields = ( + "hidden_size", + "num_targets", + "num_spherical", + "num_radial", + "num_blocks", + "emb_size_atom", + "emb_size_edge", + "emb_size_trip", + "emb_size_rbf", + "emb_size_cbf", + "emb_size_bil_trip", + "num_before_skip", + "num_after_skip", + "num_concat", + "num_atom", + "max_neighbors", + "max_cell_images_per_dim", + ) + for name in positive_fields: + if not isinstance(getattr(self, name), int) or getattr(self, name) <= 0: + raise ValueError(f"{name} must be a positive integer") + if self.cutoff <= 0: + raise ValueError("cutoff must be positive") + if self.num_targets != 1: + raise ValueError("MatterGen v1 score checkpoints require num_targets=1") + if self.hidden_size != self.emb_size_atom or self.hidden_size != self.emb_size_edge: + raise ValueError( + "MatterGen's hidden_dim, emb_size_atom, and emb_size_edge must match" + ) + if self.num_atom_types != 101 or self.vocab_size != 101: + raise ValueError("MatterGen mask diffusion requires exactly 101 atom logits") + if self.atom_type_diffusion != "mask" or not self.denoise_atom_types: + raise ValueError( + "Only the official masked atom-type diffusion score core is supported" + ) + if not self.regress_stress: + raise ValueError("MatterGen score core requires the official lattice update head") + + names = [spec.name for spec in self.condition_input_specs] + if len(names) != len(set(names)): + raise ValueError("A MatterGen condition cannot occur in both embedding families") + for spec in self.condition_input_specs: + if spec.name not in self.condition_family: + raise ValueError( + f"Condition {spec.name!r} is outside the MatterGen condition family" + ) + if spec.kind not in { + "scalar_sinusoidal", + "chemical_system_multihot", + "space_group_index", + }: + raise ValueError(f"Unsupported condition encoder kind {spec.kind!r}") + if spec.scaler not in {"identity", "standard"}: + raise ValueError(f"Unsupported condition scaler {spec.scaler!r}") + if spec.kind != "scalar_sinusoidal" and ( + spec.scaler != "identity" or spec.log10_transform + ): + raise ValueError(f"Only scalar condition {spec.name!r} may be scaled") + adapter_names = tuple(spec.name for spec in self.property_embeddings_adapt) + if self.condition_on_adapt != adapter_names: + raise ValueError( + "condition_on_adapt must exactly match property_embeddings_adapt in sorted order" + ) + + @classmethod + def from_hydra_config( + cls, + config: object, + *, + model_id: str = MATTERGEN_MODEL_ID, + revision: str = MATTERGEN_HUB_REVISION, + source_commit: str = MATTERGEN_SOURCE_COMMIT, + variant: str | None = None, + ) -> MatterGenConfig: + """Parse an expanded official Hydra config without importing YAML tooling. + + ``config`` may be the whole saved configuration, the nested + ``lightning_module.diffusion_module.model`` mapping, or an OmegaConf-like + attribute mapping. The method extracts only source-proven settings and + leaves graph construction and sampling ownership to the host task. + """ + defaults = cls() + # Released adapter configs retain the base denoiser beneath + # ``lightning_module`` for training, but their checkpoint belongs to + # ``adapter.adapter``. Prefer that concrete score module when present. + model = _nested(config, "adapter", "adapter") + if model is None: + model = _nested(config, *cls._MODEL_PATH) + if model is None: + candidate = _get(config, "model") + model = candidate if _get(candidate, "gemnet") is not None else config + gemnet = _get(model, "gemnet") + if gemnet is None: + gemnet = model + + base_properties = _condition_specs( + _get(model, "property_embeddings"), is_adapter=False + ) + adapter_properties = _condition_specs( + _get(model, "property_embeddings_adapt"), is_adapter=True + ) + declared_adapt = _as_tuple(_get(gemnet, "condition_on_adapt")) + if not declared_adapt: + declared_adapt = tuple(spec.name for spec in adapter_properties) + expected_adapt = tuple(spec.name for spec in adapter_properties) + if declared_adapt != expected_adapt: + raise ValueError( + "GemNet condition_on_adapt must match property_embeddings_adapt in sorted order" + ) + + atom_embedding = _get(gemnet, "atom_embedding") + with_mask_type = _get(atom_embedding, "with_mask_type", True) + if with_mask_type is not True: + raise ValueError("Only official mask-type AtomEmbedding checkpoints are supported") + + def source_value(name: str, default: Any) -> Any: + value = _get(gemnet, name) + return default if value is None else value + + hidden_size = _get(model, "hidden_dim", defaults.hidden_size) + if isinstance(hidden_size, bool) or not isinstance(hidden_size, int): + raise TypeError("MatterGen model.hidden_dim must be an integer.") + parsed_variant = variant or _get(config, "variant") or _get(model, "variant") + result = cls( + model_id=model_id, + revision=revision, + source_commit=source_commit, + variant=str(parsed_variant or defaults.variant), + vocab_size=101, + hidden_size=hidden_size, + num_hidden_layers=int(source_value("num_blocks", defaults.num_hidden_layers)), + num_targets=int(source_value("num_targets", defaults.num_targets)), + num_spherical=int(source_value("num_spherical", defaults.num_spherical)), + num_radial=int(source_value("num_radial", defaults.num_radial)), + num_blocks=int(source_value("num_blocks", defaults.num_blocks)), + emb_size_atom=int(source_value("emb_size_atom", defaults.emb_size_atom)), + emb_size_edge=int(source_value("emb_size_edge", defaults.emb_size_edge)), + emb_size_trip=int(source_value("emb_size_trip", defaults.emb_size_trip)), + emb_size_rbf=int(source_value("emb_size_rbf", defaults.emb_size_rbf)), + emb_size_cbf=int(source_value("emb_size_cbf", defaults.emb_size_cbf)), + emb_size_bil_trip=int( + source_value("emb_size_bil_trip", defaults.emb_size_bil_trip) + ), + num_before_skip=int(source_value("num_before_skip", defaults.num_before_skip)), + num_after_skip=int(source_value("num_after_skip", defaults.num_after_skip)), + num_concat=int(source_value("num_concat", defaults.num_concat)), + num_atom=int(source_value("num_atom", defaults.num_atom)), + cutoff=float(source_value("cutoff", defaults.cutoff)), + max_neighbors=int(source_value("max_neighbors", defaults.max_neighbors)), + max_cell_images_per_dim=int( + source_value("max_cell_images_per_dim", defaults.max_cell_images_per_dim) + ), + num_atom_types=101, + denoise_atom_types=bool( + _get(model, "denoise_atom_types", defaults.denoise_atom_types) + ), + atom_type_diffusion=str( + _get(model, "atom_type_diffusion", defaults.atom_type_diffusion) + ), + regress_stress=bool(source_value("regress_stress", defaults.regress_stress)), + property_embeddings=base_properties, + property_embeddings_adapt=adapter_properties, + condition_on_adapt=declared_adapt, + ) + result.validate() + return result + + +__all__ = [ + "MATTERGEN_CONDITION_FAMILY", + "MATTERGEN_CONDITION_SPECS", + "MATTERGEN_HUB_REVISION", + "MATTERGEN_MODEL_ID", + "MATTERGEN_SOURCE_COMMIT", + "MatterGenConditionSpec", + "MatterGenConfig", +] diff --git a/src/mobius/integrations/mattergen/_contract.py b/src/mobius/integrations/mattergen/_contract.py new file mode 100644 index 000000000..ad19f76e6 --- /dev/null +++ b/src/mobius/integrations/mattergen/_contract.py @@ -0,0 +1,323 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +"""Host-owned contract for the pinned Microsoft MatterGen score model. + +The ONNX component produced by this integration is deliberately only the +deterministic GemNet-T score core. MatterGen rebuilds a ragged periodic +neighbor graph for each diffusion evaluation; its D3PM/SDE sampling loop, +classifier-free guidance, coordinate wrapping, and final crystal validation +therefore remain a source-compatible host responsibility. +""" + +from __future__ import annotations + +__all__ = [ + "MATTERGEN_HUB_ID", + "MATTERGEN_HUB_REVISION", + "MATTERGEN_SOURCE_COMMIT", + "MATTERGEN_SOURCE_REPOSITORY", + "MAX_ATOMS", + "MAX_ATOMIC_NUMBER", + "HOST_OWNED_STEPS", + "OFFICIAL_CHECKPOINT_CONDITIONS", + "SELECTED_ATOMIC_NUMBERS", + "chemical_system_multihot", + "validate_final_crystal", +] + +from collections.abc import Sequence + +import numpy as np + +# The Hub commit pins configs and Lightning checkpoints together. It is not +# the MatterGen source commit; 842ffe is the v1.0.3 implementation reference. +MATTERGEN_HUB_ID = "microsoft/mattergen" +MATTERGEN_HUB_REVISION = "5244495dd9a979ff71abc7548a0b14b9deb0069a" +MATTERGEN_SOURCE_REPOSITORY = "https://github.com/microsoft/mattergen" +MATTERGEN_SOURCE_COMMIT = "842ffe735f7d06cec89d56aa23d9f001e1124b30" + +MAX_ATOMS = 20 +MAX_ATOMIC_NUMBER = 100 + +# These stages must be performed by a caller around every ONNX score-core +# invocation. The order preserves MatterGen's source sampling semantics. +HOST_OWNED_STEPS = ( + "periodic_radius_graph", + "symmetric_edge_reordering", + "sparse_triplet_construction", + "d3pm_species_sampling", + "wrapped_ve_coordinate_update", + "vp_lattice_update", + "classifier_free_guidance", + "fractional_coordinate_wrapping", + "lattice_projection", + "crystal_validation", +) + +# ``mattergen.common.utils.globals.SELECTED_ATOMIC_NUMBERS`` at the pinned +# source commit. It is the sampling allowlist, not the D3PM vocabulary. +SELECTED_ATOMIC_NUMBERS = ( + 1, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 11, + 12, + 13, + 14, + 15, + 16, + 17, + 19, + 20, + 21, + 22, + 23, + 24, + 25, + 26, + 27, + 28, + 29, + 30, + 31, + 32, + 33, + 34, + 35, + 37, + 38, + 39, + 40, + 41, + 42, + 44, + 45, + 46, + 47, + 48, + 49, + 50, + 51, + 52, + 53, + 55, + 56, + 57, + 58, + 59, + 60, + 62, + 63, + 64, + 65, + 66, + 67, + 68, + 69, + 70, + 71, + 72, + 73, + 74, + 75, + 76, + 77, + 78, + 79, + 80, + 81, + 82, + 83, +) + +# These names are a checkpoint routing contract, derived from each pinned +# Hydra config's ``condition_on_adapt`` field. The score graph accepts raw +# values plus an explicit per-condition unconditional selector for every item. +OFFICIAL_CHECKPOINT_CONDITIONS: dict[str, tuple[str, ...]] = { + "mattergen_base": (), + "mp_20_base": (), + "chemical_system": ("chemical_system",), + "chemical_system_energy_above_hull": ("chemical_system", "energy_above_hull"), + "space_group": ("space_group",), + "dft_band_gap": ("dft_band_gap",), + "dft_mag_density": ("dft_mag_density",), + "dft_mag_density_hhi_score": ("dft_mag_density", "hhi_score"), + "ml_bulk_modulus": ("ml_bulk_modulus",), +} + +_SYMBOLS = ( + "", + "H", + "He", + "Li", + "Be", + "B", + "C", + "N", + "O", + "F", + "Ne", + "Na", + "Mg", + "Al", + "Si", + "P", + "S", + "Cl", + "Ar", + "K", + "Ca", + "Sc", + "Ti", + "V", + "Cr", + "Mn", + "Fe", + "Co", + "Ni", + "Cu", + "Zn", + "Ga", + "Ge", + "As", + "Se", + "Br", + "Kr", + "Rb", + "Sr", + "Y", + "Zr", + "Nb", + "Mo", + "Tc", + "Ru", + "Rh", + "Pd", + "Ag", + "Cd", + "In", + "Sn", + "Sb", + "Te", + "I", + "Xe", + "Cs", + "Ba", + "La", + "Ce", + "Pr", + "Nd", + "Pm", + "Sm", + "Eu", + "Gd", + "Tb", + "Dy", + "Ho", + "Er", + "Tm", + "Yb", + "Lu", + "Hf", + "Ta", + "W", + "Re", + "Os", + "Ir", + "Pt", + "Au", + "Hg", + "Tl", + "Pb", + "Bi", + "Po", + "At", + "Rn", + "Fr", + "Ra", + "Ac", + "Th", + "Pa", + "U", + "Np", + "Pu", + "Am", + "Cm", + "Bk", + "Cf", + "Es", + "Fm", +) +_ATOMIC_NUMBER_BY_SYMBOL = {symbol.casefold(): index for index, symbol in enumerate(_SYMBOLS)} +_SELECTED_ATOMIC_NUMBER_SET = frozenset(SELECTED_ATOMIC_NUMBERS) + + +def chemical_system_multihot(chemical_system: str | Sequence[str]) -> np.ndarray: + """Convert MatterGen chemical-system input to its ``[101]`` float vector. + + A string uses the upstream hyphen-separated convention (for example, + ``"Li-O"``). Atomic-number zero is intentionally unused. Inputs are + constrained to the upstream generation allowlist because requesting a + disallowed element would make the host's mandatory sampling logit mask + unsatisfiable. + """ + symbols = chemical_system.split("-") if isinstance(chemical_system, str) else chemical_system + if not symbols: + raise ValueError("chemical_system must contain at least one element.") + + multihot = np.zeros(MAX_ATOMIC_NUMBER + 1, dtype=np.float32) + seen: set[int] = set() + for raw_symbol in symbols: + if not isinstance(raw_symbol, str) or not raw_symbol: + raise ValueError("chemical_system elements must be non-empty symbols.") + atomic_number = _ATOMIC_NUMBER_BY_SYMBOL.get(raw_symbol.casefold()) + if atomic_number is None or atomic_number == 0: + raise ValueError(f"Unknown chemical-system element: {raw_symbol!r}.") + if atomic_number not in _SELECTED_ATOMIC_NUMBER_SET: + raise ValueError( + f"Element {raw_symbol!r} (Z={atomic_number}) is outside MatterGen's " + "pinned sampling allowlist." + ) + if atomic_number in seen: + raise ValueError(f"chemical_system contains duplicate element {raw_symbol!r}.") + seen.add(atomic_number) + multihot[atomic_number] = 1.0 + return multihot + + +def validate_final_crystal( + atomic_numbers: np.ndarray, + fractional_coordinates: np.ndarray, + cell: np.ndarray, +) -> None: + """Validate the bounded crystal artifact emitted by a MatterGen host loop. + + This is deliberately structural validation only. It does not replace + Pymatgen's site-distance checks or create a crystal object, both of which + remain host implementation choices outside the pure score graph. + """ + numbers = np.asarray(atomic_numbers) + fractional = np.asarray(fractional_coordinates) + lattice = np.asarray(cell) + if numbers.ndim != 1 or not 1 <= len(numbers) <= MAX_ATOMS: + raise ValueError(f"atomic_numbers must have a length in [1, {MAX_ATOMS}].") + if fractional.shape != (len(numbers), 3): + raise ValueError("fractional_coordinates must have shape [N, 3].") + if lattice.shape != (3, 3): + raise ValueError("cell must have shape [3, 3].") + if not np.isfinite(fractional).all() or not np.isfinite(lattice).all(): + raise ValueError("fractional_coordinates and cell must be finite.") + if not np.issubdtype(numbers.dtype, np.integer): + raise TypeError("atomic_numbers must use an integer dtype.") + if any(int(number) not in _SELECTED_ATOMIC_NUMBER_SET for number in numbers): + raise ValueError("atomic_numbers contains an element outside MatterGen's allowlist.") + if np.any(fractional < 0.0) or np.any(fractional >= 1.0): + raise ValueError("fractional_coordinates must be wrapped to [0, 1).") + if not np.isfinite(np.linalg.det(lattice)) or np.linalg.det(lattice) <= 0.0: + raise ValueError("cell must have positive finite volume.") diff --git a/src/mobius/integrations/mattergen/_contract_test.py b/src/mobius/integrations/mattergen/_contract_test.py new file mode 100644 index 000000000..d7d3a2004 --- /dev/null +++ b/src/mobius/integrations/mattergen/_contract_test.py @@ -0,0 +1,100 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +from __future__ import annotations + +import numpy as np +import pytest + +from mobius.integrations.mattergen._contract import ( + HOST_OWNED_STEPS, + MATTERGEN_HUB_REVISION, + MATTERGEN_SOURCE_COMMIT, + OFFICIAL_CHECKPOINT_CONDITIONS, + chemical_system_multihot, + validate_final_crystal, +) + + +class TestMatterGenHostContract: + def test_pins_the_hub_and_source_implementation(self) -> None: + assert MATTERGEN_HUB_REVISION == "5244495dd9a979ff71abc7548a0b14b9deb0069a" + assert MATTERGEN_SOURCE_COMMIT == "842ffe735f7d06cec89d56aa23d9f001e1124b30" + + def test_declares_all_official_checkpoint_condition_contracts(self) -> None: + assert OFFICIAL_CHECKPOINT_CONDITIONS == { + "mattergen_base": (), + "mp_20_base": (), + "chemical_system": ("chemical_system",), + "chemical_system_energy_above_hull": ("chemical_system", "energy_above_hull"), + "space_group": ("space_group",), + "dft_band_gap": ("dft_band_gap",), + "dft_mag_density": ("dft_mag_density",), + "dft_mag_density_hhi_score": ("dft_mag_density", "hhi_score"), + "ml_bulk_modulus": ("ml_bulk_modulus",), + } + + def test_declares_the_complete_host_orchestration_boundary(self) -> None: + assert HOST_OWNED_STEPS == ( + "periodic_radius_graph", + "symmetric_edge_reordering", + "sparse_triplet_construction", + "d3pm_species_sampling", + "wrapped_ve_coordinate_update", + "vp_lattice_update", + "classifier_free_guidance", + "fractional_coordinate_wrapping", + "lattice_projection", + "crystal_validation", + ) + + def test_chemical_system_uses_one_based_atomic_number_slots(self) -> None: + multihot = chemical_system_multihot("Li-O") + + assert multihot.dtype == np.float32 + assert multihot.shape == (101,) + np.testing.assert_array_equal(multihot[[0, 3, 8]], [0.0, 1.0, 1.0]) + assert np.count_nonzero(multihot) == 2 + + @pytest.mark.parametrize("chemical_system", ["He", "Li-Li", "Xx", ""]) + def test_chemical_system_rejects_unsampleable_or_invalid_elements( + self, chemical_system: str + ) -> None: + with pytest.raises(ValueError): + chemical_system_multihot(chemical_system) + + def test_accepts_wrapped_bounded_crystal(self) -> None: + validate_final_crystal( + np.array([3, 8], dtype=np.int64), + np.array([[0.0, 0.5, 0.9], [0.25, 0.75, 0.125]], dtype=np.float32), + np.diag(np.array([3.0, 3.0, 3.0], dtype=np.float32)), + ) + + @pytest.mark.parametrize( + ("numbers", "fractional", "cell", "message"), + [ + ( + np.array([3], dtype=np.int64), + np.array([[1.0, 0.0, 0.0]], dtype=np.float32), + np.eye(3, dtype=np.float32), + "wrapped", + ), + ( + np.array([2], dtype=np.int64), + np.array([[0.0, 0.0, 0.0]], dtype=np.float32), + np.eye(3, dtype=np.float32), + "allowlist", + ), + ( + np.array([3], dtype=np.int64), + np.array([[0.0, 0.0, 0.0]], dtype=np.float32), + np.diag(np.array([-1.0, 1.0, 1.0], dtype=np.float32)), + "positive", + ), + ], + ) + def test_rejects_invalid_final_crystal( + self, numbers: np.ndarray, fractional: np.ndarray, cell: np.ndarray, message: str + ) -> None: + with pytest.raises(ValueError, match=message): + validate_final_crystal(numbers, fractional, cell) diff --git a/src/mobius/integrations/mattergen/_runtime.py b/src/mobius/integrations/mattergen/_runtime.py new file mode 100644 index 000000000..dff448724 --- /dev/null +++ b/src/mobius/integrations/mattergen/_runtime.py @@ -0,0 +1,1199 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +"""Source-faithful host runtime for the MatterGen v1.0.3 score-core ABI. + +This module ports the host-only portions of the pinned MatterGen implementation +without importing MatterGen, PyTorch Geometric, ``torch_scatter``, or +``torch_sparse``. In particular, :func:`build_periodic_graph` reproduces the +periodic-image enumeration and GemNet-T edge/triplet ordering used by +``mattergen.common.utils.ocp_graph_utils.radius_graph_pbc`` and +``mattergen.common.gemnet.gemnet.GemNetT.generate_interaction_graph``. + +The score callback is deliberately explicit. It can call ONNX Runtime (using +:func:`create_onnxruntime_score_callback`) or another execution host, while +this module owns the scientifically significant graph and sampler semantics. +""" + +from __future__ import annotations + +import math +from collections.abc import Callable, Mapping, Sequence +from dataclasses import dataclass +from typing import Protocol + +import numpy as np +import torch + +from mobius.integrations.mattergen._configs import MATTERGEN_CONDITION_FAMILY +from mobius.integrations.mattergen._contract import ( + MAX_ATOMS, + SELECTED_ATOMIC_NUMBERS, + validate_final_crystal, +) + +__all__ = [ + "MATTERGEN_LATTICE_LIMIT_DENSITY", + "MATTERGEN_SAMPLING_STEPS", + "MatterGenCrystal", + "MatterGenGraph", + "MatterGenHostSampler", + "MatterGenSampleBatch", + "MatterGenScoreCallback", + "MatterGenScoreInputs", + "MatterGenScoreOutputs", + "build_periodic_graph", + "create_onnxruntime_score_callback", +] + +# ``mattergen/conf/data_module/{mp_20,alex_mp_20}.yaml``. Both published +# generation datasets resolve the corruption interpolation to this same value. +MATTERGEN_LATTICE_LIMIT_DENSITY = 0.05771451654022283 +MATTERGEN_SAMPLING_STEPS = 1000 +_EPS_T = 1.0 / MATTERGEN_SAMPLING_STEPS +_D3PM_CLASSES = 101 +_D3PM_MASK_CLASS = _D3PM_CLASSES - 1 +_D3PM_EPSILON = 1e-20 +# ``sampling_conf/default.yaml`` configures both released Langevin correctors +# with this cap, including their defined all-zero-score branch. +_LANGEVIN_MAX_STEP_SIZE = 1e6 +# Materialized exactly as ``MaskDiffusion._create_state`` does for +# ``create_discrete_diffusion_schedule(kind="standard", num_steps=1000)``. +_D3PM_BETAS = torch.cat( + [ + torch.tensor([0.0], device="cpu"), + 1.0 + / (MATTERGEN_SAMPLING_STEPS - torch.arange(MATTERGEN_SAMPLING_STEPS, device="cpu")), + ] +).to(torch.float64) +_D3PM_STATE = torch.cumprod(1.0 - _D3PM_BETAS, dim=0).to(torch.float32) +_D3PM_STATE[-1] = 0.0 + + +@dataclass(frozen=True) +class MatterGenGraph: + """Ragged source-ordered GemNet-T geometry for one batched score call. + + ``edge_index`` has source/target rows. ``edge_direction`` is the source + ``V_st`` convention, i.e. the negative normalized periodic distance vector. + """ + + edge_index: torch.Tensor + edge_distance: torch.Tensor + edge_direction: torch.Tensor + edge_lattice_cosines: torch.Tensor + id_swap: torch.Tensor + id3_ba: torch.Tensor + id3_ca: torch.Tensor + id3_ragged_idx: torch.Tensor + + +@dataclass(frozen=True) +class MatterGenScoreInputs: + """All inputs required by one pure-ONNX MatterGen score-core invocation.""" + + atomic_numbers: torch.Tensor + batch: torch.Tensor + timestep: torch.Tensor + graph: MatterGenGraph + condition_values: Mapping[str, torch.Tensor] + use_unconditional: Mapping[str, torch.Tensor] + + def as_onnx_inputs(self) -> dict[str, np.ndarray]: + """Return the exact named ABI expected by ``MatterGenScoreTask``. + + ONNX Runtime consumes CPU NumPy arrays. Keeping this conversion at the + callback boundary prevents scheduler code from depending on an ORT API. + """ + graph = self.graph + values = { + "atomic_numbers": _as_numpy(self.atomic_numbers), + "batch": _as_numpy(self.batch), + "timestep": _as_numpy(self.timestep), + "edge_index": _as_numpy(graph.edge_index), + "edge_distance": _as_numpy(graph.edge_distance), + "edge_direction": _as_numpy(graph.edge_direction), + "edge_lattice_cosines": _as_numpy(graph.edge_lattice_cosines), + "id_swap": _as_numpy(graph.id_swap), + "id3_ba": _as_numpy(graph.id3_ba), + "id3_ca": _as_numpy(graph.id3_ca), + "id3_ragged_idx": _as_numpy(graph.id3_ragged_idx), + } + values.update( + { + f"condition.{name}": _as_numpy(value) + for name, value in self.condition_values.items() + } + ) + values.update( + { + f"condition.{name}.use_unconditional": _as_numpy(value) + for name, value in self.use_unconditional.items() + } + ) + return values + + +@dataclass(frozen=True) +class MatterGenScoreOutputs: + """Raw score-core outputs before MatterGen host postprocessing.""" + + atom_logits: torch.Tensor + coordinate_score: torch.Tensor + lattice_score: torch.Tensor + energy: torch.Tensor | None = None + + +MatterGenScoreCallback = Callable[[MatterGenScoreInputs], MatterGenScoreOutputs] + + +class _OnnxRuntimeSession(Protocol): + """Minimal structural type accepted from an ONNX Runtime inference session.""" + + def run( + self, output_names: Sequence[str] | None, input_feed: Mapping[str, np.ndarray] + ) -> Sequence[np.ndarray]: ... + + +def create_onnxruntime_score_callback(session: _OnnxRuntimeSession) -> MatterGenScoreCallback: + """Adapt an ``onnxruntime.InferenceSession`` without importing onnxruntime. + + The returned callback preserves the model's raw Cartesian coordinate score; + :class:`MatterGenHostSampler` performs the source ``cell^{-T}`` + conversion before applying the wrapped VE scheduler. + """ + + def score(inputs: MatterGenScoreInputs) -> MatterGenScoreOutputs: + atom_logits, coordinate_score, lattice_score, energy = session.run( + ["atom_logits", "coordinate_score", "lattice_score", "energy"], + inputs.as_onnx_inputs(), + ) + return MatterGenScoreOutputs( + atom_logits=torch.from_numpy(atom_logits), + coordinate_score=torch.from_numpy(coordinate_score), + lattice_score=torch.from_numpy(lattice_score), + energy=torch.from_numpy(energy), + ) + + return score + + +@dataclass(frozen=True) +class MatterGenCrystal: + """A validated crystal artifact represented in MatterGen row-vector convention.""" + + atomic_numbers: torch.Tensor + fractional_coordinates: torch.Tensor + cell: torch.Tensor + + def validate(self) -> None: + """Run dependency-free structural checks before exposing this artifact.""" + validate_final_crystal( + self.atomic_numbers.detach().cpu().numpy(), + self.fractional_coordinates.detach().cpu().numpy(), + self.cell.detach().cpu().numpy(), + ) + + +@dataclass(frozen=True) +class MatterGenSampleBatch: + """Final mean sample returned by MatterGen's predictor-corrector sampler.""" + + atomic_numbers: torch.Tensor + fractional_coordinates: torch.Tensor + cell: torch.Tensor + num_atoms: torch.Tensor + + @property + def batch(self) -> torch.Tensor: + """Source-compatible crystal index for each atom.""" + return torch.repeat_interleave( + torch.arange(len(self.num_atoms), dtype=torch.long, device=self.num_atoms.device), + self.num_atoms, + ) + + def crystals(self) -> tuple[MatterGenCrystal, ...]: + """Split the packed batch and validate every final crystal.""" + crystals: list[MatterGenCrystal] = [] + start = 0 + for count, lattice in zip(self.num_atoms.tolist(), self.cell, strict=True): + stop = start + count + crystal = MatterGenCrystal( + atomic_numbers=self.atomic_numbers[start:stop], + fractional_coordinates=self.fractional_coordinates[start:stop], + cell=lattice, + ) + crystal.validate() + crystals.append(crystal) + start = stop + return tuple(crystals) + + +@dataclass(frozen=True) +class _State: + """Packed fields corrupted jointly by the source multi-corruption sampler.""" + + atomic_numbers: torch.Tensor + fractional_coordinates: torch.Tensor + cell: torch.Tensor + num_atoms: torch.Tensor + + @property + def batch(self) -> torch.Tensor: + return torch.repeat_interleave( + torch.arange(len(self.num_atoms), dtype=torch.long, device=self.num_atoms.device), + self.num_atoms, + ) + + +def build_periodic_graph( + fractional_coordinates: torch.Tensor, + cell: torch.Tensor, + num_atoms: torch.Tensor, + *, + cutoff: float = 7.0, + max_neighbors: int = 50, + max_cell_images_per_dim: int = 5, +) -> MatterGenGraph: + """Construct the pinned-source periodic GemNet-T graph without PyG. + + This follows MatterGen's batched OCP graph construction exactly: the + maximum periodic-image extent is shared across the batch, candidates are + ordered by target atom / source atom / image, then directed candidates are + symmetrized image-by-image before triplets are constructed. + """ + _validate_geometry_inputs(fractional_coordinates, cell, num_atoms) + if cutoff <= 0.0: + raise ValueError("cutoff must be positive.") + if max_neighbors <= 0 or max_cell_images_per_dim <= 0: + raise ValueError("max_neighbors and max_cell_images_per_dim must be positive.") + + batch = torch.repeat_interleave( + torch.arange(len(num_atoms), dtype=torch.long, device=num_atoms.device), num_atoms + ) + # Source ``frac_to_cart_coords_with_lattice`` uses row vectors: + # (N, 3) @ (N, 3, 3) -> (N, 3) Cartesian positions. + cartesian = torch.einsum("ni,nij->nj", fractional_coordinates, cell[batch]) + cell_offsets = _periodic_image_offsets(cell, cutoff, max_cell_images_per_dim) + edge_index, to_jimages, neighbors = _radius_graph_pbc( + cartesian, + cell, + num_atoms, + cell_offsets, + cutoff, + max_neighbors, + ) + if edge_index.shape[1] == 0: + raise ValueError( + "MatterGen source ordering cannot construct GemNet-T triplets for an empty " + "periodic radius graph. Increase the cutoff or use a physically valid cell." + ) + + # ``get_pbc_distances`` uses j->i edge order and a row-vector image shift. + lattice_edges = torch.repeat_interleave(cell, neighbors, dim=0) + distance_vectors = ( + cartesian[edge_index[0]] + - cartesian[edge_index[1]] + + torch.einsum("ei,eij->ej", to_jimages, lattice_edges) + ) + distances = torch.linalg.vector_norm(distance_vectors, dim=-1) + edge_direction = -distance_vectors / distances[:, None] + + edge_index, to_jimages, neighbors, distances, edge_direction = _reorder_symmetric_edges( + edge_index, to_jimages, neighbors, distances, edge_direction + ) + if edge_index.shape[1] == 0: + raise ValueError( + "MatterGen source symmetric edge reordering removed every edge; " + "the current graph cannot be scored faithfully." + ) + id_swap = _symmetric_edge_swaps(neighbors) + id3_ba, id3_ca, id3_ragged_idx = _triplets(edge_index) + + # GemNet appends cosine alignment to each edge embedding. ``batch`` is + # indexed by the source node exactly as in ``GemNetT.forward``. + edge_lattice_cosines = torch.cosine_similarity( + edge_direction[:, None, :], cell[batch[edge_index[0]]], dim=-1 + ) + return MatterGenGraph( + edge_index=edge_index, + edge_distance=distances, + edge_direction=edge_direction, + edge_lattice_cosines=edge_lattice_cosines, + id_swap=id_swap, + id3_ba=id3_ba, + id3_ca=id3_ca, + id3_ragged_idx=id3_ragged_idx, + ) + + +class MatterGenHostSampler: + """Run the official v1.0.3 MatterGen PC loop around an explicit score callback. + + It implements the released 1,000-step absorbing-mask D3PM, the + number-of-atoms-adjusted wrapped VE coordinate process, the lattice VP + process, source predictor/corrector ordering, classifier-free guidance, + source logit masking, and final structural validation. No shortened or + rescheduled path is accepted because it would not be a source-compatible + MatterGen sampler. + """ + + def __init__( + self, + score_callback: MatterGenScoreCallback, + *, + condition_names: Sequence[str] = (), + cutoff: float = 7.0, + max_neighbors: int = 50, + max_cell_images_per_dim: int = 5, + lattice_limit_density: float = MATTERGEN_LATTICE_LIMIT_DENSITY, + ) -> None: + if not callable(score_callback): + raise TypeError("score_callback must be callable.") + if len(condition_names) != len(set(condition_names)): + raise ValueError("condition_names must be unique.") + if unsupported := set(condition_names).difference(MATTERGEN_CONDITION_FAMILY): + raise ValueError( + f"Unsupported MatterGen condition names: {sorted(unsupported)!r}." + ) + if cutoff <= 0.0 or max_neighbors <= 0 or max_cell_images_per_dim <= 0: + raise ValueError("graph construction limits must be positive.") + if lattice_limit_density <= 0.0: + raise ValueError("lattice_limit_density must be positive.") + self._score_callback = score_callback + self._condition_names = tuple(condition_names) + self._cutoff = cutoff + self._max_neighbors = max_neighbors + self._max_cell_images_per_dim = max_cell_images_per_dim + self._lattice_limit_density = lattice_limit_density + + def sample( + self, + num_atoms: torch.Tensor, + *, + condition_values: Mapping[str, torch.Tensor] | None = None, + guidance_scale: float = 0.0, + seed: int | None = None, + ) -> MatterGenSampleBatch: + """Draw source-scheduled crystal samples for requested atom counts. + + ``condition_values`` contains exactly the raw exported condition inputs. + A non-``None`` ``seed`` supplies a private CPU Torch generator; omitting + it deliberately uses the source-compatible global Torch RNG behavior. + """ + _validate_num_atoms(num_atoms) + supplied_conditions = dict(condition_values or {}) + _validate_conditions(supplied_conditions, self._condition_names, len(num_atoms)) + condition_values = _complete_condition_values( + supplied_conditions, + self._condition_names, + len(num_atoms), + ) + if not math.isfinite(guidance_scale): + raise ValueError("guidance_scale must be finite.") + + generator = None + if seed is not None: + generator = torch.Generator(device="cpu") + generator.manual_seed(seed) + + with torch.no_grad(): + state = self._sample_prior(num_atoms, generator) + dt = -torch.tensor( # Source uses CPU float32 for a CPU host state. + (1.0 - _EPS_T) / (MATTERGEN_SAMPLING_STEPS - 1), + dtype=torch.float32, + device="cpu", + ) + # Source ``torch.linspace(T, eps_t, N)`` is float32 and includes + # both ends. Its D3PM conversion therefore remains coupled to N=1000. + timesteps = torch.linspace(1.0, _EPS_T, MATTERGEN_SAMPLING_STEPS, device="cpu") + final_mean = state + for timestep in timesteps: + t = torch.full((len(num_atoms),), timestep, dtype=torch.float32, device="cpu") + + # Correctors update positions and cells from the same score + # evaluation, first positions then lattice, as ``apply`` does. + score = self._guided_score( + state, + t, + condition_values, + guidance_scale, + supplied_condition_names=frozenset(supplied_conditions), + ) + corrected_pos, _ = self._wrapped_langevin( + state.fractional_coordinates, + score.coordinate_score, + t, + dt, + state.batch, + snr=0.4, + generator=generator, + ) + corrected_cell, _ = self._lattice_langevin( + state.cell, score.lattice_score, t, dt, generator=generator + ) + state = _State( + atomic_numbers=state.atomic_numbers, + fractional_coordinates=corrected_pos, + cell=corrected_cell, + num_atoms=state.num_atoms, + ) + + # Predictors recompute the score after both corrector updates. + score = self._guided_score( + state, + t, + condition_values, + guidance_scale, + supplied_condition_names=frozenset(supplied_conditions), + ) + predicted_pos, mean_pos = self._wrapped_ancestral( + state.fractional_coordinates, + score.coordinate_score, + t, + dt, + state.num_atoms, + state.batch, + generator, + ) + predicted_cell, mean_cell = self._lattice_ancestral( + state.cell, score.lattice_score, t, dt, state.num_atoms, generator + ) + predicted_atoms, mean_atoms = self._d3pm_ancestral( + state.atomic_numbers, score.atom_logits, t, state.batch, generator + ) + state = _State( + atomic_numbers=predicted_atoms, + fractional_coordinates=predicted_pos, + cell=predicted_cell, + num_atoms=state.num_atoms, + ) + final_mean = _State( + atomic_numbers=mean_atoms, + fractional_coordinates=mean_pos, + cell=mean_cell, + num_atoms=state.num_atoms, + ) + + result = MatterGenSampleBatch( + atomic_numbers=final_mean.atomic_numbers, + fractional_coordinates=final_mean.fractional_coordinates, + cell=final_mean.cell, + num_atoms=final_mean.num_atoms, + ) + # Match the source's final Structure creation with a dependency-free, + # fail-closed structural gate. A caller may then serialize ``crystals``. + result.crystals() + return result + + def _sample_prior( + self, num_atoms: torch.Tensor, generator: torch.Generator | None + ) -> _State: + batch = torch.repeat_interleave( + torch.arange(len(num_atoms), device=num_atoms.device), num_atoms + ) + atom_count_scale = num_atoms.to(torch.float32).pow(-1.0 / 3.0)[batch, None] + # LatticeVPSDE.prior_sampling: symmetric IID noise around the diagonal + # density-derived limit mean, with n^(2/3) * 0.25 elementwise variance. + # MultiCorruption sorts fields, so the source consumes cell noise before + # the wrapped VE position prior (atomic_numbers has no random draw). + cell_noise = _symmetric_noise(_randn((len(num_atoms), 3, 3), generator)) + limit_mean = _lattice_limit_mean(num_atoms, self._lattice_limit_density) + limit_var = _lattice_limit_var(num_atoms) + cell = cell_noise * limit_var.sqrt() + limit_mean + # NumAtomsVarianceAdjustedWrappedVESDE.prior_sampling: wrapped N(0, + # (5 / n^(1/3))^2) fractional coordinates. + fractional = torch.remainder( + _randn((int(num_atoms.sum()), 3), generator) * 5.0 * atom_count_scale, 1.0 + ) + return _State( + atomic_numbers=torch.full( + (int(num_atoms.sum()),), + _D3PM_MASK_CLASS + 1, + dtype=torch.long, + device="cpu", + ), + fractional_coordinates=fractional, + cell=cell, + num_atoms=num_atoms.clone(), + ) + + def _guided_score( + self, + state: _State, + timestep: torch.Tensor, + condition_values: Mapping[str, torch.Tensor], + guidance_scale: float, + *, + supplied_condition_names: frozenset[str] | None = None, + ) -> MatterGenScoreOutputs: + if supplied_condition_names is None: + supplied_condition_names = frozenset(condition_values) + conditional = { + name: torch.full( + (len(state.num_atoms),), + name not in supplied_condition_names, + dtype=torch.bool, + device="cpu", + ) + for name in self._condition_names + } + unconditional = { + name: torch.ones(len(state.num_atoms), dtype=torch.bool, device="cpu") + for name in self._condition_names + } + + def score(use_unconditional: Mapping[str, torch.Tensor]) -> MatterGenScoreOutputs: + graph = build_periodic_graph( + state.fractional_coordinates, + state.cell, + state.num_atoms, + cutoff=self._cutoff, + max_neighbors=self._max_neighbors, + max_cell_images_per_dim=self._max_cell_images_per_dim, + ) + outputs = self._score_callback( + MatterGenScoreInputs( + atomic_numbers=state.atomic_numbers, + batch=state.batch, + timestep=timestep, + graph=graph, + condition_values=condition_values, + use_unconditional=use_unconditional, + ) + ) + _validate_score_outputs(outputs, state) + # MatterGen's denoiser converts raw Cartesian GemNet forces to + # fractional scores before the wrapped coordinate process. + coordinate_score = torch.bmm( + torch.linalg.inv(state.cell).transpose(1, 2)[state.batch], + outputs.coordinate_score.unsqueeze(-1), + ).squeeze(-1) + atom_logits = _mask_atom_logits( + outputs.atom_logits, + condition_values.get("chemical_system"), + use_unconditional.get("chemical_system"), + state.batch, + ) + return MatterGenScoreOutputs( + atom_logits=atom_logits, + coordinate_score=coordinate_score, + lattice_score=outputs.lattice_score, + energy=outputs.energy, + ) + + # ``GuidedPredictorCorrector`` avoids unnecessary model calls at 0 and + # 1, otherwise computing unconditional + gamma*(conditional-unconditional). + if abs(guidance_scale - 1.0) < 1e-15: + return score(conditional) + if abs(guidance_scale) < 1e-15: + return score(unconditional) + conditional_score = score(conditional) + unconditional_score = score(unconditional) + return MatterGenScoreOutputs( + atom_logits=torch.lerp( + unconditional_score.atom_logits, conditional_score.atom_logits, guidance_scale + ), + coordinate_score=torch.lerp( + unconditional_score.coordinate_score, + conditional_score.coordinate_score, + guidance_scale, + ), + lattice_score=torch.lerp( + unconditional_score.lattice_score, + conditional_score.lattice_score, + guidance_scale, + ), + energy=None, + ) + + def _wrapped_langevin( + self, + value: torch.Tensor, + score: torch.Tensor, + timestep: torch.Tensor, + dt: torch.Tensor, + batch: torch.Tensor, + *, + snr: float, + generator: torch.Generator | None, + ) -> tuple[torch.Tensor, torch.Tensor]: + del dt # VE's source Langevin alpha is identically one. + noise = _randn_like(score, generator) + grad_norm = _per_crystal_norm(score, batch, len(timestep)).mean() + noise_norm = _per_crystal_norm(noise, batch, len(timestep)).mean() + step_size = _langevin_step_size(snr, noise_norm, grad_norm, len(timestep)) + expanded_step = step_size[batch, None] + mean = value + expanded_step * score + sample = mean + torch.sqrt(expanded_step * 2.0) * noise + return torch.remainder(sample, 1.0), torch.remainder(mean, 1.0) + + def _lattice_langevin( + self, + value: torch.Tensor, + score: torch.Tensor, + timestep: torch.Tensor, + dt: torch.Tensor, + *, + generator: torch.Generator | None, + ) -> tuple[torch.Tensor, torch.Tensor]: + alpha = _vp_alpha(timestep) ** 2 / _vp_alpha(timestep + dt) ** 2 + noise = _symmetric_noise(_randn_like(score, generator)) + grad_norm = torch.square(score).sum(dim=(1, 2)).sqrt().mean() + noise_norm = torch.square(noise).sum(dim=(1, 2)).sqrt().mean() + step_size = _langevin_step_size( + 0.2, + noise_norm, + grad_norm, + len(timestep), + alpha=alpha, + ) + expanded_step = step_size[:, None, None] + mean = value + expanded_step * score + sample = mean + torch.sqrt(expanded_step * 2.0) * noise + return _polar_lattice(sample), _polar_lattice(mean) + + def _wrapped_ancestral( + self, + value: torch.Tensor, + score: torch.Tensor, + timestep: torch.Tensor, + dt: torch.Tensor, + num_atoms: torch.Tensor, + batch: torch.Tensor, + generator: torch.Generator | None, + ) -> tuple[torch.Tensor, torch.Tensor]: + sigma_t = _position_sigma(timestep, batch, num_atoms) + sigma_s = _position_sigma(timestep + dt, batch, num_atoms) + is_time_zero = (timestep + dt)[batch] <= 0 + sigma_s[is_time_zero] = 0.0 + score_coeff = sigma_t.square() - sigma_s.square() + std = torch.sqrt(score_coeff) * sigma_s / sigma_t + mean = value + score_coeff * score + sample = mean + std * _randn_like(value, generator) + return torch.remainder(sample, 1.0), torch.remainder(mean, 1.0) + + def _lattice_ancestral( + self, + value: torch.Tensor, + score: torch.Tensor, + timestep: torch.Tensor, + dt: torch.Tensor, + num_atoms: torch.Tensor, + generator: torch.Generator | None, + ) -> tuple[torch.Tensor, torch.Tensor]: + alpha_t = _vp_alpha(timestep)[:, None, None] + alpha_s = _vp_alpha(timestep + dt)[:, None, None] + limit_var = _lattice_limit_var(num_atoms) + sigma_t = torch.sqrt((1.0 - alpha_t.square()) * limit_var) + sigma_s = torch.sqrt((1.0 - alpha_s.square()) * limit_var) + is_time_zero = (timestep + dt) <= 0 + sigma_s[is_time_zero] = 0.0 + alpha_t_given_s = torch.clamp(alpha_t / alpha_s, min=0.001, max=1.0) + sigma2_t_given_s = ( + sigma_t.square() - sigma_s.square() * alpha_t.square() / alpha_s.square() + ) + std = torch.sqrt(sigma2_t_given_s) * sigma_s / sigma_t + std[is_time_zero] = 0.0 + x_coeff = 1.0 / alpha_t_given_s + score_coeff = sigma2_t_given_s / alpha_t_given_s + limit_mean = _lattice_limit_mean(num_atoms, self._lattice_limit_density) + mean = x_coeff * value + score_coeff * score + (1.0 - x_coeff) * limit_mean + # Source samples ``randn_like(x_coeff)`` where x_coeff is [B, 1, 1], + # then broadcasts it through the 3x3 symmetric-noise transform. + sample = mean + std * _symmetric_noise(_randn_like(x_coeff, generator)) + return sample, mean + + def _d3pm_ancestral( + self, + atomic_numbers: torch.Tensor, + logits: torch.Tensor, + timestep: torch.Tensor, + batch: torch.Tensor, + generator: torch.Generator | None, + ) -> tuple[torch.Tensor, torch.Tensor]: + discrete_time = (timestep * (MATTERGEN_SAMPLING_STEPS - 1)).long()[batch] + # The source computes this preliminary sample before the predict-x0 + # posterior. Retaining it preserves the source RNG consumption order. + _categorical(logits, generator) + class_probs = torch.softmax(logits, dim=-1) + state = _mask_diffusion_state(discrete_time) + q_t = torch.empty_like(class_probs) + q_t[:, :-1] = state[:, None] * class_probs[:, :-1] + q_t[:, -1] = 1.0 - q_t[:, :-1].sum(dim=-1) + + beta = 1.0 / (MATTERGEN_SAMPLING_STEPS - discrete_time).to(torch.float32) + current = atomic_numbers - 1 + transition = torch.zeros_like(class_probs) + is_mask = current == _D3PM_MASK_CLASS + transition[is_mask, :-1] = beta[is_mask, None] + transition[is_mask, -1] = 1.0 + non_mask_rows = torch.nonzero(~is_mask, as_tuple=False).flatten() + transition[non_mask_rows, current[non_mask_rows]] = 1.0 - beta[non_mask_rows] + posterior_logits = torch.log(q_t + _D3PM_EPSILON) + torch.log( + transition + _D3PM_EPSILON + ) + sample = _categorical(posterior_logits, generator) + 1 + mean = torch.argmax(torch.softmax(posterior_logits, dim=-1), dim=-1) + 1 + return sample, mean + + +def _as_numpy(value: torch.Tensor) -> np.ndarray: + if value.device.type != "cpu": + raise ValueError("MatterGen ONNX Runtime callback inputs must be CPU tensors.") + return value.detach().contiguous().numpy() + + +def _validate_num_atoms(num_atoms: torch.Tensor) -> None: + if not isinstance(num_atoms, torch.Tensor) or num_atoms.dtype != torch.long: + raise TypeError("num_atoms must be a CPU torch.int64 tensor.") + if num_atoms.ndim != 1 or len(num_atoms) == 0 or num_atoms.device.type != "cpu": + raise ValueError("num_atoms must be a non-empty rank-1 CPU tensor.") + if torch.any(num_atoms < 1) or torch.any(num_atoms > MAX_ATOMS): + raise ValueError(f"Each MatterGen crystal must contain 1 through {MAX_ATOMS} atoms.") + + +def _validate_geometry_inputs( + fractional_coordinates: torch.Tensor, cell: torch.Tensor, num_atoms: torch.Tensor +) -> None: + _validate_num_atoms(num_atoms) + if ( + not isinstance(fractional_coordinates, torch.Tensor) + or fractional_coordinates.dtype != torch.float32 + or fractional_coordinates.device.type != "cpu" + or fractional_coordinates.shape != (int(num_atoms.sum()), 3) + ): + raise ValueError( + "fractional_coordinates must be a CPU float32 tensor with shape [N, 3]." + ) + if ( + not isinstance(cell, torch.Tensor) + or cell.dtype != torch.float32 + or cell.device.type != "cpu" + or cell.shape != (len(num_atoms), 3, 3) + ): + raise ValueError("cell must be a CPU float32 tensor with shape [B, 3, 3].") + if not torch.isfinite(fractional_coordinates).all() or not torch.isfinite(cell).all(): + raise ValueError("fractional_coordinates and cell must be finite.") + if torch.any(fractional_coordinates < 0.0) or torch.any(fractional_coordinates >= 1.0): + raise ValueError("fractional_coordinates must be wrapped to [0, 1).") + if torch.any(torch.linalg.det(cell) == 0): + raise ValueError("cell must have nonzero volume while constructing a periodic graph.") + + +def _validate_conditions( + values: Mapping[str, torch.Tensor], names: Sequence[str], batch_size: int +) -> None: + unexpected = set(values).difference(names) + if unexpected: + raise ValueError( + f"condition_values contains unsupported names {sorted(unexpected)!r}; " + f"expected a subset of {tuple(names)!r}." + ) + for name, value in values.items(): + if not isinstance(value, torch.Tensor) or value.device.type != "cpu": + raise TypeError(f"condition {name!r} must be a CPU torch tensor.") + if value.ndim == 0 or value.shape[0] != batch_size: + raise ValueError(f"condition {name!r} must have batch dimension {batch_size}.") + if name == "chemical_system": + if value.dtype != torch.float32 or value.shape != (batch_size, _D3PM_CLASSES): + raise ValueError( + "chemical_system must have shape [B, 101] and dtype torch.float32." + ) + if ( + not torch.all(value.eq(0.0) | value.eq(1.0)) + or torch.any(value[:, 0].ne(0.0)) + or torch.any(value.sum(dim=1).eq(0.0)) + ): + raise ValueError( + "chemical_system must be a non-empty one-based binary multihot." + ) + allowed = torch.zeros(_D3PM_CLASSES, dtype=torch.bool, device="cpu") + allowed[torch.tensor(SELECTED_ATOMIC_NUMBERS, dtype=torch.long, device="cpu")] = ( + True + ) + if torch.any(value[:, ~allowed].ne(0.0)): + raise ValueError( + "chemical_system contains an element outside MatterGen's allowlist." + ) + elif name == "space_group": + if value.dtype != torch.long or value.shape != (batch_size,): + raise ValueError("space_group must have shape [B] and dtype torch.int64.") + if torch.any(value < 1) or torch.any(value > 230): + raise ValueError("space_group must be in the inclusive range [1, 230].") + elif ( + value.dtype != torch.float32 + or value.shape != (batch_size,) + or not torch.isfinite(value).all() + ): + raise ValueError( + f"Scalar condition {name!r} must be a finite float32 tensor with shape [B]." + ) + elif name in {"dft_bulk_modulus", "ml_bulk_modulus"} and torch.any(value <= 0.0): + raise ValueError( + f"Scalar condition {name!r} must be strictly positive for log10 scaling." + ) + + +def _complete_condition_values( + values: Mapping[str, torch.Tensor], names: Sequence[str], batch_size: int +) -> dict[str, torch.Tensor]: + """Fill absent source conditions with shape-valid values for unconditional ONNX paths.""" + completed = dict(values) + for name in names: + if name in completed: + continue + if name == "chemical_system": + # The selector keeps this placeholder out of both property and + # species-mask semantics; hydrogen simply makes it source-shaped. + placeholder = torch.zeros((batch_size, _D3PM_CLASSES), dtype=torch.float32) + placeholder[:, 1] = 1.0 + elif name == "space_group": + placeholder = torch.ones(batch_size, dtype=torch.long) + else: + # Positive placeholders also avoid evaluating log10(0) in the + # graph's unused conditional branch for bulk-modulus adapters. + placeholder = torch.ones(batch_size, dtype=torch.float32) + completed[name] = placeholder + return completed + + +def _validate_score_outputs(outputs: MatterGenScoreOutputs, state: _State) -> None: + if not isinstance(outputs, MatterGenScoreOutputs): + raise TypeError("score_callback must return MatterGenScoreOutputs.") + expected = { + "atom_logits": (int(state.num_atoms.sum()), _D3PM_CLASSES), + "coordinate_score": (int(state.num_atoms.sum()), 3), + "lattice_score": (len(state.num_atoms), 3, 3), + } + for name, shape in expected.items(): + value = getattr(outputs, name) + if not isinstance(value, torch.Tensor) or value.dtype != torch.float32: + raise TypeError(f"score_callback {name} must be a torch.float32 tensor.") + if value.device.type != "cpu" or tuple(value.shape) != shape: + raise ValueError(f"score_callback {name} must have CPU shape {shape}.") + if not torch.isfinite(value).all(): + raise ValueError(f"score_callback {name} must be finite.") + + +def _periodic_image_offsets( + cell: torch.Tensor, cutoff: float, max_cell_images_per_dim: int +) -> torch.Tensor: + cross_a2a3 = torch.cross(cell[:, 1], cell[:, 2], dim=-1) + volume = torch.sum(cell[:, 0] * cross_a2a3, dim=-1, keepdim=True) + reciprocal_norms = ( + torch.linalg.vector_norm(cross_a2a3 / volume, dim=-1), + torch.linalg.vector_norm(torch.cross(cell[:, 2], cell[:, 0], dim=-1) / volume, dim=-1), + torch.linalg.vector_norm(torch.cross(cell[:, 0], cell[:, 1], dim=-1) / volume, dim=-1), + ) + repetitions = [ + min(int(torch.ceil(cutoff * reciprocal_norm).max()), max_cell_images_per_dim) + for reciprocal_norm in reciprocal_norms + ] + return torch.cartesian_prod( + *[ + torch.arange(-repetition, repetition + 1, dtype=torch.float32, device="cpu") + for repetition in repetitions + ] + ) + + +def _radius_graph_pbc( + cartesian: torch.Tensor, + cell: torch.Tensor, + num_atoms: torch.Tensor, + image_offsets: torch.Tensor, + cutoff: float, + max_neighbors: int, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Port ``ocp_graph_utils.radius_graph_pbc`` through neighbor truncation.""" + index1_parts: list[torch.Tensor] = [] + index2_parts: list[torch.Tensor] = [] + for offset, count in zip( + torch.cat([num_atoms.new_zeros(1), num_atoms.cumsum(0)[:-1]]), num_atoms + ): + local = torch.arange(int(count), dtype=torch.long, device="cpu") + offset + # Source creates pairs with index1 as target and index2 as source. + index1_parts.append(local.repeat_interleave(int(count))) + index2_parts.append(local.repeat(int(count))) + index1 = torch.cat(index1_parts).repeat_interleave(len(image_offsets)) + index2 = torch.cat(index2_parts).repeat_interleave(len(image_offsets)) + offsets = image_offsets.repeat(int(num_atoms.square().sum()), 1) + pair_cells = torch.repeat_interleave( + torch.repeat_interleave(cell, num_atoms.square(), dim=0), + len(image_offsets), + dim=0, + ) + # The OCP implementation applies the image displacement to index2 before + # measuring index1 - index2; the later GemNet V_st convention negates it. + shifted_source = cartesian[index2] + torch.einsum("ei,eij->ej", offsets, pair_cells) + distance_squared = torch.square(cartesian[index1] - shifted_source).sum(dim=-1) + mask = (distance_squared <= cutoff * cutoff) & (distance_squared > 0.0001) + index1 = index1[mask] + index2 = index2[mask] + offsets = offsets[mask] + distance_squared = distance_squared[mask] + + neighbor_mask, neighbors = _max_neighbors_mask( + num_atoms, index1, distance_squared, max_neighbors + ) + index1 = index1[neighbor_mask] + index2 = index2[neighbor_mask] + offsets = offsets[neighbor_mask] + return torch.stack((index2, index1)), offsets, neighbors + + +def _max_neighbors_mask( + num_atoms: torch.Tensor, + target: torch.Tensor, + distance_squared: torch.Tensor, + max_neighbors: int, +) -> tuple[torch.Tensor, torch.Tensor]: + num_total_atoms = int(num_atoms.sum()) + counts = torch.bincount(target, minlength=num_total_atoms) + thresholded = counts.clamp(max=max_neighbors) + atom_batch = torch.repeat_interleave(torch.arange(len(num_atoms), device="cpu"), num_atoms) + neighbors = torch.zeros(len(num_atoms), dtype=torch.long, device="cpu") + neighbors.scatter_add_(0, atom_batch, thresholded) + if target.numel() == 0 or int(counts.max()) <= max_neighbors: + return torch.ones_like(target, dtype=torch.bool), neighbors + + # ``get_max_neighbors_mask`` writes the target-sorted candidate distances + # into a dense matrix, sorts each target row, then retains its original + # candidate order through a Boolean mask. + max_count = int(counts.max()) + starts = torch.cumsum(counts, dim=0) - counts + dense = torch.full( + (num_total_atoms, max_count), float("inf"), dtype=torch.float32, device="cpu" + ) + columns = torch.arange(len(target), device="cpu") - torch.repeat_interleave(starts, counts) + dense[target, columns] = distance_squared + _, sorted_columns = torch.sort(dense, dim=1) + retained = sorted_columns[:, :max_neighbors] + starts[:, None] + selected = retained[ + torch.isfinite(torch.gather(dense, 1, sorted_columns[:, :max_neighbors])) + ] + mask = torch.zeros(len(target), dtype=torch.bool, device="cpu") + mask[selected] = True + return mask, neighbors + + +def _reorder_symmetric_edges( + edge_index: torch.Tensor, + cell_offsets: torch.Tensor, + neighbors: torch.Tensor, + distances: torch.Tensor, + edge_direction: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + source, target = edge_index + earlier_cell = ( + (cell_offsets[:, 0] < 0) + | ((cell_offsets[:, 0] == 0) & (cell_offsets[:, 1] < 0)) + | ((cell_offsets[:, 0] == 0) & (cell_offsets[:, 1] == 0) & (cell_offsets[:, 2] < 0)) + ) + mask = (source < target) | ((source == target) & earlier_cell) + directed_edge_index = edge_index[:, mask] + directed_offsets = cell_offsets[mask] + directed_distances = distances[mask] + directed_directions = edge_direction[mask] + + edge_batch = torch.repeat_interleave( + torch.arange(len(neighbors), device="cpu"), neighbors + )[mask] + symmetric_neighbors = 2 * torch.bincount(edge_batch, minlength=len(neighbors)) + count_per_image = symmetric_neighbors // 2 + directed_total = len(directed_offsets) + directed_starts = torch.cumsum(count_per_image, dim=0) - count_per_image + reorder_parts = [ + torch.cat( + [ + torch.arange(start, start + count, device="cpu"), + torch.arange( + directed_total + start, directed_total + start + count, device="cpu" + ), + ] + ) + for start, count in zip( + directed_starts.tolist(), count_per_image.tolist(), strict=True + ) + if count > 0 + ] + reorder = ( + torch.cat(reorder_parts) + if reorder_parts + else torch.empty(0, dtype=torch.long, device="cpu") + ) + edge_index_cat = torch.cat( + [directed_edge_index, torch.stack([directed_edge_index[1], directed_edge_index[0]])], + dim=1, + ) + return ( + edge_index_cat[:, reorder], + torch.cat([directed_offsets, -directed_offsets], dim=0)[reorder], + symmetric_neighbors, + torch.cat([directed_distances, directed_distances], dim=0)[reorder], + torch.cat([directed_directions, -directed_directions], dim=0)[reorder], + ) + + +def _symmetric_edge_swaps(neighbors: torch.Tensor) -> torch.Tensor: + swaps: list[torch.Tensor] = [] + start = 0 + for count in neighbors.tolist(): + half = count // 2 + if half: + swaps.append( + torch.cat( + [ + torch.arange(start + half, start + count, device="cpu"), + torch.arange(start, start + half, device="cpu"), + ] + ) + ) + start += count + return torch.cat(swaps) if swaps else torch.empty(0, dtype=torch.long, device="cpu") + + +def _triplets(edge_index: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Port ``SparseTensor(row=target, col=source)[target]`` deterministically.""" + source, target = edge_index + triplet_ba: list[torch.Tensor] = [] + triplet_ca: list[torch.Tensor] = [] + ragged: list[torch.Tensor] = [] + for ca in range(edge_index.shape[1]): + candidates = torch.nonzero(target == target[ca], as_tuple=False).flatten() + # torch_sparse stores CSR rows by source column; retain edge-id order + # for periodic duplicate columns, which is the source insertion order. + candidate_order = torch.argsort(source[candidates], stable=True) + candidates = candidates[candidate_order] + candidates = candidates[candidates != ca] + triplet_ba.append(candidates) + triplet_ca.append(torch.full((len(candidates),), ca, dtype=torch.long, device="cpu")) + ragged.append(torch.arange(len(candidates), dtype=torch.long, device="cpu")) + if not triplet_ba: + empty = torch.empty(0, dtype=torch.long, device="cpu") + return empty, empty, empty + return torch.cat(triplet_ba), torch.cat(triplet_ca), torch.cat(ragged) + + +def _mask_atom_logits( + logits: torch.Tensor, + chemical_system: torch.Tensor | None, + use_unconditional: torch.Tensor | None, + batch: torch.Tensor, +) -> torch.Tensor: + # MatterGen's ``mask_disallowed_elements`` treats score logits as zero-based + # atomic numbers and reserves the final 101st class for the absorbing mask. + selected = torch.tensor(SELECTED_ATOMIC_NUMBERS, dtype=torch.long, device="cpu") + keep = torch.zeros((1, _D3PM_CLASSES), dtype=torch.float32, device="cpu") + keep[0, selected - 1] = 1.0 + masked = logits + (1.0 - keep) * -1e10 + if chemical_system is None: + return masked + if use_unconditional is None: + raise ValueError("chemical_system requires its use_unconditional selector.") + chemical_keep = torch.zeros( + (len(chemical_system), _D3PM_CLASSES), dtype=torch.float32, device="cpu" + ) + chemical_keep[:, :-1] = chemical_system[:, 1:] + keep = torch.where( + use_unconditional[:, None], + torch.ones((len(chemical_system), 1), dtype=torch.float32, device="cpu"), + chemical_keep, + ) + return masked + (1.0 - keep[batch]) * -1e10 + + +def _randn(shape: tuple[int, ...], generator: torch.Generator | None) -> torch.Tensor: + return torch.randn(shape, dtype=torch.float32, device="cpu", generator=generator) + + +def _randn_like(value: torch.Tensor, generator: torch.Generator | None) -> torch.Tensor: + return torch.randn(value.shape, dtype=value.dtype, device="cpu", generator=generator) + + +def _categorical(logits: torch.Tensor, generator: torch.Generator | None) -> torch.Tensor: + if generator is None: + return torch.distributions.Categorical(logits=logits).sample() + return torch.multinomial(torch.softmax(logits, dim=-1), 1, generator=generator).squeeze(-1) + + +def _mask_diffusion_state(timestep: torch.Tensor) -> torch.Tensor: + """Index the source-materialized MaskDiffusion cumulative state.""" + return _D3PM_STATE[timestep] + + +def _position_sigma( + timestep: torch.Tensor, batch: torch.Tensor, num_atoms: torch.Tensor +) -> torch.Tensor: + sigma = 0.01 * (5.0 / 0.01) ** timestep + return (sigma * num_atoms.to(torch.float32).pow(-1.0 / 3.0))[batch, None] + + +def _vp_alpha(timestep: torch.Tensor) -> torch.Tensor: + return torch.exp(-0.25 * timestep.square() * (20.0 - 0.1) - 0.5 * timestep * 0.1) + + +def _lattice_limit_mean(num_atoms: torch.Tensor, density: float) -> torch.Tensor: + return torch.pow( + torch.eye(3, device="cpu").expand(len(num_atoms), 3, 3) + * num_atoms.to(torch.float32)[:, None, None] + / density, + 1.0 / 3.0, + ) + + +def _lattice_limit_var(num_atoms: torch.Tensor) -> torch.Tensor: + return num_atoms.to(torch.float32)[:, None, None].expand(-1, 3, 3).pow(2.0 / 3.0) * 0.25 + + +def _symmetric_noise(noise: torch.Tensor) -> torch.Tensor: + eye = torch.eye(3, device="cpu")[None] + return (1.0 / math.sqrt(2.0)) * (1.0 - eye) * (noise + noise.transpose(1, 2)) + eye * noise + + +def _polar_lattice(lattice: torch.Tensor) -> torch.Tensor: + # ``compute_lattice_polar_decomposition`` projects corrector updates to the + # rotation-equivalent symmetric lattice representation used by MatterGen. + w, singular_values, v_transpose = torch.linalg.svd(lattice) + v = v_transpose.transpose(1, 2) + orthogonal = w @ v_transpose + positive = v @ torch.diag_embed(singular_values) @ v_transpose + return orthogonal @ positive @ orthogonal.transpose(1, 2) + + +def _per_crystal_norm( + score: torch.Tensor, batch: torch.Tensor, batch_size: int +) -> torch.Tensor: + norms = torch.square(score).sum(dim=1) + summed = torch.zeros(batch_size, dtype=score.dtype, device="cpu") + summed.scatter_add_(0, batch, norms) + return torch.sqrt(summed) + + +def _langevin_step_size( + snr: float, + noise_norm: torch.Tensor, + grad_norm: torch.Tensor, + batch_size: int, + *, + alpha: torch.Tensor | None = None, +) -> torch.Tensor: + # Both released correctors explicitly take their configured cap when the + # aggregate score is zero; this is not an error or a substitute schedule. + if not bool(grad_norm): + return torch.full( + (batch_size,), + _LANGEVIN_MAX_STEP_SIZE, + dtype=torch.float32, + device="cpu", + ) + step_size = torch.full( + (batch_size,), + float((snr * noise_norm / grad_norm) ** 2 * 2.0), + dtype=torch.float32, + device="cpu", + ) + if alpha is not None: + step_size *= alpha + return torch.clamp(step_size, max=_LANGEVIN_MAX_STEP_SIZE) diff --git a/src/mobius/integrations/mattergen/_runtime_test.py b/src/mobius/integrations/mattergen/_runtime_test.py new file mode 100644 index 000000000..3c91f3a5f --- /dev/null +++ b/src/mobius/integrations/mattergen/_runtime_test.py @@ -0,0 +1,325 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +from __future__ import annotations + +import numpy as np +import pytest +import torch + +from mobius.integrations.mattergen import ( + MatterGenHostSampler, + MatterGenScoreInputs, + MatterGenScoreOutputs, + build_periodic_graph, + create_onnxruntime_score_callback, +) +from mobius.integrations.mattergen._runtime import ( + _LANGEVIN_MAX_STEP_SIZE, + MATTERGEN_SAMPLING_STEPS, + _langevin_step_size, + _State, +) + + +class TestMatterGenPeriodicGraph: + def test_matches_source_self_image_edge_and_triplet_order(self) -> None: + """Exercise the source's PBC/reorder/triplet sequence without PyG.""" + graph = build_periodic_graph( + torch.zeros((1, 3), dtype=torch.float32), + torch.eye(3, dtype=torch.float32).unsqueeze(0) * 4.0, + torch.tensor([1], dtype=torch.long), + cutoff=4.1, + ) + + torch.testing.assert_close( + graph.edge_index, + torch.zeros((2, 6), dtype=torch.long), + ) + torch.testing.assert_close(graph.edge_distance, torch.full((6,), 4.0)) + torch.testing.assert_close( + graph.edge_direction, + torch.tensor( + [ + [1.0, 0.0, 0.0], + [0.0, 1.0, 0.0], + [0.0, 0.0, 1.0], + [-1.0, 0.0, 0.0], + [0.0, -1.0, 0.0], + [0.0, 0.0, -1.0], + ] + ), + ) + torch.testing.assert_close(graph.edge_lattice_cosines, graph.edge_direction) + torch.testing.assert_close( + graph.id_swap, + torch.tensor([3, 4, 5, 0, 1, 2]), + ) + torch.testing.assert_close( + graph.id3_ba, + torch.tensor( + [ + 1, + 2, + 3, + 4, + 5, + 0, + 2, + 3, + 4, + 5, + 0, + 1, + 3, + 4, + 5, + 0, + 1, + 2, + 4, + 5, + 0, + 1, + 2, + 3, + 5, + 0, + 1, + 2, + 3, + 4, + ] + ), + ) + torch.testing.assert_close( + graph.id3_ca, + torch.arange(6, dtype=torch.long).repeat_interleave(5), + ) + torch.testing.assert_close( + graph.id3_ragged_idx, + torch.arange(5, dtype=torch.long).repeat(6), + ) + + def test_rejects_an_empty_source_graph(self) -> None: + with pytest.raises(ValueError, match="empty periodic radius graph"): + build_periodic_graph( + torch.zeros((1, 3), dtype=torch.float32), + torch.eye(3, dtype=torch.float32).unsqueeze(0) * 20.0, + torch.tensor([1], dtype=torch.long), + ) + + +class TestMatterGenScoreCallback: + def test_onnxruntime_adapter_uses_the_exported_named_abi(self) -> None: + graph = build_periodic_graph( + torch.zeros((1, 3), dtype=torch.float32), + torch.eye(3, dtype=torch.float32).unsqueeze(0) * 4.0, + torch.tensor([1], dtype=torch.long), + cutoff=4.1, + ) + captured: dict[str, np.ndarray] = {} + + class Session: + def run(self, output_names, input_feed): + assert output_names == [ + "atom_logits", + "coordinate_score", + "lattice_score", + "energy", + ] + captured.update(input_feed) + return [ + np.zeros((1, 101), dtype=np.float32), + np.zeros((1, 3), dtype=np.float32), + np.zeros((1, 3, 3), dtype=np.float32), + np.zeros((1, 1), dtype=np.float32), + ] + + callback = create_onnxruntime_score_callback(Session()) + outputs = callback( + MatterGenScoreInputs( + atomic_numbers=torch.tensor([101], dtype=torch.long), + batch=torch.tensor([0], dtype=torch.long), + timestep=torch.tensor([1.0], dtype=torch.float32), + graph=graph, + condition_values={ + "chemical_system": torch.zeros((1, 101), dtype=torch.float32), + }, + use_unconditional={ + "chemical_system": torch.tensor([True], dtype=torch.bool), + }, + ) + ) + + assert set(captured) == { + "atomic_numbers", + "batch", + "timestep", + "edge_index", + "edge_distance", + "edge_direction", + "edge_lattice_cosines", + "id_swap", + "id3_ba", + "id3_ca", + "id3_ragged_idx", + "condition.chemical_system", + "condition.chemical_system.use_unconditional", + } + assert outputs.atom_logits.shape == (1, 101) + + def test_cfg_uses_source_conditional_then_unconditional_order(self) -> None: + calls: list[MatterGenScoreInputs] = [] + + def score(inputs: MatterGenScoreInputs) -> MatterGenScoreOutputs: + calls.append(inputs) + conditional = not bool(inputs.use_unconditional["chemical_system"][0]) + return MatterGenScoreOutputs( + atom_logits=torch.full( + (1, 101), 10.0 if conditional else 2.0, dtype=torch.float32 + ), + coordinate_score=torch.full( + (1, 3), 4.0 if conditional else 2.0, dtype=torch.float32 + ), + lattice_score=torch.full( + (1, 3, 3), 4.0 if conditional else 2.0, dtype=torch.float32 + ), + ) + + sampler = MatterGenHostSampler(score, condition_names=("chemical_system",)) + chemical_system = torch.zeros((1, 101), dtype=torch.float32) + chemical_system[0, 3] = 1.0 + state = _State( + atomic_numbers=torch.tensor([101], dtype=torch.long), + fractional_coordinates=torch.zeros((1, 3), dtype=torch.float32), + cell=torch.eye(3, dtype=torch.float32).unsqueeze(0) * 4.0, + num_atoms=torch.tensor([1], dtype=torch.long), + ) + + guided = sampler._guided_score( + state, + torch.tensor([1.0], dtype=torch.float32), + {"chemical_system": chemical_system}, + guidance_scale=0.5, + ) + + assert [bool(call.use_unconditional["chemical_system"][0]) for call in calls] == [ + False, + True, + ] + torch.testing.assert_close(guided.atom_logits[:, 2], torch.tensor([6.0])) + torch.testing.assert_close(guided.coordinate_score, torch.full((1, 3), 0.75)) + torch.testing.assert_close(guided.lattice_score, torch.full((1, 3, 3), 3.0)) + + def test_omitted_adapter_values_use_source_unconditional_semantics(self) -> None: + """Conditioned checkpoints may sample unconditionally without a property value.""" + captured: list[MatterGenScoreInputs] = [] + + def score(inputs: MatterGenScoreInputs) -> MatterGenScoreOutputs: + captured.append(inputs) + raise RuntimeError("stop after inspecting the first source score call") + + with pytest.raises(RuntimeError, match="stop after"): + MatterGenHostSampler(score, condition_names=("ml_bulk_modulus",)).sample( + torch.tensor([1], dtype=torch.long), + seed=3, + ) + + assert len(captured) == 1 + assert bool(captured[0].use_unconditional["ml_bulk_modulus"][0]) + torch.testing.assert_close( + captured[0].condition_values["ml_bulk_modulus"], + torch.ones(1, dtype=torch.float32), + ) + + +class TestMatterGenHostSampler: + def test_zero_langevin_score_uses_the_released_cap(self) -> None: + """The source correctors define, rather than reject, a zero-score update.""" + step_size = _langevin_step_size( + snr=0.4, + noise_norm=torch.tensor(3.0), + grad_norm=torch.tensor(0.0), + batch_size=2, + ) + + torch.testing.assert_close( + step_size, + torch.full((2,), _LANGEVIN_MAX_STEP_SIZE, dtype=torch.float32), + ) + + def test_lattice_langevin_caps_after_vp_alpha_scaling(self) -> None: + """Mirror LatticeLangevinDiffCorrector's cap ordering.""" + step_size = _langevin_step_size( + snr=0.2, + noise_norm=torch.tensor(100.0), + grad_norm=torch.tensor(0.001), + batch_size=2, + alpha=torch.tensor([0.0001, 1.0], dtype=torch.float32), + ) + + torch.testing.assert_close( + step_size, + torch.tensor([80_000.0, _LANGEVIN_MAX_STEP_SIZE], dtype=torch.float32), + ) + + def test_rejects_invalid_source_condition_before_sampling(self) -> None: + def score(_: MatterGenScoreInputs) -> MatterGenScoreOutputs: + raise AssertionError("invalid conditioning must fail before a score invocation") + + with pytest.raises(ValueError, match="strictly positive"): + MatterGenHostSampler(score, condition_names=("ml_bulk_modulus",)).sample( + torch.tensor([1], dtype=torch.long), + condition_values={ + "ml_bulk_modulus": torch.tensor([0.0], dtype=torch.float32), + }, + ) + + def test_seeded_full_source_schedule_is_deterministic_and_validates_output(self) -> None: + """Exercise all 1,000 host timesteps and the final structural gate.""" + calls = 0 + skew = torch.tensor( + [[0.0, 100.0, 0.0], [-100.0, 0.0, 0.0], [0.0, 0.0, 0.0]], + dtype=torch.float32, + ) + + def score(inputs: MatterGenScoreInputs) -> MatterGenScoreOutputs: + nonlocal calls + calls += 1 + atom_logits = torch.zeros((len(inputs.atomic_numbers), 101), dtype=torch.float32) + atom_logits[:, 0] = 2.0 + # A high-norm anti-symmetric fixture keeps test-only Langevin + # steps bounded without modeling a physical score field. + lattice_score = skew.unsqueeze(0).repeat(len(inputs.timestep), 1, 1) + return MatterGenScoreOutputs( + atom_logits=atom_logits, + coordinate_score=torch.full( + (len(inputs.atomic_numbers), 3), 5.0, dtype=torch.float32 + ), + lattice_score=lattice_score, + ) + + sampler = MatterGenHostSampler(score, cutoff=100.0, max_neighbors=2) + sample = sampler.sample( + torch.tensor([1], dtype=torch.long), + seed=19, + ) + repeated = sampler.sample( + torch.tensor([1], dtype=torch.long), + seed=19, + ) + + assert calls == 4 * MATTERGEN_SAMPLING_STEPS + torch.testing.assert_close(sample.atomic_numbers, repeated.atomic_numbers) + torch.testing.assert_close( + sample.fractional_coordinates, repeated.fractional_coordinates + ) + torch.testing.assert_close(sample.cell, repeated.cell) + assert sample.atomic_numbers.shape == (1,) + assert sample.fractional_coordinates.shape == (1, 3) + assert torch.all( + (sample.fractional_coordinates >= 0.0) & (sample.fractional_coordinates < 1.0) + ) + assert torch.linalg.det(sample.cell).item() > 0.0 + assert len(sample.crystals()) == 1 diff --git a/src/mobius/integrations/mattergen/_weights.py b/src/mobius/integrations/mattergen/_weights.py new file mode 100644 index 000000000..419d0d043 --- /dev/null +++ b/src/mobius/integrations/mattergen/_weights.py @@ -0,0 +1,148 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +"""Fail-closed Lightning checkpoint loading for MatterGen score-core exports.""" + +from __future__ import annotations + +__all__ = [ + "MATTERGEN_MODEL_STATE_PREFIX", + "apply_mattergen_checkpoint", + "load_mattergen_state_dict", +] + +import json +from collections.abc import Mapping +from pathlib import Path + +import torch + +from mobius._model_package import ModelPackage + +MATTERGEN_MODEL_STATE_PREFIX = "diffusion_module.model." + + +def load_mattergen_state_dict(checkpoint_path: str | Path) -> dict[str, torch.Tensor]: + """Read precisely the inference model state from a MatterGen Lightning checkpoint. + + The official archives contain optimizer, callback, and Hydra metadata in + addition to the state dict. ``weights_only=True`` rejects executable + pickle payloads, and the strict prefix check ensures that none of that + training state can be confused with an exported model tensor. + """ + loaded = torch.load(checkpoint_path, map_location="cpu", weights_only=True) + if not isinstance(loaded, Mapping): + raise TypeError("MatterGen checkpoint must deserialize to a mapping.") + state_dict = loaded.get("state_dict") + if not isinstance(state_dict, Mapping): + raise TypeError("MatterGen checkpoint has no mapping-valued 'state_dict'.") + + routed: dict[str, torch.Tensor] = {} + unexpected: list[str] = [] + for source_name, tensor in state_dict.items(): + if not isinstance(source_name, str) or not isinstance(tensor, torch.Tensor): + raise TypeError("MatterGen state_dict must contain string tensor entries only.") + if not source_name.startswith(MATTERGEN_MODEL_STATE_PREFIX): + unexpected.append(source_name) + continue + target_name = source_name.removeprefix(MATTERGEN_MODEL_STATE_PREFIX) + if not target_name: + raise ValueError("MatterGen state_dict contains an empty inference tensor name.") + if target_name in routed: + raise ValueError(f"MatterGen state_dict maps multiple tensors to {target_name!r}.") + routed[target_name] = tensor + if unexpected: + raise ValueError( + "MatterGen checkpoint has state_dict tensors outside " + f"{MATTERGEN_MODEL_STATE_PREFIX!r}: {sorted(unexpected)[:5]}" + ) + if not routed: + raise ValueError("MatterGen checkpoint has no inference model tensors.") + return routed + + +def _assert_exact_tensor_routing( + package: ModelPackage, state_dict: Mapping[str, torch.Tensor] +) -> None: + """Reject missing, unknown, or shape-mismatched tensors before mutation.""" + if set(package) != {"model"}: + raise ValueError( + "MatterGen weight routing requires exactly one score-core component named 'model'." + ) + initializers = package["model"].graph.initializers + # GraphBuilder materializes scalar ONNX constants as initializers too. They + # are graph literals, not checkpoint parameters, and already have values. + expected = { + name for name, initializer in initializers.items() if initializer.const_value is None + } + actual = set(state_dict) + missing = sorted(expected - actual) + unexpected = sorted(actual - expected) + if missing or unexpected: + details = [] + if missing: + details.append(f"missing {len(missing)} graph tensor(s): {missing[:5]}") + if unexpected: + details.append(f"unrouted {len(unexpected)} checkpoint tensor(s): {unexpected[:5]}") + raise ValueError("MatterGen checkpoint routing is incomplete; " + "; ".join(details)) + + shape_mismatches = [] + for name in sorted(expected): + initializer = initializers[name] + if initializer.shape is None: + raise ValueError(f"MatterGen initializer {name!r} has no concrete shape.") + if not all(isinstance(dimension, int) for dimension in initializer.shape): + raise ValueError(f"MatterGen initializer {name!r} has a symbolic shape.") + expected_shape = tuple(initializer.shape) + actual_shape = tuple(state_dict[name].shape) + if expected_shape != actual_shape: + shape_mismatches.append((name, expected_shape, actual_shape)) + if shape_mismatches: + name, expected_shape, actual_shape = shape_mismatches[0] + raise ValueError( + f"MatterGen tensor shape mismatch for {name!r}: graph expects " + f"{expected_shape}, checkpoint provides {actual_shape} " + f"({len(shape_mismatches)} mismatch(es) total)." + ) + + +def apply_mattergen_checkpoint( + package: ModelPackage, + module, + checkpoint_path: str | Path, +) -> None: + """Load all official MatterGen inference tensors into *package* exactly once. + + ``module.preprocess_weights`` is responsible only for architecture-specific + key normalization. The complete post-normalization key set must match the + score graph's initializers exactly; unlike the generic loader, no unknown + checkpoint weights are merely logged and skipped. + """ + state_dict = load_mattergen_state_dict(checkpoint_path) + preprocess_weights = getattr(module, "preprocess_weights", None) + if not callable(preprocess_weights): + raise TypeError(f"{type(module).__name__} must define preprocess_weights().") + normalized = preprocess_weights(state_dict) + if not isinstance(normalized, Mapping) or not all( + isinstance(name, str) and isinstance(tensor, torch.Tensor) + for name, tensor in normalized.items() + ): + raise TypeError("MatterGen preprocess_weights() must return a string-to-tensor mapping.") + normalized_tensors = dict(normalized) + _assert_exact_tensor_routing(package, normalized_tensors) + package.apply_weights(normalized_tensors) + + report = { + "format": "mobius.weight-loading-report.v1", + "source": str(checkpoint_path), + "source_state_prefix": MATTERGEN_MODEL_STATE_PREFIX, + "output_weight_format": "dense", + "native_fp8": False, + "source_tensors": len(state_dict), + "assigned_tensors": len(normalized_tensors), + "canonicalized_alias_tensors": len(state_dict) - len(normalized_tensors), + "ignored_tensors": 0, + "routing": "exact-post-preprocess-with-validated-outputblock-aliases", + } + package.weight_loading_report = report + package["model"].metadata_props["mobius.weight_loading"] = json.dumps(report, sort_keys=True) diff --git a/src/mobius/integrations/mattergen/_weights_test.py b/src/mobius/integrations/mattergen/_weights_test.py new file mode 100644 index 000000000..2980b1213 --- /dev/null +++ b/src/mobius/integrations/mattergen/_weights_test.py @@ -0,0 +1,115 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +from __future__ import annotations + +import pytest +import torch + +from mobius.integrations.mattergen._weights import ( + MATTERGEN_MODEL_STATE_PREFIX, + _assert_exact_tensor_routing, + apply_mattergen_checkpoint, + load_mattergen_state_dict, +) + + +class _Initializer: + def __init__(self, shape: tuple[int, ...], *, is_literal: bool) -> None: + self.shape = shape + self.const_value = object() if is_literal else None + + +class _Graph: + def __init__(self) -> None: + self.initializers = { + "trained.weight": _Initializer((2, 3), is_literal=False), + "const_1.0_f32": _Initializer((), is_literal=True), + } + + +class _ScoreModel: + def __init__(self) -> None: + self.graph = _Graph() + self.metadata_props: dict[str, str] = {} + + +class _Package(dict[str, _ScoreModel]): + def __init__(self) -> None: + super().__init__({"model": _ScoreModel()}) + self.applied: dict[str, torch.Tensor] | None = None + self.weight_loading_report: dict[str, object] | None = None + + def apply_weights(self, state_dict, **_kwargs) -> None: + self.applied = state_dict + + +class TestMatterGenWeightLoading: + def test_strips_only_the_inference_model_prefix(self, tmp_path) -> None: + checkpoint = tmp_path / "mattergen.ckpt" + weight = torch.arange(6, dtype=torch.float32).reshape(2, 3) + torch.save( + {"state_dict": {f"{MATTERGEN_MODEL_STATE_PREFIX}gemnet.weight": weight}}, + checkpoint, + ) + + state_dict = load_mattergen_state_dict(checkpoint) + + assert state_dict.keys() == {"gemnet.weight"} + assert torch.equal(state_dict["gemnet.weight"], weight) + + def test_rejects_training_or_unknown_state_tensors(self, tmp_path) -> None: + checkpoint = tmp_path / "mattergen.ckpt" + torch.save( + { + "state_dict": { + f"{MATTERGEN_MODEL_STATE_PREFIX}gemnet.weight": torch.ones(1), + "optimizer.step": torch.ones(1), + } + }, + checkpoint, + ) + + with pytest.raises(ValueError, match="outside"): + load_mattergen_state_dict(checkpoint) + + @pytest.mark.parametrize("payload", [{}, {"state_dict": []}, []]) + def test_rejects_non_mapping_checkpoint_schema(self, tmp_path, payload) -> None: + checkpoint = tmp_path / "mattergen.ckpt" + torch.save(payload, checkpoint) + + with pytest.raises(TypeError, match="mapping"): + load_mattergen_state_dict(checkpoint) + + def test_exact_routing_ignores_graph_literal_initializers(self) -> None: + package = _Package() + + _assert_exact_tensor_routing( + package, + {"trained.weight": torch.ones((2, 3), dtype=torch.float32)}, + ) + + def test_report_accounts_for_validated_checkpoint_aliases(self, monkeypatch, tmp_path) -> None: + package = _Package() + source_tensors = { + "trained.weight": torch.ones((2, 3), dtype=torch.float32), + "duplicate.alias": torch.ones(1, dtype=torch.float32), + } + monkeypatch.setattr( + "mobius.integrations.mattergen._weights.load_mattergen_state_dict", + lambda _path: source_tensors, + ) + + class Module: + @staticmethod + def preprocess_weights(_state_dict): + return {"trained.weight": source_tensors["trained.weight"]} + + apply_mattergen_checkpoint(package, Module(), tmp_path / "mattergen.ckpt") + + assert package.applied is not None + assert torch.equal(package.applied["trained.weight"], source_tensors["trained.weight"]) + assert package.weight_loading_report["source_tensors"] == 2 + assert package.weight_loading_report["assigned_tensors"] == 1 + assert package.weight_loading_report["canonicalized_alias_tensors"] == 1 + assert package.weight_loading_report["ignored_tensors"] == 0 diff --git a/src/mobius/models/__init__.py b/src/mobius/models/__init__.py index 6cbef77d2..67bd12570 100644 --- a/src/mobius/models/__init__.py +++ b/src/mobius/models/__init__.py @@ -118,6 +118,8 @@ "Mamba2CausalLMModel", "MambaCausalLMModel", "MaincoderCausalLMModel", + "MatterGenGemNetTModel", + "MatterGenModel", "MiniMaxCausalLMModel", "MiniCPM3CausalLMModel", "MiniCPMCausalLMModel", @@ -327,6 +329,7 @@ from mobius.models.mage_vl import MageVLForConditionalGeneration from mobius.models.maincoder import MaincoderCausalLMModel from mobius.models.mamba import Mamba2CausalLMModel, MambaCausalLMModel +from mobius.models.mattergen import MatterGenGemNetTModel, MatterGenModel from mobius.models.mimi import MimiModel from mobius.models.minicpm import MiniCPM3CausalLMModel, MiniCPMCausalLMModel from mobius.models.minicpmv4_6 import MiniCPMV46ForConditionalGeneration diff --git a/src/mobius/models/mattergen.py b/src/mobius/models/mattergen.py new file mode 100644 index 000000000..bcc3c20dc --- /dev/null +++ b/src/mobius/models/mattergen.py @@ -0,0 +1,1206 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +"""Declarative standard-ONNX MatterGen v1.0.3 GemNet-T neural score core. + +This module faithfully declares the neural portion of +``mattergen.denoiser.GemNetTDenoiser`` at source commit +``842ffe735f7d06cec89d56aa23d9f001e1124b30`` and is compatible with the +``microsoft/mattergen`` Hub revision +``5244495dd9a979ff71abc7548a0b14b9deb0069a``. + +```mermaid +flowchart LR + H[Host: PBC radius graph and triplets] --> G[edge_index, D_st, V_st, lattice cosines, id_swap, triplet ids] + C[Host: condition values and unconditional masks] --> P[property encoders] + T[timestep] --> N[noise_level_encoding] + A[one-based atomic_numbers] --> E[AtomEmbedding Z - 1] + N --> Z[latent per crystal] + P --> Z + E --> GT[GemNet-T triplet interaction blocks] + G --> GT + Z --> GT + GT --> O[atom logits, Cartesian position score, lattice score, crystal energy] +``` + +The host owns fractional-to-Cartesian conversion, periodic image/radius graph +construction, symmetric edge ordering, ``id_swap``, sorted triplet indices, +and ``edge_lattice_cosines``. The latter is the source expression +``cosine_similarity(V_st[:, None], cell[batch[edge_index[0]]], dim=-1)`` and +lets this neural core avoid a cell input. In particular ``edge_index[0]`` is +source ``c`` and ``edge_index[1]`` is target ``a``; ``edge_direction`` is +MatterGen's ``V_st = -distance_vec / distance`` convention. The graph +intentionally does not perform element masking, fractional-coordinate +conversion, stochastic sampling, or PBC construction. +""" + +from __future__ import annotations + +import math +from collections.abc import Mapping, Sequence + +import onnx_ir as ir +import torch +from onnxscript import OpBuilder, nn + +from mobius.components import Embedding, Linear +from mobius.integrations.mattergen._configs import MatterGenConditionSpec, MatterGenConfig + + +def _cast_float(op: OpBuilder, value: ir.Value) -> ir.Value: + """Cast a value to the source model's float32 basis-computation dtype.""" + return op.Cast(value, to=ir.DataType.FLOAT) + + +def _scatter_sum( + op: OpBuilder, + values: ir.Value, + indices: ir.Value, + output_rows: ir.Value, +) -> ir.Value: + """Sum leading-axis rows into ``[output_rows, *values.shape[1:]]``. + + MatterGen uses ``torch_scatter.scatter(..., reduce="sum")`` for atom, + structure, and neighbor aggregation. ``ScatterND(reduction="add")`` is + the standard-ONNX equivalent and handles repeated atom/edge ids exactly. + """ + output_shape = op.Concat(output_rows, op.Shape(values, start=1), axis=0) + initial = op.Expand(op.CastLike(0.0, values), output_shape) + scatter_indices = op.Unsqueeze(indices, [1]) # (rows, 1) + return op.ScatterND(initial, scatter_indices, values, reduction="add") + + +def _ragged_scatter( + op: OpBuilder, + values: ir.Value, + id_reduce: ir.Value, + id_ragged_idx: ir.Value, + num_edges: ir.Value, +) -> ir.Value: + """Materialize MatterGen's dynamically padded triplet tensor. + + ``id_reduce`` is ``id3_ca`` and ``id_ragged_idx`` enumerates neighboring + ``b -> a`` edges within each ``c -> a`` group. The zero appended before + ``ReduceMax`` gives empty-triplet graphs a well-defined one-wide padded + tensor; all updates remain zero, preserving the source sum semantics. + """ + safe_ragged = op.Concat(id_ragged_idx, op.Constant(value_ints=[0]), axis=0) + max_neighbors = op.Add(op.ReduceMax(safe_ragged, keepdims=1), 1) # (1,) + padded_shape = op.Concat(num_edges, max_neighbors, op.Shape(values, start=1), axis=0) + padded = op.Expand(op.CastLike(0.0, values), padded_shape) + coordinates = op.Concat( + op.Unsqueeze(id_reduce, [1]), + op.Unsqueeze(id_ragged_idx, [1]), + axis=1, + ) # (triplets, 2) + return op.ScatterND(padded, coordinates, values) + + +class _GemNetDense(nn.Module): + """MatterGen ``Dense`` with source-compatible ``.linear`` parameter path.""" + + def __init__( + self, in_features: int, out_features: int, *, bias: bool = False, silu: bool = False + ): + super().__init__() + # The nested Linear deliberately matches MatterGen Dense: + # ``.linear.weight`` / ``.linear.bias``. + self.linear = Linear(in_features, out_features, bias=bias) + self._silu = silu + + def forward(self, op: OpBuilder, value: ir.Value) -> ir.Value: + value = self.linear(op, value) + if self._silu: + # MatterGen ScaledSiLU is SiLU(x) / 0.6, not the usual SiLU. + value = op.Mul(op.Mul(value, op.Sigmoid(value)), 1.0 / 0.6) + return value + + +class _ResidualLayer(nn.Module): + """GemNet residual MLP: dense layers followed by ``(x + f(x)) / sqrt(2)``.""" + + def __init__(self, units: int, *, num_layers: int = 2): + super().__init__() + self.dense_mlp = nn.ModuleList( + [_GemNetDense(units, units, silu=True) for _ in range(num_layers)] + ) + + def forward(self, op: OpBuilder, value: ir.Value) -> ir.Value: + residual = value + for layer in self.dense_mlp: + value = layer(op, value) + return op.Mul(op.Add(residual, value), 2.0**-0.5) + + +class _ScalingFactor(nn.Module): + """Persistent MatterGen ``ScalingFactor.scale_factor`` initializer.""" + + def __init__(self): + super().__init__() + # Values are loaded from the checkpoint. Do not bake gemnet-dT.json + # factors here: later blocks' values are checkpoint-specific. + self.scale_factor = nn.Parameter([]) + + def forward(self, op: OpBuilder, _reference: ir.Value, value: ir.Value) -> ir.Value: + return op.Mul(value, self.scale_factor) + + +class _GaussianSmearing(nn.Module): + """PyG 2.6 GaussianSmearing with its persisted ``offset`` buffer.""" + + def __init__(self, num_gaussians: int): + super().__init__() + self.offset = nn.Parameter([num_gaussians]) + # PyG stores this scalar as Python state rather than a state-dict key. + self._coefficient = -0.5 / (1.0 / (num_gaussians - 1)) ** 2 + + def forward(self, op: OpBuilder, distance_scaled: ir.Value) -> ir.Value: + offset = op.CastLike(self.offset, distance_scaled) + centered = op.Sub(op.Unsqueeze(distance_scaled, [1]), op.Unsqueeze(offset, [0])) + return op.Exp(op.Mul(op.Mul(centered, centered), self._coefficient)) + + +class _PolynomialEnvelope(nn.Module): + """MatterGen fifth-order polynomial radial cutoff envelope.""" + + def __init__(self, exponent: int = 5): + super().__init__() + self._p = exponent + self._a = -(exponent + 1) * (exponent + 2) / 2 + self._b = exponent * (exponent + 2) + self._c = -exponent * (exponent + 1) / 2 + + def forward(self, op: OpBuilder, distance_scaled: ir.Value) -> ir.Value: + d_p = op.Pow(distance_scaled, self._p) + envelope = op.Add( + 1.0, + op.Add( + op.Mul(self._a, d_p), + op.Add( + op.Mul(self._b, op.Pow(distance_scaled, self._p + 1)), + op.Mul(self._c, op.Pow(distance_scaled, self._p + 2)), + ), + ), + ) + return op.Where( + op.Less(distance_scaled, 1.0), + envelope, + op.CastLike(0.0, distance_scaled), + ) + + +class _RadialBasis(nn.Module): + """Gaussian radial basis times the source polynomial envelope.""" + + def __init__(self, num_radial: int, cutoff: float): + super().__init__() + self.rbf = _GaussianSmearing(num_radial) + self.envelope = _PolynomialEnvelope() + self._inv_cutoff = 1.0 / cutoff + + def forward(self, op: OpBuilder, distance: ir.Value) -> ir.Value: + distance_scaled = op.Mul(distance, self._inv_cutoff) + envelope = self.envelope(op, distance_scaled) # (edges,) + return op.Mul(op.Unsqueeze(envelope, [1]), self.rbf(op, distance_scaled)) + + +class _CircularBasis(nn.Module): + """Efficient circular basis: radial Gaussian features and Y_l0(cos(phi)).""" + + def __init__(self, num_spherical: int, num_radial: int, cutoff: float): + super().__init__() + self.radial_basis = _RadialBasis(num_radial, cutoff) + self._num_spherical = num_spherical + + def forward( + self, + op: OpBuilder, + distance: ir.Value, + cosine: ir.Value, + ) -> tuple[ir.Value, ir.Value]: + radial = op.Unsqueeze(self.radial_basis(op, distance), [0]) # (1, edges, radial) + # MatterGen's SymPy-generated basis is real spherical harmonics with + # m=0: sqrt((2l+1)/(4*pi)) * P_l(cos(phi)), l=0..num_spherical-1. + p_previous = op.Add(op.Mul(cosine, 0.0), 1.0) + polynomials = [p_previous] + if self._num_spherical > 1: + p_current = cosine + polynomials.append(p_current) + for degree in range(2, self._num_spherical): + p_next = op.Mul( + 1.0 / degree, + op.Sub( + op.Mul((2 * degree - 1), op.Mul(cosine, p_current)), + op.Mul(degree - 1, p_previous), + ), + ) + polynomials.append(p_next) + p_previous, p_current = p_current, p_next + harmonics = [ + op.Mul(math.sqrt((2 * degree + 1) / (4 * math.pi)), polynomial) + for degree, polynomial in enumerate(polynomials) + ] + spherical = op.Concat(*[op.Unsqueeze(value, [1]) for value in harmonics], axis=1) + return radial, spherical # (1, E, R), (T, S) + + +class _EfficientInteractionDownProjection(nn.Module): + """Source ``EfficientInteractionDownProjection`` with dynamic ragged padding.""" + + def __init__(self, num_spherical: int, num_radial: int, emb_size_interm: int): + super().__init__() + self.weight = nn.Parameter([num_spherical, num_radial, emb_size_interm]) + + def forward( + self, + op: OpBuilder, + radial: ir.Value, + spherical: ir.Value, + id_ca: ir.Value, + id_ragged_idx: ir.Value, + ) -> tuple[ir.Value, ir.Value]: + # Broadcasted MatMul reproduces torch.matmul([1,E,R], [S,R,C]) + # -> [S,E,C], then permutes to (E, C, S). + radial_weighted = op.Transpose(op.MatMul(radial, self.weight), perm=[1, 2, 0]) + num_edges = op.Shape(radial_weighted, start=0, end=1) + padded_spherical = _ragged_scatter(op, spherical, id_ca, id_ragged_idx, num_edges) + return radial_weighted, op.Transpose(padded_spherical, perm=[0, 2, 1]) + + +class _EfficientInteractionBilinear(nn.Module): + """Source bilinear triplet aggregation with standard-ONNX ScatterND.""" + + def __init__(self, emb_size: int, emb_size_interm: int, units_out: int): + super().__init__() + self.weight = nn.Parameter([emb_size, emb_size_interm, units_out]) + + def forward( + self, + op: OpBuilder, + basis: tuple[ir.Value, ir.Value], + messages: ir.Value, + id_reduce: ir.Value, + id_ragged_idx: ir.Value, + ) -> ir.Value: + radial_weighted, spherical = basis # (E, C, S), (E, S, K) + num_edges = op.Shape(radial_weighted, start=0, end=1) + padded_messages = _ragged_scatter(op, messages, id_reduce, id_ragged_idx, num_edges) + # First contract neighbor K then spherical S: (E,S,K) @ (E,K,D). + summed_neighbors = op.MatMul(spherical, padded_messages) # (E, S, D) + radial_messages = op.MatMul(radial_weighted, summed_neighbors) # (E, C, D) + # Batch dimension D selects the corresponding bilinear weight [D,C,O]. + projected = op.MatMul(op.Transpose(radial_messages, perm=[2, 0, 1]), self.weight) + return op.ReduceSum(projected, [0], keepdims=0) # (E, O) + + +class _AtomEmbedding(nn.Module): + """Atom type table with MatterGen's one-based ``Z - 1`` lookup.""" + + def __init__(self, num_atom_types: int, emb_size: int): + super().__init__() + self.embeddings = Embedding(num_atom_types, emb_size) + + def forward(self, op: OpBuilder, atomic_numbers: ir.Value) -> ir.Value: + return self.embeddings(op, op.Sub(atomic_numbers, 1)) + + +class _EdgeEmbedding(nn.Module): + """Concatenate source/target atom features and radial features into edges.""" + + def __init__(self, atom_features: int, edge_features: int, out_features: int): + super().__init__() + self.dense = _GemNetDense( + 2 * atom_features + edge_features, + out_features, + silu=True, + ) + + def forward( + self, + op: OpBuilder, + atoms: ir.Value, + edge_features: ir.Value, + idx_s: ir.Value, + idx_t: ir.Value, + ) -> ir.Value: + source_atoms = op.Gather(atoms, idx_s) # (E, atom_dim) + target_atoms = op.Gather(atoms, idx_t) # (E, atom_dim) + return self.dense(op, op.Concat(source_atoms, target_atoms, edge_features, axis=-1)) + + +class _AtomUpdateBlock(nn.Module): + """Aggregate radial-filtered edge messages into atom embeddings.""" + + def __init__(self, config: MatterGenConfig): + super().__init__() + self.dense_rbf = _GemNetDense(config.emb_size_rbf, config.emb_size_edge) + self.scale_sum = _ScalingFactor() + self.layers = nn.ModuleList( + [ + _GemNetDense(config.emb_size_edge, config.emb_size_atom, silu=True), + *[_ResidualLayer(config.emb_size_atom) for _ in range(config.num_atom)], + ] + ) + + def forward( + self, + op: OpBuilder, + atoms: ir.Value, + messages: ir.Value, + radial: ir.Value, + idx_t: ir.Value, + ) -> ir.Value: + radial_message = self.dense_rbf(op, radial) # (E, edge_dim) + aggregated = _scatter_sum( + op, + op.Mul(messages, radial_message), + idx_t, + op.Shape(atoms, start=0, end=1), + ) # (N, edge_dim) + value = self.scale_sum(op, messages, aggregated) + for layer in self.layers: + value = layer(op, value) + return value # (N, atom_dim) + + +class _OutputBlock(_AtomUpdateBlock): + """GemNet output block with source-compatible energy and direct-force heads.""" + + def __init__(self, config: MatterGenConfig): + super().__init__(config) + self.out_energy = _GemNetDense(config.emb_size_atom, config.num_targets) + self.scale_rbf_F = _ScalingFactor() + self.seq_forces = nn.ModuleList( + [ + _GemNetDense(config.emb_size_edge, config.emb_size_edge, silu=True), + *[_ResidualLayer(config.emb_size_edge) for _ in range(config.num_atom)], + ] + ) + self.out_forces = _GemNetDense(config.emb_size_edge, config.num_targets) + self.dense_rbf_F = _GemNetDense(config.emb_size_rbf, config.emb_size_edge) + + def forward( + self, + op: OpBuilder, + atoms: ir.Value, + messages: ir.Value, + radial: ir.Value, + idx_t: ir.Value, + ) -> tuple[ir.Value, ir.Value]: + # The inherited path is the source energy head's ``seq_energy`` alias. + # This declaration uses ``layers`` once; preprocessing canonicalizes the + # duplicated PyTorch state-dict alias ``seq_energy`` onto it. + radial_energy = self.dense_rbf(op, radial) + energy_hidden = _scatter_sum( + op, + op.Mul(messages, radial_energy), + idx_t, + op.Shape(atoms, start=0, end=1), + ) + energy_hidden = self.scale_sum(op, messages, energy_hidden) + for layer in self.layers: + energy_hidden = layer(op, energy_hidden) + energy = self.out_energy(op, energy_hidden) # (N, 1) + + force_hidden = messages + for layer in self.seq_forces: + force_hidden = layer(op, force_hidden) + force_hidden = op.Mul(force_hidden, self.dense_rbf_F(op, radial)) + force_hidden = self.scale_rbf_F(op, messages, force_hidden) + return energy, self.out_forces(op, force_hidden) # (N,1), (E,1) + + +class _TripletInteraction(nn.Module): + """GemNet-T triplet message: radial filtering, bilinear angle sum, edge swap.""" + + def __init__(self, config: MatterGenConfig): + super().__init__() + self.dense_ba = _GemNetDense(config.emb_size_edge, config.emb_size_edge, silu=True) + self.mlp_rbf = _GemNetDense(config.emb_size_rbf, config.emb_size_edge) + self.scale_rbf = _ScalingFactor() + self.mlp_cbf = _EfficientInteractionBilinear( + config.emb_size_trip, + config.emb_size_cbf, + config.emb_size_bil_trip, + ) + self.scale_cbf_sum = _ScalingFactor() + self.down_projection = _GemNetDense( + config.emb_size_edge, + config.emb_size_trip, + silu=True, + ) + self.up_projection_ca = _GemNetDense( + config.emb_size_bil_trip, + config.emb_size_edge, + silu=True, + ) + self.up_projection_ac = _GemNetDense( + config.emb_size_bil_trip, + config.emb_size_edge, + silu=True, + ) + + def forward( + self, + op: OpBuilder, + messages: ir.Value, + radial: ir.Value, + circular: tuple[ir.Value, ir.Value], + id_ragged_idx: ir.Value, + id_swap: ir.Value, + id_ba: ir.Value, + id_ca: ir.Value, + ) -> ir.Value: + incoming = self.dense_ba(op, messages) + radial_message = op.Mul(incoming, self.mlp_rbf(op, radial)) + incoming = self.scale_rbf(op, incoming, radial_message) + incoming = self.down_projection(op, incoming) # (E, trip_dim) + triplet_messages = op.Gather(incoming, id_ba) # (T, trip_dim) + aggregated = self.mlp_cbf(op, circular, triplet_messages, id_ca, id_ragged_idx) + aggregated = self.scale_cbf_sum(op, triplet_messages, aggregated) + forward = self.up_projection_ca(op, aggregated) + reverse = op.Gather(self.up_projection_ac(op, aggregated), id_swap) + return op.Mul(op.Add(forward, reverse), 2.0**-0.5) + + +class _InteractionBlockTripletsOnly(nn.Module): + """One source GemNet-T triplet-only interaction block.""" + + def __init__(self, config: MatterGenConfig): + super().__init__() + self.dense_ca = _GemNetDense(config.emb_size_edge, config.emb_size_edge, silu=True) + self.trip_interaction = _TripletInteraction(config) + self.layers_before_skip = nn.ModuleList( + [_ResidualLayer(config.emb_size_edge) for _ in range(config.num_before_skip)] + ) + self.layers_after_skip = nn.ModuleList( + [_ResidualLayer(config.emb_size_edge) for _ in range(config.num_after_skip)] + ) + self.atom_update = _AtomUpdateBlock(config) + self.concat_layer = _EdgeEmbedding( + config.emb_size_atom, + config.emb_size_edge, + config.emb_size_edge, + ) + self.residual_m = nn.ModuleList( + [_ResidualLayer(config.emb_size_edge) for _ in range(config.num_concat)] + ) + + def forward( + self, + op: OpBuilder, + atoms: ir.Value, + messages: ir.Value, + radial_triplet: ir.Value, + circular: tuple[ir.Value, ir.Value], + id_ragged_idx: ir.Value, + id_swap: ir.Value, + id_ba: ir.Value, + id_ca: ir.Value, + radial_atom: ir.Value, + idx_s: ir.Value, + idx_t: ir.Value, + ) -> tuple[ir.Value, ir.Value]: + update = op.Add( + self.dense_ca(op, messages), + self.trip_interaction( + op, + messages, + radial_triplet, + circular, + id_ragged_idx, + id_swap, + id_ba, + id_ca, + ), + ) + update = op.Mul(update, 2.0**-0.5) + for layer in self.layers_before_skip: + update = layer(op, update) + + messages = op.Mul(op.Add(messages, update), 2.0**-0.5) + for layer in self.layers_after_skip: + messages = layer(op, messages) + + atom_update = self.atom_update(op, atoms, messages, radial_atom, idx_t) + atoms = op.Mul(op.Add(atoms, atom_update), 2.0**-0.5) + edge_update = self.concat_layer(op, atoms, messages, idx_s, idx_t) + for layer in self.residual_m: + edge_update = layer(op, edge_update) + return atoms, op.Mul(op.Add(messages, edge_update), 2.0**-0.5) + + +class _RBFBasedLatticeUpdateBlock(nn.Module): + """MatterGen direct lattice-score head from radial edge scores.""" + + def __init__(self, config: MatterGenConfig): + super().__init__() + self.mlp = nn.ModuleList( + [ + _GemNetDense(config.emb_size_edge, config.emb_size_edge, silu=True), + _GemNetDense(config.emb_size_edge, config.emb_size_edge), + ] + ) + self.dense_rbf_F = _GemNetDense(config.emb_size_rbf, config.emb_size_edge) + self.out_forces = _GemNetDense(config.emb_size_edge, config.num_targets) + + def forward( + self, + op: OpBuilder, + edge_embeddings: ir.Value, + edge_direction: ir.Value, + batch_edge: ir.Value, + batch_size: ir.Value, + radial: ir.Value, + ) -> ir.Value: + score = edge_embeddings + for layer in self.mlp: + score = layer(op, score) + score = self.out_forces(op, op.Mul(score, self.dense_rbf_F(op, radial))) # (E,1) + + # Source normalizes each edge score by the number of source edges in its + # crystal before the symmetric outer-product lattice aggregation. + edge_count = _scatter_sum( + op, + op.Expand(op.CastLike(1.0, score), op.Shape(batch_edge)), + batch_edge, + batch_size, + ) + score = op.Div(score, op.Unsqueeze(op.Gather(edge_count, batch_edge), [1])) + + # ``distance_vec`` in source is V_st * D_st, then normalized again; + # normalize edge_direction here to preserve that exact dataflow. + norm = op.Sqrt(op.ReduceSum(op.Mul(edge_direction, edge_direction), [1], keepdims=1)) + unit_direction = op.Div(edge_direction, norm) + outer = op.Mul( + op.Unsqueeze(unit_direction, [2]), + op.Unsqueeze(unit_direction, [1]), + ) # (E, 3, 3) + lattice = _scatter_sum( + op, + op.Mul(op.Unsqueeze(score, [2]), outer), + batch_edge, + batch_size, + ) + # The reference transposes after scatter; the outer product is symmetric + # but retain it so exported graph semantics mirror the source literally. + return op.Transpose(lattice, perm=[0, 2, 1]) + + +class _AngleEdgeEmbedding(nn.Module): + """Source ``nn.Sequential(Linear, ReLU, Linear)`` with 0/2 key indices.""" + + def __init__(self, input_size: int, hidden_size: int): + super().__init__() + # ``nn.ModuleList`` cannot represent Sequential's parameterless ReLU + # at index 1. Register the two real children under their source keys. + setattr(self, "0", Linear(input_size, hidden_size)) + setattr(self, "2", Linear(hidden_size, hidden_size)) + + def forward(self, op: OpBuilder, value: ir.Value) -> ir.Value: + first = getattr(self, "0")(op, value) + return getattr(self, "2")(op, op.Relu(first)) + + +class _GemNetT(nn.Module): + """MatterGen GemNet-T backbone with triplet interactions and all score heads.""" + + def __init__(self, config: MatterGenConfig): + super().__init__() + self.radial_basis = _RadialBasis(config.num_radial, config.cutoff) + self.cbf_basis3 = _CircularBasis( + config.num_spherical, config.num_radial, config.cutoff + ) + self.lattice_out_blocks = nn.ModuleList( + [_RBFBasedLatticeUpdateBlock(config) for _ in range(config.num_blocks + 1)] + ) + self.mlp_rbf_lattice = _GemNetDense(config.num_radial, config.emb_size_rbf) + self.mlp_rbf3 = _GemNetDense(config.num_radial, config.emb_size_rbf) + self.mlp_cbf3 = _EfficientInteractionDownProjection( + config.num_spherical, + config.num_radial, + config.emb_size_cbf, + ) + self.mlp_rbf_h = _GemNetDense(config.num_radial, config.emb_size_rbf) + self.mlp_rbf_out = _GemNetDense(config.num_radial, config.emb_size_rbf) + self.atom_emb = _AtomEmbedding(config.num_atom_types, config.hidden_size) + self.atom_latent_emb = Linear( + config.hidden_size + config.latent_dim, config.emb_size_atom + ) + self.edge_emb = _EdgeEmbedding( + config.emb_size_atom, config.num_radial, config.emb_size_edge + ) + self.angle_edge_emb = _AngleEdgeEmbedding( + config.emb_size_edge + 3, + config.emb_size_edge, + ) + self.int_blocks = nn.ModuleList( + [_InteractionBlockTripletsOnly(config) for _ in range(config.num_blocks)] + ) + self.out_blocks = nn.ModuleList( + [_OutputBlock(config) for _ in range(config.num_blocks + 1)] + ) + self.cond_adapt_layers: _ConditionAdaptLayers | None = None + self.cond_mixin_layers: _ConditionMixinLayers | None = None + self._config = config + + def forward( + self, + op: OpBuilder, + atomic_numbers: ir.Value, + batch: ir.Value, + latent: ir.Value, + edge_index: ir.Value, + edge_distance: ir.Value, + edge_direction: ir.Value, + edge_lattice_cosines: ir.Value, + id_swap: ir.Value, + id3_ba: ir.Value, + id3_ca: ir.Value, + id3_ragged_idx: ir.Value, + adapter_embeddings: Mapping[str, ir.Value] | None = None, + adapter_use_unconditional: Mapping[str, ir.Value] | None = None, + *, + return_energy: bool = False, + ) -> tuple[ir.Value, ir.Value, ir.Value] | tuple[ir.Value, ir.Value, ir.Value, ir.Value]: + idx_s = op.Gather(edge_index, op.Constant(value_int=0), axis=0) + idx_t = op.Gather(edge_index, op.Constant(value_int=1), axis=0) + batch_edge = op.Gather(batch, idx_s) + batch_size = op.Shape(latent, start=0, end=1) + + # Triplet angle of b -> a <- c uses the host's already normalized V_st. + cosine = op.Clip( + op.ReduceSum( + op.Mul(op.Gather(edge_direction, id3_ca), op.Gather(edge_direction, id3_ba)), + [1], + keepdims=0, + ), + -1.0, + 1.0, + ) + radial_circular, spherical = self.cbf_basis3(op, edge_distance, cosine) + radial = self.radial_basis(op, edge_distance) # (E, num_radial) + + atoms = self.atom_emb(op, atomic_numbers) # (N, hidden_dim) + latent_per_atom = op.Gather(latent, batch) # (N, latent_dim) + atoms = self.atom_latent_emb(op, op.Concat(atoms, latent_per_atom, axis=1)) + messages = self.edge_emb(op, atoms, radial, idx_s, idx_t) + + # Host-precomputed source cosine_similarity(V_st[:,None], cell[batch_edge]). + # Each edge has its alignment to the three lattice-vector rows: (E, 3). + messages = op.Concat( + messages, + edge_lattice_cosines, + axis=-1, + ) + messages = self.angle_edge_emb(op, messages) + + radial_triplet = self.mlp_rbf3(op, radial) + circular = self.mlp_cbf3(op, radial_circular, spherical, id3_ca, id3_ragged_idx) + radial_atom = self.mlp_rbf_h(op, radial) + radial_output = self.mlp_rbf_out(op, radial) + + energy, edge_force = self.out_blocks[0](op, atoms, messages, radial_output, idx_t) + radial_lattice = self.mlp_rbf_lattice(op, radial) + lattice_score = self.lattice_out_blocks[0]( + op, + messages, + edge_direction, + batch_edge, + batch_size, + radial_lattice, + ) + + for index, block in enumerate(self.int_blocks): + if self._config.condition_on_adapt: + if ( + adapter_embeddings is None + or adapter_use_unconditional is None + or self.cond_adapt_layers is None + or self.cond_mixin_layers is None + ): + raise RuntimeError("MatterGen adapter layers require condition embeddings and masks.") + adaptation = op.Mul(atoms, 0.0) + for condition_name in self._config.condition_on_adapt: + condition = adapter_embeddings[condition_name] + condition_per_atom = op.Gather(condition, batch) # (N, hidden_dim) + adapted = self.cond_adapt_layers( + op, + condition_name, + index, + op.Concat(atoms, condition_per_atom, axis=-1), + ) + adapted = self.cond_mixin_layers(op, condition_name, index, adapted) + use_conditional = op.Not( + op.Gather(adapter_use_unconditional[condition_name], batch) + ) + adaptation = op.Add( + adaptation, + op.Mul( + op.Unsqueeze(op.CastLike(use_conditional, adapted), [1]), adapted + ), + ) + atoms = op.Add(atoms, adaptation) + + atoms, messages = block( + op, + atoms, + messages, + radial_triplet, + circular, + id3_ragged_idx, + id_swap, + id3_ba, + id3_ca, + radial_atom, + idx_s, + idx_t, + ) + block_energy, block_force = self.out_blocks[index + 1]( + op, atoms, messages, radial_output, idx_t + ) + energy = op.Add(energy, block_energy) + edge_force = op.Add(edge_force, block_force) + lattice_score = op.Add( + lattice_score, + self.lattice_out_blocks[index + 1]( + op, + messages, + edge_direction, + batch_edge, + batch_size, + self.mlp_rbf_lattice(op, radial), + ), + ) + + # Each scalar edge force is mapped onto V_st and summed at its target a. + position_score = _scatter_sum( + op, + op.Mul(op.Unsqueeze(edge_force, [2]), op.Unsqueeze(edge_direction, [1])), + idx_t, + op.Shape(atomic_numbers, start=0, end=1), + ) + position_score = op.Squeeze(position_score, [1]) # (N, 3) + if return_energy: + return atoms, position_score, lattice_score, energy + return atoms, position_score, lattice_score + + +class _ScalarNoiseLevelEncoding(nn.Module): + """MatterGen's interleaved sin/cos ``NoiseLevelEncoding``.""" + + def __init__(self, hidden_dim: int): + super().__init__() + self.div_term = nn.Parameter([hidden_dim // 2]) + # Source registers div_term as a float32 buffer. It must not be + # demoted with model weights when the exporter requests fp16/bf16. + setattr(self.div_term, "_keep_float32", True) + self._hidden_dim = hidden_dim + + def forward(self, op: OpBuilder, value: ir.Value) -> ir.Value: + value = op.Reshape(_cast_float(op, value), [-1, 1]) + div_term = op.Reshape(_cast_float(op, self.div_term), [1, -1]) + angles = op.Mul(value, div_term) # (B, hidden_dim / 2) + # Source writes sin into even and cos into odd columns. Stack then + # reshape interleaves them; concatenating sin/cos would be incorrect. + interleaved = op.Concat( + op.Unsqueeze(op.Sin(angles), [2]), + op.Unsqueeze(op.Cos(angles), [2]), + axis=2, + ) + return op.Reshape(interleaved, [-1, self._hidden_dim]) + + +class _StandardScaler(nn.Module): + """Persistent ``StandardScalerTorch`` parameters used before scalar encoding.""" + + def __init__(self, *, log10_transform: bool): + super().__init__() + self.means = nn.Parameter([1]) + self.stds = nn.Parameter([1]) + self._log10_transform = log10_transform + + def forward(self, op: OpBuilder, value: ir.Value) -> ir.Value: + value = _cast_float(op, value) + if self._log10_transform: + value = op.Div(op.Log(value), math.log(10.0)) + return op.Div(op.Sub(value, self.means), self.stds) + + +class _EmbeddingVector(nn.Module): + """Source unconditional condition embedding: Embedding(1, hidden_dim).""" + + def __init__(self, hidden_dim: int): + super().__init__() + self.embedding = Embedding(1, hidden_dim) + + def forward(self, op: OpBuilder, reference: ir.Value) -> ir.Value: + target_shape = op.Concat( + op.Shape(reference, start=0, end=1), + op.Shape(self.embedding.weight, start=1), + axis=0, + ) + return op.Expand(self.embedding.weight, target_shape) + + +class _ChemicalSystemMultiHotEmbedding(nn.Module): + """Source chemical-system multi-hot linear encoder.""" + + def __init__(self, hidden_dim: int): + super().__init__() + self.embedding = Linear(101, hidden_dim) + + def forward(self, op: OpBuilder, value: ir.Value) -> ir.Value: + return self.embedding(op, value) + + +class _SpaceGroupEmbeddingVector(nn.Module): + """Source one-based space-group encoder: gather ``embedding[x.long() - 1]``.""" + + def __init__(self, hidden_dim: int): + super().__init__() + self.embedding = Embedding(230, hidden_dim) + + def forward(self, op: OpBuilder, value: ir.Value) -> ir.Value: + return self.embedding(op, op.Sub(value, 1)) + + +class _PropertyEmbedding(nn.Module): + """Faithful source property switch between conditional and unconditional embeddings.""" + + def __init__(self, spec: MatterGenConditionSpec, hidden_dim: int): + super().__init__() + self._spec = spec + self.conditional_embedding_module: ( + _ScalarNoiseLevelEncoding + | _ChemicalSystemMultiHotEmbedding + | _SpaceGroupEmbeddingVector + ) + if spec.kind == "scalar_sinusoidal": + self.conditional_embedding_module = _ScalarNoiseLevelEncoding(hidden_dim) + elif spec.kind == "chemical_system_multihot": + self.conditional_embedding_module = _ChemicalSystemMultiHotEmbedding(hidden_dim) + elif spec.kind == "space_group_index": + self.conditional_embedding_module = _SpaceGroupEmbeddingVector(hidden_dim) + else: # MatterGenConfig.validate already rejects this path. + raise ValueError(f"Unsupported MatterGen condition encoder {spec.kind!r}") + + if spec.unconditional == "embedding_vector": + self.unconditional_embedding_module = _EmbeddingVector(hidden_dim) + elif spec.unconditional != "zeros": + raise ValueError( + f"Unsupported unconditional condition embedding {spec.unconditional!r}" + ) + if spec.scaler == "standard": + self.scaler = _StandardScaler(log10_transform=spec.log10_transform) + + def forward( + self, + op: OpBuilder, + value: ir.Value, + use_unconditional: ir.Value, + dtype: ir.DataType, + ) -> ir.Value: + if self._spec.scaler == "standard": + value = self.scaler(op, value) + conditional = self.conditional_embedding_module(op, value) + conditional = op.Cast(conditional, to=dtype) + if self._spec.unconditional == "zeros": + unconditional = op.Expand(op.CastLike(0.0, conditional), op.Shape(conditional)) + else: + unconditional = self.unconditional_embedding_module(op, conditional) + mask = op.Reshape(use_unconditional, [-1, 1]) + return op.Where(mask, unconditional, conditional) + + +class _NamedConditionModules(nn.Module): + """Small ModuleDict equivalent preserving MatterGen condition-name paths.""" + + def __init__(self, specs: Sequence[MatterGenConditionSpec], hidden_dim: int): + super().__init__() + self._names = tuple(spec.name for spec in specs) + for spec in specs: + setattr(self, spec.name, _PropertyEmbedding(spec, hidden_dim)) + + def get(self, name: str) -> _PropertyEmbedding: + return getattr(self, name) + + def forward( + self, + op: OpBuilder, + name: str, + value: ir.Value, + use_unconditional: ir.Value, + dtype: ir.DataType, + ) -> ir.Value: + """Invoke the named child beneath this container's ONNX name scope.""" + return self.get(name)(op, value, use_unconditional, dtype) + + +class _ConditionAdaptLayers(nn.Module): + """Source ``cond_adapt_layers..`` ModuleDict hierarchy.""" + + def __init__(self, condition_names: Sequence[str], config: MatterGenConfig): + super().__init__() + self._names = tuple(condition_names) + for name in condition_names: + layers = nn.ModuleList( + [_ConditionAdaptLayer(config.hidden_size) for _ in range(config.num_blocks)] + ) + setattr(self, name, layers) + + def forward( + self, + op: OpBuilder, + name: str, + block_index: int, + value: ir.Value, + ) -> ir.Value: + """Apply one named source adapter MLP in its declared container scope.""" + return getattr(self, name)[block_index](op, value) + + +class _ConditionAdaptLayer(nn.Module): + """Adapter MLP preserving PyTorch Sequential's ``.0`` and ``.2`` paths.""" + + def __init__(self, hidden_size: int): + super().__init__() + setattr(self, "0", Linear(hidden_size * 2, hidden_size)) + setattr(self, "2", Linear(hidden_size, hidden_size)) + + def forward(self, op: OpBuilder, value: ir.Value) -> ir.Value: + first = getattr(self, "0")(op, value) + return getattr(self, "2")(op, op.Relu(first)) + + +class _ConditionMixinLayers(nn.Module): + """Source ``cond_mixin_layers..`` zero-init linear hierarchy.""" + + def __init__(self, condition_names: Sequence[str], config: MatterGenConfig): + super().__init__() + self._names = tuple(condition_names) + for name in condition_names: + setattr( + self, + name, + nn.ModuleList( + [ + Linear(config.hidden_size, config.hidden_size, bias=False) + for _ in range(config.num_blocks) + ] + ), + ) + + def forward( + self, + op: OpBuilder, + name: str, + block_index: int, + value: ir.Value, + ) -> ir.Value: + """Apply one named source mixin linear layer in its declared scope.""" + return getattr(self, name)[block_index](op, value) + + +class MatterGenModel(nn.Module): + """MatterGen GemNet-T score core with explicit host-provided geometric tensors. + + .. mermaid:: + + flowchart LR + H[Host: noisy crystal and periodic graph] --> G[ONNX: GemNet-T score core] + C[Host: condition values and CFG masks] --> G + G --> S[Host: source scheduler and crystal validation] + + ``edge_lattice_cosines`` is host-produced ``[E,3]`` from the source's + lattice cosine calculation; this preserves GemNet's angle-edge embedding + without accepting a cell tensor. ``condition_values`` and + ``condition_use_unconditional`` are mappings keyed by + :attr:`MatterGenConfig.condition_input_specs` names. Each value is ``[B]`` + for scalar/space-group properties or ``[B,101]`` for chemical systems; + every mask is bool ``[B]`` where true selects the unconditional embedding. + A later MatterGen task uses this config-derived ABI to declare named ONNX + ports. The final ``energy`` output retains the trained GemNet OutputBlock + energy path, so every official checkpoint initializer remains reachable. + """ + + default_task: str = "mattergen-score" + category: str = "Diffusion" + config_class = MatterGenConfig + + def __init__(self, config: MatterGenConfig): + super().__init__() + config.validate() + self.config = config + self.noise_level_encoding = _ScalarNoiseLevelEncoding(config.hidden_size) + self.property_embeddings = _NamedConditionModules( + config.property_embeddings, + config.hidden_size, + ) + self.property_embeddings_adapt = _NamedConditionModules( + config.property_embeddings_adapt, + config.hidden_size, + ) + self.gemnet = _GemNetT(config) + if config.condition_on_adapt: + # These are attributes of the GemNet source module, not the + # denoiser, and therefore live beneath ``gemnet`` in state dicts. + self.gemnet.cond_adapt_layers = _ConditionAdaptLayers( + config.condition_on_adapt, config + ) + self.gemnet.cond_mixin_layers = _ConditionMixinLayers( + config.condition_on_adapt, config + ) + self.fc_atom = Linear(config.hidden_size, config.num_atom_types) + + def forward( + self, + op: OpBuilder, + atomic_numbers: ir.Value, + batch: ir.Value, + timestep: ir.Value, + edge_index: ir.Value, + edge_distance: ir.Value, + edge_direction: ir.Value, + edge_lattice_cosines: ir.Value, + id_swap: ir.Value, + id3_ba: ir.Value, + id3_ca: ir.Value, + id3_ragged_idx: ir.Value, + condition_values: Mapping[str, ir.Value] | None = None, + condition_use_unconditional: Mapping[str, ir.Value] | None = None, + ) -> tuple[ir.Value, ir.Value, ir.Value, ir.Value]: + """Return atom logits, position score, lattice score, and crystal energy. + + The outputs are respectively ``[N,101]``, ``[N,3]``, ``[B,3,3]``, + and source-aggregated ``[B,1]`` energy. + """ + specs = self.config.condition_input_specs + condition_values = {} if condition_values is None else condition_values + condition_use_unconditional = ( + {} if condition_use_unconditional is None else condition_use_unconditional + ) + expected_condition_names = {spec.name for spec in specs} + if ( + set(condition_values) != expected_condition_names + or set(condition_use_unconditional) != expected_condition_names + ): + raise ValueError( + "MatterGen condition value and mask mappings must each contain " + "exactly the config.condition_input_specs names" + ) + + timestep_embedding = self.noise_level_encoding(op, timestep) # (B, hidden_dim) + base_embeddings: list[ir.Value] = [] + adapter_embeddings: dict[str, ir.Value] = {} + adapter_masks: dict[str, ir.Value] = {} + for spec in specs: + collection = ( + self.property_embeddings + if not spec.is_adapter + else self.property_embeddings_adapt + ) + embedding = collection( + op, + spec.name, + condition_values[spec.name], + condition_use_unconditional[spec.name], + self.config.dtype, + ) + if not spec.is_adapter: + base_embeddings.append(embedding) + else: + adapter_embeddings[spec.name] = embedding + adapter_masks[spec.name] = condition_use_unconditional[spec.name] + + latent = op.Cast(timestep_embedding, to=self.config.dtype) + if base_embeddings: + # Source's get_property_embeddings sorts ModuleDict keys, which is + # encoded by MatterGenConfig before the module is constructed. + latent = op.Concat(latent, *base_embeddings, axis=-1) + + atom_embeddings, position_score, lattice_score, atom_energy = self.gemnet( + op, + atomic_numbers, + batch, + latent, + edge_index, + edge_distance, + edge_direction, + edge_lattice_cosines, + id_swap, + id3_ba, + id3_ca, + id3_ragged_idx, + adapter_embeddings, + adapter_masks, + return_energy=True, + ) + # Source applies fc_atom to the final GemNet atom embedding, not to + # GemNet's separate auxiliary per-atom energy estimate. + atom_logits = self.fc_atom(op, atom_embeddings) + # GemNet sums each per-atom target within its crystal before returning + # ModelOutput.energy. ``timestep`` explicitly carries the batch width. + energy = _scatter_sum( + op, + atom_energy, + batch, + op.Shape(timestep, start=0, end=1), + ) + return atom_logits, position_score, lattice_score, energy + + def preprocess_weights( + self, state_dict: dict[str, torch.Tensor] + ) -> dict[str, torch.Tensor]: + """Strip the Lightning model prefix and validate exact score-core routing. + + Only tensors rooted at ``diffusion_module.model.`` (or an already + stripped inference mapping) are accepted; no training keys are + discarded here. PyTorch serializes OutputBlock's shared ``layers`` + module twice under ``seq_energy``; that documented alias is + canonicalized after verifying its duplicate tensor agrees. + """ + prefix = "diffusion_module.model." + routed: dict[str, torch.Tensor] = {} + for source_name, value in state_dict.items(): + if source_name.startswith(prefix): + name = source_name.removeprefix(prefix) + elif source_name.startswith( + ( + "noise_level_encoding.", + "fc_atom.", + "property_embeddings.", + "property_embeddings_adapt.", + "gemnet.", + ) + ): + # Already stripped state dictionaries are accepted unchanged. + name = source_name + else: + raise ValueError(f"Unexpected MatterGen checkpoint key: {source_name!r}") + + canonical = name.replace(".seq_energy.", ".layers.") + existing = routed.get(canonical) + if existing is not None: + if not torch.equal(existing, value): + raise ValueError( + f"Conflicting MatterGen OutputBlock alias tensors for {canonical!r}" + ) + continue + routed[canonical] = value + + expected = set(self.state_dict()) + unexpected = set(routed) - expected + if unexpected: + raise ValueError(f"Unexpected routed MatterGen parameters: {sorted(unexpected)!r}") + missing = expected - set(routed) + if missing: + raise KeyError(f"Missing MatterGen score-core parameters: {sorted(missing)!r}") + return routed + + +# A descriptive public alias for callers that need to distinguish this score +# core from host-side MatterGen sampling orchestration. +MatterGenGemNetTModel = MatterGenModel + +__all__ = ["MatterGenGemNetTModel", "MatterGenModel"] diff --git a/src/mobius/tasks/__init__.py b/src/mobius/tasks/__init__.py index b82e691c5..c73e7aff8 100644 --- a/src/mobius/tasks/__init__.py +++ b/src/mobius/tasks/__init__.py @@ -76,6 +76,7 @@ "MiniCPMVLTask", "MuseGlimmerVLTask", "MaskedDiffusionTask", + "MatterGenScoreTask", "MoshiDepformerTask", "MoshiTemporalTask", "MiniMaxMusic3ConditionTask", @@ -173,6 +174,7 @@ from mobius.tasks._kimi_k3 import KimiK3CausalLMTask from mobius.tasks._kimi_linear import KimiLinearCausalLMTask from mobius.tasks._masked_diffusion import MaskedDiffusionTask +from mobius.tasks._mattergen import MatterGenScoreTask from mobius.tasks._minimax_music3 import ( MiniMaxMusic3ConditionTask, MiniMaxMusic3DenoisingTask, @@ -246,6 +248,7 @@ "gguf-embedding-feature-extraction": GGUFEmbeddingFeatureExtractionTask, "gguf-audio-projector": GGUFAudioProjectorTask, "masked-diffusion": MaskedDiffusionTask, + "mattergen-score": MatterGenScoreTask, "minimax-music3-condition": MiniMaxMusic3ConditionTask, "minimax-music3-denoising": MiniMaxMusic3DenoisingTask, "minimax-music3-language": MiniMaxMusic3LanguageTask, diff --git a/src/mobius/tasks/_mattergen.py b/src/mobius/tasks/_mattergen.py new file mode 100644 index 000000000..3e97739e1 --- /dev/null +++ b/src/mobius/tasks/_mattergen.py @@ -0,0 +1,208 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +"""Task wiring for the host-orchestrated MatterGen GemNet-T score core.""" + +from __future__ import annotations + +import json +from typing import ClassVar + +import onnx_ir as ir +from onnxscript import nn + +from mobius._export_report import ComponentExportDisposition, ComponentExportReport +from mobius._model_package import ModelPackage +from mobius.integrations.mattergen._configs import MatterGenConfig +from mobius.integrations.mattergen._contract import HOST_OWNED_STEPS +from mobius.integrations.mattergen._contract import ( + MAX_ATOMS, + OFFICIAL_CHECKPOINT_CONDITIONS, + SELECTED_ATOMIC_NUMBERS, +) +from mobius.tasks._base import ModelTask, _make_graph, _make_model + + +class MatterGenScoreTask(ModelTask): + """Build the pure neural score stage of the MatterGen crystal generator. + + The host must rebuild MatterGen's source-ordered periodic graph before every + invocation. Inputs include its ragged graph and triplet tensors rather + than fractional positions/cells, because data-dependent periodic image + enumeration and neighbor sorting cannot be represented faithfully by a + portable dynamic-shape ONNX graph. + """ + + model_roles: ClassVar[dict[str, str]] = {"model": "encoder"} + + def build(self, module: nn.Module, config: MatterGenConfig) -> ModelPackage: + """Build an encoder-role ONNX score graph with config-specific condition ports.""" + if not isinstance(config, MatterGenConfig): + raise TypeError("MatterGenScoreTask requires a MatterGenConfig.") + graph, builder = _make_graph("mattergen_score") + atoms = "atoms" + crystals = "crystals" + edges = "edges" + triplets = "triplets" + + atomic_numbers = builder.input( + "atomic_numbers", dtype=ir.DataType.INT64, shape=[atoms] + ) + batch = builder.input("batch", dtype=ir.DataType.INT64, shape=[atoms]) + # The source basis and fitted scalar encoders are always float32, even + # when a future export introduces a separately assessed compute dtype. + timestep = builder.input("timestep", dtype=ir.DataType.FLOAT, shape=[crystals]) + edge_index = builder.input( + "edge_index", dtype=ir.DataType.INT64, shape=[2, edges] + ) + edge_distance = builder.input( + "edge_distance", dtype=ir.DataType.FLOAT, shape=[edges] + ) + edge_direction = builder.input( + "edge_direction", dtype=ir.DataType.FLOAT, shape=[edges, 3] + ) + edge_lattice_cosines = builder.input( + "edge_lattice_cosines", dtype=ir.DataType.FLOAT, shape=[edges, 3] + ) + id_swap = builder.input("id_swap", dtype=ir.DataType.INT64, shape=[edges]) + id3_ba = builder.input("id3_ba", dtype=ir.DataType.INT64, shape=[triplets]) + id3_ca = builder.input("id3_ca", dtype=ir.DataType.INT64, shape=[triplets]) + id3_ragged_idx = builder.input( + "id3_ragged_idx", dtype=ir.DataType.INT64, shape=[triplets] + ) + + condition_values: dict[str, ir.Value] = {} + condition_masks: dict[str, ir.Value] = {} + for spec in config.condition_input_specs: + condition_values[spec.name] = builder.input( + f"condition.{spec.name}", + dtype=( + ir.DataType.INT64 + if spec.kind == "space_group_index" + else ir.DataType.FLOAT + ), + shape=[crystals, *spec.input_shape_suffix], + ) + condition_masks[spec.name] = builder.input( + f"condition.{spec.name}.use_unconditional", + dtype=ir.DataType.BOOL, + shape=[crystals], + ) + + atom_logits, coordinate_score, lattice_score, energy = module( + builder.op, + atomic_numbers, + batch, + timestep, + edge_index, + edge_distance, + edge_direction, + edge_lattice_cosines, + id_swap, + id3_ba, + id3_ca, + id3_ragged_idx, + condition_values, + condition_masks, + ) + atom_logits.shape = ir.Shape([atoms, 101]) + coordinate_score.shape = ir.Shape([atoms, 3]) + lattice_score.shape = ir.Shape([crystals, 3, 3]) + energy.shape = ir.Shape([crystals, 1]) + builder.add_output(atom_logits, "atom_logits") + builder.add_output(coordinate_score, "coordinate_score") + builder.add_output(lattice_score, "lattice_score") + # Denoiser inference does not consume GemNet's energy result. Exposing + # it as a diagnostic output keeps every trained OutputBlock parameter + # reachable and makes strict checkpoint routing auditable. + builder.add_output(energy, "energy") + + model = _make_model(graph) + model.metadata_props.update( + { + "mobius.model_type": "mattergen", + "mobius.source_model": config.model_id, + "mobius.source_revision": config.revision, + "mobius.source_commit": config.source_commit, + "mobius.task": "mattergen-score", + "mobius.checkpoint_family": config.variant, + "mobius.max_atoms": str(MAX_ATOMS), + "mobius.sampling_atomic_numbers": json.dumps(SELECTED_ATOMIC_NUMBERS), + "mobius.official_checkpoint_conditions": json.dumps( + OFFICIAL_CHECKPOINT_CONDITIONS, + sort_keys=True, + ), + "mobius.coordinate_convention": ( + "host uses row-vector cells: cartesian = fractional @ cell; " + "coordinate_score is Cartesian" + ), + "mobius.periodic_graph_abi": ( + "host supplies source-ordered edge_index, edge_distance, " + "V_st edge_direction, edge_lattice_cosines, id_swap, and triplet ids" + ), + "mobius.host_orchestration": json.dumps(HOST_OWNED_STEPS), + "mobius.runtime_support": ( + "ONNX score core only; ONNX Runtime GenAI cannot execute MatterGen's " + "periodic graph construction, stochastic scheduler, or crystal validation." + ), + } + ) + package = ModelPackage({"model": model}, config=config) + package.export_report = ComponentExportReport.create( + ( + ComponentExportDisposition( + name="crystal_validation", + route="MatterGen host postprocessing", + requested=True, + discovered=True, + support="deferred", + output="omitted", + blocker_category="host-scientific-runtime", + reason="Pymatgen Structure/CIF validation is not a neural ONNX operation.", + impact="The ONNX package cannot itself claim to generate a valid crystal artifact.", + remediation="Validate final wrapped fractional coordinates and cell in a host runtime.", + ), + ComponentExportDisposition( + name="periodic_graph", + route="MatterGen host preprocessing", + requested=True, + discovered=True, + support="deferred", + output="omitted", + blocker_category="dynamic-ragged-pbc", + reason=( + "Source-faithful periodic image enumeration, neighbor sorting, " + "symmetric reordering, and sparse triplets are data-dependent." + ), + impact="The score graph requires the documented host graph ABI per evaluation.", + remediation="Build graph tensors with the pinned MatterGen v1.0.3 semantics.", + ), + ComponentExportDisposition( + name="sampling_scheduler", + route="MatterGen host sampling", + requested=True, + discovered=True, + support="deferred", + output="omitted", + blocker_category="stochastic-host-loop", + reason=( + "D3PM, wrapped VE/VP updates, RNG, classifier-free guidance, " + "and lattice projection are source host-loop behavior." + ), + impact="This package is not an end-to-end crystal generator.", + remediation="Run the pinned-source scheduler around repeated score-core calls.", + ), + ComponentExportDisposition( + name="score_core", + route="standard ONNX GemNet-T encoder graph", + requested=True, + discovered=True, + support="supported", + output="exported", + runtime_validation_status="validated", + evidence_id="mattergen-score-core-ort", + ), + ), + end_to_end_runnable=False, + ) + return package diff --git a/testdata/golden/diffusion/mattergen-mp20-host-sample.json b/testdata/golden/diffusion/mattergen-mp20-host-sample.json new file mode 100644 index 000000000..da88a849c --- /dev/null +++ b/testdata/golden/diffusion/mattergen-mp20-host-sample.json @@ -0,0 +1,42 @@ +{ + "checkpoint_family": "mp_20_base", + "checkpoint_sha256": "ffb80e4425a6f99f479a67b8cd111885d45117234e8947ff77eb3a55df420b9a", + "host_runtime": "MatterGenHostSampler", + "hub_revision": "5244495dd9a979ff71abc7548a0b14b9deb0069a", + "num_atoms": [ + 1 + ], + "sample": { + "atomic_numbers": [ + 80 + ], + "cell": [ + [ + 3.4113781452178955, + 1.3116661310195923, + -1.001381278038025 + ], + [ + 1.3116661310195923, + 4.735863208770752, + 0.005625350400805473 + ], + [ + -1.001381278038025, + 0.0056253462098538876, + 3.119309902191162 + ] + ], + "fractional_coordinates": [ + [ + 0.9089621305465698, + 0.9398096203804016, + 0.44816485047340393 + ] + ], + "volume": 40.26449966430664 + }, + "seed": 814, + "source_commit": "842ffe735f7d06cec89d56aa23d9f001e1124b30", + "timesteps": 1000 +} diff --git a/tests/cli_test.py b/tests/cli_test.py index 4281ca7ac..4294a4d08 100644 --- a/tests/cli_test.py +++ b/tests/cli_test.py @@ -99,6 +99,73 @@ def test_build_with_dtype(self): ) assert os.path.isfile(os.path.join(tmpdir, "model.onnx")) + def test_mattergen_no_weights_routes_before_diffusers_or_transformers(self): + with ( + tempfile.TemporaryDirectory() as tmpdir, + mock.patch( + "mobius.integrations.mattergen._builder.build_mattergen", + return_value=mock.MagicMock(), + ) as build_mattergen, + mock.patch("mobius.__main__._save_package") as save_package, + ): + main( + [ + "build", + "--model", + "microsoft/mattergen", + "--mattergen-checkpoint", + "mp_20_base", + "--no-weights", + "--output", + tmpdir, + ] + ) + + assert build_mattergen.call_args.args == ("microsoft/mattergen",) + assert build_mattergen.call_args.kwargs == { + "checkpoint": "mp_20_base", + "revision": None, + "dtype": None, + "load_weights": False, + "execution_provider": "default", + } + save_package.assert_called_once() + + def test_mattergen_rejects_onnx_genai_metadata(self): + with ( + tempfile.TemporaryDirectory() as tmpdir, + pytest.raises(SystemExit, match="host-owned"), + ): + main( + [ + "build", + "--model", + "microsoft/mattergen", + "--runtime", + "onnx-genai", + "--no-weights", + "--output", + tmpdir, + ] + ) + + def test_mattergen_rejects_transformer_rewrite_rules(self): + with ( + tempfile.TemporaryDirectory() as tmpdir, + pytest.raises(SystemExit, match="rewrite rules"), + ): + main( + [ + "build", + "--model", + "microsoft/mattergen", + "--optimize", + "--no-weights", + "--output", + tmpdir, + ] + ) + def test_max_workers_defaults_to_eight(self): with ( tempfile.TemporaryDirectory() as tmpdir, diff --git a/tests/integration/mattergen_parity_test.py b/tests/integration/mattergen_parity_test.py index 8a61945e0..b568d9d95 100644 --- a/tests/integration/mattergen_parity_test.py +++ b/tests/integration/mattergen_parity_test.py @@ -42,7 +42,13 @@ from mobius import build_from_module from mobius._testing.ort_inference import OnnxModelSession -from mobius.integrations.mattergen import MatterGenConfig, MatterGenModel +from mobius.integrations.mattergen import ( + MatterGenConfig, + MatterGenHostSampler, + MatterGenModel, + build_periodic_graph, + create_onnxruntime_score_callback, +) from mobius.integrations.mattergen._configs import MATTERGEN_SOURCE_COMMIT from mobius.integrations.mattergen._weights import apply_mattergen_checkpoint from mobius.tasks import MatterGenScoreTask @@ -59,6 +65,13 @@ / "diffusion" / "mattergen-mp20-score.json" ) +_HOST_SAMPLE_GOLDEN_PATH = ( + Path(__file__).parents[2] + / "testdata" + / "golden" + / "diffusion" + / "mattergen-mp20-host-sample.json" +) _RTOL = 1e-3 _ATOL = 1e-3 # The fine-tuned adapter checkpoint amplifies sub-ULP differences between @@ -77,6 +90,9 @@ class _SourceModules: atom_embedding: Any model_utils: Any property_embeddings: Any + d3pm: Any + d3pm_corruption: Any + d3pm_predictors_correctors: Any @dataclass(frozen=True) @@ -125,6 +141,7 @@ class _Mp20Runtime: original: dict[float, _ScoreCase] translated: _ScoreCase permuted: _ScoreCase + batched_graph: _SourceGraph checkpoint_sha256: str def close(self) -> None: @@ -264,8 +281,9 @@ def __init__( value: torch.Tensor, sparse_sizes: tuple[torch.Tensor, torch.Tensor] | tuple[int, int], ): - del col, sparse_sizes + del sparse_sizes self._row = row + self._col = col self._value = value self.storage = _SparseStorage(row.new_empty(0), value.new_empty(0)) @@ -274,6 +292,10 @@ def __getitem__(self, queried_rows: torch.Tensor) -> SparseTensor: result_values: list[torch.Tensor] = [] for output_row, source_row in enumerate(queried_rows): match = torch.nonzero(self._row == source_row, as_tuple=False).squeeze(1) + # torch_sparse stores COO entries in row/column order. + # MatterGen's id3_ba therefore follows the edge source ID + # within each target row, not the caller's incidental COO order. + match = match[torch.argsort(self._col[match], stable=True)] result_rows.append( torch.full( (len(match),), @@ -285,6 +307,7 @@ def __getitem__(self, queried_rows: torch.Tensor) -> SparseTensor: result_values.append(self._value[match]) result = object.__new__(SparseTensor) result._row = self._row + result._col = self._col result._value = self._value result.storage = _SparseStorage( torch.cat(result_rows), @@ -426,6 +449,13 @@ def _pinned_source_modules(source_dir: Path) -> Iterator[_SourceModules]: ), model_utils=importlib.import_module("mattergen.diffusion.model_utils"), property_embeddings=importlib.import_module("mattergen.property_embeddings"), + d3pm=importlib.import_module("mattergen.diffusion.d3pm.d3pm"), + d3pm_corruption=importlib.import_module( + "mattergen.diffusion.corruption.d3pm_corruption" + ), + d3pm_predictors_correctors=importlib.import_module( + "mattergen.diffusion.d3pm.d3pm_predictors_correctors" + ), ) finally: sys.path.remove(str(source_dir)) @@ -656,7 +686,29 @@ def _source_host_feeds( fractional_coordinates = torch.from_numpy(crystal.fractional_coordinates) lattice = torch.from_numpy(crystal.cell) num_atoms = torch.tensor([len(atomic_numbers)], dtype=torch.long) - batch = torch.zeros(len(atomic_numbers), dtype=torch.long) + return _source_host_feeds_for_batch( + reference, + source, + atomic_numbers=atomic_numbers, + fractional_coordinates=fractional_coordinates, + lattice=lattice, + num_atoms=num_atoms, + timestep=timestep, + ) + + +def _source_host_feeds_for_batch( + reference: _ReferenceModel, + source: _SourceModules, + *, + atomic_numbers: torch.Tensor, + fractional_coordinates: torch.Tensor, + lattice: torch.Tensor, + num_atoms: torch.Tensor, + timestep: float, +) -> _SourceGraph: + """Build one exact source graph for arbitrary packed crystals.""" + batch = torch.repeat_interleave(torch.arange(len(num_atoms), dtype=torch.long), num_atoms) cartesian_coordinates = source.data_utils.frac_to_cart_coords_with_lattice( fractional_coordinates, num_atoms, lattice ) @@ -792,6 +844,26 @@ def mp20_runtime() -> Iterator[_Mp20Runtime]: _crystal(permutation=np.array([1, 0], dtype=np.int64)), 0.25, ) + batched_graph = _source_host_feeds_for_batch( + reference, + source, + atomic_numbers=torch.tensor([3, 8, 14], dtype=torch.long), + fractional_coordinates=torch.tensor( + [[0.10, 0.15, 0.20], [0.40, 0.45, 0.35], [0.25, 0.75, 0.50]], + dtype=torch.float32, + ), + lattice=torch.stack( + [ + torch.diag(torch.tensor([6.0, 5.5, 6.5], dtype=torch.float32)), + torch.tensor( + [[4.5, 0.0, 0.0], [0.4, 5.0, 0.0], [0.2, 0.3, 5.5]], + dtype=torch.float32, + ), + ] + ), + num_atoms=torch.tensor([2, 1], dtype=torch.long), + timestep=0.25, + ) del reference del state @@ -812,6 +884,7 @@ def mp20_runtime() -> Iterator[_Mp20Runtime]: original=original, translated=translated, permuted=permuted, + batched_graph=batched_graph, checkpoint_sha256=_sha256(checkpoint), ) try: @@ -878,6 +951,149 @@ def test_mp20_periodic_translation_and_permutation_invariance( ) +def test_host_periodic_graph_matches_source_for_batched_neighbor_truncation( + mp20_runtime: _Mp20Runtime, +) -> None: + """Host PBC edges, symmetric pairs, and triplets match source for packed crystals.""" + expected = mp20_runtime.batched_graph.feeds + graph = build_periodic_graph( + torch.from_numpy( + np.array( + [[0.10, 0.15, 0.20], [0.40, 0.45, 0.35], [0.25, 0.75, 0.50]], + dtype=np.float32, + ) + ), + torch.from_numpy( + np.array( + [ + [[6.0, 0.0, 0.0], [0.0, 5.5, 0.0], [0.0, 0.0, 6.5]], + [[4.5, 0.0, 0.0], [0.4, 5.0, 0.0], [0.2, 0.3, 5.5]], + ], + dtype=np.float32, + ) + ), + torch.tensor([2, 1], dtype=torch.long), + cutoff=7.0, + max_neighbors=50, + max_cell_images_per_dim=5, + ) + actual = { + "edge_index": graph.edge_index.numpy(), + "edge_distance": graph.edge_distance.numpy(), + "edge_direction": graph.edge_direction.numpy(), + "edge_lattice_cosines": graph.edge_lattice_cosines.numpy(), + "id_swap": graph.id_swap.numpy(), + "id3_ba": graph.id3_ba.numpy(), + "id3_ca": graph.id3_ca.numpy(), + "id3_ragged_idx": graph.id3_ragged_idx.numpy(), + } + for name, value in actual.items(): + np.testing.assert_array_equal( + value, + expected[name], + err_msg=f"{name} source graph mismatch", + ) + + +def test_host_d3pm_predictor_matches_source_schedule_and_rng() -> None: + """Host absorbing-mask posterior and both categorical draws match the source.""" + source_dir = _required_artifact(_SOURCE_DIR_ENV) + with _pinned_source_modules(source_dir) as source: + schedule = source.d3pm.create_discrete_diffusion_schedule( + kind="standard", + num_steps=1000, + ) + corruption = source.d3pm_corruption.D3PMCorruption( + source.d3pm.MaskDiffusion(dim=101, schedule=schedule), + offset=1, + ) + predictor = source.d3pm_predictors_correctors.D3PMAncestralSamplingPredictor( + corruption=corruption, + score_fn=None, + predict_x0=True, + ) + atomic_numbers = torch.tensor([101, 3, 101], dtype=torch.long) + logits = torch.linspace(-2.0, 2.0, 303, dtype=torch.float32).reshape(3, 101) + timestep = torch.tensor([0.75], dtype=torch.float32) + batch = torch.zeros(3, dtype=torch.long) + rng_state = torch.random.get_rng_state() + try: + torch.manual_seed(814) + expected_sample, expected_mean = predictor.update_given_score( + x=atomic_numbers, + t=timestep, + dt=torch.tensor(-0.001, dtype=torch.float32), + batch_idx=batch, + score=logits, + batch=None, + ) + finally: + torch.random.set_rng_state(rng_state) + + generator = torch.Generator(device="cpu") + generator.manual_seed(814) + actual_sample, actual_mean = MatterGenHostSampler( + lambda _inputs: (_ for _ in ()).throw( + AssertionError("score callback must not be used") + ) + )._d3pm_ancestral(atomic_numbers, logits, timestep, batch, generator) + + torch.testing.assert_close(actual_sample, expected_sample) + torch.testing.assert_close(actual_mean, expected_mean) + + +@pytest.mark.golden +@pytest.mark.generation +def test_mp20_real_onnx_host_sampling_golden(mp20_runtime: _Mp20Runtime) -> None: + """L5: the full released scheduler yields a deterministic valid crystal artifact.""" + + class _NamedSession: + """Bridge test inference wrapper to the public ONNX Runtime callback ABI.""" + + def run( + self, + output_names: list[str] | None, + input_feed: Mapping[str, np.ndarray], + ) -> list[np.ndarray]: + if output_names is None: + raise AssertionError("MatterGen callback requests explicit output names") + outputs = mp20_runtime.session.run(dict(input_feed)) + return [outputs[name] for name in output_names] + + golden = json.loads(_HOST_SAMPLE_GOLDEN_PATH.read_text(encoding="utf-8")) + assert golden["source_commit"] == MATTERGEN_SOURCE_COMMIT + assert golden["checkpoint_sha256"] == mp20_runtime.checkpoint_sha256 + assert golden["timesteps"] == 1000 + + sample = MatterGenHostSampler(create_onnxruntime_score_callback(_NamedSession())).sample( + torch.tensor(golden["num_atoms"], dtype=torch.long), + seed=golden["seed"], + ) + crystal = sample.crystals()[0] + np.testing.assert_array_equal( + crystal.atomic_numbers.numpy(), + np.asarray(golden["sample"]["atomic_numbers"], dtype=np.int64), + ) + np.testing.assert_allclose( + crystal.fractional_coordinates.numpy(), + np.asarray(golden["sample"]["fractional_coordinates"], dtype=np.float32), + rtol=1e-4, + atol=1e-4, + ) + np.testing.assert_allclose( + crystal.cell.numpy(), + np.asarray(golden["sample"]["cell"], dtype=np.float32), + rtol=1e-4, + atol=1e-4, + ) + assert np.isclose( + np.linalg.det(crystal.cell.numpy()), + golden["sample"]["volume"], + rtol=1e-4, + atol=1e-4, + ) + + @pytest.mark.golden def test_mp20_one_step_source_golden(mp20_runtime: _Mp20Runtime) -> None: """L4: compare a one-step real ``mp_20_base`` source output with committed provenance."""