LSTM 的原理是什么?三个门分别起什么作用?
简化版
LSTM(长短期记忆网络) 是改进版 RNN,专门解决普通 RNN记不住长距离信息(梯度消失)的问题。核心是引入一条细胞状态(cell state) ——一条贯穿整个序列的「记忆传送带」,信息可以在上面较少衰减地长距离流动。再用三个门(gate) 精细控制这条传送带上的信息:① 遗忘门——决定从细胞状态里丢掉哪些旧信息;② 输入门——决定把哪些新信息写入细胞状态;③ 输出门——决定从细胞状态里输出哪些信息作为当前隐藏状态。门是 0~1 的开关(sigmoid 控制),让 LSTM 能有选择地记住、遗忘、输出信息,从而捕捉长距离依赖。
详细版
两条状态线:
- 细胞状态 C_t:长期记忆的「传送带」,跨时间步以加法方式更新,信息衰减小。
- 隐藏状态 h_t:当前时间步的输出/短期状态。
三个门(都用 sigmoid 输出 0~1 的权重):
| 门 | 作用 | 公式(简化) |
|---|---|---|
| 遗忘门 f_t | 决定丢弃多少旧记忆 | f_t = σ(W_f·[h_{t-1}, x_t]) |
| 输入门 i_t | 决定写入多少新信息 | i_t = σ(W_i·[h_{t-1}, x_t]) |
| 输出门 o_t | 决定输出多少记忆 | o_t = σ(W_o·[h_{t-1}, x_t]) |
细胞状态更新(核心):
C_t = f_t ⊙ C_{t-1} + i_t ⊙ C̃_t
└丢弃旧记忆┘ └写入新信息┘ (C̃_t 是候选新记忆,tanh 生成)
h_t = o_t ⊙ tanh(C_t) (输出门控制输出)
为什么缓解梯度消失: 细胞状态用加法更新(C_t = f_t·C_{t-1} + ...),梯度沿细胞状态回传时不像普通 RNN 那样反复乘 W 连乘衰减,形成「梯度高速公路」。
完整版教学
一、LSTM 要解决什么——普通 RNN 记不住长依赖
普通 RNN 用一个隐藏状态传递记忆,但因为梯度消失,它记不住很久之前的信息(长距离依赖丢失,详见 RNN 梯度问题专题)。根本原因是记忆在每个时间步都被同一个权重矩阵反复变换,远处信息迅速衰减。
LSTM(Long Short-Term Memory) 的设计目标就是:让重要信息能在序列里长距离、少衰减地保存下来,同时能有选择地遗忘无用信息。它的两大法宝是细胞状态(记忆传送带) 和门控机制(信息开关)。
这里的“记不住”不是普通 RNN 在表达能力上绝对不可能表示远程关系,而是训练时远处监督难以通过 BPTT 有效更新早期状态。LSTM 用独立细胞状态和可学习门控提供更稳定的主传播路径,同时允许主动遗忘过期信息。代价是四组仿射变换、两条状态和更高的串行计算成本。
二、核心创新一:细胞状态——记忆的「传送带」
LSTM 比普通 RNN 多了一条细胞状态 C_t,它像一条贯穿整个序列的传送带,专门负责长期记忆:
C_{t-1} ──────→ C_t ──────→ C_{t+1} (信息在传送带上流动,很少被剧烈改变)
关键在于细胞状态的更新是加法为主(见下面公式),信息可以在遗忘门接近 1 且写入受控时较少衰减地跨越多个时间步,而不是像普通 RNN 那样每步都被权重矩阵乘一遍。这条「传送带」是 LSTM 能记住长距离信息的物理基础,也是它缓解梯度消失的关键(第六节)。
三、核心创新二:三个门——精细控制信息
光有传送带还不够,还要能决定往传送带上加什么、删什么、取什么。LSTM 用三个门来控制。每个门都是一个 sigmoid 层,输出 0~1 之间的值,像阀门/开关——0 表示「完全关闭(不通过)」,1 表示「完全打开(全部通过)」,通过逐元素相乘来「过滤」信息。
① 遗忘门(Forget Gate)f_t——决定丢掉哪些旧记忆
f_t = σ(W_f · [h_{t-1}, x_t] + b_f)
看当前输入和上一步隐藏状态,为细胞状态的每个维度输出一个 0~1 的值:接近 0 = 「这个旧记忆该忘掉」,接近 1 = 「保留这个旧记忆」。例如读到新主语时,遗忘门可能决定忘掉旧主语的性别信息。
② 输入门(Input Gate)i_t——决定写入哪些新信息
i_t = σ(W_i · [h_{t-1}, x_t] + b_i) 决定"写入多少"
C̃_t = tanh(W_c · [h_{t-1}, x_t] + b_c) 生成"候选新记忆"
输入门决定当前的新信息里,哪些值得写进细胞状态;候选记忆 C̃_t(用 tanh 生成)是「打算写入的内容」。两者相乘 = 实际写入的新信息。
③ 输出门(Output Gate)o_t——决定输出哪些信息
o_t = σ(W_o · [h_{t-1}, x_t] + b_o)
h_t = o_t ⊙ tanh(C_t)
输出门决定从当前细胞状态里,输出哪些部分作为这一步的隐藏状态 h_t(也是这一步的输出)。
四、细胞状态怎么更新——把门组合起来
三个门协同工作,细胞状态的更新公式(LSTM 的核心):
C_t = f_t ⊙ C_{t-1} + i_t ⊙ C̃_t
└── 遗忘:丢弃旧记忆 ──┘ └── 输入:写入新记忆 ──┘
h_t = o_t ⊙ tanh(C_t) (输出:从细胞状态取输出)
直觉:新记忆 = 保留一部分旧记忆(遗忘门筛选)+ 加入一部分新信息(输入门筛选),然后输出门决定拿多少出来用。整个过程让 LSTM 能有选择地记住重要信息、遗忘无关信息、在需要时才输出——这正是处理长序列所需要的。
五、一个直觉例子
读句子「我在法国长大……我说流利的 ___」:
- 读到「法国」:输入门把「国家=法国」写入细胞状态。
- 中间无关的词:遗忘门保留「法国」这条记忆(不遗忘),输入门不写入太多干扰。
- 读到「我说流利的」:输出门把「法国」相关信息输出,帮助预测「法语」。
「法国」这条信息在细胞状态的传送带上在门控合适时跨越多个时间步较少衰减地保留下来——这就是 LSTM 捕捉长依赖的方式。
六、为什么 LSTM 能缓解梯度消失
这是高频追问。普通 RNN 梯度消失是因为记忆更新是乘法(每步乘 W),连乘导致指数衰减。LSTM 的细胞状态更新是加法为主:
C_t = f_t ⊙ C_{t-1} + i_t ⊙ C̃_t
梯度沿细胞状态反向传播时,主要路径是通过 f_t ⊙ C_{t-1} 这一项——梯度大致按遗忘门 f_t 传递,而不是反复乘同一个权重矩阵 W。当遗忘门接近 1(表示「一直记住」)时,主路径梯度可以较少衰减地长距离回传,形成一条「梯度高速公路(constant error carousel)」,从而大幅缓解梯度消失、让 LSTM 能学到长距离依赖(详见 LSTM 缓解梯度专题)。
七、从门的矩阵形状算参数量
标准 LSTM 有遗忘、输入、输出和候选四组仿射变换。输入维度 D_x=3、隐藏维度 D_h=4 时,若把输入权重、循环权重和偏置分别计数,总参数量为 4(D_xD_h+D_h²+D_h)=4(12+16+4)=128。同尺寸普通 RNN 只有 D_xD_h+D_h²+D_h=32 个循环单元参数,因此 LSTM 的门控能力以更多参数和计算为代价。
| 结构 | 仿射组数 | 本例参数量 | 状态 |
|---|---|---|---|
| vanilla RNN | 1 | 32 | h_t |
| GRU | 3 | 96 | h_t |
| LSTM | 4 | 128 | h_t,C_t |
框架常把四组权重拼成一个大矩阵以提高 GEMM 效率,所以代码里未必看到四次独立线性层。不同实现还可能有双偏置、peephole 或 projection,参数公式要按具体 API 调整。
易错点:LSTM 的“三个门”之外还有候选记忆的一组参数,所以参数量按四组仿射变换计算。
八、常见误区与追问
- 误区:LSTM 只有三个门所以只有三组参数。 候选记忆也需要一组仿射变换,标准实现通常共四组。
- 误区:细胞状态通过加法更新就绝不会梯度消失。 主路径仍连乘遗忘门,门饱和和其他路径也会影响梯度。
- 追问:为什么遗忘门偏置有时初始化为正值? 让训练初期更倾向保留记忆,但并非所有任务和框架都必须如此。
- 追问:隐藏状态和细胞状态维度必须相同吗? 标准 LSTM 常相同,带 projection 的变体可让输出隐藏维度不同。
- 追问:门值接近 0 或 1 有什么代价? 能形成明确开关,但 sigmoid 饱和会让门参数本身梯度变小。
九、加强记忆
LSTM 解决普通 RNN 记不住长依赖(梯度消失) 的问题,两大法宝:① 细胞状态 C_t——贯穿序列的「记忆传送带」,以加法更新、信息衰减小;② 三个门(sigmoid 输出 0~1 的阀门):遗忘门 f_t(丢弃哪些旧记忆)、输入门 i_t(写入哪些新信息,配 tanh 候选 C̃_t)、输出门 o_t(输出哪些作为 h_t)。核心更新:C_t = f_t⊙C_{t-1} + i_t⊙C̃_t(保留旧记忆+写入新记忆),h_t = o_t⊙tanh(C_t)。缓解梯度消失的关键:细胞状态加法更新、梯度按遗忘门传递而非反复乘 W,遗忘门接近 1 时主路径梯度衰减较慢(梯度高速公路)。让 LSTM 能有选择地记住/遗忘/输出、捕捉长距离依赖。