返回博客

/ 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。

19 minDeepSeek · FlashMLA · MLA · CUDA

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 对第 tt 个 token 生成各个 head 的 KtK_tVtV_t,并在自回归推理时把它们全部缓存下来。历史序列越长、并发请求越多,KV Cache 占用越大。

MLA 不直接缓存展开后的多头 K/V,而是先把两者联合压缩为低秩 latent:

ctKV=WDKVht,ctKVR512.c_t^{KV}=W^{DKV}h_t, \qquad c_t^{KV}\in\mathbb{R}^{512}.

内容 Key 和 Value 在数学上可以由这个 latent 上投影得到:

kt,hC=WhUKctKV,vt,hC=WhUVctKV.k_{t,h}^{C}=W_h^{UK}c_t^{KV}, \qquad v_{t,h}^{C}=W_h^{UV}c_t^{KV}.

RoPE 部分不能直接被同一种低秩表示吸收,因此 MLA 还为每个 token 单独缓存一个 64 维的 decoupled RoPE key:

ktRR64.k_t^R\in\mathbb{R}^{64}.

于是 DeepSeek-V2/V3 系列的典型 MLA Cache 不再是“每个 KV head 各存一份 K 和 V”,而是:

ctKV512ktR64576 个元素/token.\underbrace{c_t^{KV}}_{512} \oplus \underbrace{k_t^R}_{64} \quad\Longrightarrow\quad 576\text{ 个元素/token}.

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 可以写成两部分:

qt,h=[qt,hC;qt,hR],128+64=192.q_{t,h}= \left[q_{t,h}^{C};q_{t,h}^{R}\right], \qquad 128+64=192.

若按照定义计算内容部分,需要先从 ctKVc_t^{KV} 恢复每个 head 的 kt,hCk_{t,h}^{C}

(qi,hC)Tkt,hC=(qi,hC)TWhUKctKV.\left(q_{i,h}^{C}\right)^\mathsf{T}k_{t,h}^{C} = \left(q_{i,h}^{C}\right)^\mathsf{T}W_h^{UK}c_t^{KV}.

利用矩阵乘法结合律,可以把 WhUKW_h^{UK} 吸收到 Query 一侧:

(qi,hC)TWhUKctKV=((WhUK)Tqi,hC)TctKV.\left(q_{i,h}^{C}\right)^\mathsf{T}W_h^{UK}c_t^{KV} = \left(\left(W_h^{UK}\right)^\mathsf{T}q_{i,h}^{C}\right)^\mathsf{T}c_t^{KV}.

定义吸收后的 Query:

q~i,hC=(WhUK)Tqi,hCR512,\widetilde q_{i,h}^{C} = \left(W_h^{UK}\right)^\mathsf{T}q_{i,h}^{C} \in\mathbb{R}^{512},

再拼接 64 维 RoPE Query:

q~i,h=[q~i,hC;qi,hR]R576.\widetilde q_{i,h} = \left[\widetilde q_{i,h}^{C};q_{i,h}^{R}\right] \in\mathbb{R}^{576}.

这个变换通常称为 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 latent

FlashMLA 返回的也是每个 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=576dv=512,也解释了为什么直接套用假设 dk=dvd_k=d_v 的通用 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 91

kernel 根据 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_seqlenstopk_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 3

GPU 端 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 一样,不会把完整的 QKTQK^\mathsf{T} 写回显存。对一个 score block SiS_i,记录局部最大值:

mi=max(Si),m_i=\max(S_i),

局部指数和:

i=jexp(Si,jmi),\ell_i=\sum_j\exp(S_{i,j}-m_i),

以及未归一化的局部 Value 累加:

Oi=jexp(Si,jmi)Vi,j.O_i=\sum_j\exp(S_{i,j}-m_i)V_{i,j}.

合并两个 split 时,先得到共同最大值:

m=max(m1,m2),m=\max(m_1,m_2),

再将两边缩放到同一个指数基准:

=em1m1+em2m2,\ell = e^{m_1-m}\ell_1+e^{m_2-m}\ell_2, O=em1mO1+em2mO2.O = e^{m_1-m}O_1+e^{m_2-m}O_2.

最终归一化输出为:

Attention(Q,K,V)=O.\operatorname{Attention}(Q,K,V)=\frac{O}{\ell}.

这个公式既保证数值稳定,也让不同 CTA 无须共享完整 score matrix,只需交换紧凑的 softmax 状态和局部输出。Split-KV 能成立的数学基础,正是 Online Softmax 的可组合性。

8. Seesaw Scheduling:一个 Output Tile 怎样让两类计算重叠#

Dense MLA Decode 的一个关键困难来自超宽输出。一个典型 Output tile 为:

64×512.64\times512.

若使用 FP32 累加,它需要:

64×512=3276864\times512=32768

个 32-bit register。Hopper 每个 SM 约有 65536 个 32-bit register,因此一个 Output tile 已经消耗一半寄存器预算。像 FlashAttention-3 的 ping-pong schedule 那样同时常驻两个完整 Output tile,会立刻遇到资源约束。

FlashMLA 的做法是沿 Value 维把一个输出拆成两半:

O=[OL,OR],OL,ORR64×256.O=[O_L,O_R], \qquad O_L,O_R\in\mathbb{R}^{64\times256}.

两个 warpgroup 分别持有 OLO_LORO_R,同时交错处理相邻的两个 KV tile:

  • warpgroup 0 对 K0K_0 做 QK、max、exp 和左半边 PV;
  • warpgroup 1 对 K1K_1 做同样工作,并更新右半边;
  • 两边共享 Online Softmax 的 running max 与 rescale 因子;
  • 随后交叉补齐 V0RV_{0R}V1LV_{1L} 对另一半 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 SchedulingDense kernel deep dive

