diff --git a/deepmd/jax/atomic_model/base_atomic_model.py b/deepmd/jax/atomic_model/base_atomic_model.py index bed75077da..6ceb116d85 100644 --- a/deepmd/jax/atomic_model/base_atomic_model.py +++ b/deepmd/jax/atomic_model/base_atomic_model.py @@ -1,35 +1 @@ # SPDX-License-Identifier: LGPL-3.0-or-later -from typing import ( - Any, -) - -from packaging.version import ( - Version, -) - -from deepmd.jax.common import ( - ArrayAPIVariable, - to_jax_array, -) -from deepmd.jax.env import ( - flax_version, - nnx, -) -from deepmd.jax.utils.exclude_mask import ( - AtomExcludeMask, - PairExcludeMask, -) - - -def base_atomic_model_set_attr(name: str, value: Any) -> Any: - if name in {"out_bias", "out_std"}: - value = to_jax_array(value) - if value is not None: - value = ArrayAPIVariable(value) - elif Version(flax_version) >= Version("0.12.0"): - value = nnx.data(value) - elif name == "pair_excl" and value is not None: - value = PairExcludeMask(value.ntypes, value.exclude_types) - elif name == "atom_excl" and value is not None: - value = AtomExcludeMask(value.ntypes, value.exclude_types) - return value diff --git a/deepmd/jax/atomic_model/dp_atomic_model.py b/deepmd/jax/atomic_model/dp_atomic_model.py index 74a1f481ea..8db31d6d5d 100644 --- a/deepmd/jax/atomic_model/dp_atomic_model.py +++ b/deepmd/jax/atomic_model/dp_atomic_model.py @@ -1,12 +1,8 @@ # SPDX-License-Identifier: LGPL-3.0-or-later -from typing import ( - Any, -) - +import deepmd.jax.descriptor as _jax_descriptor # noqa: F401 +import deepmd.jax.fitting.fitting as _jax_fitting # noqa: F401 +import deepmd.jax.utils.exclude_mask as _jax_exclude_mask # noqa: F401 from deepmd.dpmodel.atomic_model.dp_atomic_model import DPAtomicModel as DPAtomicModelDP -from deepmd.jax.atomic_model.base_atomic_model import ( - base_atomic_model_set_attr, -) from deepmd.jax.common import ( flax_module, ) @@ -45,10 +41,6 @@ class jax_atomic_model(dpmodel_atomic_model): base_fitting_cls = BaseFitting """The base fitting class.""" - def __setattr__(self, name: str, value: Any) -> None: - value = base_atomic_model_set_attr(name, value) - return super().__setattr__(name, value) - def forward_common_atomic( self, extended_coord: jnp.ndarray, diff --git a/deepmd/jax/atomic_model/linear_atomic_model.py b/deepmd/jax/atomic_model/linear_atomic_model.py index 1453e8f495..ae9bae6c4a 100644 --- a/deepmd/jax/atomic_model/linear_atomic_model.py +++ b/deepmd/jax/atomic_model/linear_atomic_model.py @@ -3,54 +3,28 @@ Any, ) -from packaging.version import ( - Version, -) - +import deepmd.jax.atomic_model.dp_atomic_model as _jax_dp_atomic_model # noqa: F401 +import deepmd.jax.atomic_model.pairtab_atomic_model as _jax_pairtab_model # noqa: F401 +import deepmd.jax.utils.exclude_mask as _jax_exclude_mask # noqa: F401 from deepmd.dpmodel.atomic_model.linear_atomic_model import ( DPZBLLinearEnergyAtomicModel as DPZBLLinearEnergyAtomicModelDP, ) -from deepmd.jax.atomic_model.base_atomic_model import ( - base_atomic_model_set_attr, -) -from deepmd.jax.atomic_model.dp_atomic_model import ( - DPAtomicModel, -) -from deepmd.jax.atomic_model.pairtab_atomic_model import ( - PairTabAtomicModel, -) from deepmd.jax.common import ( - ArrayAPIVariable, flax_module, - to_jax_array, ) from deepmd.jax.env import ( - flax_version, jax, jnp, - nnx, ) @flax_module class DPZBLLinearEnergyAtomicModel(DPZBLLinearEnergyAtomicModelDP): def __setattr__(self, name: str, value: Any) -> None: - value = base_atomic_model_set_attr(name, value) - if name == "mapping_list": - value = [ArrayAPIVariable(to_jax_array(vv)) for vv in value] - if Version(flax_version) >= Version("0.12.0"): - value = nnx.List([nnx.data(item) for item in value]) - elif name == "zbl_weight": + if name == "zbl_weight": # discard since it's only used in tests # to fix flax.errors.TraceContextError: Cannot mutate 'FlaxModule' from different trace level return - elif name == "models": - value = [ - DPAtomicModel.deserialize(value[0].serialize()), - PairTabAtomicModel.deserialize(value[1].serialize()), - ] - if Version(flax_version) >= Version("0.12.0"): - value = nnx.List([nnx.data(item) for item in value]) return super().__setattr__(name, value) def forward_common_atomic( diff --git a/deepmd/jax/atomic_model/pairtab_atomic_model.py b/deepmd/jax/atomic_model/pairtab_atomic_model.py index 4f5a5d0ece..3db9070c0c 100644 --- a/deepmd/jax/atomic_model/pairtab_atomic_model.py +++ b/deepmd/jax/atomic_model/pairtab_atomic_model.py @@ -1,43 +1,19 @@ # SPDX-License-Identifier: LGPL-3.0-or-later -from typing import ( - Any, -) - -from packaging.version import ( - Version, -) - +import deepmd.jax.utils.exclude_mask as _jax_exclude_mask # noqa: F401 from deepmd.dpmodel.atomic_model.pairtab_atomic_model import ( PairTabAtomicModel as PairTabAtomicModelDP, ) -from deepmd.jax.atomic_model.base_atomic_model import ( - base_atomic_model_set_attr, -) from deepmd.jax.common import ( - ArrayAPIVariable, flax_module, - to_jax_array, ) from deepmd.jax.env import ( - flax_version, jax, jnp, - nnx, ) @flax_module class PairTabAtomicModel(PairTabAtomicModelDP): - def __setattr__(self, name: str, value: Any) -> None: - value = base_atomic_model_set_attr(name, value) - if name in {"tab_info", "tab_data"}: - value = to_jax_array(value) - if value is not None: - value = ArrayAPIVariable(value) - elif Version(flax_version) >= Version("0.12.0"): - value = nnx.data(value) - return super().__setattr__(name, value) - def forward_common_atomic( self, extended_coord: jnp.ndarray, diff --git a/deepmd/jax/common.py b/deepmd/jax/common.py index 668c7f5786..aacf375a74 100644 --- a/deepmd/jax/common.py +++ b/deepmd/jax/common.py @@ -1,7 +1,17 @@ # SPDX-License-Identifier: LGPL-3.0-or-later +from collections.abc import ( + Callable, +) from functools import ( wraps, ) +from importlib import ( + import_module, +) +from threading import ( + Condition, + get_ident, +) from typing import ( Any, TypeVar, @@ -9,13 +19,34 @@ ) import numpy as np +from packaging.version import ( + Version, +) +from deepmd.dpmodel.common import ( + NativeOP, +) from deepmd.jax.env import ( + flax_version, jnp, nnx, ) +class ArrayAPIVariable(nnx.Variable): + def __array__(self, *args: Any, **kwargs: Any) -> np.ndarray: + return self.value.__array__(*args, **kwargs) + + def __array_namespace__(self, *args: Any, **kwargs: Any) -> Any: + return self.value.__array_namespace__(*args, **kwargs) + + def __dlpack__(self, *args: Any, **kwargs: Any) -> Any: + return self.value.__dlpack__(*args, **kwargs) + + def __dlpack_device__(self, *args: Any, **kwargs: Any) -> Any: + return self.value.__dlpack_device__(*args, **kwargs) + + @overload def to_jax_array(array: np.ndarray) -> jnp.ndarray: ... @@ -42,6 +73,166 @@ def to_jax_array(array: np.ndarray | None) -> jnp.ndarray | None: return jnp.array(array) +_DPMODEL_TO_JAX: dict[type[Any], Callable[[Any], Any]] = {} +_AUTO_WRAPPED_CLASSES: dict[type[NativeOP], type[Any]] = {} +_FLAX_0_12 = Version("0.12.0") +_REGISTRATIONS_READY = False +_REGISTRATIONS_IN_PROGRESS = False +_REGISTRATIONS_OWNER: int | None = None +_REGISTRATIONS_COND = Condition() +_REGISTRATION_MODULES = ( + "deepmd.jax.utils.network", + "deepmd.jax.utils.exclude_mask", + "deepmd.jax.utils.type_embed", + "deepmd.jax.descriptor", + "deepmd.jax.fitting", + "deepmd.jax.atomic_model.dp_atomic_model", + "deepmd.jax.atomic_model.energy_atomic_model", + "deepmd.jax.atomic_model.dipole_atomic_model", + "deepmd.jax.atomic_model.dos_atomic_model", + "deepmd.jax.atomic_model.polar_atomic_model", + "deepmd.jax.atomic_model.property_atomic_model", + "deepmd.jax.atomic_model.pairtab_atomic_model", + "deepmd.jax.atomic_model.linear_atomic_model", + "deepmd.jax.model", +) + + +def register_dpmodel_mapping( + dpmodel_cls: type[Any], converter: Callable[[Any], Any] +) -> None: + """Register how to convert a dpmodel object to its JAX wrapper.""" + _DPMODEL_TO_JAX[dpmodel_cls] = converter + + +def _looks_like_dpmodel_object(value: Any) -> bool: + module = type(value).__module__ + return module == "deepmd.dpmodel" or module.startswith("deepmd.dpmodel.") + + +def _ensure_registrations() -> None: + global _REGISTRATIONS_IN_PROGRESS, _REGISTRATIONS_OWNER, _REGISTRATIONS_READY + + current_thread = get_ident() + with _REGISTRATIONS_COND: + if _REGISTRATIONS_READY: + return + while _REGISTRATIONS_IN_PROGRESS: + if _REGISTRATIONS_OWNER == current_thread: + return + _REGISTRATIONS_COND.wait() + if _REGISTRATIONS_READY: + return + _REGISTRATIONS_IN_PROGRESS = True + _REGISTRATIONS_OWNER = current_thread + + success = False + try: + for module in _REGISTRATION_MODULES: + import_module(module) + success = True + finally: + with _REGISTRATIONS_COND: + _REGISTRATIONS_READY = success + _REGISTRATIONS_IN_PROGRESS = False + _REGISTRATIONS_OWNER = None + _REGISTRATIONS_COND.notify_all() + + +def try_convert_module(value: Any) -> Any | None: + """Convert a registered dpmodel object to its JAX wrapper.""" + converter = _DPMODEL_TO_JAX.get(type(value)) + if converter is not None: + return converter(value) + if _looks_like_dpmodel_object(value): + _ensure_registrations() + converter = _DPMODEL_TO_JAX.get(type(value)) + if converter is not None: + return converter(value) + if isinstance(value, NativeOP): + return _auto_wrap_native_op(value) + return None + + +def _auto_wrap_native_op(value: NativeOP) -> Any: + cls = type(value) + if cls not in _AUTO_WRAPPED_CLASSES: + _AUTO_WRAPPED_CLASSES[cls] = flax_module(cls) + wrapped_cls = _AUTO_WRAPPED_CLASSES[cls] + if not (hasattr(value, "serialize") and hasattr(wrapped_cls, "deserialize")): + raise TypeError( + f"Cannot auto-wrap {cls.__name__}: " + "it must implement serialize()/deserialize() or be explicitly " + "registered via register_dpmodel_mapping()." + ) + return wrapped_cls.deserialize(value.serialize()) + + +def _use_nnx_list() -> bool: + return Version(flax_version) >= _FLAX_0_12 and hasattr(nnx, "List") + + +def _wrap_list(value: list[Any]) -> Any: + if _use_nnx_list(): + return nnx.List([nnx.data(item) for item in value]) + return value + + +def _try_convert_list(value: list[Any]) -> Any | None: + if not value: + return None + + converted = [] + changed = False + for item in value: + if isinstance(item, np.ndarray): + converted.append(ArrayAPIVariable(to_jax_array(item))) + changed = True + elif isinstance(item, (nnx.Module, nnx.Variable)): + converted.append(item) + changed = True + elif item is None: + converted.append(item) + else: + module = try_convert_module(item) + if module is None: + return None + converted.append(module) + changed = True + + if not changed: + return None + return _wrap_list(converted) + + +def dpmodel_setattr(obj: nnx.Module, name: str, value: Any) -> tuple[bool, Any]: + """Common ``__setattr__`` conversion for Flax wrappers around dpmodel objects.""" + if name in getattr(obj, "_jax_skip_auto_convert_attrs", ()): + return False, value + + if ( + isinstance(value, list) + and name in getattr(obj, "_jax_data_list_attrs", ()) + and _use_nnx_list() + ): + return False, _try_convert_list(value) or _wrap_list(value) + + if isinstance(value, np.ndarray): + return False, ArrayAPIVariable(to_jax_array(value)) + + if isinstance(value, list): + converted_list = _try_convert_list(value) + if converted_list is not None: + return False, converted_list + + if not isinstance(value, nnx.Module): + converted = try_convert_module(value) + if converted is not None: + return False, converted + + return False, value + + T = TypeVar("T") @@ -82,20 +273,32 @@ def __init_subclass__(cls, **kwargs: Any) -> None: return super().__init_subclass__(**kwargs) def __setattr__(self, name: str, value: Any) -> None: - return super().__setattr__(name, value) - - return FlaxModule - + handled, value = dpmodel_setattr(self, name, value) + if not handled: + try: + return super().__setattr__(name, value) + except ValueError as err: + msg = str(err) + if ( + Version(flax_version) >= _FLAX_0_12 + and "Cannot assign data value" in msg + and "to static attribute" in msg + ): + return super().__setattr__(name, nnx.data(value)) + raise + return None -class ArrayAPIVariable(nnx.Variable): - def __array__(self, *args: Any, **kwargs: Any) -> np.ndarray: - return self.value.__array__(*args, **kwargs) + if hasattr(FlaxModule, "deserialize"): + for base in module.__bases__: + if base in (object, NativeOP, nnx.Module): + continue + if issubclass(base, nnx.Module): + continue + if hasattr(base, "serialize") and base not in _DPMODEL_TO_JAX: - def __array_namespace__(self, *args: Any, **kwargs: Any) -> Any: - return self.value.__array_namespace__(*args, **kwargs) + def _converter(v: Any, _cls: type[Any] = FlaxModule) -> Any: + return _cls.deserialize(v.serialize()) - def __dlpack__(self, *args: Any, **kwargs: Any) -> Any: - return self.value.__dlpack__(*args, **kwargs) + _DPMODEL_TO_JAX[base] = _converter - def __dlpack_device__(self, *args: Any, **kwargs: Any) -> Any: - return self.value.__dlpack_device__(*args, **kwargs) + return FlaxModule diff --git a/deepmd/jax/descriptor/dpa1.py b/deepmd/jax/descriptor/dpa1.py index 07695b23ed..cdf1fa99a2 100644 --- a/deepmd/jax/descriptor/dpa1.py +++ b/deepmd/jax/descriptor/dpa1.py @@ -1,12 +1,7 @@ # SPDX-License-Identifier: LGPL-3.0-or-later -from typing import ( - Any, -) - -from packaging.version import ( - Version, -) - +import deepmd.jax.utils.exclude_mask as _jax_exclude_mask # noqa: F401 +import deepmd.jax.utils.network as _jax_network # noqa: F401 +import deepmd.jax.utils.type_embed as _jax_type_embed # noqa: F401 from deepmd.dpmodel.descriptor.dpa1 import DescrptBlockSeAtten as DescrptBlockSeAttenDP from deepmd.dpmodel.descriptor.dpa1 import DescrptDPA1 as DescrptDPA1DP from deepmd.dpmodel.descriptor.dpa1 import GatedAttentionLayer as GatedAttentionLayerDP @@ -17,92 +12,35 @@ NeighborGatedAttentionLayer as NeighborGatedAttentionLayerDP, ) from deepmd.jax.common import ( - ArrayAPIVariable, flax_module, - to_jax_array, ) from deepmd.jax.descriptor.base_descriptor import ( BaseDescriptor, ) -from deepmd.jax.env import ( - flax_version, - nnx, -) -from deepmd.jax.utils.exclude_mask import ( - PairExcludeMask, -) -from deepmd.jax.utils.network import ( - LayerNorm, - NativeLayer, - NetworkCollection, -) -from deepmd.jax.utils.type_embed import ( - TypeEmbedNet, -) @flax_module class GatedAttentionLayer(GatedAttentionLayerDP): - def __setattr__(self, name: str, value: Any) -> None: - if name in {"in_proj", "out_proj"}: - value = NativeLayer.deserialize(value.serialize()) - return super().__setattr__(name, value) + pass @flax_module class NeighborGatedAttentionLayer(NeighborGatedAttentionLayerDP): - def __setattr__(self, name: str, value: Any) -> None: - if name == "attention_layer": - value = GatedAttentionLayer.deserialize(value.serialize()) - elif name == "attn_layer_norm": - value = LayerNorm.deserialize(value.serialize()) - return super().__setattr__(name, value) + pass @flax_module class NeighborGatedAttention(NeighborGatedAttentionDP): - def __setattr__(self, name: str, value: Any) -> None: - if name == "attention_layers": - value = [ - NeighborGatedAttentionLayer.deserialize(ii.serialize()) for ii in value - ] - if Version(flax_version) >= Version("0.12.0"): - value = nnx.List([nnx.data(item) for item in value]) - return super().__setattr__(name, value) + pass @flax_module class DescrptBlockSeAtten(DescrptBlockSeAttenDP): - def __setattr__(self, name: str, value: Any) -> None: - if name in {"mean", "stddev"}: - value = to_jax_array(value) - if value is not None: - value = ArrayAPIVariable(value) - elif Version(flax_version) >= Version("0.12.0"): - value = nnx.data(value) - elif name in {"embeddings", "embeddings_strip"}: - if value is not None: - value = NetworkCollection.deserialize(value.serialize()) - elif Version(flax_version) >= Version("0.12.0"): - value = nnx.data(value) - elif name == "dpa1_attention": - value = NeighborGatedAttention.deserialize(value.serialize()) - elif name == "env_mat": - # env_mat doesn't store any value - pass - elif name == "emask": - value = PairExcludeMask(value.ntypes, value.exclude_types) - - return super().__setattr__(name, value) + pass @BaseDescriptor.register("dpa1") @BaseDescriptor.register("se_atten") @flax_module class DescrptDPA1(DescrptDPA1DP): - def __setattr__(self, name: str, value: Any) -> None: - if name == "se_atten": - value = DescrptBlockSeAtten.deserialize(value.serialize()) - elif name == "type_embedding": - value = TypeEmbedNet.deserialize(value.serialize()) - return super().__setattr__(name, value) + pass diff --git a/deepmd/jax/descriptor/dpa2.py b/deepmd/jax/descriptor/dpa2.py index 8da450d2ec..d06a1ff554 100644 --- a/deepmd/jax/descriptor/dpa2.py +++ b/deepmd/jax/descriptor/dpa2.py @@ -1,74 +1,19 @@ # SPDX-License-Identifier: LGPL-3.0-or-later -from typing import ( - Any, -) - -from packaging.version import ( - Version, -) - +import deepmd.jax.descriptor.dpa1 as _jax_dpa1 # noqa: F401 +import deepmd.jax.descriptor.repformers as _jax_repformers # noqa: F401 +import deepmd.jax.descriptor.se_t_tebd as _jax_se_t_tebd # noqa: F401 +import deepmd.jax.utils.network as _jax_network # noqa: F401 +import deepmd.jax.utils.type_embed as _jax_type_embed # noqa: F401 from deepmd.dpmodel.descriptor.dpa2 import DescrptDPA2 as DescrptDPA2DP -from deepmd.dpmodel.utils.network import Identity as IdentityDP -from deepmd.dpmodel.utils.network import NativeLayer as NativeLayerDP from deepmd.jax.common import ( - ArrayAPIVariable, flax_module, - to_jax_array, ) from deepmd.jax.descriptor.base_descriptor import ( BaseDescriptor, ) -from deepmd.jax.descriptor.dpa1 import ( - DescrptBlockSeAtten, -) -from deepmd.jax.descriptor.repformers import ( - DescrptBlockRepformers, -) -from deepmd.jax.descriptor.se_t_tebd import ( - DescrptBlockSeTTebd, -) -from deepmd.jax.env import ( - flax_version, - nnx, -) -from deepmd.jax.utils.network import ( - NativeLayer, -) -from deepmd.jax.utils.type_embed import ( - TypeEmbedNet, -) @BaseDescriptor.register("dpa2") @flax_module class DescrptDPA2(DescrptDPA2DP): - def __setattr__(self, name: str, value: Any) -> None: - if name in {"mean", "stddev"}: - value = to_jax_array(value) - if value is not None: - value = ArrayAPIVariable(value) - elif Version(flax_version) >= Version("0.12.0"): - value = nnx.data(value) - elif name in {"repinit"}: - value = DescrptBlockSeAtten.deserialize(value.serialize()) - elif name in {"repinit_three_body"}: - if value is not None: - value = DescrptBlockSeTTebd.deserialize(value.serialize()) - elif Version(flax_version) >= Version("0.12.0"): - value = nnx.data(value) - elif name in {"repformers"}: - value = DescrptBlockRepformers.deserialize(value.serialize()) - elif name in {"type_embedding"}: - value = TypeEmbedNet.deserialize(value.serialize()) - elif name in {"g1_shape_tranform", "tebd_transform"}: - if value is None: - if Version(flax_version) >= Version("0.12.0"): - value = nnx.data(value) - elif isinstance(value, NativeLayerDP): - value = NativeLayer.deserialize(value.serialize()) - elif isinstance(value, IdentityDP): - # IdentityDP doesn't contain any value - it's good to go - pass - else: - raise ValueError(f"Unknown layer type: {type(value)}") - return super().__setattr__(name, value) + pass diff --git a/deepmd/jax/descriptor/dpa3.py b/deepmd/jax/descriptor/dpa3.py index 226acc48db..e236e55d4b 100644 --- a/deepmd/jax/descriptor/dpa3.py +++ b/deepmd/jax/descriptor/dpa3.py @@ -1,54 +1,17 @@ # SPDX-License-Identifier: LGPL-3.0-or-later -from typing import ( - Any, -) - -from packaging.version import ( - Version, -) - +import deepmd.jax.descriptor.repflows as _jax_repflows # noqa: F401 +import deepmd.jax.utils.network as _jax_network # noqa: F401 +import deepmd.jax.utils.type_embed as _jax_type_embed # noqa: F401 from deepmd.dpmodel.descriptor.dpa3 import DescrptDPA3 as DescrptDPA3DP from deepmd.jax.common import ( - ArrayAPIVariable, flax_module, - to_jax_array, ) from deepmd.jax.descriptor.base_descriptor import ( BaseDescriptor, ) -from deepmd.jax.descriptor.repflows import ( - DescrptBlockRepflows, -) -from deepmd.jax.env import ( - flax_version, - nnx, -) -from deepmd.jax.utils.network import ( - NativeLayer, -) -from deepmd.jax.utils.type_embed import ( - TypeEmbedNet, -) @BaseDescriptor.register("dpa3") @flax_module class DescrptDPA3(DescrptDPA3DP): - def __setattr__(self, name: str, value: Any) -> None: - if name in {"mean", "stddev"}: - value = to_jax_array(value) - if value is not None: - value = ArrayAPIVariable(value) - elif Version(flax_version) >= Version("0.12.0"): - value = nnx.data(value) - elif name in {"repflows"}: - value = DescrptBlockRepflows.deserialize(value.serialize()) - elif name in {"type_embedding", "chg_embedding", "spin_embedding"}: - if value is not None: - value = TypeEmbedNet.deserialize(value.serialize()) - elif name in {"mix_cs_mlp"}: - if value is not None: - value = NativeLayer.deserialize(value.serialize()) - else: - pass - return super().__setattr__(name, value) + pass diff --git a/deepmd/jax/descriptor/hybrid.py b/deepmd/jax/descriptor/hybrid.py index b76e515c54..a4615fa0bf 100644 --- a/deepmd/jax/descriptor/hybrid.py +++ b/deepmd/jax/descriptor/hybrid.py @@ -1,38 +1,22 @@ # SPDX-License-Identifier: LGPL-3.0-or-later -from typing import ( - Any, -) - -from packaging.version import ( - Version, -) - +import deepmd.jax.descriptor.dpa1 as _jax_dpa1 # noqa: F401 +import deepmd.jax.descriptor.dpa2 as _jax_dpa2 # noqa: F401 +import deepmd.jax.descriptor.dpa3 as _jax_dpa3 # noqa: F401 +import deepmd.jax.descriptor.se_atten_v2 as _jax_se_atten_v2 # noqa: F401 +import deepmd.jax.descriptor.se_e2_a as _jax_se_e2_a # noqa: F401 +import deepmd.jax.descriptor.se_e2_r as _jax_se_e2_r # noqa: F401 +import deepmd.jax.descriptor.se_t as _jax_se_t # noqa: F401 +import deepmd.jax.descriptor.se_t_tebd as _jax_se_t_tebd # noqa: F401 from deepmd.dpmodel.descriptor.hybrid import DescrptHybrid as DescrptHybridDP from deepmd.jax.common import ( - ArrayAPIVariable, flax_module, - to_jax_array, ) from deepmd.jax.descriptor.base_descriptor import ( BaseDescriptor, ) -from deepmd.jax.env import ( - flax_version, - nnx, -) @BaseDescriptor.register("hybrid") @flax_module class DescrptHybrid(DescrptHybridDP): - def __setattr__(self, name: str, value: Any) -> None: - if name in {"nlist_cut_idx"}: - value = [ArrayAPIVariable(to_jax_array(vv)) for vv in value] - if Version(flax_version) >= Version("0.12.0"): - value = nnx.List([nnx.data(item) for item in value]) - elif name in {"descrpt_list"}: - value = [BaseDescriptor.deserialize(vv.serialize()) for vv in value] - if Version(flax_version) >= Version("0.12.0"): - value = nnx.List([nnx.data(item) for item in value]) - - return super().__setattr__(name, value) + pass diff --git a/deepmd/jax/descriptor/repflows.py b/deepmd/jax/descriptor/repflows.py index be26012a52..97db6c81a9 100644 --- a/deepmd/jax/descriptor/repflows.py +++ b/deepmd/jax/descriptor/repflows.py @@ -1,81 +1,28 @@ # SPDX-License-Identifier: LGPL-3.0-or-later from typing import ( - Any, -) - -from packaging.version import ( - Version, + ClassVar, ) +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.descriptor.repflows import ( DescrptBlockRepflows as DescrptBlockRepflowsDP, ) from deepmd.dpmodel.descriptor.repflows import RepFlowLayer as RepFlowLayerDP from deepmd.jax.common import ( - ArrayAPIVariable, flax_module, - to_jax_array, -) -from deepmd.jax.env import ( - flax_version, - nnx, -) -from deepmd.jax.utils.exclude_mask import ( - PairExcludeMask, -) -from deepmd.jax.utils.network import ( - NativeLayer, ) @flax_module class DescrptBlockRepflows(DescrptBlockRepflowsDP): - def __setattr__(self, name: str, value: Any) -> None: - if name in {"mean", "stddev"}: - value = to_jax_array(value) - if value is not None: - value = ArrayAPIVariable(value) - elif Version(flax_version) >= Version("0.12.0"): - value = nnx.data(value) - elif name in {"layers"}: - value = [RepFlowLayer.deserialize(layer.serialize()) for layer in value] - if Version(flax_version) >= Version("0.12.0"): - value = nnx.List([nnx.data(item) for item in value]) - elif name in {"edge_embd", "angle_embd"}: - value = NativeLayer.deserialize(value.serialize()) - elif name in {"env_mat_edge", "env_mat_angle"}: - # env_mat doesn't store any value - pass - elif name == "emask": - value = PairExcludeMask(value.ntypes, value.exclude_types) - else: - pass - - return super().__setattr__(name, value) + pass @flax_module class RepFlowLayer(RepFlowLayerDP): - def __setattr__(self, name: str, value: Any) -> None: - if name in { - "node_self_mlp", - "node_sym_linear", - "node_edge_linear", - "edge_self_linear", - "a_compress_n_linear", - "a_compress_e_linear", - "edge_angle_linear1", - "edge_angle_linear2", - "angle_self_linear", - }: - if value is not None: - value = NativeLayer.deserialize(value.serialize()) - elif Version(flax_version) >= Version("0.12.0"): - value = nnx.data(value) - elif name in {"n_residual", "e_residual", "a_residual"}: - value = [ArrayAPIVariable(to_jax_array(vv)) for vv in value] - if Version(flax_version) >= Version("0.12.0"): - value = nnx.List([nnx.data(item) for item in value]) - else: - pass - return super().__setattr__(name, value) + _jax_data_list_attrs: ClassVar[set[str]] = { + "n_residual", + "e_residual", + "a_residual", + } diff --git a/deepmd/jax/descriptor/repformers.py b/deepmd/jax/descriptor/repformers.py index 5701677349..24c9ee6a90 100644 --- a/deepmd/jax/descriptor/repformers.py +++ b/deepmd/jax/descriptor/repformers.py @@ -1,12 +1,10 @@ # SPDX-License-Identifier: LGPL-3.0-or-later from typing import ( - Any, -) - -from packaging.version import ( - Version, + ClassVar, ) +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.descriptor.repformers import ( Atten2EquiVarApply as Atten2EquiVarApplyDP, ) @@ -20,114 +18,39 @@ from deepmd.dpmodel.descriptor.repformers import LocalAtten as LocalAttenDP from deepmd.dpmodel.descriptor.repformers import RepformerLayer as RepformerLayerDP from deepmd.jax.common import ( - ArrayAPIVariable, flax_module, - to_jax_array, -) -from deepmd.jax.env import ( - flax_version, - nnx, -) -from deepmd.jax.utils.exclude_mask import ( - PairExcludeMask, -) -from deepmd.jax.utils.network import ( - LayerNorm, - NativeLayer, ) @flax_module class DescrptBlockRepformers(DescrptBlockRepformersDP): - def __setattr__(self, name: str, value: Any) -> None: - if name in {"mean", "stddev"}: - value = to_jax_array(value) - if value is not None: - value = ArrayAPIVariable(value) - elif Version(flax_version) >= Version("0.12.0"): - value = nnx.data(value) - elif name in {"layers"}: - value = [RepformerLayer.deserialize(layer.serialize()) for layer in value] - if Version(flax_version) >= Version("0.12.0"): - value = nnx.List([nnx.data(item) for item in value]) - elif name == "g2_embd": - value = NativeLayer.deserialize(value.serialize()) - elif name == "env_mat": - # env_mat doesn't store any value - pass - elif name == "emask": - value = PairExcludeMask(value.ntypes, value.exclude_types) - - return super().__setattr__(name, value) + pass @flax_module class Atten2Map(Atten2MapDP): - def __setattr__(self, name: str, value: Any) -> None: - if name in {"mapqk"}: - value = NativeLayer.deserialize(value.serialize()) - return super().__setattr__(name, value) + pass @flax_module class Atten2MultiHeadApply(Atten2MultiHeadApplyDP): - def __setattr__(self, name: str, value: Any) -> None: - if name in {"mapv", "head_map"}: - value = NativeLayer.deserialize(value.serialize()) - return super().__setattr__(name, value) + pass @flax_module class Atten2EquiVarApply(Atten2EquiVarApplyDP): - def __setattr__(self, name: str, value: Any) -> None: - if name in {"head_map"}: - value = NativeLayer.deserialize(value.serialize()) - return super().__setattr__(name, value) + pass @flax_module class LocalAtten(LocalAttenDP): - def __setattr__(self, name: str, value: Any) -> None: - if name in {"mapq", "mapkv", "head_map"}: - value = NativeLayer.deserialize(value.serialize()) - return super().__setattr__(name, value) + pass @flax_module class RepformerLayer(RepformerLayerDP): - def __setattr__(self, name: str, value: Any) -> None: - if name in {"linear1", "linear2", "g1_self_mlp", "proj_g1g2", "proj_g1g1g2"}: - if value is not None: - value = NativeLayer.deserialize(value.serialize()) - elif Version(flax_version) >= Version("0.12.0"): - value = nnx.data(value) - elif name in {"g1_residual", "g2_residual", "h2_residual"}: - value = [ArrayAPIVariable(to_jax_array(vv)) for vv in value] - if Version(flax_version) >= Version("0.12.0"): - value = nnx.List([nnx.data(item) for item in value]) - elif name in {"attn2g_map"}: - if value is not None: - value = Atten2Map.deserialize(value.serialize()) - elif Version(flax_version) >= Version("0.12.0"): - value = nnx.data(value) - elif name in {"attn2_mh_apply"}: - if value is not None: - value = Atten2MultiHeadApply.deserialize(value.serialize()) - elif Version(flax_version) >= Version("0.12.0"): - value = nnx.data(value) - elif name in {"attn2_lm"}: - if value is not None: - value = LayerNorm.deserialize(value.serialize()) - elif Version(flax_version) >= Version("0.12.0"): - value = nnx.data(value) - elif name in {"attn2_ev_apply"}: - if value is not None: - value = Atten2EquiVarApply.deserialize(value.serialize()) - elif Version(flax_version) >= Version("0.12.0"): - value = nnx.data(value) - elif name in {"loc_attn"}: - if value is not None: - value = LocalAtten.deserialize(value.serialize()) - elif Version(flax_version) >= Version("0.12.0"): - value = nnx.data(value) - return super().__setattr__(name, value) + _jax_data_list_attrs: ClassVar[set[str]] = { + "g1_residual", + "g2_residual", + "h2_residual", + } diff --git a/deepmd/jax/descriptor/se_atten_v2.py b/deepmd/jax/descriptor/se_atten_v2.py index a7ef4035cd..0e682d7af9 100644 --- a/deepmd/jax/descriptor/se_atten_v2.py +++ b/deepmd/jax/descriptor/se_atten_v2.py @@ -1,5 +1,8 @@ # SPDX-License-Identifier: LGPL-3.0-or-later from deepmd.dpmodel.descriptor.se_atten_v2 import DescrptSeAttenV2 as DescrptSeAttenV2DP +from deepmd.jax.common import ( + register_dpmodel_mapping, +) from deepmd.jax.descriptor.base_descriptor import ( BaseDescriptor, ) @@ -11,3 +14,9 @@ @BaseDescriptor.register("se_atten_v2") class DescrptSeAttenV2(DescrptDPA1, DescrptSeAttenV2DP): pass + + +register_dpmodel_mapping( + DescrptSeAttenV2DP, + lambda v: DescrptSeAttenV2.deserialize(v.serialize()), +) diff --git a/deepmd/jax/descriptor/se_e2_a.py b/deepmd/jax/descriptor/se_e2_a.py index 4d704a4b30..ec55ecf2e3 100644 --- a/deepmd/jax/descriptor/se_e2_a.py +++ b/deepmd/jax/descriptor/se_e2_a.py @@ -1,53 +1,17 @@ # SPDX-License-Identifier: LGPL-3.0-or-later -from typing import ( - Any, -) - -from packaging.version import ( - Version, -) - +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.descriptor.se_e2_a import DescrptSeAArrayAPI as DescrptSeADP from deepmd.jax.common import ( - ArrayAPIVariable, flax_module, - to_jax_array, ) from deepmd.jax.descriptor.base_descriptor import ( BaseDescriptor, ) -from deepmd.jax.env import ( - flax_version, - nnx, -) -from deepmd.jax.utils.exclude_mask import ( - PairExcludeMask, -) -from deepmd.jax.utils.network import ( - NetworkCollection, -) @BaseDescriptor.register("se_e2_a") @BaseDescriptor.register("se_a") @flax_module class DescrptSeA(DescrptSeADP): - def __setattr__(self, name: str, value: Any) -> None: - if name in {"dstd", "davg"}: - value = to_jax_array(value) - if value is not None: - value = ArrayAPIVariable(value) - elif Version(flax_version) >= Version("0.12.0"): - value = nnx.data(value) - elif name in {"embeddings"}: - if value is not None: - value = NetworkCollection.deserialize(value.serialize()) - elif Version(flax_version) >= Version("0.12.0"): - value = nnx.data(value) - elif name == "env_mat": - # env_mat doesn't store any value - pass - elif name == "emask": - value = PairExcludeMask(value.ntypes, value.exclude_types) - - return super().__setattr__(name, value) + pass diff --git a/deepmd/jax/descriptor/se_e2_r.py b/deepmd/jax/descriptor/se_e2_r.py index e5827c42af..a298d80503 100644 --- a/deepmd/jax/descriptor/se_e2_r.py +++ b/deepmd/jax/descriptor/se_e2_r.py @@ -1,53 +1,17 @@ # SPDX-License-Identifier: LGPL-3.0-or-later -from typing import ( - Any, -) - -from packaging.version import ( - Version, -) - +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.descriptor.se_r import DescrptSeR as DescrptSeRDP from deepmd.jax.common import ( - ArrayAPIVariable, flax_module, - to_jax_array, ) from deepmd.jax.descriptor.base_descriptor import ( BaseDescriptor, ) -from deepmd.jax.env import ( - flax_version, - nnx, -) -from deepmd.jax.utils.exclude_mask import ( - PairExcludeMask, -) -from deepmd.jax.utils.network import ( - NetworkCollection, -) @BaseDescriptor.register("se_e2_r") @BaseDescriptor.register("se_r") @flax_module class DescrptSeR(DescrptSeRDP): - def __setattr__(self, name: str, value: Any) -> None: - if name in {"dstd", "davg"}: - value = to_jax_array(value) - if value is not None: - value = ArrayAPIVariable(value) - elif Version(flax_version) >= Version("0.12.0"): - value = nnx.data(value) - elif name in {"embeddings"}: - if value is not None: - value = NetworkCollection.deserialize(value.serialize()) - elif Version(flax_version) >= Version("0.12.0"): - value = nnx.data(value) - elif name == "env_mat": - # env_mat doesn't store any value - pass - elif name == "emask": - value = PairExcludeMask(value.ntypes, value.exclude_types) - - return super().__setattr__(name, value) + pass diff --git a/deepmd/jax/descriptor/se_t.py b/deepmd/jax/descriptor/se_t.py index 6d0b026c94..2af8680b6d 100644 --- a/deepmd/jax/descriptor/se_t.py +++ b/deepmd/jax/descriptor/se_t.py @@ -1,31 +1,13 @@ # SPDX-License-Identifier: LGPL-3.0-or-later -from typing import ( - Any, -) - -from packaging.version import ( - Version, -) - +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.descriptor.se_t import DescrptSeT as DescrptSeTDP from deepmd.jax.common import ( - ArrayAPIVariable, flax_module, - to_jax_array, ) from deepmd.jax.descriptor.base_descriptor import ( BaseDescriptor, ) -from deepmd.jax.env import ( - flax_version, - nnx, -) -from deepmd.jax.utils.exclude_mask import ( - PairExcludeMask, -) -from deepmd.jax.utils.network import ( - NetworkCollection, -) @BaseDescriptor.register("se_e3") @@ -33,20 +15,4 @@ @BaseDescriptor.register("se_a_3be") @flax_module class DescrptSeT(DescrptSeTDP): - def __setattr__(self, name: str, value: Any) -> None: - if name in {"dstd", "davg"}: - value = to_jax_array(value) - if value is not None: - value = ArrayAPIVariable(value) - elif Version(flax_version) >= Version("0.12.0"): - value = nnx.data(value) - elif name in {"embeddings"}: - if value is not None: - value = NetworkCollection.deserialize(value.serialize()) - elif name == "env_mat": - # env_mat doesn't store any value - pass - elif name == "emask": - value = PairExcludeMask(value.ntypes, value.exclude_types) - - return super().__setattr__(name, value) + pass diff --git a/deepmd/jax/descriptor/se_t_tebd.py b/deepmd/jax/descriptor/se_t_tebd.py index 8e2dae782a..ce9d6a84e9 100644 --- a/deepmd/jax/descriptor/se_t_tebd.py +++ b/deepmd/jax/descriptor/se_t_tebd.py @@ -1,66 +1,25 @@ # SPDX-License-Identifier: LGPL-3.0-or-later -from typing import ( - Any, -) - -from packaging.version import ( - Version, -) - +import deepmd.jax.utils.exclude_mask as _jax_exclude_mask # noqa: F401 +import deepmd.jax.utils.network as _jax_network # noqa: F401 +import deepmd.jax.utils.type_embed as _jax_type_embed # noqa: F401 from deepmd.dpmodel.descriptor.se_t_tebd import ( DescrptBlockSeTTebd as DescrptBlockSeTTebdDP, ) from deepmd.dpmodel.descriptor.se_t_tebd import DescrptSeTTebd as DescrptSeTTebdDP from deepmd.jax.common import ( - ArrayAPIVariable, flax_module, - to_jax_array, ) from deepmd.jax.descriptor.base_descriptor import ( BaseDescriptor, ) -from deepmd.jax.env import ( - flax_version, - nnx, -) -from deepmd.jax.utils.exclude_mask import ( - PairExcludeMask, -) -from deepmd.jax.utils.network import ( - NetworkCollection, -) -from deepmd.jax.utils.type_embed import ( - TypeEmbedNet, -) @flax_module class DescrptBlockSeTTebd(DescrptBlockSeTTebdDP): - def __setattr__(self, name: str, value: Any) -> None: - if name in {"mean", "stddev"}: - value = to_jax_array(value) - if value is not None: - value = ArrayAPIVariable(value) - elif Version(flax_version) >= Version("0.12.0"): - value = nnx.data(value) - elif name in {"embeddings", "embeddings_strip"}: - if value is not None: - value = NetworkCollection.deserialize(value.serialize()) - elif name == "env_mat": - # env_mat doesn't store any value - pass - elif name == "emask": - value = PairExcludeMask(value.ntypes, value.exclude_types) - - return super().__setattr__(name, value) + pass @BaseDescriptor.register("se_e3_tebd") @flax_module class DescrptSeTTebd(DescrptSeTTebdDP): - def __setattr__(self, name: str, value: Any) -> None: - if name == "se_ttebd": - value = DescrptBlockSeTTebd.deserialize(value.serialize()) - elif name == "type_embedding": - value = TypeEmbedNet.deserialize(value.serialize()) - return super().__setattr__(name, value) + pass diff --git a/deepmd/jax/fitting/fitting.py b/deepmd/jax/fitting/fitting.py index 8fea40cd57..e73d30c715 100644 --- a/deepmd/jax/fitting/fitting.py +++ b/deepmd/jax/fitting/fitting.py @@ -1,12 +1,6 @@ # SPDX-License-Identifier: LGPL-3.0-or-later -from typing import ( - Any, -) - -from packaging.version import ( - Version, -) - +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.fitting.dipole_fitting import DipoleFitting as DipoleFittingNetDP from deepmd.dpmodel.fitting.dos_fitting import DOSFittingNet as DOSFittingNetDP from deepmd.dpmodel.fitting.ener_fitting import EnergyFittingNet as EnergyFittingNetDP @@ -17,91 +11,38 @@ PropertyFittingNet as PropertyFittingNetDP, ) from deepmd.jax.common import ( - ArrayAPIVariable, flax_module, - to_jax_array, -) -from deepmd.jax.env import ( - flax_version, - nnx, ) from deepmd.jax.fitting.base_fitting import ( BaseFitting, ) -from deepmd.jax.utils.exclude_mask import ( - AtomExcludeMask, -) -from deepmd.jax.utils.network import ( - NetworkCollection, -) - - -def setattr_for_general_fitting(name: str, value: Any) -> Any: - if name in { - "bias_atom_e", - "fparam_avg", - "fparam_inv_std", - "aparam_avg", - "aparam_inv_std", - "case_embd", - "default_fparam_tensor", - }: - value = to_jax_array(value) - if value is not None: - value = ArrayAPIVariable(value) - elif Version(flax_version) >= Version("0.12.0"): - value = nnx.data(value) - elif name == "emask": - value = AtomExcludeMask(value.ntypes, value.exclude_types) - elif name == "nets": - value = NetworkCollection.deserialize(value.serialize()) - return value @BaseFitting.register("ener") @flax_module class EnergyFittingNet(EnergyFittingNetDP): - def __setattr__(self, name: str, value: Any) -> None: - value = setattr_for_general_fitting(name, value) - return super().__setattr__(name, value) + pass @BaseFitting.register("property") @flax_module class PropertyFittingNet(PropertyFittingNetDP): - def __setattr__(self, name: str, value: Any) -> None: - value = setattr_for_general_fitting(name, value) - return super().__setattr__(name, value) + pass @BaseFitting.register("dos") @flax_module class DOSFittingNet(DOSFittingNetDP): - def __setattr__(self, name: str, value: Any) -> None: - value = setattr_for_general_fitting(name, value) - return super().__setattr__(name, value) + pass @BaseFitting.register("dipole") @flax_module class DipoleFittingNet(DipoleFittingNetDP): - def __setattr__(self, name: str, value: Any) -> None: - value = setattr_for_general_fitting(name, value) - return super().__setattr__(name, value) + pass @BaseFitting.register("polar") @flax_module class PolarFittingNet(PolarFittingNetDP): - def __setattr__(self, name: str, value: Any) -> None: - value = setattr_for_general_fitting(name, value) - if name in { - "scale", - "constant_matrix", - }: - value = to_jax_array(value) - if value is not None: - value = ArrayAPIVariable(value) - elif Version(flax_version) >= Version("0.12.0"): - value = nnx.data(value) - return super().__setattr__(name, value) + pass diff --git a/deepmd/jax/model/__init__.py b/deepmd/jax/model/__init__.py index 79d5bb2b23..08c0a4e8e7 100644 --- a/deepmd/jax/model/__init__.py +++ b/deepmd/jax/model/__init__.py @@ -1,13 +1,14 @@ # SPDX-License-Identifier: LGPL-3.0-or-later +from deepmd.jax.atomic_model.linear_atomic_model import ( + DPZBLLinearEnergyAtomicModel, +) + from .dipole_model import ( DipoleModel, ) from .dos_model import ( DOSModel, ) -from .dp_zbl_model import ( - DPZBLLinearEnergyAtomicModel, -) from .ener_model import ( EnergyModel, ) diff --git a/deepmd/jax/model/dp_model.py b/deepmd/jax/model/dp_model.py index a265229e0e..67aebd1b26 100644 --- a/deepmd/jax/model/dp_model.py +++ b/deepmd/jax/model/dp_model.py @@ -1,8 +1,4 @@ # SPDX-License-Identifier: LGPL-3.0-or-later -from typing import ( - Any, -) - from deepmd.dpmodel.model import ( DPModelCommon, ) @@ -41,11 +37,6 @@ def make_jax_dp_model_from_dpmodel( @flax_module class jax_model(dpmodel_model): - def __setattr__(self, name: str, value: Any) -> None: - if name == "atomic_model": - value = jax_atomicmodel.deserialize(value.serialize()) - return super().__setattr__(name, value) - def forward_common_atomic( self, extended_coord: jnp.ndarray, diff --git a/deepmd/jax/model/dp_zbl_model.py b/deepmd/jax/model/dp_zbl_model.py index 8d66084b8e..3a6dd6c3bb 100644 --- a/deepmd/jax/model/dp_zbl_model.py +++ b/deepmd/jax/model/dp_zbl_model.py @@ -1,11 +1,7 @@ # SPDX-License-Identifier: LGPL-3.0-or-later -from typing import ( - Any, -) - from deepmd.dpmodel.model.dp_zbl_model import DPZBLModel as DPZBLModelDP -from deepmd.jax.atomic_model.linear_atomic_model import ( - DPZBLLinearEnergyAtomicModel, +from deepmd.jax.atomic_model.linear_atomic_model import ( # noqa: F401 + DPZBLLinearEnergyAtomicModel as _DPZBLLinearEnergyAtomicModel, ) from deepmd.jax.common import ( flax_module, @@ -23,11 +19,6 @@ @BaseModel.register("zbl") @flax_module class DPZBLModel(DPZBLModelDP): - def __setattr__(self, name: str, value: Any) -> None: - if name == "atomic_model": - value = DPZBLLinearEnergyAtomicModel.deserialize(value.serialize()) - return super().__setattr__(name, value) - def forward_common_atomic( self, extended_coord: jnp.ndarray, diff --git a/deepmd/jax/utils/exclude_mask.py b/deepmd/jax/utils/exclude_mask.py index 4ae230c8dc..1604a5289e 100644 --- a/deepmd/jax/utils/exclude_mask.py +++ b/deepmd/jax/utils/exclude_mask.py @@ -1,44 +1,28 @@ # SPDX-License-Identifier: LGPL-3.0-or-later -from typing import ( - Any, -) - -from packaging.version import ( - Version, -) - from deepmd.dpmodel.utils.exclude_mask import AtomExcludeMask as AtomExcludeMaskDP from deepmd.dpmodel.utils.exclude_mask import PairExcludeMask as PairExcludeMaskDP from deepmd.jax.common import ( - ArrayAPIVariable, flax_module, - to_jax_array, -) -from deepmd.jax.env import ( - flax_version, - nnx, + register_dpmodel_mapping, ) @flax_module class AtomExcludeMask(AtomExcludeMaskDP): - def __setattr__(self, name: str, value: Any) -> None: - if name in {"type_mask"}: - value = to_jax_array(value) - if value is not None: - value = ArrayAPIVariable(value) - elif Version(flax_version) >= Version("0.12.0"): - value = nnx.data(value) - return super().__setattr__(name, value) + pass @flax_module class PairExcludeMask(PairExcludeMaskDP): - def __setattr__(self, name: str, value: Any) -> None: - if name in {"type_mask"}: - value = to_jax_array(value) - if value is not None: - value = ArrayAPIVariable(value) - elif Version(flax_version) >= Version("0.12.0"): - value = nnx.data(value) - return super().__setattr__(name, value) + pass + + +register_dpmodel_mapping( + AtomExcludeMaskDP, + lambda v: AtomExcludeMask(v.ntypes, exclude_types=list(v.get_exclude_types())), +) + +register_dpmodel_mapping( + PairExcludeMaskDP, + lambda v: PairExcludeMask(v.ntypes, exclude_types=list(v.get_exclude_types())), +) diff --git a/deepmd/jax/utils/network.py b/deepmd/jax/utils/network.py index 72d9f760eb..48d8df07d9 100644 --- a/deepmd/jax/utils/network.py +++ b/deepmd/jax/utils/network.py @@ -5,15 +5,16 @@ ) import numpy as np -from packaging.version import ( - Version, -) from deepmd.dpmodel.common import ( NativeOP, ) +from deepmd.dpmodel.utils.network import EmbeddingNet as EmbeddingNetDP +from deepmd.dpmodel.utils.network import FittingNet as FittingNetDP +from deepmd.dpmodel.utils.network import Identity as IdentityDP from deepmd.dpmodel.utils.network import LayerNorm as LayerNormDP from deepmd.dpmodel.utils.network import NativeLayer as NativeLayerDP +from deepmd.dpmodel.utils.network import NativeNet as NativeNetDP from deepmd.dpmodel.utils.network import NetworkCollection as NetworkCollectionDP from deepmd.dpmodel.utils.network import ( make_embedding_network, @@ -23,10 +24,10 @@ from deepmd.jax.common import ( ArrayAPIVariable, flax_module, + register_dpmodel_mapping, to_jax_array, ) from deepmd.jax.env import ( - flax_version, nnx, ) @@ -47,6 +48,8 @@ def __dlpack_device__(self, *args: Any, **kwargs: Any) -> Any: @flax_module class NativeLayer(NativeLayerDP): + _jax_skip_auto_convert_attrs: ClassVar[set[str]] = {"w", "b", "idt"} + def __setattr__(self, name: str, value: Any) -> None: if name in {"w", "b", "idt"}: value = to_jax_array(value) @@ -60,10 +63,7 @@ def __setattr__(self, name: str, value: Any) -> None: @flax_module class NativeNet(make_multilayer_network(NativeLayer, NativeOP)): - def __setattr__(self, name: str, value: Any) -> None: - if name in {"layers"} and Version(flax_version) >= Version("0.12.0"): - value = nnx.List(value) - return super().__setattr__(name, value) + pass class EmbeddingNet(make_embedding_network(NativeNet, NativeLayer)): @@ -76,17 +76,55 @@ class FittingNet(make_fitting_network(EmbeddingNet, NativeNet, NativeLayer)): @flax_module class NetworkCollection(NetworkCollectionDP): + _jax_data_list_attrs: ClassVar[set[str]] = {"_networks"} + NETWORK_TYPE_MAP: ClassVar[dict[str, type]] = { "network": NativeNet, "embedding_network": EmbeddingNet, "fitting_network": FittingNet, } - def __setattr__(self, name: str, value: Any) -> None: - if name in {"_networks"} and Version(flax_version) >= Version("0.12.0"): - value = nnx.List([nnx.data(item) for item in value]) - return super().__setattr__(name, value) - class LayerNorm(LayerNormDP, NativeLayer): pass + + +@flax_module +class Identity(IdentityDP): + pass + + +register_dpmodel_mapping( + NativeNetDP, + lambda v: NativeNet.deserialize(v.serialize()), +) + +register_dpmodel_mapping( + EmbeddingNetDP, + lambda v: EmbeddingNet.deserialize(v.serialize()), +) + +register_dpmodel_mapping( + FittingNetDP, + lambda v: FittingNet.deserialize(v.serialize()), +) + +register_dpmodel_mapping( + NativeLayerDP, + lambda v: NativeLayer.deserialize(v.serialize()), +) + +register_dpmodel_mapping( + LayerNormDP, + lambda v: LayerNorm.deserialize(v.serialize()), +) + +register_dpmodel_mapping( + NetworkCollectionDP, + lambda v: NetworkCollection.deserialize(v.serialize()), +) + +register_dpmodel_mapping( + IdentityDP, + lambda v: Identity(), +) diff --git a/deepmd/jax/utils/type_embed.py b/deepmd/jax/utils/type_embed.py index aff0a78a2c..eead05f978 100644 --- a/deepmd/jax/utils/type_embed.py +++ b/deepmd/jax/utils/type_embed.py @@ -1,36 +1,11 @@ # SPDX-License-Identifier: LGPL-3.0-or-later -from typing import ( - Any, -) - -from packaging.version import ( - Version, -) - +import deepmd.jax.utils.network as _jax_network # noqa: F401 from deepmd.dpmodel.utils.type_embed import TypeEmbedNet as TypeEmbedNetDP from deepmd.jax.common import ( - ArrayAPIVariable, flax_module, - to_jax_array, -) -from deepmd.jax.env import ( - flax_version, - nnx, -) -from deepmd.jax.utils.network import ( - EmbeddingNet, ) @flax_module class TypeEmbedNet(TypeEmbedNetDP): - def __setattr__(self, name: str, value: Any) -> None: - if name in {"econf_tebd"}: - value = to_jax_array(value) - if value is not None: - value = ArrayAPIVariable(value) - elif Version(flax_version) >= Version("0.12.0"): - value = nnx.data(value) - if name in {"embedding_net"}: - value = EmbeddingNet.deserialize(value.serialize()) - return super().__setattr__(name, value) + pass