第二十讲 · LSTM 与序列生成

Views: --

普通 RNN 能把过去带到现在,却很难跨越很长距离。判断“the clouds are in the sky”中的 sky,附近词就够了;判断“I grew up in France … I speak fluent French”中的 French,则必须保存很久以前的 France

长短期记忆网络(Long Short-Term Memory,LSTM)不再用一次 tanh\tanh 同时完成“记住、遗忘、输出”,而是让三个门分别控制信息流。

一、普通 RNN 为什么会忘

普通 RNN 状态为

ht=tanh(Whxxt+Whhht1+bh).h_t=\tanh(W_{hx}x_t+W_{hh}h_{t-1}+b_h).

从远处传回的梯度要反复乘 WhhW_{hh}tanh\tanh'。Sigmoid、Tanh 饱和区导数很小,连乘后梯度容易趋近零;若矩阵放大作用持续大于 1,又可能爆炸。

LSTM 增加细胞状态 CtC_t。它沿时间主要通过逐元素乘法和加法更新,让信息不必每一步都被新的非线性彻底改写。

二、LSTM 的四组候选量

把上一时刻隐藏状态与当前输入拼接:

zt=[ht1,xt].z_t=[h_{t-1},x_t].

遗忘门、输入门、输出门都经过 Sigmoid,因此每个分量在 (0,1)(0,1)

ft=σ(Wfzt+bf),f_t=\sigma(W_fz_t+b_f), it=σ(Wizt+bi),i_t=\sigma(W_iz_t+b_i), ot=σ(Wozt+bo).o_t=\sigma(W_oz_t+b_o).

候选记忆经过 Tanh,分量在 (1,1)(-1,1)

C~t=tanh(WCzt+bC).\widetilde C_t=\tanh(W_Cz_t+b_C).

直觉上:

  • ftf_t 决定旧记忆 Ct1C_{t-1} 保留多少;
  • iti_t 决定候选新信息写入多少;
  • oto_t 决定内部记忆有多少暴露为当前隐藏状态;
  • C~t\widetilde C_t 是当前准备写入的内容,不是门。

三、细胞状态与隐藏状态更新

课件第 40–42 页给出的核心公式是

Ct=ftCt1+itC~t,C_t=f_t\odot C_{t-1} +i_t\odot\widetilde C_t, ht=ottanh(Ct).h_t=o_t\odot\tanh(C_t).

第一项是保留的旧信息,第二项是写入的新信息。三个门都是向量,因此不同记忆维度可以做不同决定,而不是整块全留或全删。

一个标量计算例子

假设某一维上

Ct1=0.5,ft=0.8,it=0.3,C~t=0.4,ot=0.6.C_{t-1}=0.5, \quad f_t=0.8, \quad i_t=0.3, \quad \widetilde C_t=0.4, \quad o_t=0.6.

则新细胞状态为

Ct=0.8×0.5+0.3×0.4=0.52,C_t=0.8\times0.5+0.3\times0.4=0.52,

隐藏状态为

ht=0.6tanh(0.52)0.287.h_t=0.6\tanh(0.52)\approx0.287.

这个例子可按两条通路检查:旧记忆贡献 0.40.4,新信息贡献 0.120.12;若结果不等于二者之和,门或运算符写错了。

四、为什么它更利于长距离梯度

暂时只看状态主通路,当前状态对上一状态的直接导数为

CtCt1=ft.\frac{\partial C_t}{\partial C_{t-1}}=f_t.

跨越多步时有

CTCtk=t+1Tfk.\frac{\partial C_T}{\partial C_t} \approx\prod_{k=t+1}^{T}f_k.

网络若需要长期保存某条信息,可以让相关维度的 fkf_k 接近 11,同时让输入门接近 00,使状态沿近似恒等通路传递。相比每步都乘 WhhW_{hh} 和激活导数的普通 RNN,这条加法状态通路更容易保住梯度。

不过 LSTM 只是缓解,不是彻底消除长期依赖问题。许多 fk<1f_k<1 仍会使连乘衰减;门饱和也会让门参数梯度变小;序列计算依然必须逐步进行,难以完全并行。

五、从训练到文本生成

字符级语言模型训练时仍采用“输入当前字符,预测下一个字符”:

xt(ht,Ct)P(ct+1ct).x_t\rightarrow(h_t,C_t)\rightarrow P(c_{t+1}\mid c_{\le t}).

训练阶段常用真实上一个字符作为下一步输入,计算所有时间步的交叉熵并通过 BPTT 更新参数。生成阶段则形成闭环:

  1. 输入起始字符或提示文本;
  2. LSTM 输出下一个字符的概率分布;
  3. 从分布中选择或采样一个字符;
  4. 把选出的字符作为下一步输入;
  5. 重复直到终止标记或长度上限。

若每次都取概率最大的字符,结果稳定但容易重复;按概率采样更丰富,也更可能出错。温度 τ\tau 可调节分布:

pi=softmax(zi/τ).p_i=\operatorname{softmax}(z_i/\tau).
  • τ<1\tau<1:分布更尖锐,输出更保守;
  • τ>1\tau>1:分布更平,输出更多样也更随机。

六、课件的文章与 C 代码生成案例

课件展示了 RNN/LSTM 逐字符生成文章和 C 代码。模型并不是先学习语法树再手写程序,而是从训练文本中学习字符或词的条件分布。它可能学出缩进、括号、关键词等局部规律,生成看起来很像代码的片段。

“像代码”不等于“能运行”或“语义正确”。只用下一个字符预测训练的模型可能复制表面格式,却缺少类型、作用域和执行逻辑约束。评估生成代码至少要再做编译、测试和行为检查。

七、LSTM 与不同序列任务

LSTM 单元可以直接替换上一讲中的普通 RNN 单元,因此仍支持:

  • NNN\rightarrow N 的逐步标注或语言建模;
  • N1N\rightarrow1 的文本分类;
  • 1N1\rightarrow N 的条件生成;
  • 编码器–解码器式 NMN\rightarrow M 翻译与摘要。

区别在于每个时间步除了 hth_t,还显式传递 CtC_t。隐藏状态偏向当前对外可见表示,细胞状态偏向跨时间保存的内部记忆。

八、常见误区与 sanity check

  • CtC_thth_t 不是同一个状态;一个用于长期记忆,一个是经输出门筛选后的可见表示。
  • 遗忘门、输入门、输出门经 Sigmoid,范围为 (0,1)(0,1);候选记忆经 Tanh,范围为 (1,1)(-1,1)
  • 更新细胞状态用逐元素乘法 \odot,不是普通矩阵乘法。
  • 遗忘门为 11 只表示保留旧状态;若输入门也很大,仍会叠加新信息。
  • 梯度裁剪主要防爆炸,无法把已经消失的梯度恢复出来。
  • 训练时使用真实历史、生成时使用模型自己生成的历史,会产生暴露偏差;长文本错误可能逐步累积。

评论