Skip to content

s15 序列模型:RNN → LSTM → GRU

WARNING

🧪 Beta公测版本提示:教程主体已完成,正在优化细节,欢迎大家提Issue反馈问题或建议。

文本是有顺序的——"我爱你"和"你爱我"是两回事。序列模型专门处理这种时序数据。词向量从哪来见 文本表示;抛弃循环、改用全体互看见 Transformer。本章把 RNN 的连乘梯度账算清楚,再看 LSTM 那条加法公路为什么能让句首主语活到句末。


一、为什么序列需要专门的模型?

传统的全连接网络(MLP)和卷积网络(CNN)在处理序列数据时有根本性的局限:

MLP 的问题

  • 输入维度固定——无法处理变长序列
  • 每个输入位置独立处理——"我/爱/你"三个词分别进入三层神经元,没有时序关联
  • 参数与位置绑定——第 1 个词的权重只能学第 1 个位置的特征

CNN 的问题

  • 卷积核有固定感受野——只能看到局部上下文
  • 虽然可以通过堆叠层增大感受野,但长距离依赖仍然难以建模
  • 不是为序列专门设计的,缺乏显式的时序记忆机制

序列模型的核心需求

  1. 变长输入处理能力
  2. 参数跨时间步共享(同一套参数处理不同位置)
  3. 显式的记忆机制,能捕捉长距离依赖
  4. 输入顺序敏感

循环神经网络(RNN)通过一个优雅的循环结构同时满足了以上所有需求。


二、RNN:循环的魔力

2.1 一个 cell 是什么

外面看到的「五个蓝块」不是五套网络,是同一台小机器用了五次。这台小机器叫 cell(细胞):每次只吃当前这一个字上一格留下的记忆,吐出这一格的新记忆(以及可选的输出)。

RNNCell:(xt,ht1)ht

PyTorch 里 nn.RNNCell 就是这一格;nn.RNN 是把这一格在整段序列上自动循环。后面 LSTM / GRU 同理:LSTMCell / GRUCell 走一步,LSTM / GRU 走整段。世界模型里的 RSSM 之所以用 GRUCell 而不是 GRU,就是每一步还要插先验 / 后验,必须自己控循环。

2.2 核心公式

RNN 细胞内部只有一本账——隐藏状态 h。没有「先决定留多少、再决定写多少」,每一步都是整包搅匀:

ht=tanh(Whht1+Wxxt+b)
  • htRdh:时间步 t隐藏状态(hidden state),是网络此刻的全部记忆
  • ht1:上一格记忆——已经揉进了 x1,,xt1
  • xtRdx:当前这个字 / 这一帧
  • WhRdh×dh:记忆到记忆的权重(循环连接,五个蓝块共用这一套)
  • WxRdh×dx:当前输入写进记忆的权重
  • tanh:把每一维压到 (1,1),防止数值炸掉

核心直觉ht 是「刚才记住的」和「现在读到的」线性混合,再挤过 tanh。像人读书:每读一个词,旧印象和这个词搅在一起,变成新印象。代价是:旧印象没有原路可走,必须整包过矩阵和非线性。

demo 里对应 MyRNNCellh = tanh(W_ih(x) + W_hh(h_prev))。完整实现见 code-demo,文件在仓库 docs/applied/nlp/sequence-models/code/demo.py

2.3 时间展开(Unrolling)

同一个细胞(同一套 Wh,Wx)在不同时间步被反复调用。把时间拉开,它看起来像很深的全连接网——每一层一个时间步,但所有层共享参数

x_1 → [RNN] → h_1 → [RNN] → h_2 → [RNN] → h_3 → ... → h_T
         ↑共享W_h,Wx↑      ↑共享W_h,Wx↑

序列多长都是这一套数,模型体积不随句长增长。这就是「能处理变长输入」的来源。

RNN 时间展开

怎么读这张图:五个蓝块是同一个细胞用了五次。从上往下:字 xtWx 进记忆;从左往右:上一格记忆经 Wh 传到这一格;再往下:记忆经 Wy 变成输出 yt。底栏红箭头是训练时从右往左回传,每倒退一步都乘同一个 Wh

2.4 BPTT:梯度为什么是「乘法」在时间里走

训练时要把最后的损失 L 告诉很早的 h1,好改那一套共享的 Wh。这叫 BPTT(Backpropagation Through Time,沿时间反向传播)。算法和普通反向传播是同一套链式法则,只是这条链沿着时间排开。

Wh 的梯度要累加每一步的贡献(因为每一步都用了它):

LWh=t=1TLtWh

要让第 T 步的损失碰到第 1 步的记忆,必须把相邻两格的影响连乘起来:

