AIGC 基本功|序列并行与 Ring Attention-SP

人工智能炼丹君
2026-09-05 / 0 评论 / 10 阅读 / 正在检测是否收录...

序列并行与 Ring Attention

所属方向:分布式训练 | 难度:高阶 | 前置知识:张量并行与流水线并行、注意力显存账本、FlashAttention
关键词:序列并行、Ring Attention、长序列、通信与计算重叠、Ulysses


01. 为什么需要它

一个视频或长文被编码成 131072 个 token,hidden size 为 8192,仅一个 BF16 的 $[L,d]$ 张量就占 2 GiB。Transformer 每层不只保存一个这样的张量,还要保存归一化输入、Q/K/V、MLP 中间量、dropout 状态等;标准 attention 若显式保存 $[L,L]$ 分数矩阵,单头 BF16 就是 32 GiB,根本无法进入反向。

数据并行沿 batch 切,每个 rank 仍要处理完整序列;张量并行沿 hidden/head 切,序列长度 $L$ 仍完整复制;流水线并行沿层切,某一层的长序列峰值仍在。长上下文把瓶颈推到第四个维度:必须沿 sequence 切 token。

但“每卡拿一段 token”并不自动成立。MLP 和 LayerNorm 对每个 token 独立,本地就能算;self-attention 中,每个 query 必须看到所有 key/value。若 rank 0 只拿前 1/8 的 K/V,它算出的就成了局部 attention,数学结果已经改变。序列并行的核心问题,是怎样在不复制全部长序列、不近似 attention 的前提下交换必要信息。

目前最常见的两类答案是 Ulysses 与 Ring Attention。Ulysses 先按 sequence 分片,再用 all-to-all 把布局变成按 attention head 分片,让每个 rank 在少量 head 上看到完整 sequence;算完后再 all-to-all 变回来。Ring Attention 则让每个 rank 保留本地 query,K/V block 沿环传递,经过 $P$ 轮后本地 query 恰好见过全序列。

即使尚未阅读FlashAttention 文章,因此先补最小背景:FlashAttention 并不近似 softmax,它把 Q/K/V 分块,并用在线 softmax 的最大值、分母和加权和递推,避免把完整 $L\times L$ 注意力矩阵写回 HBM。Ring Attention 把同一个分块思想扩展到多设备:块不只从显存搬到片上 SRAM,也沿网络搬到下一个 rank。

本文预算脚本给出:$L=131072,d=8192,P=8$ 时,一个 BF16 序列张量从每卡 2 GiB 降到 0.25 GiB,理想降低 8 倍;每层每 rank 绕环接收 K/V 的单向数据约 3.5 GiB。显存不是白省的,它被换成了通信与 $P$ 轮依赖。

02. 最小可用理解

三句话建立框架:

  1. 序列并行把 $[B,L,d]$ 沿 $L$ 切成 $P$ 份,使 LayerNorm、dropout、MLP 等序列形激活每卡近似降为 $1/P$。
  2. Ulysses 用 all-to-all 在 sequence shard 与 head shard 间转置;Ring Attention 让 K/V block 绕环流动,本地 Q 用在线 softmax逐块累积精确结果。
  3. 两者没有把全注意力的 $O(L^2d)$ 算术复杂度变成线性,只是把计算与内存分散到多卡,并用通信重叠争取接近线性扩展。

如果只记一个区别:Megatron 风格的 sequence parallel 主要切 LayerNorm/dropout 等非 TP 区域的激活;Ulysses/Ring Attention 属于长上下文的 attention 并行,真正处理完整 sequence 上的全局注意力。很多系统把后者称为 context parallel,以免名称混淆。

03. 数学推导

3.1 标准 attention 的账本

令 $Q,K,V\in\mathbb{R}^{L\times d}$,为简洁省略 batch 与 head 维。scaled dot-product attention:

$$O=\operatorname{softmax}\left(\frac{QK^\top}{\sqrt d}\right)V$$

分数矩阵 $S=QK^\top/\sqrt d$ 有 $L^2$ 个元素,算力约为 $O(L^2d)$。若将序列均分到 $P$ 个 rank,每卡持有 $Q_r,K_r,V_r\in\mathbb{R}^{L/P\times d}$。逐 token 的 MLP 可以直接本地算,但 attention 的本地输出需要:

$$O_r=\operatorname{softmax}\left(\frac{Q_r[K_0;K_1;\ldots;K_{P-1}]^\top}{\sqrt d}\right)[V_0;V_1;\ldots;V_{P-1}]$$

也就是说,本地 query 仍必须与所有 rank 的 K/V 相遇。区别只在“把所有 K/V 一次聚齐”还是“逐块流过”。

3.2 在线 softmax:Ring 的数学底座

