@@ -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+
294282def 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