深度学习之长短期记忆网络LSTM
LSTM(Long Short-Term Memory,长短期记忆网络)是一种专门为弥补循环神经网络”记不住久远信息”而设计的网络结构——它给网络内部额外开了一条几乎不衰减的记忆传送带,再配上三个可学习的”闸门”,决定何时写入、何时丢弃、何时读出。自 1997 年被提出以来,LSTM 长期是语音识别、机器翻译、时间序列预测等序列任务的主力模型,深度学习社区甚至流传一句话:几乎所有基于 RNN 的惊人成果,都是用 LSTM 实现的(Karpathy, 2015)。直到 Transformer 出现后它才逐渐让位,但门控思想至今仍深刻影响着整个深度学习领域。

普通RNN存在的问题
循环神经网络(Recurrent Neural Network,RNN)是一种用于处理序列数据的神经网络。相比一般的神经网络来说,他能够处理序列变化的数据。比如某个单词的意思会因为上文提到的内容不同而有不同的含义,RNN就能够很好地解决这类问题。

在上图所示的神经网络A中,输入为$X_t$,输出为$h_t$。A上的环允许将每一步产生的信息传递到下一步中。一个RNN可以看作是同一个网络的多份副本,每一份都将信息传递到下一个副本。将环展开:

在过去的几年里,RNN在一系列的任务中都取得了令人惊叹的成就,比如语音识别,语言建模,翻译,图片标题等等。关于RNN在各个领域所取得的令人惊叹的成就。
有时,我们只需要看最近的信息,就可以完成当前的任务。比如,考虑一个语言模型,通过前面的单词来预测接下来的单词。如果我们想预测句子”the clouds are in the sky”中的最后一个单词,我们不需要更多的上下文信息——很明显下一个单词应该是sky。

然而,有时候我们需要更多的上下文信息。比如,我们想预测句子”I grew up in France… I speak fluent French”中的最后一个单词。最近的信息告诉我们,最后一个单词可能是某种语言的名字,然而如果我们想确定到底是哪种语言的话,我们需要France这个更远的上下文信息。实际上,相关信息和需要该信息的位置之间的距离可能非常的远。不幸的是,随着距离的增大,RNN对于如何将这样的信息连接起来无能为力。

普通 RNN 记不住远的信息,根源在于反向传播时梯度沿时间步指数级衰减——这就是梯度消失问题。先看 RNN 是怎么工作的。循环神经网络的核心是一个不断自我更新的隐藏状态:

$$h_t = \tanh(W \cdot [h_{t-1}, x_t] + b)$$
每一步,网络把上一时刻的隐藏状态 $h_{t-1}$ 与当前输入 $x_t$ 拼接起来,做一次线性变换,再经过 $\tanh$ 激活,得到新的隐藏状态 $h_t$。理论上,只要学出合适的权重,$h_t$ 就能携带任意久远的历史信息——第 1 步的信息通过第 2 步、第 3 步……一路传递到第 100 步。
问题出在训练阶段。训练 RNN 用的是随时间反向传播(BPTT):计算损失对权重的梯度时,需要把误差沿时间轴”倒着传”回去。信息每往回传一步,就要乘一次 $\tanh$ 的导数(最大值也只有 1)和一次权重矩阵 $W$。只要 $W$ 的特征值普遍小于 1,梯度每传一步就缩水一次,传 $k$ 步之后大约按 $c^k$ 衰减($c < 1$ 是每步的平均缩放因子)。下图展示了两种典型缩放因子下梯度的衰减速度:

当序列长度达到上百甚至上千时,距离当前时刻十几步之前的梯度就已衰减到可以忽略,模型彻底”看不到”远处的依赖。
具体表现很常见:语言模型记不住句子开头的关键信息,时间序列模型学不到长周期规律,机器翻译漏译长句的从句。与之对称的还有梯度爆炸($c > 1$ 时梯度暴涨),这类问题处理起来相对简单——用梯度裁剪把梯度的范数截断即可,因此历史上梯度消失才是更难缠的那个。
LSTM网络简介
LSTM,全称为长短期记忆网络(Long Short Term Memory networks),是一种特殊的RNN。LSTM由Hochreiter & Schmidhuber(1997)提出,许多研究者进行了一系列的工作对其改进并使之发扬光大。LSTM在许多问题上效果非常好,现在被广泛使用。
LSTM在设计上明确地避免了长期依赖的问题。所有的循环神经网络都有着重复的神经网络模块形成链的形式。在普通的RNN中,重复模块结构非常简单,例如只有一个tanh层。

