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.5 LSTM

在实践中,对于要求网络利用距离当前处理位置较远的信息的任务,训练 RNN 相当困难。尽管 RNN 可以访问整个前序序列,但隐藏状态编码的信息往往相当局部,与输入序列中最近的部分以及最近的决策更相关。然而,远距离信息对许多语言应用至关重要。考虑下面这个语言建模例子。

(14.19) The flights the airline was canceling were full.

在 airline 后面给 was 较高的概率相对容易,因为 airline 为单数一致提供了很强的局部上下文。但要给 were 赋予恰当概率就困难得多,不仅因为复数名词 flights 距离很远,还因为中间的上下文中更近的是单数名词 airline。理想情况下,网络应能在需要时保留复数 flights 的远距离信息,同时仍然正确处理序列中间的部分。

图 14.12 用于序列分类的双向 RNN。组合正向与反向过程的最终隐藏单元,以表示整个序列。这个组合表示作为后续分类器的输入。

RNN 无法传递关键信息的一个原因是:隐藏层,以及决定隐藏层取值的权重,被要求同时完成两项任务:提供对当前决策有用的信息,并更新、传递未来决策所需的信息。

训练 RNN 的第二个困难来自沿时间反向传播误差信号的需要。回忆 14.1.2 节,时间 tt 的隐藏层会参与下一时间步的计算,因此会对下一时间步的损失产生影响。结果是,在训练的反向过程中,隐藏层会根据序列长度反复相乘。这一过程经常导致梯度最终变为零,这种情况称为梯度消失问题。

为了解决这些问题,人们设计了更复杂的网络架构,显式管理随时间维护相关上下文的任务:让网络学习遗忘不再需要的信息,并记住仍将用于未来决策的信息。

RNN 最常用的扩展是长短期记忆(long short-term memory,LSTM)网络(Hochreiter and Schmidhuber, 1997)。LSTM 将上下文管理问题分为两个子问题:从上下文中删除不再需要的信息,以及添加之后决策可能需要的信息。解决这两个问题的关键是学习如何管理上下文,而不是把策略硬编码到架构中。LSTM 首先在架构中加入显式的上下文层(除了通常的循环隐藏层之外),并使用带门的专门神经单元来控制进出网络层的单元的信息流。这些门通过额外的权重实现,依次作用于输入、前一隐藏层和前一上下文层。

LSTM 中的门共享一种常见设计模式:每个门由一个前馈层、一个 sigmoid 激活函数,以及与被门控层逐点相乘组成。选择 sigmoid 作为激活函数,是因为它倾向于把输出推向 0 或 1。再结合逐点乘法,其作用类似于二值掩码:被门控层中与掩码中接近 1 的值对应的值几乎不变地通过;对应较小值的部分则基本被擦除。

首先考虑遗忘门。该门的作用是从上下文中删除不再需要的信息。遗忘门计算前一状态隐藏层与当前输入的加权和,并将结果输入 sigmoid。随后将这个掩码与上下文向量逐元素相乘,以删除上下文中不再需要的信息。两个向量的逐元素乘法用运算符 \odot 表示,有时也称 Hadamard 乘积;所得向量与两个输入向量维度相同,其第 ii 个元素是两个输入向量第 ii 个元素的乘积:

ft=σ(Ufht1+Wfxt)\mathbf {f} _ {t} = \sigma (\mathbf {U} _ {f} \mathbf {h} _ {t - 1} + \mathbf {W} _ {f} \mathbf {x} _ {t})
kt=ct1ft(14.20)\mathbf {k} _ {t} = \mathbf {c} _ {t - 1} \odot \mathbf {f} _ {t}\tag{14.20}

(14.21)

下一个任务是从前一隐藏状态和当前输入中计算实际需要提取的信息——这与我们对所有循环网络一直使用的基本计算相同。

gt=tanh(Ught1+Wgxt)(14.22)\mathbf {g} _ {t} = \tanh (\mathbf {U} _ {g} \mathbf {h} _ {t - 1} + \mathbf {W} _ {g} \mathbf {x} _ {t})\tag{14.22}

接下来生成加法门的掩码,以选择要加入当前上下文的信息。

it=σ(Uiht1+Wixt)jt=gtit(14.23)\begin{array}{l} \mathbf {i} _ {t} = \sigma (\mathbf {U} _ {i} \mathbf {h} _ {t - 1} + \mathbf {W} _ {i} \mathbf {x} _ {t}) \\ \mathbf {j} _ {t} = \mathbf {g} _ {t} \odot \mathbf {i} _ {t} \end{array}\tag{14.23}

(14.24)

然后把它加到修改后的上下文向量上,得到新的上下文向量。

ct=jt+kt(14.25)\mathbf {c} _ {t} = \mathbf {j} _ {t} + \mathbf {k} _ {t}\tag{14.25}

最后使用输出门,决定当前隐藏状态需要哪些信息(而不是决定哪些信息需要为未来决策保留)。

ot=σ(Uoht1+Woxt)(14.26)\mathbf {o} _ {t} = \sigma (\mathbf {U} _ {o} \mathbf {h} _ {t - 1} + \mathbf {W} _ {o} \mathbf {x} _ {t})\tag{14.26}
ht=ottanh(ct)(14.27)\mathbf {h} _ {t} = \mathbf {o} _ {t} \odot \operatorname{tanh} (\mathbf {c} _ {t})\tag{14.27}

图 14.13 展示了单个 LSTM 单元的完整计算。给定各个门相应的权重,LSTM 接收前一时间步的上下文层和隐藏层,以及当前输入向量作为输入;然后输出更新后的上下文向量和隐藏向量。

在每个时间步为 LSTM 提供输出的是隐藏状态 hth_t。这个输出可以作为堆叠式 RNN 中后续层的输入;在网络的最后一层,hth_t 可以用来提供 LSTM 的最终输出。

图 14.13 单个 LSTM 单元的计算图。每个单元的输入包括当前输入 xx、前一隐藏状态 ht1h_{t-1} 和前一上下文 ct1c_{t-1}。输出是新的隐藏状态 hth_t 和更新后的上下文 ctc_t

图 14.14 前馈网络、简单循环网络(SRN)和长短期记忆网络(LSTM)中使用的基本神经单元。

14.5.1 门控单元、层和网络

LSTM 使用的神经单元显然比基本前馈网络中的神经单元复杂得多。幸运的是,这种复杂性被封装在基本处理单元内部,因此我们仍能保持模块化,并轻松试验不同架构。图 14.14 展示了每种单元关联的输入和输出。

最左边的(a)是基本前馈单元:一组权重和一个激活函数决定其输出;当多个这样的单元排列成层时,层内单元之间没有连接。接着,(b)表示简单循环网络中的单元。现在它有两个输入,并配有额外的一组权重;但仍然只有一个激活函数和一个输出。

LSTM 单元增加的复杂性被封装在单元本身之中。与基本循环单元(b)相比,LSTM 对外唯一增加的复杂性,就是增加了作为输入和输出的上下文向量。

这种模块化是 LSTM 单元强大且应用广泛的关键。LSTM 单元(或 GRU 等其他变体)可以替换到 14.4 节描述的任何网络架构中。与简单 RNN 一样,使用门控单元的多层网络可以展开为深层前馈网络,并按照通常的方式使用反向传播进行训练。因此在实践中,对于任何使用循环网络的现代系统,LSTM 已经取代 RNN 成为标准单元。