Transformers are RNNs: Fast Autoregressive Transformers with Linear Attention

发表时间: 2020-06 · arXiv:2006.16236

原文: https://arxiv.org/abs/2006.16236

文章标题:Transformer即RNN:采用线性注意力的快速自回归Transformer
作者/机构:Angelos Katharopoulos, Apoorv Vyas, Nikolaos Pappas, François Fleuret


速读

一句话结论 本文通过核函数特征图和矩阵乘法结合律将自注意力机制线性化,并将其重写为循环神经网络的形式,在保持模型表现的同时,将自回归推理速度提升了数千倍,且计算与显存复杂度从序列长度的平方级降至线性级。

要解决什么问题 标准 Transformer 模型在处理长序列时面临严重的计算与显存瓶颈,其核心原因在于自注意力机制:计算每个位置的注意力输出时,需计算查询(Query)与所有键(Key)的点积并应用 Softmax,导致计算量和显存占用均与序列长度呈二次方 $O(N^2)$ 增长。这不仅限制了上下文长度,还导致自回归推理极其缓慢,因为生成新词时需重新计算与所有历史词的注意力权重。尽管已有局部敏感哈希(如 Reformer)等方法降低训练复杂度,但它们要么要求查询和键必须相同,要么未能从根本上改变自回归推理阶段每次预测成本随序列长度增长的卡点。因此,如何从底层机制打破二次方复杂度并实质性加速自回归推理,是该领域的关键阻碍。

怎么做的 本文的核心思路是摒弃 Softmax 注意力,将其替换为基于核函数特征图的点积注意力,并利用矩阵乘法结合律改变计算顺序,将复杂度降为线性 $O(N)$。在标准自注意力中,输出基于 Softmax 相似度得分加权。作者发现,只要相似度函数保证非负性,就能用核函数的特征图映射 $\phi(\cdot)$ 替代 Softmax。将查询 $Q$ 和键 $K$ 映射后,注意力公式重写为: $$V'_i = \frac{\phi(Q_i)^T \sum_{j=1}^N \phi(K_j) V_j^T}{\phi(Q_i)^T \sum_{j=1}^N \phi(K_j)}$$ 这一改变去除了非线性的 Softmax 阻碍,允许先计算 $\sum \phi(K_j) V_j^T$ 和 $\sum \phi(K_j)$。这两个求和项对所有查询 $Q_i$ 共享,计算一次即可复用,彻底绕开了 $O(N^2)$ 的注意力矩阵。为保证非负性且避免梯度消失,作者采用了激活函数 $\phi(x) = \text{elu}(x) + 1$。为支持自回归模型的因果掩码(当前位置只受历史影响),作者引入累加状态变量 $S_i$ 和 $Z_i$,将上述过程转化为递推形式: $$S_i = S_{i-1} + \phi(K_i)V_i^T$$ $$Z_i = Z_{i-1} + \phi(K_i)$$ 通过该设计,带有因果掩码的 Transformer 层在数学上等价于一个循环神经网络(RNN)。$S_i$ 和 $Z_i$ 分别充当注意力与归一化记忆状态。自回归推理时,模型无需保留所有历史键值,只需在每个时间步以常数时间更新这两个固定大小的隐藏状态。此外,为解决训练时存储所有中间状态导致显存暴增的问题,作者将梯度的反向传播推导为类似 RNN 的累加和形式,确保前向和后向计算均只需线性时间与常数级显存。

