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

如何通过层名称从预训练PyTorch模型生成新模型?

问题分析

你遇到的TypeError是因为PyTorch的load_state_dict()方法根本没有params这个参数,不能直接传要加载的层列表。另外你加载checkpoint的代码也有问题——你保存的是模型的state_dict本身,所以torch.load(PATH)得到的直接就是参数字典,不需要取checkpoint['model_state_dict']。

正确实现方法

核心思路是先从预训练的state_dict里筛选出你需要的层参数,再把这个筛选后的字典传给load_state_dict(),同时用strict=False跳过不匹配的参数,完全满足你用字符串指定层名称的需求。

完整代码示例:

import torch
import torch.nn as nn
import torch.nn.functional as F

# 原模型定义
class Net(nn.Module):
    def __init__(self):
        super(Net, self).__init__()
        self.conv1 = nn.Conv2d(3, 6, 5)
        self.pool = nn.MaxPool2d(2, 2)
        self.conv2 = nn.Conv2d(6, 16, 5)
        self.fc1 = nn.Linear(16 * 5 * 5, 120)
        self.fc2 = nn.Linear(120, 84)
        self.fc3 = nn.Linear(84, 10)

    def forward(self, x):
        x = self.pool(F.relu(self.conv1(x)))
        x = self.pool(F.relu(self.conv2(x)))
        x = x.view(-1, 16 * 5 * 5)
        x = F.relu(self.fc1(x))
        x = F.relu(self.fc2(x))
        x = self.fc3(x)
        return x

# 指定要加载的层名称
sub_layers = ['conv1.weight', 'conv1.bias','conv2.weight','conv2.bias']
PATH = '/mydrive/net'

# 加载预训练的完整参数字典
pretrained_state = torch.load(PATH)
# 筛选出目标层的参数
filtered_state = {k: v for k, v in pretrained_state.items() if k in sub_layers}

# 初始化新模型并加载筛选后的参数
new_model = Net()
new_model.load_state_dict(filtered_state, strict=False)
new_model.eval()

关键说明

  • 用字典推导式直接匹配你指定的层名称字符串,不用循环遍历连续层,完全符合你的需求。
  • strict=False必须添加:因为新模型包含未在筛选列表里的参数(比如fc1、fc2等),不加会因键不匹配抛出错误。
  • 如果你当初保存checkpoint时用了嵌套格式(比如torch.save({'model_state_dict': model.state_dict()}, PATH)),才需要用checkpoint['model_state_dict']取值,你之前的保存代码是直接存state_dict,所以不需要嵌套取值。

内容的提问来源于stack exchange,提问作者a.rose

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 22:57:39