门控循环单元


梯度问题

在上一篇日志中,我们讨论了如何在循环神经网络中计算梯度,以及矩阵连续乘积可以导致梯度消失或梯度爆炸的问题。这进一步引出了我们对于梯度问题的三个思考角度

  • 早期观测值对预测所有未来观测值具有非常重要的意义。
  • 一些词元没有相关的观测值,我们希望有一些机制来跳过隐状态表示中的此类词元
  • 序列的各个部分之间存在逻辑中断,我们希望有一些机制来重置内部状态表示

在本小节中,我们先介绍一种解决此类问题的模型,即门控循环单元(Gated Recurrent Unit, GRU)。这是一种简单且计算速度快的模型,其基于长短期记忆模型(Long-short Trem Memory, LSTM)而来。

重置门和更新门

门控循环单元与普通的循环神经网络之间的关键区别在于:前者支持隐状态的门控。这意味着模型有专门的机制来确定应该何时更新隐状态,以及应该何时重置隐状态。这些机制是可学习的,并且能够解决上面列出的问题。

例如,如果第一个词元非常重要,模型将学会在第一次观测之后不更新隐状态;同样,模型也可以学会跳过不相关的临时观测。模型还将学会在需要的时候重置隐状态。

我们先介绍重置门更新门的概念,将其视为 (0,1)(\boldsymbol{0},\boldsymbol{1}) 区间的向量。

  • 重置门:控制“可能还想记住”的过去状态的数量
  • 更新门:控制新状态中有多少个是旧状态的副本

下图描述了门控循环单元中的重置门和更新门的输入,输入是由当前时间步的输入 Xt\boldsymbol{X}_t 和前一时间步的隐状态 Ht1\boldsymbol{H}_{t-1} 给出。两个门的输出由使用 sigmoid\mathrm{sigmoid} 激活函数的两个全连接层给出。

重置门和更新门

对于给定的时间步 tt,假设输入是小批量 XtRn×d\boldsymbol{X}_t\in \mathbb{R}^{n\times d},其中样本个数为 nn,输入个数为 dd。上一个时间步的隐状态表示为 Ht1Rn×h\boldsymbol{H}_{t-1}\in\mathbb{R}^{n\times h},其中隐藏单元个数为 hh。因此可以表示当前时间步的重置门 RtRn×h\boldsymbol{R}_t\in\mathbb{R}^{n\times h} 和更新门 ZtRn×h\boldsymbol{Z}_t\in\mathbb{R}^{n\times h} 如下

Rt=σ(XtWxr+Ht1Whr+br)\boldsymbol{R}_t=\sigma(\boldsymbol{X}_t\boldsymbol{W}_{xr}+\boldsymbol{H}_{t-1}\boldsymbol{W}_{hr}+\boldsymbol{b}_r)

Zt=σ(XtWxz+Ht1Whz+br)\boldsymbol{Z}_t=\sigma(\boldsymbol{X}_t\boldsymbol{W}_{xz}+\boldsymbol{H}_{t-1}\boldsymbol{W}_{hz}+\boldsymbol{b}_r)

其中的权重参数为 Wxr,WxzRd×h,Whr,WhzRh×h\boldsymbol{W}_{xr},\boldsymbol{W}_{xz}\in\mathbb{R}^{d\times h},\boldsymbol{W}_{hr},\boldsymbol{W}_{hz}\in\mathbb{R}^{h\times h},偏置参数为 br,bzR1×h\boldsymbol{b}_r,\boldsymbol{b}_z\in\mathbb{R}^{1\times h}。我们先看重置门 Rt\boldsymbol{R}_t 的性质

  • 如果某个维度接近 0,对应重置操作,忽略旧状态
  • 如果某个维度接近 1,对应保留操作,保留旧状态

再来看更新门 Zt\boldsymbol{Z}_t 的性质

  • 如果某个维度接近 0,对应更新操作,忽略旧状态
  • 如果某个维度接近 1,对应保留操作,保留旧状态

候选隐状态

现在尝试把重置门和常规隐状态结合,定义时间步 tt 时的候选隐状态

H~t=tanh(XtWxh+(RtHt1)Whr+bh)\tilde{\boldsymbol{H}}_t=\tanh(\boldsymbol{X}_t\boldsymbol{W}_{xh}+(\boldsymbol{R}_t\circ\boldsymbol{H}_{t-1})\boldsymbol{W}_{hr}+\boldsymbol{b}_h)

上式中的 \circ 表示逐元素乘积。

根据重置门的性质,RtHt1\boldsymbol{R}_t\circ\boldsymbol{H}_{t-1} 可以看做基于重置门,对重要旧内容的过滤。候选隐状态也可以看做基于重置门过滤的循环神经网络模型,其捕捉了序列中的短期依赖关系。

候选隐状态

更新隐状态

上一步的候选隐状态仅仅过滤出了有效的记忆,现在还需要结合更新门得到最终的隐状态。时间步 tt 的隐状态 Ht\boldsymbol{H}_t 主要和旧状态 Ht1\boldsymbol{H}_{t-1} 和候选隐状态 H~t\tilde{\boldsymbol{H}}_t 有关,配合更新门 Zt\boldsymbol{Z}_t 有如下更新公式

Ht=ZtHt1+(1Zt)H~t\boldsymbol{H}_t=\boldsymbol{Z}_t\circ \boldsymbol{H}_{t-1}+(\boldsymbol{1}-\boldsymbol{Z}_t)\circ\tilde{\boldsymbol{H}}_t

这种设计可以帮助我们处理循环神经网络的梯度消失问题,并更好的捕获时间步距离长的序列依赖关系。例如,如果整个子序列的所有时间步更新门都接近 1\boldsymbol{1},则无论序列的长度如何,序列起始步的旧隐状态都将保留传递到序列结束。

更新门

总之,门控循环单元具有以下两个显著特征:

  • 重置门有助于捕获序列中的短期依赖关系
  • 更新门有助于捕获序列中的长期依赖关系

代码实现

只需要参考普通RNN模型的代码,修改其参数和模型定义函数。

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
def get_params(vocab_size, num_hiddens, device):
num_inputs = num_outputs = vocab_size

def normal(shape):
return torch.randn(size=shape, device=device)*0.01

def three():
return (normal((num_inputs, num_hiddens)),
normal((num_hiddens, num_hiddens)),
torch.zeros(num_hiddens, device=device))

