Skip to content

Commit 2901448

Browse files
njzjzCopilotgemini-code-assist[bot]
authored
refact(pt_expt): add decorator to simplify the module (#5213)
<!-- This is an auto-generated comment: release notes by coderabbit.ai --> ## Summary by CodeRabbit * **New Features** * Introduced a public adapter to expose DP classes as PyTorch modules via a decorator-based API. * **Refactor** * Converted multiple descriptor, network, and utility classes to decorator-driven PyTorch integration, simplifying initialization and attribute handling. * **Breaking Changes** * Several descriptor forward signatures expanded to accept extended topology/embedding inputs and now return expanded tuples (update call sites accordingly). <!-- end of auto-generated comment: release notes by coderabbit.ai --> --------- Signed-off-by: Jinzhe Zeng <njzjz@qq.com> Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
1 parent 156736f commit 2901448

8 files changed

Lines changed: 94 additions & 115 deletions

File tree

deepmd/pt_expt/common.py

Lines changed: 43 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,9 @@
1717
from collections.abc import (
1818
Callable,
1919
)
20+
from functools import (
21+
wraps,
22+
)
2023
from typing import (
2124
Any,
2225
overload,
@@ -292,6 +295,46 @@ def to_torch_array(array: Any) -> torch.Tensor | None:
292295
return torch.as_tensor(array, device=env.DEVICE)
293296

294297

298+
def torch_module(
299+
module: type[NativeOP],
300+
) -> type[torch.nn.Module]:
301+
"""Convert a NativeOP to a torch.nn.Module.
302+
303+
Parameters
304+
----------
305+
module : type[NativeOP]
306+
The NativeOP to convert.
307+
308+
Returns
309+
-------
310+
type[torch.nn.Module]
311+
The torch.nn.Module.
312+
313+
Examples
314+
--------
315+
>>> @torch_module
316+
... class MyModule(NativeOP):
317+
... pass
318+
"""
319+
320+
@wraps(module, updated=())
321+
class TorchModule(module, torch.nn.Module):
322+
def __init__(self, *args: Any, **kwargs: Any) -> None:
323+
torch.nn.Module.__init__(self)
324+
module.__init__(self, *args, **kwargs)
325+
326+
def __call__(self, *args: Any, **kwargs: Any) -> Any:
327+
# Ensure torch.nn.Module.__call__ drives forward() for export/tracing.
328+
return torch.nn.Module.__call__(self, *args, **kwargs)
329+
330+
def __setattr__(self, name: str, value: Any) -> None:
331+
handled, value = dpmodel_setattr(self, name, value)
332+
if not handled:
333+
super().__setattr__(name, value)
334+
335+
return TorchModule
336+
337+
295338
# Import utils to trigger dpmodel→pt_expt converter registrations
296339
# This must happen after the functions above are defined to avoid circular imports
297340
def _ensure_registrations() -> None:

deepmd/pt_expt/descriptor/se_e2_a.py

Lines changed: 3 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -1,13 +1,10 @@
11
# SPDX-License-Identifier: LGPL-3.0-or-later
2-
from typing import (
3-
Any,
4-
)
52

63
import torch
74

85
from deepmd.dpmodel.descriptor.se_e2_a import DescrptSeAArrayAPI as DescrptSeADP
96
from deepmd.pt_expt.common import (
10-
dpmodel_setattr,
7+
torch_module,
118
)
129
from deepmd.pt_expt.descriptor.base_descriptor import (
1310
BaseDescriptor,
@@ -16,20 +13,8 @@
1613

1714
@BaseDescriptor.register("se_e2_a_expt")
1815
@BaseDescriptor.register("se_a_expt")
19-
class DescrptSeA(DescrptSeADP, torch.nn.Module):
20-
def __init__(self, *args: Any, **kwargs: Any) -> None:
21-
torch.nn.Module.__init__(self)
22-
DescrptSeADP.__init__(self, *args, **kwargs)
23-
24-
def __call__(self, *args: Any, **kwargs: Any) -> Any:
25-
# Ensure torch.nn.Module.__call__ drives forward() for export/tracing.
26-
return torch.nn.Module.__call__(self, *args, **kwargs)
27-
28-
def __setattr__(self, name: str, value: Any) -> None:
29-
handled, value = dpmodel_setattr(self, name, value)
30-
if not handled:
31-
super().__setattr__(name, value)
32-
16+
@torch_module
17+
class DescrptSeA(DescrptSeADP):
3318
def forward(
3419
self,
3520
extended_coord: torch.Tensor,

deepmd/pt_expt/descriptor/se_r.py

Lines changed: 3 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -1,13 +1,10 @@
11
# SPDX-License-Identifier: LGPL-3.0-or-later
2-
from typing import (
3-
Any,
4-
)
52

63
import torch
74

85
from deepmd.dpmodel.descriptor.se_r import DescrptSeR as DescrptSeRDP
96
from deepmd.pt_expt.common import (
10-
dpmodel_setattr,
7+
torch_module,
118
)
129
from deepmd.pt_expt.descriptor.base_descriptor import (
1310
BaseDescriptor,
@@ -16,20 +13,8 @@
1613

1714
@BaseDescriptor.register("se_e2_r_expt")
1815
@BaseDescriptor.register("se_r_expt")
19-
class DescrptSeR(DescrptSeRDP, torch.nn.Module):
20-
def __init__(self, *args: Any, **kwargs: Any) -> None:
21-
torch.nn.Module.__init__(self)
22-
DescrptSeRDP.__init__(self, *args, **kwargs)
23-
24-
def __call__(self, *args: Any, **kwargs: Any) -> Any:
25-
# Ensure torch.nn.Module.__call__ drives forward() for export/tracing.
26-
return torch.nn.Module.__call__(self, *args, **kwargs)
27-
28-
def __setattr__(self, name: str, value: Any) -> None:
29-
handled, value = dpmodel_setattr(self, name, value)
30-
if not handled:
31-
super().__setattr__(name, value)
32-
16+
@torch_module
17+
class DescrptSeR(DescrptSeRDP):
3318
def forward(
3419
self,
3520
extended_coord: torch.Tensor,

deepmd/pt_expt/descriptor/se_t.py

Lines changed: 3 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -1,13 +1,10 @@
11
# SPDX-License-Identifier: LGPL-3.0-or-later
2-
from typing import (
3-
Any,
4-
)
52

63
import torch
74

85
from deepmd.dpmodel.descriptor.se_t import DescrptSeT as DescrptSeTDP
96
from deepmd.pt_expt.common import (
10-
dpmodel_setattr,
7+
torch_module,
118
)
129
from deepmd.pt_expt.descriptor.base_descriptor import (
1310
BaseDescriptor,
@@ -17,20 +14,8 @@
1714
@BaseDescriptor.register("se_e3_expt")
1815
@BaseDescriptor.register("se_at_expt")
1916
@BaseDescriptor.register("se_a_3be_expt")
20-
class DescrptSeT(DescrptSeTDP, torch.nn.Module):
21-
def __init__(self, *args: Any, **kwargs: Any) -> None:
22-
torch.nn.Module.__init__(self)
23-
DescrptSeTDP.__init__(self, *args, **kwargs)
24-
25-
def __call__(self, *args: Any, **kwargs: Any) -> Any:
26-
# Ensure torch.nn.Module.__call__ drives forward() for export/tracing.
27-
return torch.nn.Module.__call__(self, *args, **kwargs)
28-
29-
def __setattr__(self, name: str, value: Any) -> None:
30-
handled, value = dpmodel_setattr(self, name, value)
31-
if not handled:
32-
super().__setattr__(name, value)
33-
17+
@torch_module
18+
class DescrptSeT(DescrptSeTDP):
3419
def forward(
3520
self,
3621
extended_coord: torch.Tensor,

deepmd/pt_expt/descriptor/se_t_tebd.py

Lines changed: 3 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -1,34 +1,19 @@
11
# SPDX-License-Identifier: LGPL-3.0-or-later
2-
from typing import (
3-
Any,
4-
)
52

63
import torch
74

85
from deepmd.dpmodel.descriptor.se_t_tebd import DescrptSeTTebd as DescrptSeTTebdDP
96
from deepmd.pt_expt.common import (
10-
dpmodel_setattr,
7+
torch_module,
118
)
129
from deepmd.pt_expt.descriptor.base_descriptor import (
1310
BaseDescriptor,
1411
)
1512

1613

1714
@BaseDescriptor.register("se_e3_tebd_expt")
18-
class DescrptSeTTebd(DescrptSeTTebdDP, torch.nn.Module):
19-
def __init__(self, *args: Any, **kwargs: Any) -> None:
20-
torch.nn.Module.__init__(self)
21-
DescrptSeTTebdDP.__init__(self, *args, **kwargs)
22-
23-
def __call__(self, *args: Any, **kwargs: Any) -> Any:
24-
# Ensure torch.nn.Module.__call__ drives forward() for export/tracing.
25-
return torch.nn.Module.__call__(self, *args, **kwargs)
26-
27-
def __setattr__(self, name: str, value: Any) -> None:
28-
handled, value = dpmodel_setattr(self, name, value)
29-
if not handled:
30-
super().__setattr__(name, value)
31-
15+
@torch_module
16+
class DescrptSeTTebd(DescrptSeTTebdDP):
3217
def forward(
3318
self,
3419
extended_coord: torch.Tensor,

deepmd/pt_expt/descriptor/se_t_tebd_block.py

Lines changed: 26 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -1,28 +1,41 @@
11
# SPDX-License-Identifier: LGPL-3.0-or-later
2-
from typing import (
3-
Any,
4-
)
52

63
import torch
74

85
from deepmd.dpmodel.descriptor.se_t_tebd import (
96
DescrptBlockSeTTebd as DescrptBlockSeTTebdDP,
107
)
118
from deepmd.pt_expt.common import (
12-
dpmodel_setattr,
139
register_dpmodel_mapping,
10+
torch_module,
1411
)
1512

1613

17-
class DescrptBlockSeTTebd(DescrptBlockSeTTebdDP, torch.nn.Module):
18-
def __init__(self, *args: Any, **kwargs: Any) -> None:
19-
torch.nn.Module.__init__(self)
20-
DescrptBlockSeTTebdDP.__init__(self, *args, **kwargs)
21-
22-
def __setattr__(self, name: str, value: Any) -> None:
23-
handled, value = dpmodel_setattr(self, name, value)
24-
if not handled:
25-
super().__setattr__(name, value)
14+
@torch_module
15+
class DescrptBlockSeTTebd(DescrptBlockSeTTebdDP):
16+
def forward(
17+
self,
18+
nlist: torch.Tensor,
19+
coord_ext: torch.Tensor,
20+
atype_ext: torch.Tensor,
21+
atype_embd_ext: torch.Tensor | None = None,
22+
mapping: torch.Tensor | None = None,
23+
type_embedding: torch.Tensor | None = None,
24+
) -> tuple[
25+
torch.Tensor,
26+
torch.Tensor | None,
27+
torch.Tensor | None,
28+
torch.Tensor | None,
29+
torch.Tensor | None,
30+
]:
31+
return self.call(
32+
nlist,
33+
coord_ext,
34+
atype_ext,
35+
atype_embd_ext=atype_embd_ext,
36+
mapping=mapping,
37+
type_embedding=type_embedding,
38+
)
2639

2740

2841
register_dpmodel_mapping(

deepmd/pt_expt/utils/exclude_mask.py

Lines changed: 7 additions & 23 deletions
Original file line numberDiff line numberDiff line change
@@ -1,27 +1,17 @@
11
# SPDX-License-Identifier: LGPL-3.0-or-later
2-
from typing import (
3-
Any,
4-
)
52

6-
import torch
73

84
from deepmd.dpmodel.utils.exclude_mask import AtomExcludeMask as AtomExcludeMaskDP
95
from deepmd.dpmodel.utils.exclude_mask import PairExcludeMask as PairExcludeMaskDP
106
from deepmd.pt_expt.common import (
11-
dpmodel_setattr,
127
register_dpmodel_mapping,
8+
torch_module,
139
)
1410

1511

16-
class AtomExcludeMask(AtomExcludeMaskDP, torch.nn.Module):
17-
def __init__(self, *args: Any, **kwargs: Any) -> None:
18-
torch.nn.Module.__init__(self)
19-
AtomExcludeMaskDP.__init__(self, *args, **kwargs)
20-
21-
def __setattr__(self, name: str, value: Any) -> None:
22-
handled, value = dpmodel_setattr(self, name, value)
23-
if not handled:
24-
super().__setattr__(name, value)
12+
@torch_module
13+
class AtomExcludeMask(AtomExcludeMaskDP):
14+
pass
2515

2616

2717
register_dpmodel_mapping(
@@ -30,15 +20,9 @@ def __setattr__(self, name: str, value: Any) -> None:
3020
)
3121

3222

33-
class PairExcludeMask(PairExcludeMaskDP, torch.nn.Module):
34-
def __init__(self, *args: Any, **kwargs: Any) -> None:
35-
torch.nn.Module.__init__(self)
36-
PairExcludeMaskDP.__init__(self, *args, **kwargs)
37-
38-
def __setattr__(self, name: str, value: Any) -> None:
39-
handled, value = dpmodel_setattr(self, name, value)
40-
if not handled:
41-
super().__setattr__(name, value)
23+
@torch_module
24+
class PairExcludeMask(PairExcludeMaskDP):
25+
pass
4226

4327

4428
register_dpmodel_mapping(

deepmd/pt_expt/utils/network.py

Lines changed: 6 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,7 @@
2121
from deepmd.pt_expt.common import (
2222
register_dpmodel_mapping,
2323
to_torch_array,
24+
torch_module,
2425
)
2526

2627

@@ -37,6 +38,7 @@ def __array__(self, dtype: Any | None = None) -> np.ndarray:
3738
return arr.astype(dtype)
3839

3940

41+
# do not apply torch_module until its setattr working to register parameters
4042
class NativeLayer(NativeLayerDP, torch.nn.Module):
4143
def __init__(self, *args: Any, **kwargs: Any) -> None:
4244
torch.nn.Module.__init__(self)
@@ -78,15 +80,12 @@ def forward(self, x: torch.Tensor) -> torch.Tensor:
7880
return self.call(x)
7981

8082

81-
class NativeNet(make_multilayer_network(NativeLayer, NativeOP), torch.nn.Module):
83+
@torch_module
84+
class NativeNet(make_multilayer_network(NativeLayer, NativeOP)):
8285
def __init__(self, layers: list[dict] | None = None) -> None:
83-
torch.nn.Module.__init__(self)
8486
super().__init__(layers)
8587
self.layers = torch.nn.ModuleList(self.layers)
8688

87-
def __call__(self, *args: Any, **kwargs: Any) -> Any:
88-
return torch.nn.Module.__call__(self, *args, **kwargs)
89-
9089
def forward(self, x: torch.Tensor) -> torch.Tensor:
9190
return self.call(x)
9291

@@ -136,15 +135,15 @@ def forward(self, x: torch.Tensor) -> torch.Tensor:
136135
)
137136

138137

139-
class NetworkCollection(NetworkCollectionDP, torch.nn.Module):
138+
@torch_module
139+
class NetworkCollection(NetworkCollectionDP):
140140
NETWORK_TYPE_MAP: ClassVar[dict[str, type]] = {
141141
"network": NativeNet,
142142
"embedding_network": EmbeddingNet,
143143
"fitting_network": FittingNet,
144144
}
145145

146146
def __init__(self, *args: Any, **kwargs: Any) -> None:
147-
torch.nn.Module.__init__(self)
148147
self._module_networks = torch.nn.ModuleDict()
149148
super().__init__(*args, **kwargs)
150149

0 commit comments

Comments
 (0)