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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.27 12:25:02