PyTorch Lightning加载模型检查点遇报错,求解决方案及疑问解答
问题根源
你用self.save_hyperparameters(ignore=['backbone', 'loss_module'])后,这两个参数没被存到超参数文件里,而MyModel的__init__又必须要这两个参数,所以load_from_checkpoint会报错要求你手动传入。
方案1:把backbone和loss的初始化逻辑搬进__init__
直接在模型的__init__里完成backbone和loss的创建,只把可序列化的配置参数(比如backbone名称、权重类型)作为入参,这样超参数里存的都是简单配置,加载时自动初始化,还不会和checkpoint里的训练权重冲突。
代码示例:
class MyModel(pl.LightningModule): def __init__(self, backbone_name='densenet121', backbone_weights='DenseNet121_Weights.DEFAULT', lr=0.01): super().__init__() # 初始化backbone self.backbone = torch.hub.load('pytorch/vision:v0.10.0', backbone_name, weights=backbone_weights) self.backbone.classifier = nn.Linear(1024, 2) # 初始化loss self.loss_module = nn.CrossEntropyLoss() self.lr = lr # 保存所有超参数(都是可序列化的配置,没有nn.Module实例) self.save_hyperparameters() # 实例化时直接传需要的配置就行 model = MyModel(lr=0.01)
关于权重的疑问
放心,加载checkpoint时不会用初始化的预训练权重覆盖训练后的权重。load_from_checkpoint的流程是:先调用__init__创建一个新模型实例(此时用预训练权重初始化backbone),然后立刻把checkpoint里保存的所有模型权重(包括backbone训练后的权重)加载进来覆盖掉初始值。最终模型用的是checkpoint里的训练权重,不是初始的预训练权重。
方案2:加载时手动传入backbone和loss实例
如果不想改__init__,可以每次加载时先创建好和训练时一样的backbone、loss实例,传给load_from_checkpoint:
# 先复刻训练时的backbone和loss backbone = torch.hub.load('pytorch/vision:v0.10.0', 'densenet121', weights='DenseNet121_Weights.DEFAULT') loss = nn.CrossEntropyLoss() # 加载时传入这两个参数 model = MyModel.load_from_checkpoint(path_to_checkpoint, backbone=backbone, loss_module=loss)
这种方法的问题是每次加载都要重复写初始化代码,容易出错,不如方案1省心。
为什么训练时会有警告?
PyTorch Lightning不建议把nn.Module实例存到超参数里,因为模型的权重已经会单独存在checkpoint里了,再存超参数纯属冗余,而且nn.Module序列化还可能出问题,所以提示你忽略这些参数。
内容的提问来源于stack exchange,提问作者Naibed

