旧信息何时该忘、何时该写入?LSTM 的门控与加法记忆路径
从普通 RNN 的梯度连乘出发,手算 LSTM 的遗忘、写入与输出门,拆解加法记忆路径,并用 PyTorch 2.13 对齐实现与变长序列训练。
上一篇把普通循环神经网络(Recurrent Neural Network, RNN)沿时间展开后,我们看见了长程学习的症结:从损失回到很早的状态,梯度必须反复穿过同一个循环矩阵和 tanh 导数。即使前向状态仍非零,远处的训练信号也可能已经小到无法使用。
长短期记忆网络(Long Short-Term Memory, LSTM)没有取消时间递推。它做的关键改造是:把“对外工作的隐状态”和“沿时间保存的单元状态”分开,再用可学习的门决定旧信息保留多少、新信息写入多少、当前暴露多少。
本文只讲透一个核心问题:LSTM 如何把普通 RNN 的“每步整体重写”改成受控的加法记忆路径。我们会依次拆解四个信号、手算三步状态、追踪梯度,再与 PyTorch 2.13 的官方实现逐张量对齐。
01 普通 RNN 为什么很难选择性地记忆?#
普通 tanh RNN 将旧状态和新输入混在一次变换里:
假设一个客服对话在第 1 步说明“订单已经退款”,中间 40 步讨论物流,最后才问钱何时到账。模型需要同时做到:
- 长时间保留“已经退款”;
- 不把每个物流细节都同等写进有限状态;
- 在回答时取出与到账问题有关的信息。
普通 RNN 只有一个整体更新,没有独立的保留、写入和读取开关。其局部梯度还包含:
每跨一步都要再次乘矩阵和非线性导数。梯度裁剪可以压住爆炸,却不能恢复已经消失的梯度;增大隐状态维度可以增加容量,也不会自动创造稳定的长程路径。
普通 RNN:
h_(t-1) ──┐
├─► 仿射变换 ─► tanh ─► h_t
x_t ──────┘ 每一步都整体重写
LSTM:
c_(t-1) ══× 保留量 ══+ 写入量 ══► c_t 长程记忆主干
▲ ▲ │
└──── gates(x_t,h_(t-1))┘
× 输出门 ─► h_ttext粗线 c 是 LSTM 新增的单元状态(Cell State);它不是永不改变的存储器,而是一个由乘法门控制、用加法更新的状态通道。
02 一个 LSTM 单元究竟计算哪些信号?#
令当前输入 ,上一隐状态 ,上一单元状态 。PyTorch 2.13 当前采用以下方程:
每个变量的职责和形状如下:
| 信号 | 形状 | 数值范围 | 作用 |
|---|---|---|---|
[N,H] | 输入门(Input Gate),控制候选内容写入多少 | ||
[N,H] | 遗忘门(Forget Gate),控制旧状态保留多少 | ||
[N,H] | 候选记忆(Candidate Memory),提供待写入内容 | ||
[N,H] | 输出门(Output Gate),控制当前暴露多少 | ||
[N,H] | 不固定 | 单元状态,沿时间保存与累积信息 | |
[N,H] | 隐状态,传给下一步并对外提供当前表示 |
其中 是逐元素乘法(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):已保存的信息中,当前需要对外暴露哪些?
关键不是“用了更多激活函数”,而是 中出现了显式加法。旧状态可以沿第一项直接到达新状态,不必每一步都被完整压进一次新的 tanh。
04 用一个标量手算三步保留与改写#
先不计算门的仿射层,直接给出它们的输出,以隔离记忆更新。令 、:
| 时刻 | 解释 | ||||
|---|---|---|---|---|---|
| 1 | 0.90 | 0.80 | 1.00 | 0.70 | 写入一个正向事实 |
| 2 | 0.90 | 0.10 | 0.00 | 0.70 | 几乎不写入,只继续保留 |
| 3 | 0.20 | 0.70 | -1.00 | 0.70 | 大量遗忘并写入反向事实 |
第一步:
第二步没有有用新内容:
第三步出现冲突证据:
第二步把旧内容从 平滑保留到 ;第三步先把旧内容缩到 ,再写入 。这就是“门控加法”的可计算含义。
注意 和 不相等。 是内部记忆主干; 经过 tanh 和输出门,是当前提供给上层、读出头以及下一时间步门控网络的工作表示。
05 梯度为什么能沿单元状态走得更远?#
若暂时只看 的直接路径,把门值视为当前前向已确定的系数,则:
跨越多步的直接梯度路径是:
loss ─► c_T ──× f_T──► c_(T-1) ──× f_(T-1)──► ... ──× f_(k+1)──► c_k
普通 RNN 长链:每步穿过循环矩阵和 tanh 导数
LSTM 直接路径:每步主要由可学习的遗忘门决定保留比例text若 50 步的遗忘门都约为 ,直接路径还剩:
而每步局部增益为 的链只剩 。LSTM 因而能学习把某些 推近 1,让对应状态维度有一条较稳定的梯度通路。
但这不是“梯度永不消失”的证明。门本身依赖 和 ,完整导数还包含其他路径;若 长期远小于 1,乘积照样衰减;若 sigmoid 饱和,门控参数也会收到很弱的梯度。LSTM 改善了优化几何,没有消除所有长程学习困难。
06 放回 batch 后,张量和参数是什么形状?#
本文采用 batch_first=True、单层单向、无投影的基本设置:
| 名称 | 形状 | 含义 |
|---|---|---|
x | [N,T,D] | 条序列、 步、每步 维 |
h0, c0 | [1,N,H] | 初始隐状态与初始单元状态 |
output | [N,T,H] | 最后一层在每个时刻的 |
h_n,c_n | [1,N,H] | 最终隐状态与最终单元状态 |
logits | [N,C] | 序列级任务的 类未归一化分数 |
四组门通常合并为一次输入仿射和一次循环仿射:
其中:
因此一层单向 LSTM 的参数量为:
同尺寸普通 RNN 只有 个参数。LSTM 以约四倍的循环层参数和更多中间激活,换取可学习的记忆控制。
07 不调用 LSTM 封装,先写出循环本体#
下面实现单层、单向 LSTM。chunk(4, dim=-1) 的顺序必须是 PyTorch 官方约定的 i,f,g,o。
import torch
from torch import nn
class TransparentLSTM(nn.Module):
def __init__(self, input_size: int, hidden_size: int) -> None:
super().__init__()
self.input_size = input_size
self.hidden_size = hidden_size
self.weight_ih = nn.Parameter(torch.empty(4 * hidden_size, input_size))
self.weight_hh = nn.Parameter(torch.empty(4 * hidden_size, hidden_size))
self.bias_ih = nn.Parameter(torch.zeros(4 * hidden_size))
self.bias_hh = nn.Parameter(torch.zeros(4 * hidden_size))
nn.init.xavier_uniform_(self.weight_ih)
nn.init.orthogonal_(self.weight_hh)
def forward(
self,
x: torch.Tensor, # [N,T,D]
state: tuple[torch.Tensor, torch.Tensor] | None = None,
) -> tuple[
torch.Tensor,
tuple[torch.Tensor, torch.Tensor],
dict[str, torch.Tensor],
]:
assert x.ndim == 3 and x.shape[-1] == self.input_size
n, time_steps, _ = x.shape
if state is None:
h = x.new_zeros(n, self.hidden_size) # [N,H]
c = x.new_zeros(n, self.hidden_size) # [N,H]
else:
h0, c0 = state
assert h0.shape == c0.shape == (1, n, self.hidden_size)
h, c = h0[0], c0[0]
outputs = []
gate_history = {name: [] for name in ("input", "forget", "candidate", "output")}
for t in range(time_steps):
affine = (
x[:, t] @ self.weight_ih.T
+ self.bias_ih
+ h @ self.weight_hh.T
+ self.bias_hh
) # [N,4H]
a_i, a_f, a_g, a_o = affine.chunk(4, dim=-1)
i = torch.sigmoid(a_i)
f = torch.sigmoid(a_f)
g = torch.tanh(a_g)
o = torch.sigmoid(a_o)
c = f * c + i * g
h = o * torch.tanh(c)
outputs.append(h)
for name, value in zip(gate_history, (i, f, g, o), strict=True):
gate_history[name].append(value)
output = torch.stack(outputs, dim=1) # [N,T,H]
gates = {
name: torch.stack(values, dim=1) # each [N,T,H]
for name, values in gate_history.items()
}
return output, (h.unsqueeze(0), c.unsqueeze(0)), gatespython显式返回门值只用于教学和诊断。生产训练若不需要门统计,应使用框架融合实现,避免 Python 时间循环和额外激活保存拖慢吞吐。
08 与 PyTorch 2.13 官方实现逐张量对齐#
PyTorch 2.13 当前的 torch.nn.LSTM ↗ 接口为 input_size、hidden_size、num_layers、bias、batch_first、dropout、bidirectional 和 proj_size 等。下面把手写参数复制给官方层,比较每个时刻和两个最终状态:
import torch
from torch import nn
torch.manual_seed(11)
x = torch.randn(2, 5, 3) # [N=2,T=5,D=3]
manual = TransparentLSTM(input_size=3, hidden_size=4)
official = nn.LSTM(
input_size=3,
hidden_size=4,
num_layers=1,
batch_first=True,
bidirectional=False,
)
with torch.no_grad():
official.weight_ih_l0.copy_(manual.weight_ih)
official.weight_hh_l0.copy_(manual.weight_hh)
official.bias_ih_l0.copy_(manual.bias_ih)
official.bias_hh_l0.copy_(manual.bias_hh)
manual_output, (manual_hn, manual_cn), gates = manual(x)
output, (h_n, c_n) = official(x)
assert output.shape == (2, 5, 4)
assert h_n.shape == c_n.shape == (1, 2, 4)
assert gates["forget"].shape == (2, 5, 4)
torch.testing.assert_close(output, manual_output)
torch.testing.assert_close(h_n, manual_hn)
torch.testing.assert_close(c_n, manual_cn)python官方 API 还有五个必须明确的契约:
batch_first=True只改变input和output;h_0、c_0、h_n、c_n仍以层/方向维开头。output是最后一层每个时刻的 ;它不包含全部层,也不返回 序列。dropout>0只放在相邻 LSTM 层之间,最后一层后不放;num_layers=1时不会得到循环时间步 dropout。bidirectional=True令方向数 ,output最后一维变成 ;它使用未来信息,不适用于严格在线预测。proj_size>0会让隐状态/输出宽度变成投影宽度,但单元状态仍保持hidden_size;此时h_n与c_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])textPyTorch 2.13 当前的 pack_padded_sequence ↗ 在 batch_first=True 时接收 [N,T,*]。若 lengths 是张量,它必须位于 CPU;enforce_sorted=False 允许输入 batch 未按长度降序排列。
import torch
from torch import nn
from torch.nn.utils.rnn import pack_padded_sequence
class PackedLSTMClassifier(nn.Module):
def __init__(self, vocab_size: int, embed_dim: int, hidden_size: int, classes: int) -> None:
super().__init__()
self.embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=0)
self.lstm = nn.LSTM(
input_size=embed_dim,
hidden_size=hidden_size,
num_layers=1,
batch_first=True,
)
self.head = nn.Linear(hidden_size, classes)
def forward(self, tokens: torch.Tensor, lengths: torch.Tensor) -> torch.Tensor:
assert tokens.ndim == 2 and lengths.shape == (tokens.shape[0],)
x = self.embedding(tokens) # [N,T_max,D]
packed = pack_padded_sequence(
x,
lengths.cpu(),
batch_first=True,
enforce_sorted=False,
)
_, (h_n, c_n) = self.lstm(packed)
assert h_n.shape == c_n.shape
return self.head(h_n[-1]) # [N,C],单向模型最后一层
model = PackedLSTMClassifier(vocab_size=5000, embed_dim=64, hidden_size=96, classes=3)
optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4)
tokens = torch.tensor([
[8, 4, 9, 2, 7, 3, 6],
[5, 1, 2, 4, 0, 0, 0],
[7, 3, 0, 0, 0, 0, 0],
]) # [N=3,T_max=7]
lengths = torch.tensor([7, 4, 2]) # [N],CPU 整数张量
targets = torch.tensor([2, 0, 1]) # [N],类别索引
model.train()
optimizer.zero_grad(set_to_none=True)
logits = model(tokens, lengths) # [3,3],未经 softmax
loss = nn.functional.cross_entropy(logits, targets)
loss.backward()
grad_norm = nn.utils.clip_grad_norm_(
model.parameters(), max_norm=1.0, error_if_nonfinite=True
)
optimizer.step()
model.eval()
with torch.inference_mode():
probabilities = model(tokens, lengths).softmax(dim=-1) # [3,3]
predictions = probabilities.argmax(dim=-1) # [3]pythonpadding_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一条可执行的诊断路径是:
- 先用 过拟合 32 条样本,排除损失、标签与形状错误。
- 将 逐步增加到 10、30、100,画准确率而不是只看一次终值。
- 用教学版
TransparentLSTM记录forget/input/output的[N,T,H]分布。 - 对第一个输入调用
retain_grad(),记录最终损失对早期输入的梯度范数。 - 同时记录
clip_grad_norm_返回的裁剪前总范数和实际发生裁剪的比例。 - 将序列中段打乱或将第一步置零,检查预测是否按任务预期改变。
理想现象不是所有遗忘门都接近 1。若所有维度永远保留,旧内容会持续累积并挤占容量;模型应在需要跨越噪声时保留,在证据失效或被修正时遗忘。
12 最常见的“能运行,但记忆语义错了”#
- 交换门顺序。 PyTorch 参数拼接顺序是
i,f,g,o;若手写代码按其他教材的排法切片,形状完全相同但结果错误。 - 把
h_n和c_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)把记忆接口进一步压缩:没有独立的 ,而用更新门在旧隐状态和候选状态之间插值,并用重置门控制候选计算读取多少过去。
| 方法 | 状态接口 | 主要门控 | 一层循环参数量 | 直接取舍 |
|---|---|---|---|---|
| vanilla RNN | 无 | 最简单,但长程梯度路径脆弱 | ||
| LSTM | 输入、遗忘、输出门 | 控制更细,参数与状态更多 | ||
| GRU | 更新、重置门 | 约 | 接口更紧凑,少一份单元状态 |
不能仅凭“GRU 参数少”或“LSTM 门更多”预先宣布胜者。公平比较至少要控制数据切分、参数规模、训练预算、序列长度和延迟,并同时报告效果、吞吐、显存与部署状态大小。
二者仍然逐步递推,都不能在时间维像卷积或 Transformer 那样完全并行。LSTM 解决的是普通 RNN 的记忆更新与梯度路径问题,不是所有序列计算问题。
14 LSTM 会在哪些场景失败?#
- 极长而精确的检索。 门值的长乘积仍会衰减,有限维状态也可能被后续事件覆盖。
- 需要同时保留大量细节。 所有历史必须压入固定宽度 ;长文档中的多个实体会竞争容量。
- 训练吞吐受时间依赖限制。 同一层的第 步依赖第 步,长序列难以沿时间并行。
- 门饱和。 sigmoid 接近 0 或 1 时导数很小,模型可能陷入“几乎总忘”或“几乎总留”的策略。
- 不规则采样。 普通 LSTM 默认相邻步时间间隔等价,医疗事件流等任务需显式加入时间差或改用连续时间模型。
- 需要指出证据位置。 最终状态不给出可审计的来源位置,需要注意力、检索或归因机制。
- 状态管理不可靠。 在线服务中的漏重置、乱序和跨请求复用,会把模型问题放大成数据隔离问题。
增加层数、隐藏宽度或梯度裁剪阈值都不能自动解决这些限制。先把失败归因到容量、优化、计算还是状态边界,再决定结构改造。
15 今天真正需要记住什么?#
- LSTM 将对外隐状态 与长程单元状态 分开,让保留、写入和输出成为三个可学习决定。
- 核心更新 是受门控制的加法路径;它比普通 RNN 每步整体经过非线性更利于长程梯度传播。
- 直接梯度路径仍包含 ,所以 LSTM 是缓解而不是消灭梯度消失;门饱和、容量竞争和顺序计算仍存在。
- PyTorch 的门顺序、状态形状、打包长度、层间 dropout 和投影宽度都是必须测试的接口契约。
- 调试记忆要使用可控延迟任务、门分布、早期输入梯度和干预实验,不能只看最终损失。
16 思考题与小练习#
- 延续本文的标量例子,把第三步改为 ,手算 。比较原设置,解释“新证据出现”不等于模型一定会覆写旧记忆。
- 为
TransparentLSTM写测试:将参数复制到nn.LSTM,分别验证零初始状态、自定义(h0,c0)和两个 batch size 下的全部输出。然后故意交换f与g的切片,观察哪项断言最先失败。 - 构造延迟复制数据集,让距离 。在相近参数量下比较 vanilla RNN、LSTM 和 GRU 的准确率、裁剪前梯度范数、每秒样本数与状态大小,并说明仅比较最终准确率会遗漏什么。
相关工作#
- Hochreiter & Schmidhuber (1997), Long Short-Term Memory ↗:提出 LSTM 的长程误差信号与记忆单元框架。
- Gers, Schmidhuber & Cummins (2000), Learning to Forget: Continual Prediction with LSTM ↗:引入并分析遗忘门,使持续任务能主动释放旧状态。
- Graves, Mohamed & Hinton (2013), Speech Recognition with Deep Recurrent Neural Networks ↗:展示深层双向 LSTM 在语音识别中的代表性应用。
- Cho et al. (2014), Learning Phrase Representations using RNN Encoder–Decoder for Statistical Machine Translation ↗:提出带门控的编码器—解码器单元,即后来常称的 GRU。
- Greff et al. (2017), LSTM: A Search Space Odyssey ↗:系统比较 LSTM 变体,分析常见组件的实际贡献。
17 下一篇预告#
LSTM 能把较长历史压进最终状态,但当整段输入必须塞进一个定长向量时,信息瓶颈仍然存在。下一篇将进入编码器—解码器(Encoder–Decoder)与注意力机制(Attention Mechanism),追踪解码器如何在每一步直接选择不同的源位置,而不是只依赖最后一个状态。