跳转至

6D 并行下的 Attention Kernel 与通信重叠

更新日期:2026-04-15


一、各并行维度对 Attention 的影响

飞书 add-on(待手动转 mermaid / 图)

component: blk_631fefbbae02400430b8f9f4

并行 对 Attention 影响 需要特殊 kernel? 为什么
TP 多头注意力的 head 维度被均分到不同 GPU,每个 GPU 持有 H/TP 个头 否,每个 GPU 独立运行标准 FlashAttention,结束后做一次 AllReduce 合并输出 多头注意力天然 head 维度独立——各头的 QKV 投影互不依赖,切头不改变单头内的计算逻辑,只需在输出投影后汇总
PP Transformer 的层被切分到不同 stage,每个 GPU 只执行其中若干层的 Attention 否,Attention kernel 本身不受影响,层间通过点对点 (P2P Send/Recv) 传递激活张量 PP 是层间流水线,Attention 是层内计算;层边界的通信发生在 Attention 之外,kernel 内部感知不到流水线切分
DP 每个 GPU 持有完整模型副本,处理不同 micro-batch,Attention 计算完全独立 否,各 rank 的 Attention 无任何交互,只在反向传播后做梯度 AllReduce 数据并行的核心是 batch 维度切分,不同样本之间的 Attention 本就互不影响
EP 仅影响 MoE 层中的 FFN(专家路由 + All-to-All 调度),Attention 层参数不做专家切分 否,Attention 层的 QKV/Output 投影仍是 dense 参数,不参与专家路由 MoE 架构只将 FFN 替换为稀疏专家,Attention 层结构不变,因此 EP 的 token 调度和 All-to-All 通信不涉及 Attention
CP 序列维度被切分到不同 GPU,每个 GPU 仅持有 S/CP 长度的 Q 和 KV,但 Attention 需要 Q 看到全部 KV 是!必须使用 Ring Attention:各 GPU 通过环形传递 KV 块,配合 online softmax 逐步累积完整 Attention 输出 Attention 的 softmax 归一化要求每个 Q token 看到完整 KV 序列;序列被切后,单 GPU 无法独立计算正确的 softmax 分母,必须通过 Ring 通信 + online softmax 分步合并
SP 在 LayerNorm 和 Dropout 阶段将序列维度切分到 TP group 的各 GPU,降低激活内存 Attention kernel 本身不变,但进入 Attention 前需要 AllGather 拼回完整序列,Attention 输出后需要 ReduceScatter 重新切分 SP 的目标是减少激活内存(降为 1/TP),代价是在 Attention 前后各加一次集合通信;但这两次通信替换了原本 TP 就需要的 AllReduce,总通信量不增加

二、Ring Attention 深入

CP 的核心 kernel:每个 GPU 只有部分 Q 和 KV,通过环形传递 KV 完成全局 attention。参考 Ring Attention (Liu et al., 2023)

2.1 伪代码

def ring_attention(Q_local, KV_local, cp_group):
    # Q_local, K_local, V_local: [B, S/CP, H, d]
    # 每个 GPU 持有 1/CP 的序列

    cp_size = len(cp_group)

    # Online softmax 累积状态
    O_acc = zeros_like(Q_local)
    m_acc = full(Q_shape[:-1], -inf)
    l_acc = zeros(Q_shape[:-1])

    kv_current = KV_local  # 初始是本地 KV

    for step in range(cp_size):
        # 1. 启动异步发送到下一个 GPU
        send_handle = async_send(kv_current, next_rank(cp_group))
        recv_handle = async_recv(prev_rank(cp_group))

        # 2. 同时本地计算当前 KV 块的 attention
        O_step, m_step, l_step = flash_attention(Q_local, kv_current[0], kv_current[1])

        # 3. Online softmax 合并
        m_new = max(m_acc, m_step)
        alpha = exp(m_acc - m_new)
        beta = exp(m_step - m_new)
        l_acc = alpha  l_acc + beta  l_step
        O_acc = alpha  O_acc + beta  O_step  # unnormalized
        m_acc = m_new

        # 4. 等通信完成
        kv_current = await_both(send_handle, recv_handle)

    return O_acc / l_acc

2.2 关键优化:通信-计算重叠

Ring Attention 的关键是通信和计算并行。计算 step i 的 attention 时,通信 step i+1 的 KV。 时间线分析(CP=4): 如果通信时间 < 计算时间 → 通信完全被隐藏。实际中通信和计算约 1:1 → 约 30% 通信暴露(不能完全隐藏)。

2.3 Causal Mask 下的优化

Causal attention 下,位置 i 只看位置 0..i。这意味着某些 Ring step 是不必要的。 Causal mask 下,GPU 0 持有序列前 ¼。GPU 0 的 Q 只需要看前 ¼ 的 KV(自己的)→ 不需要接收后面的 KV。

优化后的 Ring 调度: 总计算量 = 1+2+3+4 = 10,是非 causal 的 16 步的 62.5%——节省了 37.5% 的计算。但 GPU 0, 1 会先完成 → 负载不均

