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. 分布式训练基础设施