Skip to article frontmatterSkip to article content
Site not loading correctly?

This may be due to an incorrect BASE_URL configuration. See the MyST Documentation for reference.

14.7 使用 RNN 的编码器—解码器模型

本节介绍编码器—解码器模型。当我们需要把一个输入序列转换为长度不同、且与输入不按词逐一对齐的输出序列时,就会使用这种模型。

已经读过第 13 章的读者,可能已经在 Transformer 架构及其机器翻译应用中见过这种模型;但由于读者的阅读顺序可能不同,有些人会先阅读 RNN 再阅读 Transformer,因此我们在这里重新介绍这一架构。

回忆序列标注任务:其中有两个序列,但它们长度相同(例如在词性标注中,每个词元都有一个对应标签);每个输入都与特定输出相关联,输出标签主要依赖局部信息。因此,要判断一个词是动词还是名词,我们主要查看该词及其相邻词。

相比之下,编码器—解码器模型尤其用于机器翻译等任务,其中输入序列与输出序列的长度可以不同,而且输入词元与输出词元之间的映射可能非常间接(有些语言把动词放在句首,另一些语言则放在句末)。第 13 章介绍过机器翻译;这里暂且指出,把英语句子映射为他加禄语或约鲁巴语句子时,两者可能包含非常不同数量的词,词的顺序也可能完全不同。

编码器—解码器网络有时也称为序列到序列网络,是一种给定输入序列、能够生成上下文恰当且长度任意的输出序列的模型。编码器—解码器网络已应用于摘要、问答和对话等非常广泛的任务,在机器翻译中尤其流行。

这些网络背后的关键思想是使用编码器网络接收输入序列,并创建其上下文化表示,通常称为上下文。随后把这个表示传给解码器,由解码器生成特定于任务的输出序列。图 14.16 展示了该架构。

图 14.16 编码器—解码器架构。上下文是输入隐藏表示的函数,解码器可以用多种方式使用它。

编码器—解码器网络由三个概念性组件构成:

  1. 编码器接收输入序列 x1:nx_{1:n},并生成相应的上下文化表示序列 h1:nh_{1:n}。LSTM、卷积网络和 Transformer 都可以用作编码器。

  2. 上下文向量 cch1:nh_{1:n} 的函数,向解码器传达输入的核心信息。

  3. 解码器接收 cc 作为输入,生成任意长度的隐藏状态序列 h1:mh_{1:m},再由此得到相应的输出状态序列 y1:my_{1:m}。与编码器一样,解码器也可以由任何种类的序列架构实现。

本节描述一个基于一对 RNN 的编码器—解码器网络;第 13 章已经介绍过用于机器翻译的编码器—解码器。我们从条件 RNN 语言模型 p(y)p(y)(即序列 yy 的概率)开始,逐步建立编码器—解码器模型的方程。

回忆在任意语言模型中,都可以如下分解概率:

p(y)=p(y1)p(y2y1)p(y3y1,y2)p(ymy1,,ym1)(14.28)p (y) = p \left(y _ {1}\right) p \left(y _ {2} \mid y _ {1}\right) p \left(y _ {3} \mid y _ {1}, y _ {2}\right) \dots p \left(y _ {m} \mid y _ {1}, \dots , y _ {m - 1}\right)\tag{14.28}

在 RNN 语言建模中,在特定时间 tt,我们通过语言模型传入前 t1t-1 个词元的前缀,使用前向推理生成一系列隐藏状态,最后得到对应于前缀最后一个词的隐藏状态。然后,以该前缀的最终隐藏状态为起点,生成下一个词元。

更正式地说,假设 gg 是 tanh 或 ReLU 一类的激活函数,它以时间 tt 的输入和时间 t1t-1 的隐藏状态为参数;softmax 的范围是可能的词表项集合。那么在时间 tt,输出 yt\mathbf{y}_t 与隐藏状态 ht\mathbf{h}_t 的计算如下:

ht=g(ht1,xt)(14.29)\mathbf {h} _ {t} = g (\mathbf {h} _ {t - 1}, \mathbf {x} _ {t})\tag{14.29}
y^t=softmax(ht)(14.30)\hat {\mathbf {y}} _ {t} = \operatorname{softmax} (\mathbf {h} _ {t})\tag{14.30}

要把这种带自回归生成的语言模型变成能够把一种语言的源文本翻译为另一种语言的目标文本的编码器—解码器翻译模型,只需做一个小改动:在源文本末尾加入句子分隔标记,然后直接拼接目标文本。

我们使用 <s> 作为句子分隔词元。考虑把英语源文本“the green witch arrived”翻译为西班牙语句子“llegó la bruja verde”(逐词释义为“arrived the witch green”)。同样也可以用问答对或文本—摘要对来说明编码器—解码器模型。

xx 表示英语源文本加分隔词元 <s>,用 yy 表示目标文本(这里是西班牙语)。于是,编码器—解码器模型按如下方式计算概率 p(yx)p(y|x)

p(yx)=p(y1x)p(y2y1,x)p(y3y1,y2,x)p(ymy1,,ym1,x)(14.31)p (y | x) = p \left(y _ {1} | x\right) p \left(y _ {2} \mid y _ {1}, x\right) p \left(y _ {3} \mid y _ {1}, y _ {2}, x\right) \dots p \left(y _ {m} \mid y _ {1}, \dots , y _ {m - 1}, x\right)\tag{14.31}

图 14.17 展示了简化版编码器—解码器模型的设置(完整模型需要注意力这一新概念,将在下一节介绍)。

