序列并行与上下文并行
序列并行(Sequence Parallelism)和上下文并行(Context Parallelism)是针对长序列训练的并行方案,突破单 GPU 的序列长度限制。
动机
标准张量并行中,Dropout 和 LayerNorm 等操作在序列维度上不切分,导致每个 GPU 仍需存储完整序列的激活值。序列并行解决这一问题:
- 序列并行 (SP):将 LayerNorm/Dropout 的激活值沿序列维度切分
- 上下文并行 (CP):将整个序列切分到不同 GPU,支持超长上下文
序列并行原理
在 Megatron-LM 中,SP 与 TP 结合使用:
python
# 标准TP:LayerNorm需要完整序列
# SP:LayerNorm沿序列切分
class SequenceParallelLayerNorm:
def forward(self, x):
# x: [batch, seq_len/tp, hidden]
# 需要全局统计量
local_mean = x.mean(dim=-2)
local_var = x.var(dim=-2)
# AllReduce 获取全局统计量
global_mean = dist.all_reduce(local_mean) / tp_size
global_var = dist.all_reduce(local_var) / tp_size
# 本地归一化
x_norm = (x - global_mean) / torch.sqrt(global_var + self.eps)
return self.gamma * x_norm + self.betaRing Attention 上下文并行
Ring Attention 将序列切分到多个 GPU,通过环形通信传递 KV 块:
python
# Ring Attention 伪代码
def ring_attention(q, k, v, rank, world_size):
"""环形注意力计算"""
seq_per_rank = q.shape[1] // world_size
local_q = q[:, rank*seq_per_rank:(rank+1)*seq_per_rank]
# 本地 KV
local_k, local_v = k[rank], v[rank]
output = torch.zeros_like(local_q)
for step in range(world_size):
# 计算当前 KV 块的注意力
scores = local_q @ local_k.transpose(-1, -2) / math.sqrt(d)
attn = softmax(scores)
output += attn @ local_v
# 环形传递 KV
next_rank = (rank + 1) % world_size
local_k = dist.send_recv(local_k, dst=next_rank, src=(rank-1)%world_size)
local_v = dist.send_recv(local_v, dst=next_rank, src=(rank-1)%world_size)
return output显存与序列长度
| 方案 | 最大序列长度 | 显存占用 | 通信 |
|---|---|---|---|
| 无并行 | L_max | O(L×H) | 无 |
| TP=4 | L_max | O(L×H/4) | AllReduce |
| SP+TP | L_max | O(L×H/4) | AllGather |
| CP=4 | 4×L_max | O(L×H/4) | P2P |
CP 与 SP 组合
CP 和 SP 可以组合使用:CP 处理序列切分,SP 处理 LayerNorm 等操作的序列维度切分。Llama 3 训练使用 TP=4 + CP=4 + DP 的组合。
CP 通信开销
Ring Attention 每步需要 P2P 通信传递 KV 块,通信量 = batch × (seq/P) × hidden × 2。对于 128K 上下文、CP=8,通信量约 2GB/step。