LSTM也有这种链状结构,不过其重复模块的结构不同。LSTM的重复模块中有4个神经网络层,并且他们之间的交互非常特别。

现在暂且不必关心细节,稍候我们会一步一步地对LSTM的各个部分进行介绍。开始之前,我们先介绍一下将用到的标记。

在上图中,每条线表示向量的传递,从一个结点的输出传递到另外结点的输入。粉红圆表示向量的元素级操作,比如相加或者相乘。黄色方框表示神经网络的层。线合并表示向量的连接,线分叉表示向量复制。
LSTM网络的实现原理
LSTM的主要思想是采用一个叫做”细胞状态(state)”的通道来贯穿整个时间序列。

细胞状态有点像是传送带,它直接穿过整个链,同时只有一些较小的线性交互。上面承载的信息可以很容易地流过而不改变。
通过精心设计”门”的结构来去除或增加信息到细胞状态的能力。门是一种让信息选择式通过的方法。它们包含一个sigmoid神经网络层和一个逐元乘法操作。

Sigmoid层输出0~1之间的值,每个值表示对应的部分信息是否应该通过。0值表示不允许信息通过,1值表示让所有信息通过。一个LSTM有3个这种门,来保护和控制元胞状态。
遗忘门
“遗忘门”决定之前状态中的信息有多少应该舍弃。它会读取$h_{t-1}$和$x_t$的内容,$\sigma$符号代表Sigmoid函数,它会输出一个0到1之间的值。其中0代表舍弃之前细胞状态$C_{t-1}$中的内容,1代表完全保留之前细胞状态$C_{t-1}$中的内容。0、1之间的值代表部分保留之前细胞状态$C_{t-1}$中的内容。

输入门
“输入门”决定什么样的信息保留在细胞状态$C_t$中,它会读取$h_{t-1}$和$x_t$的内容,$\sigma$符号代表Sigmoid函数,它会输出一个0到1之间的值。和”输入门”配合的还有另外一部分,即下图中计算tanh层的部分,这部分输入也是$h_{t-1}$和$x_t$,不过采用tanh激活函数,将这部分标记为$\tilde{c}^{(t)}$,称作为”候选状态”。

细胞状态更新
由$C_{t-1}$计算得到$C_t$。旧”细胞状态”$C_{t-1}$和”遗忘门”的结果进行计算,决定旧的”细胞状态”保留多少,忘记多少。接着”输入门”$i^{(t)}$和候选状态$\tilde{c}^{(t)}$进行计算,将所得到的结果加入到”细胞状态”中,这表示新的输入信息有多少加入到”细胞状态”中。

输出门
和其他门计算一样,它会读取$h_{t-1}$和$x_t$的内容,然后计算Sigmoid函数,得到“输出门”的值。接着把“细胞状态”通过tanh进行处理(得到一个在-1到1之间的值),并将它和输出门的结果相乘,最终得到确定输出的部分。

以上,就是LSTM的内部结构。通过门控状态来控制传输状态,记住需要长时间记忆的,忘记不重要的信息;而不像普通的RNN那样只能够“呆萌”地仅有一种记忆叠加方式。对很多需要“长期记忆”的任务来说,尤其好用。但也因为引入了很多内容,导致参数变多,也使得训练难度加大了很多。
一个时间步的完整数据流
单独看公式容易晕,把六个公式串成一条流水线就清楚了:一个时间步内,数据从输入到输出其实只走六个动作——拼接输入、一次线性变换切出四份、四路分别激活、在细胞状态线上按比例保留、再按比例写入、最后输出门控制放出。下图把每一步标了出来:

