论文题目:Generating Long Sequences with Sparse Transformers(用稀疏变换生成长序列)
预印链接地址:arXiv:1904.10509(2019)
摘要:变压器是一种功能强大的序列模型,但需要的时间和内存随序列长度呈二次增长。在本文中,我们引入了注意力矩阵的稀疏分解,将其减少到O(n√n)。我们还介绍了a)结构和初始化的变化以训练更深的网络,b)注意力矩阵的重新计算以节省内存,以及c)用于训练的快速注意力核。我们称具有这些变化的网络为稀疏变换,并表明它们可以使用数百层对数万个时间步长的序列进行建模。我们使用相同的架构从原始字节对图像、音频和文本进行建模,为Enwik8、CIFAR10和ImageNet-64的密度建模设置了新的技术状态。我们生成的无条件样本显示出全局一致性和巨大的多样性,并表明原则上可以使用自注意力机制来模拟长度为一百万或更多的序列。

Fig.1:来自ImageNet 64和古典音乐数据集的神经自回归模型的无条件样本。对于音频、图像和文本,我们使用了相同的基于自注意力机制的架构。上面的示例是在softmax温度1.0下生成的,长度为12,288和65,536。音频样本可在https://openai.com/blog/sparse-transformer上收听
Sparse Transformer:突破序列长度限制的创新架构
引言
在深度学习领域,Transformer架构凭借其强大的序列建模能力已经成为自然语言处理的主流选择。然而,标准Transformer有一个致命弱点:注意力机制的计算复杂度随序列长度呈二次方增长。这意味着当我们想要处理长文档、高分辨率图像或原始音频时,所需的内存和计算资源会迅速爆炸。
2019年,OpenAI的研究团队发表了一篇题为《Generating Long Sequences with Sparse Transformers》的论文,提出了一种创新的解决方案:Sparse Transformer(稀疏Transformer)。这项工作不仅在理论上将复杂度从O(n²)降低到O(n√n),还在图像、文本和音频等多个领域刷新了当时的最优记录。
一、问题背景:Transformer的瓶颈
1.1 自注意力机制的代价
让我们先回顾一下标准Transformer的工作原理。在自注意力层中,序列中的每个位置都需要:
- 与所有其他位置计算注意力权重
- 对所有位置的值进行加权求和
这意味着对于长度为n的序列,我们需要计算n×n的注意力矩阵。具体来说:
内存消耗 ∝ n²
计算时间 ∝ n²
当序列长度从1,000增加到10,000时,内存需求和计算时间会增加100倍!这使得Transformer难以应用于:
- 长文档:数万词的文章
- 高分辨率图像:将256×256图像展平后有65,536个像素
- 原始音频:12kHz采样率下,1秒音频就有12,000个时间步
1.2 现有方法的局限
在Sparse Transformer之前,研究者们尝试过多种方法:
- CNN架构(如PixelCNN):需要很深的网络才能扩大感受野
- 分层结构:增加了架构复杂性
- 固定长度表示:丢失了长距离信息
这些方法要么性能受限,要么只能应用于特定领域。
二、核心创新:稀疏化注意力
Sparse Transformer的核心思想是:不需要每个位置都关注所有其他位置,通过精心设计的稀疏模式,可以在几步内实现全局连接。
2.1 注意力模式分析
研究团队首先对标准Transformer学到的注意力模式进行了可视化分析。他们在CIFAR-10上训练了一个128层的网络,发现了一些有趣的现象:

- 局部连接(图2a):许多早期层学习到类似卷积的局部模式
- 行列分解(图2b):第19和20层自动学会了将注意力分解为行注意力和列注意力
- 全局模式(图2c):一些层展现出全局的、数据依赖的访问模式
- 高度稀疏(图2d):64-128层的典型层表现出高稀疏性
这些观察表明:注意力矩阵本身就倾向于稀疏,因此引入结构化稀疏性是可行的。
2.2 稀疏注意力的形式化定义
标准的全注意力定义为:
Si = {j : j ≤ i} (位置i可以关注所有之前的位置)
稀疏注意力将其分解为p个头,每个头只关注一个子集:
Si = A_i^(m) 其中 |A_i^(m)| ∝ √n
关键约束是:保持连接性 - 任意两个位置之间应该能在p步内通过某条路径连接。
2.3 两种稀疏模式

