M4-051M4: Sequences & TransformersEfficient Attention & FlashAttentionHard
Mastery:

Efficient Attention & FlashAttention: 解释 Ring Attention 与序列并行。

📐 Mathematical Definition
device i holds Qi; ring-pass Kj,Vj blocks; cost=O(L2/n) per device\text{device }i\ \text{holds }Q_i;\ \text{ring-pass }K_j,V_j\ \text{blocks};\ \text{cost}=O(L^2/n)\ \text{per device}
⚡ Executive Summary
Core Concept: 把序列维切到多设备,各设备算局部注意力并用环形传递 K/V 块,实现超长序列的分布式注意力。

📌 Key Takeaways

  • •
    序列维切分(不同于 TP 切权重、PP 切层)
  • •
    环形传递 K/V 块,边传边算、通信与计算重叠
  • •
    可支持超长序列(如 1M token)训练

📐 Mathematical Derivations

数学机理:<strong>动机</strong>——超长序列(如 1M token)的注意力矩阵 L×L 无法放入单卡显存,且计算量巨大;需要<strong>序列维并行(context parallelism)</strong>。<strong>Ring Attention(Liu 等 2023)</strong> 的做法:把序列沿长度维切成 n 份,设备 i 持有 Q_i、K_i、V_i(各 L/n 长度);然后进入<strong>环形传递</strong>循环:第 r 步设备 i 用当前的 K/V 块计算局部分块注意力(累积到输出),同时把 K/V 块传给下一个设备、并从上一个设备接收新块;经过 n 步后,每个设备都见过所有 K/V 块,故得到<strong>完整的</strong>注意力输出(数学上等价于全注意力)。<strong>关键优化</strong>——<strong>通信与计算重叠</strong>:因为每步的计算(一个 L/n × L/n 的分块注意力)耗时与通信(传一个 K/V 块)相当,故可让它们<strong>并行进行</strong>(用双缓冲),使通信几乎被完全隐藏。<strong>复杂度</strong>——每设备的计算 O((L/n)²·d),通信 O(n·(L/n)·d)=O(L·d)(传 n 次块);总显存 O(L/n)(本地的 Q/K/V 与输出)。<strong>与因果 mask 的优化</strong>——在因果(自回归)场景下,设备 i 只需接收'在自己之前'的块(因为未来的块被 mask 掉),故通信量约减半;进一步可结合<strong>zigzag/条带化分块</strong>(把序列按'块对'分配,使每对设备的负载均衡)。<strong>与 TP/PP 的关系</strong>——Ring Attention 是<strong>第四种并行维度</strong>(切序列),与 DP(切数据)、TP(切权重)、PP(切层)正交,可组合(如 TP=8 + CP=8)。

🏭 Production Trade-offs

深度剖析与工程权衡:① <strong>为什么能隐藏通信</strong>——因为分块注意力是'计算密集'的(O((L/n)²·d)),而传递一个 K/V 块是'O((L/n)·d)';计算/通信比约 L/n,只要 L/n 足够大(序列够长)就能完全重叠。这是'用足够大的局部计算掩盖通信'的经典手法。② <strong>因果场景的负载均衡</strong>——朴素切分下,设备 i 的计算量 ∝ i(因为它要处理 i 个 K/V 块);用 <strong>zigzag(条带)切分</strong>(设备 i 拿序列的 i 与 n−1−i 两段)可使各设备负载均衡。③ <strong>与序列并行(SP)的区别</strong>——SP(Megatron)切的是'非 matmul 部分'(LN/dropout)的序列维,主要省激活显存;Ring Attention/CP 切的是'注意力计算'本身,支持超长序列。两者可叠加。④ <strong>与 Flash Attention 的结合</strong>——每个设备内部用 Flash Attention 算局部分块(保持 IO 效率),设备间用环形传递;这是长上下文训练的标准组合。⑤ <strong>实际支持</strong>——Megatron-LM、DeepSpeed、以及部分训练框架已支持 context parallelism;对 1M 级上下文的训练不可或缺。⑥ <strong>面试要点</strong>——被问'超长序列怎么训练',应给出'<strong>序列维切分(CP/Ring Attention)+ 环形传递 K/V + 通信计算重叠 + zigzag 负载均衡 + 内部 Flash Attention</strong>',并说明'它是与 DP/TP/PP 正交的第四维并行';这是分布式训练的高阶问题。
⚠️ Common Interview Pitfalls
  • ✕
    把 Ring Attention 与 TP 混淆(切的是序列而非权重)
  • ✕
    忽略因果 mask 下的负载不均衡与 zigzag 切分
🎯 Interviewer Follow-ups
  • ?
    Ring Attention 与 TP/PP 的关系?
  • ?
    因果 mask 下如何减少通信?
📚

Associated Knowledge Base Guides & Mindmaps

Explore the comprehensive technical article, exam cards, and global architecture tree.

← PreviousM4-050: Efficient Attention & FlashAttention: 解释 PagedAttention 如何解决 KV Cache 碎片。📋Back to BankNext →M4-052: Efficient Attention & FlashAttention: 比较注意力优化的三条路线:IO、稀疏、近似。