资讯动态

详细讲解 FlashDecoding 与 FlashDecoding+ 的原理

发布时间:2026/8/4 10:38:36 来源:尧图企业网站定制
1. 铺垫背景Decode 阶段标准 FlashAttention 的致命瓶颈在 Prefill 阶段Query 的序列长度NQN_{Q}NQ​很大如 4096可以很好地利用 GPU 的 Tensor Core 并行计算但在 Decode 阶段Query 形状为 [B,1,H,D](Batch Size B序列长度 1Head 数 H维度 D)\text{Query 形状为 } [B, 1, H, D] \quad (\text{Batch Size } B \text{序列长度 } 1 \text{Head 数 } H \text{维度 } D)Query形状为[B,1,H,D](Batch SizeB序列长度1Head数H维度D)此时QQQ只有1 个 Token但它需要与历史 KV Cache长度NKVN_{KV}NKV​可能为 128k计算 Attention。标准 FlashAttention 在 Decode 阶段的并发危机GPU 算力饿死并行维度受限Parallelism Bound标准 FlashAttention 的并行粒度是按Batch Size (BBB)×\times×Heads (HHH)来划分 Thread Blocks 的。假设场景推理时B1,H32B1, H32B1,H32整个 GPU 只有1×32321 \times 32 321×3232个独立的并行任务。物理后果一块 H100/A100 显卡有100 个 SMStreaming Multiprocessor只有 32 个任务意味着只有 32 个 SM 在干活剩下的 70 个 SM 完全闲置空转极度的 Memory-Bound内存带宽瓶颈单个线程块需要沿着 128k 长的 KV Cache串行Sequential做循环 Tiling 累加。GPU 的内存带宽被长序列拉爆而算力利用率Occupancy Tensor Core Utilization低至可怜的 5%~10%。核心矛盾KV Cache 沿着 Sequence 维度极长但标准 FlashAttention不敢对 KV Cache / Sequence 维度做跨 SM 的并行因为最后一个 Token 的 Softmax 需要全局的最大值mmm和分母ddd2. FlashDecoding 核心原理 Sequence 维度的跨 SM 拆分FlashDecoding 的核心破局点一句话总结“利用 Online Softmax 的归一化数学性质强制对 KV Cache 的 Sequence 维度做切块Split KV扔给不同的 SM 并行计算最后做一次快速归约Reduction”[ 标准 FlashAttention (Decode 阶段) ] SM 0 ─── 处理 Batch 1, Head 1 (串行遍历 128k 全量 KV Cache......) ─── 输出 SM 1 ─── 处理 Batch 1, Head 2 (串行遍历 128k 全量 KV Cache......) ─── 输出 ... (其余 70 个 SM 闲置) [ FlashDecoding (KV 序列拆分) ] 将 128k KV Cache 切分为 16 个 Slices (每个 8k): SM 0 ─── 处理 Slice 0 (0~8k) ─── 输出局部 (O_0, m_0, d_0) ┐ SM 1 ─── 处理 Slice 1 (8k~16k) ─── 输出局部 (O_1, m_1, d_1) ├── 第二阶段: 极轻量级 Tree-Reduction ─── 最终 O ... │ (利用 Online Softmax 跨 Block 融合) SM 15 ─── 处理 Slice 15 (120k~) ─── 输出局部 (O_15, m_15, d_15)┘FlashDecoding 的两阶段执行 Pipeline阶段一Split-KV 跨 SM 并行计算Map 阶段切分策略除了按B×HB \times HB×H切分外增加一个KV Sequence 拆分因子KnumK_{num}Knum​。比如把NKV128kN_{KV}128\text{k}NKV​128k拆分为 16 个小切片Slices。并行任务数总并行任务数增加为B×H×KnumB \times H \times K_{num}B×H×Knum​例如1×32×165121 \times 32 \times 16 5121×32×16512。所有 SM 被瞬间塞满SM 片上计算每个 SM 处理属于自己的小 KV 切片在片上 SRAM 利用标准 FlashAttention 逻辑计算输出 3 个局部状态局部未归一化输出向量O~i\tilde{O}_iO~i​局部最大值mim_imi​局部分母累加和did_idi​将这 3 个极小的局部中间结果写入 Global Memory显存占用极微小仅为O(B×H×Knum×D)O(B \times H \times K_{num} \times D)O(B×H×Knum​×D)。阶段二跨 Block 快速归约Tree-Reduction 阶段发射一个极小的 Reduction Kernel。重新利用Online Softmax 的校正因子推导mglobalmax⁡(m1,m2,…,mk)m_{global} \max(m_1, m_2, \dots, m_k)mglobal​max(m1​,m2​,…,mk​)αiemi−mglobal\alpha_i e^{m_i - m_{global}}αi​emi​−mglobal​dglobal∑iαi⋅did_{global} \sum_{i} \alpha_i \cdot d_idglobal​i∑​αi​⋅di​Ofinal∑iαi⋅O~idglobalO_{final} \frac{\sum_{i} \alpha_i \cdot \tilde{O}_i}{d_{global}}Ofinal​dglobal​∑i​αi​⋅O~i​​耗时因为KnumK_{num}Knum​很小如 16 或 32归约计算量极小耗时几乎接近 0 毫秒3. FlashDecoding 的进阶突破异步与动态自适应FlashDecoding 虽然解决了 SM 占不满的问题但在工业级实际部署中仍有两个痛点固定 Split 粒度引发的负载不均与 Overhead对于短序列过度 Split 导致的 Reduction 阶段开销反而侵蚀了收益。Synchronized Barrier同步屏障开销阶段一Map与阶段二Reduce之间需要一个 Global Barrier同步所有 SM。百度与学术界等提出的FlashDecoding对此进行了深度重构与优化[ FlashDecoding 架构优化 ] │ ┌─────────────────────────────────┴─────────────────────────────────┐ ▼ ▼ 【动态 Split-KV 决策树】 【Unified Kernel 与 Asynchronous Reduction】 根据 (Batch, Head, SeqLen, Hardware SM Count) 消除独立的 Reduction Kernel 动态求解最优 Split 块数 K_num 利用 Tensor Core 与 Shared Mem 异步流水线FlashDecoding 的三大核心技术突破动态自适应 Split 策略Dynamic Load-BalancingFlashDecoding 在 Runtime 引入了一个超轻量级的代价模型Cost Model。根据当前请求的BBB、序列长度NKVN_{KV}NKV​以及目标 GPU 的 SM 物理数量动态计算出能恰好填满 SM 的最优切分块数KnumK_{num}Knum​。短序列不切或少切超长序列大幅切彻底避免了“为了切而切”的调度 Overhead。异步 Reduction 与 Kernel 融合Unified Kernel / Asynchronous ReductionFlashDecoding 将 Map 与 Reduce 逻辑融合成单个 CUDA Kernel。利用 GPU 的Atomic Operations原子操作或 SM 间的Grid-Level Barrier / Asymmetric Synchronization先算完的 SM 可以直接异步参与部分归约计算进一步消除了跨 Kernel 发射与全局同步的开销。针对 Flat Head / GQAGrouped-Query Attention的特化优化现代大模型如 LLaMA-3、Mistral广泛采用 GQA如 8 个 KV Head 对应 32 个 Query Head。FlashDecoding 针对 GQA 的 KV Cache 共享特性做成了专门的KV-Reused Layout 优化大幅提升了 Shared Memory 缓存命中率。4. 面试高频对比FlashAttention vs FlashDecoding vs FlashDecoding维维度Standard FlashAttention (V1/V2)FlashDecodingFlashDecoding主攻阶段Prefill 阶段长 Q长 K/VDecode 阶段短 Q超长 K/VDecode 阶段全场景/动态长上下文并行维度B×HB \times HB×H(Batch×\times×Heads)B×H×KnumB \times H \times K_{num}B×H×Knum​(引入 KV Sequence 维度)B×H×Dynamic(Knum)B \times H \times \text{Dynamic}(K_{num})B×H×Dynamic(Knum​)(自适应 GQA 特化)SM 利用率Decode 阶段低低于 10%Decode 阶段极高接近 100%极致全 Sequence 长度下保持 90%计算流程单 Kernel 串行 Tiling 累加两阶段Map 切片 Tree-Reduce动态 Unified 单 Kernel 异步归约核心数学片上 Online Softmax跨 Block / 跨 SM 的 Online Softmax异步原子归约 动态代价模型5. 复盘背诵口诀FlashDecoding 核心突破“Decode 阶段 Q 只有一SM 闲置算力低切分 KV 跨 SM 跑局部状态存下来Online Softmax 做归约长上下文速度飞。”一句话精炼“FlashDecoding 突破了 Decode 阶段按 Batch/Head 并行的硬性限制利用 Online Softmax 的可按块缩放特性将KV Cache 序列Sequence维度切块分发给多个 SM 并行计算最后通过毫秒级 Tree-Reduction 汇总彻底拉满 GPU 算力利用率。”

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价