Skip to content

长序列训练与上下文并行 ​

标签
AI/infra/并行策略
AI/infra/长上下文
字数
1946 字
阅读时间
8 分钟

序列长度从 2K 涨到 128K 甚至 1M 之后,出现了一类前面所有并行维度都解决不了的问题。

这一篇讲清「序列长度墙」的成因,以及三种沿序列切分的方案。

序列长度墙 ​

Attention 的计算量与激活显存都是 O(s2) —— s 翻倍,开销翻 4 倍。

维度能不能解决「单条序列太长」
DP不能 —— 每条序列仍要完整算一遍
ZeRO不能 —— 切的是状态冗余
TP不能 —— 切的是权重,s 方向的激活没变
PP不能 —— 切的是层,序列维没动
CP能 —— 沿序列维切分 Attention 计算本身

「s 翻倍开销翻 4 倍」与显存线性增长的差别,就是这条墙的来源:显存不够可以加卡,但 O(s2) 意味着加卡的速度追不上序列增长的速度。

序列长度增长时两条曲线的斜率不同:

  相对倍数
   16× ┤                              ●  Attention 开销:O(s²)
       │                            ╱
       │                          ╱
    4× ┤                  ●    ╱
       │                ╱   ●  ← s 翻倍 ⇒ 开销翻 4 倍
       │              ╱  ╱
    2× ┤         ●  ╱  ●      ○  显存 / 卡数只能线性增长
       │       ╱ ○
    1× ┤   ● ○
       └──┬─────┬─────┬──────┬──────▶ 序列长度
         2K    4K    8K    16K

  ⇒ 加卡的速度追不上序列增长的速度 —— 这就是「墙」
  ⇒ 而且 DP / ZeRO / TP / PP 四个维度都解不了它(它们都不动序列维)
     只有 CP 沿序列维切 Attention 计算本身

前置:FlashAttention ​

CP 的三种方案几乎都构建在 FlashAttention 之上 —— 因为它已经把 O(s2) 的显存降到了 O(s),剩下的问题是如何把这个 O(s) 的激活与计算沿序列维再切到多卡。

分块 + Online Softmax 的机制见 08-FlashAttention。它是 CP 的前提:如果还要存完整 s×s 矩阵,沿序列切分也无法避免每卡存部分 s×s。

三种方案 ​

Ring Attention ​

沿序列维把 Q/K/V 切到多卡,然后环形传递 KV 分片:

每卡持有自己那段 Q
  → 环形接收其他卡的 K/V 分片
  → 逐块算局部 Attention
  → 用 Online Softmax 合并
  → 继续传给下一张卡

关键设计是计算与 KV 传输重叠 —— 传下一块 K/V 的同时算当前块,把通信藏进计算里。

因果掩码带来的负载不均

加了 causal mask 后,靠后的卡要算的块更多、靠前的卡更少 —— 环形传递的前提「每卡算的块数相同」被打破了。

优化思路是 Zigzag / Striped 切分:不把序列连续切段,而是让每张卡交错持有前后两段,使各卡的因果计算量均衡。

通信原语限制
点对点环传 KV因果掩码下负载不均

Ring Attention:每卡持有自己那段 Q,KV 分片在环上转

    卡 0 ──KV──▶ 卡 1 ──KV──▶ 卡 2 ──KV──▶ 卡 3 ──┐
     ▲                                            │
     └────────────────────────────────────────────┘

  每一轮的四个动作:
    收到上一卡的 KV 分片
      → 与自己这段 Q 算局部 Attention
      → 用 Online Softmax 把结果合并进累积状态
      → 把 KV 分片传给下一卡(与计算重叠)

  ⇒ 转 N−1 轮后每卡都见过全部 KV,得到完整结果

  因果掩码下的负载不均:靠后的卡要算的块更多
     优化:Zigzag / Striped 切分 —— 让每卡交错持有前后两段

DeepSpeed-Ulysses ​

思路完全不同:先沿 Head 维切,再用 All-to-All 重排成序列维切分。

