Skip to content

Commit fc0be62

Browse files
author
Han Wang
committed
remove eval_ hooks
1 parent 19df985 commit fc0be62

3 files changed

Lines changed: 0 additions & 111 deletions

File tree

deepmd/dpmodel/atomic_model/dp_atomic_model.py

Lines changed: 0 additions & 44 deletions
Original file line numberDiff line numberDiff line change
@@ -3,8 +3,6 @@
33
Any,
44
)
55

6-
import array_api_compat
7-
86
from deepmd.dpmodel.array_api import (
97
Array,
108
)
@@ -56,10 +54,6 @@ def __init__(
5654
if hasattr(self.fitting, "reinit_exclude"):
5755
self.fitting.reinit_exclude(self.atom_exclude_types)
5856
self.type_map = type_map
59-
self.enable_eval_descriptor_hook = False
60-
self.enable_eval_fitting_last_layer_hook = False
61-
self.eval_descriptor_list: list[Array] = []
62-
self.eval_fitting_last_layer_list: list[Array] = []
6357
super().init_out_stat()
6458

6559
def fitting_output_def(self) -> FittingOutputDef:
@@ -132,37 +126,6 @@ def enable_compression(
132126
check_frequency,
133127
)
134128

135-
def set_eval_descriptor_hook(self, enable: bool) -> None:
136-
"""Set the hook for evaluating descriptor and clear the cache."""
137-
self.enable_eval_descriptor_hook = enable
138-
self.eval_descriptor_list.clear()
139-
140-
def eval_descriptor(self) -> Array:
141-
"""Evaluate the descriptor by concatenating cached results."""
142-
if not self.eval_descriptor_list:
143-
raise RuntimeError(
144-
"eval_descriptor_list is empty. "
145-
"Call set_eval_descriptor_hook(True) and perform a forward pass first."
146-
)
147-
xp = array_api_compat.array_namespace(self.eval_descriptor_list[0])
148-
return xp.concat(self.eval_descriptor_list, axis=0)
149-
150-
def set_eval_fitting_last_layer_hook(self, enable: bool) -> None:
151-
"""Set the hook for evaluating fitting last layer output and clear the cache."""
152-
self.enable_eval_fitting_last_layer_hook = enable
153-
self.fitting.set_return_middle_output(enable)
154-
self.eval_fitting_last_layer_list.clear()
155-
156-
def eval_fitting_last_layer(self) -> Array:
157-
"""Evaluate the fitting last layer output by concatenating cached results."""
158-
if not self.eval_fitting_last_layer_list:
159-
raise RuntimeError(
160-
"eval_fitting_last_layer_list is empty. "
161-
"Call set_eval_fitting_last_layer_hook(True) and perform a forward pass first."
162-
)
163-
xp = array_api_compat.array_namespace(self.eval_fitting_last_layer_list[0])
164-
return xp.concat(self.eval_fitting_last_layer_list, axis=0)
165-
166129
def forward_atomic(
167130
self,
168131
extended_coord: Array,
@@ -203,8 +166,6 @@ def forward_atomic(
203166
nlist,
204167
mapping=mapping,
205168
)
206-
if self.enable_eval_descriptor_hook:
207-
self.eval_descriptor_list.append(descriptor)
208169
ret = self.fitting(
209170
descriptor,
210171
atype,
@@ -214,11 +175,6 @@ def forward_atomic(
214175
fparam=fparam,
215176
aparam=aparam,
216177
)
217-
if self.enable_eval_fitting_last_layer_hook:
218-
assert "middle_output" in ret, (
219-
"eval_fitting_last_layer not supported for this fitting net!"
220-
)
221-
self.eval_fitting_last_layer_list.append(ret.pop("middle_output"))
222178
return ret
223179

224180
def change_type_map(

deepmd/dpmodel/model/dp_model.py

Lines changed: 0 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,6 @@
11
# SPDX-License-Identifier: LGPL-3.0-or-later
22

33

4-
from deepmd.dpmodel.array_api import (
5-
Array,
6-
)
74
from deepmd.dpmodel.descriptor.base_descriptor import (
85
BaseDescriptor,
96
)
@@ -55,19 +52,3 @@ def get_fitting_net(self) -> BaseFitting:
5552
def get_descriptor(self) -> BaseDescriptor:
5653
"""Get the descriptor."""
5754
return self.atomic_model.descriptor
58-
59-
def set_eval_descriptor_hook(self, enable: bool) -> None:
60-
"""Set the hook for evaluating descriptor."""
61-
self.atomic_model.set_eval_descriptor_hook(enable)
62-
63-
def eval_descriptor(self) -> Array:
64-
"""Evaluate the descriptor."""
65-
return self.atomic_model.eval_descriptor()
66-
67-
def set_eval_fitting_last_layer_hook(self, enable: bool) -> None:
68-
"""Set the hook for evaluating fitting last layer output."""
69-
self.atomic_model.set_eval_fitting_last_layer_hook(enable)
70-
71-
def eval_fitting_last_layer(self) -> Array:
72-
"""Evaluate the fitting last layer output."""
73-
return self.atomic_model.eval_fitting_last_layer()

source/tests/consistent/model/test_ener.py

Lines changed: 0 additions & 48 deletions
Original file line numberDiff line numberDiff line change
@@ -757,54 +757,6 @@ def test_forward_common_lower_alias(self) -> None:
757757
for key in ret_call:
758758
np.testing.assert_equal(ret_call[key], ret_fc[key])
759759

760-
def test_eval_descriptor(self) -> None:
761-
"""eval_descriptor should produce consistent results across dp and pt."""
762-
# dpmodel
763-
self.dp_model.set_eval_descriptor_hook(True)
764-
self.dp_model.call_lower(
765-
self.extended_coord,
766-
self.extended_atype,
767-
self.nlist,
768-
self.mapping,
769-
)
770-
dp_desc = self.dp_model.eval_descriptor()
771-
772-
# pt
773-
self.pt_model.set_eval_descriptor_hook(True)
774-
self.pt_model.forward_common_lower(
775-
numpy_to_torch(self.extended_coord),
776-
numpy_to_torch(self.extended_atype),
777-
numpy_to_torch(self.nlist),
778-
numpy_to_torch(self.mapping),
779-
)
780-
pt_desc = torch_to_numpy(self.pt_model.eval_descriptor())
781-
782-
np.testing.assert_allclose(dp_desc, pt_desc, rtol=1e-10, atol=1e-10)
783-
784-
def test_eval_fitting_last_layer(self) -> None:
785-
"""eval_fitting_last_layer should produce consistent results across dp and pt."""
786-
# dpmodel
787-
self.dp_model.set_eval_fitting_last_layer_hook(True)
788-
self.dp_model.call_lower(
789-
self.extended_coord,
790-
self.extended_atype,
791-
self.nlist,
792-
self.mapping,
793-
)
794-
dp_fl = self.dp_model.eval_fitting_last_layer()
795-
796-
# pt
797-
self.pt_model.set_eval_fitting_last_layer_hook(True)
798-
self.pt_model.forward_common_lower(
799-
numpy_to_torch(self.extended_coord),
800-
numpy_to_torch(self.extended_atype),
801-
numpy_to_torch(self.nlist),
802-
numpy_to_torch(self.mapping),
803-
)
804-
pt_fl = torch_to_numpy(self.pt_model.eval_fitting_last_layer())
805-
806-
np.testing.assert_allclose(dp_fl, pt_fl, rtol=1e-10, atol=1e-10)
807-
808760
def test_model_output_def(self) -> None:
809761
"""model_output_def should return the same keys and shapes on dp and pt."""
810762
dp_def = self.dp_model.model_output_def().get_data()

0 commit comments

Comments
 (0)