AI++
训练方法·硬核·10 分钟阅读·AI++ 编辑部

Ring Attention:让百万级上下文训练成为可能的环形并行

Ring Attention 把超长序列切到多张 GPU 上,每张卡只持有一段 Q/K/V,按环形拓扑轮流交换 K/V 块。 关键是让「算当前块的 attention」和「传下一块 K/V」重叠,把通信藏到计算后面。 转一圈后每张卡都见过全部 K/V,于是上下文长度能随 GPU 数线性扩展,支撑百万 token 训练。

一句话结论

Ring Attention 把超长序列的注意力计算切成「每张卡持有一段、按环形轮流传递 K/V 块、计算与通信重叠」的流水线,让上下文长度能随 GPU 数线性扩展——百万级 token 的训练在工程上变得可行。

背景问题

长上下文是大模型能力的核心方向之一(128k、1M 甚至更长),但训练长上下文模型有个硬约束:单卡装不下完整序列的 attention 中间结果

attention 的中间矩阵 $QK^\top$ 是 $L \times L$ 的。$L=128\text{k}$ 时,单层 attention 的中间矩阵就是几十 GB,单卡显存根本放不下。Flash Attention 解决了「不实例化 $L \times L$ 矩阵」的问题,但前提是 Q、K、V 都在一张卡上。当 $L$ 大到单卡连 Q、K、V 本身(哪怕分块)都装不下,或者单卡算不过来时,就必须把序列拆到多张卡上

朴素做法是「张量并行」:把 Q、K、V 按头切到多卡,每卡算自己的头,最后 all-reduce。但这对长序列有个问题——每张卡仍要看到完整序列的 K、V(因为每个头要 attend 到所有位置),通信量巨大且不解决「单卡装不下完整序列」的问题。

序列并行(sequence parallelism)的思路是把序列维度切开:每张卡只持有一段 token 的 Q、K、V。但这样算 attention 时,每张卡需要其他卡持有的 K、V——怎么高效拿到?Ring Attention 给了一个优雅的答案。

核心思路

设 $N$ 张 GPU 排成一个逻辑环(0→1→2→…→N-1→0)。把长度 $L$ 的序列切成 $N$ 段,第 $i$ 张卡持有第 $i$ 段的 $Q_i, K_i, V_i$。

目标:每张卡最终算出「自己这段 query 对全部 key/value 的 attention 输出」。

Ring Attention 的流程:

  1. 初始:卡 $i$ 持有 $Q_i, K_i, V_i$。
  2. 第一轮计算:卡 $i$ 用本地 $Q_i$ 和本地 $K_i, V_i$ 算一段 partial attention(blockwise,配合 Flash Attention 的 online softmax)。
  3. 同时:卡 $i$ 把自己的 $K_i, V_i$ 发给下一张卡 $i+1$,并从上一张卡 $i-1$ 接收 $K_{i-1}, V_{i-1}$。
  4. 重叠:发送/接收与计算并行——算当前块 attention 的同时,下一块的 K/V 正在传输。
  5. 下一轮:卡 $i$ 用 $Q_i$ 和刚收到的 $K_{i-1}, V_{i-1}$ 算 partial attention,同时继续环形传递。
  6. 转 $N-1$ 圈后:每张卡都依次见过全部 $N$ 段 K/V,partial attention 累积成完整结果。

关键:通信被计算藏起来了。如果「算一块 attention」的时间和「传一块 K/V」的时间相当,通信开销几乎为零(完全重叠)。这正是 Ring Attention 高效的根源。

关键技术点

1. Blockwise Parallel Transformer 是基础

Ring Attention 的计算单元是「blockwise attention」——不是把整段 K/V 一次算完,而是切成更小的 block 增量更新 softmax(类似 Flash Attention 的 online softmax)。这让「算一块」的粒度可以和「传一块」的粒度匹配,才能做到真正的计算-通信重叠。

2. 通信-计算重叠是核心红利