解决方案:Striped Attention (Brandon et al., 2024)——将序列 reshape 为条纹(交错分配给各 GPU),让每个 GPU 的 causal 计算量均衡。


三、Serialize/Deserialize 与 KV Cache

flowchart LR
    layer1["Layer L<br/>[B, S, H/TP, d]"]
    bound["层边界<br/>AllGather + ReduceScatter<br/>(融合为 All-to-All)"]
    layer2["Layer L+1<br/>[B, S/SP, H, d]"]

    layer1 --> bound --> layer2

    classDef stage fill:#fff,stroke:#cc785c,color:#1a1a1a;
    class layer1,bound,layer2 stage

序列并行(SP)的关键 trick:激活内存减少 TP 倍,但通信量不增加。

3.1 TP + SP 下的通信模式

TP=4, SP 开启时的数据布局: 层边界需要转换:从 [B, S, H/4, d][B, S/4, H, d]。用 AllGather(集齐 head)+ ReduceScatter(切分 sequence),实际融合为一个 "All-to-All" 式操作。

Megatron 中开启 SP:--sequence-parallel。总通信量不增加,但激活内存减少 TP 倍


四、MoE + Attention 的通信组合

MoE 层用 All-to-All (EP), Attention 层用 AllReduce (TP)。两者交替出现时的通信: 一个典型的 MoE Transformer Block 有两种通信:Attention(x) 用 TP AllReduce,MoE(x) 用 EP All-to-All。

优化 1:通信重叠 — Attention 的计算可以和 MoE 的 All-to-All 部分重叠(DMA 引擎独立于 SM)。

优化 2:减少冗余通信 — Attention→MoE 的边界可能需要 AllGather (TP) + All-to-All (EP)。优化方案是合并为单一通信。

DeepSeek-V3 的激进做法:完全不用 TP! Attention 层每个 GPU 持有完整参数(用 DP),MoE 层用 EP。代价是每 GPU 需要更多内存存 Attention 参数,但消除了 TP 的 AllReduce(每层省两次 AllReduce)。这在 MLA 下可行,因为 MLA 的 Attention 参数比标准 MHA 小很多。


五、Attention Kernel 与通信的联合优化

5.1 DeepSeek-V3 的具体优化

参考 DeepSeek-V3 Technical Report


六、长序列训练的 Attention 优化

flowchart LR
    flash["FlashAttention 2/3<br/>(O(n²) 但 IO-aware)"]
    ring["Ring Attention<br/>(分布式 KV)"]
    striped["Striped Attention<br/>(均衡 causal)"]
    sparse["Sparse Attention<br/>(O(n) 实际)"]

    flash --> ring --> striped --> sparse

    classDef stage fill:#fff,stroke:#cc785c,color:#1a1a1a;
    class flash,ring,striped,sparse stage
方案 复杂度 适合长度 备注
FlashAttention 2 O(n²) < 64k 单 GPU IO-aware,标准 baseline
FlashAttention 3 O(n²) < 64k H100 + FP8 加速
Ring Attention O(n²/CP) 64k-1M 分布式跨 GPU 切分 KV
Striped Attention O(n²/CP) 64k-1M Ring 的 causal 平衡版
DSA / NSA / 稀疏 O(n × k) > 1M 牺牲一定召回换速度

七、稀疏 Attention Kernel

对于极长序列,密集 attention 的 O(n²) 不可接受。稀疏 attention 是解决方案。

7.1 DeepSeek Sparse Attention (DSA)

# DSA: 对每个 Q, 用索引器选 top-k 相关 K, 只在这 k 个上算 attention
def dsa(Q, K, V, top_k=512):
    # 1. Indexer: 快速估计 Q 对每个 K 的重要性 (不做完整注意力)
    # 用低秩投影快速算出粗略分数
    importance = low_rank_matmul(Q, K)  # O(n²) 但矩阵很小

    # 2. Top-k 选择
    top_k_indices = importance.topk(top_k).indices  # [B, H, n, k]

    # 3. 只在 top-k 上做标准 attention
    K_selected = gather(K, top_k_indices)  # [B, H, n, k, d]
    V_selected = gather(V, top_k_indices)

    s = (Q.unsqueeze(-2) * K_selected).sum(-1)  # [B, H, n, k]
    attn = softmax(s)
    return (attn.unsqueeze(-1) * V_selected).sum(-2)

# 复杂度: O(n·k), k << n
# 例: n=128K, k=512 → 计算量减少 250x

参考文献

  • [1] Liu et al. Ring Attention. 2023. 论文

  • [2] Brandon et al. Striped Attention. 2024. 论文

  • [3] Dao et al. FlashAttention-2. 2023. 论文

  • [4] Shah et al. FlashAttention-3. 2024. 论文

  • [5] DeepSeek-V3 Technical Report. 2024. 论文

  • [6] Sequence and Context Parallelism in Megatron. DeepWiki

  • [7] Korthikanti et al. Reducing Activation Recomputation (SP). 2022. 论文


上级 · C. 分布式训练基础设施