W_xz, W_hz, b_z = three() # 更新门参数
W_xr, W_hr, b_r = three() # 重置门参数
W_xh, W_hh, b_h = three() # 候选隐状态参数
# 输出层参数
W_hq = normal((num_hiddens, num_outputs))
b_q = torch.zeros(num_outputs, device=device)
# 附加梯度
params = [W_xz, W_hz, b_z, W_xr, W_hr, b_r, W_xh, W_hh, b_h, W_hq, b_q]
for param in params:
param.requires_grad_(True)
return params

def init_gru_state(batch_size, num_hiddens, device):
return (torch.zeros((batch_size, num_hiddens), device=device), )

def gru(inputs, state, params):
W_xz, W_hz, b_z, W_xr, W_hr, b_r, W_xh, W_hh, b_h, W_hq, b_q = params
H, = state
outputs = []
for X in inputs:
Z = torch.sigmoid((X @ W_xz) + (H @ W_hz) + b_z)
R = torch.sigmoid((X @ W_xr) + (H @ W_hr) + b_r)
H_tilda = torch.tanh((X @ W_xh) + ((R * H) @ W_hh) + b_h)
H = Z * H + (1 - Z) * H_tilda
Y = H @ W_hq + b_q
outputs.append(Y)
return torch.cat(outputs, dim=0), (H,)

训练参考结果如下

1
2
3
困惑度 1.1, 17144.5 词元/秒 cuda:0
time traveller with a slight accession ofcheerfulness really thi
travelleryou can show black is white by argument said filby

GRU训练

长短期记忆网络


门控记忆元

我们在本小节进一步补充门控的概念,一边补充GRU的技术细节,另一边介绍一种新的循环神经网络模型——长短期记忆网络。长短期记忆网络引入了记忆元(Memory Cell),或简称为单元。我们姑且将其理解为与隐藏层同尺寸的额外信息,记时间步为 tt 的记忆元为 Ct\boldsymbol{C}_t。为了控制记忆元,我们需要不同的门。

  • 输出门(Output Gate):从单元中输出条目
  • 输入门(Input Gate):决定何时将数据读入单元
  • 遗忘门(Forget Gate):重置单元的内容

类似于GRU模型,对于给定的时间步 tt,假设输入是小批量 XtRn×d\boldsymbol{X}_t\in \mathbb{R}^{n\times d},其中样本个数为 nn,输入个数为 dd。上一个时间步的隐状态表示为 Ht1Rn×h\boldsymbol{H}_{t-1}\in\mathbb{R}^{n\times h},其中隐藏单元个数为 hh。因此可以表示当前时间步的输入门 ItRn×h\boldsymbol{I}_t\in\mathbb{R}^{n\times h}、输出门 OtRn×h\boldsymbol{O}_t\in\mathbb{R}^{n\times h} 和遗忘门 FtRn×h\boldsymbol{F}_t\in\mathbb{R}^{n\times h} 如下

It=σ(XtWxi+Ht1Whi+bi)\boldsymbol{I}_t=\sigma(\boldsymbol{X}_t\boldsymbol{W}_{xi}+\boldsymbol{H}_{t-1}\boldsymbol{W}_{hi}+\boldsymbol{b}_i)

Ot=σ(XtWxo+Ht1Who+bo)\boldsymbol{O}_t=\sigma(\boldsymbol{X}_t\boldsymbol{W}_{xo}+\boldsymbol{H}_{t-1}\boldsymbol{W}_{ho}+\boldsymbol{b}_o)

Ft=σ(XtWxf+Ht1Whf+bf)\boldsymbol{F}_t=\sigma(\boldsymbol{X}_t\boldsymbol{W}_{xf}+\boldsymbol{H}_{t-1}\boldsymbol{W}_{hf}+\boldsymbol{b}_f)

我们进一步介绍候选记忆元 C~tRn×d\tilde{\boldsymbol{C}}_t\in\mathbb{R}^{n\times d},其计算方式为

C~t=tanh(XtWxc+Ht1Whc+bc)\tilde{\boldsymbol{C}}_t=\tanh(\boldsymbol{X}_t\boldsymbol{W}_{xc}+\boldsymbol{H}_{t-1}\boldsymbol{W}_{hc}+\boldsymbol{b}_c)

候选记忆元

记忆元和隐状态

在GRU模型中,有一种机制来控制输入和遗忘(或跳过)。类似地,在长短期记忆网络中,也有两个门用于这样的目的:输入门 It\boldsymbol{I}_t 控制采用多少来自 C~t\tilde{\boldsymbol{C}}_t 的新数据,而遗忘门 Ot\boldsymbol{O}_t 控制保留多少过去记忆元 Ct1\boldsymbol{C}_{t-1} 的内容。当前时间步的记忆元表示为

Ct=FtCt1+ItC~t\boldsymbol{C}_t=\boldsymbol{F}_t\circ \boldsymbol{C}_{t-1}+\boldsymbol{I}_t\circ\tilde{\boldsymbol{C}}_t

如果遗忘门始终为 1\boldsymbol{1} 且输入门始终为 0\boldsymbol{0},则过去的记忆元 Ct1\boldsymbol{C}_{t-1} 将随时间被保存并传递到当前时间步。引入这种设计是为了缓解梯度消失问题,并更好地捕获序列中的长距离依赖关系。

现在需要考虑如何利用记忆元计算当前状态的隐状态,有下式

Ht=Ottanh(Ct)\boldsymbol{H}_t=\boldsymbol{O}_t\circ \tanh(\boldsymbol{C}_t)

只要输出门接近 1\boldsymbol{1},我们就能够有效地将所有记忆信息传递给预测部分;而对于输出门接近 0\boldsymbol{0} 的情况,我们只保留记忆元内的所有信息,而不需要更新隐状态。记忆元和隐状态的计算流程图如下

记忆元和隐状态

为了帮助更好理解LSTM的原理,我们重新叙述整个模型流程

  1. LSTM使用记忆元储存长期记忆,使用隐状态储存短期记忆
  2. 遗忘门 Ft\boldsymbol{F}_t 决定前期记忆 Ct1\boldsymbol{C}_{t-1} 保留多少
  3. 候选记忆元 C~t\tilde{\boldsymbol{C}}_t 由输入 Xt\boldsymbol{X}_t 和短期记忆 Ht1\boldsymbol{H}_{t-1} 决定
  4. 输入门 It\boldsymbol{I}_{t}C~t\tilde{\boldsymbol{C}}_t 共同决定当前时间步保留多少候选记忆,然后和遗忘门处理后的记忆元求和,得到 Ct\boldsymbol{C}_t
  5. 输出门 Ot\boldsymbol{O}_t 和长期记忆 Ct\boldsymbol{C}_t 共同决定下一步的隐状态,即短期记忆

