☰
Inception-ResNet 详解:v1/v2 结构对比与 PyTorch 实现
2026/10/2 18:59:24 网站建设 项目流程

开头直接说人话: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-B,Loss 经常跳到 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-v2
Stem 输出通道192320
Inception-ResNet-A 数量510
A 模块分支通道3264
Reduction-A 输出通道8961088
Inception-ResNet-B 数量1020
B 模块分支通道128256
Reduction-B 输出通道17922080
Inception-ResNet-C 数量59
C 模块分支通道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×1+3×3、1×1+3×3+3×3,另加一个 MaxPool+1×1 的池化分支。它用最朴素的方形卷积捕捉多尺度。这里的思路是:在分辨率还比较高的 35×35 阶段,先让网络从多个感受野上充分提取局部信息。

B 模块把 3×3 卷积拆成了不对称形式:3×3 等价于 1×3 + 3×1。这样在保持感受野接近的同时,参数量从 3×3=9 降到 1×3+3×1=6,约省三分之一,而且非对称卷积对图像水平/垂直方向的结构特征有更强的针对性。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 或 stride=2 的卷积,一步到位,简单粗暴。Inception-ResNet 的降采样模块更讲究:它把输入同时扔给三条或四条并行的采样路径,有卷积采样、池化采样、连续小卷积采样,最后把结果在通道维拼接。

这样做的好处很直观——不同采样方式丢失的信息不一样。MaxPool 只保留局部最大值,强响应特征被保留但细节被丢弃;stride=2 的卷积可以学习如何“压缩”,但单条路径的表达有限。多路采样拼接后,降采样层变成了“多视角融合层”,网络可以同时获得保留尖峰信息的池化分支和经过学习的卷积分支。同时因为是多支路 concat,降采样不仅没有缩小通道数,反而把通道数拉高了一大截,为下一个阶段提供更丰富的特征。这正是 Inception-ResNet 能保持很强表达力的原因之一。

2.4 v1 与 v2 的适用场景

从实验观察来说,v1 在中小数据集上(几万到几十万张图)和 v2 的差距并没有想象中那么大,但 v1 的速度和显存占用优势非常明显。v2 在 ImageNet 这种千万级数据上能拉开差距,因为它容量大、更依赖海量数据来发挥深层表达。做移动端或实时推理,v1 更合适;打比赛、刷指标,v2 结合大规模预训练收益更高。

3. PyTorch 代码实现与逐段讲解

3.1 环境与基本约定

我用的环境比较简单:Python 3.8+,PyTorch 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, stride=1, padding=0): super().__init__() self.conv = nn.Conv2d(in_ch, out_ch, kernel_size, stride, padding, bias=False) self.bn = nn.BatchNorm2d(out_ch, eps=0.001) self.relu = nn.ReLU(inplace=True) def forward(self, x): return self.relu(self.bn(self.conv(x)))

注意这里的 bias=False,因为后面紧跟 BatchNorm,卷积偏置会被 BN 的平移项吸收,留着反而多一份冗余参数。eps 我按 Google 原版习惯设成了 0.001,PyTorch 默认 BN 的 eps 是 1e-5,对小 batch 训练来说 0.001 会更稳一点,但不强制。

3.3 Inception-ResNet-A/B/C 模块实现

A 模块我实现了四个分支:1×1 卷积分支、1×1+3×3 分支、1×1+3×3+3×3 分支,以及 MaxPool+1×1 分支。四条分支 concat 后过一个 1×1 投影卷积,输出的通道数必须等于输入通道数,这样残差才能按位相加。最后乘上 scale 再加回输入,过 ReLU。

class InceptionResNetA(nn.Module): def __init__(self, in_ch, branch_ch=32, scale=0.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, padding=1), ) self.branch3 = nn.Sequential( BasicConv2d(in_ch, branch_ch, 1), BasicConv2d(branch_ch, branch_ch, 3, padding=1), BasicConv2d(branch_ch, branch_ch, 3, padding=1), ) self.branch4 = nn.Sequential( nn.MaxPool2d(3, stride=1, padding=1), BasicConv2d(in_ch, branch_ch, 1), ) self.conv = nn.Conv2d(branch_ch * 4, in_ch, 1, bias=False) self.bn = nn.BatchNorm2d(in_ch, eps=0.001) self.relu = nn.ReLU(inplace=True) 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], dim=1) 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_ch=128, scale=0.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, stride=1, padding=1), BasicConv2d(in_ch, branch_ch, 1), ) self.conv = nn.Conv2d(branch_ch * 4, in_ch, 1, bias=False) self.bn = nn.BatchNorm2d(in_ch, eps=0.001) self.relu = nn.ReLU(inplace=True) 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], dim=1) 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_ch=192, scale=0.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, stride=1, padding=1), BasicConv2d(in_ch, branch_ch, 1), ) self.conv = nn.Conv2d(branch_ch * 4, in_ch, 1, bias=False) self.bn = nn.BatchNorm2d(in_ch, eps=0.001) self.relu = nn.ReLU(inplace=True) 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], dim=1) 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_ch=3, out_ch=192): super().__init__() self.conv1 = BasicConv2d(in_ch, 32, 3, stride=2) self.conv2 = BasicConv2d(32, 32, 3, padding=1) self.conv3 = BasicConv2d(32, 64, 3, padding=1) self.pool1 = nn.MaxPool2d(3, stride=2) self.conv4 = BasicConv2d(64, 80, 1) self.conv5 = BasicConv2d(80, out_ch, 3) self.pool2 = nn.MaxPool2d(3, stride=2) 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_ch=3, out_ch=320): 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 x

