观文听傑

返回

上一篇把普通循环神经网络(Recurrent Neural Network, RNN)沿时间展开后,我们看见了长程学习的症结:从损失回到很早的状态,梯度必须反复穿过同一个循环矩阵和 tanh 导数。即使前向状态仍非零,远处的训练信号也可能已经小到无法使用。

长短期记忆网络(Long Short-Term Memory, LSTM)没有取消时间递推。它做的关键改造是:把“对外工作的隐状态”和“沿时间保存的单元状态”分开,再用可学习的门决定旧信息保留多少、新信息写入多少、当前暴露多少。

本文只讲透一个核心问题:LSTM 如何把普通 RNN 的“每步整体重写”改成受控的加法记忆路径。我们会依次拆解四个信号、手算三步状态、追踪梯度,再与 PyTorch 2.13 的官方实现逐张量对齐。

01 普通 RNN 为什么很难选择性地记忆?#

普通 tanh RNN 将旧状态和新输入混在一次变换里:

ht=tanh(Wxhxt+Whhht1+b)h_t=\tanh(W_{xh}x_t+W_{hh}h_{t-1}+b)

假设一个客服对话在第 1 步说明“订单已经退款”,中间 40 步讨论物流,最后才问钱何时到账。模型需要同时做到:

  1. 长时间保留“已经退款”;
  2. 不把每个物流细节都同等写进有限状态;
  3. 在回答时取出与到账问题有关的信息。

普通 RNN 只有一个整体更新,没有独立的保留、写入和读取开关。其局部梯度还包含:

htht1=diag(1ht2)Whh\frac{\partial h_t}{\partial h_{t-1}} =\operatorname{diag}(1-h_t^2)W_{hh}

每跨一步都要再次乘矩阵和非线性导数。梯度裁剪可以压住爆炸,却不能恢复已经消失的梯度;增大隐状态维度可以增加容量,也不会自动创造稳定的长程路径。

普通 RNN:

h_(t-1) ──┐
           ├─► 仿射变换 ─► tanh ─► h_t
x_t ──────┘                 每一步都整体重写

LSTM:

c_(t-1) ══× 保留量 ══+ 写入量 ══► c_t   长程记忆主干
             ▲            ▲          │
             └──── gates(x_t,h_(t-1))┘
                                      × 输出门 ─► h_t
text

粗线 c 是 LSTM 新增的单元状态(Cell State);它不是永不改变的存储器,而是一个由乘法门控制、用加法更新的状态通道。

02 一个 LSTM 单元究竟计算哪些信号?#

令当前输入 xtRDx_t\in\mathbb{R}^{D},上一隐状态 ht1RHh_{t-1}\in\mathbb{R}^{H},上一单元状态 ct1RHc_{t-1}\in\mathbb{R}^{H}。PyTorch 2.13 当前采用以下方程:

it=σ(Wiixt+bii+Whiht1+bhi)i_t=\sigma(W_{ii}x_t+b_{ii}+W_{hi}h_{t-1}+b_{hi}) ft=σ(Wifxt+bif+Whfht1+bhf)f_t=\sigma(W_{if}x_t+b_{if}+W_{hf}h_{t-1}+b_{hf}) gt=tanh(Wigxt+big+Whght1+bhg)g_t=\tanh(W_{ig}x_t+b_{ig}+W_{hg}h_{t-1}+b_{hg}) ot=σ(Wioxt+bio+Whoht1+bho)o_t=\sigma(W_{io}x_t+b_{io}+W_{ho}h_{t-1}+b_{ho}) ct=ftct1+itgt,ht=ottanh(ct)c_t=f_t\odot c_{t-1}+i_t\odot g_t, \qquad h_t=o_t\odot\tanh(c_t)

每个变量的职责和形状如下:

信号形状数值范围作用
iti_t[N,H](0,1)(0,1)输入门(Input Gate),控制候选内容写入多少
ftf_t[N,H](0,1)(0,1)遗忘门(Forget Gate),控制旧状态保留多少
gtg_t[N,H](1,1)(-1,1)候选记忆(Candidate Memory),提供待写入内容
oto_t[N,H](0,1)(0,1)输出门(Output Gate),控制当前暴露多少
ctc_t[N,H]不固定单元状态,沿时间保存与累积信息
hth_t[N,H](1,1)(-1,1)隐状态,传给下一步并对外提供当前表示

