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

Pytorch使用OrderedDict加载不匹配/部分模型权重的实现示例

PyTorch 自定义OrderedDict加载权重(支持部分加载)实现方案

直接调用model.load_state_dict(torch.load(PATH))报错的核心原因是该方法默认要求预训练权重的所有键、参数维度和当前模型完全一致,只要存在键不匹配、维度不符、多/少参数的情况就会抛出异常,在模型结构修改、加载第三方预训练权重、适配多卡训练权重的场景下很常见。

核心逻辑

通过手动遍历过滤预训练权重,只保留和当前模型匹配的参数存入新的OrderedDict,再加载到模型中,即可实现部分权重加载,完整可运行示例如下:

import torch
import torch.nn as nn
from collections import OrderedDict
import copy

# ---------------------- 1. 定义目标模型(实际场景为你自己需要加载权重的模型) ----------------------
class TargetModel(nn.Module):
    def __init__(self, num_classes=10):
        super().__init__()
        # 特征提取层(和预训练权重结构一致)
        self.features = nn.Sequential(
            nn.Conv2d(3, 64, kernel_size=3, padding=1),
            nn.ReLU(),
            nn.MaxPool2d(2),
            nn.Conv2d(64, 128, kernel_size=3, padding=1),
            nn.ReLU(),
            nn.MaxPool2d(2)
        )
        # 自定义分类头(预训练权重中不存在该部分参数)
        self.classifier = nn.Linear(128 * 8 * 8, num_classes)
    
    def forward(self, x):
        x = self.features(x)
        x = x.flatten(1)
        x = self.classifier(x)
        return x

# ---------------------- 2. 模拟预训练权重(实际场景替换为从本地加载pth文件即可) ----------------------
pretrain_model = nn.Sequential(
    OrderedDict([
        ('features', nn.Sequential(
            nn.Conv2d(3, 64, kernel_size=3, padding=1),
            nn.ReLU(),
            nn.MaxPool2d(2),
            nn.Conv2d(64, 128, kernel_size=3, padding=1),
            nn.ReLU(),
            nn.MaxPool2d(2)
        ))
    ])
)
# 模拟保存预训练权重到本地
torch.save(pretrain_model.state_dict(), "pretrain_weights.pth")

# ---------------------- 3. 核心自定义加载逻辑 ----------------------
# 初始化目标模型
target_model = TargetModel(num_classes=10)
# 加载预训练权重,统一转cpu处理避免设备不匹配问题
pretrain_weights = torch.load("pretrain_weights.pth", map_location="cpu")
# 初始化新的参数字典
new_state_dict = OrderedDict()

for k, v in pretrain_weights.items():
    # 可根据需求自定义过滤逻辑,比如处理多卡权重前缀:k = k.replace("module.", "")
    # 只保留目标模型中存在、且参数维度匹配的键
    if k in target_model.state_dict() and target_model.state_dict()[k].shape == v.shape:
        # deepcopy避免修改原始预训练权重的内存值,不需要保留原始权重可省略
        new_state_dict[k] = copy.deepcopy(v)
    else:
        print(f"跳过不匹配参数:{k}")

# 加载处理后的权重,strict=False允许部分参数不加载
missing_keys, unexpected_keys = target_model.load_state_dict(new_state_dict, strict=False)
print(f"未加载的模型缺失参数:{missing_keys}")
print(f"预训练中未用到的多余参数:{unexpected_keys}")

注意事项

  • 如果加载的是DDP多卡训练保存的权重,键会自带module.前缀,可在遍历的时候新增k = k.replace("module.", "")去掉前缀后再匹配
  • 仅需要加载特定层的场景,可以在遍历的时候加自定义判断逻辑,比如if k.startswith("features")只加载特征提取层
  • 打印missing_keys和unexpected_keys可确认加载结果是否符合预期,避免漏加载需要的参数

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.27 14:06:03