图 14.17 展示了英语源文本“the green witch arrived”、句子分隔词元 <s> 和西班牙语目标文本“llegó la bruja verde”。翻译源文本时,我们先对其运行网络并执行前向推理,生成隐藏状态,直到到达源文本末尾。然后开始自回归生成:在源输入末尾隐藏层以及句末标记提供的上下文中请求一个词。之后的词以先前的隐藏状态和上一个生成词的嵌入为条件。

图 14.17 使用基本 RNN 版编码器—解码器进行机器翻译时,翻译单个句子(推理阶段)。源句与目标句以分隔词元连接,解码器使用编码器最后一个隐藏状态中的上下文信息。

我们在图 14.18 中对这个模型做一些形式化和推广。(为避免混淆,需要时用上标 eedd 区分编码器和解码器的隐藏状态。)图左侧的网络元素处理输入序列 xx,构成编码器。虽然简化图只展示了编码器的单层网络,但堆叠架构才是常规做法:取堆叠顶层的输出状态作为最终表示;编码器由堆叠的双向 LSTM 组成,其中顶层正向和反向过程的隐藏状态拼接起来,为每个时间步提供上下文化表示。

图 14.18 基于 RNN 的基本编码器—解码器架构在推理阶段翻译句子的更形式化版本。编码器 RNN 的最终隐藏状态 hne\mathbf{h}_n^e 作为解码器的上下文,在解码器 RNN 中充当 h0dh_0^d,同时也提供给每个解码器隐藏状态。

编码器的全部目的就是生成输入的上下文化表示。这个表示体现在编码器的最终隐藏状态 hne\mathbf{h}_n^e 中,也称为上下文 cc,随后传给解码器。

最简单的解码器网络会取这个状态,只用它初始化解码器的第一个隐藏状态;第一个解码器 RNN 单元将 cc 作为前一隐藏状态 h0d\mathbf{h}_0^d。随后,解码器逐个元素地自回归生成输出序列,直到生成序列结束标记。每个隐藏状态都以上一个隐藏状态和上一步生成的输出为条件。

如图 14.18 所示,我们还会做更复杂的处理:让上下文向量 cc 不仅提供给第一个解码器隐藏状态,以确保在生成输出序列的过程中,上下文向量 cc 的影响不会逐渐减弱。具体做法是把 cc 作为当前隐藏状态计算的参数,使用下式:

htd=g(y^t1,ht1d,c)(14.32)\mathbf {h} _ {t} ^ {d} = g (\hat {y} _ {t - 1}, \mathbf {h} _ {t - 1} ^ {d}, \mathbf {c})\tag{14.32}

现在可以给出这种基本编码器—解码器模型中、在每个解码时间步都可获得上下文的解码器的完整方程。回忆 gg 代表某种 RNN,y^t1\hat{\mathbf{y}}_{t-1} 是上一步从 softmax 中采样出的输出的嵌入:

c=hneh0d=chtd=g(y^t1,ht1d,c)y^t= softmax (htd)(14.33)\begin{array}{r c l} \mathbf {c} & = & \mathbf {h} _ {n} ^ {e} \\ \mathbf {h} _ {0} ^ {d} & = & \mathbf {c} \\ \mathbf {h} _ {t} ^ {d} & = & g (\hat {\mathbf {y}} _ {t - 1}, \mathbf {h} _ {t - 1} ^ {d}, \mathbf {c}) \\ \hat {\mathbf {y}} _ {t} & = & \text { softmax } (\mathbf {h} _ {t} ^ {d}) \end{array}\tag{14.33}

因此,y^t\hat{\mathbf{y}}_t 是词表上的概率向量,表示每个词在时间 tt 出现的概率。为了生成文本,我们从分布 y^t\hat{\mathbf{y}}_t 中采样。例如,贪心选择就是在每个时间步选择概率最高的词来生成。第 7.6 节讨论过其他采样方法。

14.7.1 训练编码器—解码器模型

编码器—解码器架构以端到端方式训练。每个训练样本都是一对字符串组成的元组,即一个源字符串和一个目标字符串。将源—目标对与分隔词元拼接后,它们就可以作为训练数据。

对机器翻译而言,训练数据通常由句子及其译文组成,可以从标准的对齐句对数据集中取得,具体将在 13.2.2 节讨论。得到训练集后,训练过程与基于 RNN 的语言模型相同:给网络输入源文本,然后从分隔词元开始,自回归地训练模型预测下一个词,如图 14.19 所示。

注意训练(图 14.19)与推理(图 14.17)在每个时间步输出方面的差异。推理时,解码器使用自己的估计输出 y^t\hat{y}_t 作为下一个时间步 xt+1x_{t+1} 的输入。因此,随着生成更多词元,解码器往往会越来越偏离金标准目标句。因而在训练时,解码器更常使用教师强制。教师强制是指强制系统把训练数据中的金标准目标词元作为下一个输入 xt+1x_{t+1},而不是允许它依赖(可能错误的)解码器输出 y^t\hat{y}_t。这样可以加快训练。

图 14.19 基于基本 RNN 的编码器—解码器机器翻译方法的训练。注意,在解码器中,我们通常不会传播模型的 softmax 输出 y^t\hat{y}_t,而是使用教师强制,强制每个输入在训练时使用正确的金标准值。我们计算解码器中 y^\hat{y} 的 softmax 输出分布,以计算每个词元的损失;随后可以对这些损失求平均,得到句子级损失。该损失再通过解码器参数和编码器参数反向传播。