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

PyTorch修改预训练模型层后加载权重形状不匹配的解决方法

PyTorch修改预训练模型并正确加载权重方案

报错根因

触发形状不匹配断言的核心问题有两个:

  • 现有权重加载逻辑鲁棒性极差:靠参数在state_dict中的排列索引一一对应匹配,只要模型增删模块、PyTorch版本差异导致参数排序变化,就会出现键值错位。预训练权重中保留了你打算删除的norm、pre_logits、head层参数,无论你是否在模型__init__中删除这几个模块,按索引遍历都会出现形状不匹配的问题。
  • 修改后的前向传播逻辑存在代码bug:一是forward函数中调用forward_features时传入了未定义的参数z;二是每个stage的block循环中,没有将前一个block的输出作为后一个block的输入,特征计算逻辑完全错误。

正确实现步骤

1. 修正模型前向逻辑

不需要在__init__中强制删除不需要的层(只要前向传播不调用就不会生效),直接修正forward_features和forward方法即可,正确输出四个stage的特征:

def forward_features(self, x):
    # Stage 1
    x = self.patch_embed1(x)
    x = self.pos_drop(x)
    for blk in self.blocks1:
        if self.use_checkpoint:
            x = checkpoint.checkpoint(blk, x)
        else:
            x = blk(x)
    features1 = x

    # Stage 2
    x = self.patch_embed2(x)
    for blk in self.blocks2:
        if self.use_checkpoint:
            x = checkpoint.checkpoint(blk, x)
        else:
            x = blk(x)
    features2 = x

    # Stage 3
    x = self.patch_embed3(x)
    for blk in self.blocks3:
        if self.use_checkpoint:
            x = checkpoint.checkpoint(blk, x)
        else:
            x = blk(x)
    features3 = x

    # Stage 4
    x = self.patch_embed4(x)
    for blk in self.blocks4:
        if self.use_checkpoint:
            x = checkpoint.checkpoint(blk, x)
        else:
            x = blk(x)
    features4 = x

    return features1, features2, features3, features4

def forward(self, x):
    x = x[0]
    features1, features2, features3, features4 = self.forward_features(x)
    return features1, features2, features3, features4

2. 替换原有错误的权重加载逻辑

不要自己写循环按索引匹配参数,用PyTorch内置的load_state_dict方法,按键名匹配参数,同时过滤掉不需要的预训练权重,设置strict=False自动跳过不匹配的参数:

if pretrained:
    print('Loading weights...')
    weight_dict = torch.load(
        os.path.join('models', 'uniformer_small_k400_16x4.pth'),
        map_location='cpu'
    )
    # 过滤掉不需要的分类头、归一化层参数
    filtered_weights = {}
    for k, v in weight_dict.items():
        if k.startswith(('norm.', 'pre_logits.', 'head.')):
            continue
        filtered_weights[k] = v
    # 加载权重,strict=False允许存在缺失/多余的参数键
    load_res = self.featureExtractor.load_state_dict(filtered_weights, strict=False)
    # 打印加载信息确认
    print(f"Pretrain weight skipped keys: {load_res.unexpected_keys}")
    print(f"Model uninitialized keys: {load_res.missing_keys}")
    print('Loading done!')

说明

  • 这种按键名过滤+strict=False的加载方式是PyTorch修改预训练模型的标准写法,完全不会出现参数错位的问题,比自己写循环遍历断言可靠得多。
  • 如果你后续不需要使用norm、pre_logits、head层,可以直接在模型__init__方法中注释掉这三个模块的初始化代码,加载时就不会出现这几个层的未初始化提示。
  • 如果后续需要用中间层输出做下游任务,四个返回的features1到features4就是四个stage下采样后的特征图,直接使用即可。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 17:15:44