You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.01 18:40:06