Lh1=LhThThT1hT1hT2h2h1=LhTt=2Ththt1

RNN 的前向是 ht=tanh(zt)zt=Whht1+Wxxt,所以相邻两步之间那一项就是

htht1=diag(tanh(zt))Wh

「信息以乘法的方式在时间中传播」说的不是输入里写了个乘号,而是:

旧记忆对更晚记忆的影响,等于一串「tanh 再乘 Wh」连乘。 前向每走一步,旧信息被矩阵打一次折、再被 tanh 挤一次;反传要原路回去,折扣就连乘。

|tanh(z)|1,多数位置远小于 1Wh 的谱范数也常常 <1。于是每倒退一步大约再乘一个小于 1 的因子 γ。打个折扣账:

倒退步数若每步 γ=0.8直觉
10.8上一字还在
100.890.13已经淡了
200.8190.014几乎没了
500.849105句首梯度到不了句末

这就是梯度消失:不是公式写错了,是这条乘法链太长。若每次 γ>1,连乘则会梯度爆炸。长句、长轨迹里,RNN 学不会「第一句的主语管最后那个动词」,根子在这里。

逐步推导:从 ht=tanh(Whht1+) 到连乘 diag(tanh)Wh(点击展开)

ht=tanh(zt)zt=Whht1+Wxxt。对向量值 tanh 逐元素,雅可比是对角阵 diag(tanh(zt)),再右乘 Wh(因为 zht1 线性)。于是

htht1=diag(1tanh2(zt))Wh.

tanh1,等号只在 0 处。T 步连乘后,若每步谱半径 γ<1,范数以 γT 掉。这和 导数 的链式法则是同一句话,只是链沿着时间排了 T 节,而且每节共用同一个 Wh

LSTM 把「对 h 的连乘」改成「对 c 的连乘 ft」。若遗忘门接近 1,连乘不衰减,梯度可以沿细胞状态走很远。门是 sigmoid 出来的,可以学习「这一维先别忘」。

梯度消失可视化

怎么读这张图:上半是链式法则拆成「每步一个雅可比」;下半对数坐标里,普通 RNN 的 L/ht 往回走直线往下掉,LSTM 的细胞状态几乎走平。


三、LSTM:另开一条加法公路,再装三个门

LSTM(Long Short-Term Memory,Hochreiter & Schmidhuber, 1997)没有改 BPTT 这套算法,改的是细胞前向怎么走:不要让长期记忆每一步都过 Whtanh

3.1 一个 cell 上,RNN 和 LSTM 差在哪

RNN cellLSTM cell
保管的状态只有 ht两本账ct(长期笔记)和 ht(这一步对外输出)
一步接口(xt,ht1)ht(xt,ht1,ct1)(ht,ct)
旧记忆怎么变成新的整包:ht=tanh(Whht1+Wxxt)公路上ct=ftct1+itc~t
有没有开关三个门 f,i,o,每个都是 01 的旋钮
反传时相邻两步乘什么diag(tanh)Whc 上主要是遗忘门 ft

同一个时间步:RNN cell 与 LSTM cell

怎么读这张图:左栏是 RNN——ht1xt 搅匀过 tanh 就变成 ht。右栏顶上那条粗线是细胞状态 c,旧笔记乘遗忘门、新内容乘输入门,再在一起;下面四个色块是从 [ht1,xt] 拧出来的旋钮。底栏:训练回传时,RNN 连乘 tanhWh,LSTM 在公路上连乘 ft

3.2 「乘法传播」对应哪条公式,门要解决什么

上一节 2.4 BPTT 里,RNN 相邻两步乘的是 diag(tanh)Wh。有用的主语和没用的语气词都挤在同一个 h 里,每一步还被 tanh 再压一遍。我们需要的是:

  1. 一条可以几乎原样往前加的笔记 ct(不要每步搅匀)
  2. 一组学出来的 0~1 开关,决定这条笔记「擦掉哪几维、写入哪几维、对外露出哪几维」

门不是 discrete 的 0/1 电闸,是 sigmoid 拧出来的连续旋钮,这样才能对 Wf 求导。

3.3 细胞状态 ct 怎样引进来

先不管门,只看 LSTM 最狠的那一行——给记忆另开一本账,默认用加法更新:

ct=ftct1+itc~t
  • ct1:上一格的长期笔记(可以一路从句首抬过来)
  • ft:遗忘门,逐维决定旧笔记留几成( 是逐元素乘,每个记忆槽位自己的开关)
  • c~t:根据当前字新写的候选内容
  • it:输入门,决定新内容写进笔记几成

