第十九讲 · 语言模型与循环神经网络

Views: --

MLP 和 CNN 通常一次接收一个固定形状输入,但语言、语音、视频都有顺序:前面出现的内容会改变后面的含义。循环神经网络(Recurrent Neural Network,RNN)用隐藏状态把过去的信息带到当前时间步。

一、语言模型在计算什么

一句由 nn 个词组成的序列 W=(w1,,wn)W=(w_1,\ldots,w_n),语言模型可以计算整句概率

P(W)=P(w1,w2,,wn),P(W)=P(w_1,w_2,\ldots,w_n),

也可以给定前文预测下一个词

P(wtw1,,wt1).P(w_t\mid w_1,\ldots,w_{t-1}).

由条件概率定义 P(A,B)=P(A)P(BA)P(A,B)=P(A)P(B\mid A),反复展开得到课件正文和附录强调的链式法则:

P(w1,,wn)=t=1nP(wtw1,,wt1).P(w_1,\ldots,w_n) =\prod_{t=1}^n P(w_t\mid w_1,\ldots,w_{t-1}).

因此只要能在每个位置预测下一个词,就能给整句话赋概率。

二、计数、马尔可夫假设与 N-gram

最直接的方法是“出现次数相除”,但完整前文组合几乎无穷,训练语料中大量句子从未出现。N-gram 用 kk 阶马尔可夫假设截断历史:

P(wtw1,,wt1)P(wtwtk,,wt1).P(w_t\mid w_1,\ldots,w_{t-1}) \approx P(w_t\mid w_{t-k},\ldots,w_{t-1}).
  • Unigram 不看前文;
  • Bigram 只看前一个词;
  • Trigram 看前两个词。

窗口越长,理论上信息越多,但可能组合数指数增长、计数更稀疏。更根本的问题是语言有长距离依赖:“I grew up in France … I speak fluent French” 中 French 依赖很久之前的 France,固定短窗口很难捕捉。

三、从 MLP 到 RNN

普通隐藏层只看当前输入:

h=tanh(Whxx+bh).h=\tanh(W_{hx}x+b_h).

RNN 再接入上一时刻隐藏状态:

ht=tanh(Whxxt+Whhht1+bh),h_t=\tanh(W_{hx}x_t+W_{hh}h_{t-1}+b_h), ot=Wyhht+by,pt=softmax(ot).o_t=W_{yh}h_t+b_y, \qquad p_t=\operatorname{softmax}(o_t).

hth_t 是到时刻 tt 为止的历史摘要,ptp_t 是下一个符号的概率分布。最关键的是:Whx,Whh,WyhW_{hx},W_{hh},W_{yh} 在所有时间步共享,所以网络能处理不同长度序列,也不会随序列变长增加参数量。

把循环沿时间展开后,看起来像许多层串联的网络;这些“层”是同一个 RNN 单元在不同时刻的副本,不是彼此独立的参数。

四、hello 字符模型例子

课件词表为

V={h,e,l,o},\mathcal V=\{h,e,l,o\},

训练序列是 hello。输入与监督目标错开一位:

时间步输入 xtx_t目标 y^t\hat y_t
1he
2el
3ll
4lo

每个字符可用 one-hot 向量表示,例如按 [h,e,l,o][h,e,l,o] 排列,e(0,1,0,0)T(0,1,0,0)^{\mathsf T}。网络在每一步输出四个字符的概率,序列损失是各时间步交叉熵之和或平均:

L=t=1Tlogpt(y^t).L=-\sum_{t=1}^{T}\log p_t(\hat y_t).

同一个输入 l 在第 3 步要预测 l、第 4 步要预测 o。仅看当前字符无法区分,隐藏状态必须保留前文位置,这正是 RNN 相比静态 MLP 的作用。

五、一个标量状态例子

用最小模型展示“记忆”怎样进入下一步:

ht=tanh(xt+0.5ht1),h0=0.h_t=\tanh(x_t+0.5h_{t-1}), \qquad h_0=0.

输入序列为 x1=1,x2=0x_1=1,x_2=0。则

h1=tanh(1)0.762,h_1=\tanh(1)\approx0.762, h2=tanh(0+0.5×0.762)0.364.h_2=\tanh(0+0.5\times0.762) \approx0.364.

第二步输入虽然为零,状态仍非零,因为第一步的信息通过 0.5h10.5h_1 延续下来。若递归系数太小,记忆会快速衰减;若过大,反向梯度可能不稳定。

六、时间反向传播 BPTT

训练时先沿时间前向,再从最后一步向前反向,这称为 Backpropagation Through Time(BPTT)。令

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

隐藏预激活的反向量可递推为

δt=(WyhTtot+WhhTδt+1)(1ht2).\delta_t= \left( W_{yh}^{\mathsf T}\frac{\partial \ell_t}{\partial o_t} +W_{hh}^{\mathsf T}\delta_{t+1} \right) \odot(1-h_t^2).

第一项来自当前输出损失,第二项来自未来状态。因为参数跨时间共享,权重梯度必须累加所有时间步贡献:

LWhh=t=1Tδtht1T,\frac{\partial L}{\partial W_{hh}} =\sum_{t=1}^{T}\delta_t h_{t-1}^{\mathsf T}, LWhx=t=1TδtxtT.\frac{\partial L}{\partial W_{hx}} =\sum_{t=1}^{T}\delta_t x_t^{\mathsf T}.

这与卷积核在不同空间位置共享后累加梯度很相似:MLP 沿层复用计算模式,CNN 沿空间共享参数,RNN 沿时间共享参数。

长序列常使用截断 BPTT,只向前回传固定步数,以控制显存和计算;代价是更难学习超过截断窗口的依赖。

七、不同输入输出结构

RNN 不只做“每步输入、每步输出”。课件按序列长度概括了四类结构:

结构含义例子
NNN\rightarrow N每步输入、每步输出字符预测、视频逐帧分类
N1N\rightarrow1整段输入、一个输出文本情感分类
1N1\rightarrow N一个条件、生成序列图像描述、条件音乐生成
NMN\rightarrow M输入输出长度可不同机器翻译、摘要、语音识别

NMN\rightarrow M 常由编码器读完输入,再由解码器生成输出。固定长度状态会形成信息瓶颈,后来 Attention 与 Transformer 正是为更直接地访问输入序列而发展。

八、普通 RNN 的长期依赖问题

BPTT 中早期状态的梯度包含许多个 Jacobian 连乘。粗略看,它会反复乘 WhhW_{hh}tanh\tanh'

hTht=k=t+1Thkhk1.\frac{\partial h_T}{\partial h_t} =\prod_{k=t+1}^{T} \frac{\partial h_k}{\partial h_{k-1}}.

乘子范数持续小于 1 时梯度消失,持续大于 1 时梯度爆炸。梯度裁剪能缓解爆炸,却不能恢复已经消失的长期信号。LSTM 用门控和一条更接近加法的状态通路缓解这个问题,见下一讲:LSTM 与序列生成

九、常见误区与 sanity check

  • 句子概率是各条件概率的乘积,实际计算常取对数求和以避免数值下溢。
  • N-gram 的 NN 指一个片段中总词数;Bigram 预测当前词时只条件于前 1 个词。
  • 时间展开后每步使用同一组参数,不能把 Whh(1),Whh(2)W_{hh}^{(1)},W_{hh}^{(2)} 当作独立权重。
  • 共享参数的梯度要跨所有时间步累加,而不是只取最后一步。
  • 隐藏状态不是无损保存全部历史,它只是固定维摘要,普通 RNN 尤其容易遗忘长距离信息。

评论