Skip to content

Commit 1ef67ec

Browse files
author
Han Wang
committed
make output of energy model compatible among backends
1 parent c15212d commit 1ef67ec

6 files changed

Lines changed: 157 additions & 40 deletions

File tree

deepmd/dpmodel/model/ener_model.py

Lines changed: 96 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,9 @@
66
Any,
77
)
88

9+
from deepmd.dpmodel.array_api import (
10+
Array,
11+
)
912
from deepmd.dpmodel.atomic_model import (
1013
DPEnergyAtomicModel,
1114
)
@@ -48,6 +51,99 @@ def atomic_output_def(self) -> FittingOutputDef:
4851
return self.hess_fitting_def
4952
return super().atomic_output_def()
5053

54+
def call_lower(
55+
self,
56+
extended_coord: Array,
57+
extended_atype: Array,
58+
nlist: Array,
59+
mapping: Array | None = None,
60+
fparam: Array | None = None,
61+
aparam: Array | None = None,
62+
do_atomic_virial: bool = False,
63+
) -> dict[str, Array]:
64+
model_ret = self.call_common_lower(
65+
extended_coord,
66+
extended_atype,
67+
nlist,
68+
mapping,
69+
fparam=fparam,
70+
aparam=aparam,
71+
do_atomic_virial=do_atomic_virial,
72+
)
73+
model_predict = {}
74+
model_predict["atom_energy"] = model_ret["energy"]
75+
model_predict["energy"] = model_ret["energy_redu"]
76+
if self.do_grad_r("energy"):
77+
if model_ret["energy_derv_r"] is not None:
78+
model_predict["extended_force"] = model_ret["energy_derv_r"].squeeze(-2)
79+
else:
80+
model_predict["extended_force"] = model_ret["energy_derv_r"]
81+
if self.do_grad_c("energy"):
82+
derv_c_redu = model_ret.get("energy_derv_c_redu")
83+
if derv_c_redu is not None:
84+
model_predict["virial"] = derv_c_redu.squeeze(-2)
85+
else:
86+
model_predict["virial"] = derv_c_redu
87+
if do_atomic_virial:
88+
if model_ret["energy_derv_c"] is not None:
89+
model_predict["extended_virial"] = model_ret[
90+
"energy_derv_c"
91+
].squeeze(-3)
92+
else:
93+
model_predict["extended_virial"] = model_ret["energy_derv_c"]
94+
else:
95+
if model_ret.get("dforce") is not None:
96+
model_predict["dforce"] = model_ret["dforce"]
97+
if "mask" in model_ret:
98+
model_predict["mask"] = model_ret["mask"]
99+
return model_predict
100+
101+
def call(
102+
self,
103+
coord: Array,
104+
atype: Array,
105+
box: Array | None = None,
106+
fparam: Array | None = None,
107+
aparam: Array | None = None,
108+
do_atomic_virial: bool = False,
109+
) -> dict[str, Array]:
110+
model_ret = self.call_common(
111+
coord,
112+
atype,
113+
box,
114+
fparam=fparam,
115+
aparam=aparam,
116+
do_atomic_virial=do_atomic_virial,
117+
)
118+
model_predict = {}
119+
model_predict["atom_energy"] = model_ret["energy"]
120+
model_predict["energy"] = model_ret["energy_redu"]
121+
if self.do_grad_r("energy"):
122+
if model_ret.get("energy_derv_r") is not None:
123+
model_predict["force"] = model_ret["energy_derv_r"].squeeze(-2)
124+
else:
125+
model_predict["force"] = model_ret.get("energy_derv_r")
126+
if self.do_grad_c("energy"):
127+
derv_c_redu = model_ret.get("energy_derv_c_redu")
128+
if derv_c_redu is not None:
129+
model_predict["virial"] = derv_c_redu.squeeze(-2)
130+
else:
131+
model_predict["virial"] = derv_c_redu
132+
if do_atomic_virial:
133+
derv_c = model_ret.get("energy_derv_c")
134+
if derv_c is not None:
135+
model_predict["atom_virial"] = derv_c.squeeze(-3)
136+
else:
137+
model_predict["atom_virial"] = derv_c
138+
else:
139+
if model_ret.get("dforce") is not None:
140+
model_predict["force"] = model_ret["dforce"]
141+
if "mask" in model_ret:
142+
model_predict["mask"] = model_ret["mask"]
143+
return model_predict
144+
145+
forward_lower = call_lower
146+
51147
def translated_output_def(self) -> dict[str, Any]:
52148
"""Get the translated output definition.
53149

deepmd/dpmodel/model/make_model.py

Lines changed: 8 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -223,7 +223,7 @@ def enable_compression(
223223
check_frequency,
224224
)
225225

226-
def call(
226+
def call_common(
227227
self,
228228
coord: Array,
229229
atype: Array,
@@ -232,7 +232,7 @@ def call(
232232
aparam: Array | None = None,
233233
do_atomic_virial: bool = False,
234234
) -> dict[str, Array]:
235-
"""Return model prediction.
235+
"""Return model prediction with raw internal keys.
236236
237237
Parameters
238238
----------
@@ -262,7 +262,7 @@ def call(
262262
)
263263
del coord, box, fparam, aparam
264264
model_predict = model_call_from_call_lower(
265-
call_lower=self.call_lower,
265+
call_lower=self.call_common_lower,
266266
rcut=self.get_rcut(),
267267
sel=self.get_sel(),
268268
mixed_types=self.mixed_types(),
@@ -277,7 +277,7 @@ def call(
277277
model_predict = self._output_type_cast(model_predict, input_prec)
278278
return model_predict
279279

280-
def call_lower(
280+
def call_common_lower(
281281
self,
282282
extended_coord: Array,
283283
extended_atype: Array,
@@ -365,9 +365,10 @@ def forward_common_atomic(
365365
mask=atomic_ret["mask"] if "mask" in atomic_ret else None,
366366
)
367367

368-
forward_lower = call_lower
369-
forward_common = call
370-
forward_common_lower = call_lower
368+
call = call_common
369+
call_lower = call_common_lower
370+
forward_common = call_common
371+
forward_common_lower = call_common_lower
371372

372373
def get_out_bias(self) -> Array:
373374
"""Get the output bias."""

deepmd/dpmodel/model/spin_model.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -377,7 +377,7 @@ def call(
377377
coord_updated, atype_updated = self.process_spin_input(coord, atype, spin)
378378
if aparam is not None:
379379
aparam = self.expand_aparam(aparam, nloc * 2)
380-
model_predict = self.backbone_model.call(
380+
model_predict = self.backbone_model.call_common(
381381
coord_updated,
382382
atype_updated,
383383
box,
@@ -447,7 +447,7 @@ def call_lower(
447447
)
448448
if aparam is not None:
449449
aparam = self.expand_aparam(aparam, nloc * 2)
450-
model_predict = self.backbone_model.call_lower(
450+
model_predict = self.backbone_model.call_common_lower(
451451
extended_coord_updated,
452452
extended_atype_updated,
453453
nlist_updated,

deepmd/pt_expt/model/ener_model.py

Lines changed: 23 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -42,7 +42,7 @@ def forward(
4242
aparam: torch.Tensor | None = None,
4343
do_atomic_virial: bool = False,
4444
) -> dict[str, torch.Tensor]:
45-
model_ret = self.call(
45+
model_ret = self.call_common(
4646
coord,
4747
atype,
4848
box,
@@ -73,7 +73,7 @@ def _forward_lower(
7373
aparam: torch.Tensor | None = None,
7474
do_atomic_virial: bool = False,
7575
) -> dict[str, torch.Tensor]:
76-
model_ret = self.call_lower(
76+
model_ret = self.call_common_lower(
7777
extended_coord,
7878
extended_atype,
7979
nlist,
@@ -106,6 +106,26 @@ def forward_lower(
106106
fparam: torch.Tensor | None = None,
107107
aparam: torch.Tensor | None = None,
108108
do_atomic_virial: bool = False,
109+
) -> dict[str, torch.Tensor]:
110+
return self._forward_lower(
111+
extended_coord,
112+
extended_atype,
113+
nlist,
114+
mapping,
115+
fparam=fparam,
116+
aparam=aparam,
117+
do_atomic_virial=do_atomic_virial,
118+
)
119+
120+
def forward_lower_exportable(
121+
self,
122+
extended_coord: torch.Tensor,
123+
extended_atype: torch.Tensor,
124+
nlist: torch.Tensor,
125+
mapping: torch.Tensor | None = None,
126+
fparam: torch.Tensor | None = None,
127+
aparam: torch.Tensor | None = None,
128+
do_atomic_virial: bool = False,
109129
) -> torch.nn.Module:
110130
"""Trace ``_forward_lower`` into an exportable module.
111131
@@ -123,7 +143,7 @@ def forward_lower(
123143
torch.nn.Module
124144
A traced module whose ``forward`` accepts
125145
``(extended_coord, extended_atype, nlist, mapping, fparam, aparam)``
126-
and returns a dict with the same keys as ``_forward_lower``.
146+
and returns a dict with the same keys as ``forward_lower``.
127147
"""
128148
model = self
129149

source/tests/consistent/model/test_ener.py

Lines changed: 19 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -271,8 +271,8 @@ def extract_ret(self, ret: Any, backend) -> tuple[np.ndarray, ...]:
271271
# shape not matched. ravel...
272272
if backend is self.RefBackend.DP:
273273
return (
274-
ret["energy_redu"].ravel(),
275274
ret["energy"].ravel(),
275+
ret["atom_energy"].ravel(),
276276
SKIP_FLAG,
277277
SKIP_FLAG,
278278
SKIP_FLAG,
@@ -303,11 +303,11 @@ def extract_ret(self, ret: Any, backend) -> tuple[np.ndarray, ...]:
303303
)
304304
elif backend is self.RefBackend.JAX:
305305
return (
306-
ret["energy_redu"].ravel(),
307306
ret["energy"].ravel(),
308-
ret["energy_derv_r"].ravel(),
309-
ret["energy_derv_c_redu"].ravel(),
310-
ret["energy_derv_c"].ravel(),
307+
ret["atom_energy"].ravel(),
308+
ret["force"].ravel(),
309+
ret["virial"].ravel(),
310+
ret["atom_virial"].ravel(),
311311
)
312312
elif backend is self.RefBackend.PD:
313313
return (
@@ -499,7 +499,7 @@ def eval_pt_expt(self, pt_expt_obj: Any) -> Any:
499499
coord_tensor.requires_grad_(True)
500500
return {
501501
kk: vv.detach().cpu().numpy() if vv is not None else None
502-
for kk, vv in pt_expt_obj.call_lower(
502+
for kk, vv in pt_expt_obj.forward_lower(
503503
coord_tensor,
504504
pt_expt_numpy_to_torch(self.extended_atype),
505505
pt_expt_numpy_to_torch(self.nlist),
@@ -536,8 +536,8 @@ def extract_ret(self, ret: Any, backend) -> tuple[np.ndarray, ...]:
536536
# shape not matched. ravel...
537537
if backend is self.RefBackend.DP:
538538
return (
539-
ret["energy_redu"].ravel(),
540539
ret["energy"].ravel(),
540+
ret["atom_energy"].ravel(),
541541
SKIP_FLAG,
542542
SKIP_FLAG,
543543
SKIP_FLAG,
@@ -552,19 +552,19 @@ def extract_ret(self, ret: Any, backend) -> tuple[np.ndarray, ...]:
552552
)
553553
elif backend is self.RefBackend.PT_EXPT:
554554
return (
555-
ret["energy_redu"].ravel(),
556555
ret["energy"].ravel(),
557-
ret["energy_derv_r"].ravel(),
558-
ret["energy_derv_c_redu"].ravel(),
559-
ret["energy_derv_c"].ravel(),
556+
ret["atom_energy"].ravel(),
557+
ret["extended_force"].ravel(),
558+
ret["virial"].ravel(),
559+
ret["extended_virial"].ravel(),
560560
)
561561
elif backend is self.RefBackend.JAX:
562562
return (
563-
ret["energy_redu"].ravel(),
564563
ret["energy"].ravel(),
565-
ret["energy_derv_r"].ravel(),
566-
ret["energy_derv_c_redu"].ravel(),
567-
ret["energy_derv_c"].ravel(),
564+
ret["atom_energy"].ravel(),
565+
ret["extended_force"].ravel(),
566+
ret["virial"].ravel(),
567+
ret["extended_virial"].ravel(),
568568
)
569569
elif backend is self.RefBackend.PD:
570570
return (
@@ -726,8 +726,8 @@ def test_set_out_bias(self) -> None:
726726
)
727727

728728
def test_forward_common_alias(self) -> None:
729-
"""forward_common should be the same as call on dpmodel."""
730-
ret_call = self.dp_model.call(
729+
"""forward_common should be the same as call_common on dpmodel."""
730+
ret_call = self.dp_model.call_common(
731731
self.coords,
732732
self.atype,
733733
box=self.box,
@@ -741,8 +741,8 @@ def test_forward_common_alias(self) -> None:
741741
np.testing.assert_equal(ret_call[key], ret_fc[key])
742742

743743
def test_forward_common_lower_alias(self) -> None:
744-
"""forward_common_lower should be the same as call_lower on dpmodel."""
745-
ret_call = self.dp_model.call_lower(
744+
"""forward_common_lower should be the same as call_common_lower on dpmodel."""
745+
ret_call = self.dp_model.call_common_lower(
746746
self.extended_coord,
747747
self.extended_atype,
748748
self.nlist,

source/tests/pt_expt/model/test_ener_model.py

Lines changed: 9 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -143,11 +143,11 @@ def _prepare_lower_inputs(self):
143143
return ext_coord, ext_atype, nlist_t, mapping_t
144144

145145
def test_forward_lower_exportable(self) -> None:
146-
"""Test that EnergyModel.forward_lower returns an exportable module.
146+
"""Test that EnergyModel.forward_lower_exportable returns an exportable module.
147147
148-
forward_lower() uses make_fx to trace through torch.autograd.grad,
149-
decomposing the backward pass into primitive ops. The returned module
150-
can be passed directly to torch.export.export.
148+
forward_lower_exportable() uses make_fx to trace through
149+
torch.autograd.grad, decomposing the backward pass into primitive ops.
150+
The returned module can be passed directly to torch.export.export.
151151
152152
The test builds a model with numb_fparam > 0 and numb_aparam > 0 and
153153
verifies that:
@@ -184,7 +184,7 @@ def test_forward_lower_exportable(self) -> None:
184184
)
185185

186186
# --- eager reference with zero params ---
187-
ret_eager_zero = md._forward_lower(
187+
ret_eager_zero = md.forward_lower(
188188
ext_coord.requires_grad_(True),
189189
ext_atype,
190190
nlist_t,
@@ -197,7 +197,7 @@ def test_forward_lower_exportable(self) -> None:
197197
self.assertIn(key, ret_eager_zero)
198198

199199
# --- trace and export ---
200-
traced = md.forward_lower(
200+
traced = md.forward_lower_exportable(
201201
ext_coord,
202202
ext_atype,
203203
nlist_t,
@@ -262,7 +262,7 @@ def test_forward_lower_exportable(self) -> None:
262262
dtype=torch.float64,
263263
device=self.device,
264264
)
265-
ret_eager_nz = md._forward_lower(
265+
ret_eager_nz = md.forward_lower(
266266
ext_coord.requires_grad_(True),
267267
ext_atype,
268268
nlist_t,
@@ -364,13 +364,13 @@ def test_dp_consistency(self) -> None:
364364
ret_pt = md_pt(coord, self.atype, self.cell.reshape(1, 9))
365365

366366
np.testing.assert_allclose(
367-
ret_dp["energy_redu"],
367+
ret_dp["energy"],
368368
ret_pt["energy"].detach().cpu().numpy(),
369369
rtol=1e-10,
370370
atol=1e-10,
371371
)
372372
np.testing.assert_allclose(
373-
ret_dp["energy"],
373+
ret_dp["atom_energy"],
374374
ret_pt["atom_energy"].detach().cpu().numpy(),
375375
rtol=1e-10,
376376
atol=1e-10,

0 commit comments

Comments
 (0)