这项设计的价值不只是节省寄存器。QK/PV 主要使用 Tensor Core,max/exp/rescale 主要使用 CUDA Core;交错调度让两类执行单元有机会并行工作,同时给异步数据搬运留下覆盖窗口。

9. 细粒度 TMA-GEMM 流水#

一个 K tile 的形状是 64×57664\times576。FlashMLA 不等待整个 tile 搬入 Shared Memory 后才启动计算,而是把 576 维拆成 9 个 64×6464\times64 子块:

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_FIRST cache 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 数为 hqh_q,一次验证的 Query token 数为 sqs_q,KV 长度为 sks_k,K/V 维度分别为 dkd_kdvd_v。MLA Decode 的计算量近似为:

2hqsqsk(dk+dv).2h_qs_qs_k(d_k+d_v).

sks_k 远大于 Query 规模时,BF16 KV 读取量近似为:

2skdk bytes.2s_kd_k\ \text{bytes}.

因此算术强度约为:

FLOPsByteshqsqdk+dvdk2hqsq,\frac{\text{FLOPs}}{\text{Bytes}} \approx h_qs_q\frac{d_k+d_v}{d_k} \approx 2h_qs_q,

其中最后一步使用了 dk=576,dv=512d_k=576,d_v=512

DeepSeek 的典型 Decode 配置是 hq=128h_q=128;开启 MTP 或推测解码时,sqs_q 还可能大于 1。相同的一份共享 KV 会被许多 Query head 重用,因此计算量随 hqsqh_qs_q 增长,KV 读取量却不会同速增长。官方分析据此指出,在其不使用 Tensor Parallel 的 Decode 实例上,hqsq128h_qs_q\ge128 时可以进入 compute-bound 区域。Dense kernel deep dive

反过来,如果使用 Tensor Parallel 把 Query head 分散到多张卡,每张卡的 hqh_q 下降,同一个算子就更容易回到 memory-bound。瓶颈属于具体映射后的工作负载,不能只由“这是 Decode”来判断。

11. Sparse FlashMLA:只算上层选中的 token#

DeepSeek Sparse Attention 先由上层为每个 Query 选择 Top-K 历史 token,再把物理索引传给 FlashMLA:

indices: [batch, s_q, topk]

每个索引直接编码 KV Cache 的物理位置:

index=page_id×page_block_size+offset.\text{index} = \text{page\_id}\times\text{page\_block\_size} +\text{offset}.

因为 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 latent512 个 float8_e4m3512 bytes
Scale每 128 个 FP8 值共享 1 个 FP32 scale,共 4 个16 bytes
RoPE64 个 BF16,不量化128 bytes
合计656 bytes

完整 BF16 表示需要:

576×2=1152 bytes/token.576\times2=1152\text{ bytes/token}.

因此该布局的单 token Cache 大小下降约:

1656115243.1%.1-\frac{656}{1152}\approx43.1\%.

只量化前 512 维并非偶然。NoPE latent 采用 1×1281\times128 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,然后:

  1. 把结果写入自己的 Shared Memory;
  2. 通过 Distributed Shared Memory 和 st.async 发给另一个 CTA;
  3. 使用 cluster transaction barrier 同步交换;
  4. 最终让两个 CTA 都获得完整的 BF16 KV。

这种跨 CTA 交换被称为 Crossover。它没有消除反量化,而是利用“两个 CTA 需要相同 KV”这一数据复用关系,把准备工作平摊到两边。

官方在 batch_size=128h_q=128s_q=2topk=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 的关系#

两者共享同一类基本思想,但优化对象不同:

对比项FlashAttentionFlashMLA
主要目标通用 MHA/GQADeepSeek MLA 及相关 Attention 路径
典型 K/V 维度常见 64/128,通常相同Decode 中 K=576、V=512
KV 表示独立 K 与 V512 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 pipelineSeesaw、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_tableindices、有效序列长度,以及满足复用条件的 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 layer

16. 阅读源码时最值得抓住的四条主线#

FlashMLA 的源码按架构、阶段与稠密/稀疏模式展开,模板和硬件指令很多。比起从第一条 WGMMA 指令逐行阅读,更有效的顺序是:

  1. 先看 Python 接口与测试。 确认 Q、Cache、索引、输出和 LSE 的真实形状与语义;
  2. 再看 scheduler 与 combine。 理解一个逻辑请求怎样变成 split,以及 Online Softmax 状态怎样归约;
  3. 然后看数据布局。 追踪 576 维中哪些参与 QK、哪些参与 PV,FP8 scale 和 BF16 RoPE 如何排布;
  4. 最后看 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 融合
长上下文与变长 batchCTA 数不足、请求拖尾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 共享工作

这套方法可以浓缩成三个问题:

  1. 模型究竟要求保存什么数据,而不是习惯上保存什么?
  2. 哪些线性变换可以提前吸收,从而改变运行时计算形态?
  3. 当前瓶颈究竟是带宽、算力、寄存器、延迟,还是数据准备?

FlashMLA 给出的最终答案并不是一条万能 kernel,而是一组针对不同架构、阶段、精度和稀疏模式的实现。它最重要的启示也正在这里:

高性能 Attention 不是把公式翻译成 CUDA,而是从缓存表示开始,重新安排数据生命周期、并行粒度和片上资源。

参考资料#