如何在PyTorch中加载带预训练权重的部分模型
解决PyTorch迁移学习中保留全连接层随机初始化权重的问题
问题分析
你的代码存在几个关键问题:
- 新模型实例化时使用了
Classifier_model(),但定义的模型类是NewModel,类名不匹配 - 处理全连接层权重的循环逻辑无效,没有实现保留随机初始化权重的目的
trained_model_state_dict.update(new_model_state_dict)的操作逻辑颠倒,应该用预训练的卷积层权重更新新模型的参数,而非反过来
正确实现步骤
1. 确保模型类定义与实例化一致
先正确定义模型类并实例化新模型,同时加载预训练权重:
import torch import torch.nn as nn class NewModel(nn.Module): def __init__(self): super(NewModel, self).__init__() self.conv1 = nn.Conv2d(3, 64, kernel_size=3) self.conv2 = nn.Conv2d(64, 128, kernel_size=3) self.fc1 = nn.Linear(128 * 30 * 30, 512) self.fc2 = nn.Linear(512, 10) def forward(self, x): x = self.conv1(x) x = self.conv2(x) x = x.view(-1, 128 * 30 * 30) x = self.fc1(x) x = self.fc2(x) return x # 加载预训练权重文件 trained_model_path = "/content/model_weights.pth" trained_state_dict = torch.load(trained_model_path) # 实例化新模型(与类名保持一致) new_model = NewModel()
2. 过滤预训练权重,仅保留卷积层参数
从预训练模型的state_dict中移除所有以fc开头的键,只保留卷积层的权重参数:
# 过滤预训练权重,仅保留卷积层相关参数 filtered_trained_state = {k: v for k, v in trained_state_dict.items() if not k.startswith('fc')}
3. 将过滤后的权重加载到新模型
使用strict=False参数允许部分参数匹配,这样新模型的全连接层会保留初始化的随机值,不会被预训练权重覆盖:
# 加载过滤后的权重到新模型 new_model.load_state_dict(filtered_trained_state, strict=False)
4. (可选)冻结卷积层,仅训练全连接层
如果后续只需要训练全连接层,可以冻结卷积层的参数,避免它们在训练中被更新:
# 冻结卷积层参数 for param in new_model.conv1.parameters(): param.requires_grad = False for param in new_model.conv2.parameters(): param.requires_grad = False # 确保全连接层参数可训练(默认即为True,可显式设置) for param in new_model.fc1.parameters(): param.requires_grad = True for param in new_model.fc2.parameters(): param.requires_grad = True
完整可运行代码
import torch import torch.nn as nn class NewModel(nn.Module): def __init__(self): super(NewModel, self).__init__() self.conv1 = nn.Conv2d(3, 64, kernel_size=3) self.conv2 = nn.Conv2d(64, 128, kernel_size=3) self.fc1 = nn.Linear(128 * 30 * 30, 512) self.fc2 = nn.Linear(512, 10) def forward(self, x): x = self.conv1(x) x = self.conv2(x) x = x.view(-1, 128 * 30 * 30) x = self.fc1(x) x = self.fc2(x) return x # 加载预训练权重 trained_model_path = "/content/model_weights.pth" trained_state_dict = torch.load(trained_model_path) # 实例化新模型 new_model = NewModel() # 过滤预训练权重,仅保留卷积层参数 filtered_trained_state = {k: v for k, v in trained_state_dict.items() if not k.startswith('fc')} # 加载过滤后的权重到新模型 new_model.load_state_dict(filtered_trained_state, strict=False) # 可选:冻结卷积层,仅训练全连接层 for param in new_model.conv1.parameters(): param.requires_grad = False for param in new_model.conv2.parameters(): param.requires_grad = False # 验证:打印fc1的部分权重,确认是随机初始化值 print("fc1权重前5x5部分:", new_model.fc1.weight[:5, :5])
内容的提问来源于stack exchange,提问作者Formal_this
相关产品推荐
相关产品推荐

