MCPcopy Create free account
hub / github.com/apple/ml-pointersect / state_dict

Method state_dict

cdslib/core/models/base_model.py:231–266  ·  view source on GitHub ↗

Returns a dictionary that can be saved or load.

(self)

Source from the content-addressed store, hash-verified

229 self.register_parameter(name, state_dict[name])
230
231 def state_dict(self) -> T.Dict[str, T.Any]:
232 """Returns a dictionary that can be saved or load."""
233
234 to_save = dict()
235
236 # additional info to save
237 to_save["buffer_names"] = getattr(self, "buffer_names", set())
238 to_save["parameter_names"] = getattr(self, "parameter_names", set())
239 to_save["var_names_to_save"] = getattr(self, "var_names_to_save", set())
240 for var_name in to_save["var_names_to_save"]:
241 to_save[var_name] = getattr(self, var_name, None)
242
243 # additional info to load
244 to_save["var_names_to_load"] = getattr(self, "var_names_to_load", set())
245 for var_name in to_save["var_names_to_load"]:
246 to_save[var_name] = getattr(self, var_name, None)
247
248 # gather all base models
249 nets = inspect.getmembers(self, lambda v: isinstance(v, BaseModel)) # name, net
250 for name, net in nets:
251 to_save[name] = net.state_dict()
252
253 # gather all nn.modules
254 nets = inspect.getmembers(self, lambda v: isinstance(v, torch.nn.Module)) # name, net
255 for name, net in nets:
256 if isinstance(net, DDP):
257 to_save[name] = net.module.state_dict()
258 elif isinstance(net, APEX_DDP):
259 to_save[name] = net.module.state_dict()
260 else:
261 to_save[name] = net.state_dict()
262
263 # save the classname
264 to_save["classname"] = self.__class__.__name__
265
266 return to_save
267
268 def setup_for_distributed_learning(self, ddp_type: str):
269 """

Callers 2

apply_gradient_allreduceFunction · 0.45
saveMethod · 0.45

Calls

no outgoing calls

Tested by

no test coverage detected