其中 \odot 是逐元素乘法(Hadamard Product)。门和状态都是向量,不是整个单元只有一个开关:第 7 个状态维度可以选择保留,第 19 个维度可以同时覆写。

03 加法记忆路径怎样改变数据流?#

把一次更新拆开看,LSTM 先计算“保留项”和“写入项”,然后相加:

                         ┌───────────────┐
c_(t-1) [N,H] ──────────× f_t [N,H]─────┤
                         │               │
                         │               + ──► c_t [N,H]
x_t [N,D] ──┐            │               │          │
             ├─► gates ──┼─ i_t × g_t ──┘          tanh
h_(t-1)[N,H]┘            │                          │
                         └──────── o_t ─────────────× ──► h_t [N,H]
text

这条图表达了三个不同问题:

  • f_t × c_(t-1):过去的哪些维度继续留下?
  • i_t × g_t:当前产生了什么候选内容,其中多少应该写入?
  • o_t × tanh(c_t):已保存的信息中,当前需要对外暴露哪些?

关键不是“用了更多激活函数”,而是 ctc_t 中出现了显式加法。旧状态可以沿第一项直接到达新状态,不必每一步都被完整压进一次新的 tanh

04 用一个标量手算三步保留与改写#

先不计算门的仿射层,直接给出它们的输出,以隔离记忆更新。令 H=1H=1c0=0c_0=0

时刻ftf_titi_tgtg_toto_t解释
10.900.801.000.70写入一个正向事实
20.900.100.000.70几乎不写入,只继续保留
30.200.70-1.000.70大量遗忘并写入反向事实

第一步:

c1=0.90×0+0.80×1=0.80c_1=0.90\times0+0.80\times1=0.80 h1=0.70tanh(0.80)0.4648h_1=0.70\tanh(0.80)\approx0.4648

第二步没有有用新内容:

c2=0.90×0.80+0.10×0=0.72c_2=0.90\times0.80+0.10\times0=0.72 h2=0.70tanh(0.72)0.4318h_2=0.70\tanh(0.72)\approx0.4318

第三步出现冲突证据:

c3=0.20×0.72+0.70×(1)=0.556c_3=0.20\times0.72+0.70\times(-1)=-0.556 h3=0.70tanh(0.556)0.3535h_3=0.70\tanh(-0.556)\approx-0.3535

第二步把旧内容从 0.800.80 平滑保留到 0.720.72;第三步先把旧内容缩到 0.1440.144,再写入 0.70-0.70。这就是“门控加法”的可计算含义。

注意 hth_tctc_t 不相等。ctc_t 是内部记忆主干;hth_t 经过 tanh 和输出门,是当前提供给上层、读出头以及下一时间步门控网络的工作表示。

05 梯度为什么能沿单元状态走得更远?#

若暂时只看 ct1ctc_{t-1}\rightarrow c_t 的直接路径,把门值视为当前前向已确定的系数,则:

ctct1direct=ft\left.\frac{\partial c_t}{\partial c_{t-1}}\right|_{\text{direct}} =f_t

跨越多步的直接梯度路径是:

cTckdirect=t=k+1Tft\left.\frac{\partial c_T}{\partial c_k}\right|_{\text{direct}} =\prod_{t=k+1}^{T}f_t
loss ─► c_T ──× f_T──► c_(T-1) ──× f_(T-1)──► ... ──× f_(k+1)──► c_k

普通 RNN 长链:每步穿过循环矩阵和 tanh 导数
LSTM 直接路径:每步主要由可学习的遗忘门决定保留比例
text

若 50 步的遗忘门都约为 0.950.95,直接路径还剩:

0.95500.07690.95^{50}\approx0.0769

而每步局部增益为 0.50.5 的链只剩 0.5508.88×10160.5^{50}\approx8.88\times10^{-16}。LSTM 因而能学习把某些 ftf_t 推近 1,让对应状态维度有一条较稳定的梯度通路。

