系列:推理系统基础设施手记

推理系统基础设施(二):Sequence Parallelism + Ring Attention——把 1M+ Token 训成常态

0. 一句话定位

Sequence Parallelism 把长序列在 token 维度切到多卡;Ring Attention 让多卡用「环形传递 K/V 块 + 在线 softmax 累加」算出一个数学上完全等价于全量 attention 的结果,每张卡的激活内存与序列长度无关。

这是突破单卡序列长度天花板的关键手段,让 1M+ Token 的训练/推理成为常规操作。

1. 视野速览

2. 核心做法:在线 softmax + 环形 K/V 通信

假设 8 卡环,序列按 token 切成 8 段,每段长度 L/N。每张卡持自己的 query Q_i(不动),让 K/V 块沿环流动:

m_i, l_i, o_i = -inf, 0, 0     # 运行时 max / 累加和 / 输出

for step in range(N):           # 沿环走 N 步,每步拿到一个 K/V 块
    K_block, V_block = recv_from_prev()   # 来自上一张卡
    send_to_next(my_KV)                    # 同时把自己的块传给下一张卡

    s_ij = Q_i @ K_block^T / sqrt(d)               # (L/N, B)
    m_new = max(m_i, s_ij.max(-1))                 # 更新 running max
    p_ij  = exp(s_ij - m_new[:, None])
    l_new = exp(m_i - m_new) * l_i + p_ij.sum(-1)
    o_i   = exp(m_i - m_new)[:, None] * o_i + p_ij @ V_block
    m_i, l_i = m_new, l_new

O_i = o_i / l_i[:, None]    # 归一化,结果与全量 attention 完全一致

符号:Q_i = 第 i 卡本地 query;K_block, V_block = 沿环流动的 K/V 块;m_i, l_i, o_i = 在线 softmax 的运行 max / 累加和 / 部分输出。

3. 环通信结构图

Ring Attention:环形传递 K/V 块 + 在线 softmax 累加 GPU 0 Q₀(本地,不动) K/V 块ⱼ GPU 1 Q₁(本地,不动) K/V 块ⱼ GPU 2 Q₂(本地,不动) K/V 块ⱼ GPU 3 Q₃(本地,不动) K/V 块ⱼ K/V 块沿环流动(每步传一个块) GPU3 → GPU0 闭合环 在线 softmax 累加(每步): s_ij = Q_i·K_blockᵀ / √d ; m_new = max(m, s.max) ; p = exp(s − m_new) l_new = exp(m − m_new)·l + p.sum ; o_new = exp(m − m_new)·o + p·V_block → 走完 N 步 O_i = o_i / l_i 每张卡只缓存「自己的 Q + 1 个流入 K/V 块」→ 激活内存与总序列长度无关;8 卡可跑 1M token,理论扩到 100 卡 = 100M token。

图:Ring Attention 的环通信。序列按 token 切到 8 卡,K/V 块沿环逐步传递,每卡用在线 softmax 把流入的 K/V 块累加起来;走完 N 步,结果与全量 attention 数学等价。

4. 与朴素方案对比

方案激活内存能否跑 1M token备注
朴素(全量 Q×Kᵀ)O(L²)❌ 单卡必爆1M token 单卡显存直接炸
Megatron-LM SP(all-gather)临时全量 KV⚠️ 几万 tokenattention 内每卡仍要临时全量 KV
Ring Attention(本方案)与 L 无关✅ 8 卡 1M,100 卡 100M通信被 matmul 掩盖,墙钟几乎不增

通信开销被 matmul 掩盖(每 hop 通信正好盖住上一次 matmul),所以墙钟延迟几乎不增加。

5. 实测收益与代表工作

6. 具身智能关联

  1. 视频世界模型(V-JEPA 2 / Genie)需要 1M+ token 的视频片段输入,Ring Attention 是 V-JEPA 2 多机训练时的标配;
  2. 机器人 VLA 长操作日志 + 视频历史规划场景,SP+Ring 让「全操作日志一起算」成为可能;
  3. 与 HCA/CSA 协同——HCA 把 KV 再压一档,让 Ring Attention 在环里传的是压缩后的「目录块」,1M token → 8000 目录块的环通信,进一步把带宽压力降到 1/128(详见本系列 sys1 的 NUMA 链路与 fa4 的 HCA)。

7. 学习建议 / 常见坑

  1. 必须用在线 softmax + running max,否则累计误差会让结果偏;
  2. 因果 mask 下用 Zig-Zag 切块,否则后卡空转;
  3. 通信带宽是真正瓶颈——NVLink / IB 不可省(普通以太网一旦通信时间 > 计算时间,「边算边传」就退化成「等数据」);
  4. 与张量并行混用时,TP×SP 通信方向要正交,否则 all-gather 和环通信会撞车。

觉得有用?欢迎点赞、收藏,或请作者喝咖啡 ☕️

支付宝收款码

支付宝

微信收款码

微信

💬 留言

评论由 Giscus 驱动(基于 GitHub Discussions)。 当前仓库 NaphJohn/LLM-blog 尚未启用 Discussions:请在 GitHub 仓库 Settings → General → Features 勾选 Discussions 后刷新本页,评论区即自动显示。