当某一维 ft=1it=0 时,这一维就是 ct=ct1原样拷贝,不过 tanh,也不乘 Wh 这就是「信息高速公路」。有东西要记时再把 it 拧开,把 c~t 上去——加,而不是整包替换。

3.4 三个门是怎样从 [ht1,xt] 拧出来的

门不看 c 本身(经典 LSTM 如此),只看「上一刻对外说了什么」和「现在读到什么」:把 ht1xt 拼接成一根长向量,各乘一套权重,再过 sigmoid / tanh

遗忘门 — 旧笔记留几成:

ft=σ(Wf[ht1,xt]+bf)

输入门 — 新内容写几成:

it=σ(Wi[ht1,xt]+bi)

候选细胞状态 — 新内容本身(仍用 tanh 压到 (1,1)):

c~t=tanh(Wc[ht1,xt]+bc)

输出门 — 笔记对外露几成(c 是内部账本,h 才是这一步给下一层、给下一步门看的):

ot=σ(Wo[ht1,xt]+bo)ht=ottanh(ct)

σ 把值挤到 (0,1),所以叫「门」:0 关、1 开、中间半开。四个线性层在代码里常合成一次大矩阵乘,再 chunk 成四段(见 MyLSTMCell)。

读「我爱机器学习!」时可以这么想象(一维开关的卡通版):

  1. 读到「我」:输入门打开,主语写进 c 的某一维
  2. 读中间修饰:「爱」「机器」「学习」——遗忘门接近 1,主语那一维几乎原样加下去
  3. 读到「!」:也许拧小某些句法槽;输出门决定这一步的 h 要不要强调句末语气

3.5 三门公式总表

LSTM 三门详解

怎么读这张图:从左进 ct1ht1xt。橙色遗忘门乘在公路上;绿色输入门和新候选 c~ 相乘后进公路;紫色输出门从 ct 滤出 ht。黄框那行 ct=fct1+ic~ 就是加法路径。

demo 里对应的三行就是整章的核心:

python
c = f * c_prev + i * c_tilde   # 公路:留旧 + 写新
o = torch.sigmoid(o_gate)
h = o * torch.tanh(c)          # 对外只露一页笔记

3.6 门的直觉

作用直觉
遗忘门 ftft0:这一维旧笔记清掉「读到句号,清空前文句法槽」
输入门 itit1:把 c~t 写入「遇到主语,记下谁在做事」
输出门 otct 滤出 ht「答题时只抄笔记里此刻用得上的几行」

LSTM 像一个有条理的学生做笔记:遗忘门决定擦掉哪几行,输入门决定写下新知识点,输出门决定举手发言时念哪几行。笔记本本身是 c,发言内容是 h

3.7 为什么这样梯度就不易消失

对公路本身、在某一维上求导(ft 暂时看成对 ct1 常数——门由 h,x 算出来,不直接含 c):

ctct1=ft

若遗忘门学会 ft1(「这段主语还得留着」),这一步梯度是 ×1,不是 ×(diag(tanh)Wh)。从 cT 走回 c1

cTc1=fTfT1f211

长期内容可以几乎原样走回句首。这是加法公路,不是每步搅匀。

两点不要推过头:

  • 门自己的权重 Wf,Wi,Wo 仍然要经过 sigmoid / tanh 反传,那些旁路还是有非线性。LSTM 减轻的是长期内容 c 这条主干
  • ft 若长期接近 0,这一维照样断。模型要学会「该留的时候把遗忘门拧到 1」——这也是为什么常把遗忘门偏置初始化成正数,训练初期先倾向于「多记住」。

一句话:RNN 的 cell 把记忆整包乘进下一步;LSTM 的 cell 把记忆放在 c 里按元素加,三个门只是学出来的 0~1 开关。 反传仍是 BPTT,变的是这条链上每一步乘的是 Whtanh 还是 ft

3.8 常见疑问

门是离散的开/关吗? 不是。σ 输出落在 (0,1) 开区间,训练中是连续旋钮。说「打开 / 关掉」只是把靠近 1 / 靠近 0 说成开关。

LSTM 改了反向传播算法吗? 没有。还是 BPTT。改的是前向递推:多了一条 c 的加法通路,雅可比从 diag(tanh)Wh 变成(主干上)ft

为什么还要 h,不直接把 c 当输出? c 是内部笔记本,量级可以慢慢攒;对外给下一层、给下一步的门看时,先 tanh 压一压再被 ot 筛选,避免把还没整理的长期笔记整本泄露出去。下一步的三个门吃的是 ht1xt,不直接吃 ct1(经典 LSTM)。

五个蓝块和 cell 是什么关系? 蓝块 = 同一细胞的五次调用。RNN / LSTM 的差别全部发生在一块内部;展开方式、共享参数、BPTT 的「沿时间连乘」框架是一样的。