注意第②步的实现细节:f、i、g、o 四份并不是分别做四次矩阵乘法,而是把四组权重拼成一个大矩阵一次算完、再切成四份。这正是代码里 nn.Linear(input_size + hidden_size, 4 * hidden_size) 的来历——同样的运算,实现更简洁、GPU 上效率也更高。
跨多步:三个门如何配合
最后看一个跨多步的真实例子,体会三个门在一条长序列里怎么协作。处理句子 “The cat, which was very hungry, finally ate its food.”,模型要在句尾的 ate 处判断动词该用单数还是复数——主语 cat 在第 1 步出现,谓语 ate 在第 9 步才到,中间隔着 8 个词的修饰从句。普通 RNN 到第 9 步时早已把 cat 忘光;LSTM 的做法是:读到 cat 时,输入门把”主语是单数”这条信息写进细胞状态;随后读那 8 个词时,遗忘门始终接近 1(旧记忆原样保留),输入门压低(从句细节不进入主记忆);等到 ate 出现,输出门把这条记忆放出,模型据此选择单数形式。整个过程里,细胞状态中的这条信息几乎无衰减地存活了 9 步——这正是”长短期记忆”中”长期”二字的含义。
常见变体:GRU 与双向 LSTM
GRU 把 LSTM 的两个门合并成一个更新门,参数更少、训练更快;双向 LSTM 用两个方向的隐藏状态拼接,让每个位置同时看到前后文。LSTM 提出后出现了大量改进版本,其中影响最大的是 GRU 和双向结构。
GRU(门控循环单元)由 Cho 等人在 2014 年提出。它砍掉了独立的细胞状态,只保留两个门:更新门 $z_t$ 同时承担”遗忘多少旧状态”和”写入多少新状态”的职责,重置门 $r_t$ 控制候选状态对旧信息的利用程度:
$$ h_t = z_t \odot \tilde{h}_t + (1 – z_t) \odot h_{t-1}$$

