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

如何从torch.save保存的模型中恢复state_dict并解决加载报错

解决PyTorch模型state_dict键不匹配问题

问题分析

你遇到的报错核心是直接保存的模型对象的state_dict键结构,和新定义模型的键不匹配:原模型state_dict的键是0.weight这类(对应Sequential层内的层序号),而新模型里Sequential被封装在model1属性下,因此需要的键是model1.0.weight这种带前缀的格式。

解决步骤

  1. 加载原模型并提取state_dict
    先加载你之前保存的模型对象,获取它的原始state_dict:

    import torch
    import torch.nn as nn
    
    # 加载原模型文件
    model1 = torch.load('path')
    original_state_dict = model1.state_dict()
    
  2. 手动修改state_dict的键
    遍历原state_dict的所有键,给每个键添加model1.前缀,生成适配新模型的state_dict:

    new_state_dict = {}
    for key, tensor in original_state_dict.items():
        # 为原键添加"model1."前缀
        new_key = f"model1.{key}"
        new_state_dict[new_key] = tensor
    
  3. 定义匹配的模型结构并加载修改后的state_dict
    根据你提供的模型维度(输入2、隐藏层5神经元、输出1),定义对应的模型结构,再加载修改后的state_dict:

    class Model(nn.Module):
        def __init__(self):
            super().__init__()
            self.model1 = nn.Sequential(
                nn.Linear(2, 5),  # 输入层到5神经元隐藏层
                nn.ReLU(),        # 激活函数(若原模型用了其他激活,替换成对应类型即可)
                nn.Linear(5, 1)   # 隐藏层到输出层
            )
        
        def forward(self, x):
            return self.model1(x)
    
    # 实例化模型并加载适配后的state_dict
    pretrained_model = Model()
    pretrained_model.load_state_dict(new_state_dict)
    
  4. 验证模型可用性
    可以用测试张量验证模型是否正常运行:

    test_input = torch.randn(1, 2)  # 符合输入维度的测试数据
    output = pretrained_model(test_input)
    print(output.shape)  # 应输出torch.Size([1, 1]),匹配输出维度要求
    

额外注意

如果原模型的Sequential层使用了其他激活函数(比如Sigmoid、Tanh),需要将代码中的nn.ReLU()替换为对应的激活函数,否则模型结构不匹配仍会报错。

内容的提问来源于stack exchange,提问作者Acad

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.17 08:35:24