forked from deepmodeling/deepmd-kit
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathdipole_model.py
More file actions
123 lines (113 loc) · 3.85 KB
/
Copy pathdipole_model.py
File metadata and controls
123 lines (113 loc) · 3.85 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
# SPDX-License-Identifier: LGPL-3.0-or-later
from typing import (
Any,
)
from deepmd.dpmodel.array_api import (
Array,
)
from deepmd.dpmodel.atomic_model import (
DPDipoleAtomicModel,
)
from deepmd.dpmodel.model.base_model import (
BaseModel,
)
from .dp_model import (
DPModelCommon,
)
from .make_model import (
make_model,
)
DPDipoleModel_ = make_model(DPDipoleAtomicModel)
@BaseModel.register("dipole")
class DipoleModel(DPModelCommon, DPDipoleModel_):
model_type = "dipole"
def __init__(
self,
*args: Any,
**kwargs: Any,
) -> None:
DPModelCommon.__init__(self)
DPDipoleModel_.__init__(self, *args, **kwargs)
def translated_output_def(self) -> dict[str, Any]:
out_def_data = self.model_output_def().get_data()
output_def = {
"dipole": out_def_data["dipole"],
"global_dipole": out_def_data["dipole_redu"],
}
if self.do_grad_r("dipole"):
output_def["force"] = out_def_data["dipole_derv_r"]
output_def["force"].squeeze(-2)
if self.do_grad_c("dipole"):
output_def["virial"] = out_def_data["dipole_derv_c_redu"]
output_def["virial"].squeeze(-2)
output_def["atom_virial"] = out_def_data["dipole_derv_c"]
output_def["atom_virial"].squeeze(-2)
if "mask" in out_def_data:
output_def["mask"] = out_def_data["mask"]
return output_def
def call(
self,
coord: Array,
atype: Array,
box: Array | None = None,
fparam: Array | None = None,
aparam: Array | None = None,
do_atomic_virial: bool = False,
) -> dict[str, Array]:
model_ret = self.call_common(
coord,
atype,
box,
fparam=fparam,
aparam=aparam,
do_atomic_virial=do_atomic_virial,
)
if self.get_fitting_net() is not None:
model_predict = {}
model_predict["dipole"] = model_ret["dipole"]
model_predict["global_dipole"] = model_ret["dipole_redu"]
if self.do_grad_r("dipole"):
model_predict["force"] = model_ret.get("dipole_derv_r")
if self.do_grad_c("dipole"):
model_predict["virial"] = model_ret.get("dipole_derv_c_redu")
if do_atomic_virial:
model_predict["atom_virial"] = model_ret.get("dipole_derv_c")
if "mask" in model_ret:
model_predict["mask"] = model_ret["mask"]
else:
model_predict = model_ret
model_predict["updated_coord"] += coord
return model_predict
def call_lower(
self,
extended_coord: Array,
extended_atype: Array,
nlist: Array,
mapping: Array | None = None,
fparam: Array | None = None,
aparam: Array | None = None,
do_atomic_virial: bool = False,
) -> dict[str, Array]:
model_ret = self.call_common_lower(
extended_coord,
extended_atype,
nlist,
mapping,
fparam=fparam,
aparam=aparam,
do_atomic_virial=do_atomic_virial,
)
if self.get_fitting_net() is not None:
model_predict = {}
model_predict["dipole"] = model_ret["dipole"]
model_predict["global_dipole"] = model_ret["dipole_redu"]
if self.do_grad_r("dipole"):
model_predict["extended_force"] = model_ret.get("dipole_derv_r")
if self.do_grad_c("dipole"):
model_predict["virial"] = model_ret.get("dipole_derv_c_redu")
if do_atomic_virial:
model_predict["extended_virial"] = model_ret.get("dipole_derv_c")
else:
model_predict = model_ret
return model_predict
forward_lower = call_lower