理解 BigBird 的块稀疏注意力机制
BigBird 解决了标准 Transformer 模型 $O(n^2)$ 的时间和内存复杂度问题,使其能够处理长达 4096 个 token 的序列。通过将全注意力机制替换为块稀疏注意力机制,BigBird 在长文档摘要和长上下文问答等长序列任务上取得了最先进的结果。
块稀疏注意力机制
BigBird 的块稀疏注意力是对 BERT 全注意力机制的一种近似,旨在提高效率而非追求绝对优越性。它通过结合三种不同类型的注意力来减少查询 token 需要关注的 token 数量:
- 滑动注意力 (Sliding Attention):token 关注其紧邻的邻居。这捕捉了局部依赖关系,因为序列中的单词通常高度依赖于相邻的 token。
- 全局 token (Global Tokens):一小部分 token 会关注序列中的每一个其他 token,并且也会被序列中的每一个其他 token 关注。这使得信息比仅靠滑动注意力更快速地在序列中传播。
- 随机 token (Random Tokens):随机选择少量 token 来关注其他 token,从而进一步缩短信息在序列图中远程节点之间传输的距离。
虽然全注意力允许信息在单层中在任意两个 token 之间传输,但块稀疏注意力可能需要多个层才能让信息在序列中传播。全局和随机连接的结合确保了信息只需通过几层即可快速传播。
技术实现
为了在 GPU 和 TPU 上高效实现这一点,BigBird 使用了基于块的方法。序列被划分为大小为 $b$ 的块。
按块进行注意力计算
- 第一块和最后一块:第一块 ($q_1$) 和最后一块 ($q_n$) 对序列中的所有其他 token 执行正常的注意力操作。
- 中间块:对于块 $q_{3:n-2}$ 中的 token,模型收集全局、滑动和随机的 key,并仅针对这些选定的 key 计算注意力。
- 边界块:块 $q_2$ 和 $q_{n-1}$ 收集特定子集的 key(包括第一块、最后一块和附近的滑动块)以维持稀疏结构。
复杂度对比
随着序列长度的增加,BigBird 显著减轻了计算负担。与 BERT 的二次方缩放相比,BigBird 呈线性缩放:
| 注意力类型 | 序列长度 | 时间与内存复杂度 |
|---|---|---|
| 原始全注意力 (BERT) | 512 | $T$ |
| 原始全注意力 (BERT) | 1024 | $4 \times T$ |
| 原始全注意力 (BERT) | 4096 | $64 \times T$ |
| 块稀疏注意力 (BigBird) | 1024 | $2 \times T$ |
| 块稀疏注意力 (BigBird) | 4096 | $8 \times T$ |
训练策略:ITC vs. ETC
BigBird 可以使用两种不同的策略进行训练:
- 内部 Transformer 构建 (ITC):使用少量全局 token。它需要的计算量较少,且对于许多任务来说已经足够。
- 扩展 Transformer 构建 (ETC):添加额外的全局 token。这对于问答任务特别有用,例如在问答任务中,整个问题可能需要被上下文 token 全局关注。
与 Hugging Face Transformers 的集成
BigBird 可以通过 BigBirdModel 在 transformers 库中使用。关键实现细节包括:
- 序列长度要求:序列长度必须是
block_size的倍数。Hugging Face 会自动将序列填充 (padding) 到 block size 的最小倍数。 - 自动回退机制:如果序列长度太短,不足以支持所需的全局、随机和滑动 token,或者如果模型被用作解码器 (
BigBirdForCausalLM),库会自动将attention_type切换为original_full。 - 推荐用法:作者建议对于短于 1024 个 token 的序列,设置
attention_type="original_full"。
可用的检查点 (checkpoints) 包括 bigbird-roberta-base、bigbird-roberta-large 和 bigbird-base-trivia-itc。