由于长期记忆链条不依赖于任何参数,因此其有效解决了梯度消失或梯度爆炸的问题。同时由于长短期记忆的共同出现,其保证了短期记忆(隐状态链条)的权重梯度计算也不会出现消失或爆炸的情况出现。

代码实现

这里只给出模型定义部分的代码。

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
def get_lstm_params(vocab_size, num_hiddens, device):
num_inputs = num_outputs = vocab_size

def normal(shape):
return torch.randn(size=shape, device=device)*0.01

def three():
return (normal((num_inputs, num_hiddens)),
normal((num_hiddens, num_hiddens)),
torch.zeros(num_hiddens, device=device))

W_xi, W_hi, b_i = three() # 输入门参数
W_xf, W_hf, b_f = three() # 遗忘门参数
W_xo, W_ho, b_o = three() # 输出门参数
W_xc, W_hc, b_c = three() # 候选记忆元参数
# 输出层参数
W_hq = normal((num_hiddens, num_outputs))
b_q = torch.zeros(num_outputs, device=device)
# 附加梯度
params = [W_xi, W_hi, b_i, W_xf, W_hf, b_f, W_xo, W_ho, b_o, W_xc, W_hc,
b_c, W_hq, b_q]
for param in params:
param.requires_grad_(True)
return params

def init_lstm_state(batch_size, num_hiddens, device):
return (torch.zeros((batch_size, num_hiddens), device=device),
torch.zeros((batch_size, num_hiddens), device=device))

def lstm(inputs, state, params):
[W_xi, W_hi, b_i, W_xf, W_hf, b_f, W_xo, W_ho, b_o, W_xc, W_hc, b_c,
W_hq, b_q] = params
(H, C) = state
outputs = []
for X in inputs:
I = torch.sigmoid((X @ W_xi) + (H @ W_hi) + b_i)
F = torch.sigmoid((X @ W_xf) + (H @ W_hf) + b_f)
O = torch.sigmoid((X @ W_xo) + (H @ W_ho) + b_o)
C_tilda = torch.tanh((X @ W_xc) + (H @ W_hc) + b_c)
C = F * C + I * C_tilda
H = O * torch.tanh(C)
Y = (H @ W_hq) + b_q
outputs.append(Y)
return torch.cat(outputs, dim=0), (H, C)

深度循环神经网络


事实上,我们可以将多层循环神经网络堆叠在一起,通过对几个简单层的组合产生灵活的机制。这种具有多个隐藏层的循环神经网络称为深度循环神经网络,其每个隐状态都连续地传递到当前层的下一个时间步和下一层的当前时间步。

深度循环神经网络

对于一个 LL 层的深层循环神经网络模型,假设时间步 tt 有小批量输入 XtRn×d\boldsymbol{X}_t\in\mathbb{R}^{n\times d},同时将第 ll 个隐藏层的隐状态设为 Ht(l)Rn×h\boldsymbol{H}_t^{(l)}\in\mathbb{R}^{n\times h},输出层变量为 OtRn×q\boldsymbol{O}_t\in\mathbb{R}^{n\times q}

Ht(0)=Xt\boldsymbol{H}_t^{(0)}=\boldsymbol{X}_t,第 ll 个隐藏层状态使用激活函数 ϕl\phi_l,则有

Ht(l)=ϕl(Ht(l1)Wxh(l)+Ht1(l)Whh(l)+bh(l))\boldsymbol{H}_t^{(l)}=\phi_l(\boldsymbol{H}_t^{(l-1)}\boldsymbol{W}_{xh}^{(l)}+\boldsymbol{H}_{t-1}^{(l)}\boldsymbol{W}_{hh}^{(l)}+\boldsymbol{b}_h^{(l)})

而最后的输出层仅基于第 ll 个隐藏层最终的隐状态计算

Ot=Ht(L)Whq+bq\boldsymbol{O}_t=\boldsymbol{H}_t^{(L)}\boldsymbol{W}_{hq}+\boldsymbol{b}_q

与多层感知机一样,隐藏层数目 LL 和隐藏单元数目 hh 都是超参数。另外,用门控循环单元或长短期记忆网络的隐状态来代替中的隐状态进行计算,可以很容易地得到深度门控循环神经网络或深度长短期记忆神经网络。

双向循环神经网络


上下文序列

考虑下面若干个文本序列预测问题

  • 我__
  • 我__饿了
  • 我__饿了,可以吃掉整个电脑屏幕
  • __饿了

根据可获取的信息量和不同的上下文范围,这些文本序列预测答案有所不同。常规的循环神经网络通常根据上文内容进行预测,而如果我们希望在循环神经网络中拥有一种机制,使之能够根据下文内容反推上文,我们就需要修改循环神经网络的设计,称之为双向循环神经网络模型(Bidirectional RNNs)。

模型定义

双向循环神经网络添加了反向传递信息的隐藏层,以便更灵活地处理此类信息。下图描述了具有单个隐藏层的双向循环神经网络的架构,可以看出,该神经网络具有两个不同传递方向的隐状态。

双向RNN

对于时间步 tt,给定一个小批量输入 XRn×d\boldsymbol{X}\in\mathbb{R}^{n\times d},并令隐藏层的激活函数为 ϕ\phi。在这个双向架构中,令时间步的前向和反向隐状态分别为

Ht=ϕ(XtWxh+Ht1Whh+bh)Rn×h\overrightarrow{\boldsymbol{H}}_t=\phi(\boldsymbol{X}_t\overrightarrow{\boldsymbol{W}}_{xh}+\overrightarrow{\boldsymbol{H}}_{t-1}\overrightarrow{\boldsymbol{W}}_{hh}+\overrightarrow{\boldsymbol{b}}_h)\in\mathbb{R}^{n\times h}

Ht=ϕ(XtWxh+Ht1Whh+bh)Rn×h\overleftarrow{\boldsymbol{H}}_t=\phi(\boldsymbol{X}_t\overleftarrow{\boldsymbol{W}}_{xh}+\overleftarrow{\boldsymbol{H}}_{t-1}\overleftarrow{\boldsymbol{W}}_{hh}+\overleftarrow{\boldsymbol{b}}_h)\in\mathbb{R}^{n\times h}

需要注意的是,需要把前向隐状态 Ht\overrightarrow{\boldsymbol{H}}_t 和反向隐状态 Ht\overleftarrow{\boldsymbol{H}}_t 拼接起来,形成完整的隐状态 HtRn×2h\boldsymbol{H}_t\in\mathbb{R}^{n\times 2h},再传播到输出层得到输出结果 Ot\boldsymbol{O}_t

局限性

