/ DeepSeek
[DeepSeek-3] FlashMLA:把 MLA 的 KV Cache 优势变成 GPU 吞吐
从 Weight Absorption、576/512 数据布局和 Paged KV Cache 出发,系统拆解 FlashMLA 的 Split-KV、Online Softmax、Seesaw Scheduling、稀疏 Attention、FP8 KV Cache 与 Crossover。
Multi-head Latent Attention(MLA)解决了一个模型结构问题:怎样避免为每个历史 token 保存完整的多头 Key 和 Value。
但缓存变小不等于推理自然变快。MLA 的 latent 表示、解耦 RoPE、超宽 Value 维度、Paged KV Cache 和变长请求,会共同形成一种不同于常规 MHA 的计算形态。如果仍然按照通用 Attention 的方式展开 K/V,再调用多个独立算子,缓存压缩省下的带宽很可能又被中间张量和 kernel launch 开销吃掉。
FlashMLA 正是为这一步而生。它不是新的 Attention 架构,而是 DeepSeek 为 MLA 编写的一组高性能 CUDA kernel:
- 模型层面的 MLA 决定“缓存什么”;
- Weight Absorption 决定“以什么等价形式计算”;
- FlashMLA 决定“怎样在 GPU 上切块、搬运、调度和归约”。
一句话概括:
MLA 压缩 KV Cache,FlashMLA 则让压缩后的 Attention 不必重新展开,也能高效执行。
当前官方仓库同时包含 Dense MLA Decode、Sparse MLA Decode、Sparse MLA Prefill,以及面向 SM100 的 Dense MHA Prefill/Backward;其中稀疏算子服务于 DeepSeek Sparse Attention。本文重点讨论最能体现 FlashMLA 设计特点的 Decode 路径。官方仓库
1. 从 MLA 的缓存表示开始#
标准 Multi-Head Attention 对第 个 token 生成各个 head 的 和 ,并在自回归推理时把它们全部缓存下来。历史序列越长、并发请求越多,KV Cache 占用越大。
MLA 不直接缓存展开后的多头 K/V,而是先把两者联合压缩为低秩 latent:
内容 Key 和 Value 在数学上可以由这个 latent 上投影得到:
RoPE 部分不能直接被同一种低秩表示吸收,因此 MLA 还为每个 token 单独缓存一个 64 维的 decoupled RoPE key:
于是 DeepSeek-V2/V3 系列的典型 MLA Cache 不再是“每个 KV head 各存一份 K 和 V”,而是:
DeepSeek-V2 报告将这种联合低秩压缩列为 KV Cache 大幅下降的关键来源,并报告相对其对照架构约 93.3% 的缓存缩减。DeepSeek-V2 技术报告
这里有一个容易混淆的点:512 是 latent 内容维度,64 是 RoPE 维度,576 是实际进入 Decode kernel 的 Q/K 维度;它们不是三组彼此独立的 K、V。
2. 为什么原始 Query 是 192 维,Decode 输入却是 576 维#
DeepSeek 的原始 Query head 可以写成两部分:
若按照定义计算内容部分,需要先从 恢复每个 head 的 :
利用矩阵乘法结合律,可以把 吸收到 Query 一侧:
定义吸收后的 Query:
再拼接 64 维 RoPE Query:
这个变换通常称为 Weight Absorption。它避免在每个 Decode step 中为所有历史 token 展开多头 K,也把 MLA Decode 改写成一种特殊的 Multi-Query Attention:
模型定义中的 MLA
Q: 128 个 head,每个 head 为 128 NoPE + 64 RoPE
K/V: 由共享的 512 维 KV latent 上投影得到
Weight Absorption
|
v
FlashMLA Decode 看到的等价形式
Q: 128 heads x 576
K: 1 shared head x (512 latent + 64 RoPE)
V: 1 shared head x 512 latentFlashMLA 返回的也是每个 Query head 的 512 维 latent Attention 结果。模型层随后还要完成 Value 上投影等操作。换言之,FlashMLA 是 Attention 核心算子,不是完整的 MLA Layer。
3. 一个 576 维缓存为什么同时充当 K 和 V#
Dense Decode 的典型张量形状如下:
q: [batch, s_q, h_q, 576]
k_cache: [num_blocks, page_block_size, h_kv, 576]
block_table: [batch, max_blocks_per_sequence]
cache_seqlens: [batch]
out: [batch, s_q, h_q, 512]
lse: [batch, h_q, s_q]对某个 Query,参考计算可简化为:
score = q @ kv[..., :576].transpose(-1, -2)
score *= softmax_scale
prob = softmax(score, dim=-1)
out = prob @ kv[..., :512]同一块 Cache 在两个 GEMM 中扮演不同角色:
- QK 使用完整 576 维,其中前 512 维是 latent 内容,后 64 维是 RoPE;
- PV 只使用前 512 维,因为输出仍处于 latent Value 空间;
- K 和 V 不再是两块独立缓存,而是同一表示的不同切片。
这解释了 FlashMLA 测试中反复出现的 d=576 与 dv=512,也解释了为什么直接套用假设 的通用 Attention kernel 并不理想。
4. Paged KV Cache 如何进入同一个融合算子#
线上推理不会为每个请求预留一段最大长度的连续 Cache。FlashMLA 的 Dense Decode 原生接受 Paged KV Cache:
k_cache: [num_physical_blocks, 64, h_kv, 576]
block_table: [batch, max_num_blocks]例如,一个逻辑序列的连续页面可以分散在不同物理块:
Request A / logical block 0 -> physical block 18
Request A / logical block 1 -> physical block 3
Request A / logical block 2 -> physical block 91kernel 根据 block_table 完成地址转换,并通过 cache_seqlens 判断每个请求的有效长度。在官方 Dense Decode 默认测试中,page_block_size=64,同时覆盖变长序列、多个 Query token、MQA/GQA 和 causal mask。
它和 PagedAttention 的系统目标相同:按块分配和增长 Cache,降低预留浪费与显存碎片。FlashMLA 的特殊之处,是继续把以下工作融合在一条专用路径里:
Paged addressing
-> MLA QK
-> Online Softmax
-> latent PV
-> Split-KV partial result这避免了先把分页 Cache gather 成连续张量,再交给另一个 Attention kernel 的额外搬运。
5. 调用链与调度元数据#
当前 Python 入口 flash_mla_with_kvcache() 会根据 indices 是否存在选择路径:
flash_mla_with_kvcache()
|
+-- indices is None ------> dense_decode_fwd()
|
+-- indices is not None --> sparse_decode_fwd()
|
v
Tile Scheduler Metadata
|
v
Split-KV Kernel
|
v
Combine Kernel
|
v
out + lse调度元数据采用 lazy initialization。当前源码中的标准用法是:
sched_meta, _ = get_mla_metadata()
out, lse = flash_mla_with_kvcache(
q=q,
k_cache=k_cache,
block_table=block_table,
cache_seqlens=cache_seqlens,
head_dim_v=512,
tile_scheduler_metadata=sched_meta,
)get_mla_metadata() 先返回一个空的 FlashMLASchedMeta;第一次调用 Attention 时,接口才依据输入生成 GPU 调度信息。它可以在同一个 Decode step 的不同 Transformer 层之间复用,因为各层通常共享请求布局。
复用并不是无条件的。至少要保持 batch、s_q、Query/KV head 数、page block size、causal/精度模式和 Top-K 配置一致,cache_seqlens、topk_length 等影响实际工作量的值也不能变化。换到下一个 Decode step 后,序列长度已经增长,通常需要重新建立对应的调度状态。
值得注意的是,仓库 README 的示例仍保留旧版 get_mla_metadata(cache_seqlens, ...) 调用形式;当前 Python 接口为了兼容仍接受多余参数,但会忽略它们。阅读快速演进的 kernel 仓库时,测试和接口源码往往比 README 示例更接近真实行为。接口源码
6. 为什么长上下文 Decode 需要 Split-KV#
假设一个请求有数万个 KV token。如果一个 CTA 从头到尾处理整个请求,会出现三个问题:
- 请求数少时,可并行 CTA 数量不足,无法占满所有 SM;
- 变长 batch 中,短请求提前结束,长请求让少数 SM 拖尾;
- 单个 CTA 的工作时间过长,调度粒度太粗。
FlashMLA 将一个请求的 KV 区间继续切成多个 split:
Request 0 / KV[0:32768]
|-- Split 0
|-- Split 1
|-- Split 2
`-- Split 3GPU 端 Tile Scheduler 把 (request, query tile, kv split) 分配到可用 SM。每个 split 独立产生局部 Output、局部最大值和局部归一化信息,随后由 Combine kernel 合成最终结果。
Split-KV 的代价是需要中间缓冲:
out_accum: [total_num_splits, heads, queries, 512]
lse_accum: [total_num_splits, heads, queries]它用额外的临时存储和一次归约,换取更多并行度与更平衡的负载。是否应该增加 split 数,不是只由上下文长度决定,还要综合 batch、Query 数量、head 数和 GPU 的 SM 数量。
7. Online Softmax 如何让 split 结果可合并#
FlashMLA 和 FlashAttention 一样,不会把完整的 写回显存。对一个 score block ,记录局部最大值:
局部指数和:
以及未归一化的局部 Value 累加:
合并两个 split 时,先得到共同最大值:
再将两边缩放到同一个指数基准:
最终归一化输出为:
这个公式既保证数值稳定,也让不同 CTA 无须共享完整 score matrix,只需交换紧凑的 softmax 状态和局部输出。Split-KV 能成立的数学基础,正是 Online Softmax 的可组合性。
8. Seesaw Scheduling:一个 Output Tile 怎样让两类计算重叠#
Dense MLA Decode 的一个关键困难来自超宽输出。一个典型 Output tile 为:
若使用 FP32 累加,它需要:
个 32-bit register。Hopper 每个 SM 约有 65536 个 32-bit register,因此一个 Output tile 已经消耗一半寄存器预算。像 FlashAttention-3 的 ping-pong schedule 那样同时常驻两个完整 Output tile,会立刻遇到资源约束。
FlashMLA 的做法是沿 Value 维把一个输出拆成两半:
两个 warpgroup 分别持有 与 ,同时交错处理相邻的两个 KV tile:
- warpgroup 0 对 做 QK、max、exp 和左半边 PV;
- warpgroup 1 对 做同样工作,并更新右半边;
- 两边共享 Online Softmax 的 running max 与 rescale 因子;
- 随后交叉补齐 和 对另一半 Output 的贡献。
可以把大致时间线理解为:
time ----------------------------------------------------->
WG 0: QK(K0) -> softmax(P0) -> P0*V0L -> rescale -> P1*V1L
WG 1: QK(K1) -> softmax(P1) -> P1*V1R -> P0*V0R它像跷跷板一样,让两个 warpgroup 围绕同一个 Output tile 交替工作,因此被称为 Seesaw Scheduling。Dense kernel deep dive
这项设计的价值不只是节省寄存器。QK/PV 主要使用 Tensor Core,max/exp/rescale 主要使用 CUDA Core;交错调度让两类执行单元有机会并行工作,同时给异步数据搬运留下覆盖窗口。
9. 细粒度 TMA-GEMM 流水#
一个 K tile 的形状是 。FlashMLA 不等待整个 tile 搬入 Shared Memory 后才启动计算,而是把 576 维拆成 9 个 子块:
TMA copy block 0 done -> GEMM 0
TMA copy block 1 done -> GEMM 1
...
TMA copy block 8 done -> GEMM 8这种 Fine-grained TMA-GEMM Pipeline 缩短了“第一批数据就绪”到“第一条矩阵指令发射”的距离。即使算子整体处于 compute-bound,数据依赖造成的访存延迟仍可能让 Tensor Core 出现气泡,因此延迟隐藏依然重要。
同一版本还使用了:
EVICT_FIRSTcache hint,减少流式 KV 数据对 L2 的污染;- Programmatic Dependent Launch,让 Split-KV 与 Combine kernel 更紧密地衔接;
- GPU 端 Tile Scheduler,均衡变长请求与 split;
- WGMMA,执行 warpgroup 级异步矩阵乘。
官方在 H800 SXM5、CUDA 12.8 的微基准中报告 Dense Decode 在 memory-bound 配置最高约 3000 GB/s,在 compute-bound 配置最高约 660 TFLOPS。这些数字描述单算子特定配置下的吞吐,不等于完整模型的 tokens/s,也不能脱离频率、batch、序列长度和精度直接横向比较。官方仓库
10. Decode 为什么也可能 compute-bound#
“Prefill 是 compute-bound,Decode 是 memory-bound”是一个有用但不完整的经验判断。
设每个请求的 Query head 数为 ,一次验证的 Query token 数为 ,KV 长度为 ,K/V 维度分别为 和 。MLA Decode 的计算量近似为:
当 远大于 Query 规模时,BF16 KV 读取量近似为:
因此算术强度约为:
其中最后一步使用了 。
DeepSeek 的典型 Decode 配置是 ;开启 MTP 或推测解码时, 还可能大于 1。相同的一份共享 KV 会被许多 Query head 重用,因此计算量随 增长,KV 读取量却不会同速增长。官方分析据此指出,在其不使用 Tensor Parallel 的 Decode 实例上, 时可以进入 compute-bound 区域。Dense kernel deep dive
反过来,如果使用 Tensor Parallel 把 Query head 分散到多张卡,每张卡的 下降,同一个算子就更容易回到 memory-bound。瓶颈属于具体映射后的工作负载,不能只由“这是 Decode”来判断。
11. Sparse FlashMLA:只算上层选中的 token#
DeepSeek Sparse Attention 先由上层为每个 Query 选择 Top-K 历史 token,再把物理索引传给 FlashMLA:
indices: [batch, s_q, topk]每个索引直接编码 KV Cache 的物理位置:
因为 page id 已经包含在索引中,Sparse Decode 不再需要 block_table。kernel 的职责是:
load selected KV by indices
-> QK
-> Online Softmax
-> PV
-> output它不负责生成 Top-K。Token selector、稀疏索引维护和选择质量属于模型与上层推理系统。把“选哪些 token”和“怎样高效计算选中的 token”分开,是理解 Sparse FlashMLA 边界的关键。
稀疏并不自动等于更快。Top-K 太大时,节省的计算有限;Top-K 太小时,离散寻址、prologue/epilogue 和调度开销占比会上升。实际收益取决于上下文长度、Top-K、命中分布、Cache 格式和选择器本身的成本。
12. FP8 KV Cache:量化 512 维,保留 64 维#
Sparse Decode 使用一套专用 FP8 Cache 布局。每个 token 占 656 bytes:
| 区域 | 内容 | 大小 |
|---|---|---|
| NoPE latent | 512 个 float8_e4m3 | 512 bytes |
| Scale | 每 128 个 FP8 值共享 1 个 FP32 scale,共 4 个 | 16 bytes |
| RoPE | 64 个 BF16,不量化 | 128 bytes |
| 合计 | 656 bytes |
完整 BF16 表示需要:
因此该布局的单 token Cache 大小下降约:
只量化前 512 维并非偶然。NoPE latent 采用 tile 粒度的 scale,RoPE 部分则因对精度损失更敏感而保留 BF16。kernel 内部将 FP8 latent 反量化为 BF16,与 BF16 RoPE 拼接后,再以 BF16 Tensor Core 输入、FP32 累加执行 QK 和 PV。FP8 sparse deep dive
这再次说明,低精度 Cache 不是简单地给 tensor 换一个 dtype;布局、scale 粒度、敏感维度和消费该布局的 kernel 必须共同设计。
13. Crossover:当反量化比矩阵乘更慢#
H800 无法把 float8_e4m3 直接高效转换为 BF16。官方按指令吞吐估算,每个 KV token 大致需要:
- 约 34 cycles 完成 64 个 Query head 对应的 MMA;
- 约 50 cycles 完成 512 个 FP8 值的转换与缩放。
这意味着 Sparse Decode 可能不是 matrix-bound,而是 dequantization-bound:Tensor Core 已经完成工作,CUDA Core 还在准备下一批 BF16 数据。
FlashMLA 利用 MQA 的一个性质解决它:同一个 Query token 的 128 个 Query head 共享同一份 KV。于是把两个处理不同 Query head 的 CTA 组成 cluster:
CTA cluster
|-- CTA 0: query heads 0..63
`-- CTA 1: query heads 64..127每个 CTA 只加载并反量化一半 KV,然后:
- 把结果写入自己的 Shared Memory;
- 通过 Distributed Shared Memory 和
st.async发给另一个 CTA; - 使用 cluster transaction barrier 同步交换;
- 最终让两个 CTA 都获得完整的 BF16 KV。
这种跨 CTA 交换被称为 Crossover。它没有消除反量化,而是利用“两个 CTA 需要相同 KV”这一数据复用关系,把准备工作平摊到两边。
官方在 batch_size=128、h_q=128、s_q=2、topk=2048 的 H800 微基准中,报告 Crossover 将 FP8 Sparse Decode 从约 250 TFLOPS 提升到约 410 TFLOPS。FP8 sparse deep dive 这个优化依赖 Hopper 的 CTA Cluster、Distributed Shared Memory 和 cluster barrier,也说明高性能 kernel 的关键问题经常不是“少算一次乘法”,而是“谁准备数据、准备几次、怎样把结果交给消费者”。
14. FlashMLA 与 FlashAttention 的关系#
两者共享同一类基本思想,但优化对象不同:
| 对比项 | FlashAttention | FlashMLA |
|---|---|---|
| 主要目标 | 通用 MHA/GQA | DeepSeek MLA 及相关 Attention 路径 |
| 典型 K/V 维度 | 常见 64/128,通常相同 | Decode 中 K=576、V=512 |
| KV 表示 | 独立 K 与 V | 512 latent + 64 RoPE 的复合表示 |
| KV head | 多头或分组 | MLA Decode 典型为 1 个共享 head |
| 重点阶段 | 训练、Prefill 与 Decode 均有实现 | 最具代表性的是 MLA Decode |
| Paged Cache | 取决于具体接口与实现 | Dense Decode 直接纳入主路径 |
| 稀疏模式 | 不是核心统一接口 | 支持 token-level Top-K |
| 特殊调度 | 如 ping-pong、intra-warpgroup pipeline | Seesaw、Crossover |
共同点同样重要:
- 都不物化完整 Attention matrix;
- 都使用 tiling 与 Online Softmax;
- 都尽可能融合 QK、Softmax 和 PV;
- 都围绕 Tensor Core、异步搬运和片上存储组织流水。
FlashMLA 也明确受到 FlashAttention、Flash-Decoding 和 CUTLASS 的启发;仓库中的 Dense MHA Prefill 则更直接地进入了通用 MHA 的问题域。因此,把两者理解为互斥的竞争方案并不准确:FlashMLA 是在同一套 IO-aware Attention 思想上,针对 MLA 数据表示和 DeepSeek 工作负载做出的专门化。
15. FlashMLA 不负责什么#
FlashMLA 位于模型层与硬件之间,但它不是推理框架。通常不负责:
- Q/K/V projection;
- KV latent compression 和 Cache 写入;
- RoPE 计算;
- Weight Absorption 前后的线性层;
- 物理 Cache block 的分配、回收和换出;
- FP8 Cache 的量化写入;
- Sparse Top-K token 选择;
- 请求级 batching 与调度;
- Tensor Parallel 通信;
- Attention 之后的 Value 上投影。
上层需要先准备 absorbed Query、Paged KV Cache、block_table 或 indices、有效序列长度,以及满足复用条件的 scheduler metadata。FlashMLA 接管的是中间最密集、也最依赖数据布局的 Attention 计算。
这个边界可以表示为:
model / inference framework
projection -> RoPE -> cache management -> token selection
|
v
FlashMLA
address -> QK -> online softmax -> latent PV -> split combine
|
v
model / inference framework
value up-projection -> output projection -> next layer16. 阅读源码时最值得抓住的四条主线#
FlashMLA 的源码按架构、阶段与稠密/稀疏模式展开,模板和硬件指令很多。比起从第一条 WGMMA 指令逐行阅读,更有效的顺序是:
- 先看 Python 接口与测试。 确认 Q、Cache、索引、输出和 LSE 的真实形状与语义;
- 再看 scheduler 与 combine。 理解一个逻辑请求怎样变成 split,以及 Online Softmax 状态怎样归约;
- 然后看数据布局。 追踪 576 维中哪些参与 QK、哪些参与 PV,FP8 scale 和 BF16 RoPE 如何排布;
- 最后看 SM90 kernel schedule。 将寄存器、Shared Memory、warpgroup、TMA 和 CTA cluster 对应回前面的算法约束。
仓库中可优先定位这些目录:
flash_mla/flash_mla_interface.py Python API 与 Dense/Sparse 分发
csrc/api/ C++ / PyBind 调度入口
csrc/sm90/decode/dense/ Hopper Dense Decode
csrc/sm90/decode/sparse_fp8/ Hopper FP8 Sparse Decode
csrc/sm90/prefill/sparse/ Hopper Sparse Prefill
csrc/sm100/ Blackwell 相关实现
csrc/smxx/decode/ 跨架构 scheduler 与 combine
tests/ 参考语义、正确性与微基准配置
docs/ Dense 与 FP8 Sparse deep dive真正需要建立的不是“某条 CUDA 指令做了什么”的孤立记忆,而是下面这条因果链:
MLA 压缩 KV
-> Weight Absorption 形成 576/512 MQA
-> 超宽输出制造寄存器压力
-> Seesaw 拆分并交错累加
-> 长上下文需要 Split-KV
-> Online Softmax 使 split 可合并
-> 稀疏索引减少参与计算的 token
-> FP8 Cache 降低容量与带宽
-> Crossover 分摊反量化成本17. FlashMLA 最值得学习的系统方法#
FlashMLA 的价值不只是 660 TFLOPS 或 410 TFLOPS 两个峰值数字,而是展示了模型表示、数学等价变换与硬件调度如何层层咬合:
| 上层约束 | kernel 问题 | FlashMLA 的应对 |
|---|---|---|
| MLA 只缓存 latent + RoPE | 通用 MHA 布局不匹配 | 576 维 QK、512 维 PV 融合 |
| 长上下文与变长 batch | CTA 数不足、请求拖尾 | GPU Tile Scheduler + Split-KV |
| split 独立计算 | 局部 softmax 不能直接相加 | Online Softmax + Combine |
| 512 维超宽输出 | 寄存器不足以双缓冲 | Seesaw 拆分一个 Output tile |
| 576 维 K tile | 完整搬运后再算延迟过高 | 9 段 TMA-GEMM pipeline |
| Token-level sparse access | 地址离散、固定开销占比高 | 物理索引直达 Cache |
| FP8 Cache 需要反量化 | CUDA Core 跟不上 Tensor Core | 两 CTA Crossover 共享工作 |
这套方法可以浓缩成三个问题:
- 模型究竟要求保存什么数据,而不是习惯上保存什么?
- 哪些线性变换可以提前吸收,从而改变运行时计算形态?
- 当前瓶颈究竟是带宽、算力、寄存器、延迟,还是数据准备?
FlashMLA 给出的最终答案并不是一条万能 kernel,而是一组针对不同架构、阶段、精度和稀疏模式的实现。它最重要的启示也正在这里:
高性能 Attention 不是把公式翻译成 CUDA,而是从缓存表示开始,重新安排数据生命周期、并行粒度和片上资源。