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