但这不是“梯度永不消失”的证明。门本身依赖 xtx_tht1h_{t-1},完整导数还包含其他路径;若 ftf_t 长期远小于 1,乘积照样衰减;若 sigmoid 饱和,门控参数也会收到很弱的梯度。LSTM 改善了优化几何,没有消除所有长程学习困难。

06 放回 batch 后,张量和参数是什么形状?#

本文采用 batch_first=True、单层单向、无投影的基本设置:

名称形状含义
x[N,T,D]NN 条序列、TT 步、每步 DD
h0, c0[1,N,H]初始隐状态与初始单元状态
output[N,T,H]最后一层在每个时刻的 hth_t
h_n,c_n[1,N,H]最终隐状态与最终单元状态
logits[N,C]序列级任务的 CC 类未归一化分数

四组门通常合并为一次输入仿射和一次循环仿射:

[ai;af;ag;ao]=Wihxt+bih+Whhht1+bhh[a_i;a_f;a_g;a_o] =W_{ih}x_t+b_{ih}+W_{hh}h_{t-1}+b_{hh}

其中:

WihR4H×D,WhhR4H×H,bih,bhhR4HW_{ih}\in\mathbb{R}^{4H\times D},\quad W_{hh}\in\mathbb{R}^{4H\times H},\quad b_{ih},b_{hh}\in\mathbb{R}^{4H}

因此一层单向 LSTM 的参数量为:

4HD+4H2+8H=4H(D+H+2)4HD+4H^2+8H=4H(D+H+2)

同尺寸普通 RNN 只有 HD+H2+2HHD+H^2+2H 个参数。LSTM 以约四倍的循环层参数和更多中间激活,换取可学习的记忆控制。

07 不调用 LSTM 封装,先写出循环本体#

下面实现单层、单向 LSTM。chunk(4, dim=-1) 的顺序必须是 PyTorch 官方约定的 i,f,g,o

显式返回门值只用于教学和诊断。生产训练若不需要门统计,应使用框架融合实现,避免 Python 时间循环和额外激活保存拖慢吞吐。

08 与 PyTorch 2.13 官方实现逐张量对齐#

PyTorch 2.13 当前的 torch.nn.LSTM 接口为 input_sizehidden_sizenum_layersbiasbatch_firstdropoutbidirectionalproj_size 等。下面把手写参数复制给官方层,比较每个时刻和两个最终状态:

官方 API 还有五个必须明确的契约:

  • batch_first=True 只改变 inputoutputh_0c_0h_nc_n 仍以层/方向维开头。
  • output 是最后一层每个时刻的 hth_t;它不包含全部层,也不返回 ctc_t 序列。
  • dropout>0 只放在相邻 LSTM 层之间,最后一层后不放;num_layers=1 时不会得到循环时间步 dropout。
  • bidirectional=True 令方向数 R=2R=2output 最后一维变成 2H2H;它使用未来信息,不适用于严格在线预测。
  • proj_size>0 会让隐状态/输出宽度变成投影宽度,但单元状态仍保持 hidden_size;此时 h_nc_n 最后一维不同。

09 变长序列如何完成一次真实训练?#

一个 batch 的文本长度可能是 [7,4,2]。若补齐到 T_max=7 后直接取 output[:, -1],后两个样本读到的是 padding 后位置。打包序列(Packed Sequence)让 LSTM 跳过无效步,并让 h_n 对应每条序列的真实末尾。

tokens [N,T_max] ─► Embedding ─► x [N,T_max,D]
       lengths [N] ─────────────► pack_padded_sequence


                                  nn.LSTM

                               h_n [L,N,H]
                                      │ 取最后一层

                                  Linear(H,C)

                                  logits [N,C]

                           CrossEntropyLoss(logits,y[N])
text

PyTorch 2.13 当前的 pack_padded_sequencebatch_first=True 时接收 [N,T,*]。若 lengths 是张量,它必须位于 CPU;enforce_sorted=False 允许输入 batch 未按长度降序排列。

padding_idx=0 使 padding 词向量不被更新,但它本身不会让 LSTM 跳过 padding;真正跳过无效时间步的是打包。cross_entropy 接收 logits 和 int64 类别索引,不能先手动 softmax。

10 训练和流式推理的状态边界#

