AI 技术AI 推理 · 11/13
#Ring Attention#Context Parallelism#长上下文#分布式训练#FlashAttention#序列并行

AI 推理系列(十一):Ring Attention 如何让多张 GPU 共同计算一条超长序列

围绕百万 Token 长上下文的单卡容量瓶颈,解释 Ring Attention 如何沿序列维度切分 Q/K/V,让 KV Block 在多张 GPU 间环形流动并与本地计算重叠,同时分析 online softmax、Context Parallelism、Ulysses、GQA 与跨节点网络带宽之间的工程权衡。

把模型的上下文窗口从 8K 扩到 128K,首先会遇到位置编码和模型训练分布的问题;继续向百万 Token 推进时,另一个限制会迅速变得更直接:单条样本自身已经大到一张 GPU 难以容纳。即使模型参数可以通过 Tensor Parallelism 分散,单条长序列产生的激活、Attention 中间状态以及训练时需要保留的反向传播信息,仍会沿序列长度持续增长。

假设一条 128K Token 的训练样本被分到 4 张 GPU。最自然的做法是每张卡保存连续 32K Token。Linear、LayerNorm、MLP 等逐 Token 操作很容易局部完成;Attention 却不同。GPU 0 上的 Query 不只需要查看自己的 32K Token,它还必须与其他设备持有的 Key/Value 建立完整注意力关系。Ring Attention 解决的正是这个矛盾:Query 留在本地,KV Block 依次在设备之间流动,让每张卡最终都看到全局上下文,而不必同时保存全局 KV。

把序列切开以后,Attention 为什么不能只算本地块

标准 Full Attention 的核心关系可以写成 softmax(QKᵀ)V。如果序列被按 Token 维度分成四段,那么 GPU 0 持有 Q0/K0/V0,GPU 1 持有 Q1/K1/V1,以此类推。对 GPU 0 来说,仅计算 Q0 × K0ᵀ 只覆盖了局部上下文;真正的完整结果还包含 Q0K1K2K3 的关系。

最直接的方案是先 All-Gather,让每张 GPU 都得到完整 K/V,再执行本地 Q 对全局 KV 的 Attention。数学上没有问题,但显存压力又回来了:每张卡原本只保存四分之一序列,现在却需要额外容纳完整 KV。序列越长,这种复制越难接受。

Ring Attention 换了一个问题表述:与其问“怎样把全局 KV 一次复制到本地”,不如问“能不能让本地 Query 分阶段依次看到不同 KV Block”。只要每个 KV Block 最终都访问一次,并且 softmax 的全局归一化被正确维护,最终结果仍然可以与完整 Attention 对齐。

因此,Ring Attention 并没有把 Full Attention 改成只看邻居的 Sparse Attention。它改变的是数据出现的位置和时间,不是 Query 可以访问哪些 Token 的语义范围。

KV 在 Ring 中如何流动

继续使用 4 张 GPU 的例子。开始时,每张 GPU 都持有自己的 Q 和 KV:

GPU0: Q0 + KV0
GPU1: Q1 + KV1
GPU2: Q2 + KV2
GPU3: Q3 + KV3

第一轮,各设备处理本地组合:

GPU0: Q0 × KV0
GPU1: Q1 × KV1
GPU2: Q2 × KV2
GPU3: Q3 × KV3

随后 KV 沿逻辑环发送给下一张 GPU。GPU0 接收 GPU3 的 KV,GPU1 接收 GPU0 的 KV。下一轮保持 Query 不动,只替换当前正在处理的 KV Block。经过与设备数量对应的若干轮后,每个 Query Partition 都已经遍历全部 KV Partition。

flowchart LR
  A[GPU0 保留 Q0] --> B[计算 Q0 与 KV0]
  B --> C[接收 GPU3 的 KV]
  C --> D[计算 Q0 与 KV3]
  D --> E[继续接收下一块 KV]
  E --> F[遍历全部 KV Block]
  F --> G[完成 Q0 的全局 Attention]
  H[当前 KV Block] --> I[发送给下一张 GPU]
  I --> J[通信与当前计算重叠]
  J --> E

其他 GPU 同时执行同样的循环。与“把完整 KV 广播给每张卡”相比,任意时刻每个设备通常只需要持有本地状态以及当前通信中的少量 KV Block,长序列状态因此被分散到整个设备组。

