
LSTM & GRU —— 门控循环神经网络
系统讲解LSTM与GRU的门控机制设计哲学与数学原理。从LSTM的“双轨”架构(细胞状态+隐藏状态)与三门结构(遗忘门、输入门、输出门)出发,推导细胞状态更新公式如何通过按元素加法路径缓解梯度消失,并深入剖析GRU的两大简化及其参数量优势。附带讨论双向LSTM与深层堆叠LSTM的应用场景。
阅读文章ZHY's Blog
A UNIVERSE OF IDEAS · BY ZHANG HAOYI
让好奇心 点亮知识宇宙
在代码、模型与思想之间自由漫游。这里持续记录人工智能、机器学习、软件工程与成长实践,让每次阅读都成为一次新的发现。
ARTICLE NOTE
在上一篇博客中,我们系统掌握了卷积、感受野与池化等空间特征提取的核心工具。凭借这些操作,VGG将网络推至19层,GoogLeNet达到22层——一个自然而然的信念随之形成:网络越深,表达能力越强,性能理当越好。
然而,2015年一个令人不安的实验结果打破了这一直觉:在CIFAR-10上,一个56层的“纯”卷积网络,其训练误差竟显著高于20层的网络。这不是过拟合(训练误差高而非低),而是退化问题 —— 深层网络在优化层面遇到了“天花板”:梯度可以借助Batch Normalization正常流动,但求解器却难以将一堆非线性层训练成“什么都不做”的恒等映射。
本篇博客正是破解这一困局的关键章节。我们将从退化问题的本质出发,揭示为何让网络学习 H(x)=x 如此困难,而引入残差块 F(x)+x 后,目标便转换为学习 F(x)=0 —— 这一看似微小的结构变化,却使152层的ResNet在ImageNet上超越人类水平。随后,我们将从链式法则出发,严格证明残差连接如何通过恒等映射的直通路径缓解梯度消失 —— 深层特征由连乘变为连加,梯度中始终包含一个无衰减的“1”。
在此基础上,我们将详解Bottleneck块如何通过1×1→3×3→1×1的结构将参数量减少约94%,使千层网络在有限显存下成为可能;并对比ResNet(加法融合)与DenseNet(拼接融合)在梯度流动与特征复用上的本质差异。
值得注意的是,本篇对“跨层直连”的设计哲学,将直接服务于后续RNN与BPTT——当时间步长达数百时,正是残差思想启发了LSTM中的门控机制。现在,请带着“如何让梯度无损穿越深度”的疑问进入正文——理解了ResNet,您就掌握了现代所有深度架构(从Transformer到Diffusion Model)得以存在的底层基石。
在深入ResNet之前,有必要区分两个经常被混淆的概念:
| 问题 | 表现 | 根本原因 |
|---|---|---|
| 梯度消失/爆炸 | 网络无法收敛(loss不下降) | 反向传播时梯度逐层相乘,趋于0或∞ |
| 退化问题 | 网络能收敛,但更深时准确率下降 | 深层网络难以优化,求解器无法找到最优解 |
Batch Normalization(BN)的提出解决了梯度消失/爆炸问题,使得数十层的网络能够收敛。但BN并未解决退化问题——即使网络能够收敛,56层的plain network仍然比20层差。
考虑一个浅层网络已经达到了不错的性能。现在我们在这个网络上追加若干层,构成一个更深的网络。
从理论上讲,如果追加的这些层什么也不做(即实现恒等映射 y=x),那么深层网络的性能至少不会比浅层差。
但问题在于:让一个由多个非线性层(卷积+激活)组成的模块实现恒等映射,恰恰是神经网络最难做的事情之一。
为什么呢?
一个典型的卷积块包含:卷积 → BN → ReLU → 卷积 → BN。要让这个模块的输出等于输入,需要精确地调整所有卷积核的权重,使得两层卷积的复合效果恰好是恒等映射。这相当于求解一个高度非线性的方程组——在随机初始化的情况下,几乎不可能通过梯度下降恰好收敛到这样的解。
ResNet的洞见:与其让网络学习 H(x)=x,不如让网络学习 F(x)=H(x)−x,然后通过 H(x)=F(x)+x 来构造输出。
这样,实现恒等映射只需要让 F(x)=0 ——让所有卷积层的输出为0,比让它们精确地实现恒等映射要容易得多!
一个基本的残差块可以表示为:
y=F(x,{Wi})+x其中:
对于包含两层卷积的Basic Block:
F(x)=W2⋅σ(W1⋅x)其中 σ 是ReLU激活函数。
整个残差块的前向传播为:
y=W2⋅σ(W1⋅x)+x假设我们要将输入 x=5 映射到目标输出 H(x)=5.1。
如果目标从5.1变为5.2:
残差结构对输出的变化更敏感,这使得梯度更新时权重的调整幅度更大,学习效率更高。
这是理解ResNet最核心的部分。让我们从数学上证明残差连接如何缓解梯度消失。
普通网络(无残差连接):
假设第 l 层的输出为 xl+1=Fl(xl),其中 Fl 是第 l 层的非线性变换。
根据链式法则,损失 L 对第 l 层输入的梯度为:
∂xl∂L=∂xL∂L⋅i=l∏L−1∂xi∂xi+1=∂xL∂L⋅i=l∏L−1∂xi∂Fi如果每一层的雅可比矩阵 ∂Fi/∂xi 的谱范数都小于1(在饱和激活函数下很容易发生),那么连乘的结果会指数级衰减到0——这就是梯度消失。
残差网络:
残差块的前向传播为:
xl+1=xl+Fl(xl)从第 l 层到第 L 层(L>l):
xL=xl+i=l∑L−1Fi(xi)这是关键:在普通网络中,深层特征是浅层特征的连乘;而在残差网络中,深层特征是浅层特征的连加!
现在计算梯度:
∂xl∂L=∂xL∂L⋅∂xl∂xL=∂xL∂L⋅(1+i=l∑L−1∂xl∂Fi)梯度由两项组成:
🔑 核心结论:残差连接为梯度提供了一条“高速公路”——无论网络有多深,梯度至少有一条路径可以直接从输出层流到输入层,而不经过任何带参数的层。
这就是为什么ResNet可以训练到1000层以上,而普通网络在30层左右就已经无法训练了。
何恺明等人在后续论文《Identity Mappings in Deep Residual Networks》中对残差块进行了改进。
ResNet v1(原始版本):
xl+1=f(xl+F(xl))其中 f 是ReLU激活函数(后激活)。
ResNet v2(改进版本):
xl+1=xl+F(f(xl))将BN和ReLU移到残差分支内部,恒等映射路径上没有任何操作(预激活)。
为什么这样更好?
在v1中,恒等映射路径上还有一个ReLU(f),这破坏了“纯净”的恒等映射。在v2中,恒等映射路径完全无参数、无激活,信号可以真正无损地传播。
实验表明,ResNet v2的训练速度更快,泛化性能更好。
Basic Block(两个3×3卷积)在ResNet-18和ResNet-34中表现良好。但当网络加深到50层以上时,Basic Block的参数量会急剧膨胀。
以ResNet-50为例:如果全部使用Basic Block,参数量将远超硬件承载能力。
Bottleneck Block应运而生,用于深层网络(ResNet-50/101/152)。
Bottleneck Block采用**“1×1 → 3×3 → 1×1”** 的三层结构:
以输入通道256、输出通道256为例:
Basic Block(两个3×3卷积):
参数量=256×3×3×256+256×3×3×256=1,179,648Bottleneck Block(1×1→3×3→1×1):
参数量=256×1×1×64+64×3×3×64+64×1×1×256=16,384+36,864+16,384=69,632(1\times1降维)(3\times3特征提取)(1\times1升维)Bottleneck的参数量仅为Basic Block的 1,179,64869,632≈5.9% !
💡 1×1卷积的本质:1×1卷积在每个像素位置上对所有通道进行线性组合,相当于一个跨通道的全连接层。它不关心空间信息,只负责通道维度的信息整合。
1import torch2import torch.nn as nn3
4class Bottleneck(nn.Module):5 """6 ResNet Bottleneck Block7 适用于 ResNet-50/101/1528 """9 expansion = 4 # 输出通道数 = 输入通道数 * 410
11 def __init__(self, in_channels, out_channels, stride=1):12 super().__init__()13 # 1x1 降维:将通道数压缩到 out_channels14 self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=1, bias=False)15 self.bn1 = nn.BatchNorm2d(out_channels)16
17 # 3x3 空间特征提取18 self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3,19 stride=stride, padding=1, bias=False)20 self.bn2 = nn.BatchNorm2d(out_channels)21
22 # 1x1 升维:恢复到 out_channels * expansion23 self.conv3 = nn.Conv2d(out_channels, out_channels * self.expansion,24 kernel_size=1, bias=False)25 self.bn3 = nn.BatchNorm2d(out_channels * self.expansion)26
27 self.relu = nn.ReLU(inplace=True)28
29 # Shortcut路径:如果维度不匹配,用1x1卷积调整30 self.shortcut = nn.Sequential()31 if stride != 1 or in_channels != out_channels * self.expansion:32 self.shortcut = nn.Sequential(33 nn.Conv2d(in_channels, out_channels * self.expansion,34 kernel_size=1, stride=stride, bias=False),35 nn.BatchNorm2d(out_channels * self.expansion)36 )37
38 def forward(self, x):39 identity = self.shortcut(x) # 恒等映射路径40
41 # 残差路径42 out = self.relu(self.bn1(self.conv1(x)))43 out = self.relu(self.bn2(self.conv2(out)))44 out = self.bn3(self.conv3(out))45
46 # 残差 + 恒等映射47 out += identity48 out = self.relu(out)49 return out📌 设计细节:
expansion=4表示输出通道是输入通道的4倍(如64→256)- 当
stride=2时,特征图尺寸减半,shortcut路径也需要用步长为2的1×1卷积来匹配尺寸- 所有卷积层
bias=False,因为BN层已经包含了偏置项
DenseNet(密集连接卷积网络)由黄高等人于2016年提出,获得了CVPR 2017最佳论文奖。
DenseNet的基本思路与ResNet一致——建立跨层连接来改善梯度流动。但实现方式截然不同:
| ResNet | DenseNet | |
|---|---|---|
| 连接方式 | 加法(element-wise addition) | 拼接(concatenation) |
| 信息传递 | 只连接前一层的输出 | 连接前面所有层的输出 |
| 特征复用 | 隐式(通过残差学习) | 显式(所有层共享特征) |
DenseNet中,第 l 层的输入是前面所有层输出的拼接:
xl=Concat(x0,x1,…,xl−1)在ResNet中,每一层都会产生新的特征图,网络宽度(通道数)会逐渐增加。
在DenseNet中,由于每一层都能接收到前面所有层的特征,不需要每一层都学习很多特征图。DenseNet将每一层设计得很“窄”——通常每层只学习 k=12 或 32 个特征图(称为增长率(growth rate) )。
这种设计使得DenseNet在参数效率上优于ResNet。
ResNet的梯度路径:
DenseNet的梯度路径:
从梯度传播的角度看,DenseNet比ResNet更激进——它建立了 O(L2) 条连接(L 为层数),确保每一层都能直接与所有后续层通信。
| 优点 | 缺点 |
|---|---|
| 缓解梯度消失 | 显存占用大(需要保存所有中间特征图) |
| 特征重用,参数效率高 | 前向计算量较大 |
| 减轻过拟合 | 实现复杂度高于ResNet |
1class DenseLayer(nn.Module):2 """DenseNet的单个密集层"""3 def __init__(self, in_channels, growth_rate):4 super().__init__()5 # 先通过1x1卷积降维(Bottleneck设计)6 self.bn1 = nn.BatchNorm2d(in_channels)7 self.conv1 = nn.Conv2d(in_channels, 4 * growth_rate, kernel_size=1, bias=False)8 self.bn2 = nn.BatchNorm2d(4 * growth_rate)9 self.conv2 = nn.Conv2d(4 * growth_rate, growth_rate, kernel_size=3,10 padding=1, bias=False)11
12 def forward(self, x):13 # 注意:这里用拼接,不是加法!14 out = self.conv1(F.relu(self.bn1(x)))15 out = self.conv2(F.relu(self.bn2(out)))16 return torch.cat([x, out], dim=1) # 在通道维度上拼接17
18class DenseBlock(nn.Module):19 """Dense Block:包含多个Dense Layer"""20 def __init__(self, in_channels, growth_rate, num_layers):21 super().__init__()22 self.layers = nn.ModuleList()23 for i in range(num_layers):24 self.layers.append(DenseLayer(in_channels + i * growth_rate, growth_rate))25
26 def forward(self, x):27 for layer in self.layers:28 x = layer(x)29 return x| 概念 | 核心思想 | 关键公式 |
|---|---|---|
| 退化问题 | 深层网络难以优化,不是过拟合 | 56层训练误差 > 20层 |
| 残差块 | 让网络学习 F(x)=H(x)−x 而非 H(x) | H(x)=F(x)+x |
| 梯度缓解 | 恒等映射提供梯度直通路径 | ∂L/∂xl=∂L/∂xL⋅(1+∑∂Fi/∂xl) |
| Bottleneck | 1×1降维→3×3提取→1×1升维 | 参数量减少约94% |
| ResNet v2 | 恒等映射路径纯净无操作 | 预激活(Pre-activation) |
| DenseNet | 拼接而非加法,连接所有前层 | xl=Concat(x0,…,xl−1) |
ResNet的遗产远不止于计算机视觉。从GPT到BERT,从Transformer到Diffusion Model,残差连接已经成为了几乎所有现代深度架构的基础组件。可以说,没有残差连接,就没有今天的大模型时代。
理解ResNet,不仅是在理解一个卷积网络架构,更是在理解深度学习如何克服了“深度”本身带来的诅咒。
延伸阅读:
按顺序完成这组文章,循序渐进地掌握主题
发现错误、内容过时或有改进想法?欢迎告诉我
根据本文分类与标签,为你推荐可能感兴趣的内容

系统讲解LSTM与GRU的门控机制设计哲学与数学原理。从LSTM的“双轨”架构(细胞状态+隐藏状态)与三门结构(遗忘门、输入门、输出门)出发,推导细胞状态更新公式如何通过按元素加法路径缓解梯度消失,并深入剖析GRU的两大简化及其参数量优势。附带讨论双向LSTM与深层堆叠LSTM的应用场景。
阅读文章
系统讲解循环神经网络(RNN)的数学原理与随时间反向传播(BPTT)算法。从RNN的循环结构与共享参数出发,推导BPTT的梯度表达式,揭示梯度消失与梯度爆炸的数学根源,并介绍梯度裁剪作为应对梯度爆炸的工程解法。附带讨论LSTM如何通过门控机制缓解梯度消失。
阅读文章
系统讲解激活函数与权重初始化的协同演进关系。从Sigmoid/Tanh的梯度饱和问题出发,推导ReLU及其变体(LeakyReLU、PReLU、ELU、GELU、Swish)如何解决梯度消失,并深入推导Xavier初始化(适用于Sigmoid/Tanh)和Kaiming初始化(适用于ReLU)的方差守恒数学原理。
阅读文章请使用微信扫描二维码分享
当前文章会保持在原页面