From b37ddb8577a4530f0159057e6a7a0862bfe345c4 Mon Sep 17 00:00:00 2001 From: Tai An Date: Thu, 3 Sep 2026 18:52:35 -0700 Subject: [PATCH] fix(pt): pass device and dtype in SeZMDeNSFittingNet.deserialize safe_numpy_to_tensor takes `device` and `dtype` as required keyword-only arguments, so SeZMDeNSFittingNet.deserialize raised TypeError. Rebuild the state dict from the instantiated module's own state_dict, matching the other sezm_nn deserialize implementations. Signed-off-by: Anai-Guo --- deepmd/pt/model/descriptor/sezm_nn/dens.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/deepmd/pt/model/descriptor/sezm_nn/dens.py b/deepmd/pt/model/descriptor/sezm_nn/dens.py index 8dcf872188..a5f084a0a8 100644 --- a/deepmd/pt/model/descriptor/sezm_nn/dens.py +++ b/deepmd/pt/model/descriptor/sezm_nn/dens.py @@ -746,6 +746,12 @@ def deserialize(cls, data: dict[str, Any]) -> SeZMDeNSFittingNet: config = data.pop("config") variables = data.pop("@variables") obj = cls(**config) - state = {key: safe_numpy_to_tensor(value) for key, value in variables.items()} + template = obj.state_dict() + state = { + key: safe_numpy_to_tensor( + value, device=template[key].device, dtype=template[key].dtype + ) + for key, value in variables.items() + } obj.load_state_dict(state) return obj