Back to blog

Ring Attention with Blockwise Transformers for Near-Infinite Context

用于近无限上下文的环形注意力与分块 Transformer

Ring Attention with Blockwise Transformers for Near-Infinite Context

一、论文概述

项目内容
标题Ring Attention with Blockwise Transformers for Near-Infinite Context
作者Hao Liu, Matei Zaharia, Pieter Abbeel
机构UC Berkeley
论文arXiv:2310.01889
代码llm_large_context
发布2023年10月
许可-

二、核心思想

问题定义

Transformer 的自注意力机制具有二次方的内存成本,限制了其处理长序列的能力:

  • 内存瓶颈:处理 1 亿 token 需要超过 1000GB 内存(隐藏维度 1024)
  • 硬件限制:现代 GPU/TPU 通常只有不到 100GB HBM
  • 存储需求:每层输出需要存储,因为自注意力需要所有元素的交互

解决方案概述

Ring Attention 提出了一种分布式长序列处理方法:

  • 分块计算:利用分块注意力和前馈网络的计算
  • 环形通信:设备形成环形拓扑,KV 块在环中传递
  • 通信-计算重叠:KV 块的通信与注意力计算完全重叠
  • 线性扩展:上下文长度随设备数线性扩展

三、技术架构

整体框架图

Ring Attention 设计

核心公式

自注意力

Attention(Q,K,V)=softmax(QK⊤d)V\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^\top}{\sqrt{d}}\right)V

其中 Q,K,V∈Rs×dQ, K, V \in \mathbb{R}^{s \times d},ss 是序列长度,dd 是头维度。

前馈网络

FFN(x)=max⁡(0,xW1+b1)W2+b2\text{FFN}(x) = \max(0, xW_1 + b_1)W_2 + b_2

分块注意力

将序列分成 NN 个块,每个设备持有一个块:

Q=[Q0,Q1,...,QN−1],K=[K0,K1,...,KN−1],V=[V0,V1,...,VN−1]Q = [Q_0, Q_1, ..., Q_{N-1}], \quad K = [K_0, K_1, ..., K_{N-1}], \quad V = [V_0, V_1, ..., V_{N-1}]

设备 ii 持有 Qi,Ki,ViQ_i, K_i, V_i。

环形通信

在环形拓扑中,设备 ii 在计算注意力时:

  1. 使用本地 QiQ_i 和当前持有的 Kj,VjK_j, V_j 计算注意力
  2. 将 Kj,VjK_j, V_j 发送给下一个设备
  3. 从上一个设备接收 Kj−1,Vj−1K_{j-1}, V_{j-1}

环形注意力算法

for i = 0 to N-1 do  // 外循环:迭代次数
    for each device j in parallel do  // 并行执行
        k = (j - i) mod N
        // 计算 Q_j 与 K_k, V_k 的注意力
        Out_j += BlockwiseAttention(Q_j, K_k, V_k)
        // 发送 K_k, V_k 到下一个设备
        send K_k, V_k to device (j+1) mod N
        // 从上一个设备接收 K_{k-1}, V_{k-1}
        receive K_{k-1}, V_{k-1} from device (j-1) mod N
    end for
end for

内存分析

方法每层激活内存最大序列长度
标准 TransformerO(s2)O(s^2)受限于单设备内存
分块 Transformer (BPT)2bsh2bsh受限于单设备内存
Ring Attention2bsh/N2bsh/NNN 倍于单设备

其中 bb 是批量大小,ss 是序列长度,hh 是隐藏维度,NN 是设备数。

通信-计算重叠

关键条件:块计算时间 > 块传输时间

重叠策略:

  1. 设备在计算注意力时,同时发送当前 KV 块
  2. 设备在计算注意力时,同时接收下一个 KV 块
  3. 只要计算时间大于传输时间,通信开销为零

排列不变性

性质:自注意力对 KV 块的顺序具有排列不变性

Attention(Qi,[K0,K1,...,KN−1],[V0,V1,...,VN−1])=⨁j=0N−1Attention(Qi,Kj,Vj)\text{Attention}(Q_i, [K_0, K_1, ..., K_{N-1}], [V_0, V_1, ..., V_{N-1}]) = \bigoplus_{j=0}^{N-1} \text{Attention}(Q_i, K_j, V_j)

其中 ⨁\bigoplus 表示使用 lazy softmax 策略的累积。

四、核心创新

创新点说明理论/实验依据
环形注意力KV 块在环形拓扑中传递通信与计算完全重叠
分块计算利用分块注意力和前馈网络内存成本线性化
零通信开销通信-计算重叠只要计算时间 > 传输时间
线性扩展上下文长度随设备数线性扩展理论证明
精确算法不牺牲注意力计算精度排列不变性保证

五、实验结果

实验设置

配置说明
硬件TPUv4-1024
模型7B-65B 参数
序列长度最高 100M+ tokens
基线标准 Transformer, BPT

最大上下文长度

最大上下文长度

方法最大上下文长度
标准 Transformer~100K
BPT~100K
Ring Attention~100M+ (1024 设备)

结论:Ring Attention 支持比基线长 500 倍以上的序列长度。

模型 FLOPS 利用率 (MFU)

MFU 趋势

模型上下文长度MFU
7B4M~60%
13B4M~65%
65B4M~70%

结论:Ring Attention 在大模型和长序列上保持高 MFU,开销可忽略。

长程检索任务

上下文准确率

任务:长程检索(在长序列中检索特定信息)

结果:

  • 随着上下文长度增加,准确率提升
  • Ring Attention 能够利用更长的上下文信息

训练 FLOPS 成本

上下文长度相对于 4K 的 FLOPS 成本比
32K~8x
128K~32x
512K~128x
2M~512x

六、相关工作

长序列处理方法

方法关键特性局限性
FlashAttention分块计算,IO 感知受限于单设备内存
BPT分块注意力和前馈受限于单设备内存
Ring Attention环形通信,通信-计算重叠需要高速互连
Striped Attention条纹分区优化仅适用于因果注意力

序列并行方法

方法通信操作通信开销可扩展性
Megatron-LM SPAllGather线性增长受限
DeepSpeed UlyssesAll-to-All恒定好
Ring AttentionRing P2P零(重叠)最佳

七、总结

核心贡献

  1. 环形注意力:利用环形拓扑实现通信-计算重叠
  2. 分块计算:结合 BPT 实现内存线性化
  3. 零通信开销:通信完全被计算掩盖
  4. 近无限上下文:上下文长度随设备数线性扩展
  5. 大规模验证:在 TPUv4-1024 上验证百万级 token 序列

技术影响

  • 长序列训练:使百万级 token 训练成为可能
  • 分布式注意力:成为分布式注意力的标准方法
  • 广泛应用:被众多长序列模型和框架采用
  • 研究基础:启发了 Striped Attention 等后续工作

局限性

  • 硬件依赖:需要高速互连(NVLink/InfiniBand)
  • 工作负载不均衡:因果注意力中存在不均衡(由 Striped Attention 解决)
  • 通信假设:假设块计算时间 > 块传输时间
  • 实现复杂性:需要精心实现通信-计算重叠

八、参考资源