@@ -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 ,
0 commit comments