训练独立样本时,通常让每个 batch 从零状态开始;连续传感器流则可能把 (h_n,c_n) 传给下一个 chunk:

state = None
model.eval()
with torch.inference_mode():
    for x_chunk in stream:  # each [N,K,D],同一批连续会话
        output, state = model.lstm(x_chunk, state)
        h_n, c_n = state
        consume(output)
python

截断时间反向传播(Truncated BPTT)中既要传状态数值,又要在 chunk 边界切断旧计算图:

h, c = h_n.detach(), c_n.detach()
state = (h, c)
python

三个边界不能混用:

  • 独立样本之间复用状态,会造成跨用户或跨序列信息泄漏;
  • 连续流每个 chunk 清零状态,会把有效上下文硬性限制为 chunk 长度;
  • 训练连续流长期不 detach(),计算图和内存会随时间增长。

生产流式系统还要定义会话结束、超时、乱序、设备迁移和 batch 内某条流提前结束时怎样重置两份状态。LSTM 有 (h,c) 两个状态,漏重置任何一个都会留下历史。

11 怎样证明门真的在完成任务?#

只看验证损失下降,无法确认模型是否学会长程保留。可以建立一个可证伪的“延迟复制”任务:序列第 1 步给出比特,随后填充噪声,最后一步要求复原该比特。

输入:   bit  noise  noise  ...  query
标签:                              bit
距离:    <────────── Δ ───────────>
text

一条可执行的诊断路径是:

  1. 先用 Δ=3\Delta=3 过拟合 32 条样本,排除损失、标签与形状错误。
  2. Δ\Delta 逐步增加到 10、30、100,画准确率而不是只看一次终值。
  3. 用教学版 TransparentLSTM 记录 forget/input/output[N,T,H] 分布。
  4. 对第一个输入调用 retain_grad(),记录最终损失对早期输入的梯度范数。
  5. 同时记录 clip_grad_norm_ 返回的裁剪前总范数和实际发生裁剪的比例。
  6. 将序列中段打乱或将第一步置零,检查预测是否按任务预期改变。

理想现象不是所有遗忘门都接近 1。若所有维度永远保留,旧内容会持续累积并挤占容量;模型应在需要跨越噪声时保留,在证据失效或被修正时遗忘。

12 最常见的“能运行,但记忆语义错了”#

  • 交换门顺序。 PyTorch 参数拼接顺序是 i,f,g,o;若手写代码按其他教材的排法切片,形状完全相同但结果错误。
  • h_nc_n 当成同一个状态。 二者形状通常相同、职责不同,流式传递和重置必须成对进行。
  • 认为 batch_first 改变状态布局。 它只改变输入和输出;状态仍是 [L·R,N,*]
  • 变长 batch 使用 output[:, -1] 短序列读到 padding 后位置;应使用打包后的 h_n 或按真实长度索引。
  • 把 GPU lengths 直接送进打包函数。 当前官方契约要求张量形式的 lengths 位于 CPU。
  • 在单层 LSTM 上设置 dropout 就以为完成正则化。 内置 dropout 只作用于相邻循环层之间。
  • 序列分类前先 softmax。 cross_entropy 要求 logits;提前 softmax 会改变梯度并降低数值稳定性。
  • 双向模型用于在线预测。 反向分支需要未来输入,离线指标无法直接转化为实时能力。
  • 只记录裁剪后梯度。 每步都爆炸再被压平会看似稳定;必须记录裁剪前范数。
  • 把遗忘门偏置切错位置。 两个偏置向量都按四门拼接,修改前要断言切片并做前向对齐测试。
  • 把门热力图当作因果解释。 高门值只说明该坐标的数值通路强,不能单独证明某个词导致答案。

13 LSTM 与 GRU 的边界在哪里?#

门控循环单元(Gated Recurrent Unit, GRU)把记忆接口进一步压缩:没有独立的 ctc_t,而用更新门在旧隐状态和候选状态之间插值,并用重置门控制候选计算读取多少过去。

方法状态接口主要门控一层循环参数量直接取舍
vanilla RNNhth_tH(D+H+2)H(D+H+2)最简单,但长程梯度路径脆弱
LSTM(ht,ct)(h_t,c_t)输入、遗忘、输出门4H(D+H+2)4H(D+H+2)控制更细,参数与状态更多
GRUhth_t更新、重置门3H(D+H+2)3H(D+H+2)接口更紧凑,少一份单元状态