双向循环神经网络的一个关键特性是:使用来自序列两端的信息来估计输出。也就是说,我们使用来自过去和未来的观测信息来预测当前的观测。但是在对下一个词元进行预测的情况中,这样的模型并不是我们所需的。因为在预测下一个词元时,我们终究无法知道下一个词元的下文是什么,所以将不会得到很好的精度。具体地说,在训练期间,我们能够利用过去和未来的数据来估计现在空缺的词;而在测试期间,我们只有过去的数据,因此精度将会很差。

另一个严重问题是,双向循环神经网络的计算速度非常慢。其主要原因是网络的前向传播需要在双向层中进行前向和后向递归,并且网络的反向传播还依赖于前向传播的结果。因此,梯度求解将有一个非常长的链。

因此,双向循环神经网络的应用具有局限性。读者可借助工具得到双向RNN的代码,并在 The Time Machine 数据集上训练测试,观察输出结果和训练效果。

机器翻译


问题描述

机器翻译(Machine Translation)指的是将序列从一种语言自动翻译成另一种语言。在日常生活中,我们经常使用翻译软件进行语言转换,而这个过程主要依赖于一个具有高性能的机器翻译模型。

几十年来,机器翻译领域主要包括以下两种

  • 统计机器翻译:由翻译模型和语言模型等组成部分的统计分析模型
  • 神经机器翻译:基于神经网络的方法

本小节主要讨论神经机器翻译的数据集处理。

数据集

我们以Tatoeba项目的双语句子对组成的 English-Français 数据集为例。该数据集中的每一行都是制表符分隔的文本序列对,序列对由英文文本序列和翻译后的法语文本序列组成。每个文本序列可以是一个句子,也可以是包含多个句子的一个段落。

在这个将英语翻译成法语的机器翻译问题中,英语被称作源语言,法语被称作目标语言。运行下述代码即可将数据集下载到本地文件夹根目录的 data 文件夹中。

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
import os
import torch
from d2l import torch as d2l

d2l.DATA_HUB['fra-eng'] = (d2l.DATA_URL + 'fra-eng.zip',
'94646ad1522d915e7b0f9296181140edcf86a4f5')

def read_data_nmt():
"""载入“英语-法语”数据集"""
data_dir = d2l.download_extract('fra-eng')
with open(os.path.join(data_dir, 'fra.txt'), 'r',
encoding='utf-8') as f:
return f.read()

raw_text = read_data_nmt()
print(raw_text[:75])

出现下述输出则下载完成

1
2
3
4
5
6
7
Downloading ../data\fra-eng.zip from http://d2l-data.s3-accelerate.amazonaws.com/fra-eng.zip...
Go. Va !
Hi. Salut !
Run! Cours !
Run! Courez !
Who? Qui ?
Wow! Ça alors !

下载数据集后,原始文本数据需要经过几个预处理步骤。

  • 使用空格代替不间断空格
  • 使用小写字母替换大写字母
  • 并在单词和标点符号之间插入空格
1
2
3
4
5
6
7
8
9
10
11
12
13
def preprocess_nmt(text):
"""预处理“英语-法语”数据集"""
def no_space(char, prev_char):
return char in set(',.!?') and prev_char != ' '

# 使用空格替换不间断空格,使用小写字母替换大写字母
text = text.replace('\u202f', ' ').replace('\xa0', ' ').lower()
# 在单词和标点符号之间插入空格
out = [' ' + char if i > 0 and no_space(char, text[i - 1]) else char
for i, char in enumerate(text)]
return ''.join(out)

text = preprocess_nmt(raw_text)

词元化

在机器翻译中,最常见的做法是单词级词元化。定义 tokenize_nmt 函数对前 num_examples 个文本序列对进行词元化,其中每个词元要么是一个词,要么是一个标点符号。此函数最终返回两个词元列表:sourcetarget,其中

  • source[i]:是源语言第 i 个文本序列的词元列表
  • target[i] 是目标语言第 i 个文本序列的词元列表
1
2
3
4
5
6
7
8
9
10
11
12
13
def tokenize_nmt(text, num_examples=None):
"""词元化“英语-法语”数据数据集"""
source, target = [], []
for i, line in enumerate(text.split('\n')):
if num_examples and i > num_examples:
break
parts = line.split('\t')
if len(parts) == 2:
source.append(parts[0].split(' '))
target.append(parts[1].split(' '))
return source, target

source, target = tokenize_nmt(text)

词表

由于机器翻译数据集由语言对组成,因此我们可以分别为源语言和目标语言构建两个词表。使用单词级词元化时,词表大小将明显大于使用字符级词元化时的词表大小。

为了缓解这一问题,这里我们将出现次数少于2次的低频率词元 视为相同的未知 <unk> 词元。 除此之外,我们还指定了额外的特定词元,例如在小批量时用于将序列填充到相同长度的填充词 <pad>, 以及序列的开始词元 <bos> 和结束词元 <eos>。这些特殊词元在自然语言处理任务中比较常用。

1
2
src_vocab = d2l.Vocab(source, min_freq=2,
reserved_tokens=['<pad>', '<bos>', '<eos>'])

加载与训练

语言模型中的序列样本都应有一个固定的长度,无论这个样本是一个句子的一部分还是跨越了多个句子的一个片断。在机器翻译中,每个样本都是由源和目标组成的文本序列对, 其中的每个文本序列可能具有不同的长度。为了解决这一问题,我们可以通过截断填充的方法,通过指定长度的序列截断或 <pad> 填充,使得每个文本序列长度相同。

1
2
3
4
5
def truncate_pad(line, num_steps, padding_token):
"""截断或填充文本序列"""
if len(line) > num_steps:
return line[:num_steps] # 截断
return line + [padding_token] * (num_steps - len(line)) # 填充

现在我们定义函数 build_array_nmt,将文本序列转换成小批量数据集用于训练。我们将特定的 <eos> 词元添加到所有序列的末尾,用于表示序列的结束。当模型通过一个词元接一个词元地生成序列进行预测时,生成的 <eos> 词元说明完成了序列输出工作。此外,我们还记录每个文本序列的长度,统计长度时排除填充词元。

1
2
3
4
5
6
7
8
def build_array_nmt(lines, vocab, num_steps):
"""将机器翻译的文本序列转换成小批量"""
lines = [vocab[l] for l in lines]
lines = [l + [vocab['<eos>']] for l in lines]
array = torch.tensor([truncate_pad(
l, num_steps, vocab['<pad>']) for l in lines])
valid_len = (array != vocab['<pad>']).type(torch.int32).sum(1)
return array, valid_len

最后,我们可以通过数据迭代器进行训练了。