沿 Head 切(每卡拿部分 Head)
  → All-to-All 重排
  → 变成沿序列切(每卡拿部分序列)
  → 算 Attention
  → All-to-All 换回来

通信量的性质是它的优点:

项与什么相关
通信量与序列长度无关,与 Head 维度相关

限制:Head 数必须 ≥ CP 度。Head 数不够就切不动 —— 这是它的硬边界。

Ulysses:先沿 Head 切,再用 All-to-All 换成沿序列切

     切分前(沿 Head 切)                  All-to-All 重排
     ┌──────────────────────┐            ┌──────────────────────┐
     │ 卡 0:Head 0、1       │            │ 卡 0:序列段 0        │
     │ 卡 1:Head 2、3       │  ────────▶ │ 卡 1:序列段 1        │
     │ 卡 2:Head 4、5       │            │ 卡 2:序列段 2        │
     │ 卡 3:Head 6、7       │            │ 卡 3:序列段 3        │
     └──────────────────────┘            └──────────────────────┘
              算 Attention …                    然后在这里算 Attention
     ┌──────────────────────┐  ◀────────  ┌──────────────────────┐
     │ 再 All-to-All 换回来  │            │ 结果按 Head 归位      │
     └──────────────────────┘            └──────────────────────┘

  ⇒ 通信量与序列长度无关(只与 Head 维度相关)—— 这是它的优点
  ⇒ 硬边界:Head 数必须 ≥ CP 度,否则切不动

Megatron Context Parallel ​

沿序列维切分,配合 FlashAttention 分块,与 TP/SP 协同(在 Megatron 里可以同时开 TP + CP)。

项说明
切分维度序列
通信原语点对点 / AllGather
限制与 Megatron 生态绑定
通信开销与序列长度相关

三种方案对比 ​

方案切分维度通信原语限制适用场景
Ring Attention序列点对点环传 KV因果掩码负载不均超长序列、KV 大
UlyssesHead → 序列All-to-AllHead 数需 ≥ CP 度Head 多的模型
Megatron CP序列点对点 / AllGather与 Megatron 绑定Megatron 生态

两条判据:

  1. 看 Head 数 —— Head 数少(或 GQA 下 KV Head 少)时 Ulysses 的 All-to-All 方案受限,该选 Ring 或 Megatron CP
  2. 看通信模式与硬件的匹配 —— Ring 是点对点环传(对拓扑敏感度低),Ulysses 是 All-to-All(需要全互联的带宽,见 03-多卡互联与集群网络)

SP 与 CP 不是一回事 ​

这是最容易混的一处:

SP(序列并行)CP(上下文并行)
切什么非 Attention 区域的激活(LayerNorm / Dropout)Attention 计算本身
为什么TP 没覆盖到的激活冗余序列太长、O(s2) 放不下
与 TP 的关系是 TP 的补丁(同域)独立维度,可与 TP 同时开
通信AllReduce → ReduceScatter + AllGather点对点 / All-to-All

判断标准是「切的是哪一段」:只切到子模块之间的激活(LayerNorm、Dropout 那一段)是 SP;切到 Attention 内部的 QK⊤ 与 PV 是 CP。见 08-张量并行与序列并行。

CP 只解决「装得下」,还有两块在别处 ​

这一篇讲的是系统侧怎么把长序列摊到多卡上。要走通长上下文,另外两块拼图在别的目录:

缺口问题在哪
位置编码外推训 4K 的模型,位置编码没见过 32K 的位置09-位置编码 的 PI / NTK-aware / YaRN
是否真的有效分散在多卡上的序列,模型还能不能检索到中间的信息14-长上下文技术 的 NIAH / RULER 与 Lost in the Middle

CP 把显存问题变成通信问题,不改变模型对长序列的建模能力。 外推方法与结构侧的高效注意力(稀疏 / 线性 / SSM)才是能力侧的事 —— 见 14-长上下文技术。

相关 ​

参考 ​

贡献者 ​

文件历史 ​