四、GRU:LSTM 的精简版

Cho et al. (2014) 提出 GRU(Gated Recurrent Unit),把 LSTM 的三个门收成两个,并且不再单独保管 ct:长期记忆和对外输出共用一本 h。主干仍然是加法插值(所以梯度通路和 LSTM 同类),不是 RNN 那种整包 tanh

重置门(reset gate)— 控制忽略多少历史信息:

rt=σ(Wr[ht1,xt])

更新门(update gate)— 控制保留多少旧状态 vs 写入多少新状态:

zt=σ(Wz[ht1,xt])

候选隐藏状态— 用重置门过滤后的历史 + 当前输入:

h~t=tanh(Wh[rtht1,xt])

最终隐藏状态— 更新门做线性插值:

ht=(1zt)ht1+zth~t

GRU 的核心直觉是 zt(更新门)同时做了 LSTM 遗忘门和输入门的工作。当 zt0 时,htht1(保留全部历史);当 zt1 时,hth~t(完全更新为新状态)。


五、RNN vs LSTM vs GRU 对比

特性RNNLSTMGRU
门数量032
状态变量htht, ctht
梯度传播指数衰减加法路径(稳定)加法路径(稳定)
参数量2dh(dh+dx)4dh(dh+dx)3dh(dh+dx)
训练速度中等
长序列表现最好
典型场景简单时序预测机器翻译、复杂序列当 LSTM 太大时替代

RNN vs LSTM vs GRU 架构对比


六、双向 RNN

标准 RNN/LSTM/GRU 只能从左到右处理序列——t 时刻的隐藏状态只包含 t 之前的信息。但在很多 NLP 任务中,t 时刻的输出需要同时利用左右两侧的上下文。

双向 RNN(Bidirectional RNN)同时运行两个独立的循环网络:

  • 前向 RNN:从左到右处理,ht=RNN(xt,ht1)
  • 后向 RNN:从右到左处理,ht=RNN(xt,ht+1)
  • 拼接输出ht=[ht;ht]

双向 RNN 在序列标注(NER、词性标注)和文本分类中极其有效。但无法用于自回归生成(因为你无法看到"未来"的词)。


七、RNN vs Transformer:时代的交替

2017 年 Transformer 出现后,RNN 系模型在 NLP 中的主导地位逐渐被取代。但这并不意味着 RNN 不再重要:

场景选择
长序列(>2048 tokens)且追求最优效果Transformer(全局自注意力)
流式/实时处理、逐时间步推理RNN/LSTM(自然支持)
计算资源受限GRU(参数少、推理快)
时间序列预测(金融、传感器)LSTM(仍广泛使用)
学习 RNN 原理、BPTT、门控机制必须掌握(本章重点)

学习价值:RNN→LSTM→GRU→Transformer 这条技术演进路线的每一步都解决了一个明确的问题。只有理解了每一步"为什么",才能真正理解 Transformer 的注意力机制"好在哪里"。


八、本节小结

概念一句话总结
cell一步映射;RNN 是 (x,h)h,LSTM 是 (x,h,c)(h,c)
RNN同一套参数在时间上循环;记忆整包过 Whtanh
乘法传播ht/ht1=diag(tanh)Wh,连乘导致梯度消失
BPTT仍是链式法则,沿展开后的时间往回传;LSTM 没改这套算法
细胞状态 ct加法公路 ct=fct1+ic~ct/ct1=ft
遗忘 / 输入 / 输出门三个 sigmoid 旋钮:留旧、写新、对外露哪几维
GRULSTM 精简版:合并 ch,双门,仍是加法插值
双向 RNN前向+后向处理,适合标注任务
Transformers16 主题,注意力取代循环连接

下一节 s16 Attention 与 Transformer 将讨论:序列模型的 seq2seq 架构遇到什么瓶颈,注意力机制如何优雅地解决它,并最终催生了取代 RNN 的全新范式。

📥 Code

FileViewDownload
demo.pyOpenDownload
exercise.pyOpenDownload

参考

  1. Hochreiter, S. & Schmidhuber, J. (1997). Long Short-Term Memory. Neural Computation. (LSTM) [doi:10.1162/neco.1997.9.8.1735]
  2. Cho, K., et al. (2014). Learning Phrase Representations using RNN Encoder-Decoder for Statistical Machine Translation. EMNLP 2014. (GRU) [arXiv:1406.1078]
  3. Sutskever, I., Vinyals, O., & Le, Q. V. (2014). Sequence to Sequence Learning with Neural Networks. NeurIPS 2014. (Seq2Seq) [arXiv:1409.3215]