Skip to content

Commit a7dc47d

Browse files
committed
fix: _parse_state_key — int(k[0]) never executed, keys stayed as strings
The isinstance(k, (list, tuple)) branch was dead: NPZ keys are always strings. tuple(k.split('_',1)) produced ('0','weights') with str '0', not int 0. Optimizer state silently failed to match after load. Centralize parsing in _parse_state_key() used by all 4 set_state methods.
1 parent 1f87830 commit a7dc47d

1 file changed

Lines changed: 13 additions & 25 deletions

File tree

nnlib/optimizers.py

Lines changed: 13 additions & 25 deletions
Original file line numberDiff line numberDiff line change
@@ -130,11 +130,7 @@ def get_state(self) -> Dict[str, Any]:
130130
def set_state(self, state: Dict[str, Any]) -> None:
131131
"""Restores SGD mutable state."""
132132
super().set_state(state)
133-
raw = state.get("_velocity", {})
134-
self._velocity = {
135-
(int(k[0]), k[1]) if isinstance(k, (list, tuple)) else tuple(k.split("_", 1)): v
136-
for k, v in raw.items()
137-
}
133+
self._velocity = {_parse_state_key(k): v for k, v in state.get("_velocity", {}).items()}
138134

139135
def get_config(self) -> Dict[str, Any]:
140136
"""Returns SGD config."""
@@ -166,11 +162,7 @@ def get_state(self) -> Dict[str, Any]:
166162
def set_state(self, state: Dict[str, Any]) -> None:
167163
"""Restores AdaGrad mutable state."""
168164
super().set_state(state)
169-
raw = state.get("_cache", {})
170-
self._cache = {
171-
(int(k[0]), k[1]) if isinstance(k, (list, tuple)) else tuple(k.split("_", 1)): v
172-
for k, v in raw.items()
173-
}
165+
self._cache = {_parse_state_key(k): v for k, v in state.get("_cache", {}).items()}
174166

175167
def get_config(self) -> Dict[str, Any]:
176168
"""Returns AdaGrad config."""
@@ -210,11 +202,7 @@ def get_state(self) -> Dict[str, Any]:
210202
def set_state(self, state: Dict[str, Any]) -> None:
211203
"""Restores RMSprop mutable state."""
212204
super().set_state(state)
213-
raw = state.get("_cache", {})
214-
self._cache = {
215-
(int(k[0]), k[1]) if isinstance(k, (list, tuple)) else tuple(k.split("_", 1)): v
216-
for k, v in raw.items()
217-
}
205+
self._cache = {_parse_state_key(k): v for k, v in state.get("_cache", {}).items()}
218206

219207
def get_config(self) -> Dict[str, Any]:
220208
"""Returns RMSprop config."""
@@ -265,16 +253,8 @@ def get_state(self) -> Dict[str, Any]:
265253
def set_state(self, state: Dict[str, Any]) -> None:
266254
"""Restores Adam mutable state."""
267255
super().set_state(state)
268-
raw_m = state.get("_m", {})
269-
self._m = {
270-
(int(k[0]), k[1]) if isinstance(k, (list, tuple)) else tuple(k.split("_", 1)): v
271-
for k, v in raw_m.items()
272-
}
273-
raw_v = state.get("_v", {})
274-
self._v = {
275-
(int(k[0]), k[1]) if isinstance(k, (list, tuple)) else tuple(k.split("_", 1)): v
276-
for k, v in raw_v.items()
277-
}
256+
self._m = {_parse_state_key(k): v for k, v in state.get("_m", {}).items()}
257+
self._v = {_parse_state_key(k): v for k, v in state.get("_v", {}).items()}
278258

279259
def get_config(self) -> Dict[str, Any]:
280260
"""Returns Adam config."""
@@ -291,6 +271,14 @@ def get_config(self) -> Dict[str, Any]:
291271
_OPT_CLASSES: Dict[str, Type[Optimizer]] = {"SGD": SGD, "AdaGrad": AdaGrad, "RMSprop": RMSprop, "Adam": Adam}
292272

293273

274+
def _parse_state_key(k: Any) -> Tuple[int, str]:
275+
"""Parses a serialized optimizer state key back to (layer_index, param_name)."""
276+
if isinstance(k, (list, tuple)):
277+
return (int(k[0]), k[1])
278+
layer_str, name = k.split("_", 1)
279+
return (int(layer_str), name)
280+
281+
294282
def get_optimizer(opt: Union[str, Dict[str, Any], Optimizer]) -> Optimizer:
295283
"""Resolves an optimizer from string, dict, or instance."""
296284
if isinstance(opt, Optimizer):

0 commit comments

Comments
 (0)