1
2
3
4
5
6
7
8
9
10
11
12
13
def load_data_nmt(batch_size, num_steps, num_examples=600):
"""返回翻译数据集的迭代器和词表"""
text = preprocess_nmt(read_data_nmt())
source, target = tokenize_nmt(text, num_examples)
src_vocab = d2l.Vocab(source, min_freq=2,
reserved_tokens=['<pad>', '<bos>', '<eos>'])
tgt_vocab = d2l.Vocab(target, min_freq=2,
reserved_tokens=['<pad>', '<bos>', '<eos>'])
src_array, src_valid_len = build_array_nmt(source, src_vocab, num_steps)
tgt_array, tgt_valid_len = build_array_nmt(target, tgt_vocab, num_steps)
data_arrays = (src_array, src_valid_len, tgt_array, tgt_valid_len)
data_iter = d2l.load_array(data_arrays, batch_size)
return data_iter, src_vocab, tgt_vocab

seq2seq


编码器-解码器架构

上一小节讨论的机器翻译是序列转换模型的一个核心问题,其输入和输出都是长度可变的序列。为了处理这种类型的输入和输出,我们可以设计一个包含两个主要组件的架构

  • 第一个组件是一个编码器(Encoder),它接受一个长度可变的序列作为输入,并将其转换为具有固定形状的编码状态
  • 第二个组件是一个解码器(Decoder),它将固定形状的编码状态映射到长度可变的序列

这被称为编码器-解码器(Encoder-Decoder)架构。

编码器-解码器

在编码器接口中,我们只指定长度可变的序列作为编码器的输入 X

1
2
3
4
5
6
7
8
9
from torch import nn

class Encoder(nn.Module):
"""编码器-解码器架构的基本编码器接口"""
def __init__(self, **kwargs):
super(Encoder, self).__init__(**kwargs)

def forward(self, X, *args):
raise NotImplementedError

在下面的解码器接口中,我们新增一个 init_state函数, 用于将编码器的输出转换为编码后的状态。为了逐个地生成长度可变的词元序列,解码器在每个时间步都会将输入和编码后的状态映射成当前时间步的输出词元。

1
2
3
4
5
6
7
8
9
10
class Decoder(nn.Module):
"""编码器-解码器架构的基本解码器接口"""
def __init__(self, **kwargs):
super(Decoder, self).__init__(**kwargs)

def init_state(self, enc_outputs, *args):
raise NotImplementedError

def forward(self, X, state):
raise NotImplementedError

编码器-解码器架构包含了一个编码器和一个解码器,并且还拥有可选的额外的参数。在前向传播中,编码器的输出用于生成编码状态,这个状态又被解码器作为其输入的一部分。下面的代码把二者合并了起来。

1
2
3
4
5
6
7
8
9
10
11
class EncoderDecoder(nn.Module):
"""编码器-解码器架构的基类"""
def __init__(self, encoder, decoder, **kwargs):
super(EncoderDecoder, self).__init__(**kwargs)
self.encoder = encoder
self.decoder = decoder

def forward(self, enc_X, dec_X, *args):
enc_outputs = self.encoder(enc_X, *args)
dec_state = self.decoder.init_state(enc_outputs, *args)
return self.decoder(dec_X, dec_state)

序列到序列学习

seq2seq模型全称为 Sequence to Sequence 模型,是一种处理序列数据的神经网络架构,广泛应用于自然语言处理领域。它能够将一个序列转换成另一个序列,而且这两个序列的长度可以不同,这使得seq2seq模型非常适合机器翻译、文本摘要、对话系统等任务。

seq2seq模型应用到中英翻译的一个例子如下

编码器神经网络

编码器将长度可变的输入序列转换成形状固定的上下文变量 c\boldsymbol{c},并且将输入序列的信息在该上下文变量中进行编码。根据这种特性,我们可以设计一个循环神经网络来实现编码器。

假设有一个序列组成的样本,其批量大小是1,输入序列为 x1,x2,,xTx_1,x_2,\cdots,x_T。在时间步 tt,循环神经网络将词元 xtx_t 的输入特征向量 xt\boldsymbol{x}_t 和上一步的隐状态 ht1\boldsymbol{h}_{t-1} 转化为当前步的隐状态 ht\boldsymbol{h}_t。编码器通过选取合适的函数 qq 得出上下文变量 c\boldsymbol{c},即有

c=q(h1,,hT)\boldsymbol{c}=q(\boldsymbol{h}_1,\cdots,\boldsymbol{h}_T)

由于词元 xtx_t 是一个具体的文本序列,因此我们需要使用嵌入层(Embedding Layer)获取该词元的特征向量 xt\boldsymbol{x}_t 以方便计算。嵌入层的权重是一个矩阵,其函数为输入词表的大小 vocab_size,列数为特征向量的维度 embed_size。对于索引为 i 的词元,嵌入层调取权重矩阵的第 i 行以返回其特征向量。

广义上讲,Embedding是将任何类型的数据转换为向量的过程。在本课程中姑且将其视为某种数据的可训练词典或可训练哈希表,即通过阅读海量训练数据,模型在预测时不断犯错、并学习修正参数。后文中我们总默认Embedding矩阵已通过预处理得到,并称每一个输入对应的Embedding向量为特征向量。

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
class Seq2SeqEncoder(d2l.Encoder):
"""用于序列到序列学习的循环神经网络编码器"""
def __init__(self, vocab_size, embed_size, num_hiddens, num_layers,
dropout=0, **kwargs):
super(Seq2SeqEncoder, self).__init__(**kwargs)
# 嵌入层
self.embedding = nn.Embedding(vocab_size, embed_size)
self.rnn = nn.GRU(embed_size, num_hiddens, num_layers,
dropout=dropout)

def forward(self, X, *args):
# 输出'X'的形状:(batch_size,num_steps,embed_size)
X = self.embedding(X)
# 在循环神经网络模型中,第一个轴对应于时间步
X = X.permute(1, 0, 2)
# 如果未提及状态,则默认为0
output, state = self.rnn(X)
# output的形状:(num_steps,batch_size,num_hiddens)
# state的形状:(num_layers,batch_size,num_hiddens)
return output, state

接下来,我们实例化上述编码器的实现。使用一个两层门控循环单元编码器,其隐藏单元数为 1616。给定一小批量的输入序列 X(批量大小为 44,时间步为 77)。在完成所有时间步后,最后一层的隐状态的输出是一个张量,其形状为(时间步数,批量大小,隐藏单元数)。

1
2
3
4
5
encoder = Seq2SeqEncoder(vocab_size=10, embed_size=8, num_hiddens=16,
num_layers=2)
encoder.eval()
X = torch.zeros((4, 7), dtype=torch.long)
output, state = encoder(X)

