From 51ee5de7752836e56e749c49c481edb024724d68 Mon Sep 17 00:00:00 2001 From: Jinzhe Zeng Date: Wed, 8 Jul 2026 00:46:07 +0800 Subject: [PATCH 01/11] feat(jax): support DPA4 training --- deepmd/jax/descriptor/__init__.py | 4 + deepmd/jax/descriptor/dpa4.py | 283 ++++++++++++++++++ deepmd/jax/fitting/__init__.py | 4 + deepmd/jax/fitting/dpa4_ener.py | 54 ++++ deepmd/jax/model/ener_model.py | 2 + deepmd/jax/model/model.py | 49 +++ deepmd/jax/train/trainer.py | 44 ++- .../tests/consistent/descriptor/test_dpa4.py | 19 +- .../consistent/fitting/test_dpa4_ener.py | 20 +- 9 files changed, 474 insertions(+), 5 deletions(-) create mode 100644 deepmd/jax/descriptor/dpa4.py create mode 100644 deepmd/jax/fitting/dpa4_ener.py diff --git a/deepmd/jax/descriptor/__init__.py b/deepmd/jax/descriptor/__init__.py index cda2faf24d..f88a8f6589 100644 --- a/deepmd/jax/descriptor/__init__.py +++ b/deepmd/jax/descriptor/__init__.py @@ -8,6 +8,9 @@ from deepmd.jax.descriptor.dpa3 import ( DescrptDPA3, ) +from deepmd.jax.descriptor.dpa4 import ( + DescrptDPA4, +) from deepmd.jax.descriptor.hybrid import ( DescrptHybrid, ) @@ -31,6 +34,7 @@ "DescrptDPA1", "DescrptDPA2", "DescrptDPA3", + "DescrptDPA4", "DescrptHybrid", "DescrptSeA", "DescrptSeAttenV2", diff --git a/deepmd/jax/descriptor/dpa4.py b/deepmd/jax/descriptor/dpa4.py new file mode 100644 index 0000000000..b97b14c082 --- /dev/null +++ b/deepmd/jax/descriptor/dpa4.py @@ -0,0 +1,283 @@ +# SPDX-License-Identifier: LGPL-3.0-or-later +from collections.abc import ( + Mapping, + Sequence, +) +from typing import ( + Any, +) + +import numpy as np + +import deepmd.jax.utils.exclude_mask as _jax_exclude_mask # noqa: F401 +import deepmd.jax.utils.network as _jax_network # noqa: F401 +from deepmd.dpmodel.common import ( + NativeOP, +) +from deepmd.dpmodel.descriptor.dpa4 import DescrptDPA4 as DescrptDPA4DP +from deepmd.dpmodel.descriptor.dpa4_nn.activation import SwiGLU as SwiGLUDP +from deepmd.dpmodel.descriptor.dpa4_nn.grid_net import GridProduct as GridProductDP +from deepmd.dpmodel.descriptor.dpa4_nn.radial import ( + C3CutoffEnvelope as C3CutoffEnvelopeDP, +) +from deepmd.dpmodel.descriptor.dpa4_nn.radial import RadialMLP as RadialMLPDP +from deepmd.dpmodel.descriptor.dpa4_nn.so2 import SO2Linear as SO2LinearDP +from deepmd.dpmodel.descriptor.dpa4_nn.wignerd import ( + WignerDCalculator as WignerDCalculatorDP, +) +from deepmd.jax.common import ( + flax_module, + register_dpmodel_mapping, + to_jax_array, + try_convert_module, +) +from deepmd.jax.descriptor.base_descriptor import ( + BaseDescriptor, +) +from deepmd.jax.env import ( + jnp, + nnx, +) +from deepmd.jax.utils.network import ( + ArrayAPIParam, +) + + +@flax_module +class SwiGLU(SwiGLUDP): + pass + + +register_dpmodel_mapping(SwiGLUDP, lambda v: SwiGLU()) + + +@flax_module +class C3CutoffEnvelope(C3CutoffEnvelopeDP): + pass + + +register_dpmodel_mapping( + C3CutoffEnvelopeDP, + lambda v: C3CutoffEnvelope(v.rcut, v.p, precision=v.precision), +) + + +@flax_module +class RadialMLP(RadialMLPDP): + def __init__(self, *args: Any, **kwargs: Any) -> None: + super().__init__(*args, **kwargs) + self.net = nnx.List([self._convert_layer(layer) for layer in self.net]) + + @staticmethod + def _convert_layer(layer: Any) -> Any: + if isinstance(layer, nnx.Module): + return layer + if isinstance(layer, NativeOP): + converted = try_convert_module(layer) + if converted is not None: + return converted + return layer + + +register_dpmodel_mapping( + RadialMLPDP, + lambda v: RadialMLP.deserialize(v.serialize()), +) + + +@flax_module +class GridProduct(GridProductDP): + pass + + +register_dpmodel_mapping(GridProductDP, lambda v: GridProduct()) + + +@flax_module +class WignerDCalculator(WignerDCalculatorDP): + pass + + +register_dpmodel_mapping( + WignerDCalculatorDP, + lambda v: WignerDCalculator(v.lmax, eps=v.eps, precision=v.precision), +) + + +_TRAINABLE_ATTRS: dict[str, tuple[str, ...]] = { + "RMSNorm": ("adam_scale",), + "EquivariantRMSNorm": ("adam_scale", "bias"), + "ReducedEquivariantRMSNorm": ("adam_scale", "bias0"), + "ScalarRMSNorm": ("adam_scale",), + "RadialBasis": ("adam_freqs",), + "SO3Linear": ("weight", "bias"), + "FocusLinear": ("weight", "bias"), + "ChannelLinear": ("weight", "bias"), + "SO2Linear": ("weight_m0", "bias0"), + "DynamicRadialDegreeMixer": ("weight", "channel_basis"), + "SO2Convolution": ( + "adamw_attn_logit_w", + "adamw_attn_z_bias_raw", + "adamw_attn_gate_w", + "adamw_focus_compete_w", + "focus_compete_bias", + ), + "SeZMTypeEmbedding": ("adam_type_embedding",), + "SpinEmbedding": ("adam_spin_vec_weight", "adam_spin_nbr_weight"), + "EnvironmentInitialEmbedding": ("spin_scale",), + "DepthAttnRes": ("adamw_pseudo_query",), + "S2GridNet": ("residual_scale",), + "SO3GridNet": ("residual_scale",), + "DescrptDPA4": ("film_scale_strength_log", "film_shift_strength_log"), +} + +_TRAINABLE_LIST_ATTRS: dict[str, tuple[str, ...]] = { + "SO2Linear": ("weight_m",), + "SO2Convolution": ("adam_so2_layer_scales",), +} + + +def _is_array_like(value: Any) -> bool: + return hasattr(value, "shape") and hasattr(value, "dtype") + + +def _array_value(value: Any) -> Any: + if isinstance(value, nnx.Variable): + return value.value + return value + + +def _is_floating_array(value: Any) -> bool: + value = _array_value(value) + if value is None or not _is_array_like(value): + return False + return bool(jnp.issubdtype(value.dtype, jnp.floating)) + + +def _as_param(value: Any) -> Any: + if isinstance(value, ArrayAPIParam): + return value + if not _is_floating_array(value): + return value + if isinstance(value, nnx.Variable): + return ArrayAPIParam(value.value) + if isinstance(value, np.ndarray): + return ArrayAPIParam(to_jax_array(value)) + return ArrayAPIParam(value) + + +def _as_param_list(value: Any) -> Any: + if not isinstance(value, Sequence) or isinstance(value, (str, bytes)): + return value + promoted = [] + changed = False + for item in value: + new_item = _as_param(item) + promoted.append(new_item) + changed = changed or new_item is not item + if not changed: + return value + return nnx.List(promoted) if hasattr(nnx, "List") else promoted + + +def _iter_object_tree(root: Any) -> Any: + seen: set[int] = set() + + def visit(value: Any) -> Any: + if value is None or isinstance(value, (str, bytes, int, float, bool)): + return + if _is_array_like(value): + return + value_id = id(value) + if value_id in seen: + return + seen.add(value_id) + + if isinstance(value, Mapping): + for item in value.values(): + yield from visit(item) + return + if isinstance(value, Sequence): + for item in value: + yield from visit(item) + return + try: + value_dict = object.__getattribute__(value, "__dict__") + except AttributeError: + return + + yield value + for item in value_dict.values(): + yield from visit(item) + + yield from visit(root) + + +def _promote_trainable(module: Any, names: tuple[str, ...]) -> None: + if not getattr(module, "trainable", True): + return + for name in names: + if not hasattr(module, name): + continue + value = getattr(module, name) + new_value = _as_param(value) + if new_value is not value: + setattr(module, name, new_value) + + +def _promote_trainable_lists(module: Any, names: tuple[str, ...]) -> None: + if not getattr(module, "trainable", True): + return + for name in names: + if not hasattr(module, name): + continue + value = getattr(module, name) + new_value = _as_param_list(value) + if new_value is not value: + setattr(module, name, new_value) + + +def _promote_trainable_tree(module: Any) -> Any: + for submodule in _iter_object_tree(module): + names = _TRAINABLE_ATTRS.get(type(submodule).__name__) + if names is not None: + _promote_trainable(submodule, names) + list_names = _TRAINABLE_LIST_ATTRS.get(type(submodule).__name__) + if list_names is not None: + _promote_trainable_lists(submodule, list_names) + return module + + +@flax_module +class SO2Linear(SO2LinearDP): + def __init__(self, *args: Any, **kwargs: Any) -> None: + super().__init__(*args, **kwargs) + self.weight_m = _as_param_list(self.weight_m) + + @classmethod + def deserialize(cls, data: dict) -> "SO2Linear": + obj = super().deserialize(data) + obj.weight_m = _as_param_list(obj.weight_m) + return obj + + +register_dpmodel_mapping( + SO2LinearDP, + lambda v: SO2Linear.deserialize(v.serialize()), +) + + +@BaseDescriptor.register("SeZM") +@BaseDescriptor.register("sezm") +@BaseDescriptor.register("DPA4") +@BaseDescriptor.register("dpa4") +@flax_module +class DescrptDPA4(DescrptDPA4DP): + def __init__(self, *args: Any, **kwargs: Any) -> None: + super().__init__(*args, **kwargs) + _promote_trainable_tree(self) + + @classmethod + def deserialize(cls, data: dict) -> "DescrptDPA4": + obj = super().deserialize(data) + return _promote_trainable_tree(obj) diff --git a/deepmd/jax/fitting/__init__.py b/deepmd/jax/fitting/__init__.py index 77133e2bac..82cea44e1b 100644 --- a/deepmd/jax/fitting/__init__.py +++ b/deepmd/jax/fitting/__init__.py @@ -1,4 +1,7 @@ # SPDX-License-Identifier: LGPL-3.0-or-later +from deepmd.jax.fitting.dpa4_ener import ( + SeZMEnergyFittingNet, +) from deepmd.jax.fitting.fitting import ( DipoleFittingNet, DOSFittingNet, @@ -11,4 +14,5 @@ "DipoleFittingNet", "EnergyFittingNet", "PolarFittingNet", + "SeZMEnergyFittingNet", ] diff --git a/deepmd/jax/fitting/dpa4_ener.py b/deepmd/jax/fitting/dpa4_ener.py new file mode 100644 index 0000000000..b546a99afb --- /dev/null +++ b/deepmd/jax/fitting/dpa4_ener.py @@ -0,0 +1,54 @@ +# SPDX-License-Identifier: LGPL-3.0-or-later +from typing import ( + Any, + ClassVar, +) + +import deepmd.jax.utils.network as _jax_network # noqa: F401 +from deepmd.dpmodel.fitting.dpa4_ener import GLUFittingNet as GLUFittingNetDP +from deepmd.dpmodel.fitting.dpa4_ener import ( + SeZMEnergyFittingNet as SeZMEnergyFittingNetDP, +) +from deepmd.dpmodel.fitting.dpa4_ener import ( + SeZMNetworkCollection as SeZMNetworkCollectionDP, +) +from deepmd.jax.common import ( + flax_module, + register_dpmodel_mapping, +) +from deepmd.jax.fitting.base_fitting import ( + BaseFitting, +) + + +@flax_module +class GLUFittingNet(GLUFittingNetDP): + pass + + +register_dpmodel_mapping( + GLUFittingNetDP, + lambda v: GLUFittingNet.deserialize(v.serialize()), +) + + +@flax_module +class SeZMNetworkCollection(SeZMNetworkCollectionDP): + _jax_data_list_attrs: ClassVar[set[str]] = {"_networks", "networks"} + NETWORK_TYPE_MAP: ClassVar[dict[str, type]] = { + "sezm_fitting_network": GLUFittingNet, + } + + +register_dpmodel_mapping( + SeZMNetworkCollectionDP, + lambda v: SeZMNetworkCollection.deserialize(v.serialize()), +) + + +@BaseFitting.register("dpa4_ener") +@BaseFitting.register("sezm_ener") +@flax_module +class SeZMEnergyFittingNet(SeZMEnergyFittingNetDP): + def __setattr__(self, name: str, value: Any) -> None: + return super().__setattr__(name, value) diff --git a/deepmd/jax/model/ener_model.py b/deepmd/jax/model/ener_model.py index 1d3e8a1d80..626997d18e 100644 --- a/deepmd/jax/model/ener_model.py +++ b/deepmd/jax/model/ener_model.py @@ -11,6 +11,8 @@ ) +@BaseModel.register("sezm_ener") +@BaseModel.register("dpa4_ener") @BaseModel.register("ener") class EnergyModel(make_jax_dp_model_from_dpmodel(EnergyModelDP, DPAtomicModelEnergy)): pass diff --git a/deepmd/jax/model/model.py b/deepmd/jax/model/model.py index a3d067c636..73c03c5e78 100644 --- a/deepmd/jax/model/model.py +++ b/deepmd/jax/model/model.py @@ -109,6 +109,53 @@ def get_zbl_model(data: dict) -> DPZBLModel: ) +def get_sezm_model(data: dict) -> BaseModel: + """Build a DPA4/SeZM energy model from the pt-style model config.""" + data = deepcopy(data) + if "spin" in data: + raise NotImplementedError("Spin DPA4/SeZM models are not supported in JAX.") + if str(data.get("bridging_method", "none")).lower() != "none": + raise NotImplementedError("DPA4/SeZM bridging is not supported in JAX.") + if data.get("lora") is not None: + raise NotImplementedError("DPA4/SeZM LoRA is not supported in JAX.") + if data.get("use_compile"): + raise NotImplementedError("model.use_compile is not supported in JAX.") + if data.get("preset_out_bias"): + raise NotImplementedError("DPA4/SeZM preset_out_bias is not supported in JAX.") + + data.pop("type", None) + data.setdefault("descriptor", {}) + data.setdefault("fitting_net", {}) + data["descriptor"].setdefault("type", "dpa4") + data["fitting_net"].setdefault("type", "dpa4_ener") + if data["descriptor"]["type"] not in ("dpa4", "DPA4", "sezm", "SeZM"): + raise ValueError( + "Model type 'dpa4' requires a DPA4/SeZM descriptor, but got " + f"descriptor type '{data['descriptor']['type']}'." + ) + if data["fitting_net"]["type"] not in ("dpa4_ener", "sezm_ener"): + raise ValueError( + "Model type 'dpa4' requires the DPA4/SeZM energy fitting net, but got " + f"fitting_net type '{data['fitting_net']['type']}'." + ) + + descriptor_exclude_types = [ + list(pair) for pair in (data["descriptor"].get("exclude_types") or []) + ] + if "pair_exclude_types" in data: + pair_exclude_types = [list(pair) for pair in (data["pair_exclude_types"] or [])] + if descriptor_exclude_types and descriptor_exclude_types != pair_exclude_types: + raise ValueError( + "DPA4/SeZM pair_exclude_types and descriptor.exclude_types must " + "match when both are provided." + ) + else: + pair_exclude_types = descriptor_exclude_types + data["pair_exclude_types"] = pair_exclude_types + data["descriptor"]["exclude_types"] = deepcopy(pair_exclude_types) + return get_standard_model(data) + + def get_model(data: dict) -> BaseModel: """Get a model from a dictionary. @@ -125,5 +172,7 @@ def get_model(data: dict) -> BaseModel: return get_zbl_model(data) else: return get_standard_model(data) + elif model_type in ("SeZM", "sezm", "DPA4", "dpa4"): + return get_sezm_model(data) else: return BaseModel.get_class_by_type(model_type).get_model(data) diff --git a/deepmd/jax/train/trainer.py b/deepmd/jax/train/trainer.py index c77bc944b5..afa3decf09 100644 --- a/deepmd/jax/train/trainer.py +++ b/deepmd/jax/train/trainer.py @@ -598,6 +598,7 @@ def loss_fn( model_dict = _evaluate_model_dict( model, extended_coord, extended_atype, nlist, mapping, fp, ap ) + model_dict = _match_label_shapes(model_dict, label_dict) loss, _ = loss_obj( learning_rate=lr, natoms=label_dict["type"].shape[1], @@ -621,6 +622,7 @@ def loss_fn_more_loss( model_dict = _evaluate_model_dict( model, extended_coord, extended_atype, nlist, mapping, fp, ap ) + model_dict = _match_label_shapes(model_dict, label_dict) _, more_loss = loss_obj( learning_rate=lr, natoms=label_dict["type"].shape[1], @@ -831,6 +833,7 @@ def _write_checkpoint(self, ckpt_path: Path, *, step: int) -> None: else: _, single_state = nnx.split(self.models[DEFAULT_TASK_KEY]) state = single_state.to_pure_dict() + state = _drop_zero_size_array_leaves(state) if ckpt_path.is_dir(): shutil.rmtree(ckpt_path) model_def_script_cpy = deepcopy(self.model_def_script) @@ -888,11 +891,50 @@ def _evaluate_model_dict( ) model_dict["atom_energy"] = model_dict["energy"] model_dict["energy"] = model_dict["energy_redu"] - model_dict["force"] = model_dict["energy_derv_r"].squeeze(-2) + force = model_dict["energy_derv_r"].squeeze(-2) + if force.ndim == 2 or (force.ndim == 3 and force.shape[-1] != 3): + force = jnp.reshape(force, (force.shape[0], -1, 3)) + model_dict["force"] = force model_dict["virial"] = model_dict["energy_derv_c_redu"].squeeze(-2) return model_dict +def _match_label_shapes( + model_dict: dict[str, jnp.ndarray], + label_dict: dict[str, jnp.ndarray], +) -> dict[str, jnp.ndarray]: + """Match equivalent flattened model outputs to label tensor shapes.""" + force_hat = model_dict.get("force") + force = label_dict.get("force") + if ( + force_hat is not None + and force is not None + and force_hat.shape != force.shape + and force_hat.size == force.size + ): + model_dict = dict(model_dict) + model_dict["force"] = jnp.reshape(force_hat, force.shape) + return model_dict + + +_DROP_LEAF = object() + + +def _drop_zero_size_array_leaves(value: Any) -> Any: + """Drop zero-size arrays that Orbax cannot serialize.""" + if isinstance(value, dict): + filtered = {} + for key, item in value.items(): + new_item = _drop_zero_size_array_leaves(item) + if new_item is not _DROP_LEAF: + filtered[key] = new_item + return filtered + size = getattr(value, "size", None) + if size == 0: + return _DROP_LEAF + return value + + def _init_empty_state(params: Any) -> optax.EmptyState: """Initialize an empty Optax state without requiring optax.init_empty_state. diff --git a/source/tests/consistent/descriptor/test_dpa4.py b/source/tests/consistent/descriptor/test_dpa4.py index e6f3216bd4..e4a63594c0 100644 --- a/source/tests/consistent/descriptor/test_dpa4.py +++ b/source/tests/consistent/descriptor/test_dpa4.py @@ -19,6 +19,7 @@ ) from ..common import ( + INSTALLED_JAX, INSTALLED_PT, INSTALLED_PT_EXPT, CommonTest, @@ -36,6 +37,10 @@ from deepmd.pt_expt.descriptor.dpa4 import DescrptDPA4 as DescrptDPA4PTExpt else: DescrptDPA4PTExpt = None +if INSTALLED_JAX: + from deepmd.jax.descriptor.dpa4 import DescrptDPA4 as DescrptDPA4JAX +else: + DescrptDPA4JAX = None # not implemented DescrptDPA4TF = None @@ -150,7 +155,7 @@ def skip_pt(self) -> bool: skip_dp = False skip_tf = True - skip_jax = True + skip_jax = not INSTALLED_JAX or DescrptDPA4JAX is None skip_pd = True skip_pt_expt = not INSTALLED_PT_EXPT skip_array_api_strict = True @@ -159,7 +164,7 @@ def skip_pt(self) -> bool: dp_class = DescrptDPA4DP pt_class = DescrptDPA4PT pt_expt_class = DescrptDPA4PTExpt - jax_class = None + jax_class = DescrptDPA4JAX pd_class = None array_api_strict_class = None args: ClassVar[list] = [ @@ -234,6 +239,16 @@ def eval_pt_expt(self, pt_expt_obj: Any) -> Any: mixed_types=True, ) + def eval_jax(self, jax_obj: Any) -> Any: + return self.eval_jax_descriptor( + jax_obj, + self.natoms, + self.coords, + self.atype, + self.box, + mixed_types=True, + ) + def extract_ret(self, ret: Any, backend) -> tuple[np.ndarray, ...]: return (ret[0],) diff --git a/source/tests/consistent/fitting/test_dpa4_ener.py b/source/tests/consistent/fitting/test_dpa4_ener.py index ee10b54343..e8a7e5043a 100644 --- a/source/tests/consistent/fitting/test_dpa4_ener.py +++ b/source/tests/consistent/fitting/test_dpa4_ener.py @@ -15,6 +15,7 @@ ) from ..common import ( + INSTALLED_JAX, INSTALLED_PT, INSTALLED_PT_EXPT, CommonTest, @@ -40,6 +41,13 @@ from deepmd.pt_expt.utils.env import DEVICE as PT_EXPT_DEVICE else: SeZMEnerFittingPTExpt = None +if INSTALLED_JAX: + from deepmd.jax.env import ( + jnp, + ) + from deepmd.jax.fitting.dpa4_ener import SeZMEnergyFittingNet as SeZMEnerFittingJAX +else: + SeZMEnerFittingJAX = None # not implemented SeZMEnerFittingTF = None @@ -74,7 +82,7 @@ def skip_pt(self) -> bool: skip_dp = False skip_tf = True - skip_jax = True + skip_jax = not INSTALLED_JAX or SeZMEnerFittingJAX is None skip_pd = True skip_pt_expt = not INSTALLED_PT_EXPT skip_array_api_strict = True @@ -83,7 +91,7 @@ def skip_pt(self) -> bool: dp_class = SeZMEnerFittingDP pt_class = SeZMEnerFittingPT pt_expt_class = SeZMEnerFittingPTExpt - jax_class = None + jax_class = SeZMEnerFittingJAX pd_class = None array_api_strict_class = None args = fitting_sezm_ener() @@ -138,6 +146,14 @@ def eval_pt_expt(self, pt_expt_obj: Any) -> Any: .numpy() ) + def eval_jax(self, jax_obj: Any) -> Any: + return np.asarray( + jax_obj( + jnp.asarray(self.inputs), + jnp.asarray(self.atype.reshape(1, -1)), + )["energy"] + ) + def extract_ret(self, ret: Any, backend) -> tuple[np.ndarray, ...]: return (ret,) From 590823bb200dda5159263944d45d18320168df77 Mon Sep 17 00:00:00 2001 From: Jinzhe Zeng Date: Wed, 8 Jul 2026 21:30:26 +0800 Subject: [PATCH 02/11] fix(jax): address dpa4 ci feedback --- deepmd/jax/descriptor/dpa4.py | 3 ++- deepmd/jax/fitting/dpa4_ener.py | 4 +--- deepmd/jax/model/model.py | 4 ++-- deepmd/jax/utils/serialization.py | 38 +++++++++++++++++++++++++++++++ 4 files changed, 43 insertions(+), 6 deletions(-) diff --git a/deepmd/jax/descriptor/dpa4.py b/deepmd/jax/descriptor/dpa4.py index b97b14c082..ba4c87a5bd 100644 --- a/deepmd/jax/descriptor/dpa4.py +++ b/deepmd/jax/descriptor/dpa4.py @@ -66,7 +66,8 @@ class C3CutoffEnvelope(C3CutoffEnvelopeDP): class RadialMLP(RadialMLPDP): def __init__(self, *args: Any, **kwargs: Any) -> None: super().__init__(*args, **kwargs) - self.net = nnx.List([self._convert_layer(layer) for layer in self.net]) + converted = [self._convert_layer(layer) for layer in self.net] + self.net = nnx.List(converted) if hasattr(nnx, "List") else converted @staticmethod def _convert_layer(layer: Any) -> Any: diff --git a/deepmd/jax/fitting/dpa4_ener.py b/deepmd/jax/fitting/dpa4_ener.py index b546a99afb..5b2a5c6253 100644 --- a/deepmd/jax/fitting/dpa4_ener.py +++ b/deepmd/jax/fitting/dpa4_ener.py @@ -1,6 +1,5 @@ # SPDX-License-Identifier: LGPL-3.0-or-later from typing import ( - Any, ClassVar, ) @@ -50,5 +49,4 @@ class SeZMNetworkCollection(SeZMNetworkCollectionDP): @BaseFitting.register("sezm_ener") @flax_module class SeZMEnergyFittingNet(SeZMEnergyFittingNetDP): - def __setattr__(self, name: str, value: Any) -> None: - return super().__setattr__(name, value) + pass diff --git a/deepmd/jax/model/model.py b/deepmd/jax/model/model.py index 73c03c5e78..e26e68bce2 100644 --- a/deepmd/jax/model/model.py +++ b/deepmd/jax/model/model.py @@ -124,8 +124,8 @@ def get_sezm_model(data: dict) -> BaseModel: raise NotImplementedError("DPA4/SeZM preset_out_bias is not supported in JAX.") data.pop("type", None) - data.setdefault("descriptor", {}) - data.setdefault("fitting_net", {}) + data["descriptor"] = data.get("descriptor") or {} + data["fitting_net"] = data.get("fitting_net") or {} data["descriptor"].setdefault("type", "dpa4") data["fitting_net"].setdefault("type", "dpa4_ener") if data["descriptor"]["type"] not in ("dpa4", "DPA4", "sezm", "SeZM"): diff --git a/deepmd/jax/utils/serialization.py b/deepmd/jax/utils/serialization.py index 59240b41ab..e017193fc3 100644 --- a/deepmd/jax/utils/serialization.py +++ b/deepmd/jax/utils/serialization.py @@ -50,6 +50,40 @@ def _normalize_restored_state_keys( _convert_str_to_int_key(state) +_NO_ZERO_SIZE_LEAF = object() + + +def _zero_size_subtree(value: Any) -> Any: + if isinstance(value, dict): + restored = {} + for key, item in value.items(): + subtree = _zero_size_subtree(item) + if subtree is not _NO_ZERO_SIZE_LEAF: + restored[key] = subtree + return restored if restored else _NO_ZERO_SIZE_LEAF + if getattr(value, "size", None) == 0: + return value + return _NO_ZERO_SIZE_LEAF + + +def _restore_missing_zero_size_leaves(template: Any, restored: Any) -> Any: + """Reinsert zero-size leaves dropped before Orbax checkpoint saving.""" + if not isinstance(template, dict) or not isinstance(restored, dict): + return restored + restored = dict(restored) + for key, template_value in template.items(): + if key in restored: + restored[key] = _restore_missing_zero_size_leaves( + template_value, + restored[key], + ) + continue + subtree = _zero_size_subtree(template_value) + if subtree is not _NO_ZERO_SIZE_LEAF: + restored[key] = subtree + return restored + + def _state_sequence_to_numpy_list(state_value: Any) -> list[np.ndarray]: """Convert an Orbax-restored list/dict sequence to NumPy arrays.""" if isinstance(state_value, dict): @@ -369,6 +403,10 @@ def restore_model(model_params: dict, model_state: dict) -> BaseModel: abstract_model = get_model(model_params) _restore_compression_slots_from_state(abstract_model, model_state) graphdef, abstract_state = nnx.split(abstract_model) + model_state = _restore_missing_zero_size_leaves( + abstract_state.to_pure_dict(), + model_state, + ) abstract_state.replace_by_pure_dict(model_state) return nnx.merge(graphdef, abstract_state) From 1c8c03b956c5589cd3a6373b30df07b0191d3f32 Mon Sep 17 00:00:00 2001 From: Jinzhe Zeng Date: Fri, 10 Jul 2026 14:47:10 +0800 Subject: [PATCH 03/11] fix(jax): remove redundant dpa4 network import --- deepmd/jax/descriptor/dpa4.py | 1 - 1 file changed, 1 deletion(-) diff --git a/deepmd/jax/descriptor/dpa4.py b/deepmd/jax/descriptor/dpa4.py index ba4c87a5bd..6655bec50f 100644 --- a/deepmd/jax/descriptor/dpa4.py +++ b/deepmd/jax/descriptor/dpa4.py @@ -10,7 +10,6 @@ import numpy as np import deepmd.jax.utils.exclude_mask as _jax_exclude_mask # noqa: F401 -import deepmd.jax.utils.network as _jax_network # noqa: F401 from deepmd.dpmodel.common import ( NativeOP, ) From 68cb12e0d461a5359a148638eba06742e1582491 Mon Sep 17 00:00:00 2001 From: njzjz-bot Date: Sat, 11 Jul 2026 14:22:07 +0800 Subject: [PATCH 04/11] fix(jax): address dpa4 review feedback Coding-Agent: Codex Codex-Version: codex-cli 0.144.1 Model: gpt-5.6-sol Reasoning-Effort: xhigh --- deepmd/jax/descriptor/dpa4.py | 4 +- deepmd/jax/fitting/dpa4_ener.py | 1 - source/tests/jax/test_dpa4.py | 58 +++++++++++++++ source/tests/jax/test_model_factory.py | 77 ++++++++++++++++++++ source/tests/jax/test_training.py | 97 ++++++++++++++++++++++++++ 5 files changed, 235 insertions(+), 2 deletions(-) create mode 100644 source/tests/jax/test_dpa4.py diff --git a/deepmd/jax/descriptor/dpa4.py b/deepmd/jax/descriptor/dpa4.py index 6655bec50f..133bc468d2 100644 --- a/deepmd/jax/descriptor/dpa4.py +++ b/deepmd/jax/descriptor/dpa4.py @@ -9,7 +9,6 @@ import numpy as np -import deepmd.jax.utils.exclude_mask as _jax_exclude_mask # noqa: F401 from deepmd.dpmodel.common import ( NativeOP, ) @@ -113,6 +112,8 @@ class WignerDCalculator(WignerDCalculatorDP): "SO3Linear": ("weight", "bias"), "FocusLinear": ("weight", "bias"), "ChannelLinear": ("weight", "bias"), + "FrameContract": ("weight",), + "FrameExpand": ("weight",), "SO2Linear": ("weight_m0", "bias0"), "DynamicRadialDegreeMixer": ("weight", "channel_basis"), "SO2Convolution": ( @@ -132,6 +133,7 @@ class WignerDCalculator(WignerDCalculatorDP): } _TRAINABLE_LIST_ATTRS: dict[str, tuple[str, ...]] = { + "SeZMInteractionBlock": ("adam_ffn_layer_scales",), "SO2Linear": ("weight_m",), "SO2Convolution": ("adam_so2_layer_scales",), } diff --git a/deepmd/jax/fitting/dpa4_ener.py b/deepmd/jax/fitting/dpa4_ener.py index 5b2a5c6253..425268cd78 100644 --- a/deepmd/jax/fitting/dpa4_ener.py +++ b/deepmd/jax/fitting/dpa4_ener.py @@ -3,7 +3,6 @@ ClassVar, ) -import deepmd.jax.utils.network as _jax_network # noqa: F401 from deepmd.dpmodel.fitting.dpa4_ener import GLUFittingNet as GLUFittingNetDP from deepmd.dpmodel.fitting.dpa4_ener import ( SeZMEnergyFittingNet as SeZMEnergyFittingNetDP, diff --git a/source/tests/jax/test_dpa4.py b/source/tests/jax/test_dpa4.py new file mode 100644 index 0000000000..af3b876676 --- /dev/null +++ b/source/tests/jax/test_dpa4.py @@ -0,0 +1,58 @@ +# SPDX-License-Identifier: LGPL-3.0-or-later +"""Focused tests for JAX DPA4 trainable-state conversion.""" + +from deepmd.jax.descriptor.dpa4 import ( + DescrptDPA4, + _iter_object_tree, +) +from deepmd.jax.utils.network import ( + ArrayAPIParam, +) + + +def _make_trainable_descriptor() -> DescrptDPA4: + """Build a small descriptor that enables the optional trainable leaves.""" + return DescrptDPA4( + ntypes=2, + sel=4, + rcut=4.0, + channels=4, + n_radial=4, + lmax=1, + mmax=1, + n_blocks=1, + grid_branch=0, + layer_scale=True, + message_node_so3=True, + random_gamma=False, + precision="float64", + trainable=True, + seed=20260711, + ) + + +def test_optional_dpa4_weights_are_jax_parameters() -> None: + """Optional cross-grid and FFN LayerScale weights must receive gradients.""" + descriptor = _make_trainable_descriptor() + modules = list(_iter_object_tree(descriptor)) + + frame_modules = [ + module + for module in modules + if type(module).__name__ in {"FrameContract", "FrameExpand"} + ] + assert {type(module).__name__ for module in frame_modules} == { + "FrameContract", + "FrameExpand", + } + assert all(isinstance(module.weight, ArrayAPIParam) for module in frame_modules) + + interaction_blocks = [ + module for module in modules if type(module).__name__ == "SeZMInteractionBlock" + ] + assert interaction_blocks + assert all( + isinstance(scale, ArrayAPIParam) + for block in interaction_blocks + for scale in block.adam_ffn_layer_scales + ) diff --git a/source/tests/jax/test_model_factory.py b/source/tests/jax/test_model_factory.py index 75ffc519a1..da09f80a46 100644 --- a/source/tests/jax/test_model_factory.py +++ b/source/tests/jax/test_model_factory.py @@ -10,6 +10,9 @@ """ import unittest +from unittest.mock import ( + patch, +) from deepmd.jax.model.ener_model import ( EnergyModel, @@ -40,6 +43,16 @@ def _base_config() -> dict: } +def _base_sezm_config() -> dict: + """Return the smallest config needed to exercise DPA4 factory routing.""" + return { + "type": "dpa4", + "type_map": ["O", "H"], + "descriptor": {"type": "dpa4"}, + "fitting_net": {"type": "dpa4_ener"}, + } + + class TestJAXModelFactoryFittingDefault(unittest.TestCase): def test_fitting_net_without_type_defaults_to_ener(self) -> None: # fitting_net present but no "type": must default to energy. @@ -62,5 +75,69 @@ def test_explicit_fitting_type_preserved(self) -> None: self.assertIsInstance(model, EnergyModel) +class TestJAXSeZMModelFactory(unittest.TestCase): + @patch("deepmd.jax.model.model.get_standard_model", side_effect=lambda data: data) + def test_null_blocks_receive_dpa4_defaults(self, _get_standard_model) -> None: + data = _base_sezm_config() + data["descriptor"] = None + data["fitting_net"] = None + + normalized = get_model(data) + + self.assertEqual(normalized["descriptor"]["type"], "dpa4") + self.assertEqual(normalized["fitting_net"]["type"], "dpa4_ener") + + def test_rejects_unsupported_features(self) -> None: + cases = ( + ("spin", {}), + ("bridging_method", "linear"), + ("lora", {}), + ("use_compile", True), + ("preset_out_bias", [0.0]), + ) + for key, value in cases: + with self.subTest(key=key): + data = _base_sezm_config() + data[key] = value + with self.assertRaises(NotImplementedError): + get_model(data) + + def test_rejects_incompatible_descriptor_and_fitting_types(self) -> None: + data = _base_sezm_config() + data["descriptor"]["type"] = "se_e2_a" + with self.assertRaises(ValueError): + get_model(data) + + data = _base_sezm_config() + data["fitting_net"]["type"] = "ener" + with self.assertRaises(ValueError): + get_model(data) + + def test_rejects_mismatched_exclude_types(self) -> None: + data = _base_sezm_config() + data["descriptor"]["exclude_types"] = [[0, 1]] + data["pair_exclude_types"] = [[1, 1]] + + with self.assertRaises(ValueError): + get_model(data) + + @patch("deepmd.jax.model.model.get_standard_model", side_effect=lambda data: data) + def test_descriptor_exclude_types_feed_standard_model( + self, + _get_standard_model, + ) -> None: + data = _base_sezm_config() + data["descriptor"] = { + "type": "SeZM", + "exclude_types": [[0, 1]], + } + data["fitting_net"]["type"] = "sezm_ener" + + normalized = get_model(data) + + self.assertEqual(normalized["pair_exclude_types"], [[0, 1]]) + self.assertEqual(normalized["descriptor"]["exclude_types"], [[0, 1]]) + + if __name__ == "__main__": unittest.main() diff --git a/source/tests/jax/test_training.py b/source/tests/jax/test_training.py index 25d3ccdc49..a1fd2196cf 100644 --- a/source/tests/jax/test_training.py +++ b/source/tests/jax/test_training.py @@ -43,6 +43,9 @@ from deepmd.jax.train.trainer import ( DPTrainer, _copy_matching_state_tree, + _drop_zero_size_array_leaves, + _evaluate_model_dict, + _match_label_shapes, _scale_by_global_learning_rate, ) from deepmd.jax.utils.finetune import ( @@ -50,6 +53,7 @@ ) from deepmd.jax.utils.serialization import ( _normalize_restored_state_keys, + _restore_missing_zero_size_leaves, ) from deepmd.utils.compat import ( convert_optimizer_v31_to_v32, @@ -604,3 +608,96 @@ def test_jax_multitask_state_key_normalization_preserves_numeric_task_names() -> assert 1 not in state["models"] assert 0 in state["models"]["1"]["layers"] assert 0 in state["models"]["task"]["layers"] + + +def test_jax_zero_size_checkpoint_leaves_round_trip() -> None: + """Checkpoint filtering and restore must preserve every zero-size path.""" + template = { + "model": { + "empty": jnp.zeros((0, 3)), + "nested": { + "empty": jnp.zeros((2, 0)), + "weight": jnp.ones((2,)), + }, + } + } + + filtered = _drop_zero_size_array_leaves(template) + restored = _restore_missing_zero_size_leaves(template, filtered) + + assert "empty" not in filtered["model"] + assert "empty" not in filtered["model"]["nested"] + np.testing.assert_array_equal( + restored["model"]["empty"], template["model"]["empty"] + ) + np.testing.assert_array_equal( + restored["model"]["nested"]["empty"], + template["model"]["nested"]["empty"], + ) + np.testing.assert_array_equal( + restored["model"]["nested"]["weight"], + template["model"]["nested"]["weight"], + ) + + +def test_jax_match_label_shapes_reshapes_only_equivalent_force_layouts() -> None: + """Flattened force tensors reshape, while an existing layout is untouched.""" + force = jnp.arange(6).reshape(1, 2, 3) + model_dict = {"force": force} + + reshaped = _match_label_shapes(model_dict, {"force": jnp.zeros((1, 6))}) + unchanged = _match_label_shapes(model_dict, {"force": jnp.zeros((1, 2, 3))}) + + assert reshaped is not model_dict + assert reshaped["force"].shape == (1, 6) + assert unchanged is model_dict + + +def test_jax_evaluate_model_dict_normalizes_flattened_force_only() -> None: + """Model evaluation preserves canonical force tensors and expands flat ones.""" + + class FakeModel: + def __init__(self, force_derivative: jnp.ndarray) -> None: + self.force_derivative = force_derivative + + def call_common_lower(self, *args, **kwargs): + del args, kwargs + return { + "energy": jnp.zeros((1, 2, 1)), + "energy_redu": jnp.zeros((1, 1)), + "energy_derv_r": self.force_derivative, + "energy_derv_c_redu": jnp.zeros((1, 1, 9)), + } + + def model_output_def(self): + return {} + + def passthrough(model_dict, *args, **kwargs): + del args, kwargs + return dict(model_dict) + + with patch( + "deepmd.jax.train.trainer.communicate_extended_output", + side_effect=passthrough, + ): + flattened = _evaluate_model_dict( + FakeModel(jnp.arange(6).reshape(1, 1, 6)), + jnp.zeros((1, 6)), + jnp.zeros((1, 2), dtype=jnp.int32), + jnp.zeros((1, 2, 1), dtype=jnp.int32), + None, + None, + None, + ) + canonical = _evaluate_model_dict( + FakeModel(jnp.arange(6).reshape(1, 2, 1, 3)), + jnp.zeros((1, 6)), + jnp.zeros((1, 2), dtype=jnp.int32), + jnp.zeros((1, 2, 1), dtype=jnp.int32), + None, + None, + None, + ) + + assert flattened["force"].shape == (1, 2, 3) + assert canonical["force"].shape == (1, 2, 3) From 7304a6773514a68e7032ec5fa89b887377fde3a5 Mon Sep 17 00:00:00 2001 From: njzjz-bot Date: Sun, 12 Jul 2026 19:30:58 +0800 Subject: [PATCH 05/11] fix(jax): address new dpa4 review feedback Preserve DPA4 freeze policies, support PT SeZM conversion, and apply zero-size filtering to every JAX checkpoint writer. Coding-Agent: Codex Codex-Version: codex-cli 0.144.1 Model: gpt-5.6-sol Reasoning-Effort: xhigh --- deepmd/dpmodel/fitting/dpa4_ener.py | 18 ++++-- deepmd/jax/descriptor/dpa4.py | 49 +++++++++------ deepmd/jax/model/base_model.py | 65 +++++++++++++++++++- deepmd/jax/model/model.py | 8 ++- deepmd/jax/train/trainer.py | 19 +----- deepmd/jax/utils/serialization.py | 18 ++++++ source/tests/jax/test_dpa4.py | 44 ++++++++++++++ source/tests/jax/test_dpa4_conversion.py | 76 ++++++++++++++++++++++++ source/tests/jax/test_model_factory.py | 23 +++++++ 9 files changed, 276 insertions(+), 44 deletions(-) create mode 100644 source/tests/jax/test_dpa4_conversion.py diff --git a/deepmd/dpmodel/fitting/dpa4_ener.py b/deepmd/dpmodel/fitting/dpa4_ener.py index f08097f032..6f7f970947 100644 --- a/deepmd/dpmodel/fitting/dpa4_ener.py +++ b/deepmd/dpmodel/fitting/dpa4_ener.py @@ -92,11 +92,18 @@ def __init__( ) if neuron is None: neuron = [] - if isinstance(trainable, list): - trainable = all(trainable) self.in_dim = int(in_dim) self.out_dim = int(out_dim) self.neuron = [int(nn_dim) for nn_dim in neuron] + if isinstance(trainable, bool): + self.trainable = [trainable] * (len(self.neuron) + 1) + else: + self.trainable = [bool(flag) for flag in trainable] + if len(self.trainable) != len(self.neuron) + 1: + raise ValueError( + "trainable must contain one flag per hidden layer plus " + "one flag for the output layer" + ) self.activation_function = activation_function self.resnet_dt = bool(resnet_dt) self.precision = precision @@ -123,7 +130,7 @@ def __init__( resnet=False, precision=self.precision, seed=child_seed(seed, layer_idx), - trainable=trainable, + trainable=self.trainable[layer_idx], ) ) dim_in = hidden_dim @@ -139,7 +146,7 @@ def __init__( resnet=False, precision=self.precision, seed=child_seed(seed, len(self.neuron) + int(self.case_film_embd)), - trainable=trainable, + trainable=self.trainable[-1], ) def call_until_last(self, xx: Array) -> Array: @@ -181,6 +188,9 @@ def serialize(self) -> dict[str, Any]: "descriptor_dim": self.descriptor_dim, "dim_case_embd": self.dim_case_embd, "case_film_embd": self.case_film_embd, + # Preserve the effective per-layer freeze policy when backend + # wrappers rebuild this network from its serialized form. + "trainable": self.trainable.copy(), "@variables": variables, } diff --git a/deepmd/jax/descriptor/dpa4.py b/deepmd/jax/descriptor/dpa4.py index 133bc468d2..cda6507599 100644 --- a/deepmd/jax/descriptor/dpa4.py +++ b/deepmd/jax/descriptor/dpa4.py @@ -24,6 +24,7 @@ WignerDCalculator as WignerDCalculatorDP, ) from deepmd.jax.common import ( + ArrayAPIVariable, flax_module, register_dpmodel_mapping, to_jax_array, @@ -156,25 +157,27 @@ def _is_floating_array(value: Any) -> bool: return bool(jnp.issubdtype(value.dtype, jnp.floating)) -def _as_param(value: Any) -> Any: - if isinstance(value, ArrayAPIParam): +def _as_parameter_variable(value: Any, *, trainable: bool) -> Any: + """Track a floating parameter with the requested optimizer visibility.""" + variable_type = ArrayAPIParam if trainable else ArrayAPIVariable + if type(value) is variable_type: return value if not _is_floating_array(value): return value if isinstance(value, nnx.Variable): - return ArrayAPIParam(value.value) + value = value.value if isinstance(value, np.ndarray): - return ArrayAPIParam(to_jax_array(value)) - return ArrayAPIParam(value) + value = to_jax_array(value) + return variable_type(value) -def _as_param_list(value: Any) -> Any: +def _as_parameter_variable_list(value: Any, *, trainable: bool) -> Any: if not isinstance(value, Sequence) or isinstance(value, (str, bytes)): return value promoted = [] changed = False for item in value: - new_item = _as_param(item) + new_item = _as_parameter_variable(item, trainable=trainable) promoted.append(new_item) changed = changed or new_item is not item if not changed: @@ -215,38 +218,42 @@ def visit(value: Any) -> Any: yield from visit(root) -def _promote_trainable(module: Any, names: tuple[str, ...]) -> None: - if not getattr(module, "trainable", True): - return +def _promote_parameters( + module: Any, names: tuple[str, ...], *, trainable: bool +) -> None: for name in names: if not hasattr(module, name): continue value = getattr(module, name) - new_value = _as_param(value) + new_value = _as_parameter_variable(value, trainable=trainable) if new_value is not value: setattr(module, name, new_value) -def _promote_trainable_lists(module: Any, names: tuple[str, ...]) -> None: - if not getattr(module, "trainable", True): - return +def _promote_parameter_lists( + module: Any, names: tuple[str, ...], *, trainable: bool +) -> None: for name in names: if not hasattr(module, name): continue value = getattr(module, name) - new_value = _as_param_list(value) + new_value = _as_parameter_variable_list(value, trainable=trainable) if new_value is not value: setattr(module, name, new_value) def _promote_trainable_tree(module: Any) -> Any: + root_trainable = bool(getattr(module, "trainable", True)) for submodule in _iter_object_tree(module): + # A frozen descriptor freezes every descendant, including helper + # modules such as RadialBasis that do not carry a local flag. + trainable = root_trainable and bool(getattr(submodule, "trainable", True)) names = _TRAINABLE_ATTRS.get(type(submodule).__name__) if names is not None: - _promote_trainable(submodule, names) + _promote_parameters(submodule, names, trainable=trainable) list_names = _TRAINABLE_LIST_ATTRS.get(type(submodule).__name__) if list_names is not None: - _promote_trainable_lists(submodule, list_names) + _promote_parameter_lists(submodule, list_names, trainable=trainable) return module @@ -254,12 +261,16 @@ def _promote_trainable_tree(module: Any) -> Any: class SO2Linear(SO2LinearDP): def __init__(self, *args: Any, **kwargs: Any) -> None: super().__init__(*args, **kwargs) - self.weight_m = _as_param_list(self.weight_m) + self.weight_m = _as_parameter_variable_list( + self.weight_m, trainable=bool(self.trainable) + ) @classmethod def deserialize(cls, data: dict) -> "SO2Linear": obj = super().deserialize(data) - obj.weight_m = _as_param_list(obj.weight_m) + obj.weight_m = _as_parameter_variable_list( + obj.weight_m, trainable=bool(obj.trainable) + ) return obj diff --git a/deepmd/jax/model/base_model.py b/deepmd/jax/model/base_model.py index 8a5a55e8d7..1a341da902 100644 --- a/deepmd/jax/model/base_model.py +++ b/deepmd/jax/model/base_model.py @@ -1,5 +1,9 @@ # SPDX-License-Identifier: LGPL-3.0-or-later +from typing import ( + Any, +) + from deepmd.dpmodel.model.base_model import ( make_base_model, ) @@ -12,8 +16,67 @@ jax, jnp, ) +from deepmd.utils.version import ( + check_version_compatibility, +) + + +class BaseModel(make_base_model()): + """JAX model registry with adapters for regular PT SeZM checkpoints.""" + + _SEZM_MODEL_TYPES = frozenset({"sezm", "dpa4"}) + _SEZM_ATOMIC_TYPES = frozenset({"sezm_atomic"}) -BaseModel = make_base_model() + @classmethod + def deserialize(cls, data: dict[str, Any]) -> "BaseModel": + model_type = str(data.get("type", "standard")).lower() + if model_type in cls._SEZM_MODEL_TYPES: + return cls.deserialize(cls._unwrap_pt_sezm_model(data)) + if model_type in cls._SEZM_ATOMIC_TYPES: + return cls.deserialize(cls._normalize_pt_sezm_atomic(data)) + return super().deserialize(data) + + @staticmethod + def _unwrap_pt_sezm_model(data: dict[str, Any]) -> dict[str, Any]: + """Unwrap PT's model-level SeZM schema after validating its extras.""" + check_version_compatibility(int(data.get("@version", 1)), 1, 1) + if str(data.get("bridging_method", "none")).lower() not in ("none", ""): + raise NotImplementedError( + "PT SeZM/DPA4 checkpoints with bridging are not supported in JAX." + ) + if data.get("lora") is not None: + raise NotImplementedError( + "PT SeZM/DPA4 checkpoints with LoRA are not supported in JAX." + ) + atomic_model = data.get("atomic_model") + if atomic_model is None: + raise ValueError("SeZM/DPA4 model data is missing 'atomic_model'.") + return atomic_model + + @staticmethod + def _normalize_pt_sezm_atomic(data: dict[str, Any]) -> dict[str, Any]: + """Convert PT's energy-only ``sezm_atomic`` schema to ``standard``.""" + data = data.copy() + check_version_compatibility(int(data.get("@version", 2)), 3, 2) + if data.pop("dens_fitting", None) is not None: + raise NotImplementedError( + "PT SeZM/DPA4 checkpoints with a dens head are not supported in JAX." + ) + active_mode = data.pop("active_mode", None) + if active_mode not in (None, "ener"): + raise NotImplementedError( + f"PT SeZM/DPA4 active_mode {active_mode!r} is not supported in JAX." + ) + variables = data.get("@variables") + if isinstance(variables, dict): + data["@variables"] = { + key: value + for key, value in variables.items() + if key in ("out_bias", "out_std") + } + data["@version"] = 2 + data["type"] = "standard" + return data def forward_common_atomic( diff --git a/deepmd/jax/model/model.py b/deepmd/jax/model/model.py index e26e68bce2..31ddc12073 100644 --- a/deepmd/jax/model/model.py +++ b/deepmd/jax/model/model.py @@ -128,6 +128,10 @@ def get_sezm_model(data: dict) -> BaseModel: data["fitting_net"] = data.get("fitting_net") or {} data["descriptor"].setdefault("type", "dpa4") data["fitting_net"].setdefault("type", "dpa4_ener") + if data["descriptor"].get("add_chg_spin_ebd"): + raise NotImplementedError( + "DPA4/SeZM charge/spin conditioning is not supported in JAX." + ) if data["descriptor"]["type"] not in ("dpa4", "DPA4", "sezm", "SeZM"): raise ValueError( "Model type 'dpa4' requires a DPA4/SeZM descriptor, but got " @@ -142,8 +146,8 @@ def get_sezm_model(data: dict) -> BaseModel: descriptor_exclude_types = [ list(pair) for pair in (data["descriptor"].get("exclude_types") or []) ] - if "pair_exclude_types" in data: - pair_exclude_types = [list(pair) for pair in (data["pair_exclude_types"] or [])] + pair_exclude_types = [list(pair) for pair in (data.get("pair_exclude_types") or [])] + if pair_exclude_types: if descriptor_exclude_types and descriptor_exclude_types != pair_exclude_types: raise ValueError( "DPA4/SeZM pair_exclude_types and descriptor.exclude_types must " diff --git a/deepmd/jax/train/trainer.py b/deepmd/jax/train/trainer.py index afa3decf09..526e51f6bf 100644 --- a/deepmd/jax/train/trainer.py +++ b/deepmd/jax/train/trainer.py @@ -71,6 +71,7 @@ get_model, ) from deepmd.jax.utils.serialization import ( + _drop_zero_size_array_leaves, serialize_from_file, ) from deepmd.utils.argcheck import ( @@ -917,24 +918,6 @@ def _match_label_shapes( return model_dict -_DROP_LEAF = object() - - -def _drop_zero_size_array_leaves(value: Any) -> Any: - """Drop zero-size arrays that Orbax cannot serialize.""" - if isinstance(value, dict): - filtered = {} - for key, item in value.items(): - new_item = _drop_zero_size_array_leaves(item) - if new_item is not _DROP_LEAF: - filtered[key] = new_item - return filtered - size = getattr(value, "size", None) - if size == 0: - return _DROP_LEAF - return value - - def _init_empty_state(params: Any) -> optax.EmptyState: """Initialize an empty Optax state without requiring optax.init_empty_state. diff --git a/deepmd/jax/utils/serialization.py b/deepmd/jax/utils/serialization.py index e017193fc3..76b4208d9f 100644 --- a/deepmd/jax/utils/serialization.py +++ b/deepmd/jax/utils/serialization.py @@ -25,6 +25,23 @@ ) +_DROP_ZERO_SIZE_LEAF = object() + + +def _drop_zero_size_array_leaves(value: Any) -> Any: + """Remove array leaves Orbax cannot save while preserving tree containers.""" + if isinstance(value, dict): + filtered = {} + for key, item in value.items(): + new_item = _drop_zero_size_array_leaves(item) + if new_item is not _DROP_ZERO_SIZE_LEAF: + filtered[key] = new_item + return filtered + if getattr(value, "size", None) == 0: + return _DROP_ZERO_SIZE_LEAF + return value + + def _convert_str_to_int_key(item: dict) -> None: """Convert Orbax-restored numeric index keys from strings back to ints.""" for key, value in item.copy().items(): @@ -253,6 +270,7 @@ def deserialize_to_file(model_file: str, data: dict) -> None: model = BaseModel.deserialize(data["model"]) _, state = nnx.split(model) state = state.to_pure_dict() + state = _drop_zero_size_array_leaves(state) with ocp.Checkpointer( ocp.CompositeCheckpointHandler("state", "model_def_script") ) as checkpointer: diff --git a/source/tests/jax/test_dpa4.py b/source/tests/jax/test_dpa4.py index af3b876676..d25ec00f1c 100644 --- a/source/tests/jax/test_dpa4.py +++ b/source/tests/jax/test_dpa4.py @@ -5,6 +5,12 @@ DescrptDPA4, _iter_object_tree, ) +from deepmd.jax.env import ( + nnx, +) +from deepmd.jax.fitting.dpa4_ener import ( + SeZMEnergyFittingNet, +) from deepmd.jax.utils.network import ( ArrayAPIParam, ) @@ -56,3 +62,41 @@ def test_optional_dpa4_weights_are_jax_parameters() -> None: for block in interaction_blocks for scale in block.adam_ffn_layer_scales ) + + +def test_frozen_descriptor_has_no_optimizer_visible_parameters() -> None: + """The root freeze flag must demote every descendant parameter.""" + descriptor = DescrptDPA4( + ntypes=2, + sel=4, + rcut=4.0, + channels=4, + n_radial=4, + lmax=1, + mmax=1, + n_blocks=1, + grid_branch=0, + random_gamma=False, + precision="float64", + trainable=False, + seed=20260712, + ) + + assert len(nnx.to_flat_state(nnx.state(descriptor, nnx.Param))) == 0 + + +def test_frozen_fitting_stays_frozen_after_conversion_round_trip() -> None: + """Serialized GLU layers retain the all-false optimizer policy.""" + fitting = SeZMEnergyFittingNet( + ntypes=2, + dim_descrpt=4, + neuron=[4], + trainable=False, + precision="float64", + mixed_types=True, + seed=20260712, + ) + restored = SeZMEnergyFittingNet.deserialize(fitting.serialize()) + + assert restored.nets[0].trainable == [False, False] + assert len(nnx.to_flat_state(nnx.state(restored, nnx.Param))) == 0 diff --git a/source/tests/jax/test_dpa4_conversion.py b/source/tests/jax/test_dpa4_conversion.py new file mode 100644 index 0000000000..9bed095827 --- /dev/null +++ b/source/tests/jax/test_dpa4_conversion.py @@ -0,0 +1,76 @@ +# SPDX-License-Identifier: LGPL-3.0-or-later +"""End-to-end PT checkpoint conversion coverage for JAX DPA4.""" + +from copy import ( + deepcopy, +) + +import torch + +from deepmd.jax.utils.serialization import ( + deserialize_to_file as deserialize_to_jax_file, +) +from deepmd.jax.utils.serialization import ( + serialize_from_file as serialize_from_jax_file, +) +from deepmd.pt.model.model import ( + get_model as get_pt_model, +) +from deepmd.pt.train.wrapper import ( + ModelWrapper, +) +from deepmd.pt.utils.serialization import ( + serialize_from_file as serialize_from_pt_file, +) +from deepmd.utils.argcheck import ( + model_args, +) + + +def _small_dpa4_config() -> dict: + """Return a small real PT DPA4 config with zero-size state leaves.""" + return model_args().normalize_value( + { + "type": "dpa4", + "type_map": ["O", "H"], + "descriptor": { + "type": "dpa4", + "sel": 4, + "rcut": 4.0, + "channels": 4, + "n_radial": 4, + "lmax": 1, + "mmax": 1, + "n_blocks": 1, + "random_gamma": False, + "precision": "float64", + "seed": 1, + }, + "fitting_net": { + "type": "dpa4_ener", + "neuron": [4], + "precision": "float64", + "seed": 1, + }, + }, + trim_pattern="_.*", + ) + + +def test_pt_dpa4_checkpoint_converts_to_real_jax_checkpoint(tmp_path) -> None: + """The public PT schema saves and restores through Orbax without loss.""" + model_params = _small_dpa4_config() + pt_model = get_pt_model(deepcopy(model_params)).to(torch.float64) + wrapper = ModelWrapper(pt_model, model_params=deepcopy(model_params)) + pt_path = tmp_path / "dpa4.pt" + jax_path = tmp_path / "dpa4.jax" + torch.save({"model": wrapper.state_dict()}, pt_path) + + data = serialize_from_pt_file(str(pt_path)) + assert data["model"]["type"] == "SeZM" + + deserialize_to_jax_file(str(jax_path), data) + restored = serialize_from_jax_file(str(jax_path)) + + assert restored["model"]["type"] == "standard" + assert restored["model"]["descriptor"]["type"] == "SeZM" diff --git a/source/tests/jax/test_model_factory.py b/source/tests/jax/test_model_factory.py index da09f80a46..12a137f9ba 100644 --- a/source/tests/jax/test_model_factory.py +++ b/source/tests/jax/test_model_factory.py @@ -20,6 +20,9 @@ from deepmd.jax.model.model import ( get_model, ) +from deepmd.utils.argcheck import ( + model_args, +) def _base_config() -> dict: @@ -102,6 +105,11 @@ def test_rejects_unsupported_features(self) -> None: with self.assertRaises(NotImplementedError): get_model(data) + data = _base_sezm_config() + data["descriptor"]["add_chg_spin_ebd"] = True + with self.assertRaises(NotImplementedError): + get_model(data) + def test_rejects_incompatible_descriptor_and_fitting_types(self) -> None: data = _base_sezm_config() data["descriptor"]["type"] = "se_e2_a" @@ -138,6 +146,21 @@ def test_descriptor_exclude_types_feed_standard_model( self.assertEqual(normalized["pair_exclude_types"], [[0, 1]]) self.assertEqual(normalized["descriptor"]["exclude_types"], [[0, 1]]) + @patch("deepmd.jax.model.model.get_standard_model", side_effect=lambda data: data) + def test_normalized_descriptor_exclusions_override_empty_default( + self, + _get_standard_model, + ) -> None: + """Argcheck's empty model-level default is not an explicit mismatch.""" + data = _base_sezm_config() + data["descriptor"]["exclude_types"] = [[0, 1]] + data = model_args().normalize_value(data, trim_pattern="_.*") + + normalized = get_model(data) + + self.assertEqual(normalized["pair_exclude_types"], [[0, 1]]) + self.assertEqual(normalized["descriptor"]["exclude_types"], [[0, 1]]) + if __name__ == "__main__": unittest.main() From 180fff068a68f7514ccc1f6bb067ee112ddf9c2c Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Sun, 12 Jul 2026 12:27:58 +0000 Subject: [PATCH 06/11] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- deepmd/jax/utils/serialization.py | 1 - source/tests/jax/test_dpa4_conversion.py | 8 ++------ 2 files changed, 2 insertions(+), 7 deletions(-) diff --git a/deepmd/jax/utils/serialization.py b/deepmd/jax/utils/serialization.py index 0e171ee7eb..fdbd50a3ce 100644 --- a/deepmd/jax/utils/serialization.py +++ b/deepmd/jax/utils/serialization.py @@ -24,7 +24,6 @@ get_model, ) - _DROP_ZERO_SIZE_LEAF = object() diff --git a/source/tests/jax/test_dpa4_conversion.py b/source/tests/jax/test_dpa4_conversion.py index 9bed095827..703677f1be 100644 --- a/source/tests/jax/test_dpa4_conversion.py +++ b/source/tests/jax/test_dpa4_conversion.py @@ -13,15 +13,11 @@ from deepmd.jax.utils.serialization import ( serialize_from_file as serialize_from_jax_file, ) -from deepmd.pt.model.model import ( - get_model as get_pt_model, -) +from deepmd.pt.model.model import get_model as get_pt_model from deepmd.pt.train.wrapper import ( ModelWrapper, ) -from deepmd.pt.utils.serialization import ( - serialize_from_file as serialize_from_pt_file, -) +from deepmd.pt.utils.serialization import serialize_from_file as serialize_from_pt_file from deepmd.utils.argcheck import ( model_args, ) From 72aa9efedfa3123e3604de82c03a8b5566883ceb Mon Sep 17 00:00:00 2001 From: njzjz-bot Date: Sun, 12 Jul 2026 22:35:23 +0800 Subject: [PATCH 07/11] fix(pt): preserve DPA4 layer trainability Coding-Agent: Codex Codex-Version: codex-cli 0.144.1 Model: gpt-5.6-sol Reasoning-Effort: xhigh --- deepmd/pt/model/task/sezm_ener.py | 22 ++++++++++++++-------- 1 file changed, 14 insertions(+), 8 deletions(-) diff --git a/deepmd/pt/model/task/sezm_ener.py b/deepmd/pt/model/task/sezm_ener.py index 0932ec7086..f83fd339f3 100644 --- a/deepmd/pt/model/task/sezm_ener.py +++ b/deepmd/pt/model/task/sezm_ener.py @@ -243,11 +243,18 @@ def __init__( super().__init__() if neuron is None: neuron = [] - if isinstance(trainable, list): - trainable = all(trainable) self.in_dim = int(in_dim) self.out_dim = int(out_dim) self.neuron = [int(nn_dim) for nn_dim in neuron] + if isinstance(trainable, bool): + self.trainable = [trainable] * (len(self.neuron) + 1) + else: + self.trainable = [bool(flag) for flag in trainable] + if len(self.trainable) != len(self.neuron) + 1: + raise ValueError( + "trainable must contain one flag per hidden layer plus " + "one flag for the output layer" + ) self.activation_function = activation_function self.resnet_dt = bool(resnet_dt) self.precision = precision @@ -270,7 +277,7 @@ def __init__( activation_function=self.activation_function, precision=self.precision, seed=child_seed(seed, layer_idx), - trainable=trainable, + trainable=self.trainable[layer_idx], ) ) dim_in = hidden_dim @@ -285,7 +292,7 @@ def __init__( activation_function=self.activation_function, precision=self.precision, seed=child_seed(seed, len(self.neuron)), - trainable=trainable, + trainable=all(self.trainable), ) else: self.case_film = None @@ -300,12 +307,9 @@ def __init__( resnet=False, precision=self.precision, seed=child_seed(seed, len(self.neuron) + int(self.case_film_embd)), - trainable=trainable, + trainable=self.trainable[-1], ) - for param in self.parameters(): - param.requires_grad = trainable - def _apply_input_film( self, xx: torch.Tensor, @@ -395,6 +399,8 @@ def serialize(self) -> dict[str, Any]: "descriptor_dim": self.descriptor_dim, "dim_case_embd": self.dim_case_embd, "case_film_embd": self.case_film_embd, + # Keep the per-layer freeze policy stable across backend round trips. + "trainable": self.trainable.copy(), "@variables": {key: to_numpy_array(value) for key, value in state.items()}, } From 7c73a56e446cc6a200623c7d86d02a378d403369 Mon Sep 17 00:00:00 2001 From: njzjz-bot Date: Thu, 16 Jul 2026 18:41:21 +0800 Subject: [PATCH 08/11] fix(jax): reject unsupported DPA4 training modes Preserve PT DPA4 fitting freeze policies across serialization, reject inert JAX random-gamma and AMP options, register model aliases for neighbor-stat preprocessing, and keep JAX-only tests collectable without PyTorch. Coding-Agent: Codex Codex-Version: codex-cli 0.144.4 Model: gpt-5.6-sol Reasoning-Effort: xhigh --- deepmd/jax/descriptor/dpa4.py | 13 +++++++++ deepmd/jax/model/ener_model.py | 4 +++ deepmd/jax/model/model.py | 10 +++++++ deepmd/pt/model/task/sezm_ener.py | 15 ++++++++++ source/tests/jax/test_dpa4_conversion.py | 5 +++- source/tests/jax/test_model_factory.py | 37 +++++++++++++++++++++++- source/tests/pt/model/test_sezm_model.py | 26 +++++++++++++++++ 7 files changed, 108 insertions(+), 2 deletions(-) diff --git a/deepmd/jax/descriptor/dpa4.py b/deepmd/jax/descriptor/dpa4.py index cda6507599..3db2141498 100644 --- a/deepmd/jax/descriptor/dpa4.py +++ b/deepmd/jax/descriptor/dpa4.py @@ -287,6 +287,19 @@ def deserialize(cls, data: dict) -> "SO2Linear": @flax_module class DescrptDPA4(DescrptDPA4DP): def __init__(self, *args: Any, **kwargs: Any) -> None: + # The dpmodel implementation currently fixes gamma and performs no + # automatic mixed-precision policy under JAX. Reject explicit requests + # instead of silently accepting training options that have no effect. + if kwargs.get("random_gamma") is True: + raise NotImplementedError( + "DPA4 random_gamma=True is not supported in the JAX backend." + ) + if kwargs.get("use_amp") is True: + raise NotImplementedError( + "DPA4 use_amp=True is not supported in the JAX backend." + ) + kwargs.setdefault("random_gamma", False) + kwargs.setdefault("use_amp", False) super().__init__(*args, **kwargs) _promote_trainable_tree(self) diff --git a/deepmd/jax/model/ener_model.py b/deepmd/jax/model/ener_model.py index 626997d18e..17584b0f5b 100644 --- a/deepmd/jax/model/ener_model.py +++ b/deepmd/jax/model/ener_model.py @@ -13,6 +13,10 @@ @BaseModel.register("sezm_ener") @BaseModel.register("dpa4_ener") +@BaseModel.register("SeZM") +@BaseModel.register("sezm") +@BaseModel.register("DPA4") +@BaseModel.register("dpa4") @BaseModel.register("ener") class EnergyModel(make_jax_dp_model_from_dpmodel(EnergyModelDP, DPAtomicModelEnergy)): pass diff --git a/deepmd/jax/model/model.py b/deepmd/jax/model/model.py index 31ddc12073..712edcb50f 100644 --- a/deepmd/jax/model/model.py +++ b/deepmd/jax/model/model.py @@ -128,6 +128,16 @@ def get_sezm_model(data: dict) -> BaseModel: data["fitting_net"] = data.get("fitting_net") or {} data["descriptor"].setdefault("type", "dpa4") data["fitting_net"].setdefault("type", "dpa4_ener") + if data["descriptor"].get("random_gamma") is True: + raise NotImplementedError( + "DPA4 random_gamma=True is not supported in the JAX backend." + ) + if data["descriptor"].get("use_amp") is True: + raise NotImplementedError( + "DPA4 use_amp=True is not supported in the JAX backend." + ) + data["descriptor"].setdefault("random_gamma", False) + data["descriptor"].setdefault("use_amp", False) if data["descriptor"].get("add_chg_spin_ebd"): raise NotImplementedError( "DPA4/SeZM charge/spin conditioning is not supported in JAX." diff --git a/deepmd/pt/model/task/sezm_ener.py b/deepmd/pt/model/task/sezm_ener.py index f83fd339f3..b24e5c485a 100644 --- a/deepmd/pt/model/task/sezm_ener.py +++ b/deepmd/pt/model/task/sezm_ener.py @@ -310,6 +310,21 @@ def __init__( trainable=self.trainable[-1], ) + # The layer constructors retain ``trainable`` as serialization + # metadata but do not consistently apply it to newly created + # ``Parameter`` objects. Reapply the policy at this owning module so a + # deserialize round trip cannot make frozen layers optimizer-visible. + for layer, layer_trainable in zip( + self.hidden_layers, self.trainable[:-1], strict=True + ): + for param in layer.parameters(): + param.requires_grad = layer_trainable + if self.case_film is not None: + for param in self.case_film.parameters(): + param.requires_grad = all(self.trainable) + for param in self.output_layer.parameters(): + param.requires_grad = self.trainable[-1] + def _apply_input_film( self, xx: torch.Tensor, diff --git a/source/tests/jax/test_dpa4_conversion.py b/source/tests/jax/test_dpa4_conversion.py index 703677f1be..18973bc681 100644 --- a/source/tests/jax/test_dpa4_conversion.py +++ b/source/tests/jax/test_dpa4_conversion.py @@ -5,7 +5,9 @@ deepcopy, ) -import torch +import pytest + +torch = pytest.importorskip("torch") from deepmd.jax.utils.serialization import ( deserialize_to_file as deserialize_to_jax_file, @@ -39,6 +41,7 @@ def _small_dpa4_config() -> dict: "mmax": 1, "n_blocks": 1, "random_gamma": False, + "use_amp": False, "precision": "float64", "seed": 1, }, diff --git a/source/tests/jax/test_model_factory.py b/source/tests/jax/test_model_factory.py index 12a137f9ba..68db1ba848 100644 --- a/source/tests/jax/test_model_factory.py +++ b/source/tests/jax/test_model_factory.py @@ -14,6 +14,9 @@ patch, ) +from deepmd.jax.model.base_model import ( + BaseModel, +) from deepmd.jax.model.ener_model import ( EnergyModel, ) @@ -51,7 +54,11 @@ def _base_sezm_config() -> dict: return { "type": "dpa4", "type_map": ["O", "H"], - "descriptor": {"type": "dpa4"}, + "descriptor": { + "type": "dpa4", + "random_gamma": False, + "use_amp": False, + }, "fitting_net": {"type": "dpa4_ener"}, } @@ -110,6 +117,13 @@ def test_rejects_unsupported_features(self) -> None: with self.assertRaises(NotImplementedError): get_model(data) + for key in ("random_gamma", "use_amp"): + with self.subTest(descriptor_option=key): + data = _base_sezm_config() + data["descriptor"][key] = True + with self.assertRaisesRegex(NotImplementedError, key): + get_model(data) + def test_rejects_incompatible_descriptor_and_fitting_types(self) -> None: data = _base_sezm_config() data["descriptor"]["type"] = "se_e2_a" @@ -129,6 +143,27 @@ def test_rejects_mismatched_exclude_types(self) -> None: with self.assertRaises(ValueError): get_model(data) + @patch( + "deepmd.dpmodel.model.dp_model.BaseDescriptor.update_sel", + return_value=({"type": "dpa4", "sel": 16}, 0.75), + ) + def test_model_aliases_route_through_update_sel(self, update_sel) -> None: + """Neighbor-stat preprocessing recognizes every public DPA4 alias.""" + for model_type in ("dpa4", "DPA4", "sezm", "SeZM"): + with self.subTest(model_type=model_type): + local_jdata = { + "type": model_type, + "descriptor": {"type": "dpa4", "sel": "auto"}, + } + + updated, min_nbor_dist = BaseModel.update_sel( + object(), ["O", "H"], local_jdata + ) + + self.assertEqual(updated["descriptor"]["sel"], 16) + self.assertEqual(min_nbor_dist, 0.75) + self.assertEqual(update_sel.call_count, 4) + @patch("deepmd.jax.model.model.get_standard_model", side_effect=lambda data: data) def test_descriptor_exclude_types_feed_standard_model( self, diff --git a/source/tests/pt/model/test_sezm_model.py b/source/tests/pt/model/test_sezm_model.py index 494338a2da..192008cb8a 100644 --- a/source/tests/pt/model/test_sezm_model.py +++ b/source/tests/pt/model/test_sezm_model.py @@ -47,6 +47,9 @@ from deepmd.pt.model.model.sezm_property_model import ( SeZMPropertyModel, ) +from deepmd.pt.model.task.sezm_ener import ( + SeZMEnergyFittingNet, +) from deepmd.pt.train.training import ( prepare_model_for_loss, ) @@ -87,6 +90,29 @@ ) +class TestSeZMEnergyFittingTrainability(unittest.TestCase): + """Ensure frozen PT DPA4 fitting parameters survive serialization.""" + + def test_frozen_parameters_survive_round_trip(self) -> None: + fitting = SeZMEnergyFittingNet( + ntypes=2, + dim_descrpt=8, + neuron=[8], + mixed_types=True, + trainable=False, + dim_case_embd=2, + case_film_embd=True, + precision="float64", + seed=20260716, + ) + + restored = SeZMEnergyFittingNet.deserialize(fitting.serialize()) + + self.assertTrue(list(fitting.parameters())) + self.assertTrue(list(restored.parameters())) + self.assertTrue(all(not param.requires_grad for param in restored.parameters())) + + def _assert_close_with_strict_warning( actual: torch.Tensor, expected: torch.Tensor, From e5c7b1483e2dc226cdb350253072635eeddd029a Mon Sep 17 00:00:00 2001 From: njzjz-bot Date: Thu, 16 Jul 2026 18:50:05 +0800 Subject: [PATCH 09/11] fix(jax): retain full DPA4 neighbor capacity Use the full extended-atom capacity at the JAX input boundary for DPA4 so the configured sel cannot truncate in-cutoff neighbors, and cover the edge-count regression. Coding-Agent: Codex Codex-Version: codex-cli 0.144.4 Model: gpt-5.6-sol Reasoning-Effort: xhigh --- deepmd/jax/train/trainer.py | 8 ++++++++ deepmd/jax/train/validation.py | 3 +++ source/tests/jax/test_model_factory.py | 23 +++++++++++++++++++++++ 3 files changed, 34 insertions(+) diff --git a/deepmd/jax/train/trainer.py b/deepmd/jax/train/trainer.py index 450e2795e1..436ab2bb31 100644 --- a/deepmd/jax/train/trainer.py +++ b/deepmd/jax/train/trainer.py @@ -869,6 +869,7 @@ def _prepare_batch( fparam=jax_data.get("fparam", None), aparam=jax_data.get("aparam", None), pair_excl=getattr(model.atomic_model, "pair_excl", None), + conservative_nlist=type(model.get_descriptor()).__name__ == "DescrptDPA4", ) return jax_data, extended_coord, extended_atype, nlist, mapping, fp, ap @@ -1362,6 +1363,7 @@ def prepare_input( fparam: np.ndarray | None = None, aparam: np.ndarray | None = None, pair_excl: "PairExcludeMask | None" = None, + conservative_nlist: bool = False, ) -> tuple[ np.ndarray, np.ndarray, @@ -1384,6 +1386,12 @@ def prepare_input( extended_coord, extended_atype, mapping = extend_coord_with_ghosts( coord_normalized, atype, bb, rcut ) + if conservative_nlist: + # DPA4 treats ``sel`` as an initial capacity rather than a truncation + # contract. Use the full extended-atom capacity so every in-cutoff + # edge survives the dense JAX input boundary; padded entries remain + # masked by the lower descriptor path. + sel = [extended_coord.shape[1]] * len(sel) nlist = build_neighbor_list( extended_coord, extended_atype, diff --git a/deepmd/jax/train/validation.py b/deepmd/jax/train/validation.py index 9796b19a69..fcd0abce83 100644 --- a/deepmd/jax/train/validation.py +++ b/deepmd/jax/train/validation.py @@ -192,6 +192,9 @@ def predict_batch( box=box_input, fparam=fparam_input, aparam=aparam_input, + conservative_nlist=( + type(self.model.get_descriptor()).__name__ == "DescrptDPA4" + ), ) batch_output = _evaluate_model_dict( self.model, diff --git a/source/tests/jax/test_model_factory.py b/source/tests/jax/test_model_factory.py index 68db1ba848..8ab2132121 100644 --- a/source/tests/jax/test_model_factory.py +++ b/source/tests/jax/test_model_factory.py @@ -23,6 +23,9 @@ from deepmd.jax.model.model import ( get_model, ) +from deepmd.jax.train.trainer import ( + prepare_input, +) from deepmd.utils.argcheck import ( model_args, ) @@ -164,6 +167,26 @@ def test_model_aliases_route_through_update_sel(self, update_sel) -> None: self.assertEqual(min_nbor_dist, 0.75) self.assertEqual(update_sel.call_count, 4) + def test_dpa4_conservative_input_keeps_all_in_cutoff_neighbors(self) -> None: + """The JAX DPA4 input boundary must not truncate to the configured sel.""" + import numpy as np + + coord = np.asarray( + [[[0.0, 0.0, 0.0], [0.5, 0.0, 0.0], [0.0, 0.5, 0.0], [0.0, 0.0, 0.5]]] + ) + atype = np.zeros((1, 4), dtype=np.int32) + + _, _, nlist, _, _, _ = prepare_input( + rcut=2.0, + sel=[1], + coord=coord, + atype=atype, + conservative_nlist=True, + ) + + self.assertGreaterEqual(nlist.shape[-1], 4) + self.assertTrue(np.all(np.sum(nlist >= 0, axis=-1) >= 3)) + @patch("deepmd.jax.model.model.get_standard_model", side_effect=lambda data: data) def test_descriptor_exclude_types_feed_standard_model( self, From ef05aa9d77a92fb3c0dd96394b7df7f7920f7d00 Mon Sep 17 00:00:00 2001 From: njzjz-bot Date: Sat, 18 Jul 2026 14:47:16 +0800 Subject: [PATCH 10/11] test(jax): use supported DPA4 consistency options Keep cross-backend DPA4 consistency coverage within the JAX-supported feature subset by disabling the backend-specific AMP policy. Coding-Agent: Codex Codex-Version: codex-cli 0.144.4 Model: gpt-5.6-sol Reasoning-Effort: xhigh --- source/tests/consistent/descriptor/test_dpa4.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/source/tests/consistent/descriptor/test_dpa4.py b/source/tests/consistent/descriptor/test_dpa4.py index 58f7edadd3..d059d5ffdb 100644 --- a/source/tests/consistent/descriptor/test_dpa4.py +++ b/source/tests/consistent/descriptor/test_dpa4.py @@ -149,6 +149,10 @@ def data(self) -> dict: "grid_mlp": grid_mlp, "so3_readout": so3_readout, "random_gamma": False, + # JAX currently supports DPA4 without the backend-specific AMP + # policy. Keep cross-backend consistency cases within that shared + # feature subset; AMP behavior has dedicated backend tests. + "use_amp": False, "precision": precision, "trainable": False, "seed": 20251208, From 28d89e26aa372097ac01c718c0f38e8ad7a0201e Mon Sep 17 00:00:00 2001 From: njzjz-bot Date: Mon, 20 Jul 2026 23:40:23 +0800 Subject: [PATCH 11/11] refactor(jax): limit DPA4 support to descriptor Remove fitting, model, trainer, checkpoint, and model-conversion changes so the pull request contains only JAX DPA4 descriptor support and its focused tests. Coding-Agent: Codex Codex-Version: codex-cli 0.144.4 Model: gpt-5.6-sol Reasoning-Effort: xhigh --- deepmd/dpmodel/fitting/dpa4_ener.py | 18 +- deepmd/jax/fitting/__init__.py | 4 - deepmd/jax/fitting/dpa4_ener.py | 51 ------ deepmd/jax/model/base_model.py | 65 +------ deepmd/jax/model/ener_model.py | 6 - deepmd/jax/model/model.py | 63 ------- deepmd/jax/train/trainer.py | 35 +--- deepmd/jax/train/validation.py | 3 - deepmd/jax/utils/serialization.py | 55 ------ deepmd/pt/model/task/sezm_ener.py | 35 +--- .../consistent/fitting/test_dpa4_ener.py | 20 +-- source/tests/jax/test_dpa4.py | 22 +-- source/tests/jax/test_dpa4_conversion.py | 75 --------- source/tests/jax/test_model_factory.py | 158 ------------------ source/tests/jax/test_training.py | 97 ----------- source/tests/pt/model/test_sezm_model.py | 26 --- 16 files changed, 16 insertions(+), 717 deletions(-) delete mode 100644 deepmd/jax/fitting/dpa4_ener.py delete mode 100644 source/tests/jax/test_dpa4_conversion.py diff --git a/deepmd/dpmodel/fitting/dpa4_ener.py b/deepmd/dpmodel/fitting/dpa4_ener.py index 6f7f970947..f08097f032 100644 --- a/deepmd/dpmodel/fitting/dpa4_ener.py +++ b/deepmd/dpmodel/fitting/dpa4_ener.py @@ -92,18 +92,11 @@ def __init__( ) if neuron is None: neuron = [] + if isinstance(trainable, list): + trainable = all(trainable) self.in_dim = int(in_dim) self.out_dim = int(out_dim) self.neuron = [int(nn_dim) for nn_dim in neuron] - if isinstance(trainable, bool): - self.trainable = [trainable] * (len(self.neuron) + 1) - else: - self.trainable = [bool(flag) for flag in trainable] - if len(self.trainable) != len(self.neuron) + 1: - raise ValueError( - "trainable must contain one flag per hidden layer plus " - "one flag for the output layer" - ) self.activation_function = activation_function self.resnet_dt = bool(resnet_dt) self.precision = precision @@ -130,7 +123,7 @@ def __init__( resnet=False, precision=self.precision, seed=child_seed(seed, layer_idx), - trainable=self.trainable[layer_idx], + trainable=trainable, ) ) dim_in = hidden_dim @@ -146,7 +139,7 @@ def __init__( resnet=False, precision=self.precision, seed=child_seed(seed, len(self.neuron) + int(self.case_film_embd)), - trainable=self.trainable[-1], + trainable=trainable, ) def call_until_last(self, xx: Array) -> Array: @@ -188,9 +181,6 @@ def serialize(self) -> dict[str, Any]: "descriptor_dim": self.descriptor_dim, "dim_case_embd": self.dim_case_embd, "case_film_embd": self.case_film_embd, - # Preserve the effective per-layer freeze policy when backend - # wrappers rebuild this network from its serialized form. - "trainable": self.trainable.copy(), "@variables": variables, } diff --git a/deepmd/jax/fitting/__init__.py b/deepmd/jax/fitting/__init__.py index 82cea44e1b..77133e2bac 100644 --- a/deepmd/jax/fitting/__init__.py +++ b/deepmd/jax/fitting/__init__.py @@ -1,7 +1,4 @@ # SPDX-License-Identifier: LGPL-3.0-or-later -from deepmd.jax.fitting.dpa4_ener import ( - SeZMEnergyFittingNet, -) from deepmd.jax.fitting.fitting import ( DipoleFittingNet, DOSFittingNet, @@ -14,5 +11,4 @@ "DipoleFittingNet", "EnergyFittingNet", "PolarFittingNet", - "SeZMEnergyFittingNet", ] diff --git a/deepmd/jax/fitting/dpa4_ener.py b/deepmd/jax/fitting/dpa4_ener.py deleted file mode 100644 index 425268cd78..0000000000 --- a/deepmd/jax/fitting/dpa4_ener.py +++ /dev/null @@ -1,51 +0,0 @@ -# SPDX-License-Identifier: LGPL-3.0-or-later -from typing import ( - ClassVar, -) - -from deepmd.dpmodel.fitting.dpa4_ener import GLUFittingNet as GLUFittingNetDP -from deepmd.dpmodel.fitting.dpa4_ener import ( - SeZMEnergyFittingNet as SeZMEnergyFittingNetDP, -) -from deepmd.dpmodel.fitting.dpa4_ener import ( - SeZMNetworkCollection as SeZMNetworkCollectionDP, -) -from deepmd.jax.common import ( - flax_module, - register_dpmodel_mapping, -) -from deepmd.jax.fitting.base_fitting import ( - BaseFitting, -) - - -@flax_module -class GLUFittingNet(GLUFittingNetDP): - pass - - -register_dpmodel_mapping( - GLUFittingNetDP, - lambda v: GLUFittingNet.deserialize(v.serialize()), -) - - -@flax_module -class SeZMNetworkCollection(SeZMNetworkCollectionDP): - _jax_data_list_attrs: ClassVar[set[str]] = {"_networks", "networks"} - NETWORK_TYPE_MAP: ClassVar[dict[str, type]] = { - "sezm_fitting_network": GLUFittingNet, - } - - -register_dpmodel_mapping( - SeZMNetworkCollectionDP, - lambda v: SeZMNetworkCollection.deserialize(v.serialize()), -) - - -@BaseFitting.register("dpa4_ener") -@BaseFitting.register("sezm_ener") -@flax_module -class SeZMEnergyFittingNet(SeZMEnergyFittingNetDP): - pass diff --git a/deepmd/jax/model/base_model.py b/deepmd/jax/model/base_model.py index 1a341da902..8a5a55e8d7 100644 --- a/deepmd/jax/model/base_model.py +++ b/deepmd/jax/model/base_model.py @@ -1,9 +1,5 @@ # SPDX-License-Identifier: LGPL-3.0-or-later -from typing import ( - Any, -) - from deepmd.dpmodel.model.base_model import ( make_base_model, ) @@ -16,67 +12,8 @@ jax, jnp, ) -from deepmd.utils.version import ( - check_version_compatibility, -) - - -class BaseModel(make_base_model()): - """JAX model registry with adapters for regular PT SeZM checkpoints.""" - - _SEZM_MODEL_TYPES = frozenset({"sezm", "dpa4"}) - _SEZM_ATOMIC_TYPES = frozenset({"sezm_atomic"}) - @classmethod - def deserialize(cls, data: dict[str, Any]) -> "BaseModel": - model_type = str(data.get("type", "standard")).lower() - if model_type in cls._SEZM_MODEL_TYPES: - return cls.deserialize(cls._unwrap_pt_sezm_model(data)) - if model_type in cls._SEZM_ATOMIC_TYPES: - return cls.deserialize(cls._normalize_pt_sezm_atomic(data)) - return super().deserialize(data) - - @staticmethod - def _unwrap_pt_sezm_model(data: dict[str, Any]) -> dict[str, Any]: - """Unwrap PT's model-level SeZM schema after validating its extras.""" - check_version_compatibility(int(data.get("@version", 1)), 1, 1) - if str(data.get("bridging_method", "none")).lower() not in ("none", ""): - raise NotImplementedError( - "PT SeZM/DPA4 checkpoints with bridging are not supported in JAX." - ) - if data.get("lora") is not None: - raise NotImplementedError( - "PT SeZM/DPA4 checkpoints with LoRA are not supported in JAX." - ) - atomic_model = data.get("atomic_model") - if atomic_model is None: - raise ValueError("SeZM/DPA4 model data is missing 'atomic_model'.") - return atomic_model - - @staticmethod - def _normalize_pt_sezm_atomic(data: dict[str, Any]) -> dict[str, Any]: - """Convert PT's energy-only ``sezm_atomic`` schema to ``standard``.""" - data = data.copy() - check_version_compatibility(int(data.get("@version", 2)), 3, 2) - if data.pop("dens_fitting", None) is not None: - raise NotImplementedError( - "PT SeZM/DPA4 checkpoints with a dens head are not supported in JAX." - ) - active_mode = data.pop("active_mode", None) - if active_mode not in (None, "ener"): - raise NotImplementedError( - f"PT SeZM/DPA4 active_mode {active_mode!r} is not supported in JAX." - ) - variables = data.get("@variables") - if isinstance(variables, dict): - data["@variables"] = { - key: value - for key, value in variables.items() - if key in ("out_bias", "out_std") - } - data["@version"] = 2 - data["type"] = "standard" - return data +BaseModel = make_base_model() def forward_common_atomic( diff --git a/deepmd/jax/model/ener_model.py b/deepmd/jax/model/ener_model.py index 17584b0f5b..1d3e8a1d80 100644 --- a/deepmd/jax/model/ener_model.py +++ b/deepmd/jax/model/ener_model.py @@ -11,12 +11,6 @@ ) -@BaseModel.register("sezm_ener") -@BaseModel.register("dpa4_ener") -@BaseModel.register("SeZM") -@BaseModel.register("sezm") -@BaseModel.register("DPA4") -@BaseModel.register("dpa4") @BaseModel.register("ener") class EnergyModel(make_jax_dp_model_from_dpmodel(EnergyModelDP, DPAtomicModelEnergy)): pass diff --git a/deepmd/jax/model/model.py b/deepmd/jax/model/model.py index 712edcb50f..a3d067c636 100644 --- a/deepmd/jax/model/model.py +++ b/deepmd/jax/model/model.py @@ -109,67 +109,6 @@ def get_zbl_model(data: dict) -> DPZBLModel: ) -def get_sezm_model(data: dict) -> BaseModel: - """Build a DPA4/SeZM energy model from the pt-style model config.""" - data = deepcopy(data) - if "spin" in data: - raise NotImplementedError("Spin DPA4/SeZM models are not supported in JAX.") - if str(data.get("bridging_method", "none")).lower() != "none": - raise NotImplementedError("DPA4/SeZM bridging is not supported in JAX.") - if data.get("lora") is not None: - raise NotImplementedError("DPA4/SeZM LoRA is not supported in JAX.") - if data.get("use_compile"): - raise NotImplementedError("model.use_compile is not supported in JAX.") - if data.get("preset_out_bias"): - raise NotImplementedError("DPA4/SeZM preset_out_bias is not supported in JAX.") - - data.pop("type", None) - data["descriptor"] = data.get("descriptor") or {} - data["fitting_net"] = data.get("fitting_net") or {} - data["descriptor"].setdefault("type", "dpa4") - data["fitting_net"].setdefault("type", "dpa4_ener") - if data["descriptor"].get("random_gamma") is True: - raise NotImplementedError( - "DPA4 random_gamma=True is not supported in the JAX backend." - ) - if data["descriptor"].get("use_amp") is True: - raise NotImplementedError( - "DPA4 use_amp=True is not supported in the JAX backend." - ) - data["descriptor"].setdefault("random_gamma", False) - data["descriptor"].setdefault("use_amp", False) - if data["descriptor"].get("add_chg_spin_ebd"): - raise NotImplementedError( - "DPA4/SeZM charge/spin conditioning is not supported in JAX." - ) - if data["descriptor"]["type"] not in ("dpa4", "DPA4", "sezm", "SeZM"): - raise ValueError( - "Model type 'dpa4' requires a DPA4/SeZM descriptor, but got " - f"descriptor type '{data['descriptor']['type']}'." - ) - if data["fitting_net"]["type"] not in ("dpa4_ener", "sezm_ener"): - raise ValueError( - "Model type 'dpa4' requires the DPA4/SeZM energy fitting net, but got " - f"fitting_net type '{data['fitting_net']['type']}'." - ) - - descriptor_exclude_types = [ - list(pair) for pair in (data["descriptor"].get("exclude_types") or []) - ] - pair_exclude_types = [list(pair) for pair in (data.get("pair_exclude_types") or [])] - if pair_exclude_types: - if descriptor_exclude_types and descriptor_exclude_types != pair_exclude_types: - raise ValueError( - "DPA4/SeZM pair_exclude_types and descriptor.exclude_types must " - "match when both are provided." - ) - else: - pair_exclude_types = descriptor_exclude_types - data["pair_exclude_types"] = pair_exclude_types - data["descriptor"]["exclude_types"] = deepcopy(pair_exclude_types) - return get_standard_model(data) - - def get_model(data: dict) -> BaseModel: """Get a model from a dictionary. @@ -186,7 +125,5 @@ def get_model(data: dict) -> BaseModel: return get_zbl_model(data) else: return get_standard_model(data) - elif model_type in ("SeZM", "sezm", "DPA4", "dpa4"): - return get_sezm_model(data) else: return BaseModel.get_class_by_type(model_type).get_model(data) diff --git a/deepmd/jax/train/trainer.py b/deepmd/jax/train/trainer.py index 436ab2bb31..c19267250e 100644 --- a/deepmd/jax/train/trainer.py +++ b/deepmd/jax/train/trainer.py @@ -86,7 +86,6 @@ preprocess_shared_params, ) from deepmd.jax.utils.serialization import ( - _drop_zero_size_array_leaves, serialize_from_file, ) from deepmd.utils.argcheck import ( @@ -658,7 +657,6 @@ def loss_fn( model_dict = _evaluate_model_dict( model, extended_coord, extended_atype, nlist, mapping, fp, ap ) - model_dict = _match_label_shapes(model_dict, label_dict) loss, _ = loss_obj( learning_rate=lr, natoms=label_dict["type"].shape[1], @@ -682,7 +680,6 @@ def loss_fn_more_loss( model_dict = _evaluate_model_dict( model, extended_coord, extended_atype, nlist, mapping, fp, ap ) - model_dict = _match_label_shapes(model_dict, label_dict) _, more_loss = loss_obj( learning_rate=lr, natoms=label_dict["type"].shape[1], @@ -869,7 +866,6 @@ def _prepare_batch( fparam=jax_data.get("fparam", None), aparam=jax_data.get("aparam", None), pair_excl=getattr(model.atomic_model, "pair_excl", None), - conservative_nlist=type(model.get_descriptor()).__name__ == "DescrptDPA4", ) return jax_data, extended_coord, extended_atype, nlist, mapping, fp, ap @@ -934,7 +930,6 @@ def _write_checkpoint(self, ckpt_path: Path, *, step: int) -> None: else: _, single_state = nnx.split(self.models[DEFAULT_TASK_KEY]) state = single_state.to_pure_dict() - state = _drop_zero_size_array_leaves(state) if ckpt_path.is_dir(): shutil.rmtree(ckpt_path) model_def_script_cpy = deepcopy(self.model_def_script) @@ -992,32 +987,11 @@ def _evaluate_model_dict( ) model_dict["atom_energy"] = model_dict["energy"] model_dict["energy"] = model_dict["energy_redu"] - force = model_dict["energy_derv_r"].squeeze(-2) - if force.ndim == 2 or (force.ndim == 3 and force.shape[-1] != 3): - force = jnp.reshape(force, (force.shape[0], -1, 3)) - model_dict["force"] = force + model_dict["force"] = model_dict["energy_derv_r"].squeeze(-2) model_dict["virial"] = model_dict["energy_derv_c_redu"].squeeze(-2) return model_dict -def _match_label_shapes( - model_dict: dict[str, jnp.ndarray], - label_dict: dict[str, jnp.ndarray], -) -> dict[str, jnp.ndarray]: - """Match equivalent flattened model outputs to label tensor shapes.""" - force_hat = model_dict.get("force") - force = label_dict.get("force") - if ( - force_hat is not None - and force is not None - and force_hat.shape != force.shape - and force_hat.size == force.size - ): - model_dict = dict(model_dict) - model_dict["force"] = jnp.reshape(force_hat, force.shape) - return model_dict - - def share_jax_model_params( models: dict[str, BaseModel], shared_links: dict[str, Any], @@ -1363,7 +1337,6 @@ def prepare_input( fparam: np.ndarray | None = None, aparam: np.ndarray | None = None, pair_excl: "PairExcludeMask | None" = None, - conservative_nlist: bool = False, ) -> tuple[ np.ndarray, np.ndarray, @@ -1386,12 +1359,6 @@ def prepare_input( extended_coord, extended_atype, mapping = extend_coord_with_ghosts( coord_normalized, atype, bb, rcut ) - if conservative_nlist: - # DPA4 treats ``sel`` as an initial capacity rather than a truncation - # contract. Use the full extended-atom capacity so every in-cutoff - # edge survives the dense JAX input boundary; padded entries remain - # masked by the lower descriptor path. - sel = [extended_coord.shape[1]] * len(sel) nlist = build_neighbor_list( extended_coord, extended_atype, diff --git a/deepmd/jax/train/validation.py b/deepmd/jax/train/validation.py index fcd0abce83..9796b19a69 100644 --- a/deepmd/jax/train/validation.py +++ b/deepmd/jax/train/validation.py @@ -192,9 +192,6 @@ def predict_batch( box=box_input, fparam=fparam_input, aparam=aparam_input, - conservative_nlist=( - type(self.model.get_descriptor()).__name__ == "DescrptDPA4" - ), ) batch_output = _evaluate_model_dict( self.model, diff --git a/deepmd/jax/utils/serialization.py b/deepmd/jax/utils/serialization.py index fdbd50a3ce..39354cc1fe 100644 --- a/deepmd/jax/utils/serialization.py +++ b/deepmd/jax/utils/serialization.py @@ -24,22 +24,6 @@ get_model, ) -_DROP_ZERO_SIZE_LEAF = object() - - -def _drop_zero_size_array_leaves(value: Any) -> Any: - """Remove array leaves Orbax cannot save while preserving tree containers.""" - if isinstance(value, dict): - filtered = {} - for key, item in value.items(): - new_item = _drop_zero_size_array_leaves(item) - if new_item is not _DROP_ZERO_SIZE_LEAF: - filtered[key] = new_item - return filtered - if getattr(value, "size", None) == 0: - return _DROP_ZERO_SIZE_LEAF - return value - def _convert_str_to_int_key(item: dict) -> None: """Convert Orbax-restored numeric index keys from strings back to ints.""" @@ -66,40 +50,6 @@ def _normalize_restored_state_keys( _convert_str_to_int_key(state) -_NO_ZERO_SIZE_LEAF = object() - - -def _zero_size_subtree(value: Any) -> Any: - if isinstance(value, dict): - restored = {} - for key, item in value.items(): - subtree = _zero_size_subtree(item) - if subtree is not _NO_ZERO_SIZE_LEAF: - restored[key] = subtree - return restored if restored else _NO_ZERO_SIZE_LEAF - if getattr(value, "size", None) == 0: - return value - return _NO_ZERO_SIZE_LEAF - - -def _restore_missing_zero_size_leaves(template: Any, restored: Any) -> Any: - """Reinsert zero-size leaves dropped before Orbax checkpoint saving.""" - if not isinstance(template, dict) or not isinstance(restored, dict): - return restored - restored = dict(restored) - for key, template_value in template.items(): - if key in restored: - restored[key] = _restore_missing_zero_size_leaves( - template_value, - restored[key], - ) - continue - subtree = _zero_size_subtree(template_value) - if subtree is not _NO_ZERO_SIZE_LEAF: - restored[key] = subtree - return restored - - def _state_sequence_to_numpy_list(state_value: Any) -> list[np.ndarray]: """Convert an Orbax-restored list/dict sequence to NumPy arrays.""" if isinstance(state_value, dict): @@ -269,7 +219,6 @@ def deserialize_to_file(model_file: str, data: dict) -> None: model = BaseModel.deserialize(data["model"]) _, state = nnx.split(model) state = state.to_pure_dict() - state = _drop_zero_size_array_leaves(state) with ocp.Checkpointer( ocp.CompositeCheckpointHandler("state", "model_def_script") ) as checkpointer: @@ -433,10 +382,6 @@ def restore_model(model_params: dict, model_state: dict) -> BaseModel: abstract_model = get_model(model_params) _restore_compression_slots_from_state(abstract_model, model_state) graphdef, abstract_state = nnx.split(abstract_model) - model_state = _restore_missing_zero_size_leaves( - abstract_state.to_pure_dict(), - model_state, - ) abstract_state.replace_by_pure_dict(model_state) return nnx.merge(graphdef, abstract_state) diff --git a/deepmd/pt/model/task/sezm_ener.py b/deepmd/pt/model/task/sezm_ener.py index b24e5c485a..0932ec7086 100644 --- a/deepmd/pt/model/task/sezm_ener.py +++ b/deepmd/pt/model/task/sezm_ener.py @@ -243,18 +243,11 @@ def __init__( super().__init__() if neuron is None: neuron = [] + if isinstance(trainable, list): + trainable = all(trainable) self.in_dim = int(in_dim) self.out_dim = int(out_dim) self.neuron = [int(nn_dim) for nn_dim in neuron] - if isinstance(trainable, bool): - self.trainable = [trainable] * (len(self.neuron) + 1) - else: - self.trainable = [bool(flag) for flag in trainable] - if len(self.trainable) != len(self.neuron) + 1: - raise ValueError( - "trainable must contain one flag per hidden layer plus " - "one flag for the output layer" - ) self.activation_function = activation_function self.resnet_dt = bool(resnet_dt) self.precision = precision @@ -277,7 +270,7 @@ def __init__( activation_function=self.activation_function, precision=self.precision, seed=child_seed(seed, layer_idx), - trainable=self.trainable[layer_idx], + trainable=trainable, ) ) dim_in = hidden_dim @@ -292,7 +285,7 @@ def __init__( activation_function=self.activation_function, precision=self.precision, seed=child_seed(seed, len(self.neuron)), - trainable=all(self.trainable), + trainable=trainable, ) else: self.case_film = None @@ -307,23 +300,11 @@ def __init__( resnet=False, precision=self.precision, seed=child_seed(seed, len(self.neuron) + int(self.case_film_embd)), - trainable=self.trainable[-1], + trainable=trainable, ) - # The layer constructors retain ``trainable`` as serialization - # metadata but do not consistently apply it to newly created - # ``Parameter`` objects. Reapply the policy at this owning module so a - # deserialize round trip cannot make frozen layers optimizer-visible. - for layer, layer_trainable in zip( - self.hidden_layers, self.trainable[:-1], strict=True - ): - for param in layer.parameters(): - param.requires_grad = layer_trainable - if self.case_film is not None: - for param in self.case_film.parameters(): - param.requires_grad = all(self.trainable) - for param in self.output_layer.parameters(): - param.requires_grad = self.trainable[-1] + for param in self.parameters(): + param.requires_grad = trainable def _apply_input_film( self, @@ -414,8 +395,6 @@ def serialize(self) -> dict[str, Any]: "descriptor_dim": self.descriptor_dim, "dim_case_embd": self.dim_case_embd, "case_film_embd": self.case_film_embd, - # Keep the per-layer freeze policy stable across backend round trips. - "trainable": self.trainable.copy(), "@variables": {key: to_numpy_array(value) for key, value in state.items()}, } diff --git a/source/tests/consistent/fitting/test_dpa4_ener.py b/source/tests/consistent/fitting/test_dpa4_ener.py index 3982656870..3d007a9959 100644 --- a/source/tests/consistent/fitting/test_dpa4_ener.py +++ b/source/tests/consistent/fitting/test_dpa4_ener.py @@ -19,7 +19,6 @@ from ..common import ( INSTALLED_ARRAY_API_STRICT, - INSTALLED_JAX, INSTALLED_PT, INSTALLED_PT_EXPT, CommonTest, @@ -45,13 +44,6 @@ from deepmd.pt_expt.utils.env import DEVICE as PT_EXPT_DEVICE else: SeZMEnerFittingPTExpt = None -if INSTALLED_JAX: - from deepmd.jax.env import ( - jnp, - ) - from deepmd.jax.fitting.dpa4_ener import SeZMEnergyFittingNet as SeZMEnerFittingJAX -else: - SeZMEnerFittingJAX = None if INSTALLED_ARRAY_API_STRICT: import array_api_strict @@ -94,7 +86,7 @@ def skip_pt(self) -> bool: skip_dp = False skip_tf = True - skip_jax = not INSTALLED_JAX or SeZMEnerFittingJAX is None + skip_jax = True skip_pd = True skip_pt_expt = not INSTALLED_PT_EXPT skip_array_api_strict = not INSTALLED_ARRAY_API_STRICT @@ -103,7 +95,7 @@ def skip_pt(self) -> bool: dp_class = SeZMEnerFittingDP pt_class = SeZMEnerFittingPT pt_expt_class = SeZMEnerFittingPTExpt - jax_class = SeZMEnerFittingJAX + jax_class = None pd_class = None array_api_strict_class = SeZMEnerFittingStrict args = fitting_sezm_ener() @@ -158,14 +150,6 @@ def eval_pt_expt(self, pt_expt_obj: Any) -> Any: .numpy() ) - def eval_jax(self, jax_obj: Any) -> Any: - return np.asarray( - jax_obj( - jnp.asarray(self.inputs), - jnp.asarray(self.atype.reshape(1, -1)), - )["energy"] - ) - def eval_array_api_strict(self, array_api_strict_obj: Any) -> Any: return to_numpy_array( array_api_strict_obj( diff --git a/source/tests/jax/test_dpa4.py b/source/tests/jax/test_dpa4.py index d25ec00f1c..1733f6d83d 100644 --- a/source/tests/jax/test_dpa4.py +++ b/source/tests/jax/test_dpa4.py @@ -1,5 +1,5 @@ # SPDX-License-Identifier: LGPL-3.0-or-later -"""Focused tests for JAX DPA4 trainable-state conversion.""" +"""Focused tests for JAX DPA4 descriptor trainable-state conversion.""" from deepmd.jax.descriptor.dpa4 import ( DescrptDPA4, @@ -8,9 +8,6 @@ from deepmd.jax.env import ( nnx, ) -from deepmd.jax.fitting.dpa4_ener import ( - SeZMEnergyFittingNet, -) from deepmd.jax.utils.network import ( ArrayAPIParam, ) @@ -83,20 +80,3 @@ def test_frozen_descriptor_has_no_optimizer_visible_parameters() -> None: ) assert len(nnx.to_flat_state(nnx.state(descriptor, nnx.Param))) == 0 - - -def test_frozen_fitting_stays_frozen_after_conversion_round_trip() -> None: - """Serialized GLU layers retain the all-false optimizer policy.""" - fitting = SeZMEnergyFittingNet( - ntypes=2, - dim_descrpt=4, - neuron=[4], - trainable=False, - precision="float64", - mixed_types=True, - seed=20260712, - ) - restored = SeZMEnergyFittingNet.deserialize(fitting.serialize()) - - assert restored.nets[0].trainable == [False, False] - assert len(nnx.to_flat_state(nnx.state(restored, nnx.Param))) == 0 diff --git a/source/tests/jax/test_dpa4_conversion.py b/source/tests/jax/test_dpa4_conversion.py deleted file mode 100644 index 18973bc681..0000000000 --- a/source/tests/jax/test_dpa4_conversion.py +++ /dev/null @@ -1,75 +0,0 @@ -# SPDX-License-Identifier: LGPL-3.0-or-later -"""End-to-end PT checkpoint conversion coverage for JAX DPA4.""" - -from copy import ( - deepcopy, -) - -import pytest - -torch = pytest.importorskip("torch") - -from deepmd.jax.utils.serialization import ( - deserialize_to_file as deserialize_to_jax_file, -) -from deepmd.jax.utils.serialization import ( - serialize_from_file as serialize_from_jax_file, -) -from deepmd.pt.model.model import get_model as get_pt_model -from deepmd.pt.train.wrapper import ( - ModelWrapper, -) -from deepmd.pt.utils.serialization import serialize_from_file as serialize_from_pt_file -from deepmd.utils.argcheck import ( - model_args, -) - - -def _small_dpa4_config() -> dict: - """Return a small real PT DPA4 config with zero-size state leaves.""" - return model_args().normalize_value( - { - "type": "dpa4", - "type_map": ["O", "H"], - "descriptor": { - "type": "dpa4", - "sel": 4, - "rcut": 4.0, - "channels": 4, - "n_radial": 4, - "lmax": 1, - "mmax": 1, - "n_blocks": 1, - "random_gamma": False, - "use_amp": False, - "precision": "float64", - "seed": 1, - }, - "fitting_net": { - "type": "dpa4_ener", - "neuron": [4], - "precision": "float64", - "seed": 1, - }, - }, - trim_pattern="_.*", - ) - - -def test_pt_dpa4_checkpoint_converts_to_real_jax_checkpoint(tmp_path) -> None: - """The public PT schema saves and restores through Orbax without loss.""" - model_params = _small_dpa4_config() - pt_model = get_pt_model(deepcopy(model_params)).to(torch.float64) - wrapper = ModelWrapper(pt_model, model_params=deepcopy(model_params)) - pt_path = tmp_path / "dpa4.pt" - jax_path = tmp_path / "dpa4.jax" - torch.save({"model": wrapper.state_dict()}, pt_path) - - data = serialize_from_pt_file(str(pt_path)) - assert data["model"]["type"] == "SeZM" - - deserialize_to_jax_file(str(jax_path), data) - restored = serialize_from_jax_file(str(jax_path)) - - assert restored["model"]["type"] == "standard" - assert restored["model"]["descriptor"]["type"] == "SeZM" diff --git a/source/tests/jax/test_model_factory.py b/source/tests/jax/test_model_factory.py index 8ab2132121..75ffc519a1 100644 --- a/source/tests/jax/test_model_factory.py +++ b/source/tests/jax/test_model_factory.py @@ -10,25 +10,13 @@ """ import unittest -from unittest.mock import ( - patch, -) -from deepmd.jax.model.base_model import ( - BaseModel, -) from deepmd.jax.model.ener_model import ( EnergyModel, ) from deepmd.jax.model.model import ( get_model, ) -from deepmd.jax.train.trainer import ( - prepare_input, -) -from deepmd.utils.argcheck import ( - model_args, -) def _base_config() -> dict: @@ -52,20 +40,6 @@ def _base_config() -> dict: } -def _base_sezm_config() -> dict: - """Return the smallest config needed to exercise DPA4 factory routing.""" - return { - "type": "dpa4", - "type_map": ["O", "H"], - "descriptor": { - "type": "dpa4", - "random_gamma": False, - "use_amp": False, - }, - "fitting_net": {"type": "dpa4_ener"}, - } - - class TestJAXModelFactoryFittingDefault(unittest.TestCase): def test_fitting_net_without_type_defaults_to_ener(self) -> None: # fitting_net present but no "type": must default to energy. @@ -88,137 +62,5 @@ def test_explicit_fitting_type_preserved(self) -> None: self.assertIsInstance(model, EnergyModel) -class TestJAXSeZMModelFactory(unittest.TestCase): - @patch("deepmd.jax.model.model.get_standard_model", side_effect=lambda data: data) - def test_null_blocks_receive_dpa4_defaults(self, _get_standard_model) -> None: - data = _base_sezm_config() - data["descriptor"] = None - data["fitting_net"] = None - - normalized = get_model(data) - - self.assertEqual(normalized["descriptor"]["type"], "dpa4") - self.assertEqual(normalized["fitting_net"]["type"], "dpa4_ener") - - def test_rejects_unsupported_features(self) -> None: - cases = ( - ("spin", {}), - ("bridging_method", "linear"), - ("lora", {}), - ("use_compile", True), - ("preset_out_bias", [0.0]), - ) - for key, value in cases: - with self.subTest(key=key): - data = _base_sezm_config() - data[key] = value - with self.assertRaises(NotImplementedError): - get_model(data) - - data = _base_sezm_config() - data["descriptor"]["add_chg_spin_ebd"] = True - with self.assertRaises(NotImplementedError): - get_model(data) - - for key in ("random_gamma", "use_amp"): - with self.subTest(descriptor_option=key): - data = _base_sezm_config() - data["descriptor"][key] = True - with self.assertRaisesRegex(NotImplementedError, key): - get_model(data) - - def test_rejects_incompatible_descriptor_and_fitting_types(self) -> None: - data = _base_sezm_config() - data["descriptor"]["type"] = "se_e2_a" - with self.assertRaises(ValueError): - get_model(data) - - data = _base_sezm_config() - data["fitting_net"]["type"] = "ener" - with self.assertRaises(ValueError): - get_model(data) - - def test_rejects_mismatched_exclude_types(self) -> None: - data = _base_sezm_config() - data["descriptor"]["exclude_types"] = [[0, 1]] - data["pair_exclude_types"] = [[1, 1]] - - with self.assertRaises(ValueError): - get_model(data) - - @patch( - "deepmd.dpmodel.model.dp_model.BaseDescriptor.update_sel", - return_value=({"type": "dpa4", "sel": 16}, 0.75), - ) - def test_model_aliases_route_through_update_sel(self, update_sel) -> None: - """Neighbor-stat preprocessing recognizes every public DPA4 alias.""" - for model_type in ("dpa4", "DPA4", "sezm", "SeZM"): - with self.subTest(model_type=model_type): - local_jdata = { - "type": model_type, - "descriptor": {"type": "dpa4", "sel": "auto"}, - } - - updated, min_nbor_dist = BaseModel.update_sel( - object(), ["O", "H"], local_jdata - ) - - self.assertEqual(updated["descriptor"]["sel"], 16) - self.assertEqual(min_nbor_dist, 0.75) - self.assertEqual(update_sel.call_count, 4) - - def test_dpa4_conservative_input_keeps_all_in_cutoff_neighbors(self) -> None: - """The JAX DPA4 input boundary must not truncate to the configured sel.""" - import numpy as np - - coord = np.asarray( - [[[0.0, 0.0, 0.0], [0.5, 0.0, 0.0], [0.0, 0.5, 0.0], [0.0, 0.0, 0.5]]] - ) - atype = np.zeros((1, 4), dtype=np.int32) - - _, _, nlist, _, _, _ = prepare_input( - rcut=2.0, - sel=[1], - coord=coord, - atype=atype, - conservative_nlist=True, - ) - - self.assertGreaterEqual(nlist.shape[-1], 4) - self.assertTrue(np.all(np.sum(nlist >= 0, axis=-1) >= 3)) - - @patch("deepmd.jax.model.model.get_standard_model", side_effect=lambda data: data) - def test_descriptor_exclude_types_feed_standard_model( - self, - _get_standard_model, - ) -> None: - data = _base_sezm_config() - data["descriptor"] = { - "type": "SeZM", - "exclude_types": [[0, 1]], - } - data["fitting_net"]["type"] = "sezm_ener" - - normalized = get_model(data) - - self.assertEqual(normalized["pair_exclude_types"], [[0, 1]]) - self.assertEqual(normalized["descriptor"]["exclude_types"], [[0, 1]]) - - @patch("deepmd.jax.model.model.get_standard_model", side_effect=lambda data: data) - def test_normalized_descriptor_exclusions_override_empty_default( - self, - _get_standard_model, - ) -> None: - """Argcheck's empty model-level default is not an explicit mismatch.""" - data = _base_sezm_config() - data["descriptor"]["exclude_types"] = [[0, 1]] - data = model_args().normalize_value(data, trim_pattern="_.*") - - normalized = get_model(data) - - self.assertEqual(normalized["pair_exclude_types"], [[0, 1]]) - self.assertEqual(normalized["descriptor"]["exclude_types"], [[0, 1]]) - - if __name__ == "__main__": unittest.main() diff --git a/source/tests/jax/test_training.py b/source/tests/jax/test_training.py index 74e81d07f5..72e8a47ec7 100644 --- a/source/tests/jax/test_training.py +++ b/source/tests/jax/test_training.py @@ -49,9 +49,6 @@ from deepmd.jax.train.trainer import ( DPTrainer, _copy_matching_state_tree, - _drop_zero_size_array_leaves, - _evaluate_model_dict, - _match_label_shapes, _merge_descriptor_stats, _merge_fitting_param_stats, _scale_by_global_learning_rate, @@ -61,7 +58,6 @@ ) from deepmd.jax.utils.serialization import ( _normalize_restored_state_keys, - _restore_missing_zero_size_leaves, ) from deepmd.utils.compat import ( convert_optimizer_v31_to_v32, @@ -933,96 +929,3 @@ def test_jax_multitask_state_key_normalization_preserves_numeric_task_names() -> assert 1 not in state["models"] assert 0 in state["models"]["1"]["layers"] assert 0 in state["models"]["task"]["layers"] - - -def test_jax_zero_size_checkpoint_leaves_round_trip() -> None: - """Checkpoint filtering and restore must preserve every zero-size path.""" - template = { - "model": { - "empty": jnp.zeros((0, 3)), - "nested": { - "empty": jnp.zeros((2, 0)), - "weight": jnp.ones((2,)), - }, - } - } - - filtered = _drop_zero_size_array_leaves(template) - restored = _restore_missing_zero_size_leaves(template, filtered) - - assert "empty" not in filtered["model"] - assert "empty" not in filtered["model"]["nested"] - np.testing.assert_array_equal( - restored["model"]["empty"], template["model"]["empty"] - ) - np.testing.assert_array_equal( - restored["model"]["nested"]["empty"], - template["model"]["nested"]["empty"], - ) - np.testing.assert_array_equal( - restored["model"]["nested"]["weight"], - template["model"]["nested"]["weight"], - ) - - -def test_jax_match_label_shapes_reshapes_only_equivalent_force_layouts() -> None: - """Flattened force tensors reshape, while an existing layout is untouched.""" - force = jnp.arange(6).reshape(1, 2, 3) - model_dict = {"force": force} - - reshaped = _match_label_shapes(model_dict, {"force": jnp.zeros((1, 6))}) - unchanged = _match_label_shapes(model_dict, {"force": jnp.zeros((1, 2, 3))}) - - assert reshaped is not model_dict - assert reshaped["force"].shape == (1, 6) - assert unchanged is model_dict - - -def test_jax_evaluate_model_dict_normalizes_flattened_force_only() -> None: - """Model evaluation preserves canonical force tensors and expands flat ones.""" - - class FakeModel: - def __init__(self, force_derivative: jnp.ndarray) -> None: - self.force_derivative = force_derivative - - def call_common_lower(self, *args, **kwargs): - del args, kwargs - return { - "energy": jnp.zeros((1, 2, 1)), - "energy_redu": jnp.zeros((1, 1)), - "energy_derv_r": self.force_derivative, - "energy_derv_c_redu": jnp.zeros((1, 1, 9)), - } - - def model_output_def(self): - return {} - - def passthrough(model_dict, *args, **kwargs): - del args, kwargs - return dict(model_dict) - - with patch( - "deepmd.jax.train.trainer.communicate_extended_output", - side_effect=passthrough, - ): - flattened = _evaluate_model_dict( - FakeModel(jnp.arange(6).reshape(1, 1, 6)), - jnp.zeros((1, 6)), - jnp.zeros((1, 2), dtype=jnp.int32), - jnp.zeros((1, 2, 1), dtype=jnp.int32), - None, - None, - None, - ) - canonical = _evaluate_model_dict( - FakeModel(jnp.arange(6).reshape(1, 2, 1, 3)), - jnp.zeros((1, 6)), - jnp.zeros((1, 2), dtype=jnp.int32), - jnp.zeros((1, 2, 1), dtype=jnp.int32), - None, - None, - None, - ) - - assert flattened["force"].shape == (1, 2, 3) - assert canonical["force"].shape == (1, 2, 3) diff --git a/source/tests/pt/model/test_sezm_model.py b/source/tests/pt/model/test_sezm_model.py index 192008cb8a..494338a2da 100644 --- a/source/tests/pt/model/test_sezm_model.py +++ b/source/tests/pt/model/test_sezm_model.py @@ -47,9 +47,6 @@ from deepmd.pt.model.model.sezm_property_model import ( SeZMPropertyModel, ) -from deepmd.pt.model.task.sezm_ener import ( - SeZMEnergyFittingNet, -) from deepmd.pt.train.training import ( prepare_model_for_loss, ) @@ -90,29 +87,6 @@ ) -class TestSeZMEnergyFittingTrainability(unittest.TestCase): - """Ensure frozen PT DPA4 fitting parameters survive serialization.""" - - def test_frozen_parameters_survive_round_trip(self) -> None: - fitting = SeZMEnergyFittingNet( - ntypes=2, - dim_descrpt=8, - neuron=[8], - mixed_types=True, - trainable=False, - dim_case_embd=2, - case_film_embd=True, - precision="float64", - seed=20260716, - ) - - restored = SeZMEnergyFittingNet.deserialize(fitting.serialize()) - - self.assertTrue(list(fitting.parameters())) - self.assertTrue(list(restored.parameters())) - self.assertTrue(all(not param.requires_grad for param in restored.parameters())) - - def _assert_close_with_strict_warning( actual: torch.Tensor, expected: torch.Tensor,