两种结构的取舍可以看这张对比表:
| 对比项 | LSTM | GRU |
| 记忆机制 | 独立的细胞状态 $C_t$(额外一条传送带) | 无独立细胞状态,靠隐藏状态 $h_t$ |
| 门的数量 | 3 个(遗忘、输入、输出) | 2 个(更新、重置) |
| 参数量 | 更多(约 $4 \times (d_{in}+d_h) \times d_h$) | 更少(约 $3 \times (d_{in}+d_h) \times d_h$) |
| 表达能力 | 更灵活,可独立控制遗忘与写入 | 与 LSTM 相当,小数据上常更稳定 |
| 训练速度 | 较慢 | 较快 |
| 典型用途 | 长序列建模、语言模型 | 机器翻译、轻量部署、移动端 |
除此之外,实践中还常见这些变体:
- 双向 LSTM(BiLSTM):正向与反向各跑一个 LSTM,把两个方向的隐藏状态拼接作为输出,让每个位置同时利用前后文,适合句子分类、命名实体识别等任务;
- Peephole LSTM:让三个门直接”偷看”细胞状态 $C_{t-1}$,而不是只依赖 $h_{t-1}$ 与 $x_t$;
- CIFG(耦合遗忘-输入门):令 $i_t = 1 – f_t$,把两个门合并,减少参数量;
- 多层堆叠(Stacked LSTM):把多个 LSTM 层串联,高层建模更抽象的语义,低层处理细粒度模式;
- BN-LSTM:在门计算中加入批归一化,加速收敛。
一个实用的经验是:数据量不大时 GRU 通常够用且更稳,追求极致的长序列建模能力时 LSTM 更可靠;而大多数”句子级”任务现在都可以先试双向结构。
什么任务适合 LSTM
判断一个任务适不适合 LSTM,关键看数据是否”顺序敏感、长度可变、且关键信息可能出现在很远的过去”——三者占得越多,LSTM 越值得优先尝试。有了这条主线,我们就能把”为什么适用”从”因为它能处理序列”这种空泛说法,变成可操作的判定标准。
四个判定特征
下面这张表把 LSTM 擅长的任务特征和原因一一对应。你可以把它当成选型清单,逐条对照自己的业务问题:
| 特征 | 含义 | 例子 | LSTM 为何擅长 |
| 顺序敏感 | 数据排列顺序承载含义,交换顺序会改变语义 | “我打你”与”你打我”;股价先涨后跌 vs 先跌后涨 | 循环结构天然按顺序逐个处理,无需位置编码就能感知先后 |
| 变长输入 | 样本长度不一,难以用固定维度向量直接输入 | 一句话有长有短;一段音频帧数不同 | 循环步数可随输入长度自适应,无需预先裁剪或填充成固定长度 |
| 长程依赖 | 当前输出依赖几十步甚至几百步之前的信息 | 句首主语到句尾谓语的单复数;月度销量周期 | 细胞状态把梯度消失从结构上化解,信息可跨长距离稳定传递 |
| 流式输出 | 需要边读边出结果,无法等整个序列读完 | 实时语音识别、逐词翻译、在线异常检测 | 每个时刻只依赖当前输入与历史状态,天然支持在线解码 |
典型应用场景
基于上面四个特征,下表列出了 LSTM 最常见的落地领域。每一行都给出”数据形态、LSTM 解决什么问题、典型产品形态”,方便你直接映射到自己的项目:
| 领域 | 数据形态 | LSTM 解决什么 | 典型例子 |
| 语音识别 / 语音合成 | 音频帧序列 | 帧间上下文、变长序列、流式解码 | 实时语音转写、语音助手 |
| 机器翻译 / 文本生成 | 词或子词序列 | 句法长程依赖、逐词生成目标序列 | 神经机器翻译(seq2seq)、作诗/写代码模型 |
| 命名实体识别 / 词性标注 | 词序列 | 利用双向上下文做序列标注 | BiLSTM-CRF 命名实体识别 |
| 时间序列预测 / 异常检测 | 等间隔数值序列 | 周期、趋势等长程模式,减少人工特征工程 | 销量预测、设备故障预警、指标异常检测 |
| 用户行为序列 / 推荐 | 点击、购买、浏览序列 | 捕捉用户顺序偏好与会话演化 | 会话推荐、搜索路径预测 |
| 视频 / 动作识别 | 视频帧序列 | 建模时序动态,而非单帧外观 | 行为识别、手势识别 |
| 医疗生理信号 | ECG、EEG、血氧等长序列 | 长程节律异常检测 | 心律失常检测、睡眠分期 |
与 Transformer 的分工
很多人把 LSTM 和 Transformer 当成”旧”与”新”的对立,其实二者是互补关系。下面的对比表从计算、数据、部署三个维度说明它们的边界:
| 对比项 | LSTM / GRU | Transformer |
| 顺序建模方式 | 循环结构天然按时间步推进 | 需显式加入位置编码才有顺序概念 |
| 训练并行性 | 按时间步串行,大规模训练难并行 | 自注意力可在序列维度上整体并行 |
| 计算复杂度 | O(序列长度) | 自注意力 O(序列长度²) |
| 长程依赖 | 理论上可达数百步,依赖门控学习 | 全局直接关联,但长序列成本平方增长 |
| 数据量偏好 | 中小数据更稳健,不易过拟合 | 大数据下碾压式优势 |
| 流式 / 低延迟 | 天然支持在线解码,适合边缘设备 | 需要 KV cache、滑动窗口等特殊处理 |
| 典型代表任务 | 语音识别、时间序列、行为序列、低延迟场景 | 大语言模型、长文档理解、海量数据翻译 |
LSTM 并没有过时——在顺序先验强、数据规模中等、需要流式解码或低延迟部署的场景里,它依然是务实且可靠的选择;当数据量巨大、序列很长且需要全局上下文时,Transformer 是更优解。理解 LSTM 的门控与信息高速公路思想,也是理解 Transformer 的一把钥匙。
用 PyTorch 从零实现LSTM
手写一个 LSTM 单元只需约 15 行核心代码,关键是”一次线性变换 + 切分成四份”。我们把四个门/候选的线性变换合并成一个矩阵乘法,一次算完再切分,这样实现简洁、运算也高效。
先实现单元本身。输入是当前时刻的 $x_t$ 和上一时刻的状态 $(h_{t-1}, C_{t-1})$,输出新的 $(h_t, C_t)$:
import torch
import torch.nn as nn
class LSTMCell(nn.Module):
"""一个 LSTM 单元:接收 x_t 与状态 (h, c),输出新状态"""
def __init__(self, input_size, hidden_size):
super().__init__()
self.hidden_size = hidden_size
# 一次线性变换同时算出 f、i、g、o 四个分量(拼接后切分)
self.fc = nn.Linear(input_size + hidden_size, 4 * hidden_size)
def forward(self, x, state):
h_prev, c_prev = state
gates = self.fc(torch.cat([x, h_prev], dim=-1))
f, i, g, o = gates.chunk(4, dim=-1) # 按 hidden_size 切成四份
f = torch.sigmoid(f) # 遗忘门
i = torch.sigmoid(i) # 输入门
g = torch.tanh(g) # 候选记忆
o = torch.sigmoid(o) # 输出门
c = f * c_prev + i * g # 更新细胞状态
h = o * torch.tanh(c) # 计算隐藏状态
return h, (h, c)
接下来用一个简单任务验证它真的能学习:给定正弦序列 $\sin(t/4)$ 的前 64 个点,预测下一个点。这是典型的”一步时序预测”,网络必须学会利用隐藏状态记住相位信息:
import math
# 生成数据:64 个输入点,预测第 2~65 个点(逐点滚动)
seq_len = 64
steps = torch.arange(seq_len + 1).float()
xs = torch.sin(steps / 4.0)
inputs = xs[:-1].unsqueeze(0).unsqueeze(-1) # (1, 64, 1)
targets = xs[1:].unsqueeze(0) # (1, 64)
cell = LSTMCell(input_size=1, hidden_size=16)
optimizer = torch.optim.Adam(cell.parameters(), lr=1e-2)
loss_fn = nn.MSELoss()
for epoch in range(300):
h = torch.zeros(1, 16)
c = torch.zeros(1, 16)
outputs = []
for t in range(seq_len): # 逐步滚动,前一步输出作为下一步输入
h, (h, c) = cell(inputs[:, t], (h, c))
outputs.append(h)
out = torch.stack(outputs, dim=1).squeeze(-1)
loss = loss_fn(out, targets)
optimizer.zero_grad()
loss.backward()
optimizer.step()
if epoch % 60 == 0:
print(f"epoch {epoch:3d} loss = {loss.item():.4f}")
训练 300 轮后 loss 会降到 10⁻³ 量级,说明这个手写单元真的学到了正弦的周期规律。实际项目中通常直接用 PyTorch 的高层封装 nn.LSTM,它自带多层堆叠、双向、dropout 等能力:
lstm = nn.LSTM(input_size=1, hidden_size=16, num_layers=2,
bidirectional=True, batch_first=True)
out, (h_n, c_n) = lstm(inputs) # inputs: (batch, seq, input_size)
# out: (batch, seq, hidden_size * num_directions)
最后分享一个非常实用的工程技巧:把遗忘门偏置初始化为正值(如 1.0)。这样训练初期遗忘门输出接近 1,网络倾向”先记住”,再逐步学出该遗忘什么——大量实践表明这能让 LSTM 训练更快、更稳:
# 四份参数的顺序是 f、i、g、o,把遗忘门那一段偏置设为 1
for name, p in lstm.named_parameters():
if "bias_ih" in name:
n = p.size(0) // 4
p.data[n:2 * n].fill_(1.0)
使用LSTM对IMDB评论进行情感分析
from keras.datasets import imdb
from keras.layers import LSTM, Dense, Embedding
from keras.models import Sequential
from keras.preprocessing import sequence
max_features = 20000
# cut texts after this number of words (among top max_features most common words)
maxlen = 80
batch_size = 32
print("Loading data...")
(x_train, y_train), (x_test, y_test) = imdb.load_data(num_words=max_features)
print(len(x_train), "train sequences")
print(len(x_test), "test sequences")
print("Pad sequences (samples x time)")
x_train = sequence.pad_sequences(x_train, maxlen=maxlen)
x_test = sequence.pad_sequences(x_test, maxlen=maxlen)
print("x_train shape:", x_train.shape)
print("x_test shape:", x_test.shape)
print("Build model...")
model = Sequential()
model.add(Embedding(max_features, 128))
model.add(LSTM(128, dropout=0.2, recurrent_dropout=0.2))
model.add(Dense(1, activation="sigmoid"))
# try using different optimizers and different optimizer configs
model.compile(loss="binary_crossentropy", optimizer="adam", metrics=["accuracy"])
print("Train...")
model.fit(
x_train, y_train, batch_size=batch_size, epochs=15, validation_data=(x_test, y_test)
)
score, acc = model.evaluate(x_test, y_test, batch_size=batch_size)
print("Test score:", score)
print("Test accuracy:", acc)
参考链接:





