Skip to content

Commit ec2e031

Browse files
author
Han Wang
committed
implement pytorch-exportable for se_e2_a descriptor
1 parent 8787b45 commit ec2e031

13 files changed

Lines changed: 676 additions & 5 deletions

File tree

deepmd/backend/pt_expt.py

Lines changed: 126 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,126 @@
1+
# SPDX-License-Identifier: LGPL-3.0-or-later
2+
from collections.abc import (
3+
Callable,
4+
)
5+
from importlib.util import (
6+
find_spec,
7+
)
8+
from typing import (
9+
TYPE_CHECKING,
10+
ClassVar,
11+
)
12+
13+
from deepmd.backend.backend import (
14+
Backend,
15+
)
16+
17+
if TYPE_CHECKING:
18+
from argparse import (
19+
Namespace,
20+
)
21+
22+
from deepmd.infer.deep_eval import (
23+
DeepEvalBackend,
24+
)
25+
from deepmd.utils.neighbor_stat import (
26+
NeighborStat,
27+
)
28+
29+
30+
@Backend.register("pt-expt")
31+
@Backend.register("pytorch-exportable")
32+
class PyTorchExportableBackend(Backend):
33+
"""PyTorch exportable backend."""
34+
35+
name = "PyTorch Exportable"
36+
"""The formal name of the backend."""
37+
features: ClassVar[Backend.Feature] = (
38+
Backend.Feature.ENTRY_POINT
39+
| Backend.Feature.DEEP_EVAL
40+
| Backend.Feature.NEIGHBOR_STAT
41+
| Backend.Feature.IO
42+
)
43+
"""The features of the backend."""
44+
suffixes: ClassVar[list[str]] = [".pth", ".pt"]
45+
"""The suffixes of the backend."""
46+
47+
def is_available(self) -> bool:
48+
"""Check if the backend is available.
49+
50+
Returns
51+
-------
52+
bool
53+
Whether the backend is available.
54+
"""
55+
return find_spec("torch") is not None
56+
57+
@property
58+
def entry_point_hook(self) -> Callable[["Namespace"], None]:
59+
"""The entry point hook of the backend.
60+
61+
Returns
62+
-------
63+
Callable[[Namespace], None]
64+
The entry point hook of the backend.
65+
"""
66+
from deepmd.pt.entrypoints.main import main as deepmd_main
67+
68+
return deepmd_main
69+
70+
@property
71+
def deep_eval(self) -> type["DeepEvalBackend"]:
72+
"""The Deep Eval backend of the backend.
73+
74+
Returns
75+
-------
76+
type[DeepEvalBackend]
77+
The Deep Eval backend of the backend.
78+
"""
79+
from deepmd.pt.infer.deep_eval import DeepEval as DeepEvalPT
80+
81+
return DeepEvalPT
82+
83+
@property
84+
def neighbor_stat(self) -> type["NeighborStat"]:
85+
"""The neighbor statistics of the backend.
86+
87+
Returns
88+
-------
89+
type[NeighborStat]
90+
The neighbor statistics of the backend.
91+
"""
92+
from deepmd.pt.utils.neighbor_stat import (
93+
NeighborStat,
94+
)
95+
96+
return NeighborStat
97+
98+
@property
99+
def serialize_hook(self) -> Callable[[str], dict]:
100+
"""The serialize hook to convert the model file to a dictionary.
101+
102+
Returns
103+
-------
104+
Callable[[str], dict]
105+
The serialize hook of the backend.
106+
"""
107+
from deepmd.pt.utils.serialization import (
108+
serialize_from_file,
109+
)
110+
111+
return serialize_from_file
112+
113+
@property
114+
def deserialize_hook(self) -> Callable[[str, dict], None]:
115+
"""The deserialize hook to convert the dictionary to a model file.
116+
117+
Returns
118+
-------
119+
Callable[[str, dict], None]
120+
The deserialize hook of the backend.
121+
"""
122+
from deepmd.pt.utils.serialization import (
123+
deserialize_to_file,
124+
)
125+
126+
return deserialize_to_file

deepmd/dpmodel/descriptor/se_e2_a.py

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -607,7 +607,11 @@ def call(
607607
sec = self.sel_cumsum
608608

609609
ng = self.neuron[-1]
610-
gr = xp.zeros([nf * nloc, ng, 4], dtype=self.dstd.dtype)
610+
gr = xp.zeros(
611+
[nf * nloc, ng, 4],
612+
dtype=self.dstd.dtype,
613+
device=array_api_compat.device(coord_ext),
614+
)
611615
exclude_mask = self.emask.build_type_exclude_mask(nlist, atype_ext)
612616
# merge nf and nloc axis, so for type_one_side == False,
613617
# we don't require atype is the same in all frames

deepmd/pt_expt/__init__.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1 @@
1+
# SPDX-License-Identifier: LGPL-3.0-or-later
Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,8 @@
1+
# SPDX-License-Identifier: LGPL-3.0-or-later
2+
from .se_e2_a import (
3+
DescrptSeA,
4+
)
5+
6+
__all__ = [
7+
"DescrptSeA",
8+
]
Lines changed: 101 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,101 @@
1+
# SPDX-License-Identifier: LGPL-3.0-or-later
2+
from typing import (
3+
Any,
4+
)
5+
6+
import torch # noqa: TID253
7+
8+
from deepmd.dpmodel.descriptor.se_e2_a import DescrptSeAArrayAPI as DescrptSeADP
9+
from deepmd.pt.model.descriptor.base_descriptor import ( # noqa: TID253
10+
BaseDescriptor,
11+
)
12+
from deepmd.pt.utils import ( # noqa: TID253
13+
env,
14+
)
15+
from deepmd.pt.utils.exclude_mask import ( # noqa: TID253
16+
PairExcludeMask,
17+
)
18+
from deepmd.pt_expt.utils.network import (
19+
NetworkCollection,
20+
)
21+
22+
23+
@BaseDescriptor.register("se_e2_a_expt")
24+
@BaseDescriptor.register("se_a_expt")
25+
class DescrptSeA(DescrptSeADP, torch.nn.Module):
26+
def __init__(self, *args: Any, **kwargs: Any) -> None:
27+
torch.nn.Module.__init__(self)
28+
DescrptSeADP.__init__(self, *args, **kwargs)
29+
self._convert_state()
30+
31+
def __setattr__(self, name: str, value: Any) -> None:
32+
if name in {"davg", "dstd"} and "_buffers" in self.__dict__:
33+
tensor = (
34+
None if value is None else torch.as_tensor(value, device=env.DEVICE)
35+
)
36+
if name in self._buffers:
37+
self._buffers[name] = tensor
38+
return
39+
return super().__setattr__(name, tensor)
40+
if name == "embeddings" and "_modules" in self.__dict__:
41+
if value is not None and not isinstance(value, torch.nn.Module):
42+
if hasattr(value, "serialize"):
43+
value = NetworkCollection.deserialize(value.serialize())
44+
elif isinstance(value, dict):
45+
value = NetworkCollection.deserialize(value)
46+
return super().__setattr__(name, value)
47+
if name == "emask" and "_modules" in self.__dict__:
48+
if value is not None and not isinstance(value, torch.nn.Module):
49+
value = PairExcludeMask(
50+
self.ntypes, exclude_types=list(value.get_exclude_types())
51+
)
52+
return super().__setattr__(name, value)
53+
return super().__setattr__(name, value)
54+
55+
def _convert_state(self) -> None:
56+
if self.davg is not None:
57+
davg = torch.as_tensor(self.davg, device=env.DEVICE)
58+
if "davg" in self._buffers:
59+
self._buffers["davg"] = davg
60+
else:
61+
if hasattr(self, "davg"):
62+
delattr(self, "davg")
63+
self.register_buffer("davg", davg)
64+
if self.dstd is not None:
65+
dstd = torch.as_tensor(self.dstd, device=env.DEVICE)
66+
if "dstd" in self._buffers:
67+
self._buffers["dstd"] = dstd
68+
else:
69+
if hasattr(self, "dstd"):
70+
delattr(self, "dstd")
71+
self.register_buffer("dstd", dstd)
72+
if self.embeddings is not None:
73+
self.embeddings = NetworkCollection.deserialize(self.embeddings.serialize())
74+
if self.emask is not None:
75+
self.emask = PairExcludeMask(
76+
self.ntypes, exclude_types=list(self.emask.get_exclude_types())
77+
)
78+
79+
def forward(
80+
self,
81+
nlist: torch.Tensor,
82+
extended_coord: torch.Tensor,
83+
extended_atype: torch.Tensor,
84+
extended_atype_embd: torch.Tensor | None = None,
85+
mapping: torch.Tensor | None = None,
86+
type_embedding: torch.Tensor | None = None,
87+
) -> tuple[
88+
torch.Tensor,
89+
torch.Tensor | None,
90+
torch.Tensor | None,
91+
torch.Tensor | None,
92+
torch.Tensor | None,
93+
]:
94+
del extended_atype_embd, type_embedding
95+
descrpt, rot_mat, g2, h2, sw = self.call(
96+
extended_coord,
97+
extended_atype,
98+
nlist,
99+
mapping=mapping,
100+
)
101+
return descrpt, rot_mat, g2, h2, sw

deepmd/pt_expt/utils/__init__.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1 @@
1+
# SPDX-License-Identifier: LGPL-3.0-or-later

deepmd/pt_expt/utils/network.py

Lines changed: 130 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,130 @@
1+
# SPDX-License-Identifier: LGPL-3.0-or-later
2+
from typing import (
3+
Any,
4+
ClassVar,
5+
Self,
6+
)
7+
8+
import numpy as np
9+
import torch # noqa: TID253
10+
11+
from deepmd.dpmodel.common import (
12+
NativeOP,
13+
)
14+
from deepmd.dpmodel.utils.network import LayerNorm as LayerNormDP
15+
from deepmd.dpmodel.utils.network import NativeLayer as NativeLayerDP
16+
from deepmd.dpmodel.utils.network import NetworkCollection as NetworkCollectionDP
17+
from deepmd.dpmodel.utils.network import (
18+
make_embedding_network,
19+
make_fitting_network,
20+
make_multilayer_network,
21+
)
22+
from deepmd.pt.utils import ( # noqa: TID253
23+
env,
24+
)
25+
26+
27+
def _to_torch_array(value: Any) -> torch.Tensor | None:
28+
if value is None:
29+
return None
30+
if torch.is_tensor(value):
31+
return value
32+
return torch.as_tensor(value, device=env.DEVICE)
33+
34+
35+
class TorchArrayParam(torch.nn.Parameter):
36+
def __new__(cls, data: Any = None, requires_grad: bool = True) -> Self:
37+
return torch.nn.Parameter.__new__(cls, data, requires_grad)
38+
39+
def __array__(self, dtype: Any | None = None) -> np.ndarray:
40+
arr = self.detach().cpu().numpy()
41+
if dtype is None:
42+
return arr
43+
return arr.astype(dtype)
44+
45+
46+
class NativeLayer(NativeLayerDP, torch.nn.Module):
47+
def __init__(self, *args: Any, **kwargs: Any) -> None:
48+
torch.nn.Module.__init__(self)
49+
NativeLayerDP.__init__(self, *args, **kwargs)
50+
for name in ("w", "b", "idt"):
51+
if name in self._parameters or name in self._buffers:
52+
continue
53+
val = _to_torch_array(getattr(self, name))
54+
if val is None:
55+
continue
56+
if self.trainable:
57+
if hasattr(self, name) and name not in self._parameters:
58+
delattr(self, name)
59+
self.register_parameter(name, TorchArrayParam(val, requires_grad=True))
60+
else:
61+
if hasattr(self, name) and name not in self._buffers:
62+
delattr(self, name)
63+
self.register_buffer(name, val)
64+
65+
def __setattr__(self, name: str, value: Any) -> None:
66+
if name in {"w", "b", "idt"} and "_parameters" in self.__dict__:
67+
val = _to_torch_array(value)
68+
if val is None:
69+
return super().__setattr__(name, None)
70+
if getattr(self, "trainable", False):
71+
param = (
72+
value
73+
if isinstance(value, TorchArrayParam)
74+
else TorchArrayParam(val, requires_grad=True)
75+
)
76+
if name in self._parameters:
77+
self._parameters[name] = param
78+
return
79+
return super().__setattr__(name, param)
80+
if name in self._buffers:
81+
self._buffers[name] = val
82+
return
83+
return super().__setattr__(name, val)
84+
return super().__setattr__(name, value)
85+
86+
def forward(self, x: torch.Tensor) -> torch.Tensor:
87+
return self.call(x)
88+
89+
90+
class NativeNet(make_multilayer_network(NativeLayer, NativeOP), torch.nn.Module):
91+
def __init__(self, layers: list[dict] | None = None) -> None:
92+
torch.nn.Module.__init__(self)
93+
super().__init__(layers)
94+
self.layers = torch.nn.ModuleList(self.layers)
95+
96+
def forward(self, x: torch.Tensor) -> torch.Tensor:
97+
return self.call(x)
98+
99+
100+
class EmbeddingNet(make_embedding_network(NativeNet, NativeLayer)):
101+
pass
102+
103+
104+
class FittingNet(make_fitting_network(EmbeddingNet, NativeNet, NativeLayer)):
105+
pass
106+
107+
108+
class NetworkCollection(NetworkCollectionDP, torch.nn.Module):
109+
NETWORK_TYPE_MAP: ClassVar[dict[str, type]] = {
110+
"network": NativeNet,
111+
"embedding_network": EmbeddingNet,
112+
"fitting_network": FittingNet,
113+
}
114+
115+
def __init__(self, *args: Any, **kwargs: Any) -> None:
116+
torch.nn.Module.__init__(self)
117+
super().__init__(*args, **kwargs)
118+
self._module_networks = torch.nn.ModuleDict()
119+
for idx, net in enumerate(self._networks):
120+
if isinstance(net, torch.nn.Module):
121+
self._module_networks[str(idx)] = net
122+
123+
def __setitem__(self, key: int | tuple, value: Any) -> None:
124+
super().__setitem__(key, value)
125+
if isinstance(value, torch.nn.Module):
126+
self._module_networks[str(self._convert_key(key))] = value
127+
128+
129+
class LayerNorm(LayerNormDP, NativeLayer):
130+
pass

0 commit comments

Comments
 (0)