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

Method save

cdslib/core/script/base_train.py:733–786  ·  view source on GitHub ↗

Save the nn.module, optimizer, and customized save dict. Args: filename: the filename of the pth file to save epoch: the current epoch number

(self, filename: str, epoch: int = None)

Source from the content-addressed store, hash-verified

731 return checkpoint
732
733 def save(self, filename: str, epoch: int = None):
734 """Save the nn.module, optimizer, and customized save dict.
735
736 Args:
737 filename:
738 the filename of the pth file to save
739 epoch:
740 the current epoch number
741 """
742
743 to_save = dict()
744
745 # additional info to save
746 to_save["var_names_to_save"] = getattr(self, "var_names_to_save", set())
747 for var_name in to_save["var_names_to_save"]:
748 to_save[var_name] = getattr(self, var_name, None)
749
750 # additional info to load
751 to_save["var_names_to_load"] = getattr(self, "var_names_to_load", set())
752 for var_name in to_save["var_names_to_load"]:
753 to_save[var_name] = getattr(self, var_name, None)
754
755 # gather all model
756 nets = inspect.getmembers(self, lambda v: isinstance(v, BaseModel)) # name, net
757 for name, net in nets:
758 to_save[name] = net.state_dict()
759
760 # gather all nn.modules
761 nets = inspect.getmembers(self, lambda v: isinstance(v, torch.nn.Module)) # name, net
762 for name, net in nets:
763 if isinstance(net, DDP):
764 to_save[name] = net.module.state_dict()
765 else:
766 to_save[name] = net.state_dict()
767
768 # gather all optimizers
769 optimizers = inspect.getmembers(self, lambda v: isinstance(v, torch.optim.Optimizer)) # name, optimizer
770 for name, optimizer in optimizers:
771 to_save[name] = optimizer.state_dict()
772
773 # gather all automatic mixed precision scaler
774 try:
775 scalers = inspect.getmembers(self, lambda v: isinstance(v, torch.cuda.amp.GradScaler)) # name, scaler
776 for name, scaler in scalers:
777 to_save[name] = scaler.state_dict()
778 except:
779 pass
780
781 # save the classname
782 to_save["classname"] = self.__class__.__name__
783 to_save["epoch"] = epoch
784
785 # save to the file
786 torch.save(to_save, filename)
787
788 def _setup_for_distributed_learning(self):
789 """

Callers 1

runMethod · 0.95

Calls 1

state_dictMethod · 0.45

Tested by

no test coverage detected