效果如何 实验在 NVIDIA GTX 1080 Ti 或 P40 显卡上进行,涵盖序列复制、图像生成(MNIST 和 CIFAR-10)及自动语音识别(WSJ)任务。对比基线包括代表传统路线的标准 Softmax Transformer、代表最高效路线的 Reformer(基于局部敏感哈希),以及代表传统循环架构的双向 LSTM。在基准测试中,本方法与 Reformer 的资源消耗随序列长度严格线性增长,而标准 Softmax 呈二次方爆炸。在 8 层模型、序列长度 784 的 MNIST 图像自回归生成中,本方法达到与标准 Softmax 几乎相同的质量(0.83 对比 0.82 bits/dim),但得益于 RNN 式常数级状态更新,吞吐量达每秒 142 张,比标准 Softmax 快 300 多倍。在序列长度 3072 的 CIFAR-10 任务(16 层模型)中,在固定 7 天训练内,本方法完成的训练轮数是标准 Softmax 的 3 倍,困惑度更低(3.14 对比 3.26),推理速度快 4000 多倍。该方法也有代价。在非自回归的语音识别任务中,本方法的音素错误率和训练速度虽大幅优于 LSTM 和 Reformer,但标准 Softmax 依然取得了全局最低的 7.3% 错误率,表明线性化特征图的表达能力相比完整 Softmax 仍有轻微折损。此外,由于其 RNN 式推理计算成本极低,在单张图像生成(批大小为 1)时,GPU 算力无法充分利用,计算瓶颈变为序列外层循环,导致此时 CPU 运行速度反而快于 GPU。

A1 主要贡献

核心问题
标准的Transformer模型虽然在多种任务中表现出色,但其核心组件自注意力(self-attention)的计算和内存复杂性与输入序列长度 N 呈二次方关系,即 $O(N^2)$。这使得处理非常长的序列时,其计算成本过高、速度过慢,从而限制了模型的上下文长度,影响了时间连贯性和捕捉长期依赖的能力。尽管现有的一些高效Transformer方法(如稀疏分解、局部敏感哈希)在训练上降低了复杂性,但它们并未加速自回归推理过程。

研究目标
本文旨在提出一种线性Transformer模型,该模型能够显著减少内存占用,并将计算复杂度降低至与上下文长度呈线性关系,即 $O(N)$。同时,该模型需要能大幅提升自回归推理的速度。

创新点
1. 线性化自注意力机制:作者将自注意力表达为核函数特征图(kernel feature maps)的线性点积。通过利用矩阵乘法的结合律,成功将计算复杂度从 $O(N^2)$ 降低到 $O(N)$。
2. 高效的因果掩码:本文提出了一种适用于线性化注意力的因果掩码(causal masking)实现,其同样具有线性的复杂度和恒定的内存占用。
3. 揭示Transformer与RNN的关系:该线性化框架揭示了自回归Transformer与循环神经网络(RNN)之间的内在联系。基于此,作者将Transformer层重写为RNN形式,使其能够在自回归推理任务中实现数千倍的速度提升。


A3 背景知识/关键Observation/设计原则

2.1. 高效的Transformer

2.2. 理解自注意力

2.3. 线性化的Softmax


A2 方法细节

本节形式化地提出了线性Transformer。通过将传统的softmax注意力改为基于特征图的点积注意力,实现了更好的时间和内存复杂度,并获得了一个能像RNN一样以线性时间进行序列生成的因果模型。

3.1. Transformer

3.2. 线性化注意力

3.2.1. 特征图与计算成本

3.3. 因果掩码

3.3.1. 梯度计算
3.3.2. 训练与推理

3.4. Transformer即RNN


A4 实验环境


A4 实验结果

4.1. 合成任务

4.2. 图像生成

4.3. 自动语音识别


A5 结论

本文提出了线性Transformer,通过利用矩阵乘法的结合律,成功地将自注意力的计算和内存成本降低到与序列长度呈线性关系。研究表明,该模型可以与因果掩码结合使用,并保持其线性的渐进复杂度。最终,本文将Transformer模型表达为一种循环神经网络(RNN),使其能够在自回归任务上实现数千倍的推理加速。

未来工作展望
1. 深入研究Transformer与RNN的关系:这一特性为未来研究RNN和Transformer中信息的存储与检索机制开辟了多种方向。
2. 探索新的特征图:另一个值得探索的研究方向是为线性注意力选择不同的特征图。例如,使用随机傅里叶特征来近似RBF核,可能允许我们直接使用以softmax注意力预训练的模型。


A6 附录

A. 梯度推导

B. 训练过程

C. 图像生成吞吐量讨论

D. 图像生成的定性结果