对于因果语言模型,还必须遵守 causal mask。某个 Query 不能读取未来 Token,因此不同 Q/KV Block 组合中的有效区域不同。高效实现可以利用块之间的相对位置跳过完全无效的计算,但无论是否做这类优化,因果语义都不能因为分片而改变。

最容易被忽略的细节:不同 Block 的 Attention 结果不能直接相加

如果把一个 Query 对不同 KV Block 分别做 softmax,然后把各块输出直接求和,结果并不等于对完整 KV 一次做 softmax。原因是 softmax 的分母依赖所有可见 Key:每个块单独归一化,会把原本应该互相竞争的 Attention Score 分割成多个独立概率分布。

因此 Blockwise Attention 必须维护跨块的归一化统计量。直觉上,可以在处理每个新 KV Block 时更新当前行的最大值、指数和以及输出累积值。当发现更大的 score 时,旧累积结果需要按新的尺度重新缩放;随后再把当前块贡献合并进去。这样不需要物化完整 Attention Score 矩阵,也能得到与一次性 softmax 等价的结果。

这与 FlashAttention 使用 online softmax 的思想高度契合。FlashAttention 在单卡内部通过分块避免把巨大 QKᵀ 中间矩阵写回 HBM;Ring Attention 则让这些块跨设备移动。二者都依赖“Attention 可以按块流式累积,但必须保持数学上的全局归一化”这一性质。

所以“Ring Attention 就是 KV 转一圈”只描述了通信外观。真正保证结果正确的,是环形通信、分块 Attention 和稳定 softmax 累积共同组成的执行协议。

通信很多,为什么仍然可能划算

Ring 当然没有消除网络流量。每一轮设备都需要发送当前 KV Block。若 GPU 先停止计算、等待网络传输完成,再开始下一轮,跨节点环境下很容易让计算单元空转。

Ring Attention 的关键工程目标是让通信和计算重叠:当前 KV Block 正在参与 Attention kernel 时,下一个 Block 已经通过通信流开始传输。当前计算结束后,下一块最好已经就绪。理想情况下,单步执行时间更接近:

max(当前块 Attention 计算时间, 下一块 KV 通信时间)

而不是两者简单相加。

这种重叠是否有效,取决于计算与网络的比例。如果 Block 太小,单块计算很快,网络和 kernel 启动开销会变得突出;如果 Block 太大,显存需求增加,同时一次通信的粒度也变大。跨节点时还要面对网络拓扑差异:同一节点的 NVLink/NVSwitch 与跨节点 InfiniBand/RDMA 不属于同一延迟和带宽层级。

因此,设备数量增加并不会免费扩大上下文。更多设备能进一步分摊序列状态,但也意味着 Ring 步数、通信路径和拓扑协调更加重要。长上下文并行真正优化的是“在容量允许的前提下,把不可避免的通信尽量藏在不可避免的计算后面”。

Context Parallelism 把序列长度变成新的并行维度

传统大模型训练已经同时使用多种并行维度。Data Parallelism 把不同训练样本分给不同副本;Tensor Parallelism 拆分单层矩阵;Pipeline Parallelism 把不同层放在不同阶段。超长上下文增加了一个新的问题:单条样本本身也需要拆分。

Context Parallelism 通常就是沿序列维度切分网络输入和相关激活。对逐 Token 独立的算子,本地 Partition 可以直接计算;到了 Attention,需要通过设备间通信让局部 Q 获得全局上下文。Ring 是实现这种 KV 交换的一种重要方式。

并行方式主要切分对象典型收益主要新增成本
Data Parallelism不同样本扩大总训练吞吐梯度同步
Tensor Parallelism单层参数/张量单层模型可跨卡执行高频集合通信
Pipeline ParallelismTransformer 层容纳更深更大的模型Pipeline Bubble 与激活传递
Context Parallelism单条序列的 Token单样本可扩到更长上下文Attention 的跨设备 KV 通信

在实际训练集群中,CP 通常不是单独存在,而是与 DP、TP、PP 共同组成多维并行拓扑。此时“哪几张 GPU 组成 CP Group”会直接影响通信是否落在高速链路上。若把 Ring 的高频传输跨越低带宽网络层级,理论上的显存扩展可能被通信时间抵消。

这也是为什么 Context Parallelism 的优化逐渐从“算法上能不能切”走向“怎样按物理拓扑切”。集群中并非所有 GPU 两两等价,拓扑必须成为调度和并行配置的一部分。

Ring 与 Ulysses:同样切序列,通信方式不同