解码器神经网络

编码器输出的上下文变量 c\boldsymbol{c} 是由整个输入序列 x1,,xTx_1,\cdots,x_T 编码而来,因此解码器在时间步 tt' 的输出 yty_{t'} 均依赖于 c\boldsymbol{c} 和先前输出的子序列 y1,,yt1y_1,\cdots,y_{t'-1}。因此解码器的输出概率表示为

P(yty1,,yt1,c)P(y_{t'}|y_1,\cdots,y_{t'-1},\boldsymbol{c})

同样考虑设计一个循环神经网络模型。在时间步 tt',循环神经网络将上下文序列 c\boldsymbol{c}、上一个词元 yt1y_{t'-1} 的特征向量 yt1\boldsymbol{y}_{t'-1} 和上一步的隐状态 st1\boldsymbol{s}_{t-1} 转化为当前步的隐状态 st\boldsymbol{s}_t。用函数 gg 表示记为

st=g(yt1,c,st1)\boldsymbol{s}_{t'}=g(\boldsymbol{y}_{t'-1},\boldsymbol{c},\boldsymbol{s}_{t'-1})

获得当前步的隐状态后,可以使用输出层和 Softmax\mathrm{Softmax} 操作来计算当前步的条件概率分布。实现解码器时,我们直接使用编码器最后一个时间步的隐状态来初始化解码器的隐状态。这就要求使用循环神经网络实现的编码器和解码器具有相同数量的层和隐藏单元

为了进一步包含经过编码的输入序列的信息,上下文变量 c\boldsymbol{c} 在所有的时间步与解码器的输入进行拼接。为了预测输出词元的概率分布,在循环神经网络解码器的最后一层使用全连接层来变换隐状态。

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
class Seq2SeqDecoder(d2l.Decoder):
"""用于序列到序列学习的循环神经网络解码器"""
def __init__(self, vocab_size, embed_size, num_hiddens, num_layers,
dropout=0, **kwargs):
super(Seq2SeqDecoder, self).__init__(**kwargs)
self.embedding = nn.Embedding(vocab_size, embed_size)
self.rnn = nn.GRU(embed_size + num_hiddens, num_hiddens, num_layers,
dropout=dropout)
self.dense = nn.Linear(num_hiddens, vocab_size)

def init_state(self, enc_outputs, *args):
return enc_outputs[1]

def forward(self, X, state):
# 输出'X'的形状:(batch_size,num_steps,embed_size)
X = self.embedding(X).permute(1, 0, 2)
# 广播context,使其具有与X相同的num_steps
context = state[-1].repeat(X.shape[0], 1, 1)
X_and_context = torch.cat((X, context), 2)
output, state = self.rnn(X_and_context, state)
output = self.dense(output).permute(1, 0, 2)
# output的形状:(batch_size,num_steps,vocab_size)
# state的形状:(num_layers,batch_size,num_hiddens)
return output, state

接下来只需要实例化即可。

1
2
3
4
5
6
decoder = Seq2SeqDecoder(vocab_size=10, embed_size=8, num_hiddens=16,
num_layers=2)
decoder.eval()
state = decoder.init_state(encoder(X))
output, state = decoder(X, state)
output.shape, state.shape

损失函数

在每个时间步,解码器预测了输出词元的概率分布。类似于语言模型,可以使用 Softmax\mathrm{Softmax} 来获得分布,并通过计算交叉熵损失函数来进行优化。由于我们在部分词元结尾填充了无语义词元 <pad>,因此需要将这些特殊词元排除在损失函数计算之外。使用下述 sequence_mask 函数可以获取序列的有效长度。

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
def sequence_mask(X, valid_len, value=0):
"""在序列中屏蔽不相关的项"""
maxlen = X.size(1)
mask = torch.arange((maxlen), dtype=torch.float32,
device=X.device)[None, :] < valid_len[:, None]
X[~mask] = value
return X # 有效项编码为1,无关项编码为0

class MaskedSoftmaxCELoss(nn.CrossEntropyLoss):
"""带遮蔽的softmax交叉熵损失函数"""
# pred的形状:(batch_size,num_steps,vocab_size)
# label的形状:(batch_size,num_steps)
# valid_len的形状:(batch_size,)
def forward(self, pred, label, valid_len):
weights = torch.ones_like(label)
weights = sequence_mask(weights, valid_len)
self.reduction='none'
unweighted_loss = super(MaskedSoftmaxCELoss, self).forward(
pred.permute(0, 2, 1), label)
weighted_loss = (unweighted_loss * weights).mean(dim=1)
return weighted_loss

在下面的循环训练过程中,特定的序列开始词元 <bos> 和原始的输出序列(不包括序列结束词元 <eos>)拼接在一起作为解码器的输入。这被称为强制教学(Teacher Forcing),因为我们选取原始的输出序列送入解码器。

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
def train_seq2seq(net, data_iter, lr, num_epochs, tgt_vocab, device):
"""训练序列到序列模型"""
def xavier_init_weights(m):
if type(m) == nn.Linear:
nn.init.xavier_uniform_(m.weight)
if type(m) == nn.GRU:
for param in m._flat_weights_names:
if "weight" in param:
nn.init.xavier_uniform_(m._parameters[param])

net.apply(xavier_init_weights)
net.to(device)
optimizer = torch.optim.Adam(net.parameters(), lr=lr)
loss = MaskedSoftmaxCELoss()
net.train()
animator = d2l.Animator(xlabel='epoch', ylabel='loss',
xlim=[10, num_epochs])
for epoch in range(num_epochs):
timer = d2l.Timer()
metric = d2l.Accumulator(2) # 训练损失总和,词元数量
for batch in data_iter:
optimizer.zero_grad()
X, X_valid_len, Y, Y_valid_len = [x.to(device) for x in batch]
bos = torch.tensor([tgt_vocab['<bos>']] * Y.shape[0],
device=device).reshape(-1, 1)
dec_input = torch.cat([bos, Y[:, :-1]], 1) # 强制教学
Y_hat, _ = net(X, dec_input, X_valid_len)
l = loss(Y_hat, Y, Y_valid_len)
l.sum().backward() # 损失函数的标量进行“反向传播”
d2l.grad_clipping(net, 1)
num_tokens = Y_valid_len.sum()
optimizer.step()
with torch.no_grad():
metric.add(l.sum(), num_tokens)
if (epoch + 1) % 10 == 0:
animator.add(epoch + 1, (metric[0] / metric[1],))
print(f'loss {metric[0] / metric[1]:.3f}, {metric[1] / timer.stop():.1f} '
f'tokens/sec on {str(device)}')

embed_size, num_hiddens, num_layers, dropout = 32, 32, 2, 0.1
batch_size, num_steps = 64, 10
lr, num_epochs, device = 0.005, 300, d2l.try_gpu()

# 使用上一小节的数据集训练
train_iter, src_vocab, tgt_vocab = d2l.load_data_nmt(batch_size, num_steps)
encoder = Seq2SeqEncoder(len(src_vocab), embed_size, num_hiddens, num_layers,
dropout)
decoder = Seq2SeqDecoder(len(tgt_vocab), embed_size, num_hiddens, num_layers,
dropout)
net = d2l.EncoderDecoder(encoder, decoder)
train_seq2seq(net, train_iter, lr, num_epochs, tgt_vocab, device)

训练结果如下

1
loss 0.019, 12745.1 tokens/sec on cuda:0

训练

由于新版本 d2l 包中删去了编码器和解码器的类,因此若出现运行报错,请自定定义 Encoder,Decoder,删去所有引用 d2l 的类前缀 d2l.,并修改继承基类为自行定义的类。

预测与评估

为了采用一个接着一个词元的方式预测输出序列,每个解码器当前时间步的输入都将来自于前一时间步的预测词元。与训练类似,序列开始词元 <bos> 在初始时间步被输入到解码器中。整个预测过程如下图展示

预测

更多有关预测输出的策略见最后一小节。

我们可以通过与真实的标签序列进行比较来评估预测序列。Papineni 等人提出的 BLEU指标(Bilingual Evaluation Understudy)最先用于评估机器翻译的结果,但现在它已经被广泛用于测量许多应用的输出序列的质量,尤其是对于 nn 元语法。BLEU定义为

BLEU=exp(min(0,1len(label)len(predict))+n=1kwnlogpn)\mathrm{BLEU}=\exp\left(\min\left(0,1-\frac{\mathrm{len(label)}}{\mathrm{len(predict)}}\right)+\sum_{n=1}^kw_n\log p_n\right)

我们对部分符号做解释说明

  • len(label)\mathrm{len(label)}:标签序列的词元数
  • len(predict)\mathrm{len(predict)}:预测序列的词元数
  • kk:用于匹配的最长的 nn 元语法
  • wnw_n:权重,满足归一化 n=1kwn=1\displaystyle\sum_{n=1}^k w_n=1
  • pnp_nnn 元语法的精确度

这里我们进一步介绍如何计算 nn 元语法的精确度。对于一个长度为 TT 的序列 X=(x1,,xT)X=(x_1,\cdots,x_T),其 nn 元语法集合定义为

Gn(X)={(si,si+1,,si+n1)i=1,2,,Ln+1}G_n(X)=\{(s_i,s_{i+1},\cdots,s_{i+n-1})|i=1,2,\cdots,L-n+1\}

定义计数函数 CNT(g,X)\mathrm{CNT}(g,X) 表示某个 nn 元语法串 gg 在序列 XX 中出现的次数。假设我们有预测序列 Y^\hat{Y} 和参考标签序列 YY,则 nn 元语法的精确度定义为

pn=gGn(Y^)min(CNT(g,Y^),CNT(g,Y))gGn(Y^)CNT(g,Y^)p_n=\frac{\displaystyle\sum_{g\in G_n(\hat{Y})}\min\Big(\mathrm{CNT}(g,\hat{Y}),\mathrm{CNT}(g,Y)\Big)}{\displaystyle\sum_{g\in G_n(\hat{Y})}\mathrm{CNT}(g,\hat{Y})}

不难展开得到分母 gGn(Y^)CNT(g,Y^)=max(1,Y^n+1)\displaystyle\sum_{g\in G_n(\hat{Y})}\mathrm{CNT}(g,\hat{Y})=\max(1,|\hat{Y}|-n+1) ,如果我们有 MM 个参考翻译(标签)序列 Y(1),,Y(M)Y^{(1)},\cdots,Y^{(M)},则需要使用改进的精确度计算方法如下

pn=gGn(Y^)min(CNT(g,Y^),max1mMCNT(g,Y(M)))gGn(Y^)CNT(g,Y^)p_n=\frac{\displaystyle\sum_{g\in G_n(\hat{Y})}\min\Big(\mathrm{CNT}(g,\hat{Y}),\max_{1\leqslant m\leqslant M}\mathrm{CNT}(g,Y^{(M)})\Big)}{\displaystyle\sum_{g\in G_n(\hat{Y})}\mathrm{CNT}(g,\hat{Y})}

这表示我们取它在所有参考翻译中的最大出现次数,然后通过 min\min 进行裁剪。我们用一个简单的例子来解释精确度的计算,假设有预测序列 A,B,B,C,DA,B,B,C,D 和标签序列 A,B,C,D,E,FA,B,C,D,E,F。可以分别计算 n=1,2,3,4n=1,2,3,4 时的语法精确度为

  • n=1n=1 时:p1=1+1+1+151+1=45p_1=\dfrac{1+1+1+1}{5-1+1}=\dfrac45
  • n=2n=2 时:p2=1+0+1+152+1=34p_2=\dfrac{1+0+1+1}{5-2+1}=\dfrac34
  • n=3n=3 时:p3=0+0+153+1=13p_3=\dfrac{0+0+1}{5-3+1}=\dfrac13
  • n=4n=4 时:p4=0+054+1=0p_4=\dfrac{0+0}{5-4+1}=0

当预测序列与标签序列完全相同时,精确度 pn=1p_n=1,这表示预测结果良好。由于 nn 元语法越长则匹配难度越大,所以BLEU应当为更长的 nn 元语法的精确度分配更大的权重。通常而言人们习惯取 wn=1/kw_n=1/k,若忽略归一化条件则可以取 wn=1/2nw_n=1/2^n。我们选取后者的做法,并修改原来BLEU的计算公式为

BLEU=exp(min(0,1len(label)len(predict)))n=1kpn1/2n\mathrm{BLEU}=\exp\left(\min\left(0,1-\frac{\mathrm{len(label)}}{\mathrm{len(predict)}}\right)\right)\cdot\prod_{n=1}^k p_n^{1/2^n}

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
import math
import collections

def bleu(pred_seq, label_seq, k):
"""计算BLEU"""
pred_tokens, label_tokens = pred_seq.split(' '), label_seq.split(' ')
len_pred, len_label = len(pred_tokens), len(label_tokens)
score = math.exp(min(0, 1 - len_label / len_pred))
for n in range(1, k + 1):
num_matches, label_subs = 0, collections.defaultdict(int)
for i in range(len_label - n + 1):
label_subs[' '.join(label_tokens[i: i + n])] += 1
for i in range(len_pred - n + 1):
if label_subs[' '.join(pred_tokens[i: i + n])] > 0:
num_matches += 1
label_subs[' '.join(pred_tokens[i: i + n])] -= 1
score *= math.pow(num_matches / (len_pred - n + 1), math.pow(0.5, n))
return score

使用训练好的模型,就可以尝试翻译预测啦。修改 engs 和对应参考 fras 即可。

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
def predict_seq2seq(net, src_sentence, src_vocab, tgt_vocab, num_steps,
device, save_attention_weights=False):
"""序列到序列模型的预测"""
# 在预测时将net设置为评估模式
net.eval()
src_tokens = src_vocab[src_sentence.lower().split(' ')] + [
src_vocab['<eos>']]
enc_valid_len = torch.tensor([len(src_tokens)], device=device)
src_tokens = d2l.truncate_pad(src_tokens, num_steps, src_vocab['<pad>'])
# 添加批量轴
enc_X = torch.unsqueeze(
torch.tensor(src_tokens, dtype=torch.long, device=device), dim=0)
enc_outputs = net.encoder(enc_X, enc_valid_len)
dec_state = net.decoder.init_state(enc_outputs, enc_valid_len)
# 添加批量轴
dec_X = torch.unsqueeze(torch.tensor(
[tgt_vocab['<bos>']], dtype=torch.long, device=device), dim=0)
output_seq, attention_weight_seq = [], []
for _ in range(num_steps):
Y, dec_state = net.decoder(dec_X, dec_state)
# 我们使用具有预测最高可能性的词元,作为解码器在下一时间步的输入
dec_X = Y.argmax(dim=2)
pred = dec_X.squeeze(dim=0).type(torch.int32).item()
# 保存注意力权重(稍后讨论)
if save_attention_weights:
attention_weight_seq.append(net.decoder.attention_weights)
# 一旦序列结束词元被预测,输出序列的生成就完成了
if pred == tgt_vocab['<eos>']:
break
output_seq.append(pred)
return ' '.join(tgt_vocab.to_tokens(output_seq)), attention_weight_seq

engs = ['go .', "i lost .", 'he\'s calm .', 'i\'m home .']
fras = ['va !', 'j\'ai perdu .', 'il est calme .', 'je suis chez moi .']
for eng, fra in zip(engs, fras):
translation, attention_weight_seq = predict_seq2seq(
net, eng, src_vocab, tgt_vocab, num_steps, device)
print(f'{eng} => {translation}, bleu {bleu(translation, fra, k=2):.3f}')

翻译预测结果和BLEU计算如下

1
2
3
4
go . => va !, bleu 1.000
i lost . => j'ai perdu ., bleu 1.000
he's calm . => il est riche ., bleu 0.658
i'm home . => je suis en retard ?, bleu 0.447

搜索策略


为什么要优化搜索

前文提到,我们在seq2seq模型中逐个预测输出序列,直到预测序列中出现特定的序列结束词元 <eos>。下面我们来拆解这其中的一个关键问题——如何选取搜索策略?

在解码器中,当前的选择会影响后续所有步骤。一旦选择了某个词,后续解码器将基于这个词继续生成。对于任意时间步 tt',解码器输出 yty_{t'} 的概率取决于之前的输出和上下文 c\boldsymbol{c}。用 Y\mathcal{Y} 表示输出词表,其包含 <eos>,词表大小表示为 Y|\mathcal{Y}|。设输出序列的最大词元数为 TT',则机器翻译的目标就是从 O(YT)O\left(|\mathcal{Y}|^{T'}\right) 个可能的输出中寻找理想的输出。

