OpenAI Sparse Transformer

OpenAI 开发了 Sparse Transformer,这是一种旨在预测文本、图像和声音序列中下一个元素的深度神经网络。通过对注意力机制进行算法改进,该模型可以从比以往长 30 倍的序列中提取模式,从而能够使用数百层来对数万个元素进行建模。

Reducing Memory Constraints with Deep Attention

标准 Transformer 架构在每一层和每个注意力头都需要一个 $N \times N$ 的注意力矩阵。在处理原始音频或大型图像等高维数据时,这会产生显著的内存瓶颈。例如,一个 64x64x3 像素的 ImageNet 图像在深度 Transformer(64 层,4 个头)中存储矩阵需要 154 GB 的内存,远超标准 GPU 的 12-32 GB 容量。

To address this, OpenAI 实现了两种主要的优化:

  1. Checkpointing: 通过在反向传播期间从检查点重新计算注意力矩阵,最大的内存成本变得与层数无关。这允许训练具有实质上更大深度的网络;OpenAI 发现,在 CIFAR-10 等基准测试中,多达 128 层的模型表现优于较浅的网络。
  2. Operational Adjustments: 团队修改了初始化方案和 Transformer 内部的操作顺序,以支持增加的深度。

Sparse Attention and Algorithmic Complexity

为了处理即使单个注意力矩阵都变得不切实际的超大型输入,Sparse Transformer 利用了稀疏注意力模式。在这种方法中,每个输出位置仅从输入位置的一个子集(例如,$\sqrt{N}$ 个元素而非 $N$ 个)计算权重。这使得算法复杂度从 $O(N^2)$ 降低到 $O(N \sqrt{N})$。

Factorized Attention Patterns

为了确保模型保留学习动态、全局模式的能力,OpenAI 实现了注意力矩阵的二维分解。这使得网络可以通过两步稀疏注意力来关注所有位置:

  • Strided Attention: 每个位置关注其所在的行和列,有效地分解了完整的注意力操作。这对于具有二维结构的数据(如图像)特别有用。
  • Fixed Attention: 网络关注一个固定的列以及紧随最新列元素之后的内容。这种模式针对不符合二维结构的数据(如文本)进行了优化。

Experimental Results and Performance

Sparse Transformer 在以下基准数据集上实现了密度估计的新 SOTA 成绩:

  • CIFAR-10
  • Enwik8
  • Imagenet 64

OpenAI 观察到,稀疏注意力不仅比全注意力运行得更快,而且实现了更低的损失。研究人员建议,这可能是由于密集注意力中潜在的优化问题,或者是稀疏模式提供的有益归纳偏置。

Generative Capabilities across Modalities

Image Generation

该模型通过在 64x64 ImageNet 数据上的图像补全任务展示了对全局结构的掌握。由于模型是使用最大似然目标进行训练的,使用未调整的 1.0 的 softmax 温度生成的无条件样本反映了模型认为存在的图像全部分布,这偶尔可能会导致看起来很奇怪的样本。

Raw Audio Generation

通过修改位置嵌入,Sparse Transformer 可以生成原始音频波形。OpenAI 在原始古典音乐片段上训练了该模型,使其能够生成 65,000 个元素的序列,这大约对应 5 秒钟的原始音频。

Implementation and Limitations

为了便于使用稀疏注意力,OpenAI 开源了一组 block-sparse GPU 内核,这些内核可以高效地进行 query 和 key 矩阵的必要切片操作。

尽管取得了这些进展,研究人员指出,对于极高分辨率的图像或视频,自回归序列生成仍然是不切实际的。他们建议,这些优化的注意力操作可以作为未来多尺度建模方法的基元。

Sources