Skip to content

Commit 7900e78

Browse files
authored
Merge branch 'deepmodeling:master' into master
2 parents 7b8bd61 + beb50da commit 7900e78

45 files changed

Lines changed: 1098 additions & 1215 deletions

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.
Lines changed: 0 additions & 34 deletions
Original file line numberDiff line numberDiff line change
@@ -1,35 +1 @@
11
# SPDX-License-Identifier: LGPL-3.0-or-later
2-
from typing import (
3-
Any,
4-
)
5-
6-
from packaging.version import (
7-
Version,
8-
)
9-
10-
from deepmd.jax.common import (
11-
ArrayAPIVariable,
12-
to_jax_array,
13-
)
14-
from deepmd.jax.env import (
15-
flax_version,
16-
nnx,
17-
)
18-
from deepmd.jax.utils.exclude_mask import (
19-
AtomExcludeMask,
20-
PairExcludeMask,
21-
)
22-
23-
24-
def base_atomic_model_set_attr(name: str, value: Any) -> Any:
25-
if name in {"out_bias", "out_std"}:
26-
value = to_jax_array(value)
27-
if value is not None:
28-
value = ArrayAPIVariable(value)
29-
elif Version(flax_version) >= Version("0.12.0"):
30-
value = nnx.data(value)
31-
elif name == "pair_excl" and value is not None:
32-
value = PairExcludeMask(value.ntypes, value.exclude_types)
33-
elif name == "atom_excl" and value is not None:
34-
value = AtomExcludeMask(value.ntypes, value.exclude_types)
35-
return value

deepmd/jax/atomic_model/dp_atomic_model.py

Lines changed: 3 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -1,12 +1,8 @@
11
# SPDX-License-Identifier: LGPL-3.0-or-later
2-
from typing import (
3-
Any,
4-
)
5-
2+
import deepmd.jax.descriptor as _jax_descriptor # noqa: F401
3+
import deepmd.jax.fitting.fitting as _jax_fitting # noqa: F401
4+
import deepmd.jax.utils.exclude_mask as _jax_exclude_mask # noqa: F401
65
from deepmd.dpmodel.atomic_model.dp_atomic_model import DPAtomicModel as DPAtomicModelDP
7-
from deepmd.jax.atomic_model.base_atomic_model import (
8-
base_atomic_model_set_attr,
9-
)
106
from deepmd.jax.common import (
117
flax_module,
128
)
@@ -45,10 +41,6 @@ class jax_atomic_model(dpmodel_atomic_model):
4541
base_fitting_cls = BaseFitting
4642
"""The base fitting class."""
4743

48-
def __setattr__(self, name: str, value: Any) -> None:
49-
value = base_atomic_model_set_attr(name, value)
50-
return super().__setattr__(name, value)
51-
5244
def forward_common_atomic(
5345
self,
5446
extended_coord: jnp.ndarray,

deepmd/jax/atomic_model/linear_atomic_model.py

Lines changed: 4 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -3,54 +3,28 @@
33
Any,
44
)
55

6-
from packaging.version import (
7-
Version,
8-
)
9-
6+
import deepmd.jax.atomic_model.dp_atomic_model as _jax_dp_atomic_model # noqa: F401
7+
import deepmd.jax.atomic_model.pairtab_atomic_model as _jax_pairtab_model # noqa: F401
8+
import deepmd.jax.utils.exclude_mask as _jax_exclude_mask # noqa: F401
109
from deepmd.dpmodel.atomic_model.linear_atomic_model import (
1110
DPZBLLinearEnergyAtomicModel as DPZBLLinearEnergyAtomicModelDP,
1211
)
13-
from deepmd.jax.atomic_model.base_atomic_model import (
14-
base_atomic_model_set_attr,
15-
)
16-
from deepmd.jax.atomic_model.dp_atomic_model import (
17-
DPAtomicModel,
18-
)
19-
from deepmd.jax.atomic_model.pairtab_atomic_model import (
20-
PairTabAtomicModel,
21-
)
2212
from deepmd.jax.common import (
23-
ArrayAPIVariable,
2413
flax_module,
25-
to_jax_array,
2614
)
2715
from deepmd.jax.env import (
28-
flax_version,
2916
jax,
3017
jnp,
31-
nnx,
3218
)
3319

3420

3521
@flax_module
3622
class DPZBLLinearEnergyAtomicModel(DPZBLLinearEnergyAtomicModelDP):
3723
def __setattr__(self, name: str, value: Any) -> None:
38-
value = base_atomic_model_set_attr(name, value)
39-
if name == "mapping_list":
40-
value = [ArrayAPIVariable(to_jax_array(vv)) for vv in value]
41-
if Version(flax_version) >= Version("0.12.0"):
42-
value = nnx.List([nnx.data(item) for item in value])
43-
elif name == "zbl_weight":
24+
if name == "zbl_weight":
4425
# discard since it's only used in tests
4526
# to fix flax.errors.TraceContextError: Cannot mutate 'FlaxModule' from different trace level
4627
return
47-
elif name == "models":
48-
value = [
49-
DPAtomicModel.deserialize(value[0].serialize()),
50-
PairTabAtomicModel.deserialize(value[1].serialize()),
51-
]
52-
if Version(flax_version) >= Version("0.12.0"):
53-
value = nnx.List([nnx.data(item) for item in value])
5428
return super().__setattr__(name, value)
5529

5630
def forward_common_atomic(

deepmd/jax/atomic_model/pairtab_atomic_model.py

Lines changed: 1 addition & 25 deletions
Original file line numberDiff line numberDiff line change
@@ -1,43 +1,19 @@
11
# SPDX-License-Identifier: LGPL-3.0-or-later
2-
from typing import (
3-
Any,
4-
)
5-
6-
from packaging.version import (
7-
Version,
8-
)
9-
2+
import deepmd.jax.utils.exclude_mask as _jax_exclude_mask # noqa: F401
103
from deepmd.dpmodel.atomic_model.pairtab_atomic_model import (
114
PairTabAtomicModel as PairTabAtomicModelDP,
125
)
13-
from deepmd.jax.atomic_model.base_atomic_model import (
14-
base_atomic_model_set_attr,
15-
)
166
from deepmd.jax.common import (
17-
ArrayAPIVariable,
187
flax_module,
19-
to_jax_array,
208
)
219
from deepmd.jax.env import (
22-
flax_version,
2310
jax,
2411
jnp,
25-
nnx,
2612
)
2713

2814

2915
@flax_module
3016
class PairTabAtomicModel(PairTabAtomicModelDP):
31-
def __setattr__(self, name: str, value: Any) -> None:
32-
value = base_atomic_model_set_attr(name, value)
33-
if name in {"tab_info", "tab_data"}:
34-
value = to_jax_array(value)
35-
if value is not None:
36-
value = ArrayAPIVariable(value)
37-
elif Version(flax_version) >= Version("0.12.0"):
38-
value = nnx.data(value)
39-
return super().__setattr__(name, value)
40-
4117
def forward_common_atomic(
4218
self,
4319
extended_coord: jnp.ndarray,

0 commit comments

Comments
 (0)