PyTorch Lightning checkpoint无'model'属性,无法加载至nn.Module
问题分析与解决方案
核心原因
这不是文档问题,也不是你的checkpoint配置有误——PyTorch Lightning的ModelCheckpoint默认保存的是整个LightningModule的完整state_dict,其中包含了你注入的model、优化器、损失函数等所有组件的参数。因为你的TestModel是作为LightningModule的self.model属性存在的,所以对应的权重键会自动带上model.前缀。官方文档中「从checkpoint加载nn.Module」的场景,通常指的是直接保存nn.Module自身的state_dict(而非LightningModule的完整state_dict)的情况。
解决方案
方案1:修改保存逻辑,单独存储模型的state_dict
在你的LightningModule中重写on_save_checkpoint方法,将TestModel的state_dict单独存入checkpoint的model键中:
class MyPLModule(pl.LightningModule): def __init__(self, model, loss_fn, optimizer): super().__init__() self.model = model self.loss_fn = loss_fn self.optimizer = optimizer def on_save_checkpoint(self, checkpoint): # 新增单独的model state_dict条目 checkpoint["model"] = self.model.state_dict()
之后推理时就可以直接加载:
model = TestModel() checkpoint = torch.load("your_checkpoint.ckpt") model.load_state_dict(checkpoint["model"])
方案2:处理现有checkpoint的state_dict,移除前缀
如果不想重新训练,可以直接修改现有checkpoint中的权重键,去掉model.前缀后加载:
checkpoint = torch.load("your_checkpoint.ckpt") # 过滤并处理model相关的权重 model_state_dict = { key.replace("model.", ""): value for key, value in checkpoint["state_dict"].items() if key.startswith("model.") } # 加载到TestModel实例 model = TestModel() model.load_state_dict(model_state_dict)
方案3:通过LightningModule加载但只提取模型
保持依赖注入的前提下,用占位符初始化LightningModule,加载checkpoint后提取内部模型:
# 仅传入必要的TestModel实例,其他依赖传占位符(无需实际功能) pl_module = MyPLModule(model=TestModel(), loss_fn=None, optimizer=None) # 加载checkpoint,关闭严格匹配以忽略无关参数 pl_module = pl_module.load_from_checkpoint( "your_checkpoint.ckpt", model=TestModel(), strict=False ) # 提取目标模型 model = pl_module.model
内容的提问来源于stack exchange,提问作者Dan Jackson
相关产品推荐
相关产品推荐

