Returns a dictionary that can be saved or load.
(self)
| 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 | """ |
no outgoing calls
no test coverage detected