ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

从零手写PyTorch ResNet:代码详解、权重对齐与训练避坑

从零手写PyTorch ResNet:代码详解、权重对齐与训练避坑 第一次把ResNet论文摊在屏幕左边、编辑器放在右边准备动手写代码的时候我卡在一个特别不起眼的地方shortcut分支到底什么时候该挂一个1x1卷积。论文里只有一句when the dimensions increase, we perform a linear projection官方实现里就是一个if判断但那个if少写任何一个条件网络都会在第二个stage直接抛出size mismatch。后来我把整条下采样链路、权重初始化、参数量对账全部手工推了一遍才算真正把ResNet代码复现这件事吃透。这篇内容就是那次复现过程的完整整理。我会从零用PyTorch实现ResNet18/34/50/101/152重点是代码注释会写得非常细细到每一行为什么这么写、不这么写会出什么问题都讲清楚。同时补齐三件容易被忽略的事权重初始化怎么对齐原论文、怎么和torchvision官方实现逐层对拍、以及真正跑训练时才会遇到的坑。适合已经会写基础卷积网络、想把主干网络从会调用推进到能改的读者也适合正在准备多模态模型代码复现、检测分割代码复现、需要自己改Backbone结构的人。1. 手写一遍ResNet比直接调torchvision多拿到什么1.1 直接调用现成模型的三笔隐性成本torchvision.models.resnet50(pretrainedTrue)这一行确实是效率之王但它把三样东西藏起来了而这三样恰好是后面所有工作的地基。第一笔成本是结构黑盒。当你需要把ResNet作为多模态模型的视觉编码器、需要在layer3和layer4之间插入一个注意力模块、或者需要把最后的全局池化换成多尺度池化时你必须先知道layer1到layer4里每个block的输入输出通道分别是多少、哪几个block带stride。这些信息官方文档只给了一张粗糙的图实际数值得靠读源码或者自己打印。第二笔成本是排错路径断裂。真实项目里最常见的报错不是模型跑不起来而是训练loss不降梯度出现NaN加了自定义模块后shape对不上。这些问题的定位需要你对网络内部张量流动有肌肉记忆。只会调用的人这时候只能靠反复打印shape一点点试效率极低。第三笔成本是改造时的版本风险。官方实现本身在演进比如Bottleneck里stride放在第一个1x1还是第二个3x3不同时期、不同来源的实现并不一致这个差异会直接影响下采样后的特征图尺寸。自己写一遍你会对每一行负责。1.2 复现的验收标准四重对齐才算通过很多人理解的复现成功是forward能跑通、loss能下降。这个标准太松了因为一个结构写错的ResNet照样能训练只是精度低几个点你根本发现不了。我给自己定的验收标准是四重对齐形状对齐输入(1,3,224,224)逐层的输出shape必须和官方一致参数量对齐ResNet50在1000分类下必须是25,557,032差一个数就说明某处结构写错了权重键名对齐自己定义的state_dict的key集合必须和官方完全一致这样才能直接load_state_dict数值对齐加载同一份预训练权重后自己的模型和官方模型在同一个随机输入上的输出最大绝对误差应该在1e-5量级。第三条和第四条是杀手锏。只要键名集合完全一致strictTrue的加载就不会报错只要数值误差足够小就证明你的卷积顺序、BN位置、激活位置、池化配置全都对。下面几节会一步步把这两个验证做出来。1.3 动手前的环境确认版本对不上会浪费一整天环境这块踩坑的概率比你想象的高。PyTorch和torchvision是强绑定的版本错配的典型症状是import torchvision直接报undefined symbol。下面是常见版本的对应关系照着选基本不会出问题PyTorchtorchvision支持的Python2.10.163.8 - 3.112.20.173.8 - 3.122.30.183.8 - 3.122.40.193.9 - 3.12用conda建环境的时候我习惯先确认三件事python -c import torch; print(torch.__version__, torch.cuda.is_available())、python -c import torchvision; print(torchvision.__version__)、以及torch.version.cuda和驱动是否匹配。如果只是学习复现CPU版本完全够用ResNet18在224分辨率下的前向传播CPU也能秒级完成只有当你要跑完整训练时GPU才有明显收益。提示不要用pip install torch裸装它会拉一个和你的CUDA驱动可能不匹配的默认版本。按官方页面上给出的组合命令来装装完立刻跑上面那三行确认。2. 残差结构到底解决了什么问题而不只是加深了网络2.1 退化现象56层网络为什么反而比20层差ResNet之前大家已经发现了一个反直觉的现象把网络从20层加到56层训练误差反而更高。注意是训练误差不是测试误差。这就排除了过拟合——过拟合会让训练误差降低、测试误差升高而这里两个都在涨。这个现象叫退化。它说明问题不在模型容量不够而在优化器找不到那个更深的解。理论上56层网络完全可以模拟20层网络只要把多出来的36层全部学成恒等映射效果至少不会变差。既然理论上存在这样一个不劣于浅层网络的解为什么SGD找不到原因在于让一堆非线性层去逼近恒等映射本身就是个很难的优化问题。多个非线性层叠加后输入信号会被反复拉伸压缩想让它恰好等于输入需要权重精确配合这在随机初始化下几乎不可能自发发生。2.2 把H(x)拆成F(x)x之后优化目标发生了什么变化ResNet的核心改动只有一行公式。原来我们让几个卷积层直接拟合目标映射H(x)残差结构改成H(x) F(x) x即让卷积层去拟合残差F(x)最终输出再加回输入x。这一改优化难度发生了质的变化。继续用恒等映射的例子如果最优解就是恒等映射那么H(x)x对应的F(x)0。让一堆卷积层的输出逼近0比让它们逼近恒等映射容易得多——只要把权重往0推就行这是SGD非常擅长的事情。换句话说残差结构把学习一个复杂映射转化成了学习一个相对于恒等映射的偏移量而偏移量在很多层里本来就接近0。在实际训练中你会观察到训练好的ResNet里大量残差分支的权重范数确实很小。这不是训练不充分而是网络主动选择了接近恒等映射的解只在真正需要提特征的地方才让残差分支发挥作用。2.3 梯度视角那个恒为1的加法项才是关键从反向传播看残差连接的价值更直接。设损失为L考虑第l层的输出x_l前向是x_L x_l sum(F_i(x_i))沿链式法则往回传时dL/dx_l dL/dx_L * (1 d(sum F_i)/dx_l)括号里那个1来自恒等路径。它意味着无论中间多少层的梯度有多小梯度回传到任意一层时至少还保留着从最顶层直接传下来的那一份不会被连乘衰减到0。这也是为什么ResNet能把网络堆到152层甚至更深。相比之下普通卷积网络反向传播时梯度是纯粹的连乘形式几十层之后基本就消失了。理解这一点你就能明白为什么后来大量结构DenseNet、Transformer的残差连接、各种U-Net的skip都在沿用这个思路。2.4 BasicBlock与Bottleneck两种残差块的计算账ResNet实际用了两种残差块它们针对的是不同的深度需求。BasicBlock是两层3x3卷积通道数不变结构简单用在ResNet18和ResNet34上。它的参数量是2 * 3 * 3 * C * C 18C²当C64时约7.4万。Bottleneck是三层的1x1降维 - 3x3卷积 - 1x1升维结构。以输入256通道、中间宽度64为例第一个1x1把256压到64第二个3x3在64通道上做空间卷积第三个1x1再把64升到256。它的参数量是256*64 9*64*64 64*256 69632而如果直接用两层3x3在256通道上做参数量是2*9*256*256 1179648差了将近17倍。所以Bottleneck的设计动机很明确在保持表达能力的同时把深层网络的参数量和计算量压下来。这也是为什么ResNet50的参数量25.56M只比ResNet3421.80M多一点点但深度深了一倍多。版本block类型layers配置参数量224输入的FLOPsResNet18BasicBlock[2,2,2,2]11.69M1.82GResNet34BasicBlock[3,4,6,3]21.80M3.66GResNet50Bottleneck[3,4,6,3]25.56M4.11GResNet101Bottleneck[3,4,23,3]44.55M7.85GResNet152Bottleneck[3,8,36,3]60.19M11.58G注意Bottleneck有个expansion4的属性因为输出通道是中间宽度的4倍。这个属性会贯穿整个实现用来把block内部宽度换算成实际输出通道忘掉它是最常见的错误来源之一。3. 从零搭建ResNet每一行为什么这么写3.1 卷积、批归一化、激活的顺序与bias取舍先说一个几乎所有人第一次写都会忽略的细节带BatchNorm的卷积层bias应该设成False。原因是BN本身会做减均值除标准差的归一化紧接着又有可学习的缩放和平移参数γ和β。卷积的bias加在BN之前的特征上经过BN的减均值操作后会被完全抵消掉等于白加一个参数。所以# biasFalse 不是省参数而是因为后面的BN会把bias的作用完全吃掉 self.conv1 nn.Conv2d(in_channels, out_channels, kernel_size3, stridestride, padding1, biasFalse) self.bn1 nn.BatchNorm2d(out_channels)顺序上原论文走的是Conv - BN - ReLU这个顺序在PyTorch里也是社区共识。有人会问能不能换成BN - Conv - ReLU可以但那是另外的结构设计不要和标准ResNet混着用否则权重键名和数值都无法和官方对齐。nn.ReLU(inplaceTrue)里的inplace是个值得说清楚的点。它直接修改输入张量的内存省一份激活值的内存占用。在ResNet这种结构里ReLU的输入在后续不需要再用到所以inplace是安全的。但如果你在残差分支里手动加了额外的分支、或者需要拿到ReLU之前的特征做辅助损失就必须把它改成False否则会报a leaf Variable that requires grad is being used in an in-place operation或者更隐蔽的数值错误。3.2 BasicBlockstride到底该放在哪个卷积上BasicBlock的实现如下注释写得很细import torch import torch.nn as nn class BasicBlock(nn.Module): ResNet18/34 使用的基础残差块两层 3x3 卷积 # 输出通道相对于输入中间通道的倍数。 # BasicBlock 通道数不变所以是 1Bottleneck 是 4。 expansion 1 def __init__(self, inplanes, planes, stride1, downsampleNone, norm_layerNone): super().__init__() if norm_layer is None: norm_layer nn.BatchNorm2d # 第一个 3x3负责空间尺寸变化stride和通道变化 # padding1 保证 stride1 时空间尺寸不变k3, p1, s1 - 尺寸不变 self.conv1 nn.Conv2d(inplanes, planes, kernel_size3, stridestride, padding1, biasFalse) self.bn1 norm_layer(planes) # 第二个 3x3stride 固定为 1只做通道内的特征组合 # 如果把 stride 放在这里残差分支和恒等分支的尺寸依然对得上 # 但官方实现是放在 conv1 上为了权重键名和数值对齐必须保持一致 self.conv2 nn.Conv2d(planes, planes, kernel_size3, stride1, padding1, biasFalse) self.bn2 norm_layer(planes) self.relu nn.ReLU(inplaceTrue) # 下采样分支当主分支的尺寸或通道数发生变化时 # 恒等分支需要同步做变换才能相加 self.downsample downsample self.stride stride def forward(self, x): # 先把输入存下来作为恒等分支的起点 identity x out self.conv1(x) out self.bn1(out) out self.relu(out) out self.conv2(out) out self.bn2(out) # 注意这里不加 ReLU要等相加之后再加 # 如果尺寸或通道变了恒等分支也要跟着变 if self.downsample is not None: identity self.downsample(x) # 核心的一行残差相加。 # 用 而不是 out out identity可以省一次中间张量的显存 out identity out self.relu(out) return out有两个地方必须单独强调。第一个是第二个BN后面不加ReLU。很多人凭直觉写成conv-bn-relu三件套连用两遍结果是残差分支的输出被ReLU截断成非负那么F(x)x里的F(x)就永远不可能为负恒等映射F0仍然可以学到但表达能力被削弱了实测精度会掉零点几个点。原论文的加法发生在第二个BN之后、ReLU之前这个位置不能动。第二个是stride放在conv1上。这个选择和官方实现一致。历史上还有另一种变体把stride放在conv2上也就是先用stride1降通道再用stride2降尺寸两者都能跑通且shape一致但数值不同权重不能直接换用。做复现时以官方为准。3.3 Bottleneck1x1降维再升维的通道账Bottleneck多了一层需要特别小心通道数的转换class Bottleneck(nn.Module): ResNet50/101/152 使用的瓶颈残差块1x1 降维 - 3x3 - 1x1 升维 # 输出通道是中间宽度的 4 倍这个 4 在 ResNet 主体里到处都会用到 expansion 4 def __init__(self, inplanes, planes, stride1, downsampleNone, norm_layerNone): super().__init__() if norm_layer is None: norm_layer nn.BatchNorm2d # width 是 block 内部的工作通道数也就是 1x1 降维后的宽度 width planes # 第一个 1x1只做通道压缩不改变空间尺寸所以 kernel1, padding0 self.conv1 nn.Conv2d(inplanes, width, kernel_size1, biasFalse) self.bn1 norm_layer(width) # 中间的 3x3承担实际的空间特征提取stride 放在这里 self.conv2 nn.Conv2d(width, width, kernel_size3, stridestride, padding1, biasFalse) self.bn2 norm_layer(width) # 第三个 1x1把宽度升回 planes * 4 self.conv3 nn.Conv2d(width, planes * self.expansion, kernel_size1, biasFalse) self.bn3 norm_layer(planes * self.expansion) self.relu nn.ReLU(inplaceTrue) self.downsample downsample self.stride stride def forward(self, x): identity x out self.conv1(x) out self.bn1(out) out self.relu(out) out self.conv2(out) out self.bn2(out) out self.relu(out) out self.conv3(out) out self.bn3(out) # 同样不加 ReLU if self.downsample is not None: identity self.downsample(x) out identity out self.relu(out) return out注意Bottleneck的stride放在中间的3x3卷积上而BasicBlock放在第一个3x3上——两者位置不同但逻辑是一致的都放在那个真正做空间采样的3x3上。这就是常说的ResNet v1.5配置也是torchvision采用的做法。原论文最早的版本是把stride放在第一个1x1上那样会损失更多信息现在基本没人用了。提示planes这个参数在不同实现里的含义不一样。在torchvision里planes指的是block内部的宽度对Bottleneck而言是输出通道的1/4而不是输出通道。写_make_layer的时候如果把这个搞混通道数会以4的倍数滚雪球式放大最后在fc层报一个巨大的size mismatch。3.4 _make_layer与下采样分支的触发条件这是整个实现里最需要小心的一段。下采样分支的触发条件有两个用or连接def _make_layer(self, block, planes, blocks, stride1): 构建一个 stage由 blocks 个残差块串联而成 block : BasicBlock 或 Bottleneck 类 planes : block 内部宽度输出通道 planes * block.expansion blocks : 这个 stage 里残差块的数量 stride : 第一个残差块的步长控制这个 stage 是否降采样 norm_layer self._norm_layer downsample None # 条件一stride ! 1空间尺寸会变恒等分支必须跟着降采样 # 条件二inplanes ! planes * block.expansion通道数会变恒等分支必须跟着换通道 # 两个条件只要满足一个就必须挂 downsample if stride ! 1 or self.inplanes ! planes * block.expansion: downsample nn.Sequential( # 1x1 卷积同时完成通道变换和空间降采样stride 与主分支保持一致 nn.Conv2d(self.inplanes, planes * block.expansion, kernel_size1, stridestride, biasFalse), norm_layer(planes * block.expansion), ) layers [] # 第一个 block可能带 stride 和 downsample layers.append(block(self.inplanes, planes, stride, downsample, norm_layer)) # 更新 inplanes后续 block 的输入通道就变了 self.inplanes planes * block.expansion # 其余 blockstride1、downsampleNone输入输出通道一致残差直连 for _ in range(1, blocks): layers.append(block(self.inplanes, planes, norm_layernorm_layer)) return nn.Sequential(*layers)只判断stride是最经典的错误。在标准ResNet18/34/50的配置下每个stage的第一个block恰好都满足stride2所以只判断stride也能跑通你会误以为代码是对的。但一旦你尝试自定义宽度、或者把ResNet50的配置改成奇数通道通道数变化但stride1的情况就出现了然后直接报错。这个坑我在改Backbone做检测任务时踩过一次排查了两个小时。判断条件里的planes * block.expansion不能省。对Bottleneck来说stage1的输出是256通道而传进来的planes是64如果写成self.inplanes ! planes那ResNet50第一个stage会错误地挂上downsample参数量直接对不上。3.5 ResNet主体stem设计与stage堆叠主体部分把stem、四个stage、分类头串起来class ResNet(nn.Module): def __init__(self, block, layers, num_classes1000, norm_layerNone): super().__init__() if norm_layer is None: norm_layer nn.BatchNorm2d self._norm_layer norm_layer # inplanes 是一个会被 _make_layer 修改的状态变量 # 记录下一个 stage 的输入通道数 self.inplanes 64 # stem7x7 大卷积 stride 2把 224 快速降到 112 # 大卷积核在浅层能覆盖更大的感受野配合 maxpool 一共降 4 倍 self.conv1 nn.Conv2d(3, self.inplanes, kernel_size7, stride2, padding3, biasFalse) self.bn1 norm_layer(self.inplanes) self.relu nn.ReLU(inplaceTrue) # kernel3, stride2, padding1尺寸减半且不丢边角信息 self.maxpool nn.MaxPool2d(kernel_size3, stride2, padding1) # 四个 stage通道依次翻倍空间尺寸依次减半 self.layer1 self._make_layer(block, 64, layers[0], stride1) self.layer2 self._make_layer(block, 128, layers[1], stride2) self.layer3 self._make_layer(block, 256, layers[2], stride2) self.layer4 self._make_layer(block, 512, layers[3], stride2) # 自适应池化无论输入分辨率多少都压成 1x1这样模型能接受任意尺寸输入 self.avgpool nn.AdaptiveAvgPool2d((1, 1)) self.fc nn.Linear(512 * block.expansion, num_classes) # 权重初始化下一节展开讲 for m in self.modules(): if isinstance(m, nn.Conv2d): nn.init.kaiming_normal_(m.weight, modefan_out, nonlinearityrelu) elif isinstance(m, (nn.BatchNorm2d, nn.GroupNorm)): nn.init.constant_(m.weight, 1) nn.init.constant_(m.bias, 0) # 残差分支末端的 BN 权重置零 for m in self.modules(): if isinstance(m, Bottleneck) and m.bn3.weight is not None: nn.init.constant_(m.bn3.weight, 0) elif isinstance(m, BasicBlock) and m.bn2.weight is not None: nn.init.constant_(m.bn2.weight, 0) def forward(self, x): x self.conv1(x) x self.bn1(x) x self.relu(x) x self.maxpool(x) x self.layer1(x) x self.layer2(x) x self.layer3(x) x self.layer4(x) x self.avgpool(x) # flatten(1) 把 (N, C, 1, 1) 展成 (N, C)用 flatten 比 view 更稳 x torch.flatten(x, 1) x self.fc(x) return x最后是构造函数def _resnet(block, layers, **kwargs): model ResNet(block, layers, **kwargs) return model def resnet18(**kwargs): return _resnet(BasicBlock, [2, 2, 2, 2], **kwargs) def resnet34(**kwargs): return _resnet(BasicBlock, [3, 4, 6, 3], **kwargs) def resnet50(**kwargs): return _resnet(Bottleneck, [3, 4, 6, 3], **kwargs) def resnet101(**kwargs): return _resnet(Bottleneck, [3, 4, 23, 3], **kwargs) def resnet152(**kwargs): return _resnet(Bottleneck, [3, 8, 36, 3], **kwargs)4. 权重初始化和上线前的自检清单4.1 Kaiming初始化为什么必须配fan_outResNet用nn.init.kaiming_normal_(w, modefan_out, nonlinearityrelu)。这个选择不是随便定的。Kaiming初始化的推导目标是让每一层输出的方差和输入方差保持一致。它根据前一层激活值的分布来缩放权重标准差对于ReLU正确做法是把权重标准差设成sqrt(2 / fan_in)因为ReLU会把约一半的激活值压成0方差减半需要额外乘2补偿回来。那为什么modefan_outfan_in是输入通道数乘以卷积核面积fan_out是输出通道数乘以卷积核面积。理论上前向传播保方差用fan_in反向传播保方差用fan_out。在卷积层的权重形状是(out_channels, in_channels, k, k)PyTorch默认的modefan_in会取in_channels*k*k。而官方ResNet用的是fan_out这属于实践中的调优选择对深层网络的反向稳定性更友好。复现时照着官方写就对了。BN层初始化为γ1、β0这是标准操作。γ1意味着初始状态下BN近似恒等变换在归一化之后β0意味着没有额外偏移。4.2 残差分支末端BN置零让网络开局就是恒等映射上面代码里最后那段循环把每个残差块的最后一个BN层的γ初始化为0这个操作来自后来的Bag of Tricks系列工作但它在实践中的收益非常明显。原理很简单BN的输出是γ * x_hat βγ0时输出恒为0。这意味着每个残差块的输出一开始就是F(x) x 0 x x整个网络初始状态就是一堆恒等映射的串联。前面讨论过这正是ResNet希望网络从什么地方开始——一个不劣于浅层网络的起点。好处是训练初期非常稳。尤其是当你从零训练而没有加载预训练权重时加上这个零初始化前几个epoch的loss下降会明显更平滑也不需要很长的warmup去救。我在CIFAR上从零训ResNet18做过对比不加零初始化时前两个epoch loss震荡比较厉害加了之后曲线明显干净。注意一下代码里的写法它和官方实现有个细微差别。官方是把所有模块的权重先按类型初始化一遍然后在同一个循环里判断block类型再置零。我这里分成两个循环逻辑上等价但可读性更好。另外那个m.bn3.weight is not None的判断是为了兼容某些把BN换成无参数归一化的配置标准情况下可以省略。提示如果你是在加载预训练权重之后做微调这个零初始化会被权重文件覆盖掉所以不用担心它影响微调效果。它只在从零训练时起作用。4.3 形状自检用hook打印每一层的输出写完模型第一件事不是训练是自检。我习惯用一个forward hook把每个子模块的输出shape打出来import torch def trace_shapes(model, x): 注册 forward hook记录每个顶层子模块的输出形状 records {} def make_hook(name): def hook(module, inputs, output): if isinstance(output, torch.Tensor): records[name] tuple(output.shape) return hook handles [] for name, module in model.named_children(): handles.append(module.register_forward_hook(make_hook(name))) model.eval() with torch.no_grad(): model(x) for h in handles: h.remove() # 记得移除 hook否则会累积 return records model resnet50(num_classes1000) x torch.randn(1, 3, 224, 224) for name, shape in trace_shapes(model, x).items(): print(f{name:10s} - {shape})ResNet50在224输入下的期望输出是模块输出形状说明conv1(1, 64, 112, 112)7x7 stride2尺寸减半maxpool(1, 64, 56, 56)3x3 stride2再减半layer1(1, 256, 56, 56)3个Bottleneck通道升到256尺寸不变layer2(1, 512, 28, 28)第一个block stride2layer3(1, 1024, 14, 14)同上layer4(1, 2048, 7, 7)同上avgpool(1, 2048, 1, 1)全局池化fc(1, 1000)全连接分类头这里有两个容易记错的数字。第一layer1虽然叫第一层但它的输出通道是256而不是64因为有expansion4。第二layer1的尺寸保持56不变只在通道上做变换所有ResNet的stride1都给了第一个stage。记住这两点改结构时就不会乱。4.4 参数量对账差一个数字就是错def count_params(model, only_trainableTrue): if only_trainable: return sum(p.numel() for p in model.parameters() if p.requires_grad) return sum(p.numel() for p in model.parameters()) for name, builder in [(resnet18, resnet18), (resnet34, resnet34), (resnet50, resnet50), (resnet101, resnet101), (resnet152, resnet152)]: m builder(num_classes1000) print(f{name:10s} params {count_params(m):,})期望输出分别是11,689,512 / 21,797,672 / 25,557,032 / 44,549,160 / 60,192,808。如果ResNet50打出来是25,557,032基本可以确认结构完全正确如果差了64或者63这类小数通常是某个地方多挂或少挂了一个downsample如果差了几百万多半是expansion用错了。4.5 与torchvision逐层对拍形状和参数量都对上了最后做一次数值对拍这是最硬的验证import torch import torchvision def convert_state_dict(sd): 去掉 DataParallel 的 module. 前缀便于直接加载 new_sd {} for k, v in sd.items(): if k.startswith(module.): k k[7:] new_sd[k] v return new_sd # 我的实现 my_model resnet50(num_classes1000) # 官方实现 预训练权重 ref_model torchvision.models.resnet50(weightsNone) official_sd torchvision.models.resnet50( weightstorchvision.models.ResNet50_Weights.IMAGENET1K_V1).state_dict() official_sd convert_state_dict(official_sd) # strictTrue 会严格比对键名集合key 不一致会直接报错并列出差异 missing, unexpected my_model.load_state_dict(official_sd, strictTrue) my_model.eval() ref_model.load_state_dict(official_sd) ref_model.eval() torch.manual_seed(0) x torch.randn(2, 3, 224, 224) with torch.no_grad(): y_mine my_model(x) y_ref ref_model(x) print(max abs diff , (y_mine - y_ref).abs().max().item()) print(shape:, y_mine.shape)如果打印出来的最大绝对误差在1e-5量级浮点累加误差的正常范围说明你的实现和官方在结构上完全等价。如果误差是0.1或者更大几乎肯定是某个激活位置或者BN位置错了。strictTrue这一步的报错信息很值钱。如果出现Missing key(s)说明你的模型有官方没有的层Unexpected key(s)说明官方有你没有的层。常见的情况包括忘记给downsample里的BN命名、BN层的命名顺序和官方不同、或者给fc层加了额外的Dropout。这些都能通过键名差异直接定位。5. 真正跑起来之后才会踩到的坑5.1 跳过downsample条件判断的真实排查链路我在做自定义宽度实验时遇到过这么一次报错RuntimeError: The size of tensor a (64) must match the size of tensor b (128) at non-scalar dimension 1排查过程是这样的第一反应是查看是哪个加法出的问题。报错栈里指向out identity但这个加法在几十个block里都有得往前找上下文。我加了一行打印把每个block的输入输出通道和stride打出来def forward(self, x): identity x ... if out.shape ! identity.shape: print(fshape mismatch: out{tuple(out.shape)}, fidentity{tuple(identity.shape)}, stride{self.stride}) out identity打印后立刻定位到某个stage的第一个blockout是(1,128,28,28)identity是(1,64,28,28)。通道差了2倍但stride是1。这说明主分支的通道变了但downsample是None。回到_make_layer一看条件写的是if stride ! 1。而我的自定义配置里第二个stage的宽度是64也就是输出通道128expansion4输入是64通道变化了但stride是1条件没命中。修复就是把条件补全if stride ! 1 or self.inplanes ! planes * block.expansion: downsample nn.Sequential(...)这件事给我一个习惯改任何结构参数之后先跑一遍形状自检和参数量对账比等训练报错快得多。后来我把这个检查写成了一个单元测试每次改动后跑一遍几秒钟出结果。5.2 BN、batch size与学习率的三方联动BatchNorm在训练时用当前batch的统计量推理时用滑动平均的统计量。这意味着训练batch越小BN的统计量越不准模型精度越差。这不是一个可以忽略的细节。标准ResNet训练用的是batch size 256、初始学习率0.1。如果你只有一张8G显存的卡跑不了batch 256有三条路梯度累积用小batch前向反向多次累积梯度后再更新一次。这样等效batch还是256BN统计量的问题依然存在但对优化轨迹的模拟是最接近的。换成GroupNorm或者SyncBNGroupNorm不受batch size影响但要注意加载预训练权重时BN的滑动统计量会被丢掉从零训练才合适。降低学习率并延长训练小batch下学习率要按比例缩减。经验公式是lr 0.1 * batch_size / 256小batch下还可以用sqrt缩放即0.1 * sqrt(batch_size/256)后者更保守一些。另外BN的momentum默认是0.1表示滑动平均的更新比例。小batch时统计量本身噪声大可以把momentum调小到0.01让它更新得更慢更平滑。这个改动收益不大但也不会有坏处属于可用可不用的选项。还有一个常见错误是把model.eval()忘了。训练完之后做验证不切eval模式BN会用验证batch的统计量同时Dropout还在起作用验证精度会莫名偏低。我在早期项目里因为这个原因白调了两天超参。5.3 预训练权重加载的键名不对齐问题加载预训练权重时最常见的两种报错及处理方式报错原因处理方式Missing key(s): fc.weight, fc.bias自己的分类数不是1000fc层被替换过滤掉fc相关键或strictFalseUnexpected key(s): module.conv1.weight权重存的时候包了一层DataParallel去掉module.前缀Missing key(s): layer1.0.downsample.1.weightdownsample里的BN命名不一致检查downsample是否用了nn.Sequential包装size mismatch for fc.weight分类数不同加载前删掉fc键或新建fc层改分类数的标准做法是加载前先过滤def load_backbone(model, ckpt_path, num_classes): sd torch.load(ckpt_path, map_locationcpu) sd convert_state_dict(sd) # 丢掉与分类头相关的键避免 num_classes 不一致导致 size mismatch sd {k: v for k, v in sd.items() if not k.startswith(fc.)} missing, unexpected model.load_state_dict(sd, strictFalse) print(missing:, missing) print(unexpected:, unexpected) return model注意strictFalse只是让加载不报错它不会告诉你哪些权重真的用上了。生产环境一定要把missing和unexpected打印出来看一眼。如果missing里出现了layer1.0.conv1.weight这种主干层的键说明你的权重文件和模型结构不匹配这时候必须停下来查清楚不能带着随机初始化的主干继续训练。还有一点值得提醒torch.load在新版本里默认weights_onlyTrue如果权重文件里除了张量还存了别的东西可能会加载失败需要在可信文件的前提下显式设置weights_onlyFalse。5.4 把ImageNet结构硬搬到小分辨率数据集ResNet原始的stem是针对224输入设计的7x7 stride2加上3x3 stride2的maxpool一共降采样4倍。224进来变成56x56刚好算力可控。但如果你的数据集是CIFAR的32x32这套stem就完全不适合了。32进来经过stem变成8x8再经过四个stage的降采样变成0.5直接崩掉。CIFAR版的标准改法是def resnet_for_cifar(num_classes10): model ResNet(BasicBlock, [2, 2, 2, 2], num_classesnum_classes) # 换成 3x3 stride1避免在 32x32 上过度降采样 model.conv1 nn.Conv2d(3, 64, kernel_size3, stride1, padding1, biasFalse) # 直接去掉 maxpool让 layer1 工作在 32x32 model.maxpool nn.Identity() # 替换后的 conv1 需要重新初始化否则用的是 PyTorch 默认初始化 nn.init.kaiming_normal_(model.conv1.weight, modefan_out, nonlinearityrelu) return model这里有三个细节容易出错。第一nn.Identity()是去掉maxpool最干净的做法比在forward里注释掉更好因为它不影响state_dict的键名一致性也不破坏模块树。第二替换conv1之后必须重新初始化因为原来的初始化循环在构造函数里已经跑完了新建的Conv2d用的是PyTorch默认初始化。第三_make_layer里对downsample的判断会用到self.inplanes改动stem之后inplanes仍然是64不需要调整。我自己在CIFAR-10上跑过这个版本ResNet18大概能到94%左右的测试精度训练200个epoch左右单卡几小时。相比之下不改stem直接跑模型完全不收敛因为空间尺寸已经被压没了。6. 复现之后让它变成你自己的主干6.1 换分类头、改输入通道数复现通过之后模型就变成了一个可以自由改造的组件。最常见的三种改造换分类头。多分类改二分类、加Dropout、换成多标签的分类输出都只需要改fc层。做多标签时要注意输出不用softmax用sigmoid配合BCEWithLogitsLoss。改输入通道。医学影像、遥感图像里常见单通道或者四通道、十几通道的输入。做法是把conv1换成对应的输入通道数然后把预训练权重的第一层卷积核在通道维度上求平均再复制def adapt_first_conv(model, in_channels): old_conv model.conv1 new_conv nn.Conv2d(in_channels, old_conv.out_channels, kernel_sizeold_conv.kernel_size, strideold_conv.stride, paddingold_conv.padding, biasFalse) with torch.no_grad(): # 预训练权重按通道求平均复制到所有输入通道上 # 这样既保留了预训练学到的空间模式又适配了新的输入通道数 w old_conv.weight.mean(dim1, keepdimTrue) # (out,1,k,k) new_conv.weight.copy_(w.repeat(1, in_channels, 1, 1)) model.conv1 new_conv return model这个平均复制的小技巧比随机初始化效果好很多尤其是在数据量不大的场景下。去掉分类头做主干。把fc换成nn.Identity()或者在forward里直接return avgpool后的特征。注意做密集预测任务检测、分割时通常不取最后一层的7x7特征而是取layer3的输出14x14因为它的空间分辨率更高对小目标更友好。6.2 中间特征抽取与可视化验证改完结构之后验证方式也要跟着改。抽取多尺度特征做FPN或者特征融合时我建议写一个统一的helperclass ResNetBackbone(nn.Module): 把 ResNet 拆成 stem 四个 stage方便按需取中间特征 def __init__(self, backbone): super().__init__() self.stem nn.Sequential( backbone.conv1, backbone.bn1, backbone.relu, backbone.maxpool) self.layer1 backbone.layer1 self.layer2 backbone.layer2 self.layer3 backbone.layer3 self.layer4 backbone.layer4 def forward(self, x, out_indices(1, 2, 3, 4)): x self.stem(x) outs [] for i, layer in enumerate([self.layer1, self.layer2, self.layer3, self.layer4], start1): x layer(x) if i in out_indices: outs.append(x) return outs用的时候out_indices(2,3,4)就能拿到1/8、1/16、1/32三种分辨率的特征。验证时记得对每个输出打印shape确认通道数依次是512、1024、2048。想让特征图更直观可以把某一层的通道响应做个平均然后归一化到0到255存成灰度图。做完这一步你会看到浅层特征的响应基本对应边缘和纹理深层特征已经变成一张很稀疏的热力图聚焦在语义对象的位置。这个过程对理解为什么检测任务需要用多尺度特征特别有帮助。6.3 几个低成本但有效的结构改进理解标准结构之后可以尝试几个改动小、收益明确的变体ResNet-D的downsample改进。把downsample里的1x1 stride2卷积换成avgpool(stride2) 1x1(stride1)。这样做的好处是降采样不再丢失信息因为池化比带stride的卷积更平滑。改动只需要几行在分类任务上通常有零点几个点的提升。加SE模块。在每个残差块的最后残差相加之前插一个通道注意力全局池化得到通道描述向量过两层全连接得到通道权重再乘回特征。这个改动带来的参数量和计算量都很小但对通道间关系的建模有帮助。实现时要注意插入位置——放在第二个BN之后、相加之前最自然。末端BN零初始化。这个前面讲过从零训练时很值得加代码量只有两行。换掉stem。对于分辨率较小的输入把7x7 stride2换成三个3x3 stride2的串联类似ResNet-C的设计能减少浅层的信息损失代价是浅层计算量增加。这个改动在32x32到128x128的输入上收益比较明显。我个人在实际操作中的体会是先不要急着加这些模块。把标准结构完整复现、参数对账、数值对拍全部走一遍再拿标准结构在目标数据集上跑出一个baseline记录清楚每个epoch的loss和验证精度。有了这个baseline后续任何改动是不是真的有效一对比就清楚了。没有baseline就加模块最后你只会得到一堆无法归因的实验结果。最后再分享一个小技巧把形状自检、参数量对账、权重加载对拍这三段代码固定成一个test_model.py每次改结构先跑一遍。这几秒钟的检查能省掉你在训练脚本里加print、等半个小时看loss曲线、然后发现是通道写错的时间。ResNet的代码本身不长真正花时间的从来都是排查那些看起来没问题的地方。
返回列表