diff --git a/deepmd/common.py b/deepmd/common.py index 98cf2461bd..fea637328a 100644 --- a/deepmd/common.py +++ b/deepmd/common.py @@ -53,6 +53,7 @@ "tanh", "gelu", "gelu_tf", + "gelu_erf", "silu", "silut", "none", diff --git a/deepmd/dpmodel/array_api.py b/deepmd/dpmodel/array_api.py index 6d3c89d584..5c6822141d 100644 --- a/deepmd/dpmodel/array_api.py +++ b/deepmd/dpmodel/array_api.py @@ -516,6 +516,50 @@ def xp_sigmoid(x: Array) -> Array: return 1 / (1 + xp.exp(-x)) +def xp_erf(x: Array) -> Array: + """Compute the error function. + + Used by the exact (non-approximated) GELU. The array API has no ``erf``, so + each backend's own implementation is used; NumPy goes through SciPy, which + is already a core dependency. + """ + if array_api_compat.is_jax_array(x): + from deepmd.jax.env import ( + jax, + ) + + return jax.scipy.special.erf(x) + elif array_api_compat.is_torch_array(x): + import torch + + return torch.special.erf(x) + + xp = array_api_compat.array_namespace(x) + if getattr(xp, "__name__", "") == "deepmd._vendors.ndtensorflow": + import tensorflow as tf + + # The NumPy round-trip below cannot serve TensorFlow. Under + # ``tf.function`` the conversion is refused outright, and in eager mode + # it detaches the erf factor from the tape, which leaves the exact GELU + # differentiating to Phi(x) alone -- silently, and for every backend + # user of ``gelu_erf`` rather than only Uni-Mol. + # + # Imported directly rather than through ``deepmd.tf2.env`` for symmetry + # with the JAX branch above: that module raises at import time unless + # eager execution is on, and this branch has to work inside + # ``tf.function``, where it is not. + return xp.asarray(tf.math.erf(x.unwrap())) + + from scipy.special import ( + erf, + ) + + if array_api_compat.is_numpy_array(x): + return erf(x) + # array-api-strict and friends: round-trip through NumPy. + return xp.asarray(erf(np.asarray(x)), dtype=x.dtype) + + def xp_setitem_at(x: Array, mask: Array, values: Array) -> Array: """Set items at boolean mask indices. diff --git a/deepmd/dpmodel/atomic_model/__init__.py b/deepmd/dpmodel/atomic_model/__init__.py index 4d882d5e4b..c55eead5f3 100644 --- a/deepmd/dpmodel/atomic_model/__init__.py +++ b/deepmd/dpmodel/atomic_model/__init__.py @@ -45,6 +45,9 @@ from .property_atomic_model import ( DPPropertyAtomicModel, ) +from .unimol_atomic_model import ( + DPUniMolAtomicModel, +) __all__ = [ "BaseAtomicModel", @@ -54,6 +57,7 @@ "DPEnergyAtomicModel", "DPPolarAtomicModel", "DPPropertyAtomicModel", + "DPUniMolAtomicModel", "DPZBLLinearEnergyAtomicModel", "LinearEnergyAtomicModel", "PairTabAtomicModel", diff --git a/deepmd/dpmodel/atomic_model/unimol_atomic_model.py b/deepmd/dpmodel/atomic_model/unimol_atomic_model.py new file mode 100644 index 0000000000..f97ff40754 --- /dev/null +++ b/deepmd/dpmodel/atomic_model/unimol_atomic_model.py @@ -0,0 +1,106 @@ +# SPDX-License-Identifier: LGPL-3.0-or-later +"""Atomic model for Uni-Mol v1 self-supervised pretraining.""" + +from typing import ( + Any, +) + +from deepmd.dpmodel.array_api import ( + Array, +) +from deepmd.dpmodel.descriptor.unimol import ( + DescrptUniMol, +) +from deepmd.dpmodel.fitting.unimol_pretrain import ( + UniMolPretrainFitting, +) + +from .dp_atomic_model import ( + DPAtomicModel, +) + + +class DPUniMolAtomicModel(DPAtomicModel): + r"""Uni-Mol pretraining, wired at token resolution. + + The standard path hands a descriptor's five-tuple to a fitting, which + cannot carry the two virtual tokens, the pair channel or the norm + regularisers that these heads read. This model therefore overrides one + method to route the backbone's token-resolution output straight into the + heads. Nothing else about the atomic model changes. + """ + + def __init__( + self, descriptor: Any, fitting: Any, type_map: list[str], **kwargs: Any + ) -> None: + if not isinstance(descriptor, DescrptUniMol): + raise TypeError( + "DPUniMolAtomicModel needs the unimol descriptor, which is the only " + "one producing a Uni-Mol token sequence" + ) + if not isinstance(fitting, UniMolPretrainFitting): + raise TypeError("DPUniMolAtomicModel needs the unimol_pretrain fitting") + # The objective compares against distances whose virtual tokens sit at + # the origin, which is where upstream puts them and, once the transform + # has centred a frame, where the clean centroid is. Under "centroid" the + # descriptor instead places them at the centroid of the coordinates it + # is handed -- the corrupted ones -- so the two virtual columns of every + # corrupted row would be regressed against a label for a different + # position, by about the size of the noise. Centring costs nothing here + # because the transform always centres, so the only effect would be that + # silent mismatch. + if getattr(descriptor, "virtual_token_position", "origin") != "origin": + raise ValueError( + "unimol pretraining needs virtual_token_position='origin': the " + "distance target places the virtual tokens at the origin, and " + f"this descriptor places them at the " + f"{descriptor.virtual_token_position}, so the two virtual " + "columns would train against the wrong label. The corruption " + "centres every frame, so 'origin' is the centroid anyway" + ) + super().__init__(descriptor, fitting, type_map, **kwargs) + + def forward_atomic( + self, + extended_coord: Array, + extended_atype: Array, + nlist: Array, + mapping: Array | None = None, + fparam: Array | None = None, + aparam: Array | None = None, + comm_dict: dict | None = None, + charge_spin: Array | None = None, + ) -> dict[str, Array]: + """Run the backbone and its heads at token resolution. + + Parameters + ---------- + extended_coord + nf x (nall x 3) coordinates; the descriptor rejects any frame that + carries periodic images. + extended_atype + nf x nall element types, already clamped to be non-negative. + nlist + nf x nloc x nnei neighbour list, which is how real atoms are told + apart from padding. + mapping, fparam, aparam, comm_dict, charge_spin + Unused by this model. + + Returns + ------- + dict + The three head outputs plus the two norm regularisers. + """ + del mapping, fparam, aparam, comm_dict, charge_spin + backbone = self.descriptor.forward_tokens(extended_coord, extended_atype, nlist) + return self.fitting_net.call_tokens(backbone) + + def apply_out_stat(self, ret: dict[str, Array], atype: Array) -> dict[str, Array]: + """Return the head outputs untouched. + + Self-supervised targets carry no per-element bias to add back: the + element head predicts a distribution, and the coordinate and distance + heads predict geometry the data already fixes. + """ + del atype + return ret diff --git a/deepmd/dpmodel/descriptor/__init__.py b/deepmd/dpmodel/descriptor/__init__.py index ae9cc66e39..9ba61078b0 100644 --- a/deepmd/dpmodel/descriptor/__init__.py +++ b/deepmd/dpmodel/descriptor/__init__.py @@ -35,6 +35,9 @@ from .se_t_tebd import ( DescrptSeTTebd, ) +from .unimol import ( + DescrptUniMol, +) __all__ = [ "DescrptDPA1", @@ -48,5 +51,6 @@ "DescrptSeR", "DescrptSeT", "DescrptSeTTebd", + "DescrptUniMol", "make_base_descriptor", ] diff --git a/deepmd/dpmodel/descriptor/unimol.py b/deepmd/dpmodel/descriptor/unimol.py new file mode 100644 index 0000000000..8314325871 --- /dev/null +++ b/deepmd/dpmodel/descriptor/unimol.py @@ -0,0 +1,529 @@ +# SPDX-License-Identifier: LGPL-3.0-or-later +"""The Uni-Mol v1 backbone as a deepmd descriptor. + +Wraps the ported Uni-Mol encoder (:mod:`deepmd.dpmodel.descriptor.unimol_nn`) in +the descriptor interface. Unlike every other descriptor here, Uni-Mol is global +rather than local: it attends over all atom pairs with no cut-off and no smooth +envelope, so it is neither extensive nor periodic and its forces are not +conserved. It is meant for molecular property and pretraining work. + +Nothing stops a configuration from pairing it with an energy fitting, and +nothing could usefully: the descriptor does not see the fitting. Such a model +would train and evaluate, but its forces would be neither smooth at any cutoff +nor conserved, and its energy would not be extensive, so it should not be used +as a potential energy surface. + +Uni-Mol's own vocabulary is kept, because the released weights are indexed by +it: four special tokens, then 26 elements, then ``[MASK]``. A deepmd +``type_map`` is mapped onto those ids, and an element outside the vocabulary +becomes ``[UNK]``. +""" + +import array_api_compat + +from deepmd.dpmodel.array_api import ( + Array, +) +from deepmd.dpmodel.common import ( + NativeOP, + cast_precision, +) +from deepmd.dpmodel.descriptor.base_descriptor import ( + BaseDescriptor, +) +from deepmd.dpmodel.descriptor.unimol_nn import ( + GaussianLayer, + NonLinearHead, + TransformerEncoderWithPair, +) +from deepmd.dpmodel.utils.network import ( + NativeLayer, +) +from deepmd.dpmodel.utils.seed import ( + child_seed, +) +from deepmd.utils.version import ( + check_version_compatibility, +) + +# Uni-Mol's molecular dictionary, in file order. The order is load-bearing: it +# fixes the embedding rows, the edge-type ids (t_i * V + t_j) and the output +# width of the element head in the released checkpoints. +UNIMOL_SPECIAL_TOKENS = ["[PAD]", "[CLS]", "[SEP]", "[UNK]"] +UNIMOL_ELEMENTS = [ + "C", "N", "O", "S", "H", "Cl", "F", "Br", "I", "Si", "P", "B", "Na", "K", + "Al", "Ca", "Sn", "As", "Hg", "Fe", "Zn", "Cr", "Se", "Gd", "Au", "Li", +] # fmt: skip +UNIMOL_MASK_TOKEN = "[MASK]" + + +def unimol_vocabulary() -> list[str]: + """Return the 31 Uni-Mol tokens in checkpoint order.""" + return [*UNIMOL_SPECIAL_TOKENS, *UNIMOL_ELEMENTS, UNIMOL_MASK_TOKEN] + + +@BaseDescriptor.register("unimol") +class DescrptUniMol(NativeOP, BaseDescriptor): + r"""Uni-Mol v1 transformer backbone. + + Every atom attends to every other atom; geometry enters only through + pairwise distances, expanded in a Gaussian basis whose affine parameters are + specific to the ordered element pair, and injected as a per-head attention + bias. Two virtual tokens wrap each molecule, and the running sum of + per-layer attention logits is the pair representation that the pretraining + heads read. + + Parameters + ---------- + type_map : list[str] + Element names, mapped onto the Uni-Mol vocabulary. + encoder_layers : int + Number of transformer blocks. + encoder_embed_dim : int + Width of the node representation. + encoder_ffn_embed_dim : int + Width of the feed-forward hidden layer. + encoder_attention_heads : int + Number of attention heads, which is also the width of the pair channel. + max_atoms : int + Largest molecule accepted, which fixes ``sel``. + max_seq_len : int + Upstream's sequence guard. Nothing in the forward pass consults it, but + :meth:`get_rcut` reports a radius derived from it, so it is not inert. + activation_function : str + Activation of the blocks and heads. Uni-Mol uses the exact GELU. + dropout, emb_dropout, attention_dropout, activation_dropout : float + Dropout rates. They are applied here, by the shared encoder, when the + arrays are torch tensors and the module is in training mode. Inference + is the identity, and training on any other array namespace raises + ``NotImplementedError`` rather than silently dropping the + regularisation (see ``unimol_nn.encoder.dropout``). + no_final_head_layer_norm : bool + Skip the layer norm on the pair delta. Upstream builds that norm unless + its weight is negative. + single_precision_basis : bool + Evaluate the Gaussian basis in fp32, which is what upstream does. + single_precision_distance : bool + Round the pairwise distances to fp32 before the Gaussian basis. Upstream + precomputes its distance matrix in fp32 in the data pipeline, so this + reproduces its numbers; the default computes them in the working + precision, which is more accurate and is what gradients flow through. + virtual_token_position : str + Where the two virtual tokens sit. ``"centroid"`` places them at the + centroid of the real atoms, which keeps the sequence translation + invariant. ``"origin"`` places them at the origin, reproducing upstream + exactly for data that its own pipeline has already centred. + precision : str + Floating-point precision of the parameters. + seed : int, optional + Random seed for initialization. + """ + + def __init__( + self, + type_map: list[str], + encoder_layers: int = 15, + encoder_embed_dim: int = 512, + encoder_ffn_embed_dim: int = 2048, + encoder_attention_heads: int = 64, + max_atoms: int = 256, + max_seq_len: int = 512, + activation_function: str = "gelu_erf", + dropout: float = 0.1, + emb_dropout: float = 0.1, + attention_dropout: float = 0.1, + activation_dropout: float = 0.0, + no_final_head_layer_norm: bool = False, + single_precision_basis: bool = True, + single_precision_distance: bool = False, + virtual_token_position: str = "centroid", + gaussian_kernels: int = 128, + precision: str = "float64", + seed: int | list[int] | None = None, + type_map_tokens: list[str] | None = None, + **kwargs: float | str | bool, + ) -> None: + del kwargs + self.type_map = list(type_map) + self.encoder_layers = encoder_layers + self.encoder_embed_dim = encoder_embed_dim + self.encoder_ffn_embed_dim = encoder_ffn_embed_dim + self.encoder_attention_heads = encoder_attention_heads + self.max_atoms = max_atoms + self.max_seq_len = max_seq_len + self.activation_function = activation_function + self.dropout = dropout + self.emb_dropout = emb_dropout + self.attention_dropout = attention_dropout + self.activation_dropout = activation_dropout + self.no_final_head_layer_norm = no_final_head_layer_norm + self.single_precision_basis = single_precision_basis + self.single_precision_distance = single_precision_distance + if virtual_token_position not in ("centroid", "origin"): + raise ValueError( + f"virtual_token_position must be centroid or origin, got {virtual_token_position}" + ) + self.virtual_token_position = virtual_token_position + self.gaussian_kernels = gaussian_kernels + self.precision = precision + self.vocabulary = type_map_tokens or unimol_vocabulary() + self.ntokens = len(self.vocabulary) + + import numpy as np + + token_of = {sym: i for i, sym in enumerate(self.vocabulary)} + self.pad_idx = token_of["[PAD]"] + self.bos_idx = token_of["[CLS]"] + self.eos_idx = token_of["[SEP]"] + self.unk_idx = token_of["[UNK]"] + # type id -> Uni-Mol token id; unknown elements fall back to [UNK]. + self.type_to_token = np.array( + [token_of.get(sym, self.unk_idx) for sym in self.type_map], dtype=np.int64 + ) + + # The token embedding is trained, so it is a layer rather than a bare + # array: a bare array becomes a buffer on the torch backends, and a + # buffer never receives a gradient. + self.embed_tokens = NativeLayer( + self.ntokens, + encoder_embed_dim, + bias=False, + precision=precision, + seed=child_seed(seed, 0), + ) + self.gbf = GaussianLayer( + gaussian_kernels, + self.ntokens**2, + single_precision_basis=single_precision_basis, + precision=precision, + seed=child_seed(seed, 1), + ) + self.gbf_proj = NonLinearHead( + gaussian_kernels, + encoder_attention_heads, + activation_function, + precision=precision, + seed=child_seed(seed, 2), + ) + self.encoder = TransformerEncoderWithPair( + encoder_layers=encoder_layers, + embed_dim=encoder_embed_dim, + ffn_embed_dim=encoder_ffn_embed_dim, + attention_heads=encoder_attention_heads, + emb_dropout=emb_dropout, + dropout=dropout, + attention_dropout=attention_dropout, + activation_dropout=activation_dropout, + max_seq_len=max_seq_len, + activation_function=activation_function, + no_final_head_layer_norm=no_final_head_layer_norm, + precision=precision, + seed=child_seed(seed, 3), + ) + + # ------------------------------------------------------------------ + # capability queries + # ------------------------------------------------------------------ + def get_rcut(self) -> float: + """All pairs are neighbours, so the radius is effectively unbounded. + + The number is derived from ``max_seq_len`` only to scale with the + largest sequence configured; nothing compares against it as a real + cut-off, because this descriptor has none. + """ + return float(self.max_seq_len) * 1e3 + + def get_rcut_smth(self) -> float: + """No smoothing exists; smoothing starts where the cut-off is.""" + return self.get_rcut() + + def get_sel(self) -> list[int]: + """One entry, because the neighbour list is type-blind.""" + return [self.max_atoms - 1] + + def get_ntypes(self) -> int: + """Number of element types.""" + return len(self.type_map) + + def get_type_map(self) -> list[str]: + """Element names.""" + return self.type_map + + def get_dim_out(self) -> int: + """Width of the node representation.""" + return self.encoder_embed_dim + + def get_dim_emb(self) -> int: + """Width of the pair channel, which is the head count.""" + return self.encoder_attention_heads + + def mixed_types(self) -> bool: + """The neighbour list is not split by type.""" + return True + + def has_message_passing(self) -> bool: + """No ghost-atom exchange.""" + return False + + def need_sorted_nlist_for_lower(self) -> bool: + """All pairs take part, so their order does not matter.""" + return False + + def get_env_protection(self) -> float: + """No environment-matrix protection is used.""" + return 0.0 + + def supports_edge_parallel(self) -> bool: + """Global attention cannot be split across ranks.""" + return False + + def dense_lower_supports_comm(self) -> bool: + """There is no ghost-communication path.""" + return False + + def compression_needs_min_nbor_dist(self) -> bool: + """Compression is not supported, so no neighbour statistics are needed.""" + return False + + def compute_input_stats(self, merged, path=None) -> None: # noqa: ANN001 + """No environment statistics: the basis is learned, not normalized.""" + + def set_stat_mean_and_stddev(self, mean, stddev) -> None: # noqa: ANN001 + """Stat-free descriptor; nothing to assign.""" + + def get_stat_mean_and_stddev(self) -> tuple[list, list]: + """Stat-free descriptor; no statistics to report.""" + return [], [] + + def share_params(self, base_class, shared_level, resume=False) -> None: # noqa: ANN001 + """Parameter sharing is implemented by the PyTorch-Exportable wrapper.""" + raise NotImplementedError + + def change_type_map( + self, type_map: list[str], model_with_new_type_stat: object = None + ) -> None: + """Remap the element names onto Uni-Mol tokens. + + The token embedding itself is indexed by Uni-Mol token, not by deepmd + type, so only the lookup table changes. + """ + import numpy as np + + token_of = {sym: i for i, sym in enumerate(self.vocabulary)} + self.type_map = list(type_map) + self.type_to_token = np.array( + [token_of.get(sym, self.unk_idx) for sym in self.type_map], dtype=np.int64 + ) + + @classmethod + def update_sel(cls, train_data, type_map, local_jdata: dict) -> tuple[dict, None]: # noqa: ANN001 + """``sel`` follows from ``max_atoms``, so neighbour statistics are moot.""" + return local_jdata.copy(), None + + # ------------------------------------------------------------------ + # forward + # ------------------------------------------------------------------ + def build_tokens(self, coord_ext: Array, atype_ext: Array, nlist: Array) -> dict: + """Turn a padded deepmd frame into Uni-Mol's token sequence. + + Real atoms are identified from the neighbour list rather than from + ``atype``: by the time a descriptor is called, virtual atoms have been + clamped to type 0 and are indistinguishable from a real first element, + whereas the neighbour list still shows them as empty rows. + + Returns + ------- + dict + ``tokens``, ``coord``, ``padding_mask``, ``real_mask`` and + ``n_real``; the sequence is ``[CLS] atoms [SEP] pad...``. + """ + xp = array_api_compat.array_namespace(coord_ext) + dev = array_api_compat.device(coord_ext) + nf, nloc = nlist.shape[0], nlist.shape[1] + coord = xp.reshape(coord_ext, (nf, -1, 3)) + nall = coord.shape[1] + if nall != nloc: + raise ValueError( + "the unimol descriptor needs every atom to be local: it attends " + "over all pairs, so it supports neither periodic images nor the " + "ghost-atom layout that freezing and parallel evaluation assume " + f"(got {nall} extended atoms for {nloc} local atoms)" + ) + + real_mask = xp.any(nlist >= 0, axis=-1) + n_real = xp.sum(xp.astype(real_mask, xp.int64), axis=-1) + if bool(xp.any(n_real < 2)): + raise ValueError( + "the unimol descriptor needs at least two real atoms per frame; " + "single-atom frames cannot be told apart from padding" + ) + + real = xp.astype(real_mask, coord.dtype) + # BOS and EOS sit at the centroid of the real atoms, which keeps the + # sequence translation invariant. Upstream centres the coordinates in + # its data pipeline and then places both at the origin, so the two agree + # whenever the data went through that transform. + if self.virtual_token_position == "centroid": + centroid = ( + xp.sum(coord * real[..., None], axis=1) + / xp.astype(n_real, coord.dtype)[:, None] + ) + else: + centroid = xp.zeros((nf, 3), dtype=coord.dtype, device=dev) + + # The lookup table is plain integer data rather than a parameter, so it + # does not travel with the module and has to be placed explicitly. + token_table = xp.asarray(self.type_to_token, device=dev) + atom_tokens = xp.take(token_table, xp.reshape(atype_ext, (-1,)), axis=0) + atom_tokens = xp.reshape(atom_tokens, (nf, nloc)) + atom_tokens = xp.where( + real_mask, atom_tokens, xp.full_like(atom_tokens, self.pad_idx) + ) + + nt = nloc + 2 + positions = xp.arange(nt, dtype=xp.int64, device=dev)[None, :] + eos_at = (n_real + 1)[:, None] + bos_row = xp.full((nf, 1), self.bos_idx, dtype=atom_tokens.dtype, device=dev) + pad_row = xp.full((nf, 1), self.pad_idx, dtype=atom_tokens.dtype, device=dev) + tokens = xp.concat([bos_row, atom_tokens, pad_row], axis=1) + tokens = xp.where( + positions == eos_at, xp.full_like(tokens, self.eos_idx), tokens + ) + + virtual = centroid[:, None, :] + coord_full = xp.concat([virtual, coord, virtual], axis=1) + at_virtual = (positions == 0) | (positions == eos_at) + coord_full = xp.where( + at_virtual[..., None], + xp.broadcast_to(virtual, coord_full.shape), + coord_full, + ) + + padding_mask = xp.astype(tokens == self.pad_idx, coord.dtype) + return { + "tokens": tokens, + "coord": coord_full, + "padding_mask": padding_mask, + "real_mask": real_mask, + "n_real": n_real, + } + + @cast_precision + def forward_tokens(self, coord_ext: Array, atype_ext: Array, nlist: Array) -> dict: + """Run the backbone and return everything at token resolution. + + The five-tuple of :meth:`call` cannot carry the two virtual tokens or + the norm regularisers, which the pretraining heads need, so the + pretraining path goes through here instead. + """ + xp = array_api_compat.array_namespace(coord_ext) + seq = self.build_tokens(coord_ext, atype_ext, nlist) + tokens, coord, padding_mask = seq["tokens"], seq["coord"], seq["padding_mask"] + nf, nt = tokens.shape + + # The embedding is a parameter, so it is indexed directly: wrapping it + # with asarray would detach it from the gradient. + emb = xp.take(self.embed_tokens.w, xp.reshape(tokens, (-1,)), axis=0) + emb = xp.reshape(emb, (nf, nt, self.encoder_embed_dim)) + emb = xp.astype(emb, coord.dtype) + + diff = coord[:, :, None, :] - coord[:, None, :, :] + dist = xp.sqrt(xp.sum(diff**2, axis=-1)) + if self.single_precision_distance: + # Upstream's data pipeline stores the distance matrix in fp32, and + # the Gaussian basis is narrow enough for that rounding to matter. + dist = xp.astype(xp.astype(dist, xp.float32), dist.dtype) + edge_type = tokens[:, :, None] * self.ntokens + tokens[:, None, :] + + bias = self.gbf_proj(self.gbf(dist, edge_type)) + bias = xp.reshape(xp.permute_dims(bias, (0, 3, 1, 2)), (-1, nt, nt)) + + # ``training`` exists only once the PyTorch-Exportable wrapper has made + # this a torch module; on the array-API path it is always inference. + x, pair_rep, delta_pair_rep, x_norm, delta_pair_norm = self.encoder( + emb, bias, padding_mask, training=bool(getattr(self, "training", False)) + ) + # Upstream clears the -inf that padding leaves on the pair channel + # before any head reads it (unimol/models/unimol.py:221). + pair_rep = xp.where(xp.isinf(pair_rep), xp.zeros_like(pair_rep), pair_rep) + return { + **seq, + "node_ebd": x, + "pair_rep": pair_rep, + "delta_pair_rep": delta_pair_rep, + "x_norm": x_norm, + "delta_pair_norm": delta_pair_norm, + } + + def call( + self, + coord_ext: Array, + atype_ext: Array, + nlist: Array, + mapping: Array | None = None, + fparam: Array | None = None, + comm_dict: dict | None = None, + charge_spin: Array | None = None, + ) -> tuple[Array, None, None, None, None]: + """Return the per-atom representation. + + The two virtual tokens are dropped so that the output lines up with the + local atoms. Slots the backbone does not produce are ``None``, which the + atomic model accepts; the pretraining path uses :meth:`forward_tokens`. + """ + del mapping, fparam, comm_dict, charge_spin + out = self.forward_tokens(coord_ext, atype_ext, nlist) + nloc = nlist.shape[1] + return out["node_ebd"][:, 1 : nloc + 1, :], None, None, None, None + + # ------------------------------------------------------------------ + # serialization + # ------------------------------------------------------------------ + def serialize(self) -> dict: + """Serialize the descriptor.""" + return { + "@class": "Descriptor", + "type": "unimol", + "@version": 1, + "type_map": self.type_map, + "encoder_layers": self.encoder_layers, + "encoder_embed_dim": self.encoder_embed_dim, + "encoder_ffn_embed_dim": self.encoder_ffn_embed_dim, + "encoder_attention_heads": self.encoder_attention_heads, + "max_atoms": self.max_atoms, + "max_seq_len": self.max_seq_len, + "activation_function": self.activation_function, + "dropout": self.dropout, + "emb_dropout": self.emb_dropout, + "attention_dropout": self.attention_dropout, + "activation_dropout": self.activation_dropout, + "no_final_head_layer_norm": self.no_final_head_layer_norm, + "single_precision_basis": self.single_precision_basis, + "single_precision_distance": self.single_precision_distance, + "virtual_token_position": self.virtual_token_position, + "gaussian_kernels": self.gaussian_kernels, + "precision": self.precision, + "type_map_tokens": self.vocabulary, + "embed_tokens": self.embed_tokens.serialize(), + "gbf": self.gbf.serialize(), + "gbf_proj": self.gbf_proj.serialize(), + "encoder": self.encoder.serialize(), + } + + @classmethod + def deserialize(cls, data: dict) -> "DescrptUniMol": + """Deserialize the descriptor.""" + data = data.copy() + check_version_compatibility(data.pop("@version"), 1, 1) + data.pop("@class", None) + data.pop("type", None) + embed_tokens = data.pop("embed_tokens") + gbf = data.pop("gbf") + gbf_proj = data.pop("gbf_proj") + encoder = data.pop("encoder") + obj = cls(**data) + obj.embed_tokens = NativeLayer.deserialize(embed_tokens) + obj.gbf = GaussianLayer.deserialize(gbf) + obj.gbf_proj = NonLinearHead.deserialize(gbf_proj) + obj.encoder = TransformerEncoderWithPair.deserialize(encoder) + return obj diff --git a/deepmd/dpmodel/descriptor/unimol_nn/__init__.py b/deepmd/dpmodel/descriptor/unimol_nn/__init__.py new file mode 100644 index 0000000000..06a0a8de85 --- /dev/null +++ b/deepmd/dpmodel/descriptor/unimol_nn/__init__.py @@ -0,0 +1,26 @@ +# SPDX-License-Identifier: LGPL-3.0-or-later +"""Neural-network building blocks of the Uni-Mol v1 backbone.""" + +from .encoder import ( + GaussianLayer, + NonLinearHead, + SelfMultiheadAttention, + TransformerEncoderLayer, + TransformerEncoderWithPair, +) +from .heads import ( + DistanceHead, + MaskLMHead, + coord_update, +) + +__all__ = [ + "DistanceHead", + "GaussianLayer", + "MaskLMHead", + "NonLinearHead", + "SelfMultiheadAttention", + "TransformerEncoderLayer", + "TransformerEncoderWithPair", + "coord_update", +] diff --git a/deepmd/dpmodel/descriptor/unimol_nn/encoder.py b/deepmd/dpmodel/descriptor/unimol_nn/encoder.py new file mode 100644 index 0000000000..5d6e909269 --- /dev/null +++ b/deepmd/dpmodel/descriptor/unimol_nn/encoder.py @@ -0,0 +1,638 @@ +# SPDX-License-Identifier: LGPL-3.0-or-later +"""Uni-Mol v1 transformer encoder, ported to the array-API dpmodel layer. + +The maths follows Uni-Mol (https://github.com/deepmodeling/Uni-Mol) at commit +90f52c4 and the Uni-Core modules it builds on +(https://github.com/dptech-corp/Uni-Core) at commit ace6fae, both MIT licensed: + + Copyright (c) DP Technology + This source code is licensed under the MIT license found in the LICENSE + file in the root directory of that source tree. + +Parts of Uni-Core derive in turn from fairseq (Copyright (c) Facebook, Inc. and +its affiliates, MIT licensed). + +Ported files and the classes taken from them: + +======================================== ========================================== +upstream classes +======================================== ========================================== +unicore/modules/multihead_attention.py :class:`SelfMultiheadAttention` +unicore/modules/transformer_encoder_layer.py :class:`TransformerEncoderLayer` +unimol/models/transformer_encoder_with_pair.py :class:`TransformerEncoderWithPair` +unimol/models/unimol.py :class:`GaussianLayer`, :class:`NonLinearHead` +======================================== ========================================== + +The port keeps upstream's exact order of operations, so that a checkpoint +converted from Uni-Mol reproduces upstream outputs to floating-point rounding. +Two upstream quirks are reproduced deliberately and are marked in the code: the +truncated ``pi`` in the Gaussian basis, and the fp32 cast inside the norm +regularisers. +""" + +import array_api_compat + +from deepmd.dpmodel.common import ( + NativeOP, +) +from deepmd.dpmodel.utils.network import ( + LayerNorm, + NativeLayer, +) +from deepmd.dpmodel.utils.seed import ( + child_seed, +) + +# Uni-Mol's Gaussian basis uses a truncated pi (unimol/models/unimol.py:393-397). +# Keeping it is required for bitwise agreement with the released weights. +UNIMOL_PI = 3.14159 + + +def softmax(x, axis: int = -1): # noqa: ANN001, ANN201 + """Numerically stable softmax over ``axis``. + + Rows that are entirely ``-inf`` would produce NaN, as they do upstream. + Uni-Mol never builds such a row: the BOS column is never masked. + """ + xp = array_api_compat.array_namespace(x) + x_max = xp.max(x, axis=axis, keepdims=True) + e = xp.exp(x - x_max) + return e / xp.sum(e, axis=axis, keepdims=True) + + +def dropout(x, p: float, training: bool): # noqa: ANN001, ANN201 + """Apply dropout, but only while training. + + deepmd has no dropout anywhere else, and the array API has no random + numbers, so this dispatches to torch when a training step actually needs + it. Inference, which is what the array-API backends are for, is the + identity. Training on a non-torch backend is refused rather than silently + dropping the regularisation, which would be a quiet parity bug. + """ + if not training or p <= 0.0: + return x + if array_api_compat.is_torch_array(x): + import torch + + return torch.nn.functional.dropout(x, p=p, training=True) + raise NotImplementedError( + "dropout during training is only implemented for the PyTorch backends; " + "the array-API path is for inference" + ) + + +def norm_loss(x, eps: float = 1e-10, tolerance: float = 1.0): # noqa: ANN001, ANN201 + """Hinge on the deviation of the row norm from ``sqrt(dim)``. + + Mirrors ``norm_loss`` in unimol/models/transformer_encoder_with_pair.py:101. + Upstream evaluates this in fp32 because it pretrains a pure-fp16 model; the + cast is reproduced so that the value matches upstream exactly. + """ + xp = array_api_compat.array_namespace(x) + x = xp.astype(x, xp.float32) + max_norm = x.shape[-1] ** 0.5 + norm = xp.sqrt(xp.sum(x**2, axis=-1) + eps) + error = xp.abs(norm - max_norm) - tolerance + return xp.where(error > 0, error, xp.zeros_like(error)) + + +def masked_mean(mask, value, axis=-1, eps: float = 1e-10): # noqa: ANN001, ANN201 + """Mean of ``value`` over ``mask``, then mean over what is left. + + Mirrors ``masked_mean`` in transformer_encoder_with_pair.py:108. The ``eps`` + in the denominator is what makes an all-padding row return 0 rather than NaN. + """ + xp = array_api_compat.array_namespace(value) + mask = xp.astype(mask, value.dtype) + num = xp.sum(mask * value, axis=axis) + den = eps + xp.sum(mask, axis=axis) + return xp.mean(num / den) + + +class NonLinearHead(NativeOP): + """Two-layer head, ``linear1 -> activation -> linear2``. + + Mirrors unimol/models/unimol.py:347. + """ + + def __init__( + self, + input_dim: int, + out_dim: int, + activation_function: str = "gelu_erf", + hidden: int | None = None, + precision: str = "float64", + seed: int | list[int] | None = None, + ) -> None: + hidden = hidden or input_dim + self.input_dim = input_dim + self.out_dim = out_dim + self.activation_function = activation_function + self.hidden = hidden + self.precision = precision + self.linear1 = NativeLayer( + input_dim, + hidden, + activation_function=activation_function, + precision=precision, + seed=child_seed(seed, 0), + ) + self.linear2 = NativeLayer( + hidden, + out_dim, + activation_function=None, + precision=precision, + seed=child_seed(seed, 1), + ) + + def call(self, x): # noqa: ANN001, ANN201 + return self.linear2(self.linear1(x)) + + def serialize(self) -> dict: + return { + "input_dim": self.input_dim, + "out_dim": self.out_dim, + "hidden": self.hidden, + "activation_function": self.activation_function, + "precision": self.precision, + "linear1": self.linear1.serialize(), + "linear2": self.linear2.serialize(), + } + + @classmethod + def deserialize(cls, data: dict) -> "NonLinearHead": + data = data.copy() + linear1 = data.pop("linear1") + linear2 = data.pop("linear2") + obj = cls(**data) + obj.linear1 = NativeLayer.deserialize(linear1) + obj.linear2 = NativeLayer.deserialize(linear2) + return obj + + +class GaussianLayer(NativeOP): + """Gaussian radial basis with a per-edge-type affine on the distance. + + Mirrors unimol/models/unimol.py:400. ``means`` and ``stds`` are shared by all + edge types; ``mul`` and ``bias`` hold one scalar per ordered element pair. + """ + + def __init__( + self, + k: int = 128, + edge_types: int = 1024, + single_precision_basis: bool = True, + precision: str = "float64", + seed: int | list[int] | None = None, + ) -> None: + self.k = k + self.edge_types = edge_types + # Upstream evaluates the basis in fp32 because it pretrains an fp16 + # model (unimol/models/unimol.py:418-420). Keeping that is what makes + # the released weights reproduce; set False to keep the working dtype. + self.single_precision_basis = single_precision_basis + self.precision = precision + # All four tables are trained. They are layers rather than bare arrays + # because a bare array becomes a buffer on the torch backends, and a + # buffer never receives a gradient: the whole geometry pathway would + # sit frozen at its initial values. Each gets its own seed, or they + # would draw identical numbers wherever their shapes agree. + self.means = NativeLayer( + 1, k, bias=False, precision=precision, seed=child_seed(seed, 0) + ) + self.stds = NativeLayer( + 1, k, bias=False, precision=precision, seed=child_seed(seed, 1) + ) + self.mul = NativeLayer( + edge_types, 1, bias=False, precision=precision, seed=child_seed(seed, 2) + ) + self.bias = NativeLayer( + edge_types, 1, bias=False, precision=precision, seed=child_seed(seed, 3) + ) + + def call(self, dist, edge_type): # noqa: ANN001, ANN201 + """Expand ``dist`` (nf x nt x nt) into ``k`` Gaussians per atom pair.""" + xp = array_api_compat.array_namespace(dist) + flat = xp.reshape(edge_type, (-1,)) + # The tables are parameters: they already live on the right device, and + # re-wrapping them with asarray would cut them out of the gradient. + mul = xp.reshape(xp.take(self.mul.w, flat, axis=0), (*edge_type.shape, 1)) + bias = xp.reshape(xp.take(self.bias.w, flat, axis=0), (*edge_type.shape, 1)) + mul = xp.astype(mul, dist.dtype) + bias = xp.astype(bias, dist.dtype) + x = mul * dist[..., None] + bias + x = xp.repeat(x, self.k, axis=-1) + work = xp.float32 if self.single_precision_basis else x.dtype + x = xp.astype(x, work) + mean = xp.astype(xp.reshape(self.means.w, (-1,)), work) + std = xp.abs(xp.astype(xp.reshape(self.stds.w, (-1,)), work)) + 1e-5 + a = (2 * UNIMOL_PI) ** 0.5 + out = xp.exp(-0.5 * (((x - mean) / std) ** 2)) / (a * std) + return xp.astype(out, dist.dtype) + + def serialize(self) -> dict: + """Serialize the basis.""" + return { + "k": self.k, + "edge_types": self.edge_types, + "single_precision_basis": self.single_precision_basis, + "precision": self.precision, + "means": self.means.serialize(), + "stds": self.stds.serialize(), + "mul": self.mul.serialize(), + "bias": self.bias.serialize(), + } + + @classmethod + def deserialize(cls, data: dict) -> "GaussianLayer": + """Deserialize the basis.""" + data = data.copy() + tables = {key: data.pop(key) for key in ("means", "stds", "mul", "bias")} + obj = cls(**data) + for key, value in tables.items(): + setattr(obj, key, NativeLayer.deserialize(value)) + return obj + + +class SelfMultiheadAttention(NativeOP): + """Self-attention with a fused QKV projection and an additive pair bias. + + Mirrors unicore/modules/multihead_attention.py:12. Only the ``return_attn`` + path of upstream is kept, because Uni-Mol always takes it: the pre-softmax + logits are the pair representation that the next layer biases with. + """ + + def __init__( + self, + embed_dim: int, + num_heads: int, + dropout: float = 0.1, + bias: bool = True, + scaling_factor: float = 1.0, + precision: str = "float64", + seed: int | list[int] | None = None, + ) -> None: + assert embed_dim % num_heads == 0, "embed_dim must be divisible by num_heads" + self.embed_dim = embed_dim + self.num_heads = num_heads + self.dropout = dropout + self.head_dim = embed_dim // num_heads + self.scaling = (self.head_dim * scaling_factor) ** -0.5 + self.precision = precision + self.in_proj = NativeLayer( + embed_dim, + embed_dim * 3, + bias=bias, + precision=precision, + seed=child_seed(seed, 0), + ) + self.out_proj = NativeLayer( + embed_dim, + embed_dim, + bias=bias, + precision=precision, + seed=child_seed(seed, 1), + ) + + def call(self, query, attn_bias, training: bool = False): # noqa: ANN001, ANN201 + """Return the attended values and the biased pre-softmax logits. + + Parameters + ---------- + query + nf x nt x embed_dim. + attn_bias + (nf * num_heads) x nt x nt, already carrying ``-inf`` on padded keys. + + Returns + ------- + tuple + ``(output, attn_weights)``, shaped nf x nt x embed_dim and + (nf * num_heads) x nt x nt. + """ + xp = array_api_compat.array_namespace(query) + nf, nt, _ = query.shape + qkv = self.in_proj(query) + q, k, v = ( + qkv[..., i * self.embed_dim : (i + 1) * self.embed_dim] for i in range(3) + ) + + def split_heads(t): # noqa: ANN001, ANN202 + t = xp.reshape(t, (nf, nt, self.num_heads, self.head_dim)) + t = xp.permute_dims(t, (0, 2, 1, 3)) + return xp.reshape(t, (nf * self.num_heads, nt, self.head_dim)) + + q = split_heads(q) * self.scaling + k = split_heads(k) + v = split_heads(v) + + attn_weights = q @ xp.permute_dims(k, (0, 2, 1)) + # Upstream adds the bias in place and returns the biased logits, which + # become the next layer's bias (multihead_attention.py:100-103). + attn_weights = attn_weights + attn_bias + attn = dropout(softmax(attn_weights, axis=-1), self.dropout, training) + o = attn @ v + o = xp.reshape(o, (nf, self.num_heads, nt, self.head_dim)) + o = xp.permute_dims(o, (0, 2, 1, 3)) + o = xp.reshape(o, (nf, nt, self.embed_dim)) + return self.out_proj(o), attn_weights + + def serialize(self) -> dict: + return { + "embed_dim": self.embed_dim, + "num_heads": self.num_heads, + "dropout": self.dropout, + "precision": self.precision, + "in_proj": self.in_proj.serialize(), + "out_proj": self.out_proj.serialize(), + } + + @classmethod + def deserialize(cls, data: dict) -> "SelfMultiheadAttention": + data = data.copy() + in_proj = data.pop("in_proj") + out_proj = data.pop("out_proj") + obj = cls(**data) + obj.in_proj = NativeLayer.deserialize(in_proj) + obj.out_proj = NativeLayer.deserialize(out_proj) + return obj + + +class TransformerEncoderLayer(NativeOP): + """Pre-layer-norm transformer block that also returns attention logits. + + Mirrors unicore/modules/transformer_encoder_layer.py:15. Uni-Mol never sets + ``post_ln``, so only the pre-LN order is implemented. + """ + + def __init__( + self, + embed_dim: int = 768, + ffn_embed_dim: int = 3072, + attention_heads: int = 8, + dropout: float = 0.1, + attention_dropout: float = 0.1, + activation_dropout: float = 0.0, + activation_function: str = "gelu_erf", + precision: str = "float64", + seed: int | list[int] | None = None, + ) -> None: + self.embed_dim = embed_dim + self.ffn_embed_dim = ffn_embed_dim + self.attention_heads = attention_heads + self.dropout = dropout + self.attention_dropout = attention_dropout + self.activation_dropout = activation_dropout + self.activation_function = activation_function + self.precision = precision + self.self_attn = SelfMultiheadAttention( + embed_dim, + attention_heads, + dropout=attention_dropout, + precision=precision, + seed=child_seed(seed, 0), + ) + self.self_attn_layer_norm = LayerNorm( + embed_dim, precision=precision, seed=child_seed(seed, 1) + ) + self.fc1 = NativeLayer( + embed_dim, + ffn_embed_dim, + activation_function=activation_function, + precision=precision, + seed=child_seed(seed, 2), + ) + self.fc2 = NativeLayer( + ffn_embed_dim, + embed_dim, + activation_function=None, + precision=precision, + seed=child_seed(seed, 3), + ) + self.final_layer_norm = LayerNorm( + embed_dim, precision=precision, seed=child_seed(seed, 4) + ) + + def call(self, x, attn_bias, training: bool = False): # noqa: ANN001, ANN201 + residual = x + x = self.self_attn_layer_norm(x) + x, attn_weights = self.self_attn(x, attn_bias=attn_bias, training=training) + x = dropout(x, self.dropout, training) + x = residual + x + + residual = x + x = self.final_layer_norm(x) + # fc1 carries the activation, so the activation dropout sits between + # the two linears, as upstream has it. + x = dropout(self.fc1(x), self.activation_dropout, training) + x = dropout(self.fc2(x), self.dropout, training) + x = residual + x + return x, attn_weights + + def serialize(self) -> dict: + return { + "embed_dim": self.embed_dim, + "ffn_embed_dim": self.ffn_embed_dim, + "attention_heads": self.attention_heads, + "dropout": self.dropout, + "attention_dropout": self.attention_dropout, + "activation_dropout": self.activation_dropout, + "activation_function": self.activation_function, + "precision": self.precision, + "self_attn": self.self_attn.serialize(), + "self_attn_layer_norm": self.self_attn_layer_norm.serialize(), + "fc1": self.fc1.serialize(), + "fc2": self.fc2.serialize(), + "final_layer_norm": self.final_layer_norm.serialize(), + } + + @classmethod + def deserialize(cls, data: dict) -> "TransformerEncoderLayer": + data = data.copy() + parts = { + key: data.pop(key) + for key in ( + "self_attn", + "self_attn_layer_norm", + "fc1", + "fc2", + "final_layer_norm", + ) + } + obj = cls(**data) + obj.self_attn = SelfMultiheadAttention.deserialize(parts["self_attn"]) + obj.self_attn_layer_norm = LayerNorm.deserialize(parts["self_attn_layer_norm"]) + obj.fc1 = NativeLayer.deserialize(parts["fc1"]) + obj.fc2 = NativeLayer.deserialize(parts["fc2"]) + obj.final_layer_norm = LayerNorm.deserialize(parts["final_layer_norm"]) + return obj + + +class TransformerEncoderWithPair(NativeOP): + """Stack of pre-LN blocks that carries a pair representation. + + Mirrors unimol/models/transformer_encoder_with_pair.py:14. Each block's + pre-softmax logits become the next block's bias, so the pair representation + is the running sum of per-layer logits. Returns the node representation, the + final pair representation, the pair delta, and the two norm regularisers. + """ + + def __init__( + self, + encoder_layers: int = 6, + embed_dim: int = 768, + ffn_embed_dim: int = 3072, + attention_heads: int = 8, + emb_dropout: float = 0.1, + dropout: float = 0.1, + attention_dropout: float = 0.1, + activation_dropout: float = 0.0, + max_seq_len: int = 256, + activation_function: str = "gelu_erf", + no_final_head_layer_norm: bool = False, + precision: str = "float64", + seed: int | list[int] | None = None, + ) -> None: + self.encoder_layers = encoder_layers + self.embed_dim = embed_dim + self.ffn_embed_dim = ffn_embed_dim + self.attention_heads = attention_heads + self.emb_dropout = emb_dropout + self.dropout = dropout + self.attention_dropout = attention_dropout + self.activation_dropout = activation_dropout + self.max_seq_len = max_seq_len + self.activation_function = activation_function + self.no_final_head_layer_norm = no_final_head_layer_norm + self.precision = precision + self.emb_layer_norm = LayerNorm( + embed_dim, precision=precision, seed=child_seed(seed, 0) + ) + self.final_layer_norm = LayerNorm( + embed_dim, precision=precision, seed=child_seed(seed, 1) + ) + self.final_head_layer_norm = ( + None + if no_final_head_layer_norm + else LayerNorm( + attention_heads, precision=precision, seed=child_seed(seed, 2) + ) + ) + self.layers = [ + TransformerEncoderLayer( + embed_dim=embed_dim, + ffn_embed_dim=ffn_embed_dim, + attention_heads=attention_heads, + dropout=dropout, + attention_dropout=attention_dropout, + activation_dropout=activation_dropout, + activation_function=activation_function, + precision=precision, + # Each block draws its own numbers; sharing one seed would make + # every layer start bitwise identical. + seed=child_seed(seed, 3 + index), + ) + for index in range(encoder_layers) + ] + + def call(self, emb, attn_mask, padding_mask, training: bool = False): # noqa: ANN001, ANN201 + """Run the stack. + + Parameters + ---------- + emb + nf x nt x embed_dim token embedding. + attn_mask + (nf * heads) x nt x nt bias from the Gaussian basis. + padding_mask + nf x nt, 1 on padded positions. + + Returns + ------- + tuple + ``(x, pair_rep, delta_pair_rep, x_norm, delta_pair_rep_norm)``. + """ + xp = array_api_compat.array_namespace(emb) + nf, nt = emb.shape[0], emb.shape[1] + x = dropout(self.emb_layer_norm(emb), self.emb_dropout, training) + pad = xp.astype(padding_mask, x.dtype) + x = x * (1 - pad[..., None]) + + # Upstream merges padding into the bias as -inf on padded key columns, + # so the attention itself never sees a key_padding_mask. It does this + # in place, which aliases the saved input bias; the delta below is + # therefore taken against the already-filled bias, and the resulting + # NaN on padded pairs is overwritten with 0 right after. + neg_inf = xp.asarray(float("-inf"), dtype=attn_mask.dtype) + key_pad = xp.astype(padding_mask, xp.bool)[:, None, None, :] + attn_mask = xp.reshape(attn_mask, (nf, -1, nt, nt)) + attn_mask = xp.where(key_pad, neg_inf, attn_mask) + input_attn_mask = attn_mask + attn_mask = xp.reshape(attn_mask, (-1, nt, nt)) + + for layer in self.layers: + x, attn_mask = layer(x, attn_bias=attn_mask, training=training) + + token_mask = 1.0 - pad + x_norm = masked_mean(token_mask, norm_loss(x)) + + x = self.final_layer_norm(x) + + # Padded pairs are dropped from the delta anyway. Upstream reaches that + # by computing -inf minus -inf, getting NaN, and overwriting it with 0; + # zeroing both operands first gives the same values without the NaN. + out_pair = xp.reshape(attn_mask, (nf, -1, nt, nt)) + zero = xp.zeros_like(out_pair) + delta_pair_repr = xp.where(key_pad, zero, out_pair) - xp.where( + key_pad, zero, input_attn_mask + ) + pair_rep = xp.permute_dims(out_pair, (0, 2, 3, 1)) + delta_pair_repr = xp.permute_dims(delta_pair_repr, (0, 2, 3, 1)) + + pair_mask = token_mask[..., None] * token_mask[..., None, :] + delta_pair_repr_norm = masked_mean( + pair_mask, norm_loss(delta_pair_repr), axis=(-1, -2) + ) + + if self.final_head_layer_norm is not None: + delta_pair_repr = self.final_head_layer_norm(delta_pair_repr) + + return x, pair_rep, delta_pair_repr, x_norm, delta_pair_repr_norm + + def serialize(self) -> dict: + return { + "encoder_layers": self.encoder_layers, + "embed_dim": self.embed_dim, + "ffn_embed_dim": self.ffn_embed_dim, + "attention_heads": self.attention_heads, + "emb_dropout": self.emb_dropout, + "dropout": self.dropout, + "attention_dropout": self.attention_dropout, + "activation_dropout": self.activation_dropout, + "max_seq_len": self.max_seq_len, + "activation_function": self.activation_function, + "no_final_head_layer_norm": self.no_final_head_layer_norm, + "precision": self.precision, + "emb_layer_norm": self.emb_layer_norm.serialize(), + "final_layer_norm": self.final_layer_norm.serialize(), + "final_head_layer_norm": None + if self.final_head_layer_norm is None + else self.final_head_layer_norm.serialize(), + "layers": [layer.serialize() for layer in self.layers], + } + + @classmethod + def deserialize(cls, data: dict) -> "TransformerEncoderWithPair": + data = data.copy() + emb_ln = data.pop("emb_layer_norm") + final_ln = data.pop("final_layer_norm") + head_ln = data.pop("final_head_layer_norm") + layers = data.pop("layers") + obj = cls(**data) + obj.emb_layer_norm = LayerNorm.deserialize(emb_ln) + obj.final_layer_norm = LayerNorm.deserialize(final_ln) + obj.final_head_layer_norm = ( + None if head_ln is None else LayerNorm.deserialize(head_ln) + ) + obj.layers = [TransformerEncoderLayer.deserialize(layer) for layer in layers] + return obj diff --git a/deepmd/dpmodel/descriptor/unimol_nn/heads.py b/deepmd/dpmodel/descriptor/unimol_nn/heads.py new file mode 100644 index 0000000000..e96db45952 --- /dev/null +++ b/deepmd/dpmodel/descriptor/unimol_nn/heads.py @@ -0,0 +1,209 @@ +# SPDX-License-Identifier: LGPL-3.0-or-later +"""The three Uni-Mol v1 pretraining heads. + +Ported from Uni-Mol (https://github.com/deepmodeling/Uni-Mol) at commit 90f52c4, +MIT licensed: + + Copyright (c) DP Technology + This source code is licensed under the MIT license found in the LICENSE + file in the root directory of that source tree. + +``MaskLMHead`` predicts the element of every corrupted atom, ``coord_update`` +denoises the coordinates through the pair channel, and ``DistanceHead`` +predicts the clean pairwise distances. Together with the two norm regularisers +of the encoder these make up the five pretraining objectives. +""" + +import array_api_compat + +from deepmd.dpmodel.common import ( + NativeOP, +) +from deepmd.dpmodel.utils.network import ( + LayerNorm, + NativeLayer, +) +from deepmd.dpmodel.utils.seed import ( + child_seed, +) + +__all__ = ["DistanceHead", "MaskLMHead", "coord_update"] + + +class MaskLMHead(NativeOP): + """Element prediction for the corrupted atoms. + + Mirrors unimol/models/unimol.py:292. The output projection is a free + parameter, not tied to the token embedding: upstream takes the weight of a + throw-away ``nn.Linear`` and keeps it (``:301-303``). + """ + + def __init__( + self, + embed_dim: int, + output_dim: int, + activation_function: str = "gelu_erf", + precision: str = "float64", + seed: int | list[int] | None = None, + ) -> None: + self.embed_dim = embed_dim + self.output_dim = output_dim + self.activation_function = activation_function + self.precision = precision + self.dense = NativeLayer( + embed_dim, + embed_dim, + activation_function=activation_function, + precision=precision, + seed=child_seed(seed, 0), + ) + self.layer_norm = LayerNorm( + embed_dim, precision=precision, seed=child_seed(seed, 1) + ) + self.out_proj = NativeLayer( + embed_dim, + output_dim, + bias=True, + precision=precision, + seed=child_seed(seed, 2), + ) + + def call(self, features, masked_tokens=None): # noqa: ANN001, ANN201 + """Project the selected positions onto the vocabulary. + + Parameters + ---------- + features + nf x nt x embed_dim node representation. + masked_tokens + nf x nt boolean mask; only these positions are projected, which is + what upstream does to save memory. + + Returns + ------- + Array + n_masked x output_dim logits, or nf x nt x output_dim if no mask. + """ + xp = array_api_compat.array_namespace(features) + if masked_tokens is not None: + # sole index: the array API allows a boolean mask only on its own + features = features[xp.astype(masked_tokens, xp.bool)] + return self.out_proj(self.layer_norm(self.dense(features))) + + def serialize(self) -> dict: + return { + "embed_dim": self.embed_dim, + "output_dim": self.output_dim, + "activation_function": self.activation_function, + "precision": self.precision, + "dense": self.dense.serialize(), + "layer_norm": self.layer_norm.serialize(), + "out_proj": self.out_proj.serialize(), + } + + @classmethod + def deserialize(cls, data: dict) -> "MaskLMHead": + data = data.copy() + parts = {k: data.pop(k) for k in ("dense", "layer_norm", "out_proj")} + obj = cls(**data) + obj.dense = NativeLayer.deserialize(parts["dense"]) + obj.layer_norm = LayerNorm.deserialize(parts["layer_norm"]) + obj.out_proj = NativeLayer.deserialize(parts["out_proj"]) + return obj + + +class DistanceHead(NativeOP): + """Pairwise distance prediction from the pair representation. + + Mirrors unimol/models/unimol.py:370. It reads the final attention logits, + not the pair delta, and symmetrises its output. + """ + + def __init__( + self, + heads: int, + activation_function: str = "gelu_erf", + precision: str = "float64", + seed: int | list[int] | None = None, + ) -> None: + self.heads = heads + self.activation_function = activation_function + self.precision = precision + self.dense = NativeLayer( + heads, + heads, + activation_function=activation_function, + precision=precision, + seed=child_seed(seed, 0), + ) + self.layer_norm = LayerNorm( + heads, precision=precision, seed=child_seed(seed, 1) + ) + self.out_proj = NativeLayer( + heads, 1, bias=True, precision=precision, seed=child_seed(seed, 2) + ) + + def call(self, pair_rep): # noqa: ANN001, ANN201 + """Map nf x nt x nt x heads onto a symmetric nf x nt x nt matrix.""" + xp = array_api_compat.array_namespace(pair_rep) + x = self.out_proj(self.layer_norm(self.dense(pair_rep))) + x = xp.reshape(x, x.shape[:3]) + return (x + xp.matrix_transpose(x)) * 0.5 + + def serialize(self) -> dict: + return { + "heads": self.heads, + "activation_function": self.activation_function, + "precision": self.precision, + "dense": self.dense.serialize(), + "layer_norm": self.layer_norm.serialize(), + "out_proj": self.out_proj.serialize(), + } + + @classmethod + def deserialize(cls, data: dict) -> "DistanceHead": + data = data.copy() + parts = {k: data.pop(k) for k in ("dense", "layer_norm", "out_proj")} + obj = cls(**data) + obj.dense = NativeLayer.deserialize(parts["dense"]) + obj.layer_norm = LayerNorm.deserialize(parts["layer_norm"]) + obj.out_proj = NativeLayer.deserialize(parts["out_proj"]) + return obj + + +def coord_update(coord, delta_pair_rep, padding_mask, pair2coord_proj): # noqa: ANN001, ANN201 + r"""Denoise coordinates through the pair channel. + + Mirrors unimol/models/unimol.py:229-245, the form introduced by upstream + #211: the normaliser counts every non-padding token, BOS and EOS included, + and pairs touching padding are zeroed before the sum. + + .. math:: + + \hat{x}_i = x_i + \sum_j \frac{x_j - x_i}{N} c_{ij} + + Parameters + ---------- + coord + nf x nt x 3 input (noisy) coordinates. + delta_pair_rep + nf x nt x nt x heads pair delta from the encoder. + padding_mask + nf x nt, 1 on padded positions. + pair2coord_proj + Head mapping heads -> 1. + + Returns + ------- + Array + nf x nt x 3 updated coordinates. + """ + xp = array_api_compat.array_namespace(coord) + pad = xp.astype(padding_mask, coord.dtype) + atom_num = xp.reshape(xp.sum(1 - pad, axis=1), (-1, 1, 1, 1)) + delta_pos = coord[:, None, :, :] - coord[:, :, None, :] + attn_probs = pair2coord_proj(delta_pair_rep) + update = delta_pos / atom_num * attn_probs + pair_coords_mask = (1 - pad)[..., None] * (1 - pad)[:, None, :] + update = update * pair_coords_mask[..., None] + return coord + xp.sum(update, axis=2) diff --git a/deepmd/dpmodel/fitting/unimol_pretrain.py b/deepmd/dpmodel/fitting/unimol_pretrain.py new file mode 100644 index 0000000000..bb804d60ea --- /dev/null +++ b/deepmd/dpmodel/fitting/unimol_pretrain.py @@ -0,0 +1,375 @@ +# SPDX-License-Identifier: LGPL-3.0-or-later +"""The Uni-Mol v1 pretraining head set, as a deepmd fitting. + +Ported from Uni-Mol (https://github.com/deepmodeling/Uni-Mol) at commit 90f52c4, +MIT licensed: + + Copyright (c) DP Technology + This source code is licensed under the MIT license found in the LICENSE + file in the root directory of that source tree. + +Three heads sit on the Uni-Mol backbone: the element head reads the node +representation, the coordinate head reads the pair delta, and the distance head +reads the pair representation. None of the three is reducible to a frame total +and none is differentiated with respect to the coordinates, because this task +denoises structures rather than predicting a potential energy surface. + +The distance head predicts a full row per atom, and upstream's objective counts +the two virtual tokens among the columns, so the output is padded to +``max_atoms + 2`` columns and the loss masks it back down. +""" + +import array_api_compat + +from deepmd.dpmodel.array_api import ( + Array, +) +from deepmd.dpmodel.common import ( + NativeOP, + safe_cast_array, +) +from deepmd.dpmodel.descriptor.unimol_nn import ( + DistanceHead, + MaskLMHead, + NonLinearHead, + coord_update, +) +from deepmd.dpmodel.fitting.base_fitting import ( + BaseFitting, +) +from deepmd.dpmodel.output_def import ( + FittingOutputDef, + OutputVariableDef, +) +from deepmd.dpmodel.utils.seed import ( + child_seed, +) +from deepmd.utils.version import ( + check_version_compatibility, +) + + +@BaseFitting.register("unimol_pretrain") +class UniMolPretrainFitting(NativeOP, BaseFitting): + r"""The three self-supervised heads of Uni-Mol v1. + + Parameters + ---------- + ntypes : int + Number of element types of the model. + dim_descrpt : int + Width of the node representation produced by the backbone. + n_token : int + Size of the Uni-Mol vocabulary, which is the width of the element head. + attention_heads : int + Width of the pair channel, which both pair-reading heads consume. + max_atoms : int + Largest molecule accepted, which fixes the column count of the distance + head at ``max_atoms + 2``. + activation_function : str + Activation of the heads; Uni-Mol uses the exact GELU. + mask_token_head : bool + Build the element head. + coord_head : bool + Build the coordinate head. + dist_head : bool + Build the distance head. + precision : str + Floating-point precision of the parameters. + seed : int, optional + Random seed for initialization. + """ + + def __init__( + self, + ntypes: int, + dim_descrpt: int, + n_token: int = 31, + attention_heads: int = 64, + max_atoms: int = 256, + activation_function: str = "gelu_erf", + mask_token_head: bool = True, + coord_head: bool = True, + dist_head: bool = True, + precision: str = "float64", + seed: int | list[int] | None = None, + type_map: list[str] | None = None, + **kwargs: float | str | bool, + ) -> None: + del kwargs + self.type_map = list(type_map) if type_map is not None else None + self.ntypes = ntypes + self.dim_descrpt = dim_descrpt + self.n_token = n_token + self.attention_heads = attention_heads + self.max_atoms = max_atoms + self.activation_function = activation_function + self.precision = precision + self.lm_head = ( + MaskLMHead( + dim_descrpt, + n_token, + activation_function, + precision, + child_seed(seed, 0), + ) + if mask_token_head + else None + ) + self.pair2coord_proj = ( + NonLinearHead( + attention_heads, + 1, + activation_function, + hidden=attention_heads, + precision=precision, + seed=child_seed(seed, 1), + ) + if coord_head + else None + ) + self.dist_head = ( + DistanceHead( + attention_heads, activation_function, precision, child_seed(seed, 2) + ) + if dist_head + else None + ) + + def get_type_map(self) -> list[str]: + """Element names, as configured on the model.""" + return self.type_map if self.type_map is not None else [] + + def change_type_map( + self, type_map: list[str], model_with_new_type_stat: object = None + ) -> None: + """Adopt a new element list. + + The heads are indexed by Uni-Mol token, not by deepmd type, so nothing + but the recorded names changes. + """ + del model_with_new_type_stat + self.type_map = list(type_map) + self.ntypes = len(type_map) + + def compute_input_stats(self, merged, stat_file_path=None, **kwargs) -> None: # noqa: ANN001, ANN003 + """No input statistics: the heads read a learned representation. + + Nothing here is normalized against the training set, so there is + nothing to accumulate. + """ + + def get_dim_fparam(self) -> int: + """No frame parameters: the objective reads structure only.""" + return 0 + + def get_dim_aparam(self) -> int: + """No atomic parameters.""" + return 0 + + def has_default_fparam(self) -> bool: + """There are no frame parameters, so there is no default either.""" + return False + + def get_default_fparam(self): # noqa: ANN201 + """There are no frame parameters.""" + return None + + def get_sel_type(self) -> list[int]: + """Every element takes part in the objective.""" + return [] + + def reinit_exclude(self, exclude_types: list[int] | None = None) -> None: + """Type exclusion is meaningless here: every atom is predicted.""" + if exclude_types: + raise NotImplementedError( + "unimol_pretrain predicts every atom and does not support " + "excluded types" + ) + + def set_case_embd(self, case_idx: int) -> None: + """Case embeddings are a multi-task feature this fitting does not use.""" + raise NotImplementedError("unimol_pretrain does not support case embeddings") + + def output_def(self) -> FittingOutputDef: + """Declare the three head outputs. + + All three are per-atom, none reduces to a frame total, and none is + differentiated with respect to coordinates or cell. + """ + variables = [] + if self.lm_head is not None: + variables.append( + OutputVariableDef( + "token_logits", + [self.n_token], + reducible=False, + r_differentiable=False, + c_differentiable=False, + ) + ) + if self.pair2coord_proj is not None: + variables.append( + OutputVariableDef( + "coord_update", + [3], + reducible=False, + r_differentiable=False, + c_differentiable=False, + ) + ) + if self.dist_head is not None: + variables.append( + OutputVariableDef( + "pair_dist", + [self.max_atoms + 2], + reducible=False, + r_differentiable=False, + c_differentiable=False, + ) + ) + # The two regularisers are frame scalars, but only per-atom variables + # survive the atomic-output machinery, so each is broadcast over the + # local atoms and the loss averages it back with the real-atom mask. + for name in ("x_norm", "delta_pair_norm"): + variables.append( + OutputVariableDef( + name, + [1], + reducible=False, + r_differentiable=False, + c_differentiable=False, + ) + ) + return FittingOutputDef(variables) + + def call_tokens(self, backbone: dict[str, Array]) -> dict[str, Array]: + """Run the heads on the token-resolution backbone output. + + Parameters + ---------- + backbone + The dictionary returned by + :meth:`deepmd.dpmodel.descriptor.unimol.DescrptUniMol.forward_tokens`. + + Returns + ------- + dict + ``token_logits`` and ``coord_update`` carry one row per local atom, + with the two virtual tokens dropped; ``pair_dist`` keeps the virtual + columns, padded out to ``max_atoms + 2``; the two norm regularisers + are passed through for the loss. + """ + # The backbone hands its output back at the global precision, because + # its own forward is wrapped in ``cast_precision``. These heads may be + # configured at a different one, so the dictionary is cast here and the + # results cast back on the way out. ``cast_precision`` cannot do it: it + # casts arrays it is handed directly, and this argument is a dictionary. + backbone = { + kk: safe_cast_array(vv, "global", self.precision) + for kk, vv in backbone.items() + } + xp = array_api_compat.array_namespace(backbone["node_ebd"]) + node = backbone["node_ebd"] + nf, nt = node.shape[0], node.shape[1] + nloc = nt - 2 + out = {} + for name in ("x_norm", "delta_pair_norm"): + value = xp.astype(xp.reshape(backbone[name], (1, 1, 1)), node.dtype) + out[name] = ( + xp.zeros( + (nf, nloc, 1), + dtype=node.dtype, + device=array_api_compat.device(node), + ) + + value + ) + if self.lm_head is not None: + logits = self.lm_head(node) + out["token_logits"] = logits[:, 1 : nloc + 1, :] + if self.pair2coord_proj is not None: + updated = coord_update( + backbone["coord"], + backbone["delta_pair_rep"], + backbone["padding_mask"], + self.pair2coord_proj, + ) + out["coord_update"] = updated[:, 1 : nloc + 1, :] + if self.dist_head is not None: + dist = self.dist_head(backbone["pair_rep"])[:, 1 : nloc + 1, :] + width = self.max_atoms + 2 + if dist.shape[-1] > width: + raise ValueError( + f"a frame of {nloc} atoms exceeds max_atoms={self.max_atoms}; " + "the distance head declares a fixed width, so larger frames " + "cannot be expressed. Convert the data with a matching " + "max_atoms, or raise it here" + ) + if dist.shape[-1] < width: + pad = xp.zeros( + (nf, nloc, width - dist.shape[-1]), + dtype=dist.dtype, + device=array_api_compat.device(dist), + ) + dist = xp.concat([dist, pad], axis=-1) + out["pair_dist"] = dist + return { + kk: safe_cast_array(vv, self.precision, "global") for kk, vv in out.items() + } + + def call(self, descriptor: Array, atype: Array, **kwargs) -> dict[str, Array]: # noqa: ANN003 + """Not reachable: the heads need token-resolution inputs. + + The pretraining path goes through :meth:`call_tokens`, driven by the + Uni-Mol atomic model, because the two virtual tokens and the pair + channel do not fit the descriptor's five-tuple. + """ + raise NotImplementedError( + "unimol_pretrain reads token-resolution backbone output; it is driven " + "through call_tokens by the unimol atomic model" + ) + + def serialize(self) -> dict: + """Serialize the fitting.""" + return { + "@class": "Fitting", + "type": "unimol_pretrain", + "@version": 1, + "ntypes": self.ntypes, + "type_map": self.type_map, + "dim_descrpt": self.dim_descrpt, + "n_token": self.n_token, + "attention_heads": self.attention_heads, + "max_atoms": self.max_atoms, + "activation_function": self.activation_function, + "precision": self.precision, + "mask_token_head": self.lm_head is not None, + "coord_head": self.pair2coord_proj is not None, + "dist_head": self.dist_head is not None, + "lm_head": None if self.lm_head is None else self.lm_head.serialize(), + "pair2coord_proj": None + if self.pair2coord_proj is None + else self.pair2coord_proj.serialize(), + "dist_head_net": None + if self.dist_head is None + else self.dist_head.serialize(), + } + + @classmethod + def deserialize(cls, data: dict) -> "UniMolPretrainFitting": + """Deserialize the fitting.""" + data = data.copy() + check_version_compatibility(data.pop("@version"), 1, 1) + data.pop("@class", None) + data.pop("type", None) + parts = { + k: data.pop(k) for k in ("lm_head", "pair2coord_proj", "dist_head_net") + } + obj = cls(**data) + if parts["lm_head"] is not None: + obj.lm_head = MaskLMHead.deserialize(parts["lm_head"]) + if parts["pair2coord_proj"] is not None: + obj.pair2coord_proj = NonLinearHead.deserialize(parts["pair2coord_proj"]) + if parts["dist_head_net"] is not None: + obj.dist_head = DistanceHead.deserialize(parts["dist_head_net"]) + return obj diff --git a/deepmd/dpmodel/loss/__init__.py b/deepmd/dpmodel/loss/__init__.py index 115f1f6b03..605949aaa9 100644 --- a/deepmd/dpmodel/loss/__init__.py +++ b/deepmd/dpmodel/loss/__init__.py @@ -14,6 +14,9 @@ from deepmd.dpmodel.loss.tensor import ( TensorLoss, ) +from deepmd.dpmodel.loss.unimol import ( + UniMolLoss, +) __all__ = [ "DOSLoss", @@ -21,4 +24,5 @@ "EnergySpinLoss", "PropertyLoss", "TensorLoss", + "UniMolLoss", ] diff --git a/deepmd/dpmodel/loss/loss.py b/deepmd/dpmodel/loss/loss.py index 926c6e28c5..91a1c2137c 100644 --- a/deepmd/dpmodel/loss/loss.py +++ b/deepmd/dpmodel/loss/loss.py @@ -51,6 +51,28 @@ def call( def label_requirement(self) -> list[DataRequirementItem]: """Return data label requirements needed for this loss calculation.""" + def frame_transform(self, type_map: list[str], stream: str = "default"): # noqa: ANN201 + """Return a per-frame data transform this objective needs, or None. + + Self-supervised objectives build their own labels by corrupting the + input, which has to happen while the data is read rather than inside + the loss. A trainer installs whatever this returns on the datasets of + the corresponding task. Supervised losses need nothing and return None. + + Parameters + ---------- + type_map : list[str] + Element names of the model, which a transform needs in order to map + elements onto types. + stream : str + Which dataset the transform is for. An objective whose randomness + must not be shared between datasets derives its draw sequence from + this, so the caller passes something stable and distinct per + dataset; the trainer uses ``/training`` and + ``/validation``. + """ + return None + @property def supports_ragged_batches(self) -> bool: """Whether this objective accepts a flat per-node batch axis.""" diff --git a/deepmd/dpmodel/loss/unimol.py b/deepmd/dpmodel/loss/unimol.py new file mode 100644 index 0000000000..08c464ffc6 --- /dev/null +++ b/deepmd/dpmodel/loss/unimol.py @@ -0,0 +1,338 @@ +# SPDX-License-Identifier: LGPL-3.0-or-later +"""The Uni-Mol v1 molecular-pretraining objective. + +Ported from Uni-Mol (https://github.com/deepmodeling/Uni-Mol) at commit 90f52c4, +MIT licensed: + + Copyright (c) DP Technology + This source code is licensed under the MIT license found in the LICENSE + file in the root directory of that source tree. + +Five terms, with upstream's default weights from its README pretraining recipe: +element prediction (1), coordinate denoising (5), distance prediction (10), and +the two norm regularisers (0.01 each). The regularisers are produced by the +backbone, so the loss only weights them. +""" + +import array_api_compat + +from deepmd.dpmodel.array_api import ( + Array, +) +from deepmd.dpmodel.loss.loss import ( + Loss, +) +from deepmd.utils.data import ( + DataRequirementItem, +) +from deepmd.utils.version import ( + check_version_compatibility, +) + +# Upstream normalises the distance target with these two constants +# (unimol/losses/unimol.py:17-18). They are hard-coded there, not fitted. +DIST_MEAN = 6.312581655060595 +DIST_STD = 3.3899264663911888 + + +def _smooth_l1(pred: Array, label: Array, beta: float = 1.0) -> Array: + r"""Mean smooth L1, matching ``F.smooth_l1_loss(reduction="mean")``. + + .. math:: + + \ell_\beta(e)=\begin{cases} + e^2/(2\beta),&|e|<\beta,\\ + |e|-\beta/2,&|e|\ge\beta. + \end{cases} + """ + xp = array_api_compat.array_namespace(pred) + diff = xp.abs(pred - label) + elementwise = xp.where(diff < beta, 0.5 * diff**2 / beta, diff - 0.5 * beta) + if 0 in elementwise.shape: + # The number of corrupted atoms is rounded stochastically, so a frame + # can draw none at all, and a batch of one such frame leaves nothing to + # average. Upstream returns NaN there, which would go on to poison every + # weight; zero is the only finite answer. Non-empty batches, which is + # every batch upstream ever trained on, are untouched. + return xp.zeros( + (), dtype=elementwise.dtype, device=array_api_compat.device(elementwise) + ) + return xp.mean(elementwise) + + +def _frame_scalar(value: Array, mask: Array | None) -> Array: + """Recover a frame scalar that was broadcast over the local atoms. + + The two norm regularisers are frame quantities, but the model can only emit + per-atom variables, so each real atom carries the same value and padded ones + are zeroed. Averaging over the real atoms returns the original scalar. + """ + xp = array_api_compat.array_namespace(value) + if value.ndim <= 1: + return xp.reshape(value, ()) + per_atom = xp.reshape(value, value.shape[:2]) + if mask is None: + return xp.mean(per_atom) + weights = xp.astype(mask, per_atom.dtype) + total = xp.sum(weights) + # A frame of nothing but padding would divide zero by zero; the other two + # reductions in this file already refuse to, so this one does too. + return xp.sum(per_atom * weights) / xp.where(total > 0, total, xp.ones_like(total)) + + +def _token_mask_from_atoms(mask: Array, ncol: int) -> Array: + """Mark the non-padding token columns: BOS, the real atoms, then EOS.""" + xp = array_api_compat.array_namespace(mask) + n_real = xp.sum(xp.astype(mask, xp.int64), axis=-1) + positions = xp.arange(ncol, dtype=xp.int64, device=array_api_compat.device(mask))[ + None, : + ] + return xp.astype(positions < (n_real + 2)[:, None], xp.int64) + + +def _clean_distances(coord_target: Array, mask: Array, ncol: int) -> Array: + """Pairwise distances from the clean coordinates, virtual tokens included. + + Storing this as a label would cost O(natoms^2) per frame, so it is derived + here instead. The two virtual tokens sit at the origin, which is the + centroid of the clean coordinates because the transform centres them. + """ + xp = array_api_compat.array_namespace(coord_target) + nf = coord_target.shape[0] + real = xp.astype(mask, coord_target.dtype)[..., None] + atoms = coord_target * real + dev = array_api_compat.device(coord_target) + zero = xp.zeros((nf, 1, 3), dtype=coord_target.dtype, device=dev) + tokens = xp.concat([zero, atoms, zero], axis=1) + if tokens.shape[1] < ncol: + pad = xp.zeros( + (nf, ncol - tokens.shape[1], 3), dtype=coord_target.dtype, device=dev + ) + tokens = xp.concat([tokens, pad], axis=1) + diff = atoms[:, :, None, :] - tokens[:, None, :, :] + return xp.sqrt(xp.sum(diff**2, axis=-1)) + + +def _masked_nll(logits: Array, target: Array, pad_idx: int) -> Array: + """Negative log likelihood over the selected positions. + + Upstream evaluates ``log_softmax`` in fp32 (``losses/unimol.py:36``) because + it pretrains an fp16 model; the cast is kept so the value matches. + """ + xp = array_api_compat.array_namespace(logits) + logits = xp.astype(logits, xp.float32) + x_max = xp.max(logits, axis=-1, keepdims=True) + shifted = logits - x_max + log_probs = shifted - xp.log(xp.sum(xp.exp(shifted), axis=-1, keepdims=True)) + target = xp.reshape(target, (-1,)) + keep = target != pad_idx + picked = xp.take_along_axis(log_probs, xp.reshape(target, (-1, 1)), axis=1) + picked = xp.reshape(picked, (-1,)) + picked = xp.where(keep, picked, xp.zeros_like(picked)) + count = xp.astype(xp.sum(xp.astype(keep, logits.dtype)), logits.dtype) + # An empty selection would divide zero by zero; see :func:`_smooth_l1`. + return -xp.sum(picked) / xp.where(count > 0, count, xp.ones_like(count)) + + +@Loss.register("unimol") +class UniMolLoss(Loss): + r"""Uni-Mol v1 self-supervised pretraining loss. + + .. math:: + + L = w_t L_\text{token} + w_c L_\text{coord} + w_d L_\text{dist} + + w_x L_{\|x\|} + w_p L_{\|\Delta p\|} + + Every term is a flat mean over the positions it covers, so molecules with + more corrupted atoms weigh more, exactly as upstream. + + Parameters + ---------- + masked_token_loss : float + Weight of the element-prediction term. + masked_coord_loss : float + Weight of the coordinate-denoising term. + masked_dist_loss : float + Weight of the distance-prediction term. + x_norm_loss : float + Weight of the node-norm regulariser. + delta_pair_repr_norm_loss : float + Weight of the pair-delta-norm regulariser. + beta : float + The transition point of the smooth L1 used by the coordinate and + distance terms. + pad_idx : int + Token id that marks padding, excluded from every term. + """ + + def __init__( + self, + masked_token_loss: float = 1.0, + masked_coord_loss: float = 5.0, + masked_dist_loss: float = 10.0, + x_norm_loss: float = 0.01, + delta_pair_repr_norm_loss: float = 0.01, + beta: float = 1.0, + pad_idx: int = 0, + mask_prob: float = 0.15, + leave_unmasked_prob: float = 0.05, + random_token_prob: float = 0.05, + noise_type: str = "uniform", + noise: float = 1.0, + data_seed: int = 1, + **kwargs: float, + ) -> None: + self.masked_token_loss = masked_token_loss + self.masked_coord_loss = masked_coord_loss + self.masked_dist_loss = masked_dist_loss + self.x_norm_loss = x_norm_loss + self.delta_pair_repr_norm_loss = delta_pair_repr_norm_loss + self.beta = beta + self.pad_idx = pad_idx + # The corruption settings live here because the objective owns them: + # the labels are whatever the corruption produced. + self.mask_prob = mask_prob + self.leave_unmasked_prob = leave_unmasked_prob + self.random_token_prob = random_token_prob + self.noise_type = noise_type + self.noise = noise + self.data_seed = data_seed + + def call( + self, + learning_rate: float, + natoms: int, + model_dict: dict[str, Array], + label_dict: dict[str, Array], + mae: bool = False, + ) -> tuple[Array, dict[str, Array]]: + """Evaluate the five terms and their weighted sum.""" + del learning_rate, natoms, mae + mask = model_dict.get("mask") + token_target = label_dict["unimol_token_target"] + xp = array_api_compat.array_namespace(token_target) + masked = token_target != self.pad_idx + more_loss = {} + loss = None + + def add(term: Array, weight: float, name: str) -> None: + nonlocal loss + more_loss[name] = term + loss = weight * term if loss is None else loss + weight * term + + if self.masked_token_loss > 0: + # The model emits one row per local atom; the objective covers only + # the corrupted ones, so they are gathered here. + logits = model_dict["token_logits"] + if logits.ndim == 3: + logits = logits[masked] + add( + _masked_nll(logits, token_target[masked], self.pad_idx), + self.masked_token_loss, + "token_loss", + ) + if self.masked_coord_loss > 0: + coord_pred = model_dict["coord_update"][masked] + coord_label = label_dict["unimol_coord_target"][masked] + add( + _smooth_l1( + xp.reshape(coord_pred, (-1, 3)), + xp.reshape(coord_label, (-1, 3)), + self.beta, + ), + self.masked_coord_loss, + "coord_loss", + ) + if self.masked_dist_loss > 0: + # Rows are the corrupted atoms; columns are every non-padding token, + # BOS, EOS and the diagonal included (losses/unimol.py:159-180). + ncol = model_dict["pair_dist"].shape[-1] + token_mask = label_dict.get("unimol_token_mask") + if token_mask is None: + token_mask = _token_mask_from_atoms(mask, ncol) + dist_target = label_dict.get("unimol_dist_target") + if dist_target is None: + dist_target = _clean_distances( + label_dict["unimol_coord_target"], mask, ncol + ) + pair_mask = masked[..., None] & xp.astype(token_mask, xp.bool)[:, None, :] + dist_label = (dist_target[pair_mask] - DIST_MEAN) / DIST_STD + add( + _smooth_l1(model_dict["pair_dist"][pair_mask], dist_label, self.beta), + self.masked_dist_loss, + "dist_loss", + ) + if self.x_norm_loss > 0: + add( + _frame_scalar(model_dict["x_norm"], mask), + self.x_norm_loss, + "x_norm_loss", + ) + if self.delta_pair_repr_norm_loss > 0: + add( + _frame_scalar(model_dict["delta_pair_norm"], mask), + self.delta_pair_repr_norm_loss, + "delta_pair_norm_loss", + ) + return loss, more_loss + + def frame_transform(self, type_map: list[str], stream: str = "default"): # noqa: ANN201 + """Build Uni-Mol's corruption, which also produces the labels. + + ``stream`` has to differ per dataset. The draw sequence standing in for + the epoch is derived from ``(data_seed, stream)`` and then kept per + process, so two datasets handed the same label share one generator and a + validation pass advances the corruption training is about to see. A + fresh object per dataset is not enough on its own -- the label is what + separates them. + """ + from deepmd.dpmodel.utils.unimol_transform import ( + make_unimol_data_transform, + ) + + return make_unimol_data_transform( + type_map, + seed=self.data_seed, + stream=stream, + mask_prob=self.mask_prob, + leave_unmasked_prob=self.leave_unmasked_prob, + random_token_prob=self.random_token_prob, + noise_type=self.noise_type, + noise=self.noise, + ) + + @property + def label_requirement(self) -> list[DataRequirementItem]: + """Labels produced by the Uni-Mol data transform, not by a simulation.""" + return [ + DataRequirementItem("unimol_token_target", ndof=1, atomic=True, must=True), + DataRequirementItem("unimol_coord_target", ndof=3, atomic=True, must=True), + ] + + def serialize(self) -> dict: + """Serialize the loss module.""" + return { + "@class": "UniMolLoss", + "@version": 1, + "masked_token_loss": self.masked_token_loss, + "masked_coord_loss": self.masked_coord_loss, + "masked_dist_loss": self.masked_dist_loss, + "x_norm_loss": self.x_norm_loss, + "delta_pair_repr_norm_loss": self.delta_pair_repr_norm_loss, + "beta": self.beta, + "pad_idx": self.pad_idx, + "mask_prob": self.mask_prob, + "leave_unmasked_prob": self.leave_unmasked_prob, + "random_token_prob": self.random_token_prob, + "noise_type": self.noise_type, + "noise": self.noise, + "data_seed": self.data_seed, + } + + @classmethod + def deserialize(cls, data: dict) -> "UniMolLoss": + """Deserialize the loss module.""" + data = data.copy() + check_version_compatibility(data.pop("@version"), 1, 1) + data.pop("@class") + return cls(**data) diff --git a/deepmd/dpmodel/model/__init__.py b/deepmd/dpmodel/model/__init__.py index 462ee802d2..ada9aa58e0 100644 --- a/deepmd/dpmodel/model/__init__.py +++ b/deepmd/dpmodel/model/__init__.py @@ -48,6 +48,9 @@ from .spin_model import ( SpinModel, ) +from .unimol_pretrain_model import ( + UniMolPretrainModel, +) __all__ = [ "DOSModel", @@ -61,5 +64,6 @@ "PolarModel", "PropertyModel", "SpinModel", + "UniMolPretrainModel", "make_model", ] diff --git a/deepmd/dpmodel/model/unimol_pretrain_model.py b/deepmd/dpmodel/model/unimol_pretrain_model.py new file mode 100644 index 0000000000..53cf8a723e --- /dev/null +++ b/deepmd/dpmodel/model/unimol_pretrain_model.py @@ -0,0 +1,121 @@ +# SPDX-License-Identifier: LGPL-3.0-or-later +"""Model wrapper for Uni-Mol v1 self-supervised pretraining.""" + +from typing import ( + Any, +) + +import array_api_compat + +from deepmd.dpmodel.array_api import ( + Array, +) +from deepmd.dpmodel.atomic_model import ( + DPUniMolAtomicModel, +) +from deepmd.dpmodel.common import ( + NativeOP, +) +from deepmd.dpmodel.model.base_model import ( + BaseModel, +) +from deepmd.dpmodel.output_def import ( + OutputVariableDef, +) + +from .dp_model import ( + DPModelCommon, +) +from .make_model import ( + make_model, +) + + +def _reject_periodic(box) -> None: # noqa: ANN001 + """Refuse a periodic cell, with an explanation rather than an allocation.""" + if box is None: + return + xp = array_api_compat.array_namespace(box) + if bool(xp.any(box != 0)): + raise ValueError( + "the unimol descriptor is molecular and does not support periodic " + "boundaries; pass box=None" + ) + + +DPUniMolPretrainModel_ = make_model(DPUniMolAtomicModel, T_Bases=(NativeOP, BaseModel)) + + +@BaseModel.register("unimol_pretrain") +class UniMolPretrainModel(DPModelCommon, DPUniMolPretrainModel_): + r"""Uni-Mol v1 molecular pretraining. + + Predicts the element of every corrupted atom, denoises the coordinates and + predicts the clean pairwise distances, and reports the two norm + regularisers of the backbone. Nothing here reduces to a frame total and + nothing is differentiated with respect to the coordinates: this is + representation learning, not a potential energy surface. + """ + + def __init__(self, *args: Any, **kwargs: Any) -> None: + DPModelCommon.__init__(self) + DPUniMolPretrainModel_.__init__(self, *args, **kwargs) + + def call( + self, + coord: Array, + atype: Array, + box: Array | None = None, + fparam: Array | None = None, + aparam: Array | None = None, + do_atomic_virial: bool = False, + charge_spin: Array | None = None, + ) -> dict[str, Array]: + """Evaluate the pretraining heads on a frame. + + Raises + ------ + ValueError + If a periodic cell is supplied. Uni-Mol is molecular: it has no + cut-off, so a cell would ask the neighbour-list builder for an + astronomical number of images before any other check could fire. + """ + _reject_periodic(box) + model_ret = self.call_common( + coord, + atype, + box, + fparam=fparam, + aparam=aparam, + do_atomic_virial=do_atomic_virial, + charge_spin=charge_spin, + ) + return {k: v for k, v in model_ret.items() if v is not None} + + def call_lower( + self, + extended_coord: Array, + extended_atype: Array, + nlist: Array, + mapping: Array | None = None, + fparam: Array | None = None, + aparam: Array | None = None, + do_atomic_virial: bool = False, + charge_spin: Array | None = None, + ) -> dict[str, Array]: + """Evaluate the pretraining heads on an extended frame.""" + model_ret = self.call_common_lower( + extended_coord, + extended_atype, + nlist, + mapping, + fparam=fparam, + aparam=aparam, + do_atomic_virial=do_atomic_virial, + charge_spin=charge_spin, + ) + return {k: v for k, v in model_ret.items() if v is not None} + + def translated_output_def(self) -> dict[str, OutputVariableDef]: + """The head outputs, passed through under their own names.""" + return dict(self.model_output_def().get_data()) diff --git a/deepmd/dpmodel/utils/lmdb_data.py b/deepmd/dpmodel/utils/lmdb_data.py index 0396640a4a..07ea27370d 100644 --- a/deepmd/dpmodel/utils/lmdb_data.py +++ b/deepmd/dpmodel/utils/lmdb_data.py @@ -690,6 +690,12 @@ class LmdbDecodeConfig: Registered data requirements keyed by field name. dataset Dataset identifier used in frame-level diagnostics. + frame_transform + Optional callable applied to every decoded frame, as + ``transform(frame, frame_index)``. Self-supervised objectives corrupt + their inputs and derive their labels here, before the model runs, which + is the only place that works for a backend whose loss never sees the + model. ``None``, the default, leaves decoding unchanged. """ ntypes: int @@ -697,6 +703,7 @@ class LmdbDecodeConfig: type_remap: np.ndarray | None data_requirements: dict[str, Any] dataset: str = "" + frame_transform: Callable[[dict[str, Any], int], dict[str, Any]] | None = None def _requirement_dtype(requirement: Any) -> np.dtype: @@ -967,6 +974,12 @@ def decode_lmdb_frame( "fid", } ) + if config.frame_transform is not None: + # Ahead of the requirement checks below: a self-supervised transform is + # what produces the fields those checks look for, by corrupting the + # input it was handed. + frame = config.frame_transform(frame, original_key) + for key in list(frame): if key.startswith("find_") or key in structural_keys or key in requirements: continue @@ -2598,6 +2611,17 @@ def print_summary(self, name: str, prob: Any) -> None: def set_noise(self, noise_settings: dict[str, Any]) -> None: """No-op for now.""" + def set_frame_transform( + self, transform: Callable[[dict[str, Any], int], dict[str, Any]] | None + ) -> None: + """Install a per-frame transform, or remove it with ``None``. + + The transform runs on every decoded frame, in whichever process decodes + it, and receives ``(frame, frame_index)``. Self-supervised training uses + it to corrupt inputs and derive labels before the model runs. + """ + self._decode_config.frame_transform = transform + # --- Properties --- @property diff --git a/deepmd/dpmodel/utils/network.py b/deepmd/dpmodel/utils/network.py index 18b89730eb..80a7e591ac 100644 --- a/deepmd/dpmodel/utils/network.py +++ b/deepmd/dpmodel/utils/network.py @@ -25,6 +25,7 @@ Array, xp_add_at, xp_bincount, + xp_erf, xp_setitem_at, xp_sigmoid, ) @@ -350,6 +351,15 @@ def fn(x): # noqa: ANN001, ANN202 * (1 + xp.tanh(xp.sqrt(xp.asarray(2 / xp.pi)) * (x + 0.044715 * x**3))) ) + return fn + elif activation_function == "gelu_erf": + + def fn(x): # noqa: ANN001, ANN202 + xp = array_api_compat.array_namespace(x) + # Exact GELU, x * Phi(x). deepmd's "gelu"/"gelu_tf" are the tanh + # approximation, which differs from this by up to 4.7e-4 per element. + return 0.5 * x * (1 + xp_erf(x / xp.sqrt(xp.asarray(2.0, dtype=x.dtype)))) + return fn elif activation_function == "relu6": diff --git a/deepmd/dpmodel/utils/unimol_transform.py b/deepmd/dpmodel/utils/unimol_transform.py new file mode 100644 index 0000000000..c9d441289d --- /dev/null +++ b/deepmd/dpmodel/utils/unimol_transform.py @@ -0,0 +1,628 @@ +# SPDX-License-Identifier: LGPL-3.0-or-later +"""Data-side transforms of Uni-Mol v1 molecular pretraining. + +Ported from Uni-Mol (https://github.com/deepmodeling/Uni-Mol) at commit 90f52c4, +MIT licensed: + + Copyright (c) DP Technology + This source code is licensed under the MIT license found in the LICENSE + file in the root directory of that source tree. + +Upstream expresses each step as a lazy dataset wrapper +(``unimol/data/*_dataset.py``); here they are plain functions over one frame, so +that a deepmd data loader can call them. The random draws keep upstream's order +and its per-sample seeding, ``hash((seed, epoch, index)) % 1e6``, so a frame +comes out corrupted exactly as upstream would corrupt it. + +Corruption has to happen here rather than inside a loss, because the +PyTorch-Exportable backend runs the model before the loss ever sees a frame. + +The legacy ``numpy.random`` interface is used on purpose, against deepmd's +usual preference for ``np.random.Generator``: upstream seeds the global +legacy PRNG, and a Generator draws a different stream, which would give +different masks and different noise for the same seed. Every such call is +marked with ``# noqa: NPY002``. +""" + +import contextlib +import hashlib +import os +from collections.abc import ( + Callable, + Iterator, + Sequence, +) + +import numpy as np + +__all__ = [ + "UniMolFrameTransform", + "add_bos_eos", + "center_coordinates", + "crop_atoms", + "edge_type", + "make_unimol_data_transform", + "mask_points", + "numpy_seed", + "pair_distance", + "remove_hydrogen", + "reset_epoch_streams", + "sample_conformer", + "unimol_frame_transform", +] + + +@contextlib.contextmanager +def numpy_seed(seed: int | None, *addl_seeds: int) -> Iterator[None]: + """Seed the legacy NumPy PRNG, then restore the previous state. + + Mirrors ``unimol/data/data_utils.py:9``. Note the modulus is ``1e6`` here, + while Uni-Core's own copy uses ``1e8``; the Uni-Mol transforms call this one. + Hashing a tuple of ints is stable across ``PYTHONHASHSEED``. + """ + if seed is None: + yield + return + if len(addl_seeds) > 0: + seed = int(hash((seed, *addl_seeds)) % 1e6) + state = np.random.get_state() # noqa: NPY002 + np.random.seed(seed) # noqa: NPY002 + try: + yield + finally: + np.random.set_state(state) # noqa: NPY002 + + +def sample_conformer( + conformers: Sequence[np.ndarray], seed: int, epoch: int, index: int +) -> np.ndarray: + """Draw one conformer from the pool. + + Mirrors ``ConformerSampleDataset``. Upstream appends an RDKit 2D conformer to + the pool while loading; in deepmd that belongs to the offline converter, so + the pool arrives here already complete and this step needs no RDKit. + """ + with numpy_seed(seed, epoch, index): + sample_idx = np.random.randint(len(conformers)) # noqa: NPY002 + return np.asarray(conformers[sample_idx], dtype=np.float32) + + +def remove_hydrogen( + atoms: np.ndarray, + coordinates: np.ndarray, + remove_hydrogen: bool = False, + remove_polar_hydrogen: bool = False, +) -> tuple[np.ndarray, np.ndarray]: + """Apply the hydrogen policy. + + Mirrors ``RemoveHydrogenDataset``. ``remove_hydrogen`` drops every H; + ``remove_polar_hydrogen`` drops only the trailing run of H, which is what + Uni-Mol calls polar hydrogens. Upstream maps ``only_polar`` to this pair: + -1 keeps all, 0 removes all, 1 removes the trailing run. + """ + atoms = np.asarray(atoms) + if remove_hydrogen: + keep = atoms != "H" + atoms, coordinates = atoms[keep], coordinates[keep] + if not remove_hydrogen and remove_polar_hydrogen: + end_idx = 0 + for i, atom in enumerate(atoms[::-1]): + if atom != "H": + break + end_idx = i + 1 + if end_idx != 0: + atoms, coordinates = atoms[:-end_idx], coordinates[:-end_idx] + return atoms, coordinates.astype(np.float32) + + +def crop_atoms( + atoms: np.ndarray, + coordinates: np.ndarray, + seed: int, + epoch: int, + index: int, + max_atoms: int = 256, +) -> tuple[np.ndarray, np.ndarray]: + """Randomly keep at most ``max_atoms`` atoms. + + Mirrors ``CroppingDataset``. The subset is drawn without replacement and + without regard to where the atoms are in space. + """ + if max_atoms and len(atoms) > max_atoms: + with numpy_seed(seed, epoch, index): + keep = np.random.choice(len(atoms), max_atoms, replace=False) # noqa: NPY002 + atoms, coordinates = np.asarray(atoms)[keep], coordinates[keep] + return atoms, coordinates.astype(np.float32) + + +def center_coordinates(coordinates: np.ndarray) -> np.ndarray: + """Move the centroid to the origin. Mirrors ``NormalizeDataset``.""" + return (coordinates - coordinates.mean(axis=0)).astype(np.float32) + + +def mask_points( + tokens: np.ndarray, + coordinates: np.ndarray, + *, + num_types: int, + special_indices: Sequence[int], + pad_idx: int, + mask_idx: int, + seed: int, + epoch: int, + index: int, + mask_prob: float = 0.15, + leave_unmasked_prob: float = 0.05, + random_token_prob: float = 0.05, + noise_type: str = "uniform", + noise: float = 1.0, +) -> dict[str, np.ndarray]: + """Corrupt a frame the way Uni-Mol pretraining does. + + Mirrors ``MaskPointsDataset.__getitem_cached__``. Of the selected atoms, 90% + become ``[MASK]``, 5% become a random element and 5% are left alone; all + three go into the loss targets. Coordinate noise lands on the masked and the + randomly replaced atoms, not on the ones left alone. + + The order of the random draws matters and is kept: the rounding draw, the + selection, the two split draws, the noise, then the random elements. + + Returns + ------- + dict + ``tokens`` and ``coordinates`` are the corrupted inputs; ``targets`` + holds the true element at every selected position and ``pad_idx`` + elsewhere. + """ + assert 0.0 < mask_prob < 1.0 + assert 0.0 <= random_token_prob <= 1.0 + assert 0.0 <= leave_unmasked_prob <= 1.0 + assert random_token_prob + leave_unmasked_prob <= 1.0 + + weights = None + if random_token_prob > 0.0: + weights = np.ones(num_types, dtype=np.float64) + weights[list(special_indices)] = 0 + weights = weights / weights.sum() + + if noise_type == "trunc_normal": + + def noise_f(n): # noqa: ANN001, ANN202 + return np.clip( + np.random.randn(n, 3) * noise, # noqa: NPY002 + a_min=-noise * 2.0, + a_max=noise * 2.0, + ) + elif noise_type == "normal": + + def noise_f(n): # noqa: ANN001, ANN202 + return np.random.randn(n, 3) * noise # noqa: NPY002 + elif noise_type == "uniform": + + def noise_f(n): # noqa: ANN001, ANN202 + return np.random.uniform(low=-noise, high=noise, size=(n, 3)) # noqa: NPY002 + elif noise_type == "none": + + def noise_f(n): # noqa: ANN001, ANN202 + return 0.0 + else: + # Upstream's fall-through silently adds no noise at all, which turns a + # misspelt setting into a run that trains on clean coordinates. + raise ValueError( + f"unknown noise_type {noise_type!r}; it must be one of " + "'uniform', 'normal', 'trunc_normal' or 'none'" + ) + + with numpy_seed(seed, epoch, index): + sz = len(tokens) + assert sz > 0 + # A random addend rounds the count probabilistically, so a small + # molecule can end up with no masked atom at all. + num_mask = int(mask_prob * sz + np.random.rand()) # noqa: NPY002 + mask_idc = np.random.choice(sz, num_mask, replace=False) # noqa: NPY002 + mask = np.full(sz, False, dtype=bool) + mask[mask_idc] = True + + targets = np.full(len(mask), pad_idx, dtype=np.int64) + targets[mask] = np.asarray(tokens)[mask] + + rand_or_unmask_prob = random_token_prob + leave_unmasked_prob + if rand_or_unmask_prob > 0.0: + rand_or_unmask = mask & (np.random.rand(sz) < rand_or_unmask_prob) # noqa: NPY002 + if random_token_prob == 0.0: + unmask, rand_mask = rand_or_unmask, None + elif leave_unmasked_prob == 0.0: + unmask, rand_mask = None, rand_or_unmask + else: + unmask_prob = leave_unmasked_prob / rand_or_unmask_prob + decision = np.random.rand(sz) < unmask_prob # noqa: NPY002 + unmask = rand_or_unmask & decision + rand_mask = rand_or_unmask & (~decision) + else: + unmask = rand_mask = None + + if unmask is not None: + mask = mask ^ unmask + + new_tokens = np.copy(np.asarray(tokens)) + new_tokens[mask] = mask_idx + + num_mask = mask.astype(np.int32).sum() + new_coord = np.copy(coordinates) + new_coord[mask, :] += noise_f(num_mask) + + if rand_mask is not None: + num_rand = rand_mask.sum() + if num_rand > 0: + new_tokens[rand_mask] = np.random.choice(num_types, num_rand, p=weights) # noqa: NPY002 + + return { + "tokens": new_tokens.astype(np.int64), + "targets": targets.astype(np.int64), + "coordinates": new_coord.astype(np.float32), + } + + +def add_bos_eos( + tokens: np.ndarray, + coordinates: np.ndarray, + targets: np.ndarray, + bos_idx: int, + eos_idx: int, + pad_idx: int, +) -> tuple[np.ndarray, np.ndarray, np.ndarray]: + """Wrap a frame in the two virtual tokens. + + Mirrors ``PrependTokenDataset``/``AppendTokenDataset`` as Uni-Mol calls them + (``tasks/unimol.py:192-206``). Both sit at the origin, which after centering + is the centroid; both take ``pad_idx`` as target, so neither enters the loss. + """ + tokens = np.concatenate([[bos_idx], tokens, [eos_idx]]).astype(np.int64) + targets = np.concatenate([[pad_idx], targets, [pad_idx]]).astype(np.int64) + zero = np.zeros((1, 3), dtype=coordinates.dtype) + coordinates = np.concatenate([zero, coordinates, zero], axis=0) + return tokens, coordinates, targets + + +def pair_distance(coordinates: np.ndarray) -> np.ndarray: + """All-pairs Euclidean distance. Mirrors ``DistanceDataset``.""" + diff = coordinates[:, None, :] - coordinates[None, :, :] + return np.sqrt((diff**2).sum(axis=-1)).astype(np.float32) + + +def edge_type(tokens: np.ndarray, num_types: int) -> np.ndarray: + """Ordered element-pair id, ``t_i * num_types + t_j``. + + Mirrors ``EdgeTypeDataset``. Special tokens take part, so the table has + ``num_types ** 2`` entries. + """ + tokens = np.asarray(tokens) + return (tokens[:, None] * num_types + tokens[None, :]).astype(np.int64) + + +_EPOCH_STREAMS: dict[tuple[int, int], np.random.Generator] = {} + + +def reset_epoch_streams() -> None: + """Forget the per-process draw sequences. + + The generators live in the process so that they survive a transform being + re-created for every batch in a decoder worker. That also means a second + run inside one interpreter continues the first run's sequence rather than + repeating it, which would make the seed reproducible only for whichever run + happened to be first. A trainer calls this as it installs its transforms. + """ + _EPOCH_STREAMS.clear() + + +def _stream_entropy(value: int | str | None) -> int: + """A stable non-negative integer for a seed or a stream label. + + ``hash`` will not do: it is randomized per process for strings, which is + the very thing that would make a run irreproducible. + """ + if value is None: + return 0 + if isinstance(value, int): + return abs(int(value)) + digest = hashlib.blake2b(str(value).encode(), digest_size=8).digest() + return int.from_bytes(digest, "big") + + +def _next_epoch(stream: int, seed: int | None) -> int: + """Draw the number that stands in for upstream's epoch. + + Upstream seeds each sample with ``(seed, epoch, index)``, so a molecule is + corrupted differently in every epoch. There is no epoch to read here, and no + counter would do: decoding runs in worker processes that receive a fresh + copy of the transform for every batch, so anything the transform carries is + reset over and over and the corruption freezes. The generator therefore + lives in the process, keyed by the transform's stream, and each call draws + the next number from it. + + The generator is seeded from the configured seed and the stream alone, so + one process decoding a run twice draws the same sequence both times, given + the reset the trainer performs between runs. + + With more than one decoder process (``DP_LMDB_NUM_WORKERS > 1``) every + worker seeds an identical generator and therefore draws the same sequence + of epoch numbers; frames still differ from one another because the index + enters the per-sample seed, but the draws are no longer decorrelated + between workers, and which frame meets which epoch depends on how the + batches were distributed. The reproducibility the seed offers is for a + single decoding process. + """ + # The process id keys the cache so a forked child builds its own generator + # rather than inheriting a half-consumed one. It deliberately does not enter + # the seed: it changes from run to run, and mixing it in would break the + # reproducibility this seed is supposed to provide. + key = (stream, os.getpid()) + rng = _EPOCH_STREAMS.get(key) + if rng is None: + entropy = [stream] if seed is None else [seed, stream] + _EPOCH_STREAMS[key] = rng = np.random.default_rng(entropy) + return int(rng.integers(1 << 62)) + + +class UniMolFrameTransform: + """The per-frame corruption a deepmd data reader installs. + + The reader hands over one already-converted frame, so the conformer draw, + the hydrogen policy and the size cap are behind us; what remains is centring + and the corruption itself. The frame comes back with corrupted coordinates + and element types, plus the two labels the objective needs. The distance + target is not stored, because it would cost O(natoms^2) per frame; the loss + derives it from the clean coordinates. + + The atom count is left alone. Cropping here would not work: a frame's atom + count and the batch layout are settled before the transform runs, so a + shorter frame would not match the batch it belongs to. The converter applies + the size cap instead. + + This is a class rather than a closure because the reader may decode in + worker processes, which pickle whatever the decoder configuration carries. + + Parameters + ---------- + type_map : Sequence[str] + Element names of the model. It must contain ``mask_token``, since a + masked atom has to be expressible as a type. + seed : int + Together with the frame index and the epoch draw, this seeds the + corruption, the way upstream seeds it per sample and epoch. + mask_token : str + Name of the pseudo-element standing for ``[MASK]``. + **mask_kwargs + Passed to :func:`mask_points`. + """ + + def __init__( + self, + type_map: Sequence[str], + *, + seed: int = 1, + mask_token: str = "[MASK]", + stream: str = "default", + **mask_kwargs: float | str, + ) -> None: + from deepmd.dpmodel.descriptor.unimol import ( + unimol_vocabulary, + ) + + vocabulary = unimol_vocabulary() + token_of = {sym: i for i, sym in enumerate(vocabulary)} + type_map = list(type_map) + if mask_token not in type_map: + raise ValueError( + f"the model type_map must contain {mask_token!r} for Uni-Mol " + "pretraining, because masked atoms are carried as a pseudo-element" + ) + if "epoch" in mask_kwargs: + raise TypeError( + "the epoch is not fixed at build time: the transform draws a new " + "one every time it sees a frame, so that a molecule is corrupted " + "differently each time it comes round" + ) + type_index = {sym: i for i, sym in enumerate(type_map)} + self.vocabulary = vocabulary + self.seed = seed + self.mask_kwargs = mask_kwargs + self.pad = token_of["[PAD]"] + self.mask_idx = token_of[mask_token] + self.type_to_token = np.array( + [token_of.get(sym, token_of["[UNK]"]) for sym in type_map], dtype=np.int64 + ) + self.token_to_type = np.array( + [type_index.get(sym, 0) for sym in vocabulary], dtype=np.int64 + ) + # An element Uni-Mol has no token for tokenizes to [UNK], and [UNK] has + # no element to come back to: the atom would return as [MASK] if it were + # corrupted, and as whatever [UNK] mapped to if it were not. Refuse the + # frame instead of quietly changing an ordinary atom. + self.type_map = type_map + self.untokenizable = np.array( + [i for i, sym in enumerate(type_map) if sym not in token_of], + dtype=np.int64, + ) + # Upstream draws a replacement over all 26 of its elements. A model + # whose type_map covers fewer of them could not express the others, and + # mapping them onto [MASK] would quietly turn a random-element atom into + # a masked one, so they are excluded from the draw instead. With the + # full element set, which is what the example configures, nothing is + # excluded and the distribution is upstream's. + specials = [ + token_of[s] for s in ("[PAD]", "[CLS]", "[SEP]", "[UNK]", mask_token) + ] + self.excluded = [ + *specials, + *( + i + for i, sym in enumerate(vocabulary) + if i not in specials and sym not in type_index + ), + ] + # Identifies this transform's draw sequence within a process, so that + # the training and the validation set do not share one. Derived from the + # configured seed and the caller's label rather than drawn from entropy: + # the same configuration has to corrupt the same way twice. + self.stream = int( + np.random.SeedSequence( + [_stream_entropy(seed), _stream_entropy(stream)] + ).generate_state(1)[0] + ) + + def __call__(self, frame: dict, index: int) -> dict: + """Corrupt one frame and attach the labels the objective reads.""" + coord = np.asarray(frame["coord"], dtype=np.float64).reshape(-1, 3) + atype = np.asarray(frame["atype"], dtype=np.int64).reshape(-1) + if self.untokenizable.size and bool(np.any(np.isin(atype, self.untokenizable))): + unknown = sorted( + {self.type_map[t] for t in atype[np.isin(atype, self.untokenizable)]} + ) + raise ValueError( + f"frame {index} holds element(s) {unknown} that Uni-Mol's " + "vocabulary cannot express; leave those molecules out, or train " + "on a type_map that Uni-Mol covers" + ) + coord = center_coordinates(coord) + tokens = self.type_to_token[atype] + + corrupted = mask_points( + tokens, + coord, + num_types=len(self.vocabulary), + special_indices=self.excluded, + pad_idx=self.pad, + mask_idx=self.mask_idx, + seed=self.seed, + epoch=_next_epoch(self.stream, self.seed), + index=index, + **self.mask_kwargs, + ) + frame = dict(frame) + frame["coord"] = corrupted["coordinates"].astype(np.float64) + # Every token that can come out of the corruption is one the type_map + # expresses: the input was checked above, and the random replacement + # draws from the expressible elements only. + frame["atype"] = self.token_to_type[corrupted["tokens"]] + # Targets stay in Uni-Mol token space, which is what the element head + # predicts over; unselected atoms carry the padding id. + frame["unimol_token_target"] = corrupted["targets"].astype(np.int64) + frame["unimol_coord_target"] = coord.astype(np.float64) + frame["find_unimol_token_target"] = np.float32(1.0) + frame["find_unimol_coord_target"] = np.float32(1.0) + return frame + + +def make_unimol_data_transform( + type_map: Sequence[str], + *, + seed: int = 1, + mask_token: str = "[MASK]", + stream: str = "default", + **mask_kwargs: float | str, +) -> Callable[[dict, int], dict]: + """Build the per-frame transform a deepmd data reader installs. + + See :class:`UniMolFrameTransform`, which this returns. + + Parameters + ---------- + type_map : Sequence[str] + Element names of the model, including ``mask_token``. + seed : int + Seeds the corruption, together with the frame index and the epoch draw. + mask_token : str + Name of the pseudo-element standing for ``[MASK]``. + **mask_kwargs + Passed to :func:`mask_points`. + + Returns + ------- + callable + ``transform(frame, index) -> frame``, matching the reader's hook. + """ + return UniMolFrameTransform( + type_map, seed=seed, mask_token=mask_token, stream=stream, **mask_kwargs + ) + + +def unimol_frame_transform( + atoms: Sequence[str], + conformers: Sequence[np.ndarray], + *, + vocab: dict[str, int], + num_types: int, + special_indices: Sequence[int], + pad_idx: int, + bos_idx: int, + eos_idx: int, + mask_idx: int, + unk_idx: int, + seed: int, + epoch: int, + index: int, + remove_hydrogen_: bool = False, + remove_polar_hydrogen: bool = False, + max_atoms: int = 256, + max_seq_len: int = 512, + **mask_kwargs: float | str, +) -> dict[str, np.ndarray]: + """Run the whole Uni-Mol pretraining chain for one molecule. + + Mirrors ``UniMolTask.load_dataset`` (``tasks/unimol.py:140-245``): sample a + conformer, apply the hydrogen policy, crop, centre, tokenize, corrupt, then + wrap in BOS/EOS and build the distance matrix and edge types. The inputs are + built from the corrupted coordinates and the targets from the clean ones. + """ + coordinates = sample_conformer(conformers, seed, epoch, index) + atoms_arr = np.asarray(atoms) + atoms_arr, coordinates = remove_hydrogen( + atoms_arr, coordinates, remove_hydrogen_, remove_polar_hydrogen + ) + atoms_arr, coordinates = crop_atoms( + atoms_arr, coordinates, seed, epoch, index, max_atoms + ) + coordinates = center_coordinates(coordinates) + + tokens = np.asarray([vocab.get(str(a), unk_idx) for a in atoms_arr], dtype=np.int64) + assert 0 < len(tokens) < max_seq_len + + corrupted = mask_points( + tokens, + coordinates, + num_types=num_types, + special_indices=special_indices, + pad_idx=pad_idx, + mask_idx=mask_idx, + seed=seed, + epoch=epoch, + index=index, + **mask_kwargs, + ) + + src_tokens, src_coord, tokens_target = add_bos_eos( + corrupted["tokens"], + corrupted["coordinates"], + corrupted["targets"], + bos_idx, + eos_idx, + pad_idx, + ) + clean_coord = np.concatenate( + [ + np.zeros((1, 3), dtype=coordinates.dtype), + coordinates, + np.zeros((1, 3), dtype=coordinates.dtype), + ], + axis=0, + ) + return { + "src_tokens": src_tokens, + "src_coord": src_coord, + "src_distance": pair_distance(src_coord), + "src_edge_type": edge_type(src_tokens, num_types), + "tokens_target": tokens_target, + "coord_target": clean_coord, + "distance_target": pair_distance(clean_coord), + } diff --git a/deepmd/pd/utils/utils.py b/deepmd/pd/utils/utils.py index 158b76f0df..1849a4b2d4 100644 --- a/deepmd/pd/utils/utils.py +++ b/deepmd/pd/utils/utils.py @@ -224,6 +224,9 @@ def forward(self, x: paddle.Tensor) -> paddle.Tensor: return F.relu(x) elif self.activation.lower() == "gelu" or self.activation.lower() == "gelu_tf": return F.gelu(x, approximate=True) + elif self.activation.lower() == "gelu_erf": + # Exact GELU; the two names above are the tanh approximation. + return F.gelu(x, approximate=False) elif self.activation.lower() == "tanh": return paddle.tanh(x) elif self.activation.lower() == "relu6": diff --git a/deepmd/pt/utils/utils.py b/deepmd/pt/utils/utils.py index 9f95c59adc..878cf0d713 100644 --- a/deepmd/pt/utils/utils.py +++ b/deepmd/pt/utils/utils.py @@ -198,6 +198,10 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: return F.relu(x) elif self.activation.lower() == "gelu" or self.activation.lower() == "gelu_tf": return F.gelu(x, approximate="tanh") + elif self.activation.lower() == "gelu_erf": + # Exact GELU. "gelu"/"gelu_tf" above are the tanh approximation, + # which differs from this by up to 4.7e-4 per element. + return F.gelu(x, approximate="none") elif self.activation.lower() == "tanh": return torch.tanh(x) elif self.activation.lower() == "relu6": diff --git a/deepmd/pt_expt/descriptor/__init__.py b/deepmd/pt_expt/descriptor/__init__.py index f2718a6df0..672c65c77c 100644 --- a/deepmd/pt_expt/descriptor/__init__.py +++ b/deepmd/pt_expt/descriptor/__init__.py @@ -46,6 +46,9 @@ from .se_t_tebd import ( DescrptSeTTebd, ) +from .unimol import ( + DescrptUniMol, +) __all__ = [ "BaseDescriptor", @@ -60,4 +63,5 @@ "DescrptSeR", "DescrptSeT", "DescrptSeTTebd", + "DescrptUniMol", ] diff --git a/deepmd/pt_expt/descriptor/unimol.py b/deepmd/pt_expt/descriptor/unimol.py new file mode 100644 index 0000000000..abad236ba2 --- /dev/null +++ b/deepmd/pt_expt/descriptor/unimol.py @@ -0,0 +1,41 @@ +# SPDX-License-Identifier: LGPL-3.0-or-later + +from deepmd.dpmodel.descriptor.unimol import DescrptUniMol as DescrptUniMolDP +from deepmd.pt_expt.common import ( + torch_module, +) +from deepmd.pt_expt.descriptor.base_descriptor import ( + BaseDescriptor, +) + + +@BaseDescriptor.register("unimol") +@torch_module +class DescrptUniMol(DescrptUniMolDP): + """Uni-Mol v1 backbone on the PyTorch-Exportable backend.""" + + def share_params( + self, + base_class: "DescrptUniMol", + shared_level: int, + model_prob: float = 1.0, + resume: bool = False, + ) -> None: + """Share parameters with ``base_class`` for multi-task training. + + Level 0 shares the whole backbone, level 1 only the token embedding. + There are no environment statistics to merge, so ``model_prob`` and + ``resume`` play no part. + """ + del model_prob, resume + assert self.__class__ == base_class.__class__, ( + "Only descriptors of the same type can share params!" + ) + if shared_level == 0: + for key in ("gbf", "gbf_proj", "encoder"): + self._modules[key] = base_class._modules[key] + self.embed_tokens = base_class.embed_tokens + elif shared_level == 1: + self.embed_tokens = base_class.embed_tokens + else: + raise NotImplementedError diff --git a/deepmd/pt_expt/fitting/__init__.py b/deepmd/pt_expt/fitting/__init__.py index 8217f64bd2..63bcc6c2ce 100644 --- a/deepmd/pt_expt/fitting/__init__.py +++ b/deepmd/pt_expt/fitting/__init__.py @@ -23,6 +23,9 @@ from .property_fitting import ( PropertyFittingNet, ) +from .unimol_pretrain import ( + UniMolPretrainFitting, +) __all__ = [ "BaseFitting", @@ -33,4 +36,5 @@ "PolarFitting", "PropertyFittingNet", "SeZMEnergyFittingNet", + "UniMolPretrainFitting", ] diff --git a/deepmd/pt_expt/fitting/unimol_pretrain.py b/deepmd/pt_expt/fitting/unimol_pretrain.py new file mode 100644 index 0000000000..0ad5c121a8 --- /dev/null +++ b/deepmd/pt_expt/fitting/unimol_pretrain.py @@ -0,0 +1,18 @@ +# SPDX-License-Identifier: LGPL-3.0-or-later + +from deepmd.dpmodel.fitting.unimol_pretrain import ( + UniMolPretrainFitting as UniMolPretrainFittingDP, +) +from deepmd.pt_expt.common import ( + torch_module, +) + +from .base_fitting import ( + BaseFitting, +) + + +@BaseFitting.register("unimol_pretrain") +@torch_module +class UniMolPretrainFitting(UniMolPretrainFittingDP): + """The Uni-Mol pretraining heads on the PyTorch-Exportable backend.""" diff --git a/deepmd/pt_expt/loss/__init__.py b/deepmd/pt_expt/loss/__init__.py index 77350a8cb0..57e7e4eee9 100644 --- a/deepmd/pt_expt/loss/__init__.py +++ b/deepmd/pt_expt/loss/__init__.py @@ -14,6 +14,9 @@ from deepmd.pt_expt.loss.tensor import ( TensorLoss, ) +from deepmd.pt_expt.loss.unimol import ( + UniMolLoss, +) __all__ = [ "DOSLoss", @@ -21,4 +24,5 @@ "EnergySpinLoss", "PropertyLoss", "TensorLoss", + "UniMolLoss", ] diff --git a/deepmd/pt_expt/loss/unimol.py b/deepmd/pt_expt/loss/unimol.py new file mode 100644 index 0000000000..e456cb877f --- /dev/null +++ b/deepmd/pt_expt/loss/unimol.py @@ -0,0 +1,6 @@ +# SPDX-License-Identifier: LGPL-3.0-or-later +from deepmd.dpmodel.loss.unimol import ( + UniMolLoss, +) + +__all__ = ["UniMolLoss"] diff --git a/deepmd/pt_expt/model/__init__.py b/deepmd/pt_expt/model/__init__.py index b0d7ca1402..22a1c2e7db 100644 --- a/deepmd/pt_expt/model/__init__.py +++ b/deepmd/pt_expt/model/__init__.py @@ -42,6 +42,9 @@ from .spin_ener_model import ( SpinEnergyModel, ) +from .unimol_pretrain_model import ( + UniMolPretrainModel, +) __all__ = [ "BaseModel", @@ -56,6 +59,7 @@ "PolarModel", "PropertyModel", "SpinEnergyModel", + "UniMolPretrainModel", "get_model", "make_hessian_model", ] diff --git a/deepmd/pt_expt/model/unimol_pretrain_model.py b/deepmd/pt_expt/model/unimol_pretrain_model.py new file mode 100644 index 0000000000..f454ef179c --- /dev/null +++ b/deepmd/pt_expt/model/unimol_pretrain_model.py @@ -0,0 +1,97 @@ +# SPDX-License-Identifier: LGPL-3.0-or-later +from typing import ( + Any, +) + +import torch + +from deepmd.dpmodel.atomic_model import ( + DPUniMolAtomicModel, +) +from deepmd.dpmodel.model.dp_model import ( + DPModelCommon, +) +from deepmd.dpmodel.output_def import ( + OutputVariableDef, +) + +from .make_model import ( + make_model, +) +from .model import ( + BaseModel, +) + +DPUniMolPretrainModel_ = make_model(DPUniMolAtomicModel, T_Bases=(BaseModel,)) + + +@BaseModel.register("unimol_pretrain") +class UniMolPretrainModel(DPModelCommon, DPUniMolPretrainModel_): + """Uni-Mol v1 self-supervised pretraining on the PyTorch-Exportable backend.""" + + def __init__(self, *args: Any, **kwargs: Any) -> None: + DPModelCommon.__init__(self) + DPUniMolPretrainModel_.__init__(self, *args, **kwargs) + + def forward( + self, + coord: torch.Tensor, + atype: torch.Tensor, + box: torch.Tensor | None = None, + fparam: torch.Tensor | None = None, + aparam: torch.Tensor | None = None, + do_atomic_virial: bool = False, + charge_spin: torch.Tensor | None = None, + ) -> dict[str, torch.Tensor]: + """Evaluate the pretraining heads on a frame. + + Raises + ------ + ValueError + If a periodic cell is supplied; Uni-Mol is molecular. + """ + if box is not None and bool(torch.any(box != 0)): + raise ValueError( + "the unimol descriptor is molecular and does not support periodic " + "boundaries; pass box=None" + ) + model_ret = self.forward_common( + coord, + atype, + box, + fparam=fparam, + aparam=aparam, + do_atomic_virial=do_atomic_virial, + charge_spin=charge_spin, + ) + return {k: v for k, v in model_ret.items() if v is not None} + + def forward_lower( + self, + extended_coord: torch.Tensor, + extended_atype: torch.Tensor, + nlist: torch.Tensor, + mapping: torch.Tensor | None = None, + fparam: torch.Tensor | None = None, + aparam: torch.Tensor | None = None, + do_atomic_virial: bool = False, + comm_dict: dict[str, torch.Tensor] | None = None, + charge_spin: torch.Tensor | None = None, + ) -> dict[str, torch.Tensor]: + """Evaluate the pretraining heads on an extended frame.""" + model_ret = self.forward_common_lower( + extended_coord, + extended_atype, + nlist, + mapping, + fparam=fparam, + aparam=aparam, + do_atomic_virial=do_atomic_virial, + comm_dict=comm_dict, + charge_spin=charge_spin, + ) + return {k: v for k, v in model_ret.items() if v is not None} + + def translated_output_def(self) -> dict[str, OutputVariableDef]: + """The head outputs, passed through under their own names.""" + return dict(self.model_output_def().get_data()) diff --git a/deepmd/pt_expt/train/training.py b/deepmd/pt_expt/train/training.py index aecfc56b6a..84eef4b5cd 100644 --- a/deepmd/pt_expt/train/training.py +++ b/deepmd/pt_expt/train/training.py @@ -96,6 +96,7 @@ EnergySpinLoss, PropertyLoss, TensorLoss, + UniMolLoss, ) from deepmd.pt_expt.model import ( get_model, @@ -343,6 +344,10 @@ def get_loss( loss_params["var_name"] = var_name loss_params["intensive"] = intensive return PropertyLoss(**loss_params) + elif loss_type == "unimol": + # Self-supervised: it takes no learning rate and no model geometry, + # because its targets come from the corruption it defines itself. + return UniMolLoss(**loss_params) else: raise ValueError(f"Unsupported loss type for pt_expt: {loss_type}") @@ -1861,6 +1866,14 @@ def __init__( self.loss = self.losses if self.multi_task else self.losses[DEFAULT_TASK_KEY] # Data requirements --------------------------------------------------- + # Draw sequences for input-corrupting objectives live in the process, so + # a second run in one interpreter would otherwise continue the first + # run's sequence instead of repeating it. + from deepmd.dpmodel.utils.unimol_transform import ( + reset_epoch_streams, + ) + + reset_epoch_streams() self.valid_numb_batch_by_task: dict[str, int] = {} for model_key in self.model_keys: data_requirement = list(self.losses[model_key].label_requirement) @@ -1872,6 +1885,34 @@ def __init__( self.validation_data_by_task[model_key].add_data_requirements( data_requirement ) + # A self-supervised objective builds its labels by corrupting the + # input, which has to happen as the data is read. + for split, dataset in ( + ("training", self.training_data_by_task[model_key]), + ("validation", self.validation_data_by_task[model_key]), + ): + if dataset is None: + continue + # A fresh transform per dataset is not enough: an objective that + # corrupts its input derives its draw sequence from this label, + # so two datasets given the same one share a generator and a + # validation pass advances the corruption training is about to + # see. The task name is in it too, so two Uni-Mol tasks in one + # multi-task run do not collide either. + frame_transform = self.losses[model_key].frame_transform( + self.model_params_by_task[model_key]["type_map"], + stream=f"{model_key}/{split}", + ) + if frame_transform is None: + break + if not hasattr(dataset, "set_frame_transform"): + raise ValueError( + f"the {self.losses[model_key].__class__.__name__} " + "objective corrupts its input as the data is read, " + "which this dataset type does not support; convert " + "the data to LMDB first" + ) + dataset.set_frame_transform(frame_transform) if self.multi_task: valid_params = ( training_params["data_dict"][model_key].get("validation_data", {}) @@ -2252,6 +2293,7 @@ def update_finetune_bias( torch.optim.Adam if opt_type == "Adam" else torch.optim.AdamW, lr=initial_lr, betas=adam_betas, + eps=float(optimizer_params["adam_eps"]), weight_decay=weight_decay, ) else: diff --git a/deepmd/pt_expt/utils/lmdb_dataset.py b/deepmd/pt_expt/utils/lmdb_dataset.py index 9af91be82c..f16de39a01 100644 --- a/deepmd/pt_expt/utils/lmdb_dataset.py +++ b/deepmd/pt_expt/utils/lmdb_dataset.py @@ -211,6 +211,10 @@ def add_data_requirements( self._reader.add_data_requirement(data_requirement) self._refresh_stat_groups() + def set_frame_transform(self, transform) -> None: # noqa: ANN001 + """Install a per-frame transform on the underlying reader.""" + self._reader.set_frame_transform(transform) + def close(self) -> None: """Cancel prefetched work and release decoder processes.""" iterator = getattr(self, "_batch_iterator", None) diff --git a/deepmd/pt_expt/utils/network.py b/deepmd/pt_expt/utils/network.py index adc7d4f326..651c8bb4f8 100644 --- a/deepmd/pt_expt/utils/network.py +++ b/deepmd/pt_expt/utils/network.py @@ -181,6 +181,9 @@ def _torch_activation(x: torch.Tensor, name: str) -> torch.Tensor: return torch.relu(x) elif name in ("gelu", "gelu_tf"): return torch.nn.functional.gelu(x, approximate="tanh") + elif name == "gelu_erf": + # Exact GELU; the two names above are the tanh approximation. + return torch.nn.functional.gelu(x, approximate="none") elif name == "relu6": return torch.clamp(x, min=0.0, max=6.0) elif name == "softplus": diff --git a/deepmd/tf/common.py b/deepmd/tf/common.py index 1a37769dfe..aa9abbc09f 100644 --- a/deepmd/tf/common.py +++ b/deepmd/tf/common.py @@ -186,6 +186,31 @@ def silut(x: tf.Tensor) -> tf.Tensor: return silut +def gelu_erf(x: tf.Tensor) -> tf.Tensor: + """Exact Gaussian Error Linear Unit. + + Unlike :func:`gelu` and :func:`gelu_tf`, which are the tanh approximation, + this evaluates ``x * Phi(x)`` through the error function. The two forms + differ by up to 4.7e-4 per element. + + Parameters + ---------- + x : tf.Tensor + float Tensor to perform activation + + Returns + ------- + tf.Tensor + `x` with the exact GELU activation applied + + References + ---------- + Original paper + https://arxiv.org/abs/1606.08415 + """ + return 0.5 * x * (1.0 + tf.math.erf(x / tf.sqrt(tf.cast(2.0, x.dtype)))) + + ACTIVATION_FN_DICT = { "relu": tf.nn.relu, "relu6": tf.nn.relu6, @@ -194,6 +219,7 @@ def silut(x: tf.Tensor) -> tf.Tensor: "tanh": tf.nn.tanh, "gelu": gelu, "gelu_tf": gelu_tf, + "gelu_erf": gelu_erf, "silu": silu, "silut": get_silut("silut"), "linear": lambda x: x, diff --git a/deepmd/utils/argcheck.py b/deepmd/utils/argcheck.py index 65ab8a85a2..f5de3a02bd 100644 --- a/deepmd/utils/argcheck.py +++ b/deepmd/utils/argcheck.py @@ -2753,6 +2753,144 @@ def descrpt_se_a_mask_args() -> list[Argument]: ] +doc_descrpt_unimol = ( + "The Uni-Mol v1 transformer backbone. Every atom attends to every other atom, so the " + "descriptor is neither local nor extensive and does not support periodic boundaries; it is " + "intended for molecular property and self-supervised pretraining work." +) +doc_fitting_unimol_pretrain = ( + "The three self-supervised heads of Uni-Mol v1 pretraining: element prediction, coordinate " + "denoising and pairwise distance prediction." +) + + +@descrpt_args_plugin.register( + "unimol", doc=supported_backends("pt_expt") + doc_descrpt_unimol +) +def descrpt_unimol_args() -> list[Argument]: + doc_seed = "Random seed for parameter initialization" + doc_precision = f"The precision of the parameters, supported options are {list_to_doc(PRECISION_DICT.keys())} Default follows the interface precision." + doc_encoder_layers = "Number of transformer blocks." + doc_encoder_embed_dim = "Width of the node representation." + doc_encoder_ffn_embed_dim = "Width of the feed-forward hidden layer." + doc_encoder_attention_heads = ( + "Number of attention heads, which is also the width of the pair channel." + ) + doc_max_atoms = "Largest molecule accepted, which fixes the neighbor count." + doc_max_seq_len = ( + "Upstream's sequence-length guard, including the two virtual tokens." + ) + doc_activation_function = ( + f"The activation function. Uni-Mol uses the exact GELU, `gelu_erf`. " + f"Supported: {list_to_doc(ACTIVATION_FN_DICT.keys())}" + ) + doc_dropout = "Dropout on the residual branches, applied during training." + doc_emb_dropout = "Dropout on the token embedding, applied during training." + doc_attention_dropout = ( + "Dropout on the attention probabilities, applied during training." + ) + doc_activation_dropout = "Dropout after the feed-forward activation." + doc_no_final_head_layer_norm = "Skip the layer norm on the pair delta. Upstream builds it unless its loss weight is negative." + doc_single_precision_basis = "Evaluate the Gaussian basis in fp32, as upstream does. Required to reproduce the released weights." + doc_single_precision_distance = ( + "Round pairwise distances to fp32 before the basis, matching upstream's precomputed distance " + "matrix. The default computes them in the working precision, which is more accurate." + ) + doc_virtual_token_position = ( + "Where the two virtual tokens sit: `centroid` of the real atoms, which keeps the sequence " + "translation invariant, or `origin`, which reproduces upstream for pre-centred data." + ) + doc_gaussian_kernels = "Number of Gaussian radial basis functions." + return [ + Argument( + "encoder_layers", int, optional=True, default=15, doc=doc_encoder_layers + ), + Argument( + "encoder_embed_dim", + int, + optional=True, + default=512, + doc=doc_encoder_embed_dim, + ), + Argument( + "encoder_ffn_embed_dim", + int, + optional=True, + default=2048, + doc=doc_encoder_ffn_embed_dim, + ), + Argument( + "encoder_attention_heads", + int, + optional=True, + default=64, + doc=doc_encoder_attention_heads, + ), + Argument("max_atoms", int, optional=True, default=256, doc=doc_max_atoms), + Argument("max_seq_len", int, optional=True, default=512, doc=doc_max_seq_len), + Argument( + "activation_function", + str, + optional=True, + default="gelu_erf", + doc=doc_activation_function, + ), + Argument("dropout", float, optional=True, default=0.1, doc=doc_dropout), + Argument("emb_dropout", float, optional=True, default=0.1, doc=doc_emb_dropout), + Argument( + "attention_dropout", + float, + optional=True, + default=0.1, + doc=doc_attention_dropout, + ), + Argument( + "activation_dropout", + float, + optional=True, + default=0.0, + doc=doc_activation_dropout, + ), + Argument( + "no_final_head_layer_norm", + bool, + optional=True, + default=False, + doc=doc_no_final_head_layer_norm, + ), + Argument( + "single_precision_basis", + bool, + optional=True, + default=True, + doc=doc_single_precision_basis, + ), + Argument( + "single_precision_distance", + bool, + optional=True, + default=False, + doc=doc_single_precision_distance, + ), + Argument( + "virtual_token_position", + str, + optional=True, + default="centroid", + doc=doc_virtual_token_position, + ), + Argument( + "gaussian_kernels", + int, + optional=True, + default=128, + doc=doc_gaussian_kernels, + ), + Argument("precision", str, optional=True, default="default", doc=doc_precision), + Argument("seed", [int, list, None], optional=True, doc=doc_seed), + ] + + def descrpt_variant_type_args(exclude_hybrid: bool = False) -> Variant: doc_descrpt_type = "The type of the descriptor." @@ -3313,6 +3451,52 @@ def fitting_dipole() -> list[Argument]: # YWolfeee: Delete global polar mode, merge it into polar mode and use loss setting to support. +@fitting_args_plugin.register( + "unimol_pretrain", doc=supported_backends("pt_expt") + doc_fitting_unimol_pretrain +) +def fitting_unimol_pretrain() -> list[Argument]: + doc_seed = "Random seed for parameter initialization" + doc_precision = f"The precision of the parameters, supported options are {list_to_doc(PRECISION_DICT.keys())} Default follows the interface precision." + doc_n_token = ( + "Size of the Uni-Mol vocabulary, which is the width of the element head." + ) + doc_attention_heads = ( + "Width of the pair channel, which the two pair-reading heads consume." + ) + doc_max_atoms = ( + "Largest molecule accepted, which fixes the distance head's column count." + ) + doc_activation_function = f"The activation function of the heads. Supported: {list_to_doc(ACTIVATION_FN_DICT.keys())}" + doc_mask_token_head = "Build the element-prediction head." + doc_coord_head = "Build the coordinate-denoising head." + doc_dist_head = "Build the distance-prediction head." + return [ + Argument("n_token", int, optional=True, default=31, doc=doc_n_token), + Argument( + "attention_heads", int, optional=True, default=64, doc=doc_attention_heads + ), + Argument("max_atoms", int, optional=True, default=256, doc=doc_max_atoms), + Argument( + "activation_function", + str, + optional=True, + default="gelu_erf", + doc=doc_activation_function, + ), + Argument( + "mask_token_head", + bool, + optional=True, + default=True, + doc=doc_mask_token_head, + ), + Argument("coord_head", bool, optional=True, default=True, doc=doc_coord_head), + Argument("dist_head", bool, optional=True, default=True, doc=doc_dist_head), + Argument("precision", str, optional=True, default="default", doc=doc_precision), + Argument("seed", [int, list, None], optional=True, doc=doc_seed), + ] + + def fitting_variant_type_args() -> Variant: doc_descrpt_type = "The type of the fitting." @@ -4307,6 +4491,10 @@ def _check_lr_args(data: dict[str, Any]) -> bool: def optimizer_adam() -> list[Argument]: doc_adam_beta1 = "Adam beta1 coefficient for first moment decay." doc_adam_beta2 = "Adam beta2 coefficient for second moment decay." + doc_adam_eps = ( + "Adam epsilon, added for numerical stability. The default is PyTorch's own; " + "recipes carried over from other frameworks sometimes assume a different one." + ) doc_weight_decay = ( "Weight decay coefficient for Adam, applied as an L2 penalty to gradients." ) @@ -4325,6 +4513,13 @@ def optimizer_adam() -> list[Argument]: default=0.999, doc=supported_backends("tf", "pt", "pd", "tf2") + doc_adam_beta2, ), + Argument( + "adam_eps", + float, + optional=True, + default=1e-8, + doc=supported_backends("pt_expt") + doc_adam_eps, + ), Argument( "weight_decay", float, @@ -4339,6 +4534,10 @@ def optimizer_adam() -> list[Argument]: def optimizer_adamw() -> list[Argument]: doc_adam_beta1 = "AdamW beta1 coefficient for first moment decay." doc_adam_beta2 = "AdamW beta2 coefficient for second moment decay." + doc_adam_eps = ( + "AdamW epsilon, added for numerical stability. The default is PyTorch's own; " + "recipes carried over from other frameworks sometimes assume a different one." + ) doc_weight_decay = "Decoupled weight decay coefficient for the AdamW optimizer." return [ Argument( @@ -4355,6 +4554,13 @@ def optimizer_adamw() -> list[Argument]: default=0.999, doc=supported_backends("pt", "pd", "tf2") + doc_adam_beta2, ), + Argument( + "adam_eps", + float, + optional=True, + default=1e-8, + doc=supported_backends("pt_expt") + doc_adam_eps, + ), Argument( "weight_decay", float, @@ -5332,6 +5538,98 @@ def loss_tensor() -> list[Argument]: ] +@loss_args_plugin.register("unimol", doc=supported_backends("pt_expt")) +def loss_unimol() -> list[Argument]: + doc_masked_token_loss = "Weight of the element-prediction term." + doc_masked_coord_loss = "Weight of the coordinate-denoising term." + doc_masked_dist_loss = "Weight of the distance-prediction term." + doc_x_norm_loss = "Weight of the node-norm regularizer." + doc_delta_pair_repr_norm_loss = "Weight of the pair-delta-norm regularizer." + doc_beta = ( + "Transition point of the smooth L1 used by the coordinate and distance terms." + ) + doc_mask_prob = "Expected fraction of atoms selected for corruption." + doc_leave_unmasked_prob = ( + "Fraction of the selected atoms left with their true element, still predicted." + ) + doc_random_token_prob = "Fraction of the selected atoms given a random element." + doc_noise_type = ( + "Coordinate noise distribution: 'uniform', 'normal', 'trunc_normal' or 'none'." + ) + doc_noise = "Scale of the coordinate noise, in the units of the coordinates." + doc_data_seed = ( + "Seed of the corruption, combined with the frame index and a per-visit " + "draw so that a molecule is corrupted differently each time it comes " + "round. A run is reproducible from it only when one process decodes the " + "data (DP_LMDB_NUM_WORKERS=0); with decoder workers the draws follow how " + "frames were distributed." + ) + return [ + Argument( + "masked_token_loss", + [float, int], + optional=True, + default=1.0, + doc=doc_masked_token_loss, + ), + Argument( + "masked_coord_loss", + [float, int], + optional=True, + default=5.0, + doc=doc_masked_coord_loss, + ), + Argument( + "masked_dist_loss", + [float, int], + optional=True, + default=10.0, + doc=doc_masked_dist_loss, + ), + Argument( + "x_norm_loss", + [float, int], + optional=True, + default=0.01, + doc=doc_x_norm_loss, + ), + Argument( + "delta_pair_repr_norm_loss", + [float, int], + optional=True, + default=0.01, + doc=doc_delta_pair_repr_norm_loss, + ), + Argument("beta", [float, int], optional=True, default=1.0, doc=doc_beta), + Argument( + "mask_prob", [float, int], optional=True, default=0.15, doc=doc_mask_prob + ), + Argument( + "leave_unmasked_prob", + [float, int], + optional=True, + default=0.05, + doc=doc_leave_unmasked_prob, + ), + Argument( + "random_token_prob", + [float, int], + optional=True, + default=0.05, + doc=doc_random_token_prob, + ), + Argument( + "noise_type", + str, + optional=True, + default="uniform", + doc=doc_noise_type, + ), + Argument("noise", [float, int], optional=True, default=1.0, doc=doc_noise), + Argument("data_seed", int, optional=True, default=1, doc=doc_data_seed), + ] + + def loss_variant_type_args() -> Variant: doc_loss = "The type of the loss. When the fitting type is `ener`, the loss type should be set to `ener`, its legacy alias `ener_hess`, `dens` (Only DPA4/SeZM supported), or left unset. Hessian supervision is configured through `start_pref_h` and `limit_pref_h` on the `ener` loss. When the fitting type is `property`, the loss type should be set to `property`. When the fitting type is `dipole` or `polar`, the loss type should be set to `tensor`." diff --git a/deepmd/utils/unimol_checkpoint.py b/deepmd/utils/unimol_checkpoint.py new file mode 100644 index 0000000000..e1b6b8b692 --- /dev/null +++ b/deepmd/utils/unimol_checkpoint.py @@ -0,0 +1,222 @@ +# SPDX-License-Identifier: LGPL-3.0-or-later +"""Import released Uni-Mol v1 checkpoints into deepmd. + +The released files (``mol_pre_all_h_220816.pt`` and ``mol_pre_no_h_220816.pt``, +MIT licensed, from the Uni-Mol repository and its Hugging Face mirror) hold a +single ``model`` entry with 211 float32 tensors and no training state, so they +load with ``weights_only=True``. + +The parameter names line up one to one with the ported modules, but the arrays +do not: deepmd stores a linear weight as ``(num_in, num_out)`` and applies it as +``x @ w``, which is the transpose of ``torch.nn.Linear.weight``, and it calls +layer-norm parameters ``w``/``b`` rather than ``weight``/``bias``. Every weight +is therefore renamed and transposed here rather than loaded directly. +""" + +from typing import ( + TYPE_CHECKING, + Any, +) + +import numpy as np + +if TYPE_CHECKING: + from deepmd.dpmodel.descriptor.unimol import ( + DescrptUniMol, + ) + +__all__ = [ + "UNIMOL_V1_BASE_ARCHITECTURE", + "descriptor_from_unimol_checkpoint", + "load_unimol_state_dict", + "split_unimol_state_dict", +] + +# Uni-Mol's ``unimol_base`` architecture (unimol/models/unimol.py:423-442). +UNIMOL_V1_BASE_ARCHITECTURE = { + "encoder_layers": 15, + "encoder_embed_dim": 512, + "encoder_ffn_embed_dim": 2048, + "encoder_attention_heads": 64, + "activation_function": "gelu_erf", + "dropout": 0.1, + "emb_dropout": 0.1, + "attention_dropout": 0.1, + "activation_dropout": 0.0, + "max_seq_len": 512, +} + +_BACKBONE_PREFIXES = ("embed_tokens", "gbf.", "gbf_proj.", "encoder.") +_HEAD_PREFIXES = ("lm_head.", "pair2coord_proj.", "dist_head.") + + +def load_unimol_state_dict(path: str) -> dict[str, np.ndarray]: + """Read a released checkpoint into NumPy arrays. + + Parameters + ---------- + path : str + Local path to the ``.pt`` file. Nothing is downloaded here; see + :func:`descriptor_from_unimol_checkpoint` for the optional download. + + Returns + ------- + dict[str, np.ndarray] + The ``model`` entry of the checkpoint. + """ + try: + import torch + except ImportError as e: + raise ImportError( + "reading a Uni-Mol checkpoint needs PyTorch; install it with " + "`pip install torch`" + ) from e + obj = torch.load(path, map_location="cpu", weights_only=True) + state = obj["model"] if "model" in obj else obj + return {k: v.numpy() for k, v in state.items()} + + +def split_unimol_state_dict( + state_dict: dict[str, np.ndarray], +) -> tuple[dict[str, np.ndarray], dict[str, np.ndarray]]: + """Split a checkpoint into backbone and pretraining-head parameters. + + Returns + ------- + tuple + ``(backbone, heads)``. The heads are the element, coordinate and + distance heads, which belong to the fitting rather than the descriptor. + """ + backbone = {k: v for k, v in state_dict.items() if k.startswith(_BACKBONE_PREFIXES)} + heads = {k: v for k, v in state_dict.items() if k.startswith(_HEAD_PREFIXES)} + unknown = set(state_dict) - set(backbone) - set(heads) + if unknown: + raise ValueError(f"unexpected parameters in the checkpoint: {sorted(unknown)}") + return backbone, heads + + +def _set_table(layer: Any, value: np.ndarray) -> None: + """Assign a lookup table, which is stored as a layer weight.""" + layer.w = np.ascontiguousarray(value).astype(layer.w.dtype) + + +def _set_linear(layer: Any, state: dict[str, np.ndarray], prefix: str) -> None: + """Assign a torch Linear onto a deepmd layer, transposing the weight.""" + layer.w = np.ascontiguousarray(state[prefix + ".weight"].T).astype(layer.w.dtype) + bias = state.get(prefix + ".bias") + if bias is not None and layer.b is not None: + layer.b = bias.astype(layer.b.dtype) + + +def _set_layer_norm(layer: Any, state: dict[str, np.ndarray], prefix: str) -> None: + """Assign a torch LayerNorm onto a deepmd layer norm.""" + layer.w = state[prefix + ".weight"].astype(layer.w.dtype) + layer.b = state[prefix + ".bias"].astype(layer.b.dtype) + + +def _check_architecture( + descriptor: "DescrptUniMol", state_dict: dict[str, np.ndarray] +) -> None: + """Refuse a descriptor whose shape does not match the checkpoint. + + Architecture overrides are easy to get wrong, and a mismatch would + otherwise pass silently: too few layers would ignore the rest of the + checkpoint, and a different width would be caught only as a dtype or shape + error somewhere deep in the first forward pass. + """ + layers = len( + {k.split(".")[2] for k in state_dict if k.startswith("encoder.layers.")} + ) + if layers != len(descriptor.encoder.layers): + raise ValueError( + f"the checkpoint holds {layers} encoder layers but the descriptor " + f"was built with {len(descriptor.encoder.layers)}" + ) + expected = { + "embed_tokens.weight": descriptor.embed_tokens.w.shape, + "encoder.layers.0.fc1.weight": descriptor.encoder.layers[0].fc1.w.shape[::-1], + "gbf.mul.weight": descriptor.gbf.mul.w.shape, + "gbf.means.weight": descriptor.gbf.means.w.shape, + } + for key, shape in expected.items(): + found = tuple(state_dict[key].shape) + if found != tuple(shape): + raise ValueError( + f"the checkpoint's {key} has shape {found}, but the descriptor " + f"expects {tuple(shape)}; check the architecture overrides and " + "the type_map against the checkpoint" + ) + + +def apply_unimol_backbone( + descriptor: "DescrptUniMol", state_dict: dict[str, np.ndarray] +) -> None: + """Load backbone parameters into a descriptor, in place.""" + _check_architecture(descriptor, state_dict) + # The embedding and the four basis tables are layers, so their values live + # in ``w``; they carry no transpose because they are lookup tables rather + # than projections. + _set_table(descriptor.embed_tokens, state_dict["embed_tokens.weight"]) + for name in ("means", "stds", "mul", "bias"): + _set_table(getattr(descriptor.gbf, name), state_dict[f"gbf.{name}.weight"]) + _set_linear(descriptor.gbf_proj.linear1, state_dict, "gbf_proj.linear1") + _set_linear(descriptor.gbf_proj.linear2, state_dict, "gbf_proj.linear2") + + encoder = descriptor.encoder + _set_layer_norm(encoder.emb_layer_norm, state_dict, "encoder.emb_layer_norm") + _set_layer_norm(encoder.final_layer_norm, state_dict, "encoder.final_layer_norm") + if encoder.final_head_layer_norm is not None: + _set_layer_norm( + encoder.final_head_layer_norm, state_dict, "encoder.final_head_layer_norm" + ) + for i, layer in enumerate(encoder.layers): + prefix = f"encoder.layers.{i}." + _set_linear(layer.self_attn.in_proj, state_dict, prefix + "self_attn.in_proj") + _set_linear(layer.self_attn.out_proj, state_dict, prefix + "self_attn.out_proj") + _set_layer_norm( + layer.self_attn_layer_norm, state_dict, prefix + "self_attn_layer_norm" + ) + _set_linear(layer.fc1, state_dict, prefix + "fc1") + _set_linear(layer.fc2, state_dict, prefix + "fc2") + _set_layer_norm(layer.final_layer_norm, state_dict, prefix + "final_layer_norm") + + +def descriptor_from_unimol_checkpoint( + path: str, + type_map: list[str] | None = None, + precision: str = "float64", + **overrides: Any, +) -> "DescrptUniMol": + """Build a descriptor carrying the released Uni-Mol v1 weights. + + Parameters + ---------- + path : str + Local path to the checkpoint. + type_map : list[str], optional + Element names for the model. Defaults to Uni-Mol's own 26 elements. + precision : str + Precision to hold the parameters in. + **overrides + Architecture overrides on top of ``UNIMOL_V1_BASE_ARCHITECTURE``. + + Returns + ------- + DescrptUniMol + A descriptor whose backbone reproduces upstream. + """ + from deepmd.dpmodel.descriptor.unimol import ( + UNIMOL_ELEMENTS, + DescrptUniMol, + ) + + state_dict = load_unimol_state_dict(path) + backbone, _ = split_unimol_state_dict(state_dict) + arch = {**UNIMOL_V1_BASE_ARCHITECTURE, **overrides} + descriptor = DescrptUniMol( + type_map=type_map if type_map is not None else list(UNIMOL_ELEMENTS), + precision=precision, + **arch, + ) + apply_unimol_backbone(descriptor, backbone) + return descriptor diff --git a/deepmd/utils/unimol_data.py b/deepmd/utils/unimol_data.py new file mode 100644 index 0000000000..712997e7d0 --- /dev/null +++ b/deepmd/utils/unimol_data.py @@ -0,0 +1,306 @@ +# SPDX-License-Identifier: LGPL-3.0-or-later +"""Convert Uni-Mol pretraining data into a deepmd LMDB dataset. + +Uni-Mol ships its molecular pretraining set as a single LMDB file whose values +are pickled dicts, ``{"atoms": [element symbols], "coordinates": [n x 3 arrays], +"smi": str}``, with roughly ten RDKit conformers per molecule. deepmd reads a +different LMDB layout, so the data is converted once, offline, rather than +taught to the training loop. + +One conformer becomes one frame, which is what makes the conversion streaming +and lets ordinary frame sampling stand in for Uni-Mol's per-epoch conformer +draw. Frames of the same molecule share a system id, so they can be kept +together if needed. + +Upstream also appends a two-dimensional RDKit conformer to the pool while +loading. That belongs here rather than in the training loop, so the data path +never needs RDKit; pass ``add_2d_conformer=True`` to reproduce it. +""" + +import argparse +import logging +import os +import pickle +import shutil +from collections.abc import ( + Iterator, +) +from typing import ( + Any, +) + +import numpy as np + +__all__ = ["convert_unimol_lmdb", "read_unimol_lmdb"] + +log = logging.getLogger(__name__) + + +def _encode_array(arr: np.ndarray) -> dict[str, Any]: + """Encode an array the way a deepmd LMDB frame stores it. + + Three keys, matching the datasets already published in this format; the + reader takes the dtype and shape from here and casts on the way out. + """ + return { + "type": str(arr.dtype), + "shape": list(arr.shape), + "data": arr.tobytes(), + } + + +def read_unimol_lmdb(path: str) -> Iterator[dict[str, Any]]: + """Stream molecules out of a Uni-Mol LMDB file. + + .. warning:: + Upstream stores each record as a Python pickle, and unpickling runs + whatever the file says. Convert only files you obtained from a source you + trust; the conversion this module performs is one-off, and the deepmd + dataset it writes is msgpack, which carries no such risk. + + Parameters + ---------- + path : str + Path to the ``.lmdb`` file, which upstream writes as a single file + rather than a directory. + + Yields + ------ + dict + ``atoms``, ``coordinates`` and ``smi`` as stored upstream. + """ + import lmdb + + env = lmdb.open( + path, subdir=False, readonly=True, lock=False, readahead=False, meminit=False + ) + try: + with env.begin() as txn: + cursor = txn.cursor() + for _, value in cursor: + # See the warning above: this executes whatever the file + # says, because that is the format upstream published. + yield pickle.loads(value) + finally: + env.close() + + +def _conformers(record: dict[str, Any], add_2d_conformer: bool) -> list[np.ndarray]: + conformers = [np.asarray(c, dtype=np.float32) for c in record["coordinates"]] + if add_2d_conformer: + from rdkit import ( + Chem, + ) + from rdkit.Chem import ( + AllChem, + ) + + mol = Chem.AddHs(Chem.MolFromSmiles(record["smi"])) + AllChem.Compute2DCoords(mol) + coords = mol.GetConformer().GetPositions().astype(np.float32) + conformers.append(coords[: len(record["atoms"])]) + return conformers + + +def convert_unimol_lmdb( + src: str, + dst: str, + type_map: list[str] | None = None, + add_2d_conformer: bool = False, + max_molecules: int | None = None, + max_conformers: int | None = None, + max_atoms: int | None = 256, + map_size: int = 1024**4, +) -> dict[str, int]: + """Write a deepmd LMDB dataset from a Uni-Mol one. + + Parameters + ---------- + src : str + The Uni-Mol ``.lmdb`` file. + dst : str + Directory for the deepmd dataset, replaced if it exists. + type_map : list[str], optional + Element names for the output. Defaults to Uni-Mol's own 26 elements. + add_2d_conformer : bool + Append the RDKit two-dimensional conformer to each pool, as upstream + does while loading. Needs RDKit. + max_molecules : int, optional + Stop after this many molecules, which is useful for a trial run. + max_conformers : int, optional + Keep at most this many conformers per molecule. + max_atoms : int, optional + Crop molecules larger than this to a random subset of that many atoms, + as upstream does. It happens here rather than while training, because a + frame's atom count and the batch layout are settled before any per-frame + transform runs. + map_size : int + Maximum size of the output database. + + Returns + ------- + dict + Counts of molecules, frames and skipped records. + """ + import lmdb + import msgpack + + from deepmd.dpmodel.descriptor.unimol import ( + UNIMOL_ELEMENTS, + ) + + names = list(type_map) if type_map is not None else list(UNIMOL_ELEMENTS) + index_of = {sym: i for i, sym in enumerate(names)} + + # Build beside the destination and move it into place at the end. A source + # error, a malformed record or an RDKit failure part way through would + # otherwise have destroyed a dataset that took hours to write and left + # nothing, or half of something, in its place. + staging = f"{dst}.partial" + if os.path.exists(staging): + shutil.rmtree(staging) + env = lmdb.open(staging, map_size=map_size) + fmt = "012d" + frame_idx = 0 + molecules = 0 + skipped = 0 + frame_system_ids: list[int] = [] + frame_nlocs: list[int] = [] + + try: + with env.begin(write=True) as txn: + for record in read_unimol_lmdb(src): + if max_molecules is not None and molecules >= max_molecules: + break + atoms = [str(a) for a in record["atoms"]] + if len(atoms) < 2 or any(a not in index_of for a in atoms): + # The descriptor cannot tell a single real atom from + # padding, and an unmapped element would silently become + # [UNK]; skip both rather than write something misleading. + skipped += 1 + continue + if max_atoms is not None and len(atoms) > max_atoms: + keep = np.sort( + np.random.default_rng(molecules).choice( + len(atoms), max_atoms, replace=False + ) + ) + atoms = [atoms[i] for i in keep] + else: + keep = None + atom_types = np.array([index_of[a] for a in atoms], dtype=np.int64) + atom_numbs = [int((atom_types == i).sum()) for i in range(len(names))] + pool = _conformers(record, add_2d_conformer) + if max_conformers is not None: + pool = pool[:max_conformers] + for coords in pool: + if keep is not None and coords.shape[0] >= int(keep[-1]) + 1: + coords = coords[keep] + if coords.shape[0] != len(atoms): + skipped += 1 + continue + frame = { + "atom_numbs": atom_numbs, + # int32 types and float32 coordinates, which is what the + # datasets already published in this format use, and + # what the source holds: upstream generated these + # conformers in single precision, so widening them here + # would store zeros. The reader casts to the precision + # the model asks for. + "atom_types": _encode_array(atom_types.astype(np.int32)), + # No cell at all: molecules are not periodic, and a zero + # cell would be taken for a real one and inverted. + "coords": _encode_array(coords.astype(np.float32)), + } + txn.put( + format(frame_idx, fmt).encode(), + msgpack.packb(frame, use_bin_type=True), + ) + frame_system_ids.append(molecules) + frame_nlocs.append(len(atoms)) + frame_idx += 1 + molecules += 1 + metadata = { + "nframes": frame_idx, + "frame_idx_fmt": fmt, + "type_map": names, + "frame_system_ids": frame_system_ids, + "frame_nlocs": frame_nlocs, + "system_info": {"nframes": frame_idx}, + } + txn.put(b"__metadata__", msgpack.packb(metadata, use_bin_type=True)) + except BaseException: + env.close() + shutil.rmtree(staging, ignore_errors=True) + raise + env.close() + if frame_idx == 0: + shutil.rmtree(staging, ignore_errors=True) + raise ValueError( + f"{src} produced no usable frame, so there is nothing to train on; " + f"{skipped} record(s) were skipped for holding one atom or an " + "element outside the type_map" + ) + # Move the old dataset aside rather than deleting it first. Two renames + # still leave one instant with nothing at `dst`, so this does not promise + # that `dst` is always readable -- it promises that a complete dataset is + # always somewhere on disk, at `dst` or beside it, and that an interrupted + # run can be recovered rather than having lost data. + previous = f"{dst}.replaced" + if os.path.exists(previous): + if os.path.exists(dst): + # A leftover from a run that finished: `dst` is the current dataset + # and this is the superseded one. + shutil.rmtree(previous) + else: + # A run died between the two renames below, which makes this the + # only surviving copy. Deleting it here -- before the new + # conversion has proved itself -- is how both copies get lost. + os.rename(previous, dst) + had_previous = os.path.exists(dst) + if had_previous: + os.rename(dst, previous) + try: + os.rename(staging, dst) + except BaseException: + if had_previous: + os.rename(previous, dst) + raise + if had_previous: + shutil.rmtree(previous, ignore_errors=True) + return {"molecules": molecules, "frames": frame_idx, "skipped": skipped} + + +def main(args: list[str] | None = None) -> None: + """Command line entry point.""" + parser = argparse.ArgumentParser( + description="Convert a Uni-Mol pretraining LMDB into a deepmd LMDB dataset." + ) + parser.add_argument("src", help="the Uni-Mol .lmdb file") + parser.add_argument("dst", help="output directory for the deepmd dataset") + parser.add_argument( + "--add-2d-conformer", + action="store_true", + help="append the RDKit 2D conformer to every pool, as upstream does", + ) + parser.add_argument("--max-molecules", type=int, default=None) + parser.add_argument("--max-conformers", type=int, default=None) + parsed = parser.parse_args(args) + counts = convert_unimol_lmdb( + parsed.src, + parsed.dst, + add_2d_conformer=parsed.add_2d_conformer, + max_molecules=parsed.max_molecules, + max_conformers=parsed.max_conformers, + ) + logging.basicConfig(level=logging.INFO) + log.info( + "converted %d molecules into %d frames (%d records skipped)", + counts["molecules"], + counts["frames"], + counts["skipped"], + ) + + +if __name__ == "__main__": + main() diff --git a/doc/model/index.rst b/doc/model/index.rst index d08b932059..c6db98d169 100644 --- a/doc/model/index.rst +++ b/doc/model/index.rst @@ -13,6 +13,7 @@ Model dpa3 dpa4 dpa4c + unimol train-hybrid sel train-energy diff --git a/doc/model/unimol.md b/doc/model/unimol.md new file mode 100644 index 0000000000..ae2e88c54f --- /dev/null +++ b/doc/model/unimol.md @@ -0,0 +1,169 @@ +# Descriptor Uni-Mol {{ pytorch_icon }} {{ dpmodel_icon }} + +> [!NOTE] +> **Supported backends**: PyTorch-Exportable {{ pytorch_icon }}, DP {{ dpmodel_icon }} + +Uni-Mol is a molecular representation model: a transformer over all atom pairs, +in which geometry enters only through pairwise distances. It was pretrained on +about 209 million RDKit conformers with three self-supervised objectives and no +energies or forces at all. + +This is a port of Uni-Mol v1, faithful enough to load the released weights and +to reproduce the published objective. It exists for two reasons: to make +Uni-Mol's data and objectives available to multi-task training alongside +DFT-labelled data, and to give molecular property work a pretrained backbone. + +> [!IMPORTANT] +> Uni-Mol is **not** a potential energy surface model. It attends over every +> atom pair with no cut-off and no smooth envelope, so it is not extensive, it +> does not support periodic boundaries, and its forces are neither smooth nor +> conserved. The descriptor refuses any frame whose atoms are not all local, +> which rules out periodic images and the ghost-atom layout that freezing and +> parallel evaluation assume, so it is not available for molecular dynamics or +> frozen deployment. Nothing prevents a configuration from pairing it with an +> energy fitting; such a model would run, but it would not be a usable +> potential energy surface. + +Two further limits are worth knowing before configuring a run. The two virtual +tokens sit at the centroid of the real atoms by default, which keeps the +sequence translation invariant; set `virtual_token_position` to `origin` to +reproduce upstream exactly on data its own pipeline has centred. And the +descriptor decides on the data it is given -- rejecting periodic frames and +frames of fewer than two atoms -- so it cannot be traced or compiled, and +`enable_compile` is not available for it. + +## Architecture + +The backbone is a 15-layer pre-layer-norm transformer of width 512 with 64 +attention heads. Distances are expanded in 128 Gaussians whose affine +parameters depend on the ordered element pair, projected to one bias per head, +and added to the attention logits. Each layer's pre-softmax logits become the +next layer's bias, so the running sum of logits is a pair representation that +the pretraining heads read. Two virtual tokens wrap each molecule. + +Uni-Mol's own 31-token vocabulary is kept, because the released weights are +indexed by it: four special tokens, 26 elements, then `[MASK]`. A `type_map` is +mapped onto those ids, and an element outside the vocabulary becomes `[UNK]`. + +## Pretraining objective + +Fifteen percent of the atoms are selected. Of those, 90% become `[MASK]`, 5% +become a random element and 5% are left alone; all three are predicted. The +masked and randomly replaced atoms also have uniform noise of ±1 Å added to +each coordinate component. Five terms are minimized: + +| Term | Weight | What it predicts | +| --------------- | ------ | ----------------------------------------------- | +| element | 1 | the true element of every selected atom | +| coordinate | 5 | the clean coordinates, through the pair channel | +| distance | 10 | the clean pairwise distances | +| node norm | 0.01 | keeps node norms near $\sqrt{512}$ | +| pair-delta norm | 0.01 | keeps pair-delta norms near $\sqrt{64}$ | + +Corruption happens in the data pipeline rather than inside the loss, which is +both what upstream does and what the PyTorch-Exportable backend requires, since +it runs the model before the loss sees a frame. The objective owns the settings +and hands the trainer the transform it needs, so a training run installs it +automatically; the masking rate, the 90/5/5 split, the noise and the seed are +all configurable under `loss`. + +A masked atom is carried as a `[MASK]` pseudo-element, so the model's +`type_map` has to declare it alongside the elements. + +Each dataset gets its own corruption, and the validation set is corrupted afresh +on every pass, as upstream does. The validation loss therefore measures the model +against a different set of masked atoms each time and is not comparable +step-to-step the way a fixed validation set would be; read it as a trend. + +## Training + +```sh +dp --pt-expt train examples/unimol/pretrain/input.json +``` + +The dataset has to be an LMDB one, because the corruption happens as frames are +read; `deepmd.utils.unimol_data` below produces it. Give its path as a string +under `systems`, not as a list, which is how LMDB datasets are addressed. + +## Using the released weights + +The released checkpoints, `mol_pre_all_h_220816.pt` and +`mol_pre_no_h_220816.pt`, hold weights only and load with `weights_only=True`: + +```py +from deepmd.utils.unimol_checkpoint import descriptor_from_unimol_checkpoint + +descriptor = descriptor_from_unimol_checkpoint("mol_pre_all_h_220816.pt") +``` + +Parameter names line up one to one with the ported modules, but the arrays do +not: deepmd stores a linear weight as `(num_in, num_out)` and applies it as +`x @ w`, the transpose of `torch.nn.Linear.weight`. The converter renames and +transposes every weight, so there is no "just add a prefix" path. + +## Converting the pretraining data + +The released archive holds `train.lmdb` and `valid.lmdb`; the example +configuration expects both, converted separately: + +```sh +python -m deepmd.utils.unimol_data ligands/train.lmdb ./unimol_train --add-2d-conformer +python -m deepmd.utils.unimol_data ligands/valid.lmdb ./unimol_valid --add-2d-conformer +``` + +Upstream stores each molecule as a Python pickle, so the converter unpickles +whatever the file contains; run it only on data from a source you trust. What it +writes is msgpack, which carries no such risk, and the conversion happens once. + +The result is an LMDB dataset in the same layout as the other datasets +published in this format: zero-padded twelve-digit keys, one msgpack frame each, +and a `__metadata__` entry. Coordinates are stored as `float32` and types as +`int32`, which is both what that format uses and what the source holds, since +upstream generated these conformers in single precision; the reader casts to +whatever precision the model asks for. Molecules carry no cell at all. + +One conformer becomes one frame, so ordinary frame sampling stands in for +upstream's per-epoch conformer draw. The two-dimensional RDKit conformer that +upstream appends while loading is added at conversion time, behind that flag, +so the training data path never needs RDKit. + +## How closely this matches upstream + +Every component is checked against tensors dumped from upstream running +unmodified on the same inputs. The data-side transforms agree bitwise, which +means the random stream itself is reproduced, down to which atoms are masked +and what noise each one receives. The encoder, fed upstream's own attention +bias, agrees to fp64 rounding. + +What remains is upstream's own use of fp32 in three places, which sets the +floor at about 1e-7 relative on the whole objective: + +- the Gaussian basis is evaluated in fp32, reproduced by default and + switchable with `single_precision_basis`; +- the distance matrix is precomputed in fp32 by upstream's data pipeline, + while the descriptor computes distances in the working precision, which is + more accurate and is what gradients flow through; `single_precision_distance` + reproduces upstream's numbers instead; +- `log_softmax` and both norm regularisers are evaluated in fp32, which is + reproduced. + +The example trains in single precision, which is the default DPA models train +in. This backbone is a transformer rather than a potential energy surface, so +double precision buys nothing and costs several times the training time. + +Training trajectories cannot be reproduced exactly in any case: upstream +pretrained a pure fp16 model with fused kernels and its own Adam variant, which +places epsilon differently from PyTorch's. The example carries upstream's +optimizer values, including `adam_eps`, so the recipe matches even though the +trajectory cannot. + +Uni-Mol uses the exact error-function GELU, available here as `gelu_erf`. +deepmd's `gelu` and `gelu_tf` are the tanh approximation, which differs by up +to 4.7e-4 per element. + +## Attribution + +The ported code follows Uni-Mol (commit `90f52c4`) and the Uni-Core modules it +builds on (commit `ace6fae`), both MIT licensed, Copyright (c) DP Technology. +Parts of Uni-Core derive in turn from fairseq, Copyright (c) Facebook, Inc. and +its affiliates, also MIT licensed. Each ported file records its provenance. diff --git a/examples/unimol/pretrain/input.json b/examples/unimol/pretrain/input.json new file mode 100644 index 0000000000..b1f4b66745 --- /dev/null +++ b/examples/unimol/pretrain/input.json @@ -0,0 +1,113 @@ +{ + "_comment": "Uni-Mol v1 self-supervised molecular pretraining. Convert the data first with: python -m deepmd.utils.unimol_data ligands.lmdb ./unimol_train --add-2d-conformer", + "model": { + "type_map": [ + "C", + "N", + "O", + "S", + "H", + "Cl", + "F", + "Br", + "I", + "Si", + "P", + "B", + "Na", + "K", + "Al", + "Ca", + "Sn", + "As", + "Hg", + "Fe", + "Zn", + "Cr", + "Se", + "Gd", + "Au", + "Li", + "[MASK]" + ], + "descriptor": { + "type": "unimol", + "encoder_layers": 15, + "encoder_embed_dim": 512, + "encoder_ffn_embed_dim": 2048, + "encoder_attention_heads": 64, + "max_atoms": 256, + "activation_function": "gelu_erf", + "_comment": "dropout acts during training only; the two precision switches decide whether a run reproduces upstream's published numbers or takes the more accurate path", + "dropout": 0.1, + "emb_dropout": 0.1, + "attention_dropout": 0.1, + "single_precision_basis": true, + "single_precision_distance": false, + "seed": 1, + "_comment_max_atoms": "the size cap is applied by the converter, before batching; the descriptor and the distance head only declare the width they expect", + "precision": "float32", + "_comment_precision": "single precision, which is what DPA models train in by default; the backbone is a transformer, not a potential energy surface, so float64 buys nothing here and costs several times the training time", + "virtual_token_position": "origin", + "_comment_virtual_token_position": "the objective's distance target puts the virtual tokens at the origin, which is where the centred frame's centroid is; the pretraining path requires it" + }, + "fitting_net": { + "type": "unimol_pretrain", + "n_token": 31, + "attention_heads": 64, + "max_atoms": 256, + "activation_function": "gelu_erf", + "seed": 1, + "precision": "float32" + }, + "_comment_type_map": "[MASK] is the pseudo-element a corrupted atom is carried as; the objective needs it" + }, + "learning_rate": { + "type": "wsd", + "start_lr": 0.0001, + "stop_lr": 1e-12, + "warmup_steps": 10000, + "decay_type": "linear", + "decay_phase_ratio": 1.0, + "_comment": "upstream warms up over 10k steps then decays linearly to zero over 1M; WSD requires a positive stop_lr, so it goes to 1e-12" + }, + "loss": { + "type": "unimol", + "masked_token_loss": 1.0, + "masked_coord_loss": 5.0, + "masked_dist_loss": 10.0, + "x_norm_loss": 0.01, + "delta_pair_repr_norm_loss": 0.01, + "mask_prob": 0.15, + "leave_unmasked_prob": 0.05, + "random_token_prob": 0.05, + "noise_type": "uniform", + "noise": 1.0, + "data_seed": 1 + }, + "training": { + "training_data": { + "systems": "./unimol_train", + "batch_size": 16, + "_comment": "an LMDB dataset path, given as a string; build it with deepmd.utils.unimol_data" + }, + "validation_data": { + "systems": "./unimol_valid", + "batch_size": 16, + "_comment": "convert the released valid.lmdb the same way as train.lmdb" + }, + "numb_steps": 1000000, + "seed": 1, + "disp_file": "lcurve.out", + "disp_freq": 100, + "save_freq": 10000 + }, + "optimizer": { + "type": "Adam", + "adam_beta1": 0.9, + "adam_beta2": 0.99, + "adam_eps": 1e-06, + "weight_decay": 0.0001, + "_comment": "upstream's recipe; its own Adam places eps differently, so this matches the value but not the placement" + } +} diff --git a/source/tests/common/dpmodel/test_unimol.py b/source/tests/common/dpmodel/test_unimol.py new file mode 100644 index 0000000000..65830dc747 --- /dev/null +++ b/source/tests/common/dpmodel/test_unimol.py @@ -0,0 +1,646 @@ +# SPDX-License-Identifier: LGPL-3.0-or-later +"""Parity tests for the Uni-Mol v1 port. + +Every expected value in ``unimol_v1_golden.npz`` was produced by running +upstream Uni-Mol (commit 90f52c4) unmodified on CPU, over molecules 0 to 3 of +the example data shipped with that repository, at seed 1 and epoch 1. The +generator is ``golden/dump_golden.py`` in the development task directory; it +needs the upstream sources on ``PYTHONPATH`` and is not run by the test suite. + +Tolerances are not arbitrary. Upstream makes three fp32 choices that deepmd +does not have to make, and they set the floor: + +* the Gaussian basis is evaluated in fp32, +* the distance matrix is precomputed in fp32 by its data pipeline, +* ``log_softmax`` and both norm regularisers are evaluated in fp32. + +With those matched, agreement is at the fp32 rounding level, around 1e-7 +relative. The encoder itself, fed upstream's own attention bias, agrees to fp64 +rounding, which is what the strictest test below asserts. +""" + +import os +import unittest + +import numpy as np + +from deepmd.dpmodel.descriptor.unimol import ( + UNIMOL_ELEMENTS, + DescrptUniMol, + unimol_vocabulary, +) +from deepmd.dpmodel.descriptor.unimol_nn import ( + DistanceHead, + GaussianLayer, + MaskLMHead, + NonLinearHead, + TransformerEncoderWithPair, + coord_update, +) +from deepmd.dpmodel.fitting.unimol_pretrain import ( + UniMolPretrainFitting, +) +from deepmd.dpmodel.loss.unimol import ( + UniMolLoss, +) +from deepmd.dpmodel.utils.unimol_transform import ( + mask_points, + unimol_frame_transform, +) + +GOLDEN = os.path.join(os.path.dirname(__file__), "unimol_v1_golden.npz") +SEED, EPOCH, INDICES = 1, 1, [0, 1] +SMALL = {"layers": 2, "dim": 32, "ffn": 64, "heads": 4, "k": 128, "vocab": 31} + + +def _set_linear(layer, weights, prefix, weight_key=None, bias_key=None): + layer.w = np.ascontiguousarray(weights[weight_key or prefix + ".weight"].T) + key = bias_key or prefix + ".bias" + if key in weights and layer.b is not None: + layer.b = weights[key].copy() + + +def _set_layer_norm(layer, weights, prefix): + layer.w = weights[prefix + ".weight"].copy() + layer.b = weights[prefix + ".bias"].copy() + + +class UniMolGoldenMixin: + @classmethod + def setUpClass(cls) -> None: + cls.golden = np.load(GOLDEN) + cls.weights = { + k[len("small_weights/") :]: v + for k, v in cls.golden.items() + if k.startswith("small_weights/") + } + cls.n_real = (cls.golden["input/src_tokens"] != 0).sum(axis=1) - 2 + + def build_encoder(self): + enc = TransformerEncoderWithPair( + encoder_layers=SMALL["layers"], + embed_dim=SMALL["dim"], + ffn_embed_dim=SMALL["ffn"], + attention_heads=SMALL["heads"], + activation_function="gelu_erf", + ) + w = self.weights + _set_layer_norm(enc.emb_layer_norm, w, "encoder.emb_layer_norm") + _set_layer_norm(enc.final_layer_norm, w, "encoder.final_layer_norm") + _set_layer_norm(enc.final_head_layer_norm, w, "encoder.final_head_layer_norm") + for i, layer in enumerate(enc.layers): + p = f"encoder.layers.{i}." + _set_linear(layer.self_attn.in_proj, w, p + "self_attn.in_proj") + _set_linear(layer.self_attn.out_proj, w, p + "self_attn.out_proj") + _set_layer_norm(layer.self_attn_layer_norm, w, p + "self_attn_layer_norm") + _set_linear(layer.fc1, w, p + "fc1") + _set_linear(layer.fc2, w, p + "fc2") + _set_layer_norm(layer.final_layer_norm, w, p + "final_layer_norm") + return enc + + def build_descriptor(self, **overrides): + kwargs = { + "type_map": [*UNIMOL_ELEMENTS, "[MASK]"], + "encoder_layers": SMALL["layers"], + "encoder_embed_dim": SMALL["dim"], + "encoder_ffn_embed_dim": SMALL["ffn"], + "encoder_attention_heads": SMALL["heads"], + "max_atoms": int(self.n_real.max()) + 1, + "activation_function": "gelu_erf", + "virtual_token_position": "origin", + } + kwargs.update(overrides) + desc = DescrptUniMol(**kwargs) + w = self.weights + desc.embed_tokens.w = w["embed_tokens.weight"].copy() + for key in ("means", "stds", "mul", "bias"): + getattr(desc.gbf, key).w = w[f"gbf.{key}.weight"].copy() + _set_linear(desc.gbf_proj.linear1, w, "gbf_proj.linear1") + _set_linear(desc.gbf_proj.linear2, w, "gbf_proj.linear2") + enc = self.build_encoder() + desc.encoder = enc + return desc + + def deepmd_inputs(self): + """Rebuild the deepmd-side inputs for the molecules in the golden batch.""" + vocab = unimol_vocabulary() + type_map = [*UNIMOL_ELEMENTS, "[MASK]"] + type_of = {sym: i for i, sym in enumerate(type_map)} + src_tokens = self.golden["input/src_tokens"] + src_coord = self.golden["input/src_coord"].astype(np.float64) + nf, nloc = len(self.n_real), int(self.n_real.max()) + coord = np.zeros((nf, nloc, 3)) + atype = np.zeros((nf, nloc), dtype=np.int64) + nlist = np.full((nf, nloc, nloc - 1), -1, dtype=np.int64) + for f in range(nf): + n = int(self.n_real[f]) + coord[f, :n] = src_coord[f, 1 : n + 1] + atype[f, :n] = [type_of[vocab[t]] for t in src_tokens[f, 1 : n + 1]] + for i in range(n): + others = [j for j in range(n) if j != i] + nlist[f, i, : len(others)] = others + return coord, atype, nlist + + +class TestUniMolTransform(UniMolGoldenMixin, unittest.TestCase): + """The data-side corruption must reproduce upstream's random stream.""" + + def test_frame_transform_matches_upstream(self) -> None: + vocab = {sym: i for i, sym in enumerate(unimol_vocabulary())} + specials = [vocab[s] for s in ("[PAD]", "[CLS]", "[SEP]", "[UNK]", "[MASK]")] + for row, index in enumerate(INDICES): + atoms = self.golden[f"raw/{index}/atoms"] + conformers = self.golden[f"raw/{index}/conformers"] + got = unimol_frame_transform( + atoms, + conformers, + vocab=vocab, + num_types=len(vocab), + special_indices=specials, + pad_idx=vocab["[PAD]"], + bos_idx=vocab["[CLS]"], + eos_idx=vocab["[SEP]"], + mask_idx=vocab["[MASK]"], + unk_idx=vocab["[UNK]"], + seed=SEED, + epoch=EPOCH, + index=index, + max_atoms=256, + mask_prob=0.15, + leave_unmasked_prob=0.05, + random_token_prob=0.05, + noise_type="uniform", + noise=1.0, + ) + n = int(self.n_real[row]) + 2 + with self.subTest(molecule=index): + # Tokens, targets and both coordinate arrays are bitwise equal, + # which is what proves the random stream itself is reproduced. + np.testing.assert_array_equal( + got["src_tokens"], self.golden["input/src_tokens"][row, :n] + ) + np.testing.assert_array_equal( + got["tokens_target"], self.golden["target/tokens_target"][row, :n] + ) + np.testing.assert_array_equal( + got["src_edge_type"], + self.golden["input/src_edge_type"][row, :n, :n], + ) + np.testing.assert_array_equal( + got["src_coord"], self.golden["input/src_coord"][row, :n] + ) + np.testing.assert_array_equal( + got["coord_target"], self.golden["target/coord_target"][row, :n] + ) + # scipy's distance_matrix and a sqrt of summed squares differ in + # the last fp32 place. + np.testing.assert_allclose( + got["src_distance"], + self.golden["input/src_distance"][row, :n, :n], + atol=1e-5, + ) + + def test_masking_statistics(self) -> None: + """About 15% of atoms are selected, and only replaced ones are moved. + + The ported corruption is run here rather than read off the fixture, so + a change in the port can fail this test. + """ + vocab = {sym: i for i, sym in enumerate(unimol_vocabulary())} + specials = [vocab[s] for s in ("[PAD]", "[CLS]", "[SEP]", "[UNK]", "[MASK]")] + rng = np.random.default_rng(0) + selected_counts = [] + for index in range(24): + size = 40 + tokens = rng.integers(4, 30, size=size) + coords = rng.normal(size=(size, 3)).astype(np.float32) + out = mask_points( + tokens, + coords, + num_types=len(vocab), + special_indices=specials, + pad_idx=0, + mask_idx=vocab["[MASK]"], + seed=1, + epoch=1, + index=index, + ) + picked = out["targets"] != 0 + selected_counts.append(int(picked.sum())) + moved = np.abs(out["coordinates"] - coords).max(axis=-1) > 0 + # Only selected atoms move, and an atom left with its own element + # is still predicted. + self.assertTrue(bool(np.all(~moved | picked))) + np.testing.assert_array_equal( + out["targets"][picked], np.asarray(tokens)[picked] + ) + mean_fraction = float(np.mean(selected_counts)) / 40 + self.assertGreater(mean_fraction, 0.10) + self.assertLess(mean_fraction, 0.20) + + for row, index in enumerate(INDICES): + n = int(self.n_real[row]) + targets = self.golden["target/tokens_target"][row, 1 : n + 1] + selected = int((targets != 0).sum()) + self.assertGreater(selected, 0) + self.assertLessEqual(selected, max(1, int(0.5 * n))) + noisy = self.golden["input/src_coord"][row, 1 : n + 1] + clean = self.golden["target/coord_target"][row, 1 : n + 1] + moved = np.abs(noisy - clean).max(axis=-1) > 0 + # Every moved atom is a selected atom; the unchanged 5% are not moved. + self.assertTrue(bool(np.all(~moved | (targets != 0)))) + self.assertLessEqual(int(moved.sum()), selected) + + +class TestUniMolEncoder(UniMolGoldenMixin, unittest.TestCase): + """The backbone, isolated from upstream's fp32 choices.""" + + def test_encoder_matches_upstream_in_fp64(self) -> None: + enc = self.build_encoder() + emb = self.golden["small_fp64/enc_in/emb"].astype(np.float64) + bias = self.golden["small_fp64/enc_in/attn_bias"].astype(np.float64) + pad = self.golden["small_fp64/enc_in/padding_mask"].astype(np.float64) + x, pair, delta, x_norm, delta_norm = enc(emb, bias.copy(), pad) + # Fed upstream's own bias, everything is fp64 end to end. + np.testing.assert_allclose( + x, self.golden["small_fp64/enc_out/x"], rtol=1e-11, atol=1e-11 + ) + finite = np.isfinite(self.golden["small_fp64/enc_out/pair_rep"]) + np.testing.assert_allclose( + np.asarray(pair)[finite], + self.golden["small_fp64/enc_out/pair_rep"][finite], + rtol=1e-11, + atol=1e-11, + ) + np.testing.assert_allclose( + delta, + self.golden["small_fp64/enc_out/delta_pair_rep"], + rtol=1e-10, + atol=1e-10, + ) + np.testing.assert_allclose( + float(x_norm), + float(self.golden["small_fp64/enc_out/x_norm"]), + rtol=1e-6, + atol=1e-8, + ) + np.testing.assert_allclose( + float(delta_norm), + float(self.golden["small_fp64/enc_out/delta_pair_norm"]), + rtol=1e-6, + atol=1e-8, + ) + + def test_gaussian_basis_and_projection(self) -> None: + gbf = GaussianLayer(SMALL["k"], SMALL["vocab"] ** 2) + for key in ("means", "stds", "mul", "bias"): + getattr(gbf, key).w = self.weights[f"gbf.{key}.weight"].copy() + proj = NonLinearHead(SMALL["k"], SMALL["heads"], "gelu_erf", hidden=SMALL["k"]) + _set_linear(proj.linear1, self.weights, "gbf_proj.linear1") + _set_linear(proj.linear2, self.weights, "gbf_proj.linear2") + dist = self.golden["input/src_distance"].astype(np.float64) + bias = proj(gbf(dist, self.golden["input/src_edge_type"])) + nt = dist.shape[1] + bias = np.ascontiguousarray(np.transpose(bias, (0, 3, 1, 2))).reshape( + -1, nt, nt + ) + # The basis is evaluated in fp32 upstream, so this is one fp32 ulp. + np.testing.assert_allclose( + bias, self.golden["small_fp64/enc_in/attn_bias"], rtol=1e-6, atol=1e-6 + ) + + +class TestUniMolNormRegularisers(unittest.TestCase): + """The hinge itself, which the golden values cannot constrain. + + On the small random model both regularisers sit at exactly zero, because + the node norms happen to fall inside the tolerance. Comparing against zero + says nothing about the formula, so it is checked directly here. + """ + + def test_hinge_is_zero_inside_the_tolerance_and_grows_outside(self) -> None: + from deepmd.dpmodel.descriptor.unimol_nn.encoder import ( + norm_loss, + ) + + dim = 16 + root = dim**0.5 + unit = np.ones((1, 1, dim)) / root # norm 1 + # Upstream's tolerance is 1: a norm within 1 of sqrt(dim) costs nothing. + inside = unit * root + np.testing.assert_allclose(norm_loss(inside), 0.0, atol=1e-6) + np.testing.assert_allclose(norm_loss(unit * (root + 0.5)), 0.0, atol=1e-6) + # Beyond it the cost is the excess, either side. + np.testing.assert_allclose(norm_loss(unit * (root + 3.0)), 2.0, atol=1e-5) + np.testing.assert_allclose(norm_loss(unit * (root - 3.0)), 2.0, atol=1e-5) + + def test_masked_mean_ignores_padding_and_survives_an_empty_row(self) -> None: + from deepmd.dpmodel.descriptor.unimol_nn.encoder import ( + masked_mean, + ) + + value = np.array([[1.0, 2.0, 99.0], [4.0, 99.0, 99.0]]) + mask = np.array([[1.0, 1.0, 0.0], [1.0, 0.0, 0.0]]) + # Per row: 1.5 and 4.0, then the mean over rows. + np.testing.assert_allclose(float(masked_mean(mask, value)), 2.75, atol=1e-12) + # An all-padding row returns zero rather than dividing by zero, which + # is what upstream's epsilon in the denominator is for. + empty = masked_mean(np.zeros((1, 3)), np.ones((1, 3))) + self.assertTrue(np.isfinite(float(empty))) + np.testing.assert_allclose(float(empty), 0.0, atol=1e-6) + + +class TestUniMolHeads(UniMolGoldenMixin, unittest.TestCase): + """The three pretraining heads.""" + + def test_heads_match_upstream(self) -> None: + w = self.golden + lm = MaskLMHead(SMALL["dim"], SMALL["vocab"], "gelu_erf") + _set_linear(lm.dense, self.weights, "lm_head.dense") + _set_layer_norm(lm.layer_norm, self.weights, "lm_head.layer_norm") + _set_linear( + lm.out_proj, + self.weights, + None, + weight_key="lm_head.weight", + bias_key="lm_head.bias", + ) + dist_head = DistanceHead(SMALL["heads"], "gelu_erf") + _set_linear(dist_head.dense, self.weights, "dist_head.dense") + _set_layer_norm(dist_head.layer_norm, self.weights, "dist_head.layer_norm") + _set_linear(dist_head.out_proj, self.weights, "dist_head.out_proj") + p2c = NonLinearHead(SMALL["heads"], 1, "gelu_erf", hidden=SMALL["heads"]) + _set_linear(p2c.linear1, self.weights, "pair2coord_proj.linear1") + _set_linear(p2c.linear2, self.weights, "pair2coord_proj.linear2") + + x = w["small_fp64/enc_out/x"].astype(np.float64) + pair = w["small_fp64/enc_out/pair_rep"].astype(np.float64) + pair = np.where(np.isneginf(pair), 0.0, pair) + np.testing.assert_allclose( + lm(x, w["input/masked_tokens"]), + w["small_fp64/model/logits"], + rtol=1e-11, + atol=1e-11, + ) + np.testing.assert_allclose( + dist_head(pair), + w["small_fp64/model/encoder_distance"], + rtol=1e-11, + atol=1e-11, + ) + np.testing.assert_allclose( + coord_update( + w["input/src_coord"].astype(np.float64), + w["small_fp64/enc_out/delta_pair_rep"].astype(np.float64), + w["small_fp64/enc_in/padding_mask"].astype(np.float64), + p2c, + ), + w["small_fp64/model/encoder_coord"], + rtol=1e-11, + atol=1e-11, + ) + + +class TestUniMolDescriptor(UniMolGoldenMixin, unittest.TestCase): + """The descriptor, driven through deepmd-shaped inputs.""" + + def test_forward_matches_upstream(self) -> None: + desc = self.build_descriptor() + coord, atype, nlist = self.deepmd_inputs() + out = desc.forward_tokens(coord, atype, nlist) + ref = self.golden["small_fp64/enc_out/x"].astype(np.float64) + for f in range(len(self.n_real)): + n = int(self.n_real[f]) + with self.subTest(molecule=f): + np.testing.assert_allclose( + out["node_ebd"][f, 1 : n + 1], + ref[f, 1 : n + 1], + rtol=1e-6, + atol=1e-6, + ) + np.testing.assert_array_equal( + out["tokens"][f, : n + 2], + self.golden["input/src_tokens"][f, : n + 2], + ) + + def test_call_drops_the_virtual_tokens(self) -> None: + desc = self.build_descriptor() + coord, atype, nlist = self.deepmd_inputs() + node, rot, g2, h2, sw = desc.call(coord, atype, nlist) + self.assertEqual(node.shape, (coord.shape[0], coord.shape[1], SMALL["dim"])) + # The values must be the token-level ones with BOS and EOS removed, + # not merely an array of the right shape. + tokens = desc.forward_tokens(coord, atype, nlist)["node_ebd"] + np.testing.assert_allclose(node, tokens[:, 1 : coord.shape[1] + 1, :], atol=0) + self.assertGreater(float(np.abs(np.asarray(node)).max()), 0.0) + self.assertIsNone(rot) + self.assertIsNone(g2) + self.assertIsNone(h2) + self.assertIsNone(sw) + + def test_rejects_periodic_and_tiny_frames(self) -> None: + desc = self.build_descriptor() + coord, atype, nlist = self.deepmd_inputs() + with self.assertRaisesRegex(ValueError, "every atom to be local"): + desc.forward_tokens(np.concatenate([coord, coord], axis=1), atype, nlist) + lonely = np.full_like(nlist[:, :1, :], -1) + with self.assertRaisesRegex(ValueError, "two real atoms"): + desc.forward_tokens(coord[:, :1], atype[:, :1], lonely) + + def test_serialize_round_trip(self) -> None: + desc = self.build_descriptor() + clone = DescrptUniMol.deserialize(desc.serialize()) + coord, atype, nlist = self.deepmd_inputs() + np.testing.assert_allclose( + clone.forward_tokens(coord, atype, nlist)["node_ebd"], + desc.forward_tokens(coord, atype, nlist)["node_ebd"], + rtol=1e-14, + atol=1e-14, + ) + + +class TestUniMolLoss(UniMolGoldenMixin, unittest.TestCase): + """The five-term objective.""" + + def _labels(self, nloc, ncol): + nf = len(self.n_real) + labels = { + "unimol_token_target": np.zeros((nf, nloc), dtype=np.int64), + "unimol_coord_target": np.zeros((nf, nloc, 3)), + "unimol_dist_target": np.zeros((nf, nloc, ncol)), + "unimol_token_mask": np.zeros((nf, ncol), dtype=np.int64), + } + for f in range(nf): + n = int(self.n_real[f]) + labels["unimol_token_target"][f, :n] = self.golden["target/tokens_target"][ + f, 1 : n + 1 + ] + labels["unimol_coord_target"][f, :n] = self.golden["target/coord_target"][ + f, 1 : n + 1 + ] + labels["unimol_dist_target"][f, :n, : n + 2] = self.golden[ + "target/distance_target" + ][f, 1 : n + 1, : n + 2] + labels["unimol_token_mask"][f, : n + 2] = 1 + return labels + + def test_terms_match_upstream(self) -> None: + desc = self.build_descriptor() + fitting = UniMolPretrainFitting( + ntypes=len(desc.get_type_map()), + dim_descrpt=SMALL["dim"], + n_token=SMALL["vocab"], + attention_heads=SMALL["heads"], + max_atoms=desc.max_atoms, + activation_function="gelu_erf", + ) + w = self.weights + _set_linear(fitting.lm_head.dense, w, "lm_head.dense") + _set_layer_norm(fitting.lm_head.layer_norm, w, "lm_head.layer_norm") + _set_linear( + fitting.lm_head.out_proj, + w, + None, + weight_key="lm_head.weight", + bias_key="lm_head.bias", + ) + _set_linear(fitting.pair2coord_proj.linear1, w, "pair2coord_proj.linear1") + _set_linear(fitting.pair2coord_proj.linear2, w, "pair2coord_proj.linear2") + _set_linear(fitting.dist_head.dense, w, "dist_head.dense") + _set_layer_norm(fitting.dist_head.layer_norm, w, "dist_head.layer_norm") + _set_linear(fitting.dist_head.out_proj, w, "dist_head.out_proj") + + coord, atype, nlist = self.deepmd_inputs() + pred = fitting.call_tokens(desc.forward_tokens(coord, atype, nlist)) + labels = self._labels(coord.shape[1], fitting.max_atoms + 2) + total, more = UniMolLoss().call(1.0, 0, pred, labels) + + expected = { + "token_loss": "loss_token", + "coord_loss": "loss_coord", + "dist_loss": "loss_dist", + "x_norm_loss": "loss_x_norm", + "delta_pair_norm_loss": "loss_delta_pair_norm", + } + for mine, ref in expected.items(): + with self.subTest(term=mine): + np.testing.assert_allclose( + float(more[mine]), + float(self.golden[f"small_fp64/loss/{ref}"]), + rtol=1e-5, + atol=1e-8, + ) + np.testing.assert_allclose( + float(total), + float(self.golden["small_fp64/loss/loss_total"]), + rtol=1e-5, + atol=1e-8, + ) + + def test_an_empty_selection_stays_finite(self) -> None: + """A frame can draw no corrupted atom, and then there is nothing to average. + + The count is rounded stochastically, so a small molecule in a batch of + one reaches this. Every term would be a mean over an empty set; NaN + there would spread to every weight at the next backward pass. + """ + nf, nloc, ncol = 1, 4, 8 + rng = np.random.default_rng(0) + pred = { + "token_logits": rng.normal(size=(nf, nloc, SMALL["vocab"])), + "coord_update": rng.normal(size=(nf, nloc, 3)), + "pair_dist": rng.normal(size=(nf, nloc, ncol)), + "x_norm": np.zeros((nf, nloc, 1)), + "delta_pair_norm": np.zeros((nf, nloc, 1)), + "mask": np.ones((nf, nloc), dtype=np.int64), + } + labels = { + # every position is padding, which is what "nothing selected" means + "unimol_token_target": np.zeros((nf, nloc), dtype=np.int64), + "unimol_coord_target": rng.normal(size=(nf, nloc, 3)), + } + total, more = UniMolLoss().call(1.0, 0, pred, labels) + self.assertTrue(np.isfinite(float(total))) + for name, term in more.items(): + with self.subTest(term=name): + self.assertTrue(np.isfinite(float(term))) + + def test_derived_labels_match_explicit_ones(self) -> None: + """Storing the distance target would cost O(natoms^2) per frame. + + The loss derives it, and the token column mask, from the clean + coordinates and the real-atom mask instead. Both routes must agree. + """ + nf, nloc = len(self.n_real), int(self.n_real.max()) + ncol = nloc + 2 + rng = np.random.default_rng(0) + labels = self._labels(nloc, ncol) + mask = np.zeros((nf, nloc), dtype=np.int64) + for f in range(nf): + mask[f, : int(self.n_real[f])] = 1 + pred = { + "token_logits": rng.normal(size=(nf, nloc, SMALL["vocab"])), + "coord_update": rng.normal(size=(nf, nloc, 3)), + "pair_dist": rng.normal(size=(nf, nloc, ncol)), + "x_norm": np.zeros((nf, nloc, 1)), + "delta_pair_norm": np.zeros((nf, nloc, 1)), + "mask": mask, + } + explicit, _ = UniMolLoss().call(1.0, 0, pred, labels) + derived, _ = UniMolLoss().call( + 1.0, + 0, + pred, + { + k: v + for k, v in labels.items() + if k not in ("unimol_dist_target", "unimol_token_mask") + }, + ) + # The explicit target is upstream's, stored in fp32; the derived one is + # computed in the working precision, so they part company there. + np.testing.assert_allclose(float(derived), float(explicit), rtol=1e-7) + + def test_serialize_round_trip(self) -> None: + loss = UniMolLoss(masked_coord_loss=3.0, beta=0.5) + clone = UniMolLoss.deserialize(loss.serialize()) + self.assertEqual(clone.masked_coord_loss, 3.0) + self.assertEqual(clone.beta, 0.5) + + +class TestUniMolExample(unittest.TestCase): + """The shipped example must stay valid as the arguments evolve. + + It is checked here rather than in the shared example test, because that one + also requires the referenced dataset to exist in the repository, and this + example points at data the user converts from upstream. + """ + + def test_example_configuration_is_valid(self) -> None: + import json + from pathlib import ( + Path, + ) + + from deepmd.utils.argcheck import ( + normalize, + ) + + path = ( + Path(__file__).parents[4] + / "examples" + / "unimol" + / "pretrain" + / "input.json" + ) + self.assertTrue(path.is_file(), f"missing example: {path}") + config = normalize(json.loads(path.read_text())) + self.assertEqual(config["model"]["descriptor"]["type"], "unimol") + self.assertEqual(config["model"]["fitting_net"]["type"], "unimol_pretrain") + self.assertEqual(config["loss"]["type"], "unimol") + # A corrupted atom is carried as this pseudo-element, so the map needs it. + self.assertIn("[MASK]", config["model"]["type_map"]) + + +if __name__ == "__main__": + unittest.main() diff --git a/source/tests/common/dpmodel/test_unimol_data.py b/source/tests/common/dpmodel/test_unimol_data.py new file mode 100644 index 0000000000..6a885e7cb1 --- /dev/null +++ b/source/tests/common/dpmodel/test_unimol_data.py @@ -0,0 +1,382 @@ +# SPDX-License-Identifier: LGPL-3.0-or-later +"""The Uni-Mol to deepmd data conversion.""" + +import os +import pickle +import shutil +import tempfile +import unittest +import unittest.mock + +import numpy as np + +from deepmd.dpmodel.descriptor.unimol import ( + UNIMOL_ELEMENTS, +) +from deepmd.dpmodel.loss.unimol import ( + UniMolLoss, +) +from deepmd.dpmodel.utils.lmdb_data import ( + LmdbDataReader, +) +from deepmd.dpmodel.utils.unimol_transform import ( + make_unimol_data_transform, +) +from deepmd.utils.unimol_data import ( + convert_unimol_lmdb, + read_unimol_lmdb, +) + + +def _write_unimol_lmdb(path: str, molecules: list[dict]) -> None: + """Write a file in the upstream layout: one pickle per molecule.""" + import lmdb + + env = lmdb.open(path, subdir=False, map_size=1 << 24) + with env.begin(write=True) as txn: + for i, mol in enumerate(molecules): + txn.put(f"{i}".encode(), pickle.dumps(mol)) + env.close() + + +class TestUniMolDataConversion(unittest.TestCase): + def setUp(self) -> None: + self.tmp = tempfile.mkdtemp() + rng = np.random.default_rng(0) + # Large enough that the 15% selection actually selects something: with + # four atoms it usually selects none, and every assertion about the + # corruption would hold vacuously. + big = ["C", "N", "O", "H", "C", "C", "N", "O", "H", "H", "C", "F", "S", "C"] + self.molecules = [ + { + "atoms": big, + "coordinates": [ + rng.normal(size=(len(big), 3)).astype(np.float32) * 2.0 + for _ in range(3) + ], + "smi": "CNO", + }, + { + "atoms": ["C", "C", "H"], + "coordinates": [ + rng.normal(size=(3, 3)).astype(np.float32) for _ in range(2) + ], + "smi": "CC", + }, + # A single atom cannot be told apart from padding downstream, and an + # unmapped element would silently become [UNK]; both are skipped. + { + "atoms": ["C"], + "coordinates": [rng.normal(size=(1, 3)).astype(np.float32)], + "smi": "C", + }, + { + "atoms": ["C", "Xx"], + "coordinates": [rng.normal(size=(2, 3)).astype(np.float32)], + "smi": "C", + }, + ] + self.src = os.path.join(self.tmp, "mol.lmdb") + _write_unimol_lmdb(self.src, self.molecules) + + def tearDown(self) -> None: + shutil.rmtree(self.tmp, ignore_errors=True) + + def test_reader_streams_the_upstream_layout(self) -> None: + records = list(read_unimol_lmdb(self.src)) + self.assertEqual(len(records), len(self.molecules)) + self.assertEqual(records[0]["atoms"], self.molecules[0]["atoms"]) + self.assertEqual(len(records[0]["coordinates"]), 3) + + def test_conversion_round_trips_through_the_deepmd_reader(self) -> None: + dst = os.path.join(self.tmp, "converted") + counts = convert_unimol_lmdb(self.src, dst, map_size=1 << 24) + # One frame per conformer, and the two unusable records are skipped. + self.assertEqual(counts["frames"], 5) + self.assertEqual(counts["molecules"], 2) + self.assertEqual(counts["skipped"], 2) + + reader = LmdbDataReader(dst, list(UNIMOL_ELEMENTS)) + self.assertEqual(len(reader), 5) + frame = reader[0] + coord = np.asarray(frame["coord"]).reshape(-1, 3) + np.testing.assert_allclose( + coord, self.molecules[0]["coordinates"][0].astype(np.float64), atol=1e-6 + ) + symbols = [UNIMOL_ELEMENTS[i] for i in np.asarray(frame["atype"]).reshape(-1)] + self.assertEqual(symbols, list(self.molecules[0]["atoms"])) + # Molecules are not periodic, so no cell is written at all. A zero cell + # would not do: the neighbour-list builder takes any cell at face value + # and inverts it. + self.assertNotIn("box", frame) + + def test_transform_corrupts_frames_in_the_data_path(self) -> None: + """The reader hook is what makes self-supervised training possible.""" + type_map = [*UNIMOL_ELEMENTS, "[MASK]"] + dst = os.path.join(self.tmp, "for_training") + convert_unimol_lmdb(self.src, dst, type_map=type_map, map_size=1 << 24) + reader = LmdbDataReader(dst, type_map) + plain = reader[0] + reader.set_frame_transform(make_unimol_data_transform(type_map, seed=1)) + corrupted = reader[0] + + self.assertEqual( + sorted(set(corrupted) - set(plain)), + [ + "find_unimol_coord_target", + "find_unimol_token_target", + "unimol_coord_target", + "unimol_token_target", + ], + ) + target = np.asarray(corrupted["unimol_token_target"]) + clean = np.asarray(corrupted["unimol_coord_target"]).reshape(-1, 3) + noisy = np.asarray(corrupted["coord"]).reshape(-1, 3) + # Only selected atoms may move, and the clean target is centred. + moved = np.abs(noisy - clean).max(axis=-1) > 0 + self.assertTrue(bool(np.all(~moved | (target != 0)))) + # Centring is exact only to fp32, because the transform keeps upstream's + # fp32 coordinates. + np.testing.assert_allclose(clean.mean(axis=0), np.zeros(3), atol=1e-6) + # Masked atoms are carried as the pseudo-element. Of the selected + # atoms, 90% are masked and 5% take a random element; both are moved, + # so the masked ones are a subset of the moved ones. + is_mask = np.asarray(corrupted["atype"]) == type_map.index("[MASK]") + self.assertLessEqual(int(is_mask.sum()), int(moved.sum())) + self.assertTrue(bool(np.all(~is_mask | moved))) + # The fixture is large enough that something is actually corrupted, so + # the assertions above are not vacuous. + self.assertGreater(int((target != 0).sum()), 0) + self.assertGreater(int(moved.sum()), 0) + + def test_corruption_changes_between_visits(self) -> None: + """Upstream redraws every epoch; a frozen mask would be memorised.""" + type_map = [*UNIMOL_ELEMENTS, "[MASK]"] + dst = os.path.join(self.tmp, "revisited") + convert_unimol_lmdb(self.src, dst, type_map=type_map, map_size=1 << 24) + reader = LmdbDataReader(dst, type_map) + reader.set_frame_transform(make_unimol_data_transform(type_map, seed=1)) + first = np.asarray(reader[0]["unimol_token_target"]).copy() + second = np.asarray(reader[0]["unimol_token_target"]).copy() + self.assertFalse(np.array_equal(first, second)) + + def test_the_transform_survives_a_worker_process(self) -> None: + """LMDB decoding runs in spawned workers, which pickle the decoder config. + + Each batch sends a fresh copy, so anything the transform carries is + reset over and over. A counter would therefore freeze the corruption, + and a closure would not have made it across at all. + """ + type_map = [*UNIMOL_ELEMENTS, "[MASK]"] + transform = make_unimol_data_transform(type_map, seed=1) + frame = { + "coord": np.array( + [[0.0, 0.0, 0.0], [1.5, 0.0, 0.0], [0.0, 1.5, 0.0], [0.0, 0.0, 1.5]] * 4 + ), + "atype": np.array([0, 1, 2, 3] * 4, dtype=np.int64), + } + drawn = set() + for _ in range(12): + revived = pickle.loads(pickle.dumps(transform)) + drawn.add(revived(frame, 0)["unimol_token_target"].tobytes()) + self.assertGreater(len(drawn), 1) + + def test_the_objective_supplies_its_own_transform(self) -> None: + """A trainer installs whatever the loss declares, and nothing else. + + Supervised losses return None here, so the data path is untouched for + them; this objective returns the corruption that produces its labels. + """ + from deepmd.dpmodel.loss.property import ( + PropertyLoss, + ) + + type_map = [*UNIMOL_ELEMENTS, "[MASK]"] + self.assertIsNone( + PropertyLoss(task_dim=1, var_name="property").frame_transform(type_map) + ) + + loss = UniMolLoss(mask_prob=0.2) + self.assertEqual( + [r.key for r in loss.label_requirement], + ["unimol_token_target", "unimol_coord_target"], + ) + dst = os.path.join(self.tmp, "through_the_loss") + convert_unimol_lmdb(self.src, dst, type_map=type_map, map_size=1 << 24) + reader = LmdbDataReader(dst, type_map) + before = set(reader[0]) + reader.set_frame_transform(loss.frame_transform(type_map)) + after = reader[0] + for key in ("unimol_token_target", "unimol_coord_target"): + self.assertIn(key, set(after) - before) + self.assertEqual( + len(np.asarray(after["unimol_token_target"])), len(after["atype"]) + ) + + def test_an_element_outside_unimol_vocabulary_is_refused(self) -> None: + """Rewriting it as [MASK] would quietly corrupt an ordinary atom. + + Uni-Mol knows 26 elements. A model whose type_map goes beyond them + tokenizes the extras as [UNK], and [UNK] has no element to come back to, + so such an atom would return as [MASK] once it happened to be selected. + The frame is refused whatever the draw does. + """ + wider = ["C", "N", "O", "H", "Mg", "[MASK]"] + transform = make_unimol_data_transform(wider, seed=1) + frame = { + "coord": np.array( + [[0.0, 0.0, 0.0], [1.5, 0.0, 0.0], [0.0, 1.5, 0.0], [0.0, 0.0, 1.5]] + ), + # the third atom is magnesium, which Uni-Mol has no token for + "atype": np.array([0, 1, 4, 3], dtype=np.int64), + } + for _ in range(8): + with self.assertRaisesRegex(ValueError, "cannot express"): + transform(frame, 0) + + # Without it, the same frame goes through. + ordinary = make_unimol_data_transform(["C", "N", "O", "H", "[MASK]"], seed=1) + ordinary( + { + "coord": frame["coord"], + "atype": np.array([0, 1, 2, 3], dtype=np.int64), + }, + 0, + ) + + def test_the_same_seed_corrupts_the_same_way_twice(self) -> None: + """Two runs of one configuration have to agree. + + The number standing in for the epoch used to come from OS entropy, so + the documented guarantee -- reproducible when a single process decodes + -- did not actually hold. It is derived from the seed now. + + The generator lives in a module-level table so that it survives the + transform being re-created for every batch in a worker. Clearing that + table is what a fresh process looks like. + """ + from deepmd.dpmodel.utils import unimol_transform as transform_module + + type_map = [*UNIMOL_ELEMENTS, "[MASK]"] + frame = { + "coord": np.arange(48, dtype=np.float64).reshape(16, 3), + "atype": np.zeros(16, dtype=np.int64), + } + + def one_run(): + transform_module._EPOCH_STREAMS.clear() + transform = make_unimol_data_transform(type_map, seed=1, stream="training") + return [ + transform(dict(frame), 0)["unimol_token_target"].tolist() + for _ in range(3) + ] + + first, second = one_run(), one_run() + self.assertEqual(first, second) + # and it is not reproducible by being frozen: the visits differ + self.assertGreater(len({tuple(v) for v in first}), 1) + # a different seed gives a different stream + transform_module._EPOCH_STREAMS.clear() + other = make_unimol_data_transform(type_map, seed=2, stream="training") + self.assertNotEqual( + first[0], other(dict(frame), 0)["unimol_token_target"].tolist() + ) + + def test_training_and_validation_do_not_share_a_stream(self) -> None: + type_map = [*UNIMOL_ELEMENTS, "[MASK]"] + frame = { + "coord": np.arange(48, dtype=np.float64).reshape(16, 3), + "atype": np.zeros(16, dtype=np.int64), + } + train = make_unimol_data_transform(type_map, seed=1, stream="training") + valid = make_unimol_data_transform(type_map, seed=1, stream="validation") + self.assertNotEqual( + train(dict(frame), 0)["unimol_token_target"].tolist(), + valid(dict(frame), 0)["unimol_token_target"].tolist(), + ) + + def test_transform_requires_the_mask_pseudo_element(self) -> None: + with self.assertRaisesRegex(ValueError, r"\[MASK\]"): + make_unimol_data_transform(list(UNIMOL_ELEMENTS)) + + def _convert(self, dst, **kwargs): + return convert_unimol_lmdb(self.src, dst, map_size=1 << 24, **kwargs) + + def test_replacing_a_dataset_keeps_one_on_disk(self) -> None: + """The second conversion must not leave the destination empty.""" + dst = os.path.join(self.tmp, "replaced_twice") + self._convert(dst) + first = os.path.getsize(os.path.join(dst, "data.mdb")) + self._convert(dst) + self.assertTrue(os.path.isdir(dst)) + self.assertGreater(os.path.getsize(os.path.join(dst, "data.mdb")), 0) + # and the backup is cleaned up once the new one is in place + self.assertFalse(os.path.exists(dst + ".replaced")) + self.assertGreater(first, 0) + + def test_a_stranded_backup_survives_a_failed_publish(self) -> None: + """A run that died between the two renames leaves the only copy aside. + + Deleting that copy and then failing to put the new one in place loses + both. The publish is two renames, so the window is the second one + failing -- which is what this forces, because nothing else reaches it. + """ + dst = os.path.join(self.tmp, "stranded") + self._convert(dst) + # exactly the state an interrupted publish leaves behind + os.rename(dst, dst + ".replaced") + self.assertFalse(os.path.exists(dst)) + + real_rename = os.rename + + def fail_on_publish(a, b): + # let the restore through, refuse the staging -> dst move + if str(a).endswith(".partial"): + raise OSError("simulated failure publishing the new dataset") + return real_rename(a, b) + + with unittest.mock.patch("os.rename", side_effect=fail_on_publish): + with self.assertRaises(OSError): + self._convert(dst) + + survivor = dst if os.path.exists(dst) else dst + ".replaced" + self.assertTrue( + os.path.exists(survivor), "the only copy of the dataset was destroyed" + ) + self.assertGreater(os.path.getsize(os.path.join(survivor, "data.mdb")), 0) + + def test_a_failed_conversion_leaves_the_old_dataset_in_place(self) -> None: + """Rolling back has to put the previous dataset back at its name.""" + dst = os.path.join(self.tmp, "kept_on_failure") + self._convert(dst) + before = os.path.getsize(os.path.join(dst, "data.mdb")) + + # a source that cannot be read: the conversion raises before publishing + with self.assertRaises(Exception): + convert_unimol_lmdb( + os.path.join(self.tmp, "missing.lmdb"), dst, map_size=1 << 24 + ) + self.assertTrue(os.path.isdir(dst)) + self.assertEqual(os.path.getsize(os.path.join(dst, "data.mdb")), before) + + def test_an_empty_conversion_does_not_replace_a_good_dataset(self) -> None: + """Refusing an empty result must not cost the dataset already there.""" + dst = os.path.join(self.tmp, "kept_on_empty") + self._convert(dst) + before = os.path.getsize(os.path.join(dst, "data.mdb")) + with self.assertRaisesRegex(ValueError, "no usable frame"): + self._convert(dst, max_molecules=0) + self.assertTrue(os.path.isdir(dst)) + self.assertEqual(os.path.getsize(os.path.join(dst, "data.mdb")), before) + + def test_limits_are_respected(self) -> None: + dst = os.path.join(self.tmp, "limited") + counts = convert_unimol_lmdb( + self.src, dst, max_molecules=1, max_conformers=2, map_size=1 << 24 + ) + self.assertEqual(counts["molecules"], 1) + self.assertEqual(counts["frames"], 2) + + +if __name__ == "__main__": + unittest.main() diff --git a/source/tests/common/dpmodel/unimol_v1_golden.npz b/source/tests/common/dpmodel/unimol_v1_golden.npz new file mode 100644 index 0000000000..6dd3801eab Binary files /dev/null and b/source/tests/common/dpmodel/unimol_v1_golden.npz differ diff --git a/source/tests/consistent/test_activation.py b/source/tests/consistent/test_activation.py index 1c388b7305..9704ebe60f 100644 --- a/source/tests/consistent/test_activation.py +++ b/source/tests/consistent/test_activation.py @@ -1,6 +1,9 @@ # SPDX-License-Identifier: LGPL-3.0-or-later import sys import unittest +from importlib.util import ( + find_spec, +) import numpy as np @@ -111,6 +114,95 @@ def test_tf2_consistent_with_ref(self) -> None: test = get_activation_fn_dp(self.activation)(input) np.testing.assert_allclose(self.ref, to_numpy_array(test), atol=1e-7) + def test_the_tf2_namespace_name_still_resolves(self) -> None: + """Anchor the literal the tf2 dispatch compares against. + + ``xp_erf`` and its neighbours in ``deepmd/dpmodel/array_api.py`` select + the TensorFlow path by comparing ``xp.__name__`` against the import + path of the vendored namespace -- the idiom this tree uses in + ``utils/nlist.py`` and ``utils/default_neighbor_list.py`` too, nine + occurrences across three modules. + + A literal like that fails silently if the module is ever renamed or + moved: the comparison simply stops matching, the code falls back to the + NumPy round-trip, and the gradient is lost again with every + forward-value test still passing. This asserts the name resolves to a + real module, so a rename breaks a test instead of a derivative. + + Deliberately not gated on ``INSTALLED_TF2``: ``find_spec`` does not + import the module, so this runs everywhere, including the ordinary runs + where the tf2 cases skip. A guard that skips alongside the thing it + guards would protect nothing. + """ + self.assertIsNotNone( + find_spec("deepmd._vendors.ndtensorflow"), + "the namespace name that array_api.py dispatches on no longer " + "resolves; the TensorFlow branches there are now dead code", + ) + + @unittest.skipUnless(INSTALLED_TF2, "TensorFlow 2 is not installed") + def test_tf2_gradient_consistent_with_ref(self) -> None: + """The derivative has to survive the dispatch, not only the value. + + A backend that falls through to a NumPy conversion still returns the + right number, because the conversion happens after the value is + computed -- but it detaches the term from the tape, so the derivative + comes out wrong with nothing to show for it. Comparing forward values + alone cannot see that, which is why this compares the gradient. + + The input is deliberately a narrow range rather than the wide random + sample the value tests use: for the exact GELU the missing term is + ``x * phi(x)``, which vanishes for large ``|x|`` and would hide the + very defect this pins. + + The point count is even so that the grid straddles zero without + landing on it. ``relu`` and ``relu6`` have a kink there, where the + derivative does not exist: autodiff reports the subgradient 0 while a + central difference reports 0.5, and neither is wrong. Comparing them + at that point tests nothing about dispatch, which is what this is for. + """ + from deepmd._vendors import ndtensorflow as ndtf + from deepmd.tf2.env import ( + tf, + ) + + probe = np.linspace(-3.0, 3.0, 24) + raw = tf.constant(probe, dtype=tf.float64) + with tf.GradientTape() as tape: + tape.watch(raw) + out = get_activation_fn_dp(self.activation)(ndtf.asarray(raw)) + total = tf.reduce_sum(out.unwrap()) + grad = tape.gradient(total, raw) + self.assertIsNotNone( + grad, f"{self.activation} left no gradient path on the tf2 backend" + ) + # central differences on the reference implementation + eps = 1e-6 + ref = get_activation_fn_dp(self.activation) + expected = (ref(probe + eps) - ref(probe - eps)) / (2 * eps) + np.testing.assert_allclose(to_numpy_array(grad), expected, rtol=1e-5, atol=1e-6) + + @unittest.skipUnless(INSTALLED_TF2, "TensorFlow 2 is not installed") + def test_tf2_activation_is_traceable_in_graph_mode(self) -> None: + """Every activation must survive ``tf.function``. + + A NumPy conversion is refused outright on a graph tensor, so an + activation that reaches one cannot be trained or frozen on tf2 at all. + """ + from deepmd._vendors import ndtensorflow as ndtf + from deepmd.tf2.env import ( + tf, + ) + + activation = self.activation + + @tf.function + def traced(values): + return get_activation_fn_dp(activation)(ndtf.asarray(values)).unwrap() + + traced_out = traced(tf.constant(self.random_input, dtype=tf.float64)) + np.testing.assert_allclose(self.ref, traced_out.numpy(), atol=1e-7) + @unittest.skipUnless(INSTALLED_PD, "Paddle is not installed") def test_pd_consistent_with_ref(self): if INSTALLED_PD: diff --git a/source/tests/pt_expt/model/test_unimol.py b/source/tests/pt_expt/model/test_unimol.py new file mode 100644 index 0000000000..8ca7052efb --- /dev/null +++ b/source/tests/pt_expt/model/test_unimol.py @@ -0,0 +1,838 @@ +# SPDX-License-Identifier: LGPL-3.0-or-later +"""The Uni-Mol model on the PyTorch-Exportable backend. + +Checks that the registered path builds, that it agrees with the array-API +implementation on the same weights, and that the objective still reproduces the +values dumped from upstream. The golden archive and how it was produced are +described in ``source/tests/common/dpmodel/test_unimol.py``. +""" + +import os +import shutil +import tempfile +import unittest + +import numpy as np +import torch + +from deepmd.dpmodel.descriptor.unimol import ( + UNIMOL_ELEMENTS, +) +from deepmd.dpmodel.descriptor.unimol import DescrptUniMol as DescrptUniMolDP +from deepmd.dpmodel.descriptor.unimol import ( + unimol_vocabulary, +) +from deepmd.dpmodel.fitting.unimol_pretrain import ( + UniMolPretrainFitting as UniMolPretrainFittingDP, +) +from deepmd.pt_expt.loss.unimol import ( + UniMolLoss, +) +from deepmd.pt_expt.model import ( + get_model, +) +from deepmd.utils.argcheck import ( + normalize, +) + +GOLDEN = os.path.join( + os.path.dirname(__file__), "..", "..", "common", "dpmodel", "unimol_v1_golden.npz" +) +SMALL = {"layers": 2, "dim": 32, "ffn": 64, "heads": 4, "vocab": 31} + + +def _assign(module, name, array): + """Put a NumPy array into a torch-backed slot, keeping its kind.""" + current = getattr(module, name) + tensor = torch.as_tensor(np.ascontiguousarray(array), dtype=current.dtype) + if isinstance(current, torch.nn.Parameter): + tensor = torch.nn.Parameter(tensor, requires_grad=current.requires_grad) + setattr(module, name, tensor) + + +class TestUniMolPtExpt(unittest.TestCase): + @classmethod + def setUpClass(cls) -> None: + cls.golden = np.load(GOLDEN) + cls.weights = { + k[len("small_weights/") :]: v + for k, v in cls.golden.items() + if k.startswith("small_weights/") + } + cls.n_real = (cls.golden["input/src_tokens"] != 0).sum(axis=1) - 2 + cls.type_map = [*UNIMOL_ELEMENTS, "[MASK]"] + + def inputs(self): + vocab = unimol_vocabulary() + type_of = {sym: i for i, sym in enumerate(self.type_map)} + src_tokens = self.golden["input/src_tokens"] + src_coord = self.golden["input/src_coord"].astype(np.float64) + nf, nloc = len(self.n_real), int(self.n_real.max()) + coord = np.zeros((nf, nloc, 3)) + atype = np.zeros((nf, nloc), dtype=np.int64) + nlist = np.full((nf, nloc, nloc - 1), -1, dtype=np.int64) + for f in range(nf): + n = int(self.n_real[f]) + coord[f, :n] = src_coord[f, 1 : n + 1] + atype[f, :n] = [type_of[vocab[t]] for t in src_tokens[f, 1 : n + 1]] + for i in range(n): + others = [j for j in range(n) if j != i] + nlist[f, i, : len(others)] = others + return coord, atype, nlist + + def config(self, nloc, single_precision_basis: bool = True): + return { + "model": { + "type_map": self.type_map, + "descriptor": { + "type": "unimol", + "encoder_layers": SMALL["layers"], + "encoder_embed_dim": SMALL["dim"], + "encoder_ffn_embed_dim": SMALL["ffn"], + "encoder_attention_heads": SMALL["heads"], + "max_atoms": nloc, + "virtual_token_position": "origin", + "single_precision_basis": single_precision_basis, + "precision": "float64", + }, + "fitting_net": { + "type": "unimol_pretrain", + "n_token": SMALL["vocab"], + "attention_heads": SMALL["heads"], + "max_atoms": nloc, + "precision": "float64", + }, + }, + "learning_rate": {"type": "exp", "start_lr": 1e-4, "stop_lr": 1e-6}, + "loss": {"type": "unimol"}, + "training": { + "training_data": {"systems": ["x"]}, + "numb_steps": 1, + "seed": 1, + }, + } + + def load_weights(self, descriptor, fitting, torch_backend: bool) -> None: + w = self.weights + put = ( + _assign + if torch_backend + else (lambda m, n, a: setattr(m, n, np.ascontiguousarray(a))) + ) + + def lin(layer, wkey, bkey): + put(layer, "w", w[wkey].T) + if bkey in w and layer.b is not None: + put(layer, "b", w[bkey]) + + def ln(layer, prefix): + put(layer, "w", w[prefix + ".weight"]) + put(layer, "b", w[prefix + ".bias"]) + + put(descriptor.embed_tokens, "w", w["embed_tokens.weight"]) + for key in ("means", "stds", "mul", "bias"): + put(getattr(descriptor.gbf, key), "w", w[f"gbf.{key}.weight"]) + lin( + descriptor.gbf_proj.linear1, + "gbf_proj.linear1.weight", + "gbf_proj.linear1.bias", + ) + lin( + descriptor.gbf_proj.linear2, + "gbf_proj.linear2.weight", + "gbf_proj.linear2.bias", + ) + ln(descriptor.encoder.emb_layer_norm, "encoder.emb_layer_norm") + ln(descriptor.encoder.final_layer_norm, "encoder.final_layer_norm") + ln(descriptor.encoder.final_head_layer_norm, "encoder.final_head_layer_norm") + for i, layer in enumerate(descriptor.encoder.layers): + p = f"encoder.layers.{i}." + lin( + layer.self_attn.in_proj, + p + "self_attn.in_proj.weight", + p + "self_attn.in_proj.bias", + ) + lin( + layer.self_attn.out_proj, + p + "self_attn.out_proj.weight", + p + "self_attn.out_proj.bias", + ) + ln(layer.self_attn_layer_norm, p + "self_attn_layer_norm") + lin(layer.fc1, p + "fc1.weight", p + "fc1.bias") + lin(layer.fc2, p + "fc2.weight", p + "fc2.bias") + ln(layer.final_layer_norm, p + "final_layer_norm") + lin(fitting.lm_head.dense, "lm_head.dense.weight", "lm_head.dense.bias") + ln(fitting.lm_head.layer_norm, "lm_head.layer_norm") + lin(fitting.lm_head.out_proj, "lm_head.weight", "lm_head.bias") + lin( + fitting.pair2coord_proj.linear1, + "pair2coord_proj.linear1.weight", + "pair2coord_proj.linear1.bias", + ) + lin( + fitting.pair2coord_proj.linear2, + "pair2coord_proj.linear2.weight", + "pair2coord_proj.linear2.bias", + ) + lin(fitting.dist_head.dense, "dist_head.dense.weight", "dist_head.dense.bias") + ln(fitting.dist_head.layer_norm, "dist_head.layer_norm") + lin( + fitting.dist_head.out_proj, + "dist_head.out_proj.weight", + "dist_head.out_proj.bias", + ) + + def build_torch_model(self, nloc, single_precision_basis: bool = True): + model = get_model(normalize(self.config(nloc, single_precision_basis))["model"]) + self.load_weights( + model.atomic_model.descriptor, + model.atomic_model.fitting_net, + torch_backend=True, + ) + model.eval() + return model + + def test_registered_path_builds(self) -> None: + coord, _, _ = self.inputs() + model = self.build_torch_model(coord.shape[1]) + self.assertIsInstance(model, torch.nn.Module) + self.assertEqual(type(model).__name__, "UniMolPretrainModel") + self.assertEqual( + sorted(model.translated_output_def().keys()), + [ + "coord_update", + "delta_pair_norm", + "mask", + "pair_dist", + "token_logits", + "x_norm", + ], + ) + + def test_a_non_default_precision_runs(self) -> None: + """Torch refuses to multiply a float64 activation by a float32 weight. + + The backbone returns its output at the global precision, because its own + forward casts back, so heads configured at another precision are handed + the wrong dtype. Every other test here pins float64, which is the global + precision and hides it. NumPy hides it too, by upcasting silently. + """ + import copy + + coord, atype, nlist = self.inputs() + config = copy.deepcopy(self.config(coord.shape[1])) + for part in ("descriptor", "fitting_net"): + config["model"][part]["precision"] = "float32" + model = get_model(normalize(config)["model"]) + model.eval() + # The other tests here overwrite every parameter with a CPU array while + # loading weights, which quietly moves the model; this one keeps the + # model where it was built, so the inputs go to it. + device = next(model.parameters()).device + out = model.forward_lower( + torch.as_tensor(coord, device=device), + torch.as_tensor(atype, device=device), + torch.as_tensor(nlist, device=device), + ) + for name in ("token_logits", "coord_update", "pair_dist"): + with self.subTest(output=name): + self.assertTrue(torch.isfinite(out[name]).all()) + + def test_the_default_virtual_token_position_is_refused(self) -> None: + """The default configuration, reached the way a user reaches it. + + Under "centroid" the descriptor puts the virtual tokens at the centroid + of the coordinates it is handed, which during pretraining are the + corrupted ones, while the distance target puts them at the origin. The + two virtual columns of every corrupted row would then train against a + label for a different position, off by about the size of the noise. + + The key is *omitted* here rather than set to "centroid" by hand. Those + are different paths: writing the value tests the guard, while leaving it + out tests what a user who never heard of the option actually gets, which + is whatever argcheck fills in. A test that only wrote the value would + keep passing if the default moved out from under it. + """ + import copy + + config = copy.deepcopy(self.config(8)) + del config["model"]["descriptor"]["virtual_token_position"] + normalized = normalize(config) + # argcheck fills the key, and what it fills is the position this + # objective cannot score + self.assertEqual( + normalized["model"]["descriptor"]["virtual_token_position"], "centroid" + ) + with self.assertRaisesRegex(ValueError, "virtual_token_position='origin'"): + get_model(normalized["model"]) + + def test_writing_the_wrong_position_by_hand_is_refused_too(self) -> None: + """The same guard, reached by stating the value instead of defaulting.""" + import copy + + config = copy.deepcopy(self.config(8)) + config["model"]["descriptor"]["virtual_token_position"] = "centroid" + with self.assertRaisesRegex(ValueError, "virtual_token_position='origin'"): + get_model(normalize(config)["model"]) + + def test_the_norm_regularisers_ignore_the_padding(self) -> None: + """The regulariser is recovered through the real-atom mask. + + A batch holds frames of different sizes, so the short ones are padded, + and the atomic model zeroes every output it hands back at the padded + rows. The norm regularisers are scalars broadcast over the atoms, so + averaging one over *all* the rows divides it by the fraction of the + frame that is real -- a number that depends on which other molecules + happen to share the batch. + + ``mask`` is what prevents that. This pins both halves: that the padded + rows really are zero, and that dropping the mask really would change + the answer. + + ``delta_pair_norm`` is the probe rather than ``x_norm`` because + ``x_norm`` is a hinge -- upstream penalises only the part of + ``|‖h‖ - sqrt(d)|`` past a tolerance of 1.0 -- so behind a LayerNorm it + sits at exactly zero and would compare equal either way, proving + nothing. + """ + coord, atype, nlist = self.inputs() + # A padded row is marked by a negative type, which is how a mixed-size + # batch reaches the model (deepmd/dpmodel/utils/lmdb_data.py). Padding + # it with type 0 instead would make the short frames look full of + # carbon, and this test would prove nothing. + atype = atype.copy() + for frame in range(atype.shape[0]): + atype[frame, int(self.n_real[frame]) :] = -1 + nloc = coord.shape[1] + model = self.build_torch_model(nloc) + out = model.forward_lower( + torch.as_tensor(coord), + torch.as_tensor(atype), + torch.as_tensor(nlist), + ) + self.assertIn("mask", out) + mask = out["mask"] + value = out["delta_pair_norm"].detach() + + short = [f for f in range(mask.shape[0]) if int(mask[f].sum()) < nloc] + # the batch has to actually contain a padded frame, or this proves nothing + self.assertTrue(short, "fixture carries no padded frame") + + for frame in short: + with self.subTest(frame=frame): + n_real = int(mask[frame].sum()) + rows = value[frame].reshape(nloc) + self.assertTrue(bool((rows[mask[frame] == 0] == 0).all())) + masked_mean = float(rows[mask[frame] == 1].mean()) + unmasked_mean = float(rows.mean()) + # every real row carries the same scalar, so the masked mean is + # that scalar + self.assertAlmostEqual( + float(rows[mask[frame] == 1][0]), masked_mean, places=10 + ) + # and the unmasked one is it diluted by exactly the padding + self.assertAlmostEqual( + unmasked_mean, masked_mean * n_real / nloc, places=6 + ) + self.assertLess(unmasked_mean, masked_mean) + + def test_matches_the_array_api_implementation(self) -> None: + """Same weights, same numbers, once the fp32 basis is out of the way. + + Upstream evaluates the Gaussian basis in fp32, and NumPy and Torch round + that last place differently, which is worth about 2e-8 relative on the + node representation. Turning that off isolates the two implementations + from each other, and then they agree to fp64 rounding. + """ + coord, atype, nlist = self.inputs() + nloc = coord.shape[1] + model = self.build_torch_model(nloc, single_precision_basis=False) + ret = model.forward_lower( + torch.as_tensor(coord.reshape(len(self.n_real), -1)), + torch.as_tensor(atype), + torch.as_tensor(nlist), + None, + ) + + cfg = normalize(self.config(nloc, single_precision_basis=False))["model"] + descriptor = DescrptUniMolDP( + type_map=self.type_map, + **{k: v for k, v in cfg["descriptor"].items() if k != "type"}, + ) + fitting = UniMolPretrainFittingDP( + ntypes=len(self.type_map), + dim_descrpt=SMALL["dim"], + **{k: v for k, v in cfg["fitting_net"].items() if k != "type"}, + ) + self.load_weights(descriptor, fitting, torch_backend=False) + reference = fitting.call_tokens(descriptor.forward_tokens(coord, atype, nlist)) + + for key in ("token_logits", "coord_update", "pair_dist"): + with self.subTest(output=key): + np.testing.assert_allclose( + ret[key].detach().cpu().numpy(), + np.asarray(reference[key]), + rtol=1e-12, + atol=1e-12, + ) + # The two regularisers are evaluated in fp32, because upstream + # evaluates them there, so the two backends round the last place + # differently. + for key in ("x_norm", "delta_pair_norm"): + with self.subTest(output=key): + np.testing.assert_allclose( + ret[key].detach().cpu().numpy(), + np.asarray(reference[key]), + rtol=1e-6, + atol=1e-8, + ) + + def test_single_precision_basis_costs_one_fp32_place(self) -> None: + """With upstream's fp32 basis the two backends part company measurably. + + The gap is one fp32 unit in the last place in the basis, amplified by + the stack. It is recorded here so that the looser tolerance elsewhere + has a stated cause rather than being tuned until tests pass. + """ + coord, atype, nlist = self.inputs() + nloc = coord.shape[1] + exact = self.build_torch_model(nloc, single_precision_basis=False) + upstream_like = self.build_torch_model(nloc, single_precision_basis=True) + args = ( + torch.as_tensor(coord.reshape(len(self.n_real), -1)), + torch.as_tensor(atype), + torch.as_tensor(nlist), + None, + ) + a = exact.forward_lower(*args)["token_logits"].detach().cpu().numpy() + b = upstream_like.forward_lower(*args)["token_logits"].detach().cpu().numpy() + gap = np.abs(a - b).max() / np.abs(a).max() + self.assertLess(gap, 1e-5) + self.assertGreater(gap, 1e-12) + + def test_objective_matches_upstream(self) -> None: + coord, atype, nlist = self.inputs() + nloc = coord.shape[1] + model = self.build_torch_model(nloc) + ret = model.forward_lower( + torch.as_tensor(coord.reshape(len(self.n_real), -1)), + torch.as_tensor(atype), + torch.as_tensor(nlist), + None, + ) + ncol = model.atomic_model.fitting_net.max_atoms + 2 + nf = len(self.n_real) + labels = { + "unimol_token_target": torch.zeros((nf, nloc), dtype=torch.int64), + "unimol_coord_target": torch.zeros((nf, nloc, 3), dtype=torch.float64), + "unimol_dist_target": torch.zeros((nf, nloc, ncol), dtype=torch.float64), + "unimol_token_mask": torch.zeros((nf, ncol), dtype=torch.int64), + } + for f in range(nf): + n = int(self.n_real[f]) + labels["unimol_token_target"][f, :n] = torch.as_tensor( + self.golden["target/tokens_target"][f, 1 : n + 1].astype(np.int64) + ) + labels["unimol_coord_target"][f, :n] = torch.as_tensor( + self.golden["target/coord_target"][f, 1 : n + 1].astype(np.float64) + ) + labels["unimol_dist_target"][f, :n, : n + 2] = torch.as_tensor( + self.golden["target/distance_target"][f, 1 : n + 1, : n + 2].astype( + np.float64 + ) + ) + labels["unimol_token_mask"][f, : n + 2] = 1 + labels = {k: v.to(ret["token_logits"].device) for k, v in labels.items()} + + total, more = UniMolLoss().call(1.0, 0, ret, labels) + expected = { + "token_loss": "loss_token", + "coord_loss": "loss_coord", + "dist_loss": "loss_dist", + "x_norm_loss": "loss_x_norm", + "delta_pair_norm_loss": "loss_delta_pair_norm", + } + for mine, ref in expected.items(): + with self.subTest(term=mine): + np.testing.assert_allclose( + float(more[mine].detach()), + float(self.golden[f"small_fp64/loss/{ref}"]), + rtol=1e-5, + atol=1e-8, + ) + np.testing.assert_allclose( + float(total.detach()), + float(self.golden["small_fp64/loss/loss_total"]), + rtol=1e-5, + atol=1e-8, + ) + + def test_dropout_is_active_only_in_training(self) -> None: + """Uni-Mol regularises with dropout at three sites; deepmd had none.""" + coord, atype, nlist = self.inputs() + model = self.build_torch_model(coord.shape[1]) + args = ( + torch.as_tensor(coord.reshape(len(self.n_real), -1)), + torch.as_tensor(atype), + torch.as_tensor(nlist), + None, + ) + model.eval() + first = model.forward_lower(*args)["token_logits"].detach().clone() + second = model.forward_lower(*args)["token_logits"].detach() + torch.testing.assert_close(first, second, rtol=0, atol=0) + + model.train() + torch.manual_seed(0) + a = model.forward_lower(*args)["token_logits"].detach() + torch.manual_seed(1) + b = model.forward_lower(*args)["token_logits"].detach() + self.assertFalse(bool(torch.allclose(a, b))) + model.eval() + + # The rates are what drive it: with every rate at zero, training mode + # is deterministic again. Without this, a single hard-coded dropout + # call would satisfy the test above. + quiet = self.build_torch_model(coord.shape[1]) + for module in ( + quiet.atomic_model.descriptor, + quiet.atomic_model.descriptor.encoder, + ): + for attr in ( + "dropout", + "emb_dropout", + "attention_dropout", + "activation_dropout", + ): + if hasattr(module, attr): + setattr(module, attr, 0.0) + for layer in quiet.atomic_model.descriptor.encoder.layers: + layer.dropout = 0.0 + layer.attention_dropout = 0.0 + layer.activation_dropout = 0.0 + layer.self_attn.dropout = 0.0 + quiet.train() + torch.manual_seed(0) + c = quiet.forward_lower(*args)["token_logits"].detach() + torch.manual_seed(1) + d = quiet.forward_lower(*args)["token_logits"].detach() + torch.testing.assert_close(c, d, rtol=0, atol=0) + quiet.eval() + + def test_upper_path_builds_its_own_neighbour_list(self) -> None: + """What a user actually calls: coordinates and types, no neighbour list.""" + coord, atype, _ = self.inputs() + model = self.build_torch_model(coord.shape[1]) + small = 12 + ret = model.forward( + torch.as_tensor(coord[:1, :small].reshape(1, -1)), + torch.as_tensor(atype[:1, :small]), + None, + ) + self.assertEqual(tuple(ret["token_logits"].shape), (1, small, SMALL["vocab"])) + self.assertEqual(tuple(ret["coord_update"].shape), (1, small, 3)) + self.assertEqual(tuple(ret["mask"].shape), (1, small)) + + def test_periodic_cell_is_refused(self) -> None: + """A cell has to be caught here, before the neighbour list is built. + + The descriptor has no cut-off, so a cell would send the neighbour-list + builder looking for an astronomical number of periodic images and the + run would die on allocation rather than on a readable error. + """ + coord, atype, _ = self.inputs() + model = self.build_torch_model(coord.shape[1]) + small = 8 + box = torch.eye(3, dtype=torch.float64).reshape(1, 9) * 20.0 + with self.assertRaisesRegex(ValueError, "periodic"): + model.forward( + torch.as_tensor(coord[:1, :small].reshape(1, -1)), + torch.as_tensor(atype[:1, :small]), + box.to(next(model.parameters()).device), + ) + + def test_gradients_flow(self) -> None: + coord, atype, nlist = self.inputs() + model = self.build_torch_model(coord.shape[1]) + ret = model.forward_lower( + torch.as_tensor(coord.reshape(len(self.n_real), -1)), + torch.as_tensor(atype), + torch.as_tensor(nlist), + None, + ) + ret["token_logits"].sum().backward() + named = dict(model.named_parameters()) + # The heads alone would satisfy "some parameter has a gradient", so the + # backbone is named explicitly: the embedding, the distance basis and + # the first block all have to be reached. + backbone = [ + name + for name in named + if "embed_tokens" in name or ".gbf." in name or "layers.0." in name + ] + self.assertGreater(len(backbone), 5) + for name in backbone: + with self.subTest(parameter=name): + grad = named[name].grad + self.assertIsNotNone(grad, f"{name} received no gradient") + self.assertTrue( + bool(torch.any(grad != 0)), f"{name} got a zero gradient" + ) + + +class TestUniMolCheckpointImport(unittest.TestCase): + """The documented way to use the released weights. + + The golden archive carries the small model's weights under upstream's own + names, so the importer can be driven with exactly the layout a released + checkpoint has, without shipping a 190 MB file. + """ + + def test_imports_upstream_named_weights(self) -> None: + from deepmd.dpmodel.descriptor.unimol import DescrptUniMol as DescrptUniMolDP + from deepmd.utils.unimol_checkpoint import ( + apply_unimol_backbone, + split_unimol_state_dict, + ) + + golden = np.load(GOLDEN) + state = { + k[len("small_weights/") :]: v + for k, v in golden.items() + if k.startswith("small_weights/") + } + backbone, heads = split_unimol_state_dict(state) + self.assertTrue(all(not k.startswith("lm_head") for k in backbone)) + self.assertTrue(any(k.startswith("lm_head") for k in heads)) + with self.assertRaisesRegex(ValueError, "unexpected"): + split_unimol_state_dict({**state, "something.unexpected": np.zeros(1)}) + + descriptor = DescrptUniMolDP( + type_map=[*UNIMOL_ELEMENTS, "[MASK]"], + encoder_layers=SMALL["layers"], + encoder_embed_dim=SMALL["dim"], + encoder_ffn_embed_dim=SMALL["ffn"], + encoder_attention_heads=SMALL["heads"], + max_atoms=16, + ) + apply_unimol_backbone(descriptor, backbone) + # Lookup tables arrive as they are; projections arrive transposed, + # because deepmd applies a linear weight the other way round. + np.testing.assert_allclose( + descriptor.embed_tokens.w, state["embed_tokens.weight"], atol=0 + ) + np.testing.assert_allclose( + descriptor.gbf.mul.w, state["gbf.mul.weight"], atol=0 + ) + np.testing.assert_allclose( + descriptor.encoder.layers[0].fc1.w, + state["encoder.layers.0.fc1.weight"].T, + atol=0, + ) + np.testing.assert_allclose( + descriptor.encoder.layers[0].fc1.b, + state["encoder.layers.0.fc1.bias"], + atol=0, + ) + + def test_rejects_a_descriptor_the_checkpoint_does_not_fit(self) -> None: + """A wrong architecture override must not load silently.""" + from deepmd.dpmodel.descriptor.unimol import DescrptUniMol as DescrptUniMolDP + from deepmd.utils.unimol_checkpoint import ( + apply_unimol_backbone, + split_unimol_state_dict, + ) + + golden = np.load(GOLDEN) + state = { + k[len("small_weights/") :]: v + for k, v in golden.items() + if k.startswith("small_weights/") + } + backbone, _ = split_unimol_state_dict(state) + wrong_depth = DescrptUniMolDP( + type_map=[*UNIMOL_ELEMENTS, "[MASK]"], + encoder_layers=SMALL["layers"] + 1, + encoder_embed_dim=SMALL["dim"], + encoder_ffn_embed_dim=SMALL["ffn"], + encoder_attention_heads=SMALL["heads"], + max_atoms=16, + ) + with self.assertRaisesRegex(ValueError, "encoder layers"): + apply_unimol_backbone(wrong_depth, backbone) + + wrong_width = DescrptUniMolDP( + type_map=[*UNIMOL_ELEMENTS, "[MASK]"], + encoder_layers=SMALL["layers"], + encoder_embed_dim=SMALL["dim"] * 2, + encoder_ffn_embed_dim=SMALL["ffn"], + encoder_attention_heads=SMALL["heads"], + max_atoms=16, + ) + with self.assertRaisesRegex(ValueError, "expects"): + apply_unimol_backbone(wrong_width, backbone) + + +class TestUniMolTraining(unittest.TestCase): + """A run from a configuration file, which is how the feature is used. + + Everything between the configuration and the first step is exercised here: + the loss factory, the accessors the atomic model calls on any fitting, the + reader hook that produces the labels, and the absence of a cell on + molecular frames. Each of those was broken at some point and no unit test + would have shown it. + """ + + @classmethod + def setUpClass(cls) -> None: + import pickle + + import lmdb + + from deepmd.utils.unimol_data import ( + convert_unimol_lmdb, + ) + + cls.tmp = tempfile.mkdtemp() + cls.type_map = [*UNIMOL_ELEMENTS, "[MASK]"] + rng = np.random.default_rng(0) + molecules = [ + { + "atoms": ["C", "N", "O", "H", "C", "H"], + "coordinates": [rng.normal(size=(6, 3)).astype(np.float32) * 2.0], + "smi": "CNOC", + } + for _ in range(8) + ] + src = os.path.join(cls.tmp, "mol.lmdb") + env = lmdb.open(src, subdir=False, map_size=1 << 24) + with env.begin(write=True) as txn: + for i, mol in enumerate(molecules): + txn.put(f"{i}".encode(), pickle.dumps(mol)) + env.close() + cls.data = os.path.join(cls.tmp, "converted") + convert_unimol_lmdb(src, cls.data, type_map=cls.type_map, map_size=1 << 24) + + @classmethod + def tearDownClass(cls) -> None: + shutil.rmtree(cls.tmp, ignore_errors=True) + + def test_a_configuration_trains(self) -> None: + from deepmd.pt_expt.entrypoints.main import ( + get_trainer, + ) + + config = { + "model": { + "type_map": self.type_map, + "descriptor": { + "type": "unimol", + "encoder_layers": 1, + "encoder_embed_dim": 16, + "encoder_ffn_embed_dim": 32, + "encoder_attention_heads": 2, + "max_atoms": 16, + "virtual_token_position": "origin", + "seed": 1, + }, + "fitting_net": { + "type": "unimol_pretrain", + "attention_heads": 2, + "max_atoms": 16, + "seed": 1, + }, + }, + "learning_rate": {"type": "exp", "start_lr": 1e-3, "stop_lr": 1e-4}, + "loss": {"type": "unimol"}, + "training": { + # LMDB datasets are addressed with a string, not a list. + "training_data": {"systems": self.data, "batch_size": 2}, + "numb_steps": 4, + "seed": 1, + "disp_freq": 10, + "save_freq": 100, + "disp_file": os.path.join(self.tmp, "lcurve.out"), + "save_ckpt": os.path.join(self.tmp, "model.ckpt"), + }, + } + from deepmd.utils.compat import ( + update_deepmd_input, + ) + + trainer = get_trainer(normalize(update_deepmd_input(config, warning=False))) + self.assertEqual(type(trainer.model).__name__, "UniMolPretrainModel") + # The objective installed its own corruption on the dataset. + frame = trainer.training_data._reader[0] + self.assertIn("unimol_token_target", frame) + trainer.run() + + def test_the_trainer_gives_each_dataset_its_own_corruption(self) -> None: + """Through the trainer, not by handing the labels over myself. + + The draw sequence is derived from a label the caller supplies. Passing + two distinct labels in a test proves nothing about production, because + the question is whether the *trainer* passes two -- and for a while it + did not: both datasets took the default, shared one generator, and a + validation pass advanced the corruption training was about to see. + """ + from deepmd.pt_expt.entrypoints.main import ( + get_trainer, + ) + from deepmd.utils.compat import ( + update_deepmd_input, + ) + + config = { + "model": { + "type_map": self.type_map, + "descriptor": { + "type": "unimol", + "encoder_layers": 1, + "encoder_embed_dim": 16, + "encoder_ffn_embed_dim": 32, + "encoder_attention_heads": 2, + "max_atoms": 16, + "virtual_token_position": "origin", + "seed": 1, + }, + "fitting_net": { + "type": "unimol_pretrain", + "attention_heads": 2, + "max_atoms": 16, + "seed": 1, + }, + }, + "learning_rate": {"type": "exp", "start_lr": 1e-3, "stop_lr": 1e-4}, + "loss": {"type": "unimol"}, + "training": { + "training_data": {"systems": self.data, "batch_size": 2}, + "validation_data": {"systems": self.data, "batch_size": 2}, + "numb_steps": 1, + "seed": 1, + "disp_freq": 10, + "save_freq": 100, + "disp_file": os.path.join(self.tmp, "lcurve_streams.out"), + "save_ckpt": os.path.join(self.tmp, "model_streams.ckpt"), + }, + } + trainer = get_trainer(normalize(update_deepmd_input(config, warning=False))) + + def installed(dataset): + return dataset._reader._decode_config.frame_transform + + train = installed(trainer.training_data) + valid = installed(trainer.validation_data) + self.assertIsNotNone(train) + self.assertIsNotNone(valid) + # different objects is the easy half; different draw sequences is the + # half that was broken + self.assertIsNot(train, valid) + self.assertNotEqual( + train.stream, + valid.stream, + "the trainer gave both datasets the same corruption stream", + ) + + +if __name__ == "__main__": + unittest.main()