不加重叠的序列并行(如朴素 all-to-all)会让通信成为瓶颈——$L$ 越长、卡越多,通信量越大。Ring Attention 把通信摊到 $N$ 步、每步和计算并行,通信时间被「藏」进计算时间,前提是单块计算量 ≥ 单块通信量。这在 attention(计算密集)上通常成立,所以效果显著。

3. 与 Flash Attention 深度耦合

Ring Attention 的内核就是「跨卡的 Flash Attention」。每张卡收到一段 K/V 后,用 Flash Attention 的 online softmax 把它累积进当前结果。所以 Ring Attention 的实现严重依赖 Flash Attention 的分块接口——这也是为什么 FA1/FA2 普及后 Ring Attention 才真正好用。

4. 负载均衡与块大小调优

  • 块大小影响通信/计算比:太大→计算时间长但通信次数少;太小→通信启动开销占比大。需根据网络带宽和算力调。
  • 卡数扩展:理论上线性——$N$ 张卡能支持 $N$ 倍长的上下文。但环上传递延迟随 $N$ 线性增加,$N$ 很大时延迟可能侵蚀重叠红利。
  • 拓扑匹配:物理拓扑最好和逻辑环一致(如 NVLink 全互联或环),跨节点时网络带宽是瓶颈。

5. 变体:Ulysses 与 Context Parallelism

Ring Attention 不是唯一的序列并行方案:

  • Ulysses:用 all-to-all 把「按头切」换成「按序列切」,通信是 all-to-all(一次性)而非环形(多步)。在头数多、卡数适中时更高效;卡数远多于头数时退化。
  • Megatron Context Parallelism:工业化的 Ring Attention 变体,配合 TP/PP 一起调度,是 Llama-3 等长上下文训练的实战方案。

它们各有适用区间,Ring Attention 在「卡数多、序列极长」时仍是主流选择。

效果与局限

效果:

  • 原论文与后续工业实践:上下文长度随 GPU 数近线性扩展,百万级 token 训练在数百卡规模上可行。
  • 通信开销在计算密集的 attention 上被有效隐藏,扩展效率高。
  • 与 Flash Attention、ZeRO、TP/PP 等可组合,是长上下文训练栈的关键一环。
  • 已被主流框架支持(Megatron-LM、DeepSpeed、PyTorch FSDP 的 CP 等)。

局限:

  • 只解决 attention 的序列并行。FFN 层、嵌入层、loss 计算等仍需其他并行策略配合,完整的长上下文训练是组合方案。
  • 对网络敏感。跨节点环的带宽若不如 NVLink,通信可能藏不住,扩展效率下降。
  • 小模型 / 短序列无收益。序列不够长时,单卡 + Flash Attention 就够,Ring Attention 的通信开销反而拖累。
  • 负载不均:如果序列内某些位置计算量不同(如带 padding 或混合长度),环上各卡负载可能不均。
  • 与 PagedAttention 等推理优化关系不大:Ring Attention 主要是训练侧技术,推理时长上下文仍靠 MLA、KV cache 压缩等。

谁该关心

  • 做长上下文训练的团队:Ring Attention / Context Parallelism 是百万 token 训练的基建,理解它的通信-计算重叠是调优分布式训练的前提。
  • 大模型 infra 工程师:它是 TP/PP/DP/CP 这套并行组合里 CP 的代表,掌握它才能设计完整的 4D 并行。
  • 研究分布式系统的人:Ring Attention 是「计算-通信重叠」范式的经典案例,对设计其他环形/流水线并行有启发。
  • 追求超长上下文能力的模型团队:上下文长度是模型能力的重要维度,Ring Attention 决定了你能训多长。

一个判断:Ring Attention 不会单独存在——它永远是「Flash Attention + 序列并行 + TP/PP/DP」这套组合的一部分。它的价值在于把「序列长度」这个维度从单卡约束解放出来,让长上下文从「实验室 demo」变成「可工程化训练」。对绝大多数应用,128k 上下文已经够用,Ring Attention 是 infra 团队的必修课而非业务团队的关注点;但一旦要做 1M+ 上下文,它就是绕不开的核心。

长上下文序列并行通信优化分布式训练