不能仅凭“GRU 参数少”或“LSTM 门更多”预先宣布胜者。公平比较至少要控制数据切分、参数规模、训练预算、序列长度和延迟,并同时报告效果、吞吐、显存与部署状态大小。

二者仍然逐步递推,都不能在时间维像卷积或 Transformer 那样完全并行。LSTM 解决的是普通 RNN 的记忆更新与梯度路径问题,不是所有序列计算问题。

14 LSTM 会在哪些场景失败?#

  • 极长而精确的检索。 门值的长乘积仍会衰减,有限维状态也可能被后续事件覆盖。
  • 需要同时保留大量细节。 所有历史必须压入固定宽度 (h,c)(h,c);长文档中的多个实体会竞争容量。
  • 训练吞吐受时间依赖限制。 同一层的第 tt 步依赖第 t1t-1 步,长序列难以沿时间并行。
  • 门饱和。 sigmoid 接近 0 或 1 时导数很小,模型可能陷入“几乎总忘”或“几乎总留”的策略。
  • 不规则采样。 普通 LSTM 默认相邻步时间间隔等价,医疗事件流等任务需显式加入时间差或改用连续时间模型。
  • 需要指出证据位置。 最终状态不给出可审计的来源位置,需要注意力、检索或归因机制。
  • 状态管理不可靠。 在线服务中的漏重置、乱序和跨请求复用,会把模型问题放大成数据隔离问题。

增加层数、隐藏宽度或梯度裁剪阈值都不能自动解决这些限制。先把失败归因到容量、优化、计算还是状态边界,再决定结构改造。

15 今天真正需要记住什么?#

  1. LSTM 将对外隐状态 hth_t 与长程单元状态 ctc_t 分开,让保留、写入和输出成为三个可学习决定。
  2. 核心更新 ct=ftct1+itgtc_t=f_t\odot c_{t-1}+i_t\odot g_t 是受门控制的加法路径;它比普通 RNN 每步整体经过非线性更利于长程梯度传播。
  3. 直接梯度路径仍包含 ft\prod f_t,所以 LSTM 是缓解而不是消灭梯度消失;门饱和、容量竞争和顺序计算仍存在。
  4. PyTorch 的门顺序、状态形状、打包长度、层间 dropout 和投影宽度都是必须测试的接口契约。
  5. 调试记忆要使用可控延迟任务、门分布、早期输入梯度和干预实验,不能只看最终损失。

16 思考题与小练习#

  1. 延续本文的标量例子,把第三步改为 f3=0.95,i3=0.05,g3=1,o3=0.7f_3=0.95,i_3=0.05,g_3=-1,o_3=0.7,手算 c3,h3c_3,h_3。比较原设置,解释“新证据出现”不等于模型一定会覆写旧记忆。
  2. TransparentLSTM 写测试:将参数复制到 nn.LSTM,分别验证零初始状态、自定义 (h0,c0) 和两个 batch size 下的全部输出。然后故意交换 fg 的切片,观察哪项断言最先失败。
  3. 构造延迟复制数据集,让距离 Δ{5,20,80}\Delta\in\{5,20,80\}。在相近参数量下比较 vanilla RNN、LSTM 和 GRU 的准确率、裁剪前梯度范数、每秒样本数与状态大小,并说明仅比较最终准确率会遗漏什么。

相关工作#

17 下一篇预告#

LSTM 能把较长历史压进最终状态,但当整段输入必须塞进一个定长向量时,信息瓶颈仍然存在。下一篇将进入编码器—解码器(Encoder–Decoder)与注意力机制(Attention Mechanism),追踪解码器如何在每一步直接选择不同的源位置,而不是只依赖最后一个状态。

旧信息何时该忘、何时该写入?LSTM 的门控与加法记忆路径
https://zwjcode.cn/blog/lstm-gates-additive-memory-path
作者
发布于 2026年9月5日
版权协议 CC BY-NC-SA 4.0
评论加载似乎遇到了问题,请尝试刷新页面。