第一部分:为什么需要 RNN?(序列建模问题)
传统 CNN 和全连接网络(FC)的输入和输出大小是固定 的(比如 224x224 的图片输入,输出 1000 个类别)。
但现实中有很多任务输入或输出是可变长度的序列(如文本、语音、视频帧)。
RNN 就是为了处理可变长度序列数据而设计的。根据输入输出长度的不同,分为四种模式(参考图 1):
-
1对1 (One-to-one): 传统神经网络。
-
1对多 (One-to-many): 固定大小输入,变长输出(例如:图像描述,输入一张图,输出一段文字)。
-
多对1 (Many-to-one): 变长输入,固定输出(例如:视频分类,输入多个视频帧,输出一个动作标签)。
-
多对多 (Many-to-many): 变长输入,变长输出,且长度通常相等(例如:逐帧视频分类 、机器翻译,这节课的核心)。
第二部分:RNN 的核心原理
1. "展开图"与循环结构
RNN 有一个隐藏状态(Hidden State) ,它像一个"内存条",保存了过去所有时间步的信息。
展开结构(Unrolled RNN) 是理解 RNN 的最佳方式:我们把时间维度铺开,每一个 RNN 单元不仅接收当前输入 ,还接收上一个单元的输出
。
text
[输入] x1 ---> [RNN] ---> [输出] y1 ↑ ↑ | └── (携带上一步的信息 h1,传给下一步) (x2) ---> [RNN] ---> [输出] y2 ... 以此类推
2. 数学公式
-
隐藏状态更新公式:
-
:旧的隐藏状态(记忆)。
-
:当前时刻的输入。
-
:带有参数 W 的函数(如 tanh 或 ReLU)。
-
关键点: 所有时间步共享同一组权重 W,参数量不随序列长度增加!
-
-
输出公式:
- 使用另一组权重
将隐藏状态转换为输出。注意,隐藏状态和输出的维度可以不同。
- 使用另一组权重
基础 RNN(Vanilla RNN)具体公式:
-
为什么用 tanh?因为它的值域在 -1, 1 之间,有界且零中心化,可以避免数值爆炸。
-
h0 初始化为零向量(或学习得到的参数)。
第三部分:实战举例------手写 RNN 检测"连续 1"
任务: 输入一串 0 和 1,如果当前输入和上一个输入都是 1 ,则输出 1,否则输出 0。
输入序列:[0, 1, 0, 1, 1, 1, 0, 1, 1],期望输出:[0, 0, 0, 0, 1, 1, 0, 0, 1]
手动构造一个 RNN:
-
设定隐藏状态
为 3 维向量:
[当前值, 上一个值, 常数 1]。(作为初学者,可以想象它像一个小账本)。 -
初始化 h0=0,0,1(假设最开始看到了两个 0)。
-
使用 ReLU 作为激活函数(简化计算:大于 0 保留,小于 0 变成 0)。
-
构造权重矩阵(参考图 7):
-
(输入变换):将 x 送入第 1 个位置,即
W_xh = [[1], [0], [0]]。 -
(状态转移):将上一次的"当前值"移入"上一个值"的位置,即
W_hh = [[0,0,0], [1,0,0], [0,0,1]]。
-
-
计算逻辑(以一个具体步骤演示):
假设上一步状态
(上一次输入是1),当前输入
。
(记忆已更新)。
-
输出: 设置
=1,1,−1,则
=ReLU((1∗1)+(1∗1)+(1∗(−1)))=ReLU(1)=1。
💡 结论: 只要构造出合适的权重,普通的 RNN 就能完成这个逻辑任务!但在实际应用中,我们不需要手写权重,而是通过梯度下降自动学习 WW。
第四部分:如何训练 RNN?(难点与解决方案)
1. 多对多任务的损失计算与梯度更新
对于多对多任务,每个时间步都有一个损失 LtLt,总损失是它们之和:L=。
因为所有时间步共享权重 W,所以在反向传播(BPTT)时,我们需要把每个时间步算出的 W 的梯度全部加在一起。
2. 梯度消失问题与截断 BPTT
-
致命缺陷: 当序列很长时,反向传播跨越的时间步很多,如果使用 tanh 或 Sigmoid,它们的导数往往小于 1,连乘之后(∏tanh′)梯度会迅速趋近于 0(梯度消失)。这意味着网络"记不住"很早期的信息(长距离依赖)。
-
解决方案1:截断 BPTT (Truncated Backpropagation through time)
-
你不需要在整个序列上反向传播(耗时且爆内存)。
-
做法: 设置一个窗口(如 N=10),前向传播一直往前跑(记忆保留),但反向传播只倒退最近的 10 步。这样既节省了内存,又近似了梯度。
-
(注:分布式训练时,每个 GPU 算梯度后加起来更新权重,原理类似。)
-
3. 解决方案2:LSTM(长短期记忆网络)
LSTM 是为了解决梯度消失问题而设计的。(参考图13)它的核心是引入一条**"高速公路"** ,叫细胞状态(),沿着这条路径走,信息不经过任何非线性激活函数,只能通过"门"来加减,从而保证梯度能无损流过。
LSTM 公式解析(4个门):
通过拼接输入和隐状态,乘以一个大的权重矩阵得到 i,f,o,g:
-
遗忘门 f :决定上一时刻的
中有多少保留下来。(Sigmoid 函数,0 是忘掉,1 是保留)
-
输入门 i :决定当前的新信息 g 有多少写入到
中。
-
新候选值 g:生成当前的候选记忆。(tanh 函数)
-
细胞状态更新:
=f⊙
+i⊙g(加法操作让梯度畅通无阻!)
-
输出门 o :决定最终输出
的多少。
=o⊙tanh(
)
第五部分:RNN 的高级用法与模型细节
1. 字符级语言模
-
任务: 根据已有字符预测下一个字符。例如:输入 "h",输出 "e";输入 "e",输出 "l"...
-
输入层处理: 我们不会直接用 One-hot 向量做矩阵乘法(因为太稀疏),而是使用嵌入层(Embedding Layer),本质上就是一个大小可学习的矩阵,相当于查表提取密集特征。
-
采样(Sampling): 在测试时,我们用网络预测出下一个字符的概率分布,然后采样一个字符,将其作为下一次的输入,循环往复,直到遇到"结束符"。
2. 多层 RNN 与应用
-
多层 RNN: 将 RNN 层堆叠起来(深度维度),每一层处理上一层的数据,增加模型的表达能力。
-
应用:
-
图像描述: 提取 CNN 倒数第二层的特征,作为 h0 输入到 RNN 中。
-
视觉问答: 将问题和图片特征融合,输入 RNN,输出答案概率。
-
优点: 能处理任意长序列,模型参数量固定。
-
缺点: 无法并行训练(前一步的输出是后一步的输入,必须串行),速度极慢;容易丢失早期信息。
-
🎁 总结备忘
-
输入输出: 只要想到了"变长序列",第一反应就是 RNN 族。
-
共享参数: RNN 的参数在不同时间步是完全共享的,这是和 CNN 共享卷积核有异曲同工之妙的地方。
-
实际现状: 虽然 RNN/LSTM 理论很重要,但在现代 NLP 领域(如 ChatGPT),它们基本已被 Transformer 取代,因为 Transformer 支持高度并行训练,且能通过注意力机制直接捕捉长距离依赖。
-
联想记忆: 你可以把 RNN 想象成**"一面倒的动态长卷"** ,把 LSTM 的细胞状态想象成**"一条笔直的传送带"**,门控就像是控制传送带上物品加减的开关。