DeepSpeed-Ulysses 代表另一种经典序列并行思路。它开始时同样让不同 GPU 持有不同 Token Partition,但在进入 Attention 前通过 All-to-All 重新排列 Q/K/V,使每张设备获得完整序列范围上的一部分 Attention Head。于是局部设备可以独立处理自己负责的 Head,结束后再通过一次通信恢复序列布局。

两种路线没有简单的绝对优劣,适配条件取决于模型 Head 数、设备数量和互联结构。

路线Attention 阶段的数据移动方式设备本地主要持有什么典型约束
Ring Attention / Ring CPKV Block 按轮次点对点流动本地 Q + 当前 KV Block需要良好的计算通信重叠,Ring 拓扑敏感
UlyssesQ/K/V 经 All-to-All 重排完整序列的一部分 Head并行度受 Head 划分和集合通信效率影响
简单 All-Gather KV每张卡收集完整 KV本地 Q + 全局 KV实现直接,但长序列显存复制成本高

工程系统还可能混合多种策略。例如设备规模较大时,一部分维度使用 Ulysses,一部分维度使用 Ring;节点内与节点间采用不同通信策略。目标不是坚持某个算法名称,而是在目标序列长度、GPU 数量、Head 配置和网络拓扑下减少不可隐藏的通信时间。

GQA 为什么会顺带降低 Context Parallelism 的通信压力

前面介绍 GQA 时,重点是减少 KV Head 数,从而降低 KV Cache 和 Decode 阶段的访存压力。在 Context Parallelism 中,这个结构还有另一层价值:Ring 或其他 CP 通信传输的核心对象正是 K/V,因此 KV Head 越少,单个 Token 需要交换的 KV 数据通常也越少。

这说明模型结构和分布式系统并不是两套独立优化。GQA 最初看起来是 Attention 架构选择,但它会改变通信张量的尺寸;FlashAttention 看起来是单卡 kernel 优化,但 block 计算速度决定通信是否容易被隐藏;RoPE Scaling 看起来只涉及位置表示,却决定模型是否值得处理如此长的输入。

对于百万 Token 训练,一个更完整的技术栈往往是:位置编码先保证模型能够表示目标长度;高效 Attention kernel 降低单卡块计算的 IO 开销;Context Parallelism 解决单条序列状态无法放入单卡的问题;GQA/MQA 等结构继续减少需要保存和传输的 KV;更上层再与 TP、PP、DP 组合扩展整个模型和训练吞吐。

所以不能把 Ring Attention 单独理解成“长上下文开关”。它只解决长上下文系统中的一个明确边界:单样本序列维度如何跨设备分布并完成完整 Attention。

它扩展的是容量,不会消除 Full Attention 的计算量

Ring Attention 最容易被误解成一种“超长上下文加速算法”。更准确地说,它首先是一种容量与并行化方案。把长度为 N 的序列切到更多设备,可以降低每张卡持有的序列状态和部分激活,但 Full Attention 的总关系计算并不会因此从二次复杂度突然变成线性复杂度。

设备增加后,总工作被分散到更多 GPU,墙钟时间可能因并行而下降,单卡内存也会降低;但整个集群完成的乘加总量仍与完整 Attention 的语义范围相关。上下文继续扩大时,计算成本依然快速增长。到了这个阶段,若业务真正需要更便宜的百万 Token Attention,还可能需要窗口化、稀疏化、检索、状态压缩或其他改变计算范围的路线。

这也是“能训练百万 Token”和“百万 Token 训练成本合理”之间的差别。Ring Attention 可以把原本单卡无法执行的问题变成集群可执行问题,却不会替你回答这个训练任务是否值得消耗相应的集群算力。

原始工作以及后续长序列系统展示的极长 Context 实验,应理解为扩展能力和系统可行性的证明。实际生产训练需要根据目标任务、有效长程依赖、硬件成本和训练时间来确定是否值得把 Context 推到同样规模。

跨节点后,网络会成为必须单独观测的瓶颈

在单机 8 卡环境中,设备间可能拥有高带宽互联,Ring 的点对点交换相对容易隐藏。扩展到多机后,KV Block 会进入跨节点网络,最慢链路可能决定整个 CP Group 的步进速度。此时只观察 GPU utilization 很难定位问题,因为 GPU 空闲可能只是等待远端 KV。

生产训练至少需要同时观测三类信号。第一类是 Attention kernel 时间,包括当前块实际计算时间以及不同序列长度下的利用率。第二类是通信时间和 overlap 比例,确认发送/接收是否真正与计算并发,而不是在时间线上首尾相接。第三类是拓扑和流量,包括节点内与跨节点带宽、不同 Rank 的等待差异以及是否出现单链路热点。

