Skip to content

序列并行与上下文并行

序列并行(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.beta

Ring 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_maxO(L×H)
TP=4L_maxO(L×H/4)AllReduce
SP+TPL_maxO(L×H/4)AllGather
CP=44×L_maxO(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。

相关资源

最近更新