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

如何在PyTorch中为修改后的预训练模型加载权重时使用strict=False

解决方案

你可以直接替换掉原来手动遍历权重的冗余逻辑,改用PyTorch原生的load_state_dict方法,这个方法本身就支持strict=False参数,完美适配你修改过模型结构的场景。修改后的完整代码如下:

import os
import torch
import torch.nn as nn

class MyClass(nn.Module):
    def __init__(self, pretrained=False):
        super(MyClass, self).__init__()
        self.encoder = S3D_featureExtractor_multi_output()
        if pretrained:
            weight_dict = torch.load(os.path.join('models', 'weights.pt'))
            # 使用load_state_dict并设置strict=False
            missing_keys, unexpected_keys = self.encoder.load_state_dict(weight_dict, strict=False)
            # 可选:打印不匹配的键,方便排查
            if missing_keys:
                print(f"模型中存在但权重里没有的层: {missing_keys}")
            if unexpected_keys:
                print(f"权重里存在但模型中没有的层: {unexpected_keys}")
            print('Loading finished!')
                
    def forward(self, x):
        a, b = self.encoder(x)
        return a, b

关键修改说明:

  1. 删掉了原来手动遍历权重字典、逐个复制张量的代码,这部分逻辑完全可以被load_state_dict替代,更简洁高效
  2. 调用self.encoder.load_state_dict(weight_dict, strict=False)时,strict=False会自动忽略两种不匹配情况:
    • 当前模型有但预训练权重里没有的层(missing_keys)
    • 预训练权重里有但当前模型已经删掉/修改的层(unexpected_keys)
  3. 可选保留打印不匹配键的逻辑,方便你确认哪些层没有被加载,排查模型修改是否符合预期

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.24 20:36:27