论文提出了两种主要的2D稀疏分解方案(p=2):
跨步注意力(Strided Attention)
- 第一个头:关注前l个位置(局部窗口)
- 第二个头:每隔l个位置关注一次(跨步访问)
数学定义:
A_i^(1) = {t, t+1, ..., i} 其中 t = max(0, i-l)
A_i^(2) = {j : (i-j) mod l = 0}
优势:对于有周期性结构的数据(如图像、某些音乐)效果很好 适用场景:CIFAR-10、ImageNet、音频
固定注意力(Fixed Attention)
- 第一个头:关注同一块内的位置
- 第二个头:关注特定的"汇总单元"
数学定义:
A_i^(1) = {j : ⌊j/l⌋ = ⌊i/l⌋}
A_i^(2) = {j : j mod l ∈ {t, t+1, ..., l}} 其中 t = l-c
优势:对于没有明显周期结构的数据(如文本)效果更好 适用场景:Enwik8文本数据集
参数说明:
- l(stride):通常选择128或256,接近√n
- c:对于固定模式,通常选择8、16或32
三、架构增强:让深度网络可训练
除了稀疏注意力,论文还引入了多项技术改进:
3.1 改进的残差结构

采用pre-activation设计:
H_k = H_{k-1} + resblock(H_{k-1})
resblock(H) = attention(norm(H)) + feedforward(norm(H + attention(norm(H))))
关键改进:
- 在注意力和前馈网络之前进行层归一化
- 每个残差块直接接收来自输出的梯度
- 权重初始化按1/√(2N)缩放,N为层数
这使得训练128层甚至更深的网络成为可能。
3.2 梯度检查点(Gradient Checkpointing)
对于自注意力层,内存使用与序列长度的关系是:
内存 = O(n² × 层数)
通过在反向传播时重新计算注意力权重,可以将内存需求降低到:
内存 = O(n × 层数)
权衡:增加约33%的计算时间,但能将可训练序列长度从4,096提升到16,384。
3.3 高效的稀疏内核
论文实现了自定义GPU内核来高效计算稀疏注意力:
- 局部窗口:直接计算
- 跨步访问:通过矩阵转置转换为局部窗口
- 固定位置:聚合成块计算
- 优化:融合softmax,从不计算上三角部分
3.4 混合精度训练
使用FP16进行前向和反向传播,但以FP32存储权重。配合:
- 动态损失缩放
- 采样时将queries和keys转换为FP32(防止溢出)
四、实验结果:多领域的突破
4.1 CIFAR-10图像生成
配置:
- 128层,2个注意力头
- 维度d=256
- 跨步模式,stride=48
- 59M参数
结果:
- 2.80 bits per dim(之前最优:2.85)
- 超越了PixelSNAIL和Image Transformer

4.2 Enwik8文本建模
配置:
- 30层,8个注意力头
- 维度d=512
- 固定模式,stride=128, c=32
- 上下文长度:12,288 tokens
- 95M参数
结果:
- 0.99 bits per byte
- 与277M参数的Transformer-XL持平
- 上下文长度是标准Transformer的3倍
长距离依赖验证:

随着评估时最小上下文长度的增加,性能单调提升:
- 6,144 tokens → 0.9952 bpb
- 12,160 tokens → 0.9908 bpb
这证明了网络确实在有效利用长距离依赖。
4.3 ImageNet 64×64生成
配置:
- 48层,16个注意力头
- 维度d=512
- 跨步模式,stride=128
- 152M参数
- 训练:64个V100 GPU,7天
结果:
- 3.44 bits per dim(之前最优:3.52)
- 超越了SPN和Parallel Multiscale

生成的图像展现出:
- 全局一致性:物体结构完整
- 多样性:各种场景、物体、风格
- 高质量:没有明显的稀疏模式伪影
4.4 原始音频生成
测试了模型对极长序列的扩展能力:

| 序列长度 | 参数量 | Bits per byte |
|---|---|---|
| 65,536 (~5秒) | 152M | 1.97 |
| 262,144 (~22秒) | 25M | 2.17 |
| 1,048,576 (~87秒) | 3M | 2.99 |
关键发现:
- 序列长度增加4倍,模型容量需减少约8倍
- 原则上可以处理百万级时间步的序列
- 5秒样本展现出全局连贯性(节奏、音色、演奏风格)
4.5 稀疏模式的效率对比

