ARTICLE DETAIL

资讯详情

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

残差连接与shortcut全解:原理、代码实现与跨架构调参

残差连接与shortcut全解:原理、代码实现与跨架构调参 第一次在代码里看到 shortcut 这个词大多数人脑子里蹦出来的是快捷键——比如有人搜 db commander shortcut 怎么用想找的其实是软件操作里的按键组合。但在神经网络的结构图里shortcut 指的是另一回事一条把输入原封不动送到后面的跳跃连接也就是我们常说的残差连接。这两个词经常被混着用其实说的是同一类东西的不同侧面。我打算把残差连接这件事从头到尾捋一遍。为什么会有人发明它它到底解决了什么实际问题代码里怎么写才不会踩坑参数量和显存怎么算以及它在 Transformer、U-Net 这些完全不同的架构里是怎么被改造的。不管你是刚开始看 ResNet 论文的新手还是已经能默写 Bottleneck 但说不清 pre-LN 和 post-LN 差别的老手这篇内容应该都能捞到点东西。我会尽量把每个为什么讲透而不是只丢一段代码让你抄。1. shortcut这个词的两副面孔从快捷键到跳跃连接1.1 一个词两种语境在软件工具里shortcut 是路径的缩写本质上是绕过完整流程、直接抵达目标的快捷方式。这个语义其实和神经网络里的用法高度一致——残差连接干的事情就是让信息绕过若干层非线性变换直接流到后面的层去。区别只在于软件里的 shortcut 是为了省用户的时间网络里的 shortcut 是为了省梯度的力气。所以当你看到 shortcut connection、skip connection、residual connection 这几个说法时基本可以当作同义词处理。严格来说skip connection 是更早、更宽泛的叫法Highway Network、U-Net 里的跨层拼接都能算residual connection 特指 He 等人在 2015 年那篇 ResNet 论文里提出的形式化写法强调学习的是残差而 shortcut 更多出现在论文的图和代码实现里指代那条恒等分支本身。有意思的是不同框架里这个分支的命名也不统一。PyTorch 的官方 ResNet 实现里用的是identity和downsample两个变量名TensorFlow 的 Keras 实现里常用shortcut有些开源仓库干脆叫residual。名字叫法无所谓关键是你要清楚它指向的是哪一条路径。1.2 深层网络里那道绕不过去的坎2015 年之前业界已经形成了一个朴素认知网络越深表达能力越强。但真去堆的时候问题就来了。VGG 把网络堆到 19 层已经算是当时的极限再往上加层训练误差不降反升。注意这里说的不是测试误差是训练误差——这意味着根本不是过拟合的问题模型在训练集上就学不动了。这个现象被叫做退化问题degradation problem。当时的解释路径有两条一条是梯度消失/爆炸另一条是优化困难。但 BatchNorm 的引入其实已经很大程度上缓解了梯度量级的问题退化却依然存在。He 等人的洞察是一个 56 层的网络理论上完全可以退化成一个 20 层网络——把多出来的 36 层全部学成恒等映射f(x)x不就行了如果 20 层能达到某个精度56 层至少不该更差。可实测下来就是更差说明问题出在让一堆非线性层逼近恒等映射这件事本身很难。这就是残差连接的破题点既然拟合恒等映射这么费劲那就把它变成默认行为。让网络学F(x) H(x) - x前向输出变成y F(x) x。当最优解接近恒等映射时网络只需要把F的权重压向 0 就够了这比让一堆卷积核和激活函数互相抵消要容易得多。我第一次真正理解这个逻辑是看到论文里那张退化问题的曲线图——56 层训练误差高于 20 层而且是在训练集上。那一刻才意识到深了更难训不是玄学是有具体机制在里面的。1.3 残差连接的最小定义抛开所有实现细节残差连接的形式化定义只有一行y F(x, W) x其中x是输入F是待学习的残差函数通常是两三个卷积层或一个注意力模块W是它的参数。要求F(x)和x的维度一致才能逐元素相加。如果两者维度不一致就退一步用线性投影y F(x, W) W_s · xW_s一般是一个 1×1 卷积CNN或者线性层Transformer负责把通道数和空间尺寸对齐。在 ResNet 的四个 stage 切换处也就是分辨率减半、通道数翻倍的位置用的就是这个形式。这里有个细节值得停一下为什么是加法不是乘法也不是拼接加法在反向传播时有个漂亮的数学性质下一节会展开讲。而拼接concatenation是另一条路线DenseNet 用的是它代价是特征图通道数会随着层数线性增长显存吃得很凶。乘法就是门控路了Highway Network 走的是这条。选择加法本质上是选了一条信息保底通道。不管F学成什么样x至少能原样传下去。这个保底性质是后面所有讨论的地基。2. 残差连接为什么有效把梯度的高速公路修起来2.1 恒等映射比逼近恒等映射容易得多先把这个论点讲清楚因为它是整个残差思想的核心。假设我们有一组层最优的变换恰好是恒等的跳过去。对于普通堆叠的卷积层这意味着σ(W_n σ(...σ(W_1 x)...)) x网络必须让每一层的权重都精确地凑出一种正负抵消的效果这在连续参数空间里是个很苛刻的约束而且越深越苛刻。用了残差结构之后目标变成F(x, W) 0。理论上只要把最后一层的权重初始化得足够接近 0就已经很接近最优解了。也就是说残差结构把必须学会什么变成了可以什么都不学。神经网络的参数空间里权重全部趋近 0 的区域是一个平凡解优化器滑进去毫不费力。这个特性带来的直接后果是增加的层至少不会让模型更差。深层残差网络可以随时关掉多余层退化成浅层网络的表现。这是残差连接最本质的贡献也是为什么 ResNet 能一口气堆到 100 层以上还稳得住。顺带说一句这也是为什么有些论文发现残差网络的某些层在训练后残差分支的响应非常小——它们确实被关掉了。2.2 反向传播中的加法项从梯度流角度看残差连接的效果更容易量化。对y F(x) x求导∂y/∂x 1 ∂F(x)/∂x反向传播时上游传回来的梯度g会以这样的方式流向x∂L/∂x ∂L/∂y · (1 ∂F/∂x)关键就在那个1。它跟F的参数没有任何关系是一条恒定存在的通路。不管∂F/∂x有多小——哪怕因为连乘而衰减到 1e-10——梯度里始终有一个∂L/∂y的量级保底传下去。这就好比在一个层层衰减的信号链路上每两级之间都加了一根直连导线。信号可以在导线里几乎无损耗地传到最前面。不过要提醒一点这个无损耗是有条件的。如果那个1所在的通路上被插入了归一化层、激活函数或者 dropout保底性质就被破坏了。这是后面第 5 节要重点讲的坑。我见过不少人在 shortcut 上加了一个ReLU或者BN理由是怕数值不稳定结果训练效果反而变差。原因就是这条无参数通路被污染了梯度必须穿过非线性才能回流高速公路变成了乡间小路。2.3 三种shortcut的选型与代价实践中能见到的 shortcut 主要有三类各有适用场景类型形式参数量适用场景备注恒等连接y F(x) x0输入输出维度一致首选无额外开销投影连接y F(x) W_s xC_in × C_out维度变化处用 1×1 卷积stride 对齐零填充连接对通道做零填充后相加0显存紧张效果略逊于投影第三种是 ResNet 原论文里讨论过的方案。下采样时通道数不够的部分用零补齐它的好处是零参数坏处是相当于在残差分支上掩盖了一部分通道的信息。原论文的消融实验显示零填充和投影的差距不算大但投影更稳。还有一个更实际的考量投影分支本身的参数量有时并不小。以 ResNet 的 stage 切换为例输入 256 通道、输出 512 通道、1×1 卷积、stride 2参数量是256 × 512 131072。相比之下同一个 bottleneck 块的三个卷积层加起来才约 69632 个参数。也就是说一个投影 shortcut 的参数量可能接近甚至超过它保护的那个残差块本身。所以在 stage 数量不多的情况下用投影最划算如果是极深的网络可以考虑只在必要处用投影。另外还有第四类门控连接y T(x)·F(x) (1-T(x))·xHighway Network 用的是这个。T是一个 sigmoid 输出的门控。它比残差连接更灵活但也更难训因为门控本身需要学习。后来的实践基本证明把T固定成 1 就够了多余的灵活性没带来多少收益。3. 从零手写残差块代码、参数与显存账3.1 最小可运行版本先用最少的代码把核心结构写出来方便后面加东西。下面是一个标准的 BasicBlock对应 ResNet-18/34import torch import torch.nn as nn class BasicBlock(nn.Module): expansion 1 def __init__(self, in_ch, out_ch, stride1): super().__init__() self.conv1 nn.Conv2d(in_ch, out_ch, 3, stridestride, padding1, biasFalse) self.bn1 nn.BatchNorm2d(out_ch) self.conv2 nn.Conv2d(out_ch, out_ch, 3, stride1, padding1, biasFalse) self.bn2 nn.BatchNorm2d(out_ch) self.relu nn.ReLU(inplaceTrue) # 维度对齐要么是恒等要么是 1x1 投影 if stride ! 1 or in_ch ! out_ch: self.shortcut nn.Sequential( nn.Conv2d(in_ch, out_ch, 1, stridestride, biasFalse), nn.BatchNorm2d(out_ch) ) else: self.shortcut nn.Identity() def forward(self, x): identity self.shortcut(x) out self.conv1(x) out self.bn1(out) out self.relu(out) out self.conv2(out) out self.bn2(out) out out identity # 相加发生在第二个 BN 之后、ReLU 之前 out self.relu(out) return out有几个位置细节必须说清楚因为它们直接决定训练是否稳第一biasFalse。因为后面紧跟 BatchNormBN 自己会减均值卷积的偏置项会被完全抵消掉留着纯粹浪费参数。这一点在几乎所有现代 CNN 实现里都是默认做法。第二相加的位置。经典版本ResNet v1是BN → 相加 → ReLUReLU 在相加之后。而 pre-activation 版本ResNet v2把顺序改成了BN → ReLU → Conv相加之后不再接任何非线性。这两者的差别在下文 3.1 之后和 4.1 都会提到先记住结论v2 的恒等通路更干净深层时更容易训。第三nn.Identity()不能省。有些写法直接写self.shortcut None前向里判空。功能上一样但nn.Identity()在打印模型结构时更清晰导出 ONNX 之类的格式时也不容易出问题。3.2 瓶颈结构怎么省参数ResNet-50 及以上用的是 Bottleneck思路是先用 1×1 把通道压下来做一次 3×3 卷积再用 1×1 升回去。核心结构长这样class Bottleneck(nn.Module): expansion 4 def __init__(self, in_ch, mid_ch, stride1): super().__init__() out_ch mid_ch * self.expansion self.conv1 nn.Conv2d(in_ch, mid_ch, 1, biasFalse) self.bn1 nn.BatchNorm2d(mid_ch) self.conv2 nn.Conv2d(mid_ch, mid_ch, 3, stridestride, padding1, biasFalse) self.bn2 nn.BatchNorm2d(mid_ch) self.conv3 nn.Conv2d(mid_ch, out_ch, 1, biasFalse) self.bn3 nn.BatchNorm2d(out_ch) self.relu nn.ReLU(inplaceTrue) if stride ! 1 or in_ch ! out_ch: self.shortcut nn.Sequential( nn.Conv2d(in_ch, out_ch, 1, stridestride, biasFalse), nn.BatchNorm2d(out_ch) ) else: self.shortcut nn.Identity() def forward(self, x): identity self.shortcut(x) out self.relu(self.bn1(self.conv1(x))) out self.relu(self.bn2(self.conv2(out))) out self.bn3(self.conv3(out)) out out identity return self.relu(out)为什么中间通道要压到输出的四分之一算一笔账就明白了。假设输入输出都是 256 通道同时做两个选择方案 A直接用两个 3×3 卷积。参数量是256 × 256 × 9 × 2 1179648。方案 BBottleneck中间压到 64 通道。参数量是256×64 64×64×9 64×256 16384 36864 16384 69632。差了将近 17 倍。这就是 Bottleneck 能在同样显存预算下把网络推到 152 层的直接原因。压缩比通常取 4也有取 2 或 8 的变体取决于任务对通道容量的需求。压缩比太大中间层表达能力不足太小省不下参数。4 是在 ImageNet 上被反复验证过的甜点值。3.3 参数与FLOPs的手工核算接着上面算。一个完整的 Bottleneck输入 256输出 256不改变分辨率的参数构成是conv1256 × 64 × 1 × 1 16384conv264 × 64 × 3 × 3 36864conv364 × 256 × 1 × 1 16384BN 参数2 × 64 2 × 64 2 × 256 768shortcutnn.Identity()0 个参数合计约70400个参数。三个卷积占了 69632BN 只占 1% 左右。所以粗算网络规模时忽略 BN 参数完全没问题。再看 FLOPs。一个 3×3 卷积在 H×W 特征图上的乘加次数是C_in × C_out × 9 × H × W。以 14×14 特征图、256→64→64→256 的 Bottleneck 为例conv1256 × 64 × 1 × 14 × 14 ≈ 3.21Mconv264 × 64 × 9 × 14 × 14 ≈ 7.22Mconv364 × 256 × 1 × 14 × 14 ≈ 3.21M总计约13.6M MACs也就是 27M FLOPs 左右。相比之下如果不用 Bottleneck直接两个 3×3 的 256 通道卷积算下来是256 × 256 × 9 × 14 × 14 × 2 ≈ 231M MACs差了 17 倍和参数量的比例一致因为卷积计算量基本由参数量和特征图尺寸决定。这就是为什么很多轻量化论文喜欢拿 Bottleneck 开刀——它把大部分计算压在了通道数最少的那一层。后续的 MobileNet 用深度可分离卷积本质上是同一个思路的不同实现把计算往低通道、低秩的方向压。3.4 训练侧的几个小开关结构写对了只是第一步训练时还有几个开关要拨。Stochastic Depth随机深度。训练时以概率p随机丢弃残差分支只保留 shortcut。这相当于在训练一个有2^N种深度的隐式集成模型。推理时全部保留并按期望做缩放。它的实现非常轻class StochasticDepth(nn.Module): def __init__(self, p0.1): super().__init__() self.p p def forward(self, x): if not self.training or self.p 0.0: return x keep 1.0 - self.p mask torch.empty(x.size(0), 1, 1, 1, devicex.device).bernoulli_(keep) return x / keep * mask # 用法在残差分支末端套一层 out self.bn3(self.conv3(out)) out self.stochastic_depth(out) out out identity有意思的是正是因为有了 shortcut 这条保底通路丢掉整个残差分支网络依然能跑。这在没有残差连接的网络里是不可想象的——随机扔掉一层卷积输出直接崩掉。所以 Stochastic Depth 反过来也是一个佐证残差连接确实提供了信息通路的冗余性。它能涨点、能加快收敛代价是几乎为零。Layer Scale。做法是给残差分支乘一个可学习的缩放系数初始值设得很小比如 1e-4 或 1e-6。这让训练初期网络的行为近似于一个浅层模型随着训练推进逐渐打开深层分支。它在 ViT 和 ConvNeXt 这类深层 Transformer 里效果明显。Drop Path vs Dropout。在残差分支上用 dropout 时记得加在分支末端而不是中间层且推理时不能忘记关闭。dropout 加在 BN 之前的卷积层后是常见做法但要注意和 BN 的统计量之间的相互影响——训练和推理时的方差估计会不一致深层网络里这个偏差会被累积。相对安全的做法是只在残差分支的最后、相加之前加。4. 残差连接的迁移战场从CNN到Transformer再到扩散模型4.1 Transformer里的pre-LN与post-LNTransformer 的每个子层都是x x Sublayer(x)的形式残差连接是它的骨架。但归一化层的位置有讲究这就是 pre-LN 和 post-LN 的经典之争。Post-LN 是原论文的写法x LayerNorm(x Sublayer(x))Pre-LN 是后来被广泛采用的写法x x Sublayer(LayerNorm(x))差别在哪看残差通路的干净程度。Post-LN 里从输出回传到输入的梯度必须穿过最后一个 LayerNorm而这个 LayerNorm 在深层网络中会显著改变梯度尺度。层数一多靠近输出的层梯度偏大、靠近输入的层梯度偏小导致必须用 warmup 才能训起来——这就是原版 Transformer 训练时前几千步学习率要慢慢爬升的原因之一。Pre-LN 把归一化放进残差分支内部主通路上从头到尾没有任何算子梯度可以一路无阻地回传。代价是最终输出的数值尺度会随着层数累积变大所以在最后一个 block 之后通常还要补一个 LayerNorm。现在的开源实现里pre-LN 已经是默认选择GPT 系列、LLaMA 系列走的都是这条路。我自己在做小规模实验时最大的感受是pre-LN 几乎不需要调 warmup而 post-LN 不 warmup 就会直接发散。这不是玄学是梯度通路结构决定的。4.2 U-Net与扩散模型U-Net 里的 skip connection 是另一种用法编码器某一层的特征直接拼接到解码器对应层。注意这里是拼接不是相加。原因在于两者的目的不同——残差连接是为了保住梯度通路U-Net 的跳跃是为了把编码器里保留下来的高分辨率细节送给解码器让上采样不至于太模糊。拼接保住了全部信息相加则会混合在一起通道数还对不上。在扩散模型里残差连接被用到了极致。整个 U-Net 主干由几十个 ResBlock 堆成每个 block 内部又是两层卷积加残差。更关键的是时间步编码的注入方式时间嵌入向量经过两层 MLP 后通过一个仿射变换scale shift作用在残差分支的中间特征上class ResBlockWithTime(nn.Module): def __init__(self, in_ch, out_ch, time_dim): super().__init__() self.norm1 nn.GroupNorm(8, in_ch) self.conv1 nn.Conv2d(in_ch, out_ch, 3, padding1) self.time_mlp nn.Sequential( nn.SiLU(), nn.Linear(time_dim, out_ch * 2) # 输出 scale 和 shift ) self.norm2 nn.GroupNorm(8, out_ch) self.conv2 nn.Conv2d(out_ch, out_ch, 3, padding1) self.shortcut nn.Conv2d(in_ch, out_ch, 1) if in_ch ! out_ch else nn.Identity() def forward(self, x, t_emb): h self.conv1(nn.functional.silu(self.norm1(x))) scale, shift self.time_mlp(t_emb).chunk(2, dim-1) h self.norm2(h) * (1 scale[:, :, None, None]) shift[:, :, None, None] h self.conv2(nn.functional.silu(h)) return h self.shortcut(x)注意这里有个实现细节scale写作1 scale用的是零初始化的线性层训练开始时这个调制相当于恒等映射让网络先学会基本的去噪再逐步引入时间条件。这个技巧和 Layer Scale 是同一类思想——让新增的模块初始时什么都不做把训练初期的稳定性交给残差通路去兜底。扩散模型还有一个有意思的地方由于去噪任务需要预测的是噪声残差本身整个模型在某种意义上就是在学一个大号的残差函数。这算是残差思想在任务层面的又一次呼应。4.3 图网络、RNN与MLP-Mixer残差连接的适用性远超 CNN。图神经网络里GCN 的每一层写作H σ(A_hat H W)堆深了同样会遇到过平滑问题——所有节点特征收敛到同一个值。加上残差之后变成H σ(A_hat H W) H节点的自身特征被保留下来深层 GCN 才变得可用。这个改动几乎是所有深层图网络论文的标配。RNN 里的残差连接稍微别扭一点因为时间维度本身就有梯度通路加残差主要解决的是深度方向多层 RNN 堆叠的问题。做法通常是h_t RNN(x_t, h_{t-1}) x_t需要维度对齐或者层与层之间加残差。效果没有 CNN 里那么立竿见影因为 LSTM 的门控机制本身已经部分承担了信息保底的功能。MLP-Mixer、ConvNeXt 这类混合架构基本是把 Transformer 的 block 结构原样搬过来残差 归一化 一个空间混合模块 残差 归一化 一个通道混合模块。可以说只要涉及深层堆叠残差连接就已经是默认配置了没有它基本训不动。5. 踩坑实录残差连接最容易翻车的五个地方5.1 维度对不上这是最常见的报错RuntimeError: The size of tensor a (64) must match the size of tensor b (128) at non-singleton dimension 1。原因无非三种通道数不匹配、空间尺寸不匹配、batch 维度被意外操作。排查顺序建议这样先看 stride。如果残差分支里有 stride2 的卷积shortcut 必须同步下采样否则空间尺寸会差一倍。再看通道数如果 conv 的输出通道和输入不同shortcut 必须用 1×1 卷积投影。最后检查有没有中间层的 reshape 或 slice 操作破坏了形状。一个容易漏掉的情况是 padding 设置。3×3 卷积配padding1时输入输出尺寸一致如果某个分支忘了设 padding或者设成了 0尺寸就会差 2。这种错误在浅层网络里可能被下一层的自适应池化掩盖到了深层才暴露出来排查起来很费时间。另外提醒一句用assert在 forward 里做形状断言比等到加法那行才报错要容易定位得多。生产代码里加一两行断言代价极小。5.2 在通路上乱加算子前面提过shortcut 通路上加 ReLU、BN、Dropout 会破坏梯度保底性质。但实践中还有人会在 shortcut 上加别的比如加MaxPool做下采样而不是用 stride 卷积加AvgPool加一层可学习的缩放这个有时是有益的见 6.3用MaxPool做下采样的问题是它不可学习且会丢弃非最大值的信息用 stride 卷积则能在下采样的同时调整通道。原论文里两种都试过stride 卷积略好。用平均池化的话在通道数一致时效果接近恒等但会引入额外的平滑。还有一种隐性污染如果 shortcut 用的是nn.Sequential而里面忘了关relu或者复制粘贴时把一个激活层带进来了这种错误不会报错只会让效果悄悄变差。建议写完之后打印一次模型结构逐层确认 shortcut 分支里只有投影卷积和它的 BN。5.3 数值不稳定与NaN残差连接本身是提升数值稳定性的但在混合精度训练下会出问题。典型场景是残差分支因为梯度爆炸产生了inf而 shortcut 上某个位置正好是 0inf 0还是inf反过来如果分支输出是NaN加任何东西都还是NaN。更隐蔽的是0 × inf这类组合在 fp16 的动态缩放loss scaling机制下容易出现。应对措施有几条第一用torch.cuda.amp时打开梯度裁剪clip_grad_norm_(model.parameters(), 1.0)是个保守但有效的起点。第二注意 BN 或 LayerNorm 的 epsilon 不要设得太小1e-5是默认值1e-8在 fp16 下会溢出。第三如果用了 GroupNorm分组数要能整除通道数否则会静默地产生错误结果有些版本的 PyTorch 不报错只给警告。还有一个常被忽略的点学习率太大时残差分支的输出尺度会瞬间涨上去而 shortcut 提供的保底并不能阻止F(x)自己炸掉。这时缩短 warmup、降低峰值学习率比调结构更有效。5.4 残差分支被喂太饱或饿死训练完看残差分支的响应范数会发现有的层响应很强有的几乎为零。这两种情况都有讲究。响应接近零说明这一层基本被关掉了网络选择绕过它。如果只是少数几层属于正常现象是网络在做自适应深度。如果整片整片地被关掉可能是初始化尺度太小或者学习率过低导致分支从未被有效激活。可以试着调大初始化标准差或者检查是否有过强的权重衰减。响应过强则容易导致训练后期 loss 震荡。这时残差分支已经主导了输出恒等通路被淹没。典型原因是 Layer Scale 的初始值设太大或者没做 warmup 直接上了大学习率。缓解办法是给残差分支加一个固定的缩放系数比如1/√NN 是网络总层数这是 GPT-2 论文里的做法简单有效。判断方法很直接在前向里挂个 hook记录每个残差分支输出的 L2 范数和对应 shortcut 输入的范数做比值。健康的网络里这个比值通常分布在 0.1 到 2 之间极少出现全零或全是一个数量级以上的情况。5.5 问题速查表把上面这些整理成一个速查表出问题时可以按行对照现象可能原因排查方法处理方式加法处形状报错stride/通道不一致打印每步张量形状加 1×1 投影或对齐 stride训练 loss 不降shortcut 上有非线性检查模型结构移除 shortcut 上的 ReLU/BN深层时误差反升未用残差或顺序错对比浅层基线改成 v2 顺序BN-ReLU-Conv混合精度下 NaNfp16 溢出关闭 amp 对比加梯度裁剪、调大 epsilon某些层梯度为 0分支被完全抑制打印梯度范数调初始化、增大学习率后期 loss 震荡分支响应过强打印分支/shortcut 范数比加残差缩放或 Layer Scale6. 调参与初始化的经验值6.1 初始化与gamma缩放残差网络的初始化有几个流传很广的小技巧。最著名的是在 ResNet v2 论文里提到的把每个残差块最后一个 BN 的gamma初始化为 0。这样在训练开始时每个残差块的输出都是 0整个网络从恒等映射开始等价于一个浅层网络。随着训练推进gamma逐渐从 0 长起来网络的有效深度逐步增加。这个做法在 PyTorch 里需要手动处理def zero_init_residual(model): for m in model.modules(): if isinstance(m, Bottleneck): nn.init.constant_(m.bn3.weight, 0) elif isinstance(m, BasicBlock): nn.init.constant_(m.bn2.weight, 0)配合 He 初始化也叫 Kaiming 初始化效果最稳。He 初始化的标准差是√(2/n_in)其中n_in是该层的输入连接数也就是C_in × k × k。这个系数 2 是针对 ReLU 推导出来的——ReLU 会砍掉一半的激活值所以需要把方差翻倍来补偿。卷积层的权重初始化用 He normalBN 的 weight 初始化为 1、bias 为 0。如果用了 Layer Scale初始值一般取1e-4到1e-6具体看网络深度——层数越多值越小。6.2 学习率、warmup与batch size的联动残差结构让深层网络的优化变得容易但它并不免除学习率调度的重要性。经验规律有几条用 SGD 训练 CNN 时初始学习率 0.1 配 batch 256 是一个经典起点。batch 翻倍时学习率线性放大linear scaling rule这个规则在残差网络里依然适用。warmup 步数取总步数的 5% 左右或者固定 500~2000 步。pre-LN 结构可以省掉 warmuppost-LN 必须留着。用 AdamW 时学习率典型值在 1e-4 到 5e-4 之间weight decay 取 0.05这是 ViT 论文里的配置对 CNN 也适用。注意 AdamW 的 weight decay 不要作用在 BN 的 gamma 和 bias 上这些参数上加正则反而有害。还有一个容易被忽视的点残差连接的存在让网络对学习率的容错度更高但一旦残差分支的响应超过 shortcut 太多这个容错度就会消失。所以学习率的上限本质上是残差分支输出尺度不爆炸的上限。6.3 残差缩放的几个流派当网络深到几十上百层时直接y F(x) x会让输出的方差随层数累积放大。于是有了几种缩放方案固定缩放y F(x) × 1/√N xN 是残差块总数。GPT-2 用的是这个。零参数代价是推理时也没法恢复。可学习缩放Layer Scaley F(x) × diag(λ) xλ是可学习向量初始化为小值。ViT、ConvNeXt v2 用的是这个。多了一点参数每个通道一个但灵活性更高。零初始化加缩放ReZeroy F(x) × α xα初始化为 0。训练初期整个网络是恒等映射然后逐渐展开。它在浅层 Transformer 上效果很好深层时需要配合更细致的调参。Fixup 初始化不引入任何缩放系数而是通过精心设计每一层的初始化方差来抵消方差累积。比如把残差分支里最后一个卷积的权重初始化为 0。选哪个我的经验是网络不超过 50 层直接恒等映射就行不用折腾。50 到 100 层加上1/√N的固定缩放几乎无成本。100 层以上或者 Transformer 结构上 Layer Scale。ReZero 和 Fixup 更适合有理论基础的研究场景工程里用得不多。值得强调的是这些缩放方案都是在恒等通路之外做的加法不会破坏∂y/∂x里的那个1。这一点必须守住否则前面所有的努力都白费。最后分享一个我个人排查残差网络问题的习惯把模型里每一个残差块的输入范数、分支输出范数、相加后范数都 hook 下来跑一个前向画三条曲线对比。健康的网络里相加后的范数应该略大于输入范数而不是数量级的跳变如果出现某一层范数突然掉了两个数量级那一层几乎肯定有问题往前后各看两层就能定位。这个土办法帮我省下的调试时间比任何论文都多。
返回列表