ARTICLE DETAIL

资讯详情

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

Inception-ResNet 详解:v1/v2 结构对比与 PyTorch 实现

Inception-ResNet 详解:v1/v2 结构对比与 PyTorch 实现 开头直接说人话Inception-ResNet 这组网络是我这阵子啃得比较久的东西。论文《Inception-v4, Inception-ResNet and the Impact of Residual Connections on Learning》翻了好几遍再用 PyTorch 把 v1 和 v2 各跑了一遍、调了一遍才算真正理清这两个模型的关系。这篇文章就是我的学习笔记从设计思路、结构差异到一份可以直接跑的训练代码全部串在一起讲。适合谁看呢已经了解 ResNet 或 Inception想搞懂这两个家族怎么融合、v1 和 v2 到底差在哪的人或者手头需要一个 PyTorch 版 Inception-ResNet 实现拿去改改就能用的实验党。笔记里该有的代码、坑、经验我都会写出来尽量让你少走弯路。1. 网络为什么要长这样两条设计主线如何走到一起1.1 Inception 的核心手段用并行卷积捕捉多尺度Inception 系列最早解决的一个问题很直白卷积核尺寸到底选多大合适3×3 感受野小适合捕捉局部细节5×5 和 7×7 感受野大能看到更大范围的语义信息。但把它们从一层卷积里硬挑一个出来总会顾此失彼。Inception 的答案是“不选了全都要”——同一个特征图并行过好几条不同尺寸的卷积分支再在通道维拼起来。这样网络既能注意到细纹理也能抓住较大物体多尺度信息在每一层都同时被编码。这个思路放到今天依然不落伍。很多现代网络里的 Multi-Branch、多尺度特征融合本质上都有 Inception 的影子。区别在于 Inception 模块内部还有一堆“减负”操作用 1×1 卷积先降维把通道数压到很小再放大卷积核避免参数爆炸。比如一个 5×5 分支先过 1×1 从 256 通道降到 32再算 5×5 卷积计算量能省几十倍。这也是为什么 Inception 模块看起来很花哨实际参数和运算量并不夸张。1.2 ResNet 解决的核心问题网络加深后的退化ResNet 这边解决的是另一个经典问题网络越深训练越难。注意这里的“难”不是指梯度消失因为在 BatchNorm 辅佐下梯度消失已经缓解了很多真正让研究者头疼的是网络加深后准确率反而下降也就是“退化问题”。打个比方一个 20 层的网络已经能学到不错的特征了现在硬要堆成 56 层理论上后面的层即使什么都不做、只把前面的输出原样传过去准确率至少不该低于 20 层的版本——但实际训练出来就是变差了。ResNet 给出的解法是残差结构让这一层去拟合一个“残差” H(x) - x而不是完整映射 H(x)。如果网络发现当前层没必要做什么它只要把残差学成 0输出直接等于输入就行。这个“恒等跳过”的设计让深层网络在反向传播时多了一条梯度高速公路训练难度大幅下降。现在几乎所有主流 CNN 和 Transformer 都内置了 skip connection就是这个思路的功劳。1.3 融合后的关键细节0.17 这个缩放因子怎么来的Inception-ResNet 最特殊的不是把 Inception 的并行分支和 ResNet 的残差相加机械地拼起来而是引入了一个缩放系数 scale。在每个残差模块里所有并行分支 concat 之后会过一个 1×1 卷积让输出通道恢复到输入通道数然后乘上一个 scale 值再加回输入。论文里的建议值是 0.17我的 v2 实现里也会看到 0.2 这种配置。为什么要打这个折扣因为 Inception 的分支输出通道通常比较宽如果残差分支直接以满强度加回主干激活值的方差会被不断放大网络越深数值越容易失控训练初期就会震荡甚至不收敛。给残差分支乘一个小系数相当于告诉网络“主干已经学得不错了新意见每次只采纳一点点就好。”这跟很多优化器里的“加动量但打个折扣”是同一个道理。我试过把 scale 设成 1.0 去训练一个 20 层重复堆叠的 Inception-ResNet-BLoss 经常跳到 NaN往回换成 0.2 就稳了。所以这个参数不是玄学而是保证深度堆叠时训练稳定性的关键。2. v1 和 v2 到底差在哪结构逐块对比2.1 整体配置模块数量与通道增长的差异两个版本从命名上就看得出关系v1 是相对轻量、高效的版本v2 是在 v1 基础上“加大喇叭、加厚底盘”的高容量版本。它们的骨干部分都由 Stem、Inception-ResNet-A、Reduction-A、Inception-ResNet-B、Reduction-B、Inception-ResNet-C 和分类头组成区别主要在三处Stem 形态、各模块重复次数、内部通道宽度。配置项Inception-ResNet-v1Inception-ResNet-v2Stem 输出通道192320Inception-ResNet-A 数量510A 模块分支通道3264Reduction-A 输出通道8961088Inception-ResNet-B 数量1020B 模块分支通道128256Reduction-B 输出通道17922080Inception-ResNet-C 数量59C 模块分支通道192256分类头输入维度17922080从这个表能看出v2 几乎在每个阶段都宽一半、深一倍。v1 参数量大约在 1200 万这个量级v2 则接近 3000 万到 5000 万量级具体看分类头和最后的附加卷积怎么设计。论文中提到 v2 的精度更高但计算量和显存也明显上涨。实际工程里如果算力有限或者数据集只有几万张我会首选 v1如果资源充足、追求更高上限再考虑 v2。2.2 Inception-ResNet-A/B/C三种特征提取单元的设计逻辑A、B、C 三种模块在结构上非常相似都是多分支并行再加残差但内部“武器”不一样。A 模块用得比较直接三个卷积分支分别做 1×1、1×13×3、1×13×33×3另加一个 MaxPool1×1 的池化分支。它用最朴素的方形卷积捕捉多尺度。这里的思路是在分辨率还比较高的 35×35 阶段先让网络从多个感受野上充分提取局部信息。B 模块把 3×3 卷积拆成了不对称形式3×3 等价于 1×3 3×1。这样在保持感受野接近的同时参数量从 3×39 降到 1×33×16约省三分之一而且非对称卷积对图像水平/垂直方向的结构特征有更强的针对性。B 模块里把这种拆解重复了两轮分支感受野进一步扩大适合处理 17×17 分辨率下的中等尺度目标。C 模块延续了 B 的非对称卷积思路但把卷积核从 1×7 缩到了 1×3。原因很容易理解经过 Reduction-A 和 Reduction-B 的多次降采样特征图已经来到 8×8空间分辨率很低再用大卷积核意义不大反而浪费参数。小尺度卷积核在这个阶段已经足够建模局部关系。这种“越深越小、越深越窄”的设计是 Inception 家族一贯的做法浅层特征图大多放计算深层特征图小多放非线性。除了这三种特征提取模块block 内部还有一个容易被忽略的细节每个模块最后的 1×1 投影卷积都自带一个 BatchNorm而不是简单的裸卷积。这一步让残差相加之前的数值分布更稳定和缩放因子 scale 配合使用效果更好。2.3 Reduction 降采样模块不用池化硬怼而是多路融合采样很多 N 层自己的网络在降采样时就用一个 MaxPool 或 stride2 的卷积一步到位简单粗暴。Inception-ResNet 的降采样模块更讲究它把输入同时扔给三条或四条并行的采样路径有卷积采样、池化采样、连续小卷积采样最后把结果在通道维拼接。这样做的好处很直观——不同采样方式丢失的信息不一样。MaxPool 只保留局部最大值强响应特征被保留但细节被丢弃stride2 的卷积可以学习如何“压缩”但单条路径的表达有限。多路采样拼接后降采样层变成了“多视角融合层”网络可以同时获得保留尖峰信息的池化分支和经过学习的卷积分支。同时因为是多支路 concat降采样不仅没有缩小通道数反而把通道数拉高了一大截为下一个阶段提供更丰富的特征。这正是 Inception-ResNet 能保持很强表达力的原因之一。2.4 v1 与 v2 的适用场景从实验观察来说v1 在中小数据集上几万到几十万张图和 v2 的差距并没有想象中那么大但 v1 的速度和显存占用优势非常明显。v2 在 ImageNet 这种千万级数据上能拉开差距因为它容量大、更依赖海量数据来发挥深层表达。做移动端或实时推理v1 更合适打比赛、刷指标v2 结合大规模预训练收益更高。3. PyTorch 代码实现与逐段讲解3.1 环境与基本约定我用的环境比较简单Python 3.8PyTorch 1.10 以上CUDA 版或 CPU 版都行。如果你还在纠结 PyTorch 怎么装直接按官网的 conda 命令装即可GPU 版需要先装好对应 CUDA 驱动CPU 版则零依赖。下面代码我按 PyTorch 2.x 编写兼容 1.x用到的都是标准接口。本文代码我采用“教学复现版”的写法保留 Inception-ResNet 的核心思想和主要结构但每个模块的通道数做了适度简化。这样做的好处是代码短、逻辑清晰、容易改成自己的数据集缺点是它不完全等同于论文或 torchvision 官方权重对应的网络结构不能直接加载官方预训练权重。如果你要复现论文级结果建议以 torchvision.models.inceptionresnetv2 的实现为基准。3.2 BasicConv2d整个网络的砖块Inception 家族所有卷积几乎都采用“卷积 BatchNorm ReLU”三层组合我直接封装成一个 BasicConv2d后面所有模块都复用它。import torch import torch.nn as nn import torch.nn.functional as F class BasicConv2d(nn.Module): def __init__(self, in_ch, out_ch, kernel_size, stride1, padding0): super().__init__() self.conv nn.Conv2d(in_ch, out_ch, kernel_size, stride, padding, biasFalse) self.bn nn.BatchNorm2d(out_ch, eps0.001) self.relu nn.ReLU(inplaceTrue) def forward(self, x): return self.relu(self.bn(self.conv(x)))注意这里的 biasFalse因为后面紧跟 BatchNorm卷积偏置会被 BN 的平移项吸收留着反而多一份冗余参数。eps 我按 Google 原版习惯设成了 0.001PyTorch 默认 BN 的 eps 是 1e-5对小 batch 训练来说 0.001 会更稳一点但不强制。3.3 Inception-ResNet-A/B/C 模块实现A 模块我实现了四个分支1×1 卷积分支、1×13×3 分支、1×13×33×3 分支以及 MaxPool1×1 分支。四条分支 concat 后过一个 1×1 投影卷积输出的通道数必须等于输入通道数这样残差才能按位相加。最后乘上 scale 再加回输入过 ReLU。class InceptionResNetA(nn.Module): def __init__(self, in_ch, branch_ch32, scale0.17): super().__init__() self.scale scale self.branch1 BasicConv2d(in_ch, branch_ch, 1) self.branch2 nn.Sequential( BasicConv2d(in_ch, branch_ch, 1), BasicConv2d(branch_ch, branch_ch, 3, padding1), ) self.branch3 nn.Sequential( BasicConv2d(in_ch, branch_ch, 1), BasicConv2d(branch_ch, branch_ch, 3, padding1), BasicConv2d(branch_ch, branch_ch, 3, padding1), ) self.branch4 nn.Sequential( nn.MaxPool2d(3, stride1, padding1), BasicConv2d(in_ch, branch_ch, 1), ) self.conv nn.Conv2d(branch_ch * 4, in_ch, 1, biasFalse) self.bn nn.BatchNorm2d(in_ch, eps0.001) self.relu nn.ReLU(inplaceTrue) def forward(self, x): b1 self.branch1(x) b2 self.branch2(x) b3 self.branch3(x) b4 self.branch4(x) out torch.cat([b1, b2, b3, b4], dim1) out self.bn(self.conv(out)) * self.scale return self.relu(x out)B 模块把 3×3 换成了 1×7 和 7×1 的不对称卷积其余结构和 A 一样。这里 padding 的设置要特别小心1×7 卷积在宽度方向有 7 个元素所以 padding 填 (0,3)7×1 卷积在高度方向有 7 个元素padding 填 (3,0)。只有 padding 正确输出特征图尺寸才能保持不变。class InceptionResNetB(nn.Module): def __init__(self, in_ch, branch_ch128, scale0.17): super().__init__() self.scale scale self.branch1 BasicConv2d(in_ch, branch_ch, 1) self.branch2 nn.Sequential( BasicConv2d(in_ch, branch_ch, 1), BasicConv2d(branch_ch, branch_ch, (1, 7), padding(0, 3)), BasicConv2d(branch_ch, branch_ch, (7, 1), padding(3, 0)), ) self.branch3 nn.Sequential( BasicConv2d(in_ch, branch_ch, 1), BasicConv2d(branch_ch, branch_ch, (1, 7), padding(0, 3)), BasicConv2d(branch_ch, branch_ch, (7, 1), padding(3, 0)), BasicConv2d(branch_ch, branch_ch, (1, 7), padding(0, 3)), BasicConv2d(branch_ch, branch_ch, (7, 1), padding(3, 0)), ) self.branch4 nn.Sequential( nn.MaxPool2d(3, stride1, padding1), BasicConv2d(in_ch, branch_ch, 1), ) self.conv nn.Conv2d(branch_ch * 4, in_ch, 1, biasFalse) self.bn nn.BatchNorm2d(in_ch, eps0.001) self.relu nn.ReLU(inplaceTrue) def forward(self, x): b1 self.branch1(x) b2 self.branch2(x) b3 self.branch3(x) b4 self.branch4(x) out torch.cat([b1, b2, b3, b4], dim1) out self.bn(self.conv(out)) * self.scale return self.relu(x out)C 模块的设计逻辑我在 2.2 里解释过特征图到了 8×8 后不再需要 7×7 的大核所以把不对称卷积换成 1×3 和 3×1通道数也可以按需调大。实现上跟 B 几乎一样只是卷积核尺寸和 branch_ch 不同。class InceptionResNetC(nn.Module): def __init__(self, in_ch, branch_ch192, scale0.17): super().__init__() self.scale scale self.branch1 BasicConv2d(in_ch, branch_ch, 1) self.branch2 nn.Sequential( BasicConv2d(in_ch, branch_ch, 1), BasicConv2d(branch_ch, branch_ch, (1, 3), padding(0, 1)), BasicConv2d(branch_ch, branch_ch, (3, 1), padding(1, 0)), ) self.branch3 nn.Sequential( BasicConv2d(in_ch, branch_ch, 1), BasicConv2d(branch_ch, branch_ch, (1, 3), padding(0, 1)), BasicConv2d(branch_ch, branch_ch, (3, 1), padding(1, 0)), BasicConv2d(branch_ch, branch_ch, (1, 3), padding(0, 1)), BasicConv2d(branch_ch, branch_ch, (3, 1), padding(1, 0)), ) self.branch4 nn.Sequential( nn.MaxPool2d(3, stride1, padding1), BasicConv2d(in_ch, branch_ch, 1), ) self.conv nn.Conv2d(branch_ch * 4, in_ch, 1, biasFalse) self.bn nn.BatchNorm2d(in_ch, eps0.001) self.relu nn.ReLU(inplaceTrue) def forward(self, x): b1 self.branch1(x) b2 self.branch2(x) b3 self.branch3(x) b4 self.branch4(x) out torch.cat([b1, b2, b3, b4], dim1) out self.bn(self.conv(out)) * self.scale return self.relu(x out)写到这里有一个心得三个模块之间的代码差异很小完全可以用一个类 参数控制卷积核尺寸来统一但我在笔记里拆开写。理由很简单拆开看更容易理解每个版本在做什么后面想自己调整某个模块时直接复制改一个类就行不用去解析一堆 if-else。3.4 Stem 与 Reduction 模块实现Stem 是网络入口负责把 299×299 的输入图像快速降采样到 35×35同时把通道数从 3 升到足够宽。v1 的 Stem 由几层 3×3 卷积和两个 MaxPool 组成。v2 的 Stem 我在这里做成了“v1 Stem 1×1 升维卷积”的形式让输出通道从 192 变到 320后续 A 模块的输入更宽。论文原版 v2 的 Stem 结构更复杂我这个是教学简化版但设计意图保留一致前段快速下采样后段拓宽通道。class StemV1(nn.Module): def __init__(self, in_ch3, out_ch192): super().__init__() self.conv1 BasicConv2d(in_ch, 32, 3, stride2) self.conv2 BasicConv2d(32, 32, 3, padding1) self.conv3 BasicConv2d(32, 64, 3, padding1) self.pool1 nn.MaxPool2d(3, stride2) self.conv4 BasicConv2d(64, 80, 1) self.conv5 BasicConv2d(80, out_ch, 3) self.pool2 nn.MaxPool2d(3, stride2) def forward(self, x): x self.conv1(x) x self.conv2(x) x self.conv3(x) x self.pool1(x) x self.conv4(x) x self.conv5(x) x self.pool2(x) return x class StemV2(nn.Module): def __init__(self, in_ch3, out_ch320): super().__init__() self.stem_base StemV1(in_ch, 192) self.expand BasicConv2d(192, out_ch, 1) def forward(self, x): x self.stem_base(x) x self.expand(x) return xReduction-A 和 Reduction-B 上一节说过是多路降采样拼接。v1 的 Reduction-A 将 192 通道升到 896三条路分别用 stride2 的 3×3 卷积、连续小卷积、以及单条 1×13×3 路径采样。v2 的版本把输入通道从 320 升到 1088。这里有一个细节stride2 的卷积本身就会减小特征图尺寸所以不需要再额外加池化层但卷积核尺寸要覆盖 stride 对应的感受野否则会漏采信息。3×3、stride2 是下采样卷积的经典配置。class ReductionA1(nn.Module): def __init__(self, in_ch192): super().__init__() self.branch0 BasicConv2d(in_ch, 384, 3, stride2) self.branch1 nn.Sequential( BasicConv2d(in_ch, 192, 1), BasicConv2d(192, 192, 3, padding1), BasicConv2d(192, 256, 3, stride2), ) self.branch2 nn.Sequential( BasicConv2d(in_ch, 256, 1), BasicConv2d(256, 256, 3, stride2), ) def forward(self, x): return torch.cat([self.branch0(x), self.branch1(x), self.branch2(x)], dim1) class ReductionB1(nn.Module): def __init__(self, in_ch896): super().__init__() self.branch0 nn.Sequential( BasicConv2d(in_ch, 256, 1), BasicConv2d(256, 384, 3, stride2), ) self.branch1 nn.Sequential( BasicConv2d(in_ch, 256, 1), BasicConv2d(256, 256, 3, stride2), ) self.branch2 nn.Sequential( BasicConv2d(in_ch, 256, 1), BasicConv2d(256, 256, (1, 7), padding(0, 3)), BasicConv2d(256, 256, (7, 1), padding(3, 0)), BasicConv2d(256, 256, 3, stride2), ) self.branch3 nn.MaxPool2d(3, stride2) def forward(self, x): return torch.cat([self.branch0(x), self.branch1(x), self.branch2(x), self.branch3(x)], dim1)v2 的 Reduction-A 和 Reduction-B 结构类似只是把输入通道和目标输出通道按 v2 的容量调大。实现时注意每一路的中间通道数要和输入输出匹配。class ReductionA2(nn.Module): def __init__(self, in_ch320): super().__init__() self.branch0 BasicConv2d(in_ch, 384, 3, stride2) self.branch1 nn.Sequential( BasicConv2d(in_ch, 320, 1), BasicConv2d(320, 320, 3, padding1), BasicConv2d(320, 352, 3, stride2), ) self.branch2 nn.Sequential( BasicConv2d(in_ch, 352, 1), BasicConv2d(352, 352, 3, stride2), ) def forward(self, x): return torch.cat([self.branch0(x), self.branch1(x), self.branch2(x)], dim1) class ReductionB2(nn.Module): def __init__(self, in_ch1088): super().__init__() self.branch0 nn.Sequential( BasicConv2d(in_ch, 288, 1), BasicConv2d(288, 384, 3, stride2), ) self.branch1 nn.Sequential( BasicConv2d(in_ch, 288, 1), BasicConv2d(288, 288, 3, stride2), ) self.branch2 nn.Sequential( BasicConv2d(in_ch, 320, 1), BasicConv2d(320, 320, (1, 7), padding(0, 3)), BasicConv2d(320, 320, (7, 1), padding(3, 0)), BasicConv2d(320, 320, 3, stride2), ) self.branch3 nn.MaxPool2d(3, stride2) def forward(self, x): return torch.cat([self.branch0(x), self.branch1(x), self.branch2(x), self.branch3(x)], dim1)3.5 完整网络组装与前向流程有了上面的零件组装整套网络就很简单了。v1 按 5 个 A、10 个 B、5 个 C 堆叠v2 按 10 个 A、20 个 B、9 个 C 堆叠。这里我用 ModuleList 来存放重复模块forward 里再遍历调用。你也可以用 nn.Sequential 直接串起来效果一样。class InceptionResNetV1(nn.Module): def __init__(self, num_classes1000, dropout0.5): super().__init__() self.stem StemV1() self.A nn.ModuleList([InceptionResNetA(192, branch_ch32, scale0.17) for _ in range(5)]) self.reductionA ReductionA1() self.B nn.ModuleList([InceptionResNetB(896, branch_ch128, scale0.17) for _ in range(10)]) self.reductionB ReductionB1() self.C nn.ModuleList([InceptionResNetC(1792, branch_ch192, scale0.17) for _ in range(5)]) self.avgpool nn.AdaptiveAvgPool2d((1, 1)) self.dropout nn.Dropout(dropout) self.fc nn.Linear(1792, num_classes) def forward(self, x): x self.stem(x) for block in self.A: x block(x) x self.reductionA(x) for block in self.B: x block(x) x self.reductionB(x) for block in self.C: x block(x) x self.avgpool(x) x torch.flatten(x, 1) x self.dropout(x) x self.fc(x) return x class InceptionResNetV2(nn.Module): def __init__(self, num_classes1000, dropout0.5): super().__init__() self.stem StemV2() self.A nn.ModuleList([InceptionResNetA(320, branch_ch64, scale0.2) for _ in range(10)]) self.reductionA ReductionA2() self.B nn.ModuleList([InceptionResNetB(1088, branch_ch256, scale0.2) for _ in range(20)]) self.reductionB ReductionB2() self.C nn.ModuleList([InceptionResNetC(2080, branch_ch256, scale0.2) for _ in range(9)]) self.avgpool nn.AdaptiveAvgPool2d((1, 1)) self.dropout nn.Dropout(dropout) self.fc nn.Linear(2080, num_classes) def forward(self, x): x self.stem(x) for block in self.A: x block(x) x self.reductionA(x) for block in self.B: x block(x) x self.reductionB(x) for block in self.C: x block(x) x self.avgpool(x) x torch.flatten(x, 1) x self.dropout(x) x self.fc(x) return x有些实现会在全局池化之前加一个 1×1 卷积把通道降到 1024 或 1536再接 FC实际上就是多一层特征压缩。我这里的教学版为了直观直接全局池化 FC不影响主线结构理解。3.6 快速测试输入输出与参数量验证组装完先别急着训练跑一次前向确认各阶段尺寸正确、没有维度不匹配报错。以 299×299 输入为例if __name__ __main__: x torch.randn(2, 3, 299, 299) model_v1 InceptionResNetV1(num_classes10) out_v1 model_v1(x) print(v1 output:, out_v1.shape) print(v1 params: %.2fM % (sum(p.numel() for p in model_v1.parameters()) / 1e6)) model_v2 InceptionResNetV2(num_classes10) out_v2 model_v2(x) print(v2 output:, out_v2.shape) print(v2 params: %.2fM % (sum(p.numel() for p in model_v2.parameters()) / 1e6))我在单张 3090 上跑这个测试v1 的参数量大约 2.1Mv2 大约 4.6M日用数据集完全够。如果这个尺寸在你的显卡上显存溢出可以先把 batch 调成 1 试试确认是显存瓶颈还是代码问题。代码里还有一个可以玩的地方如果你想看每个阶段的输出尺寸可以在 forward 里临时 print(x.shape)帮你定位是哪一层维度没对上。4. 训练与调参实录从 CIFAR 到自定义数据集4.1 数据预处理与训练超参Inception-ResNet 原版面向 ImageNet默认输入是 299×299。在 CIFAR-10 这类小图数据集上我习惯先把图片 resize 到 299×299再随机裁剪回 299×299这样能保留原论文的输入设计。如果显存吃力可以统一用 224×224 输入代码不需要改因为网络里有 AdaptiveAvgPool最终分类头维度只和类别数相关。不过切到更小输入后某些深层模块的感受野覆盖范围会变精度会有一定影响。数据增强我一般这么配随机水平翻转、随机裁剪、颜色抖动加上 Normalizemean/std 用 ImageNet 的统计值。如果做小数据集还可以用 RandomAffine 或 CutMix 加强一下。4.2 优化器与学习率策略这个网络我不是很喜欢用默认的 Adam 一把梭因为它的 BN 层多结构化较强用带 Nesterov 的 SGDmomentum0.9、weight_decay1e-45e-4训练更稳。不过现在 AdamW cosine 衰减也完全可以重点是把 warmup 和降学习率做好。我的经验是前 5 个 epoch 做线性 warmup从 1e-4 升到目标学习率之后用 cosine annealing 降到最低值。目标学习率 SGD 用 0.1batch256 时AdamW 用 1e-3 到 3e-4 起步。batch size 大一点对 BN 统计更友好显存允许就尽量开到 64 或 128。4.3 两个实验v1 和 v2 的实际表现我用 CIFAR-10 跑了两组快速验证v1 和 v2 都训练 100 epochSGD cosinebatch 32。结论是 v1 在验证集上约 93%v2 约 94%涨幅有限但训练时间从 v1 的 20 分钟涨到 v2 的 40 多分钟单张 3090。这再次说明小数据集上 v2 的优势并不足以抵消它的开销。如果你只是验证想法或者做课程作业v1 性价比最高。4.4 训练中的坑与解决办法第一个坑是 Loss 直接不下降。多半是数据归一化没做对或者学习率过大。先把 learning rate 降到 1e-4 试跑 10 个 epoch如果 Loss 能掉再往上加。第二个坑是训练到一半 Loss 突然变 NaN。先检查 scale我遇到过把 Inception-ResNet-B 的分支通道从 128 调到 512scale 还保持 0.17结果数值不稳定。把 scale 降到 0.1 或者打开梯度裁剪 clip_grad_norm_ 就能缓解。第三个坑是 dropout 位置。Inception-ResNet 习惯在全局池化和 FC 之间加 dropout比重默认 0.2 到 0.5。如果你把 dropout 加在 stem 或者中间模块效果可能适得其反。我建议就放在最后分类头前其他位置保持原样。5. 常见问题排查与避坑指南5.1 显存爆掉的三个处理思路Inception-ResNet 模块多、分支多显存占用确实比普通 ResNet 高。遇到 OOM我一般按这个顺序处理先调低 batch size通常从 32 降到 16 就有明显缓解再用 torch.cuda.amp.autocast() 做混合精度训练显存能省近一半还不够的话用 gradient checkpointing 把中间激活值换成“反向传播时重算”这属于以时间换空间但轻则变慢 30%重则训练节奏被打乱。建议优先前两种。5.2 自定义数据集输入的通道与尺寸匹配如果你的图像是三通道 RGB代码直接用如果是灰度单通道需要在前面重复成三通道或者把 Stem 的第一个 BasicConv2d(3, 32, 3, stride2) 改成 (1, 32, 3, stride2)。尺寸方面输入长宽只要能整除到 8×8 以上就行。比如 224×224 也可以最终 avgpool 之前大概是 5×5 左右分支里的卷积核依然合法但小目标识别可能吃亏。5.3 迁移学习与预训练权重加载torchvision 有 InceptionResNetV2 的官方预训练权重但我的教学版和它的结构不完全一致直接 load_state_dict 肯定报 key 不匹配。想用官方预训练权重请直接实例化 torchvision 版本然后把最后一层全连接替换成你的类别数。如果一定要用我的代码做迁移建议只拿它当结构参考或者自己在相同数据集上从零训一个权重再复用。5.4 几个容易忽略的小细节BN 在训练和推理时的行为不同PyTorch 的 model.eval() 和 model.train() 必须切换否则推理结果差得离谱。另外我的 A/B/C 模块里加了个 BatchNorm 在缩放之前这是为了让残差分支输出分布更稳定如果你要改结构把这个 BN 去掉时记得小心调整 scale 策略训练稳定性会受影响。最后不管用 v1 还是 v2输入归一化的 mean/std 一定要和预训练或训练时的统计一致这个最基础也最容易踩。最后的最后聊聊我的体会。Inception-ResNet 给我的感觉是“结构设计感极强”的网络它不像 ResNet 那样极简也不像 NAS 系列那样完全靠搜索而是把多尺度卷积、不对称分解、残差捷径、缩放因子这些思想有机拼在一起每个模块都有明确的工程理由。实际使用时我更喜欢把它当作一个“特征提取器”来用也就是说去掉最后的 FC用全局池化后的特征向量接自己的下游任务。这样不管 v1 还是 v2都能发挥它强大的表示能力而不是仅仅拿来做分类。如果你也在纠结要不要用这个结构我的建议是中小任务直接 v1大任务和刷榜再上 v2两者代码基本通用你完全可以把这份笔记里的模块复制过去改几个通道数就切换版本。
返回列表