当通信时间持续高于块计算时间时,Ring 的隐藏条件已经破坏。可以考虑调整 CP Group 的物理映射、块大小、通信并发方式,或者减少 KV 通信量。若问题来自跨节点带宽本身,继续提高单卡 Attention kernel FLOPS 未必能改善端到端训练步时间。

相反,如果通信长期能够完全隐藏,而 Attention kernel 占据主要时间,则说明系统更接近计算受限。此时 FlashAttention 类 kernel、低精度计算和更适合硬件的分块策略可能比进一步优化网络更有价值。

什么时候值得引入 Context Parallelism

是否启用 CP,首先应看单条样本长度,而不是集群有多少 GPU。如果训练数据主要在 2K~8K Token,单卡本来可以轻松容纳,并且 Attention 不是显存主瓶颈,那么引入跨设备序列通信只会增加额外复杂度。

当单条样本已经让激活或 Attention 状态成为单卡限制,CP 才真正解决不可替代的问题。例如长代码仓库、长视频产生的大量视觉 Token、超长 Agent 轨迹或其他自然长序列都可能进入这个区域。此时应该先确定目标 Context 是否真的需要 Full Attention,再决定 CP 规模。

上线训练配置前,可以按以下顺序验证:先在较短序列上确认分布式结果与非 CP 基线在数值误差允许范围内一致;随后逐步增加 Context,观察每卡显存是否近似按预期下降;再通过 profiler 检查 Attention 和 P2P/集合通信时间线;最后在目标节点规模下测试真实网络拓扑,而不是只用单机结果外推多机性能。

如果序列能放下但训练步时间随 CP 数量增加反而恶化,需要先查通信是否暴露在关键路径上。如果数值结果异常,应检查 causal mask、分块边界、softmax 跨块累积以及反向传播的数据依赖,而不是先把问题归因于网络。

长上下文并行正在从“能切开”走向“按拓扑切得更合理”

当设备规模进一步增加,单一逻辑 Ring 很难代表真实集群。节点内可能是高速 NVLink/NVSwitch,节点间则通过 InfiniBand;不同层级的最优消息粒度和通信模式并不相同。近期研究继续探索层次化 Sequence Parallelism,本质上是在承认一个事实:分布式 Attention 的算法应该理解物理网络层级,而不是把所有 Rank 当成等价点。

与此同时,单卡 Attention kernel 也持续变化。新的 GPU 架构会改变片上存储、矩阵计算和异步数据搬运能力,因此“计算一块 Attention 需要多久”这个值本身也在变化。若 kernel 变得更快,而网络没有同步提升,过去能够被隐藏的通信可能重新暴露;反过来,通信硬件升级也会改变最合理的分块和 CP 规模。

因此 Ring Attention 最值得保留的不是某个固定拓扑参数,而是一种系统设计方法:把一条超长序列拆成局部状态,让全局 Attention 依靠可流式组合的块计算恢复,再让通信尽可能与这些块计算重叠。

在 AI 推理系列前面的技术中,RoPE 解决“位置能否表示到那么远”,FlashAttention 优化“单张 GPU 内怎样高效计算 Attention”,GQA 降低“每个 Token 需要多少 KV 状态”。当这些技术仍无法让一条超长序列塞进单卡时,Ring Attention / Context Parallelism 才接过下一棒:把上下文本身变成一种分布式并行维度,让整个 GPU 集群共同完成同一条序列的 Attention。真正的工程边界则始终没有变化——能分开保存只是第一步,能在通信成本可控的情况下把它重新算成一个完整、正确且可扩展的 Attention,才是这类系统真正困难的地方。

资料来源

  1. Ring Attention with Blockwise Transformers for Near-Infinite Context
  2. Context Parallel Package — Megatron Core
  3. DeepSpeed Ulysses: System Optimizations for Enabling Training of Extreme Long Sequence Transformer Models
  4. HSAP: A Hierachical Sequence-aware Parallelism for Hybrid-Context Generative Models
  5. FlashAttention-4: Algorithm and Kernel Pipelining Co-Design for Asymmetric Hardware Scaling
  6. 14.7 长上下文技术:从理论到工程实践 | 大模型原理与架构 | LLM Internals
  7. 深度学习的分布式训练与集合通信(三)-技术干货-昇腾社区
  8. 【论文分享】| 序列并行视角下的各类研究