Reduction-A 和 Reduction-B 上一节说过,是多路降采样拼接。v1 的 Reduction-A 将 192 通道升到 896,三条路分别用 stride=2 的 3×3 卷积、连续小卷积、以及单条 1×1+3×3 路径采样。v2 的版本把输入通道从 320 升到 1088。这里有一个细节:stride=2 的卷积本身就会减小特征图尺寸,所以不需要再额外加池化层;但卷积核尺寸要覆盖 stride 对应的感受野,否则会漏采信息。3×3、stride=2 是下采样卷积的经典配置。

class ReductionA1(nn.Module): def __init__(self, in_ch=192): super().__init__() self.branch0 = BasicConv2d(in_ch, 384, 3, stride=2) self.branch1 = nn.Sequential( BasicConv2d(in_ch, 192, 1), BasicConv2d(192, 192, 3, padding=1), BasicConv2d(192, 256, 3, stride=2), ) self.branch2 = nn.Sequential( BasicConv2d(in_ch, 256, 1), BasicConv2d(256, 256, 3, stride=2), ) def forward(self, x): return torch.cat([self.branch0(x), self.branch1(x), self.branch2(x)], dim=1) class ReductionB1(nn.Module): def __init__(self, in_ch=896): super().__init__() self.branch0 = nn.Sequential( BasicConv2d(in_ch, 256, 1), BasicConv2d(256, 384, 3, stride=2), ) self.branch1 = nn.Sequential( BasicConv2d(in_ch, 256, 1), BasicConv2d(256, 256, 3, stride=2), ) 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, stride=2), ) self.branch3 = nn.MaxPool2d(3, stride=2) def forward(self, x): return torch.cat([self.branch0(x), self.branch1(x), self.branch2(x), self.branch3(x)], dim=1)

v2 的 Reduction-A 和 Reduction-B 结构类似,只是把输入通道和目标输出通道按 v2 的容量调大。实现时注意每一路的中间通道数要和输入输出匹配。

class ReductionA2(nn.Module): def __init__(self, in_ch=320): super().__init__() self.branch0 = BasicConv2d(in_ch, 384, 3, stride=2) self.branch1 = nn.Sequential( BasicConv2d(in_ch, 320, 1), BasicConv2d(320, 320, 3, padding=1), BasicConv2d(320, 352, 3, stride=2), ) self.branch2 = nn.Sequential( BasicConv2d(in_ch, 352, 1), BasicConv2d(352, 352, 3, stride=2), ) def forward(self, x): return torch.cat([self.branch0(x), self.branch1(x), self.branch2(x)], dim=1) class ReductionB2(nn.Module): def __init__(self, in_ch=1088): super().__init__() self.branch0 = nn.Sequential( BasicConv2d(in_ch, 288, 1), BasicConv2d(288, 384, 3, stride=2), ) self.branch1 = nn.Sequential( BasicConv2d(in_ch, 288, 1), BasicConv2d(288, 288, 3, stride=2), ) 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, stride=2), ) self.branch3 = nn.MaxPool2d(3, stride=2) def forward(self, x): return torch.cat([self.branch0(x), self.branch1(x), self.branch2(x), self.branch3(x)], dim=1)

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_classes=1000, dropout=0.5): super().__init__() self.stem = StemV1() self.A = nn.ModuleList([InceptionResNetA(192, branch_ch=32, scale=0.17) for _ in range(5)]) self.reductionA = ReductionA1() self.B = nn.ModuleList([InceptionResNetB(896, branch_ch=128, scale=0.17) for _ in range(10)]) self.reductionB = ReductionB1() self.C = nn.ModuleList([InceptionResNetC(1792, branch_ch=192, scale=0.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_classes=1000, dropout=0.5): super().__init__() self.stem = StemV2() self.A = nn.ModuleList([InceptionResNetA(320, branch_ch=64, scale=0.2) for _ in range(10)]) self.reductionA = ReductionA2() self.B = nn.ModuleList([InceptionResNetB(1088, branch_ch=256, scale=0.2) for _ in range(20)]) self.reductionB = ReductionB2() self.C = nn.ModuleList([InceptionResNetC(2080, branch_ch=256, scale=0.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_classes=10) 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_classes=10) 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.1M,v2 大约 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,最终分类头维度只和类别数相关。不过切到更小输入后,某些深层模块的感受野覆盖范围会变,精度会有一定影响。

数据增强我一般这么配:随机水平翻转、随机裁剪、颜色抖动,加上 Normalize,mean/std 用 ImageNet 的统计值。如果做小数据集,还可以用 RandomAffine 或 CutMix 加强一下。

4.2 优化器与学习率策略

这个网络我不是很喜欢用默认的 Adam 一把梭,因为它的 BN 层多,结构化较强,用带 Nesterov 的 SGD(momentum=0.9、weight_decay=1e-4~5e-4)训练更稳。不过现在 AdamW + cosine 衰减也完全可以,重点是把 warmup 和降学习率做好。

我的经验是前 5 个 epoch 做线性 warmup,从 1e-4 升到目标学习率,之后用 cosine annealing 降到最低值。目标学习率 SGD 用 0.1(batch=256 时),AdamW 用 1e-3 到 3e-4 起步。batch size 大一点对 BN 统计更友好,显存允许就尽量开到 64 或 128。

4.3 两个实验:v1 和 v2 的实际表现

我用 CIFAR-10 跑了两组快速验证:v1 和 v2 都训练 100 epoch,SGD + cosine,batch 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 调到 512,scale 还保持 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, stride=2) 改成 (1, 32, 3, stride=2)。尺寸方面,输入长宽只要能整除到 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,两者代码基本通用,你完全可以把这份笔记里的模块复制过去,改几个通道数就切换版本。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询