在Enwik8(12,288上下文)上:
- Dense Attention: 1.00 bpb, 1.31秒/迭代
- Fixed Sparse: 0.99 bpb, 0.55秒/迭代(快2.4倍,性能更好)
- Strided Sparse: 1.13 bpb, 0.35秒/迭代
在CIFAR-10(3,072上下文)上:
- Dense Attention: 2.82 bpb, 0.54秒/迭代
- Fixed Sparse: 2.85 bpb, 0.47秒/迭代
- Strided Sparse: 2.80 bpb, 0.38秒/迭代(快1.4倍,性能最好)
重要发现:稀疏模式不仅更快,在某些情况下性能还更好!这可能表明:
- 稀疏性提供了有用的归纳偏置
- 密集注意力可能存在优化问题
五、实现细节和训练技巧
5.1 位置编码策略
根据数据类型选择不同的位置编码:
图像(有天然2D结构):
- 使用数据嵌入
- n_emb = 3(行、列、通道)
文本和音频(无明显2D结构):
- 使用注意力嵌入
- n_emb = 2(对应stride矩阵的行列索引)
5.2 优化超参数
- 优化器:Adam
- Warmup:5000步线性warmup
- 梯度裁剪:1.0
- 权重衰减:0.01
- 学习率调度:余弦退火
- Dropout:根据数据集调整(0.01-0.40)
5.3 多头注意力的使用
三种整合方式:
-
交替使用:每个残差块使用一种模式
attention(X) = W_p · attend(X, A^(r mod p)) -
合并头:单个头关注两种模式的并集
attention(X) = W_p · attend(X, ∪ A^(m)) -
多头并行(最常用):
attention(X) = W_p · [attend(X, A^(i))]_{i=1}^{n_h}
六、关键洞察和启示
6.1 稀疏性的两面性
论文发现:
- 固定模式更适合文本:因为文本没有周期性结构
- 跨步模式更适合图像/音频:因为这些数据有空间或时间规律
这表明:最优的稀疏模式应该匹配数据的固有结构。
6.2 模型容量与序列长度的权衡
从音频实验可见:
序列长度 × 模型容量^4 ≈ 常数
这意味着:
- 要处理更长序列,必须接受更小的模型
- 但即使是3M参数的模型,也能学到一定的结构
6.3 深度的重要性
128层网络的成功表明:
- 通过适当的架构设计,可以训练极深的Transformer
- Pre-activation和梯度检查点是关键技术
- 深度可能比宽度更重要(对于某些任务)
七、局限性和未来方向
7.1 当前的局限
- 预定义模式:不像图2d那样数据依赖
- 内存仍然是瓶颈:虽有改进,但百万级序列仍需小模型
- 模式选择:需要针对数据类型手动选择
7.2 后续影响
Sparse Transformer开启了高效注意力机制的研究浪潮:
- Longformer(2020):滑动窗口+全局token
- BigBird(2020):随机+窗口+全局的混合模式
- Performer(2021):基于核方法的线性注意力
- Flash Attention(2022):IO优化的精确注意力
八、总结
Sparse Transformer通过以下创新,成功将Transformer扩展到数万甚至百万时间步:
- 理论创新:稀疏分解将复杂度从O(n²)降至O(n√n)
- 架构创新:改进的残差结构和初始化支持极深网络
- 工程创新:梯度检查点、自定义内核、混合精度训练
核心贡献:
- ✅ 在CIFAR-10、Enwik8、ImageNet 64刷新SOTA
- ✅ 证明可处理12,288 tokens的文本(之前:4,096)
- ✅ 证明可处理65,536时间步的音频并保持质量
- ✅ 原则上可扩展到百万级序列
最重要的启示:
注意力不需要密集 - 精心设计的稀疏模式可以在保持性能的同时大幅降低成本。
这项工作不仅解决了Transformer的一个核心瓶颈,也为后续的高效注意力机制研究奠定了基础。在大模型时代,这些技术对于处理长上下文、降低推理成本具有重要意义。
Sparse Transformers:用稀疏变换生成长序列&spm=1001.2101.3001.5002&articleId=156885539&d=1&t=3&u=4742294d9bb84121bfa37ab9408beda3)
1742

被折叠的 条评论
为什么被折叠?



