从单卡到集群:为什么需要把序列切开
训练一个能处理百万级 token 上下文的 Transformer 模型,首先遇到的不是算法问题,而是单张 GPU 的显存墙。以 Llama 2-70B 为例,在 GQA、FP16、batch=1、80 层、8 个 KV 头、head_dim=128 的配置下,缓存 4K token 的 KV 约需 1.25 GiB,缓存 128K token 则膨胀到约 40 GiB。当序列长度达到百万级时,仅 KV 缓存就超过 300 GiB,远超单卡 HBM 容量。
标准自注意力的计算量与序列长度 N 的平方成正比,处理 128K token 的注意力计算量是 4K 的 1024 倍。即使 FlashAttention 通过分块计算避免了 O(N²) 的 HBM 读写,纯计算量本身仍然是 O(N²d)。当 N 达到百万级,一个完整稠密注意力层已是数十 PFLOPs 量级,单张 GPU 无法在合理时间内完成。
直观的应对方案是模型并行:把权重切到多卡上。但模型并行解决的是参数量过大问题,序列长度带来的 KV 缓存和计算量依然落在每张卡上。另一种思路是序列并行:把一条长序列切成多段,每张卡只负责一段的 Q、K、V,然后通过通信让每张卡看到完整的 KV。这正是 Ring Attention 的出发点。
在线 softmax:为什么分块计算不会破坏注意力结果
理解 Ring Attention 之前,必须先理解 FlashAttention 的在线 softmax 机制。标准 softmax 需要先计算所有分数的最大值和指数和,这意味着必须拿到完整的一行分数才能归一化。但在分块计算中,每个块只拿到部分分数,如何保证最终结果与全局计算一致?
在线 softmax 的直觉是:每处理一个新块,就用新块的最大值修正之前累积的指数和,再增量更新输出。具体来说,维护三个标量:当前全局最大值 m、累积指数和 l、以及加权输出 o。当新块到来时,先计算新块内的最大值 m_new,如果 m_new 大于 m,就把之前的 l 和 o 按 exp(m - m_new) 缩放,再累加新块的贡献。这一过程在数学上等价于一次性计算完整 softmax,但允许逐块推进。
这一性质是 Ring Attention 能够跨设备工作的基础。每个设备只需要本地 Q 与当前持有的 KV 块计算局部注意力,然后用在线 softmax 逐步聚合来自其他设备的 KV 块结果。只要最终遍历过所有 KV 块,聚合结果就与全局注意力完全一致,不引入任何近似。
环形通信与计算重叠:一次完整的注意力步骤
Ring Attention 把序列长度 N 均匀分配到 P 个设备上,每个设备持有 N/P 个 token 的 Q、K、V。所有设备组成一个逻辑环。计算过程分为 P 步,每一步中,每个设备用本地 Q 与当前持有的 K、V 块计算局部注意力,同时将 K、V 块沿环发送给下一个设备,并异步接收上一个设备发来的新 K、V 块。
下面用一个贯穿全文的场景来具体说明。假设我们要训练一个处理 128K token 上下文的模型,使用 8 张 GPU,每张卡负责 16K token。初始状态下,设备 0 持有序列第 0~16K 的 Q₀、K₀、V₀,设备 1 持有第 16K~32K 的 Q₁、K₁、V₁,依此类推。
第一步:设备 0 用 Q₀ 与本地 K₀、V₀ 计算局部注意力,同时将 K₀、V₀ 发送给设备 1,并接收设备 7 发来的 K₇、V₇。第二步:设备 0 用 Q₀ 与刚收到的 K₇、V₇ 计算注意力,同时将 K₇、V₇ 转发给设备 1,并接收设备 6 发来的 K₆、V₆。如此循环,经过 8 步后,设备 0 的 Q₀ 已经与所有位置的 K、V 交互过,得到最终输出 O₀。其他设备同理。
flowchart TD
A[设备 0: Q0, K0, V0] --> B[Step 1: 计算 Q0 与 K0,V0 的局部注意力]
B --> C[发送 K0,V0 至设备 1]
C --> D[接收设备 7 的 K7,V7]
D --> E[Step 2: 计算 Q0 与 K7,V7 的局部注意力]
E --> F[转发 K7,V7 至设备 1]
F --> G[接收设备 6 的 K6,V6]
G --> H[...]
H --> I[Step 8: 计算 Q0 与 K1,V1 的局部注意力]
I --> J[输出 O0]
图中的流程只展示了设备 0 的视角。实际上所有设备并行执行相同的步骤,每个设备都在用自己的 Q 与不断轮转的 KV 块计算注意力。通信与计算的重叠发生在每一步内部:当设备正在计算当前 KV 块的注意力时,下一个 KV 块正在网络上传输。只要单块计算时间大于传输时间,通信开销就被完全隐藏。在典型的大模型训练中,注意力计算是计算密集的,这一条件通常成立。
与序列并行、Ulysses 的对比:何时选择环形通信
Ring Attention 并非唯一的序列并行方案。Megatron-LM 的上下文并行(Context Parallelism)本质与 Ring Attention 相同,都基于 FlashAttention 的分块计算和 KV 环传。DeepSpeed Ulysses 则采用了完全不同的通信模式:它在序列维度分片计算 Q、K、V 后,通过 All-to-All 通信将数据重组为按注意力头维度分片,每个设备负责全部序列的部分头,计算完成后再通过 All-to-All 还原回序列分片。
下表从通信量、通信模式、对注意力头的约束和适用场景四个维度对比这三种方案。
| 方案 | 通信量(单层前向) | 通信模式 | 对注意力头的要求 | 适用场景 |
|---|---|---|---|---|
| Ring Attention / CP | O(N) 的 P2P 传输,P 步 | 环形 Send/Recv,可与计算重叠 | 无特殊要求 | 设备数较少、计算可掩盖通信时 |
| DeepSpeed Ulysses | O(N/P) 的 All-to-All,两步 | All-to-All,负载均衡但不可与计算重叠 | 注意力头数必须被序列并行度整除 | 设备数较多、通信带宽高时 |
| 朴素序列并行(无 FlashAttention) | O(N) 的 AllGather,两步 | AllGather + ReduceScatter,不可重叠 | 无特殊要求 | 已过时,被 CP/Ulysses 取代 |
Ulysses 的通信量更小,但 All-to-All 操作通常无法与计算有效重叠,且要求注意力头数整除并行度,限制了灵活性。Ring Attention 的通信量随设备数线性增长,但通过计算重叠可以隐藏延迟,且对模型结构无额外约束。在实际部署中,选择往往取决于集群拓扑:节点内 NVLink 带宽高、延迟低,适合 Ring Attention 的 P2P 通信;跨节点 InfiniBand 带宽相对有限,Ulysses 的 All-to-All 可能更高效。
负载不均衡:因果注意力下的计算浪费
Ring Attention 在双向注意力(如 BERT)中工作良好,因为每个 Q 都需要与所有 KV 交互,每步计算量均匀。但在自回归语言模型中,因果注意力掩码使得每个 Q 只能看到自身及之前的 KV。这导致 Ring Attention 的每一步中,部分设备在计算注定被掩码掉的位置,造成计算浪费和负载不均。
回到 8 卡 128K 的场景。设备 0 持有序列最开头的 16K token,它的 Q₀ 只需要与 K₀、V₀ 交互,后续所有 KV 块都在未来位置,会被因果掩码完全屏蔽。设备 7 持有序列末尾的 16K token,它的 Q₇ 需要与所有 8 个 KV 块交互。在 Ring Attention 的 8 步中,设备 0 只在第一步做有效计算,其余 7 步都是无效的;设备 7 则 8 步全有效。整个环的延迟由最慢的设备决定,即设备 7,而其他设备的大量计算被浪费。
Striped Attention 针对这一问题提出了改进:将序列按交错方式分片,而不是连续分块。例如,设备 0 持有 token 0、8、16、…,设备 1 持有 token 1、9、17、…,以此类推。这样每个设备持有的 Q 分散在整个序列中,每一步都有有效计算和无效计算,负载趋于均衡。代价是实现复杂度增加,且需要更细粒度的通信调度。
通信拓扑的再思考:从环到树
Ring Attention 的通信步数与设备数 P 成正比。当 P 很大时(例如跨节点百卡级),即使单步延迟很小,累积的通信延迟也可能成为瓶颈。Tree Attention 从自注意力的数学形式出发,提出了一种对数级通信步数的方案。
Tree Attention 观察到,标准注意力可以表示为某个能量函数的梯度,而该能量函数中的 logsumexp 操作满足结合律。这意味着多个设备上的局部注意力结果可以通过树形归约(类似 AllReduce)聚合,通信步数降为 O(log P)。在实现上,Tree Attention 先让每个设备本地计算自己分片的注意力分子和分母(经过数值稳定化),然后通过树形通信聚合这些部分结果,最后各设备独立完成归一化。
这一方案在跨节点场景下优势明显。现代 GPU 集群通常具有两层带宽结构:节点内 NVLink 带宽高达 900 GB/s,节点间 InfiniBand 带宽通常在 50~400 GB/s。Tree Attention 可以感知拓扑,在节点内使用高带宽环形归约,在节点间使用树形归约减少跨节点通信量。实验显示,在百卡规模下,Tree Attention 的解码延迟相比 Ring Attention 可降低数倍。
但 Tree Attention 的工程实现更复杂,需要高效的 Allreduce 原语支持,且对数值稳定性要求更高。目前 Ring Attention 仍然是工业界更成熟的选择,已被用于训练百万 token 上下文的模型。
部署边界与未解决的问题
Ring Attention 将上下文长度扩展到了设备数倍的单卡极限,但它并不改变注意力计算本身的平方复杂度。训练一个百万 token 上下文的模型,即使使用 256 张 GPU,每张卡仍需处理 4096 token 的注意力计算,总计算量依然巨大。实际部署中,往往需要结合其他优化:
- FlashAttention 在单设备内减少 HBM 读写,是 Ring Attention 的基础。
- 渐进式长度扩展:先用短序列预训练,再用长序列微调,避免从头训练的巨大成本。
- 位置编码外推:RoPE 配合 NTK 缩放或 YaRN,将训练长度高效扩展到数倍。
在可观测性方面,需要关注每步的计算时间与通信时间的比值。如果计算时间显著大于通信时间,说明重叠有效;如果两者接近,可能需要调整分块大小或通信调度。NVIDIA Nsight Systems 等工具可以可视化 kernel 执行与通信操作的 timeline,帮助定位瓶颈。
Ring Attention 的失效模式主要出现在以下情况:
- 小模型或短序列:计算量不足以掩盖通信延迟,环传开销成为主导。
- 因果注意力下的负载不均:Striped Attention 可以缓解,但实现复杂,且可能引入额外的通信开销。
- 跨节点带宽不足:当节点间带宽远低于节点内时,Ring Attention 的 P2P 通信可能阻塞,此时 Ulysses 或 Tree Attention 可能更优。
长上下文训练的真正挑战不仅在于让模型“看到”更多 token,更在于让模型有效利用这些信息。Lost in the Middle 现象表明,即使模型能够处理 128K 上下文,当关键信息位于中间位置时,检索准确率可能大幅下降。Ring Attention 解决了“能不能训练”的问题,但“训练出来的模型好不好用”还需要数据策略、位置编码和训练范式的协同演进。