对一行 query,把 key/value 分成若干块。处理前 $t$ 个块后,保存行最大值 $m_t$、指数和 $\ell_t$ 与未归一化加权和 $u_t$。新块分数为 $s$,块内最大值为 $m_b=\max(s)$,更新全局最大值:

$$m_{t+1}=\max(m_t,m_b)$$

为了把旧累计量换到新的指数基准,定义 $\alpha=\exp(m_t-m_{t+1})$,则:

$$\ell_{t+1}=\alpha\ell_t+\sum_j\exp(s_j-m_{t+1})$$

$$u_{t+1}=\alpha u_t+\sum_j\exp(s_j-m_{t+1})v_j$$

处理完所有块后:

$$o=\frac{u_T}{\ell_T}$$

每次最大值变大时,旧分母与旧分子同时乘 $\alpha$,所以比值语义不变。算法只需保存 $m,\ell,u$ 和当前 score tile,不需保存全矩阵。它与减最大值的稳定 softmax 完全等价,而不是近似。

3.3 Ring Attention 的 $P$ 轮

第 $r$ 个 rank 固定本地 $Q_r$,初始持有 $K_r,V_r$。第 0 轮计算本地块,第 1 轮把 K/V 发给下一个 rank并接收上一个 rank 的块,如此循环。第 $t$ 轮 rank $r$ 处理来源 $(r-t)\bmod P$ 的 K/V;$P$ 轮后每个 $Q_r$ 看过所有块。

本地计算量近似:

$$C_{\mathrm{rank}}=O\left(\frac{L}{P}\cdot L\cdot d\right)=O\left(\frac{L^2d}{P}\right)$$

每 rank 需要接收 $P-1$ 次大小约 $2(L/P)d\,b$ 的 K/V,其中 $b$ 是每元素字节数:

$$V_{\mathrm{ring}}=2\frac{P-1}{P}Ldb$$

当 $L$ 与 $P$ 同比例增长、每卡本地 token 数保持不变时,单轮的 Q/KV block 大小、计算量和传输量近似不变,增长的是轮数和每卡总计算量。是否能覆盖通信取决于单轮 block 是否足够大、有效算力与链路带宽;弱扩展本身不会自动改善单轮的重叠条件。若本地 block 太小、网络慢或 kernel 没有异步双缓冲,通信会裸露出来。

训练反向不能只把前向环倒放那么简单:每个 rank 还要计算对本地 Q、流动 K/V 的梯度,并把属于原拥有者的 dK/dV 沿环累积或归还。checkpoint 与重计算还能减少保存量,但增加计算。

3.4 因果 mask 如何穿过环

自回归 attention 要求 query 位置 $i$ 只能看 key 位置 $j\le i$。分块后,根据全局位置判断:

  • K/V block 完全位于 Q block 未来:整块跳过;
  • 完全位于过去:无需 mask;
  • 与 Q block 重叠:块内施加三角 mask。

跳过未来块可省计算,但不同 rank 的有效工作量会不均。某些实现使用 zigzag 或重新排列 token,使每卡同时持有前后位置,平衡 causal attention 工作量。若只看非因果公式估算吞吐,部署到 causal LLM 时可能偏差很大。

3.5 Ulysses 的 all-to-all 转置

设输入布局是 sequence-sharded:

$$[B,L/P,H,D_h]$$

$H$ 是 head 数,$D_h$ 是每头维度。Ulysses all-to-all 把它转成:

$$[B,L,H/P,D_h]$$

每卡序列完整、head 只剩 $H/P$,于是能在本地执行普通 attention;输出再 all-to-all 回 sequence shard。它通常要求 $H$ 能被 $P$ 整除,若 GQA 只有很少 KV heads,约束会更紧。Ring 不要求按 head 数切,但要经历 $P$ 个顺序环步骤。

Ulysses 的 collective 容易利用成熟 all-to-all,但跨节点 all-to-all 对网络争用敏感;Ring 只与邻居点对点通信,更贴合环形拓扑,也更容易细粒度重叠。没有一种方法在所有硬件和 head 配置上恒优。

3.6 Megatron Sequence Parallel 不等于长上下文 attention

传统 TP 的 column/row parallel 边界往往让某些激活在 TP rank 上复制。Megatron sequence parallel 把 LayerNorm、dropout 等逐 token 操作沿 sequence 切分,并用 reduce-scatter 与 all-gather 衔接 TP 区域,使这些激活内存近似降为 $1/P$。它与 TP process group 绑定,主要减少非 TP 区域的重复激活。

但 attention 若仍按 head 做 TP,每个 head 仍可能看到完整 sequence。因此配置项 sequence_parallel=true 并不等于已经支持百万 token。需要进一步的 context parallel、Ulysses 或 Ring Attention 才能沿 attention 的上下文维扩展。

