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)
| 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 | """ |