传统策略

让我们先看一个简单的贪心策略。前文提到过,解码器通过自回归形式输出条件概率 P(yty1,,yt1,c)P(y_{t'}|y_1,\cdots,y_{t'-1},\boldsymbol{c}),因此我们对于每一个时间步 tt' 贪心地从 Y\mathcal{Y} 中找出具有最高条件概率的词元,即有

yt=argmaxyYP(yy1,,yt1,c)y_{t'}=\arg\max_{y\in\mathcal{Y}} P(y|y_1,\cdots,y_{t'-1},\boldsymbol{c})

当序列中包含 <eos> 或达到最大序列长度 TT' 时停止,输出完成。

然而,打过信息竞赛的朋友们都知道,贪心策略知识一味地追求局部最优解,这就无法保证全局最优解性。一般的,我们追求的理想输出序列应该是最大化条件概率乘积链 t=1TP(yty1,,yt1,c)\displaystyle\prod_{t'=1}^{T'} P(y_{t'}|y_1,\cdots,y_{t'-1},\boldsymbol{c}),而贪心策略显然无法保证这种最优性构造。

老一辈的做法表示,我们可以使用穷举法遍历每一个可能性,这就带来了 Θ(YT)\Theta\left(|\mathcal{Y}|^{T'}\right) 的时间复杂度,显然在大规模数据量下不可取。因此,我们需要寻找另一种平衡精度和计算量的搜索方法。

束搜索

束搜索(Bean Search)是贪心搜索的一个改进版本,其定义超参数 kk 称为束宽(Beam Size)。在时间步 t=1t'=1,我们选取具有条件概率前 kk 高的 kk 个词元。这 kk 个词元分别是 kk 个候选输出序列的第一个词元。在随后的每个时间步中,基于上一时间步的 kk 个候选输出序列,继续从 kYk|\mathcal{Y}| 个可能的选择中选取具有最高条件概率的 kk 个候选输出序列。下图展示了 k=2k=2 时的束搜索示例。

束搜索

由于我们不确定 <eos> 的位置,因此任意长度不超过 TT' 的束搜索结果都可以作为候选序列。我们选取使得下式左侧最大的序列作为输出序列,即有

1LαlogP(y1,,yLc)=1Lαt=1LlogP(yty1,,yt1,c)\frac{1}{L^{\alpha}}\log P(y_1,\cdots,y_L|\boldsymbol{c})=\frac{1}{L^{\alpha}}\sum_{t'=1}^L\log P(y_{t'}|y_1,\cdots,y_{t'-1},\boldsymbol{c})

其中 LL 为候选序列的长度,通常取 α=0.75\alpha=0.75 以惩罚长序列。束搜索的计算量为 O(kYT)O(k|\mathcal{Y}|T'),介于贪心搜索和穷举搜索之间。实际上,贪心搜索可以看作一种束宽为 11 的特殊类型的束搜索。通过灵活地选择束宽,束搜索可以在正确率和计算代价之间进行权衡。