news 2026/8/4 10:38:33

详细讲解 FlashDecoding 与 FlashDecoding+ 的原理

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
详细讲解 FlashDecoding 与 FlashDecoding+ 的原理

1. 铺垫背景:Decode 阶段标准 FlashAttention 的致命瓶颈

在 Prefill 阶段,Query 的序列长度NQN_{Q}NQ很大(如 4096),可以很好地利用 GPU 的 Tensor Core 并行计算;但在 Decode 阶段:

Query 形状为 [B,1,H,D](Batch Size B,序列长度 1,Head 数 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,序列长度1HeadH,维度D)

此时QQQ只有1 个 Token,但它需要与历史 KV Cache(长度NKVN_{KV}NKV,可能为 128k)计算 Attention。

标准 FlashAttention 在 Decode 阶段的并发危机(GPU 算力饿死):

  1. 并行维度受限(Parallelism Bound)
  • 标准 FlashAttention 的并行粒度是按Batch Size (BBB)×\times×Heads (HHH)来划分 Thread Blocks 的。
  • 假设场景:推理时B=1,H=32B=1, H=32B=1,H=32,整个 GPU 只有1×32=321 \times 32 = 321×32=32个独立的并行任务。
  • 物理后果:一块 H100/A100 显卡有100+ 个 SM(Streaming Multiprocessor)!只有 32 个任务意味着只有 32 个 SM 在干活,剩下的 70+ 个 SM 完全闲置空转
  1. 极度的 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和分母ddd


2. 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。比如把NKV=128kN_{KV}=128\text{k}NKV=128k拆分为 16 个小切片(Slices)。
  • 并行任务数:总并行任务数增加为B×H×KnumB \times H \times K_{num}B×H×Knum(例如1×32×16=5121 \times 32 \times 16 = 5121×32×16=512)。所有 SM 被瞬间塞满!
  • SM 片上计算:每个 SM 处理属于自己的小 KV 切片,在片上 SRAM 利用标准 FlashAttention 逻辑计算,输出 3 个局部状态
  1. 局部未归一化输出向量:O~i\tilde{O}_iO~i
  2. 局部最大值:mim_imi
  3. 局部分母累加和: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 的校正因子推导

mglobal=max⁡(m1,m2,…,mk)m_{global} = \max(m_1, m_2, \dots, m_k)mglobal=max(m1,m2,,mk)

αi=emi−mglobal\alpha_i = e^{m_i - m_{global}}αi=emimglobal

dglobal=∑iαi⋅did_{global} = \sum_{i} \alpha_i \cdot d_idglobal=iαidi

Ofinal=∑iαi⋅O~idglobalO_{final} = \frac{\sum_{i} \alpha_i \cdot \tilde{O}_i}{d_{global}}Ofinal=dglobaliαiO~i

  • 耗时:因为KnumK_{num}Knum很小(如 16 或 32),归约计算量极小,耗时几乎接近 0 毫秒

3. FlashDecoding+ 的进阶突破:异步与动态自适应

FlashDecoding 虽然解决了 SM 占不满的问题,但在工业级实际部署中仍有两个痛点:

  1. 固定 Split 粒度引发的负载不均与 Overhead:对于短序列,过度 Split 导致的 Reduction 阶段开销反而侵蚀了收益。
  2. 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+ 的三大核心技术突破:

  1. 动态自适应 Split 策略(Dynamic Load-Balancing)
  • FlashDecoding+ 在 Runtime 引入了一个超轻量级的代价模型(Cost Model)
  • 根据当前请求的BBB、序列长度NKVN_{KV}NKV以及目标 GPU 的 SM 物理数量,动态计算出能恰好填满 SM 的最优切分块数KnumK_{num}Knum
  • 短序列不切或少切,超长序列大幅切,彻底避免了“为了切而切”的调度 Overhead。
  1. 异步 Reduction 与 Kernel 融合(Unified Kernel / Asynchronous Reduction)
  • FlashDecoding+ 将 Map 与 Reduce 逻辑融合成单个 CUDA Kernel
  • 利用 GPU 的Atomic Operations(原子操作)或 SM 间的Grid-Level Barrier / Asymmetric Synchronization:先算完的 SM 可以直接异步参与部分归约计算,进一步消除了跨 Kernel 发射与全局同步的开销。
  1. 针对 Flat Head / GQA(Grouped-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/V)Decode 阶段(短 Q,超长 K/V)Decode 阶段(全场景/动态长上下文)
并行维度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 算力利用率。”

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/8/4 10:37:07

三角洲掉帧?这款电脑录屏软件完全不损帧率

打三角洲行动的都知道,这游戏对电脑配置要求不低。好不容易跑满帧率,开个录屏直接掉到三十几帧,操作全糊了,精彩击杀回放一看全是马赛克——别提多糟心。我本身是个老FPS玩家,从CS打到瓦罗兰特再打到三角洲&#xff0c…

作者头像 李华
网站建设 2026/8/4 10:35:33

嵌入式通信协议设计:从字节流到可靠数据交换的工程实践

你有没有过这样的经历:电赛项目里,传感器数据时有时无,控制指令偶尔失灵,屏幕上的波形跳得跟心电图一样,但就是找不到原因。你怀疑过硬件,换过线,甚至重新焊了板子,最后发现&#xf…

作者头像 李华
网站建设 2026/8/4 10:32:17

3步掌握Steam创意工坊下载:WorkshopDL完全使用指南

3步掌握Steam创意工坊下载:WorkshopDL完全使用指南 【免费下载链接】WorkshopDL WorkshopDL - The Best Steam Workshop Downloader 项目地址: https://gitcode.com/gh_mirrors/wo/WorkshopDL 你是否在Epic Games Store或GOG平台购买了游戏,却无法…

作者头像 李华
网站建设 2026/8/4 10:30:22

山海万灵 HarmonyOS 文化知识实战(06):知识图谱节点与推荐关系

在文化知识应用中,读者从一只神兽继续阅读时,下一张卡片不能只按列表顺序出现。山海万灵把神兽、地域和展厅放进同一组稳定标识:推荐项既带有目标节点,也带有“同展厅关联”“同区域关联”等可读原因。页面据此展示推荐卡&#xf…

作者头像 李华
网站建设 2026/8/4 10:29:20

相空间轨迹分类实战解析

from enum import Enum from dataclasses import dataclass from typing import List, Optional, Tuple import numpy as npclass PhaseSpacePattern(Enum):"""三类相空间形态枚举标签"""STABLE_ELLIPSE "stable_ellipse"BIFURCATION…

作者头像 李华