FlashMLA 新内核深度解析:Seesaw 调度、细粒度 TMA 流水线与 H800 上的 660 TFlops MLA 解码
发布时间:2026/9/15 18:28:05 锦皓数字建站

FlashMLA 新内核深度解析Seesaw 调度、细粒度 TMA 流水线与 H800 上的 660 TFlops MLA 解码【免费下载链接】FlashMLAFlashMLA: Efficient Multi-head Latent Attention Kernels项目地址: https://gitcode.com/GitHub_Trending/fl/FlashMLA本篇技术指南以 FlashMLA 仓库官方博客 docs/20250422-new-kernel-deep-dive.md 为核心骨架深入剖析 2025.04.22 发布的新版 MLA 解码内核从理论层面论证解码阶段注意力为何会落入计算受限compute-bound区间到寄存器约束下诞生的 Seesaw 调度 数学变换再到细粒度 TMA 流水线、缓存提示、Programmatic Dependent Launch 与 Tile Scheduler 等实现细节。读完本文你将掌握新内核从瓶颈分析、调度设计到源码实现splitkv_mla.cuh的完整技术脉络并能在本仓库中用测试与基准脚本复现其性能。背景为什么需要一个新的 MLA 内核FlashMLA 的旧版内核已经在 MLA 解码上取得了亮眼成绩在内存受限memory-bound配置下达到约 3000 GB/s 带宽在计算受限配置下达到约 580 TFlops。而新版内核的目标是把计算受限场景的数字进一步推高最高可达约660 TFlops——相比旧版约提升 5%~15%。官方 README 的 News 中同样确认了这一性能更新并强调新版本接口与旧版本完全兼容升级即可直接获得性能提升。需要先澄清一个直觉上的矛盾MLA 是解码阶段单 token 生成的注意力内核而解码通常被认为是典型的访存密集场景。新版内核却明确以计算受限为主要优化对象并且设计决策几乎全部围绕如何喂饱 Tensor Core 展开。下面先从算法层面回答为什么 MLA 解码核会计算受限。MLA 算法瓶颈的理论分析为什么解码核也会计算受限GPU 内核可以粗略分为两类受浮点运算能力FLOPs限制的计算受限内核与受显存带宽限制的内存受限内核。判定方法就是计算每字节访存对应的浮点运算量FLOPs/byte再与 GPU 的实际能力对比。设 q 头数为 $h_q$每个请求的 q token 数为 $s_q$若关闭 MTP / 投机解码则为 1每个请求的 kv token 数为 $s_k\ (s_k \gg h_q s_q)$K、V 的头维度分别为 $d_k$、$d_v$。则运算量约为 $2 (h_q s_q \cdot d_k \cdot s_k h_q s_q \cdot s_k \cdot d_v) 2 h_q s_q s_k (d_kd_v)$访存量字节约为 $\mathop{\text{sizeof}}(\text{bfloat16}) \times (h_q s_q d_k s_k d_k h_q s_q d_v) \approx 2s_k d_k$因为 $s_k \gg h_q s_q$KV 缓存读入占主导且 MLA 中 K 与 V 共用同一份缓存因此计算-访存比约为 $h_q s_q \cdot \frac{d_kd_v}{d_k} \approx 2 h_q s_q$。NVIDIA H800 SXM5 的峰值显存带宽为 3.35 TB/s峰值算力为 990 TFlops但由于降频本例中时钟降到约 1600 MHz实际可用峰值约865 TFlops。由此得到临界条件当 $h_qs_q \ge \frac{1}{2} \cdot \frac{865}{3.35} 128$ 时内核是计算受限的否则为内存受限。DeepSeek 的在线推理系统在解码实例上不使用张量并行Tensor Parallel即 $h_q 128$因此 MLA 解码内核天然落入计算受限区间。这就是新内核一切优化的出发点必须让 Tensor Core 尽可能忙起来。新内核的高层调度设计寄存器约束为什么 FlashAttention-3 的 ping-pong 不能直接套用要充分压榨计算资源需要做到两件事一是把 CUDA Core 运算与 Tensor Core 运算重叠起来二是把访存与计算重叠起来让 Tensor Core 持续处于忙碌状态。FlashAttention-3 论文提出的 ping-pong 调度与 warpgroup 内 GEMM-softmax 流水正是为此设计的但这里无法直接套用根本原因是寄存器资源约束。由于 WGMMAwarpgroup 级矩阵指令的要求输出矩阵每一轮 mainloop 中都要被缩放并累加类似 FlashAttention 的在线 softmax 算法必须存放在寄存器中。一个 $64 \times 512$ 的输出矩阵要占用 32,768 个 32 位寄存器而每个 SM 总共只有 65,536 个 32 位寄存器意味着每个 SM 只能容纳一个输出矩阵——这直接排除了准备两个输出矩阵、让它们交替被 CUDA Core 和 Tensor Core 处理的经典乒乓方案。原文档在此特意留了一个思考题读者可以暂停一下想想有没有比官方更好的方案。从源码可以进一步印证这份资源预算在 config.h 中BLOCK_SIZE_M 64、HEAD_DIM_V 512即每个 CTA 处理 64 个 q 头行× 512 维输出。64 × 512的 fp32 累加器共 32,768 个值若由单个 128 线程的 warpgroup 独占平均每线程 256 个寄存器几乎触及 SM90 的每线程寄存器上限——这也是文档所说对一个 warpgroup 而言太大的原因。Seesaw 调度单输出矩阵的 ping-pong 变体官方给出的方案是在 FlashAttention 在线 softmax 与累加之外再做一层数学变换。每一步取两个 KV 块记为 $K_0, K_1, V_0, V_1$因为输出矩阵独占 32,768 个寄存器对一个 warpgroup 太多于是把输出矩阵纵向一分为二为 $O_L$ 与 $O_R$各 $64 \times 256$同理把 $V_0, V_1$ 拆成 $V_{0L}, V_{0R}, V_{1L}, V_{1R}$各 $64 \times 256$。$O_L$ 常驻 warpgroup 0 的寄存器$O_R$ 常驻 warpgroup 1 的寄存器。完整的 12 步算法如下为简洁起见假设单 q 头方括号内数字表示执行该步的 warpgroup维护一个跨两个 warpgroup 共享的运行最大值 $m$初始化为 $-\infty$与输出矩阵 $\vec o_L, \vec o_R$初始化为 0。[0] 计算 $\vec p_0 \vec q K_0^\intercal / qk_scale$。[1] 计算 $\vec p_1 \vec q K_1^\intercal / qk_scale$。[0] 计算 $mp_0 \max(\vec p_0)$、$m_new_0 \max(m, mp_0)$、$scale_0 \exp(m_new_0 - m)$并更新 $m \gets m_new_0$。[0] 对 $\vec p_0$ 做 softmax$\vec p_0 \gets \exp(\vec p_0 - m_new_0)$。[0] 更新 $\vec o_L \gets \vec o_L \cdot scale_0 \vec p_0 V_{0L}$。[1] 计算 $mp_1 \max(\vec p_1)$、$m_new_1 \max(m, mp_1)$、$scale_1 \exp(m_new_1 - m)$并更新 $m \gets m_new_1$。[1] 对 $\vec p_1$ 做 softmax$\vec p_1 \gets \exp(\vec p_1 - m_new_1)$。[1] 更新 $\vec o_R \gets \vec o_R \cdot (scale_0 \cdot scale_1) \vec p_1 V_{1R}$。[0] 更新 $\vec p_0 \gets \vec p_0 \cdot scale_1$。[1] 更新 $\vec o_R \gets \vec o_R \vec p_0 V_{0R}$。[0] 更新 $\vec o_L \gets \vec o_L \cdot scale_1 \vec p_1 V_{1L}$。这一调度可以看作只用一份输出矩阵的 ping-pong 变体官方称之为seesaw跷跷板调度数学上与 FlashAttention 的在线 softmax 算法完全等价。它的价值在于CUDA Core 与 Tensor Core 重叠两个 warpgroup 交替执行 softmax/缩放CUDA Core 工作与 PV GEMMTensor Core 工作互相掩盖延迟访存与计算重叠一旦某份数据不再被需要就可以立即发射对应的 TMATensor Memory Accelerator指令让下一次拷贝与当前计算并行。完整的调度时序图见下注意在 MLA 中$K$ 与 $V$ 是同一份数据的不同名字即 KV 缓存同时充当 K 与 VFlashMLA 新内核完整调度时序图seesaw 调度在源码中seesaw 调度正是由flash_fwd_splitkv_mla_kernelsplitkv_mla.cuh里两个分支的wg0_subroutine/wg1_subroutine实现的warpgroup 0 负责 $P_0 QK_0^\intercal$ 与 $O_L$ 的 PV 累加warpgroup 1 负责 $P_1 QK_1^\intercal$ 与 $O_R$ 的 PV 累加NamedBarrier如sScale0Ready、sScale1Ready、sP0Ready、rO1sP0sV0RIssued与cute::warpgroup_wait精确编排了两个 warpgroup 之间的数据依赖而共享内存中的sScale0/sScale1正是算法步骤里跨 warpgroup 传递的 scale 因子sM则是共享的运行最大值 $m$。值得注意的一个工程细节文档中的 $\exp$ 在实现里全部换成了以 2 为底的exp2f并配合params.scale_softmax_log2使用见 params.h 与wg0_bunch_0中的exp2f(rP0(i)*scale_softmax_log2 - new_max)。这是因为硬件上没有通用的 $\exp$ 指令换底到 2 可以把指数运算折叠进 FMA 中这也是代码中 LSE 的累加统一采用log2f的原因。另一个细节是共享内存的重用sP1复用了sQ的第 8 个 64×64 tile输出暂存区sO_addr则与sK0/sK1重叠见 splitkv_mla.cuh把宝贵的共享内存用到极致。技术细节新内核虽然面向计算受限场景带宽不再是瓶颈但内存延迟依然不可忽视如果数据在需要使用的那一刻还没就绪就只能空等。针对这一点新内核采用了两种手段。细粒度 TMA copy-GEMM 流水线对于一个 $64 \times 576$ 的 K 块不是一次性整块拷贝而是拆成 9 次 TMA 拷贝每次搬一个 $64 \times 64$ 的 tile。GEMM 不必等 9 次拷贝全部完成——第一片 TMA 拷贝一结束就可以启动第一片 GEMM以此类推从而大幅提升对内存延迟的容忍度。源码与此一一对应在 splitkv_mla.cuh 中launch_kv_tiles_copy_tma按模板参数START_HEAD_DIM_TILE_IDX → END_HEAD_DIM_TILE_IDX递归地为 576 维 K 头维度逐片发射 64×64 的 TMA 拷贝每个 K 块对应 9 个独立的 TMA barrierbarriers_K0[9]/barriers_K1[9]见 splitkv_mla.cuhqkt_gemm_one_tile_sQ/qkt_gemm_one_tile_rQ逐个 barrierwait后立即做该片的 $QK^\intercal$warpgroup_cooperative_qkt_gemm用PHASE_IDX0/1/2把 9 片 $QK^\intercal$ 的计算拆到两个 warpgroup 上流水执行PHASE-0 由 warpgroup 0 算前 4 片PHASE-1 由 warpgroup 1 算完 K1 的全部 9 片PHASE-2 再由 warpgroup 0 补算 K0 的后 5 片见 splitkv_mla.cuh——这正是 seesaw 主循环中拷贝与计算咬合的微观体现。缓存提示Cache Hints对 TMA 拷贝使用cute::TMA::CacheHintSm90::EVICT_FIRST提示可以提升 L2 缓存命中率实验证实。在 splitkv_mla.cuh 的 KV tile 拷贝与 Q 的 TMA 加载同文件 L704中都可以看到该缓存提示参数。综合以上两项优化新内核在 H800 SXM5 上可以达到约 80% 的 Tensor Core 利用率以降频后的理论峰值为基准与3 TB/s 的显存带宽。代价是在内存受限场景下比旧版双缓冲 ping-pong 版本慢约 2%官方认为这个代价可以接受。Programmatic Dependent Launch重叠 splitkv_mla 与 combine解码时每条请求的 KV 序列可能很长需要做 split-KV把一个请求的 KV 块按 SM 切分每个SM part负责一段产出部分输出 $O_{accum}$ 与部分 LSE随后由combine 内核把这些部分结果按在线 softmax 的规则合并为最终输出。两个内核之间原本存在天然依赖官方使用Programmatic Dependent LaunchPDL让 combine 内核在 splitkv_mla 尚未完全结束时就开始启动并预取数据实现两核重叠。源码证据清晰MLA 内核与 combine 内核均通过cudaLaunchKernelEx并携带cudaLaunchAttributeProgrammaticStreamSerialization属性发射见 splitkv_mla.cuh 与 combine.cuMLA 内核在完成最后一个请求后调用cudaTriggerProgrammaticLaunchCompletion()通知下游可以开始splitkv_mla.cuhcombine 内核在真正消费数据前调用cudaGridDependencySynchronize()等待上游完成combine.cu并用__ldg之外的普通加载预取o_accum数据注释明确说明__ldg与 PDL 不兼容。combine 内核本身也体现了精心设计的数值处理每个 warp 负责一个 q 头先在各 split 的 LSE 上取最大值max_lse再以exp2f(local_lse - max_lse)为权重对各 split 的部分输出做加权累加最终在 log2 域合成全局 LSE见 combine.cu数值稳定性与效率兼顾。Tile Scheduler跨 SM 的负载均衡新内核实现了一个 tile scheduler把请求 × KV 块的任务单元分配给各 SM保证 SM 之间负载均衡。其元数据核心是 params.h 中的DecodingSchedMetastruct DecodingSchedMeta { int begin_req_idx, end_req_idx; // 起始/结束请求索引均含 int begin_block_idx, end_block_idx; // 起始/结束块索引含不含 int begin_split_idx; // 起始 split 索引 int is_first_req_splitted, is_last_req_splitted; int _pad[1]; };调度本身由单 warp 的get_mla_metadata_kernelget_decoding_sched_meta.cu完成先按seqlens_k密集或topk稀疏统计每个请求的块数再以payload ceil_div(total_num_blocks, num_sm_parts) fixed_overhead_num_blocks为每个 SM part 的固定工作量上限贪心地逐请求切分块区间必要时把单个请求跨多个 SM part 切分split并记录num_splits_ptr供 combine 阶段对齐使用。这套元数据由 Python 侧get_mla_metadata在首次调用时触发生成之后在同一sched_meta上复用形状与cache_seqlens不变即可细节见 flash_mla_interface.py。性能表现与权衡把全部优化叠加后新内核在 H800 SXM5CUDA 12.8上的表现可以总结为场景旧版新版内存受限带宽~3000 GB/s~3000 GB/s略慢约 2%计算受限算力~580 TFlops最高 ~660 TFlops提升 5%~15%此外新内核达到约 80% 的 Tensor Core 利用率相对降频后峰值与 3 TB/s 带宽。这些数字来自官方文档与 README 的公开声明均为 H800 SXM5 CUDA 12.8 环境下的实测结果不同 GPU、驱动与时钟策略下的绝对数值会有差异但计算受限场景显著受益、内存受限场景基本持平的结论在架构层面是稳健的。如何在仓库中验证与复现运行正确性测试需 SM90 GPU、CUDA 12.8 与 PyTorch 2.0安装方式见 READMEpython tests/test_flash_mla_dense_decoding.py该测试默认配置test_flash_mla_dense_decoding.py正是文档分析的典型场景h_q128、h_kv1、d576、dv512、block_size64并通过reference_torch与 PyTorch 参考实现逐元素对比out与lse保证 seesaw 调度与在线 softmax 的数学等价性可被直接验证。运行性能基准bench_flash_mla.py# 单独测 FlashMLA输出 TFLOPS 与 GB/s python benchmark/bench_flash_mla.py --one --target flash_mla # 与 PyTorch 参考实现对比 python benchmark/bench_flash_mla.py --compare --baseline torch --target flash_mla基准脚本内置的shape_configs覆盖 batch128、seqlen 从 1024 到 32768、h_q128、h_kv1、d576、dv512的配置并以FLOPS s_q * total_seqlens * h_q * (d dv) * 2与bytes (total_seqlens * h_kv * d ...) * sizeof(dtype)分别折算 TFLOPS 与 GB/s见 bench_flash_mla.py可以直接对照本文与文档中给出的性能区间。接入业务代码则使用与旧版完全兼容的接口见 README 与 flash_mla_interface.pyfrom flash_mla import get_mla_metadata, flash_mla_with_kvcache tile_scheduler_metadata, num_splits get_mla_metadata( cache_seqlens, s_q * h_q // h_kv, h_kv, h_q, is_fp8, topk, ) for i in range(num_layers): o_i, lse_i flash_mla_with_kvcache( q_i, kvcache_i, block_table, cache_seqlens, dv, tile_scheduler_metadata, num_splits, is_causal, is_fp8_kvcache, indices, )其中tile_scheduler_metadata即本文介绍的 Tile Scheduler 元数据载体s_q为每个 q 序列的 token 数关闭 MTP/投机解码时应为 1h_kv为 KV 头数h_q为查询头数。若希望深入阅读调度主循环的完整实现建议以 splitkv_mla.cuhseesaw 主循环与 PDL、combine.cu部分结果合并与 get_decoding_sched_meta.cu负载均衡三条主线为索引。总结FlashMLA 新内核的核心贡献可以概括为一条完整的技术链路理论定位通过 FLOPs/byte 比分析确认 MLA 解码在 DeepSeek 部署形态无张量并行、$h_q128$下处于计算受限区间调度创新受 WGMMA 寄存器约束启发提出单输出矩阵的 seesaw 调度用两个 warpgroup 交错执行 CUDA Core 与 Tensor Core 工作数学上与在线 softmax 等价工程优化细粒度 9 片 TMA 流水、EVICT_FIRST缓存提示、PDL 重叠 splitkv_mla 与 combine、基于DecodingSchedMeta的 tile scheduler 负载均衡最终效果计算受限场景从 ~580 TFlops 提升至最高 ~660 TFlops5%~15%Tensor Core 利用率约 80%内存受限场景仅慢约 2%且接口完全向后兼容。致谢FlashMLA 的算法与调度设计深受 FlashAttention、Flash-Decoding 与 CUTLASS 及其背后众多项目的启发官方文档对上述工作表示感谢。本文的源码级佐证全部来自当前仓库主循环见 csrc/sm90/decode/dense/splitkv_mla.cuh块尺寸与头维度定义见 csrc/sm90/decode/dense/config.h合并与调度元数据见 csrc/smxx/decode/combine/combine.cu 与 csrc/smxx/decode/get_decoding_sched_meta/get_decoding_sched_meta.cu参数结构见 csrc/params.h。引用如引用本文所述的新内核工作可参考官方文档提供的书目信息Jiashi Li 与 Shengyu Liu 于 2025 年发布的《FlashMLA: Efficient MLA decoding kernels》。对应的完整 BibTeX 条目可在 docs/20250422-new-kernel-deep-dive.md 的 Citation 一节中找到。【免费下载链接】FlashMLAFlashMLA: Efficient Multi-head Latent Attention Kernels项目地址: https://gitcode.com/GitHub_Trending/fl/FlashMLA创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
锦
锦皓数字建站
深耕本土企业品牌数字化升级,专注原创端正雅致商务官网,从视觉设计到稳定运维全程保驾护航。