长序列训练与上下文并行
序列长度从 2K 涨到 128K 甚至 1M 之后,出现了一类前面所有并行维度都解决不了的问题。
这一篇讲清「序列长度墙」的成因,以及三种沿序列切分的方案。
序列长度墙
Attention 的计算量与激活显存都是
| 维度 | 能不能解决「单条序列太长」 |
|---|---|
| DP | 不能 —— 每条序列仍要完整算一遍 |
| ZeRO | 不能 —— 切的是状态冗余 |
| TP | 不能 —— 切的是权重, |
| PP | 不能 —— 切的是层,序列维没动 |
| CP | 能 —— 沿序列维切分 Attention 计算本身 |
「
翻倍开销翻 4 倍」与显存线性增长的差别,就是这条墙的来源:显存不够可以加卡,但 意味着加卡的速度追不上序列增长的速度。
序列长度增长时两条曲线的斜率不同:
相对倍数
16× ┤ ● Attention 开销:O(s²)
│ ╱
│ ╱
4× ┤ ● ╱
│ ╱ ● ← s 翻倍 ⇒ 开销翻 4 倍
│ ╱ ╱
2× ┤ ● ╱ ● ○ 显存 / 卡数只能线性增长
│ ╱ ○
1× ┤ ● ○
└──┬─────┬─────┬──────┬──────▶ 序列长度
2K 4K 8K 16K
⇒ 加卡的速度追不上序列增长的速度 —— 这就是「墙」
⇒ 而且 DP / ZeRO / TP / PP 四个维度都解不了它(它们都不动序列维)
只有 CP 沿序列维切 Attention 计算本身前置:FlashAttention
CP 的三种方案几乎都构建在 FlashAttention 之上 —— 因为它已经把
分块 + Online Softmax 的机制见 08-FlashAttention。它是 CP 的前提:如果还要存完整
三种方案
Ring Attention
沿序列维把
每卡持有自己那段 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 大 |
| Ulysses | Head → 序列 | All-to-All | Head 数需 ≥ CP 度 | Head 多的模型 |
| Megatron CP | 序列 | 点对点 / AllGather | 与 Megatron 绑定 | Megatron 生态 |
两条判据:
- 看 Head 数 —— Head 数少(或 GQA 下 KV Head 少)时 Ulysses 的 All-to-All 方案受限,该选 Ring 或 Megatron CP
- 看通信模式与硬件的匹配 —— Ring 是点对点环传(对拓扑敏感度低),Ulysses 是 All-to-All(需要全互联的带宽,见 03-多卡互联与集群网络)
SP 与 CP 不是一回事
这是最容易混的一处:
| SP(序列并行) | CP(上下文并行) | |
|---|---|---|
| 切什么 | 非 Attention 区域的激活(LayerNorm / Dropout) | Attention 计算本身 |
| 为什么 | TP 没覆盖到的激活冗余 | 序列太长、 |
| 与 TP 的关系 | 是 TP 的补丁(同域) | 独立维度,可与 TP 同时开 |
| 通信 | AllReduce → ReduceScatter + AllGather | 点对点 / All-to-All |
判断标准是「切的是哪一段」:只切到子模块之间的激活(LayerNorm、Dropout 那一段)是 SP;切到 Attention 内部的
与 是 CP。见 08-张量并行与序列并行。
CP 只解决「装得下」,还有两块在别处
这一篇讲的是系统侧怎么把长序列摊到多卡上。要走通长上下文,另外两块拼图在别的目录:
| 缺口 | 问题 | 在哪 |
|---|---|---|
| 位置编码外推 | 训 4K 的模型,位置编码没见过 32K 的位置 | 09-位置编码 的 PI / NTK-aware / YaRN |
| 是否真的有效 | 分散在多卡上的序列,模型还能不能检索到中间的信息 | 14-长上下文技术 的 NIAH / RULER 与 Lost in the Middle |
CP 把显存问题变成通信问题,不改变模型对长序列的建模能力。 外推方法与结构侧的高效注意力(稀疏 / 线性 / SSM)才是能力侧的事 —— 见 14-长上下文技术。
相关
- 08-FlashAttention —— CP 的算法前提,也是 Ring Attention 的基础
- 08-张量并行与序列并行 —— SP 与 CP 的区分
- 11-Self-Attention 机制 ——
与显存算例 - 03-多卡互联与集群网络 —— All-to-All 对拓扑的要求
- 10-KV Cache 与推理优化 —— 长上下文在推理侧的对应问题
- 14-长上下文技术 —— 能力侧:外推方法、稀疏 / 线性注意力 / SSM、长上下文评估
- 09-位置编码 —— RoPE 外推(PI / NTK-aware / YaRN / Dynamic NTK)
参考
- https://caomaolufei.github.io/AIInfraGuide/guides/模块三-分布式训练/第9章-长序列训练与上下文并行
- https://arxiv.org/abs/2310.01889
- https://arxiv.org/abs/2309.14509
- https://arxiv.org/abs/2205.14135
YJ