一个序列张量的每卡存储随 P 下降,而完整前向绕环接收的 K/V 字节增加并趋于上限

一个序列张量的每卡存储随 P 下降,而完整前向绕环接收的 K/V 字节增加并趋于上限。右图是接收量,不重复计发送量,尚未包含反向通信。

04. 代码实现

ring_attention_sim.py 只用标准库,实现完整 softmax attention 与分块在线递推。Q/K/V 均为 $[4,2]$,每个 K/V block 两行,相当于两轮 ring:

Q=K=V shape=(4, 2), kv_block=2, ring_steps=2
full[0]=[1.713718770, 1.717687378]
ring[0]=[1.713718770, 1.717687378]
max_abs_diff=0.000e+00

本组小样本在打印精度内相同;算法等价性来自第 3.2 节的不变式,浮点实现一般仍存在求和顺序引起的舍入差异。脚本没有真正启动多进程;它把 K/V block 依次喂给同一递推器,模拟每个 rank 收到环上块后的本地数学。生产实现还要异步 send/recv、双缓冲、causal mask 与 backward。

sp_budget.py 对 $L=131072,d=8192,H=64,P=8$ 输出:

L=131072, d=8192, heads=64, sp=8, dtype_bytes=2
one [L,d] tensor: replicated=2.00 GiB, sequence-sharded/rank=0.25 GiB
ideal memory reduction for sequence-shaped activations=8.0x
Ulysses head-divisible=True (64//8=8 heads/rank)
Ring Attention rounds=8, local query tokens=16384
Ring K/V traffic per rank per layer (one direction, ideal)=3.50 GiB

2 GiB 只是一个张量,不是整层总显存;0.25 GiB 是理想静态 shard,也没算当前通信块、输出、梯度与工作区。3.5 GiB 对应 $2(P-1)Ldb/P$,是每层、每 rank、单方向的 K/V 接收量。实际网络还包含反向与协议开销。

05. 工业级实现对照

以 2026-09 为准,DeepSpeed 的 deepspeed/sequence/layer.py 中 DistributedAttention 用自定义 all-to-all 在 scatter 轴与 gather 轴间转换 Q/K/V,调用 local_attention 后再变回去。源码明确检查总 head 数必须能被 sequence parallel size 整除,并提供 stream 与 overlap handle 管理通信重叠。

这个工业实现比最小公式多处理 batch 维位置、不同 Q/K/V 布局、rotary position embedding、异步 stream、autograd 的反向 all-to-all 与进程组。布局索引错一位可能 shape 仍合法但语义错误,所以系统必须统一采用明确的张量维约定。

Megatron-LM 的模型并行配置中 sequence_parallel 用于并行 LayerNorm 与 dropout;context_parallel_size 则面向上下文切分。当前 Megatron Core 还支持多种 context parallel 通信方式。读配置时不要只看“SP”缩写,要追到具体张量在哪个轴分片以及 attention 内部采取哪种交换。

Ring Attention 的工业 kernel 通常建立两个 K/V buffer:当前块参与计算时,下一个块在独立通信 stream 接收;计算结束交换 buffer。要真正隐藏通信,必须满足:

$$T_{\mathrm{attention\ block}}\ge T_{\mathrm{send/recv}}$$

此外还要确保计算与 NCCL 没争抢同一资源到互相拖慢。仅在代码里使用 async op 不代表 timeline 上已经重叠,必须用 profiler 验证。

FlashAttention 提供单卡块级 IO 优化,Ring/Ulysses 提供多卡数据布局。两者不是竞争关系:每个 rank 的 local attention 往往正由 FlashAttention kernel 完成。前者减少 HBM 读写,后者减少单卡需要常驻的全局序列。

checkpoint 也必须保存并行元数据。若训练时改变 SP/CP 度,参数本身或许未沿 sequence 切,但 optimizer、随机数状态和位置相关缓存可能变化。恢复前应由框架的 distributed checkpoint 层重新映射,而不是直接让每个 rank 读取旧编号文件。

06. 代价与边界

序列并行省的是序列形激活,不一定省参数和优化器状态;它通常要与 ZeRO/FSDP、TP、PP 组合。完整世界规模可能写成 $W=d_p t_p p_p c_p$,其中 $c_p$ 是 context parallel 度。每多一根轴,都新增 process group、拓扑映射与整除约束。

Ring Attention 保留全注意力的二次计算量。卡数增加能把每卡计算降到 $1/P$,但集群总 FLOPs 仍为 $O(L^2d)$;当上下文翻倍、卡数不变时,总计算约四倍。它突破的是单卡内存和 wall-clock 可扩展性,不是算法复杂度。

通信也不是“免费”。Ring 的每 rank 字节量随 $L$ 线性增长且有 $P$ 轮延迟链;Ulysses 有前后 all-to-all,容易受跨节点网络与并发任务干扰。若 $L$ 尚短,通信与布局转换可能比 attention 本身更贵,普通 TP 更合适。

常见边界包括:

  • Ulysses 的 head 或 KV head 不能被 SP 度整除,GQA/MQA 尤其容易受限;
  • causal mask 造成各 block 计算不均,需要 load-balanced ring;
  • position encoding 必须使用全局位置,不能把每 rank 的局部 token 从 0 重新编号;
  • dropout 的随机数要在切分前后保持统计与重算一致,否则并行度改变会造成难解释的差异;
  • variable-length packed sequence 需要携带边界,不能让不同样本互相 attention;
  • 视频 token 在时间、空间上的排列会影响 causal 或局部结构,切块策略不能只按连续内存方便决定。

在低带宽以太网集群上,增加 context parallel 度可能只是把 OOM 变成通信瓶颈。应先用 activation checkpoint、FlashAttention、减小 microbatch 等单机方法,确认仍由单卡序列容量限制,再引入跨设备 SP。

6.1 从单张量扩展到整层激活账本

一个 $[L,d]$ BF16 张量占 $2Ld$ 字节,但训练层内可能同时存在 residual、norm 输入、Q/K/V、attention output 与 MLP 中间值。设所有与序列线性相关且必须保存到反向的张量合计为 $cLd$ 个元素,序列分片后的理想每卡保存:

$$M_{\mathrm{linear/rank}}\approx\frac{cLdb}{P}$$

$c$ 由架构、融合 kernel 与 checkpoint 策略决定,不能用固定常数套所有模型。SwiGLU 会产生门控与 value 两个支路;视频 DiT 还可能有条件分支、cross-attention 与调制参数。最可靠方法是从 profiler 的 allocation timeline 识别跨 backward 存活的张量。

attention tile 与通信 buffer 是额外峰值。Ring 双缓冲至少需要当前和下一份 K/V block;FlashAttention kernel 还需要 score tile、softmax statistic 和 workspace。于是:

$$M_{\mathrm{peak/rank}}=M_{\mathrm{linear/rank}}+M_{\mathrm{local\ attention}}+M_{\mathrm{double\ buffer}}+M_{\mathrm{runtime}}$$

当 $P$ 继续增加时,第一项下降,但 buffer、runtime 和最小 kernel workspace 不会等比下降,显存收益逐渐偏离 $1/P$。用总 allocated 除卡数预测极限会过于乐观。

6.2 反向为何比前向更难

前向中,本地 Q 只需依次看过所有 K/V block。反向要得到 $dQ_r$、$dK_j$ 与 $dV_j$。当来源 $j$ 的 K/V 到达 rank $r$ 时,本地可计算它对 $dQ_r$ 的贡献,也会产生属于来源 rank $j$ 的 $dK_j,dV_j$ 部分。后两者必须跨所有 query shard 求和并归还拥有者。

因此每个块需要身份信息,环中既可能传 K/V,也可能传对应梯度累计。为了减少前向保存,backward 常重算局部 score 与概率,依赖前向保存的行最大值和归一化分母。若随机 dropout 参与 attention,重算还必须恢复相同随机 mask。

理论通信量必须同时列前向与反向。只用 $2(P-1)Ldb/P$ 报告 K/V 前向流量会低估训练网络负担;真实实现还包括 dK/dV、可能的 dQ 归约、控制同步和其他模块的 collective。性能分析应以 profiler 中每层总通信字节与裸露时间为准。

6.3 通信重叠需要满足哪些条件

把一个本地 query block 与一个 K/V block 的 attention 计算时间记为 $T_c$,邻居传输一个 K/V block 的时间记为 $T_n=\alpha+S/\beta$,其中 $\alpha$ 是链路延迟,$\beta$ 是有效带宽,$S$ 是消息字节。双缓冲后的理想每轮时间:

$$T_{\mathrm{round}}\approx\max(T_c,T_n)$$

只有 $T_c\ge T_n$,网络才大部分隐藏。减小 block 能降低临时显存,却增加轮数和延迟占比;增大 block 提高 GEMM 效率,却增大 buffer。block size 是计算、带宽、延迟与显存的共同旋钮。

实现还需避免隐式同步。例如在循环中把 GPU 标量取回 CPU、过早 wait 通信 handle、复用尚未完成接收的 buffer,都会把异步流水串行化。正确做法是独立 stream、event 依赖与 ping-pong buffer,并在 GPU timeline 上确认 send/recv 与 attention kernel 真正重叠。

环的物理顺序同样重要。逻辑 rank 相邻却跨交换机,会让每轮走慢链路;合理 ring 应优先沿 NVLink 域,再通过少量跨节点边连接。多维集群还可以分层:节点内用 Ulysses all-to-all,节点间用 Ring,或反过来根据硬件选择,这正是统一序列并行方法的动机。

6.4 长度、head 与并行轴的联合约束

设 context parallel 度为 $c_p$、tensor parallel 度为 $t_p$。若 TP 已把 head 切为 $H/t_p$,Ulysses 再沿 head 分 $c_p$,就需要本地可见 head 数继续可分:

$$H\ \bmod\ (t_pc_p)=0$$

具体约束取决于框架是在同一 head 轴叠加还是采用不同布局。GQA 中应将 $H_{KV}$ 单独代入;只有 8 个 KV head 的模型很难做大 head-parallel。某些系统复制 K/V 以放松整除,但通信和显存模型随之改变。

sequence length 通常也要求能被 $c_p$ 与 block size 整除。不整除时可以 padding 或不等长 shard。padding 实现简单但会浪费二次 attention 计算;不等长 shard 又使通信消息与计算不均。对长度差异大的数据,先按 token 数分桶往往比依赖动态不等长 ring 更稳定。

TP 与 CP 同时使用还会出现多种 collective。attention 投影可能在 TP 组 all-reduce,context attention 在 CP 组 send/recv,层外梯度还在 DP 组 reduce-scatter。若各组在不同 rank 上以不一致顺序发起 collective,可能死锁。框架必须给跨组通信建立确定顺序或独立 stream 依赖。

6.5 视频与多模态序列的特殊问题

视频 token 常按时间、空间高、空间宽展平。连续切 sequence 可能让每个 rank 持有若干完整帧,也可能持有同一帧的空间条带;两种布局对位置编码、局部 attention、数据增强和负载平衡影响不同。若模型有 temporal attention 与 spatial attention 分解结构,最优切分轴也可能随层变化。

全局位置必须从原始坐标生成。RoPE 若把每个 rank 的局部索引重新从 0 开始,不会报 shape 错,却让不同块位置混叠。二维或三维 RoPE 更要携带时间与空间坐标,而不是只传一个线性 offset。跨模态序列还要保持文本、图像、视频 segment 的 mask 与边界。

视频长度和分辨率经常动态变化。同一 global batch 内若 token 数差异大,padding 使最长样本决定所有 rank 工作量;packed sequence 可减少浪费,却要求 block 不跨样本做 attention。调度系统最好按总 token 数而非样本数组 batch,并把实际有效 token 纳入 loss normalization。

因果结构也不总是标准下三角。视频生成可能使用双向视觉 attention、文本到视频 cross-attention 或分块因果时间 mask。Ring 优化中“未来块可跳过”的判断必须来自真实 mask 结构,不能硬编码 LLM 假设。

6.6 正确性验收与故障定位

第一层验收是 shape:每个边界记录 global shape、local shape、分片轴、global offset、padding 与拥有者。第二层是数学:小尺寸 FP64 单卡作为 oracle,逐项比较 forward 和 dQ/dK/dV。第三层是随机性:开启 dropout 与 checkpoint 后,在固定 seed 下比较统计和可复现范围。

输出不一致时,可按轮 dump 每个 rank 处理的 K/V 来源编号。若少一个或重复一个,是 ring 调度;若所有来源齐全但数值错,检查在线 softmax 的旧累计重标定;若仅 causal 错,检查全局位置与块分类;若 forward 对而 backward 错,检查 dK/dV 归还与跨 query shard 求和。

OOM 时不要只减 block。先区分 persistent activation、通信双 buffer、kernel workspace 与碎片。若 activation 占主导,提高 CP 或 checkpoint 有效;若 buffer 占主导,减 block 才有效;若碎片主导,可调分配器或稳定 shape。错误手段可能降低吞吐却完全不改变峰值来源。

性能验收同时做 strong scaling 与 weak scaling。strong scaling 固定全局 $L$,增加卡数,观察端到端时间能否下降;容量弱扩展固定每卡 local token,令 $L\propto P$:精确全注意力的每卡计算 $L^2/P$ 此时随 $P$ 线性增长,不能期望每卡时间不变;应比较实测时间与这一增长基线,以及通信隐藏率。Ring Attention 可随设备数扩展可处理上下文,但总计算也随之增加。

最后把精度也纳入验收。在线 softmax 很稳定,但不同 block 顺序改变浮点加法;BF16、FP8 attention 更明显。应比较最终任务指标、loss 轨迹和梯度,而不是要求多卡输出 bitwise 等同。若误差远超单纯归约顺序,再排查 scale、mask 与布局。

6.7 Ulysses、Ring 与混合方案如何选

先看 head 约束。若本地 Q head 与 KV head 都能被目标 SP 度整除,且集群 all-to-all 带宽高,Ulysses 路径直接,local attention 可复用成熟 kernel。若 KV head 很少、SP 度很大,Ring 不沿 head 切,通常更容易扩展。再看拓扑:全连接高速交换更适合 all-to-all;邻接带宽强、跨组带宽弱时,环形点对点更自然。

再看本地块计算。长序列、大 head dimension 使每轮计算足以隐藏 K/V 传输,Ring 受益;较短序列或大量小 head 时,$P$ 轮延迟可能突出。Ulysses collective 次数少,但消息全局交换可能造成拥塞。决策依据应是每轮计算时间、有效链路带宽和 collective 实测,不是理论总字节一项。

混合方案把设备组织成二维网格。例如节点内 8 卡做 Ulysses,节点间若干组做 Ring;这样 head 整除只限制节点内度数,跨节点使用邻接通信。代价是布局更复杂、process group 更多,Q/K/V 需要在两个维度间保持一致顺序。只有单一方案确实受限时才值得引入。

还要把 TP 纳入考虑。如果 TP 已经切 head,Ulysses 可用 head 更少;有时降低 TP、提高 CP 更适合长序列,有时模型单层容量又要求 TP 不能降。最优并行度会随训练长度变化,短序列预训练与长序列继续训练未必使用同一布局。

6.8 推理场景与训练不同

prefill 处理整段 prompt,计算形态接近训练前向,序列并行能分摊长 prompt attention。decode 每步只有少量新 query,却要读取不断增长的 KV cache;此时瓶颈常是 KV 带宽与跨卡延迟,而非大块 GEMM。训练中优秀的 Ring block,在逐 token decode 里可能太细、同步轮数太多。

推理还涉及请求级并发。可以把不同请求分给数据并行副本,也可把一个超长请求沿 context 分片。前者吞吐高、单请求受单卡长度限制;后者支持超长上下文、占用多个设备。调度器应根据请求长度动态选择,而不是所有请求固定占一个 CP 组。

KV cache 的拥有关系必须稳定。若每 rank 保存一段 sequence,新增 token 的 K/V 放在哪一段、何时重平衡、beam search 如何复制,都要定义。Paged KV cache 与 context parallel 叠加时,逻辑页、物理 GPU 与全局位置形成三层映射,错误可能表现为偶发内容质量下降而非崩溃。

因此本文通信公式主要用于训练或 prefill,不能直接拿来预测 decode tokens/s。推理基准要分别报告 time-to-first-token、inter-token latency、并发吞吐、KV cache 容量和长短请求混合表现。

6.9 上线前检查清单

上线前确认:sequence、head、KV head 与 block 的整除或 padding 规则明确;每个 shard 保存全局位置;causal/packed mask 按全局索引生成;每轮 K/V 来源不重不漏;在线 softmax 保存并正确重标定最大值、分母与分子;backward 的 dK/dV 回到原拥有者;dropout 重算保持随机一致。

系统侧确认 ring 物理顺序符合拓扑,通信与计算 timeline 真重叠,buffer 生命周期无覆盖,所有 rank collective 顺序一致,checkpoint 记录分片元数据,变更 CP 度可以恢复。性能侧同时做 strong/weak scaling,并在真实长度分布而非单一最大长度上测试。

最后保留小尺寸 oracle。每次升级 attention kernel、通信库、RoPE 或 mask 实现,都自动生成随机 Q/K/V,与 FP64 单卡结果比较 forward 和全部梯度;再跑包含因果、变长、GQA、极端分数和非整除长度的边界集。长序列错误代价很高,小 oracle 是最便宜的保险。

6.10 最后的边界:更长不等于更有用

系统能够处理百万 token,只说明容量与运行时间可接受,不保证模型会利用远距离信息。长上下文实验还要检查位置外推、训练长度分布、检索准确率与有效注意范围。并行系统解决“算得出来”,数据与模型设计决定“学得会不会”。

因此扩展长度时同时保留短上下文质量基线,按距离分桶评估信息召回,并报告有效 token 而非 padding 后长度。否则可能用大量 GPU 计算模型并未使用的上下文。

07. 经典论文脉络

  • FlashAttention(arXiv:2205.14135):用 IO-aware tiling 与在线 softmax 实现精确 attention,避免物化完整分数矩阵,是多卡分块 attention 的本地 kernel 基础。
  • Reducing Activation Recomputation in Large Transformer Models(arXiv:2205.05198):提出 Megatron sequence parallel 与选择性重计算,减少 TP 区域之外的重复激活。
  • DeepSpeed Ulysses(arXiv:2309.14509):用 all-to-all 在 sequence 与 head 布局间转换,使超长序列 attention 可扩展。
  • Ring Attention with Blockwise Transformers(arXiv:2310.01889):让 K/V 块沿环流动并与分块计算重叠,把可处理上下文随设备数扩展。
  • USP(arXiv:2405.07719):统一 Ulysses 与 Ring 两类序列并行,针对模型结构与网络拓扑组合两者。

这条演进线从“单卡不保存 $L^2$ 矩阵”,走到“多卡不复制 $L$ 激活”,再走到按 head 结构和网络拓扑组合通信方式。算法、kernel 与系统三层缺一不可。

08. 常见误解

误解一:序列并行把 attention 复杂度从二次降成线性。 精确全注意力的集群总计算仍为 $O(L^2d)$;它只把工作分摊,并让内存近似随本地 token 数增长。

误解二:每卡切一段 token 后直接做本地 attention。 那只允许 token 看同一分片,会改变模型。必须让 query 见到全局 K/V,或明确接受局部/稀疏近似。

误解三:Ring Attention 是近似注意力。 在线 softmax 递推在精确算术下与完整 softmax 等价;差别来自浮点归约顺序,不是丢弃连接。

误解四:开启 Megatron sequence_parallel 就解决长上下文。 它主要减少 LayerNorm、dropout 等激活复制;长上下文 attention 通常还需 context parallel。

误解五:异步 send/recv 就等于通信被隐藏。 只有计算时间覆盖传输、stream 真能并发且无资源争抢,timeline 上才会重叠。

误解六:Ulysses 的 SP 度只受 GPU 数限制。 head 数、KV head 数与 tensor layout 必须可切;GQA 模型常比标准 MHA 更早碰到上限。

09. 动手验证

先运行 ring_attention_sim.py,把 block_size 改成 1、2、4。预期 max_abs_diff 始终接近机器精度;再把 score 全部加 1000,稳定递推仍有限,而直接 exp(score) 会溢出。这同时验证精确性与减最大值的重要性。

为脚本加入 causal mask,使用 query/key 的全局索引屏蔽 $j>i$。预期分块结果与完整下三角 attention 一致。然后故意用局部索引 mask,会发现后续 rank 错误允许关注未来或错误屏蔽过去,这是真实分布式实现中很隐蔽的 bug。

运行 sp_budget.py,固定每卡 token 为 16384,让 $L$ 与 ranks 同比例从 1、2、4、8 增长。每卡序列张量保持 0.25 GiB,而 ring K/V 流量与全局 $L$ 增长。这个弱扩展实验说明“可放下更长序列”并不等于“通信不增长”,只是计算增长可能覆盖通信。

在真实集群比较 FlashAttention 单卡、Ulysses、Ring/context parallel:固定模型与 global sequence,记录峰值显存、attention kernel 时间、collective/点对点时间、重叠率与 MFU。分别在节点内和跨节点运行;预期最佳方法随 NVLink、InfiniBand、head 数和 local block 大小变化。

最后做数值验收:关闭 dropout,用同一 Q/K/V 比较单卡与多卡输出及 dQ/dK/dV。BF16 下用合理容差,同时单独检查每个 rank 拼接后的全局顺序。输出对而梯度错,通常说明 backward 中流动 K/V 梯度没有正确归还拥有者。

10. 延伸阅读

读懂本文的完整前置路径是:

选 Ulysses 还是 Ring,不应从名字出发。先列出 $L,H,H_{KV},d$、网络拓扑、每卡内存与 causal 结构;再计算 local shape、通信字节、轮数和可重叠窗口;最后用 profiler 验证。长上下文系统真正的能力,是让数学等价、内存边界和物理网络三者同时对齐。

附录:完整代码

09 节用到的脚本全文如下(ring_attention_sim.py、sp_budget.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。

ring_attention_sim.py

#!/usr/bin/env python3
"""Exact blockwise attention recurrence used by ring-style attention; stdlib only."""

import math


def dot(a, b):
    return sum(x * y for x, y in zip(a, b))


def full_attention(q, k, v):
    outputs = []
    scale = math.sqrt(len(q[0]))
    for query in q:
        scores = [dot(query, key) / scale for key in k]
        peak = max(scores)
        weights = [math.exp(score - peak) for score in scores]
        denom = sum(weights)
        outputs.append([sum(w * value[j] for w, value in zip(weights, v)) / denom for j in range(len(v[0]))])
    return outputs


def blockwise_attention(q, k, v, block_size):
    scale = math.sqrt(len(q[0]))
    peak = [-math.inf] * len(q)
    denom = [0.0] * len(q)
    numer = [[0.0] * len(v[0]) for _ in q]
    for start in range(0, len(k), block_size):
        kb, vb = k[start:start + block_size], v[start:start + block_size]
        for row, query in enumerate(q):
            scores = [dot(query, key) / scale for key in kb]
            new_peak = max(peak[row], max(scores))
            correction = math.exp(peak[row] - new_peak) if peak[row] != -math.inf else 0.0
            weights = [math.exp(score - new_peak) for score in scores]
            denom[row] = correction * denom[row] + sum(weights)
            numer[row] = [correction * old + sum(w * value[j] for w, value in zip(weights, vb)) for j, old in enumerate(numer[row])]
            peak[row] = new_peak
    return [[x / denom[row] for x in numer[row]] for row in range(len(q))]


q = [[1.0, 0.0], [0.0, 1.0], [1.0, 1.0], [-1.0, 1.0]]
k = [[1.0, 0.0], [0.0, 1.0], [1.0, 1.0], [1.0, -1.0]]
v = [[1.0, 2.0], [2.0, 0.0], [0.0, 3.0], [4.0, 1.0]]
full = full_attention(q, k, v)
streamed = blockwise_attention(q, k, v, block_size=2)
max_diff = max(abs(a - b) for ra, rb in zip(full, streamed) for a, b in zip(ra, rb))
print("Q=K=V shape=(4, 2), kv_block=2, ring_steps=2")
print("full[0]=[" + ", ".join(f"{x:.9f}" for x in full[0]) + "]")
print("ring[0]=[" + ", ".join(f"{x:.9f}" for x in streamed[0]) + "]")
print(f"max_abs_diff={max_diff:.3e}")

sp_budget.py

#!/usr/bin/env python3
"""Compare idealized per-rank sequence-memory and method constraints."""

import argparse


parser = argparse.ArgumentParser()
parser.add_argument("--tokens", type=int, default=131072)
parser.add_argument("--hidden", type=int, default=8192)
parser.add_argument("--heads", type=int, default=64)
parser.add_argument("--ranks", type=int, default=8)
args = parser.parse_args()
if args.tokens % args.ranks:
    parser.error("tokens must be divisible by ranks")

tensor_gib = args.tokens * args.hidden * 2 / 1024**3
local_gib = tensor_gib / args.ranks
print(f"L={args.tokens}, d={args.hidden}, heads={args.heads}, sp={args.ranks}, dtype_bytes=2")
print(f"one [L,d] tensor: replicated={tensor_gib:.2f} GiB, sequence-sharded/rank={local_gib:.2f} GiB")
print(f"ideal memory reduction for sequence-shaped activations={args.ranks:.1f}x")
print(f"Ulysses head-divisible={args.heads % args.ranks == 0} ({args.heads}//{args.ranks}={args.heads // args.ranks} heads/rank)")
print(f"Ring Attention rounds={args.ranks}, local query tokens={args.tokens // args.ranks}")
print(f"Ring K/V traffic per rank per layer (one direction, ideal)={2*tensor_gib*(args.ranks-1)/args.ranks:.2f} GiB")

make_figures.py

"""Regenerate this article's deterministic teaching figure (numpy + matplotlib)."""
from pathlib import Path
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
OUT=Path(__file__).resolve().parents[1]/"figures"
OUT.mkdir(exist_ok=True)
plt.rcParams.update({"font.sans-serif":["PingFang SC","Arial Unicode MS","DejaVu Sans"],"axes.unicode_minus":False})
p=np.array([1,2,4,8,16,32]); L=131072;d=8192;b=2;g=1024**3
fig,axes=plt.subplots(1,2,figsize=(10,4.4))
axes[0].plot(p,np.full(p.shape,L*d*b/g),"--",label="Replicated")
axes[0].plot(p,L*d*b/p/g,"o-",label="Sequence shard")
axes[0].set(xscale="log",xlabel="Ranks P",ylabel="One tensor per rank (GiB)");axes[0].legend()
axes[1].plot(p,2*(p-1)/p*L*d*b/g,"o-")
axes[1].set(xscale="log",xlabel="Ranks P",ylabel="Forward K/V received per rank (GiB)")
for ax in axes:ax.grid(alpha=.25)
fig.suptitle("Ring attention budget: L=131072, d=8192, BF16");fig.tight_layout()
fig.savefig(OUT/'ring_budget.png',dpi=170)
plt.close(fig)
print(OUT/'ring_budget.png')


更多 AIGC 论文解读,关注微信公众号「人工智能炼丹君」

每日更新 · 论文精选 · 深度解读 · 技术脉络

微信搜索 人工智能炼丹君 或扫描下方二维码关注

扫码关注「人工智能炼丹君」

0

评论 (0)

取消
粤ICP备2021042327号