第二十讲 · LSTM 与序列生成
普通 RNN 能把过去带到现在,却很难跨越很长距离。判断“the clouds are in the sky”中的 sky,附近词就够了;判断“I grew up in France … I speak fluent French”中的 French,则必须保存很久以前的 France。
长短期记忆网络(Long Short-Term Memory,LSTM)不再用一次 同时完成“记住、遗忘、输出”,而是让三个门分别控制信息流。
一、普通 RNN 为什么会忘
普通 RNN 状态为
从远处传回的梯度要反复乘 和 。Sigmoid、Tanh 饱和区导数很小,连乘后梯度容易趋近零;若矩阵放大作用持续大于 1,又可能爆炸。
LSTM 增加细胞状态 。它沿时间主要通过逐元素乘法和加法更新,让信息不必每一步都被新的非线性彻底改写。
二、LSTM 的四组候选量
把上一时刻隐藏状态与当前输入拼接:
遗忘门、输入门、输出门都经过 Sigmoid,因此每个分量在 :
候选记忆经过 Tanh,分量在 :
直觉上:
- 决定旧记忆 保留多少;
- 决定候选新信息写入多少;
- 决定内部记忆有多少暴露为当前隐藏状态;
- 是当前准备写入的内容,不是门。
三、细胞状态与隐藏状态更新
课件第 40–42 页给出的核心公式是
第一项是保留的旧信息,第二项是写入的新信息。三个门都是向量,因此不同记忆维度可以做不同决定,而不是整块全留或全删。
一个标量计算例子
假设某一维上
则新细胞状态为
隐藏状态为
这个例子可按两条通路检查:旧记忆贡献 ,新信息贡献 ;若结果不等于二者之和,门或运算符写错了。
四、为什么它更利于长距离梯度
暂时只看状态主通路,当前状态对上一状态的直接导数为
跨越多步时有
网络若需要长期保存某条信息,可以让相关维度的 接近 ,同时让输入门接近 ,使状态沿近似恒等通路传递。相比每步都乘 和激活导数的普通 RNN,这条加法状态通路更容易保住梯度。
不过 LSTM 只是缓解,不是彻底消除长期依赖问题。许多 仍会使连乘衰减;门饱和也会让门参数梯度变小;序列计算依然必须逐步进行,难以完全并行。
五、从训练到文本生成
字符级语言模型训练时仍采用“输入当前字符,预测下一个字符”:
训练阶段常用真实上一个字符作为下一步输入,计算所有时间步的交叉熵并通过 BPTT 更新参数。生成阶段则形成闭环:
- 输入起始字符或提示文本;
- LSTM 输出下一个字符的概率分布;
- 从分布中选择或采样一个字符;
- 把选出的字符作为下一步输入;
- 重复直到终止标记或长度上限。
若每次都取概率最大的字符,结果稳定但容易重复;按概率采样更丰富,也更可能出错。温度 可调节分布:
- :分布更尖锐,输出更保守;
- :分布更平,输出更多样也更随机。
六、课件的文章与 C 代码生成案例
课件展示了 RNN/LSTM 逐字符生成文章和 C 代码。模型并不是先学习语法树再手写程序,而是从训练文本中学习字符或词的条件分布。它可能学出缩进、括号、关键词等局部规律,生成看起来很像代码的片段。
“像代码”不等于“能运行”或“语义正确”。只用下一个字符预测训练的模型可能复制表面格式,却缺少类型、作用域和执行逻辑约束。评估生成代码至少要再做编译、测试和行为检查。
七、LSTM 与不同序列任务
LSTM 单元可以直接替换上一讲中的普通 RNN 单元,因此仍支持:
- 的逐步标注或语言建模;
- 的文本分类;
- 的条件生成;
- 编码器–解码器式 翻译与摘要。
区别在于每个时间步除了 ,还显式传递 。隐藏状态偏向当前对外可见表示,细胞状态偏向跨时间保存的内部记忆。
八、常见误区与 sanity check
- 与 不是同一个状态;一个用于长期记忆,一个是经输出门筛选后的可见表示。
- 遗忘门、输入门、输出门经 Sigmoid,范围为 ;候选记忆经 Tanh,范围为 。
- 更新细胞状态用逐元素乘法 ,不是普通矩阵乘法。
- 遗忘门为 只表示保留旧状态;若输入门也很大,仍会叠加新信息。
- 梯度裁剪主要防爆炸,无法把已经消失的梯度恢复出来。
- 训练时使用真实历史、生成时使用模型自己生成的历史,会产生暴露偏差;长文本错误可能逐步累积。