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
相关产品推荐
相关产品推荐

