如何通过层名称从预训练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
相关产品推荐
相关产品推荐

