在大模型推理与训练日益追求极致性能的今天,如何高效处理长序列注意力计算已成为系统架构设计的核心挑战。FlashAttention 凭借其精妙的切分策略,不仅突破了显存瓶颈,更为分布式并行计算提供了全新思路。本文将带你从底层公式出发,逐步拆解其切分逻辑,并深入探讨在高并发场景下的分布式落地实践。
一、核心公式与 Tiling 切分逻辑:打破显存墙的基石
传统的注意力机制计算遵循宏观公式:O = Softmax(Q K^T) V。其中 Q、K、V 的维度均为 [N, d](N 为序列长度),输出 O 同样为 [N, d]。然而,当 N 达到数万甚至数十万时,中间产生的 N × N 注意力分数矩阵会瞬间撑爆显存,成为典型的高并发场景下的性能杀手。
FlashAttention 的核心洞察在于:GPU 的 SRAM(高速缓存)虽然容量小,但带宽极高;而 HBM(高带宽显存)容量大,但访问速度相对较慢。 因此,算法不应一次性生成巨大的注意力矩阵,而应将其“化整为零”。
具体来说,FlashAttention 采用了 Tiling(分块)策略:
- Q 切分(外循环): 将 Q 按行切分为
T_r个块,每块大小为B_r × d,记为Q_1, Q_2, ...。 - K, V 切分(内循环): 将 K 和 V 按行切分为
T_c个块,每块大小为B_c × d,记为K_1, V_1, ...。这里 K 的切分对应注意力矩阵的列方向。
通过这种双层循环结构,FlashAttention 确保了每次参与计算的数据块都能被完整地放入 SRAM 中,从而避免了频繁的 HBM 读写。以下是其核心算法流程的伪代码示意:
for i in 1 to Tr: # 外循环:遍历 Query 块 (加载 Qi 到 SRAM)
# 初始化局部累加器 O_i, l_i (sum), m_i (max)
for j in 1 to Tc: # 内循环:遍历 Key/Value 块 (加载 Kj, Vj 到 SRAM)
1. 计算分数: S_ij = Qi * Kj^T
2. 更新统计量 (Online Softmax): 更新局部 max 和 sum
3. 计算局部结果: P_ij = Softmax(S_ij)
4. 累加到 O_i: O_i = O_i + P_ij * Vj (注意这里有 rescale)
# 内循环结束,O_i 计算完成,写回 HBM
✅ 实践建议: 在实际系统架构设计中,块大小 B_r 和 B_c 的选择至关重要。通常需要根据 SRAM 的实际容量和寄存器压力进行调优,以达到计算与访存的最佳平衡。
二、Step-by-Step 数值推演:揭开 Online Softmax 的神秘面纱
为了更直观地理解切分后的计算过程,我们用一个极简的数值例子来模拟。假设序列长度 N = 4,块大小 B = 2。那么 Q 被切成 Q_1, Q_2,K 和 V 被切成 K_1, V_1 和 K_2, V_2。我们的目标是计算输出 O_1(对应前两个 Token)。
- Step 1:加载 Q₁(外循环启动)
从 HBM 读取Q_1(Token 0, 1)到 SRAM。此时将输出O_1初始化为 0,并初始化最大值向量m和归一化因子l。 - Step 2:加载 K₁, V₁(内循环第一轮)
从 HBM 读取K_1, V_1(Token 0, 1)。计算局部分数S_11 = Q_1 × K_1^T(一个 2×2 矩阵)。假设得到S_11 = [[10, 20], [10, 10]]。此时进行局部 Softmax 计算,更新当前最大值m_new = [20, 10],并计算局部概率P_11。随后更新输出:O_1 = P_11 × V_1。此时O_1仅包含 Token 0,1 对 Token 0,1 的注意力结果。 - Step 3:加载 K₂, V₂(内循环第二轮)—— 关键步骤
从 HBM 读取K_2, V_2(Token 2, 3)。注意:Q₁ 依然驻留在 SRAM 中,无需重新加载! 计算分数S_12 = Q_1 × K_2^T。假设得到S_12 = [[30, 5], [5, 5]]。此时触发 Online Softmax 更新机制:对比上一轮的最大值m_old = [20, 10]和当前局部最大值[30, 5]。对于第一行,新最大值 30 大于旧值 20,这意味着上一轮计算出的O_1第一行贡献偏大,需要乘以e^(20-30)进行缩放(Rescale)。最后将缩放后的旧输出与新的局部贡献相加:O_1 = Rescale(O_1) + P_12 × V_2。 - Step 4:写回
内循环结束,所有 K, V 块均已遍历。此时O_1即为最终结果,将其写回 HBM。
⚠️ 注意事项: Online Softmax 的精髓在于“边走边看”,它通过动态调整历史累积值的权重,确保了最终结果与一次性计算全局 Softmax 的数值等价性。这种流式计算模式极大地降低了显存占用。
三、从单机到分布式:Ring Attention 与 FlashDecoding 架构演进
在单 GPU 上,上述双层循环是串行执行的。但在分布式训练或推理场景中,我们可以将循环拆解,利用多设备并行来进一步加速。这不仅是微服务架构思想在底层计算领域的体现,更是实现高可用、高吞吐系统的关键。
场景 A:Ring Attention —— 长文本 Prefill 的分布式利器
Ring Attention 是 FlashAttention 分布式版本的经典实现。其核心逻辑是将内循环(遍历 K, V 块)转化为跨设备的“接力传球”。
假设有 2 个计算核心(Core 0, Core 1),序列长度 N=4。Core 0 持有 Token 0,1 的数据(Q₀, K₀, V₀),Core 1 持有 Token 2,3 的数据(Q₁, K₁, V₁)。目标是让 Core 0 计算 Q₀ 对全量 K 的注意力。
- Phase 1(本地计算): Core 0 计算
Attn(Q₀, K₀, V₀),Core 1 计算Attn(Q₁, K₁, V₁)。此时片上网络(NoC)空闲。 - Phase 2(通信与计算重叠): 通过 NoC 进行环形通信。Core 0 将 K₀, V₀ 发送给 Core 1,同时接收来自 Core 1 的 K₁, V₁。在接收数据的同时,Core 0 立即计算
Attn(Q₀, K₁, V₁),并利用 Online Softmax 将结果与 Phase 1 的结果合并。
在片内,Memory Block 之间通过 NoC 互联,带宽极高,这使得 Ring Attention 能够实现近乎线性的加速比。对于长文本 Prefill 场景,这无疑是突破单卡算力瓶颈的最佳方案。
场景 B:FlashDecoding —— 解码阶段的 Split-K 并行
在 Decoding 阶段,由于 Q 只有 1 行(当前生成的 Token),FlashAttention 的外循环(按 Q 切分)无法并行,导致 GPU 利用率低下。FlashDecoding 提出了 Split-K 策略来解决这一问题。
其核心思想是:既然 Q 切不动,那就切 K。
- 广播 Q: 将当前的 Query(1 × d)广播给所有计算核心。
- 并行计算 (Map): 假设有 100 个核心,KV Cache 被分散在 100 个 Block 中。每个核心独立计算 Q 与本地 K, V 块的注意力,得到局部的输出
O_partial和统计量(最大值 m, 归一化因子 l)。例如 Core 0 得到Partial_O_0,Core 99 得到Partial_O_99。 - 树状归约 (Reduce): 这 100 个局部结果
Partial_O需要通过树状归约进行合并。利用 Online Softmax 公式,两两合并,最终得到全局唯一的输出 O。
✅ 架构启示: FlashDecoding 的 Split-K 策略完美诠释了分布式系统中“分而治之”的思想。通过将长序列的 KV Cache 切分到不同核心,再通过高效的归约通信合并结果,极大地提升了高并发解码场景下的吞吐量。
[AFFILIATE_SLOT_2]四、总结与展望
FlashAttention 的切分机制不仅是算法层面的优化,更是对现代 GPU 存储层次结构的深刻理解。从微观的 Tiling 分块与 Online Softmax 数值稳定技巧,到宏观的 Ring Attention 与 FlashDecoding 分布式并行策略,它为我们展示了如何通过系统架构创新来释放硬件潜能。
"FlashAttention 采用了**双层分块(Tiling)**策略:
- 外层循环切分 Query,负责决定输出的行。
- 内层循环切分 Key/Value,负责在 SRAM 中流式计算并累加结果。
- 利用 Online Softmax 技巧,保证了分块计算后能还原出精确的全局 Softmax 结果,避免了 N × N N \times N N×N 矩阵的显存读写。
分布式/多核架构时,这种切分方式可以自然地映射为并行策略:
- 针对 Prefill (长文本): 采用 Ring Attention 模式。将内层循环的 KV 加载变成片上网络的数据流动。每个 Core 固定处理一部分 Q,让 KV Block 在 Core 之间流转,实现计算与通信的完美重叠。
- 针对 Decoding (生成): 由于 Q 很小,采用 Split-K (FlashDecoding) 模式。将 KV Cache 物理打散到所有 Block,广播 Q,让所有 Core 并行计算部分 Attention,最后通过片上归约树(Reduction Tree)合并结果。
这种软硬协同的切分,能最大化利用片内的高带宽和多核算力。"
展望未来,随着模型序列长度的持续增长和微服务架构在 AI 推理中的普及,类似 FlashAttention 这种软硬件协同设计的思想将变得愈发重要。掌握其切分逻辑,将有助于我们构建更加高可用、高性能的 AI 基础设施。
浙公网安备 33010602011771号