FlashMLA 注意力内核源码走读:656 字节 KV 缓存背后的完整链路
【免费下载链接】FlashMLAFlashMLA: Efficient Multi-head Latent Attention Kernels项目地址: https://gitcode.com/GitHub_Trending/fl/FlashMLA
FlashMLA 注意力内核库是 DeepSeek 面向多头潜在注意力(MLA)的高性能 GPU 实现,运行于 Hopper/Blackwell 架构,驱动 DeepSeek-V3 系列模型的解码与预填充。它最鲜明的两个数字:每个 token 的 FP8 KV 缓存只占 656 字节,密集解码内核在 H800 SXM5 上跑到 660 TFLOPS。下面这次 FlashMLA 源码解析不按模块罗列,而是跟随一次解码请求,沿“调度 → 数据加载 → 计算 → 合并”的路径把整条链路走一遍。
调度:解码请求开工前的“预分单”
解码阶段每个请求只有 1 个 q token,却要扫一条很长的 KV 缓存——128K 上下文的请求,工作量是 1K 请求的一百倍。如果运行时临时把请求分给各个 SM(流式多处理器,GPU 的并行执行基本单元),负载必然严重不均。
FlashMLA 的做法是把分单提前做:
- 一个小型内核(
run_get_decoding_sched_meta_kernel)先把“(请求, KV 块区间)”这类作业单元摊给所有 SM,生成 tile scheduler 元数据。它解决的是 SM 空转问题,换来的是各 SM 工作量大致相等。 - 主内核
splitkv_mla只负责按元数据领作业,不做运行时仲裁,省去了调度本身的开销。 - 长序列会被切成多段(split-KV),不同段分给不同 SM 并行算,最后再合并。代价是多一个合并步骤,换来的是单请求延迟随上下文长度近似线性增长被摊平。
- 元数据在形状与序列长度不变时可以跨调用复用,避免每层重复计算。
相关实现可看 内核源码 与 新内核深入文档。
解码阶段调用示例
from flash_mla import get_mla_metadata, flash_mla_with_kvcache sched_meta, num_splits = get_mla_metadata( cache_seqlens, s_q * h_q // h_kv, h_kv, h_q, is_fp8, topk) o, lse = flash_mla_with_kvcache( q, kvcache, block_table, cache_seqlens, 512, sched_meta, num_splits, False, is_fp8_kvcache, indices)用大白话说:先一次性算好“活怎么分”,之后每个解码步把 query、KV 缓存和稀疏索引递进去,直接拿到注意力输出o和用于合并的lse。
数据加载:FP8 KV 缓存的 656 字节布局
DeepSeek-V3.2 把上下文从 64K 拉到 128K,单个 128K token 请求的 BF16 KV 缓存要 576 × 2 × 62 × 128 × 1024 ≈ 8.72 GiB,小 batch 下极易 OOM。FlashMLA 的答案是细粒度量化:对每个 token KV 的前 512 维做 1×128 的 tile 级量化,压缩后的 FP8 KV 缓存(FP8 为 8 位浮点格式,float8_e4m3)把占用近乎减半,同时保住了精度。
逐字段看字节布局
| 字段 | 值 | 含义 |
|---|---|---|
| 量化 NoPE 部分 | 512 字节 = 512 个float8_e4m3 | 576 维 KV 的前 512 维,压到 8 位存储 |
| 缩放因子 | 16 字节 = 4 个float32 | 每 128 个 FP8 值共享一个缩放因子,tile 级量化 |
| RoPE 部分 | 128 字节 = 64 个bfloat16 | 后 64 维,对精度损失敏感,故意不量化 |
| 合计 | 656 字节 | 一个 token 的 KV 缓存 |
- 内核先把 512 个 FP8 反量化回 bfloat16,与 64 个 RoPE 值拼成完整 576 维向量,矩阵乘法全程用 bfloat16 做、float32 累加。代价是一次反量化开销,换来的是存储减半、计算无损。
- 加载侧把 64×576 的 K 块拆成 9 次 TMA 复制(TMA 即张量内存加速器,NVIDIA 的异步数据搬运引擎,类似 DMA),每次 64×64;某片一到就发对应 GEMM,用流水掩盖显存延迟。
- TMA 复制附带
EVICT_FIRST缓存提示,把用完的数据标记为“优先逐出”,给后续复用的数据腾 L2 空间,实测提高了 L2 命中率。
计算环节一:Crossover 机制,把反量化开销砍半
这一节是整条链路上最“疼”的地方,按挑战、方案、收益拆开看。
为什么反量化成了瓶颈
H800 无法直接把float8_e4m3转成bfloat16,一个 token 的反量化要走四步:FP8→half→float32→bfloat16,再乘缩放因子。按 NVIDIA 官方吞吐数据折算,每 token 至少约50 个周期;而 Tensor Core(专门做矩阵乘加的硬件单元)处理 64 个 query 头对应的 MMA 只要 64 × (576+512) × 2 / 4096 ≈34 个周期。50 > 34,内核处于反量化受限状态,Tensor Core 被迫等数据。
两个 CTA 分摊一份 KV
- 关键事实:MQA 模式下,同一 query token 的 128 个 query 头读的是同一份 K/V。每个 CTA(CUDA 线程块,被调度到 SM 上并行运行的一组线程)只负责 64 个头,正好可以分工。
- 用 Hopper 的 CTA cluster(CTA 集群,一组可互相直接访问对方共享内存的 CTA)发射 2 个 CTA。
- 每个 CTA 用 128 位宽的
__ldg宽加载,只取半份量化 K/V,反量化自己这半份,写入自己的共享内存。 - 同时用
st.async异步把这半份写进对方的共享内存,再靠 cluster 事务栅栏同步。 - 同步结束后,两个 CTA 的共享内存里都有完整的反量化 K/V,各取所需开始 MMA。
收益很直观:每 CTA 反量化量减半,吞吐翻倍。没有 Crossover 的旧版 FP8 稀疏解码只有 250 TFLOPS,加上后到 410 TFLOPS,过程细节见 Hopper FP8 稀疏解码深入文档。
计算环节二:Seesaw 调度,用一块输出矩阵错峰交接
FlashAttention-3 的 ping-pong 调度需要两套输出矩阵交替推进,但这里放不下:一个 64×512 的输出矩阵要占 32,768 个 32 位寄存器,而单个 SM 总共只有 65,536 个——一套占半,两套必然溢出。
💡 你可以把它理解成:工厂只有一张大工作台放成品(寄存器),订单却源源不断。办法是把台面从中间劈成左右两半,派两队人(warpgroup)错开干活:A 队的装配线(Tensor Core)忙着算 K₁ 时,A 队的收尾活(CUDA Core 上的 softmax 与重缩放)由 B 队的空档顶上,反之亦然。两队像跷跷板两头此起彼伏,所以叫 Seesaw 调度。
拆成流程,一轮交接大约 5 步:
- 把 64×512 输出矩阵纵向劈成 o_L、o_R(各 64×256),分别放在两个 warpgroup 的寄存器里。
- 两队并行算两个 KV 块 K₀、K₁ 的 QKᵀ,得到注意力分数 p₀、p₁。
- A 队先做 p₀ 的在线 softmax(更新 running max 与 scale₀),并更新自己那一半:o_L ← o_L·scale₀ + p₀·V₀L。
- 交叉交接:B 队按复合缩放更新 o_R 并累加 p₁·V₁R;同时 A 队补 o_R 的 p₀·V₀R 部分,B 队对称地补 o_L。
- 如此循环推进。它在数学上等价于 FlashAttention 的在线 softmax,却只用一块输出矩阵。
它解决的问题正是 34 周期 vs 50 周期之外的另一半矛盾:CUDA Core 的 softmax 工作与 Tensor Core 的矩阵乘法互相等待。错峰之后两者充分重叠,同时数据一用完就能发出下一块的 TMA 复制,把访存也盖进计算窗口。最终实测达到约80% 的 Tensor Core 利用率(相对降频后的理论峰值)与3 TB/s带宽。
合并与性能账本
splitkv_mla算完的每一段各自留下部分输出与 lse,combine内核用 lse 把它们归一成最终结果。- 两个内核通过 Programmatic Dependent Launch(程序化依赖启动,让后一个内核在前一个收尾阶段就开始准备的启动机制)重叠执行,省掉一次完整的内核切换间隙。
- tile scheduler 与 PDL 组合的效果:SM 之间不抢活、内核之间不空转,长上下文下尤其明显。
性能账本(H800 SXM5 除非另注)
| 场景 | 指标 | 数值 | 一句话说明 |
|---|---|---|---|
| 密集解码(计算受限) | TFLOPS | 660 | 新版内核较旧版提升 5%~15% |
| 密集解码(访存受限) | 带宽 | 3000 GB/s | 逼近 H800 约 3.35 TB/s 的理论上限 |
| 密集解码(新内核) | Tensor Core 利用率 / 带宽 | ~80% / 3 TB/s | 利用率相对降频后理论峰值 |
| FP8 稀疏解码(topk=2048) | TFLOPS | 410 | batch=128、128 头的计算受限配置 |
| FP8 稀疏解码(topk=32768) | TFLOPS | 460 | topk 更大,前后处理占比下降 |
| FP8 稀疏解码(无 Crossover) | TFLOPS | 250 | 对照基线,量化 Crossover 的收益 |
| 稀疏预填充(H800 / B200) | TFLOPS | 640 / 1450 | 驱动 DeepSeek-V3.2 的稀疏注意力 |
| 稠密 MHA 预填充(B200) | TFLOPS | 前向 1460 / 反向 1000 | NVIDIA 报告值 |
回头看这条链路:tile scheduler 把活摊匀,FP8 加 Crossover 解掉反量化瓶颈,Seesaw 把 Tensor Core 喂饱,combine 收口合并。对做 LLM 推理加速的人而言,“低精度存、高精度算、调度上错峰交接”这三招,是这份源码里最值得直接搬走的部分。
【免费下载链接】FlashMLAFlashMLA: Efficient Multi-head Latent Attention Kernels项目地址: https://gitcode.com/GitHub_Trending/fl/FlashMLA
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考