ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

为什么 FlashMLA 的解码内核能跑到 660 TFLOPS——内核里的 5 个关键设计

为什么 FlashMLA 的解码内核能跑到 660 TFLOPS——内核里的 5 个关键设计 为什么 FlashMLA 的解码内核能跑到 660 TFLOPS——内核里的 5 个关键设计【免费下载链接】FlashMLAFlashMLA: Efficient Multi-head Latent Attention Kernels项目地址: https://gitcode.com/GitHub_Trending/fl/FlashMLAFlashMLA 是 DeepSeek 开源的 MLAMulti-head Latent Attention注意力内核库支撑 DeepSeek-V3 系列模型的推理H800 上密集 MLA 解码最高 660 TFLOPS稀疏注意力预填充最高 640 TFLOPS。如果你是做 CUDA 内核开发、或者想弄懂 LLM 推理里解码到底卡在算力还是带宽的工程师这个仓库值得逐行读。全景速览仓库里有什么仓库分硬件两条线组织内核源码Python 侧只做薄封装分组位置说明稀疏解码FP8 KV 缓存csrc/sm90/decode/sparse_fp8/SM90/SM100MQA 模式FP8 KV bfloat16 计算密集解码csrc/sm90/decode/dense/SM90bfloat16稀疏预填充csrc/sm90/prefill/sparse/SM90/SM100token 级稀疏密集 MHA 前向/反向csrc/sm100/prefill/dense/SM100基于 CUTLASSPython 接口flash_mla/get_mla_metadata、flash_mla_with_kvcache等官方技术文档docs/两篇内核深度长文理解实现的最好入口关键依赖PyTorch ≥ 2.0、CUDA ≥ 12.8SM100 内核需 12.9、CUTLASS子模块。支持矩阵摘自 README内核GPU 架构KV 格式密集解码SM90bfloat16稀疏解码SM90 / SM100FP8稀疏预填充SM90 / SM100bfloat16密集 MHA 预填充SM100—技术内核5 个设计决策每个设计都按为什么需要 → 怎么实现 → 带来什么收益展开。设计 1先判断瓶颈——解码注意力居然是算力受限为什么需要优化方向由瓶颈决定解码就是访存密集是常见直觉但 MLA 不成立。怎么实现做一次算术强度账。FLOPs 约为2·h_q·s_q·s_k·(d_kd_v)访存量约为2·s_k·d_k字节比值约2·h_q·s_q。H800 的实际峰值带宽约 3.35 TB/s、考虑降频后算力约 865 TFlops比值 258/2 ≈ 129——只要h_q·s_q ≥ 128就是算力受限。DeepSeek 线上解码不开张量并行h_q恰好是 128落在算力受限一侧。收益整个内核围绕喂饱 Tensor Core设计见后文调度与流水线而不是堆带宽技巧。设计 2FP8 KV 缓存——给 656 字节的 token 布局算账 为什么需要DeepSeek-V3.2 把上下文翻倍到 128K 后单请求 KV 缓存在 bfloat16 下高达 8.72 GiBdocs 里的算式576 × 2 × 62 × 128 × 1024要么 OOM要么 batch 上不去。怎么实现每个 token 的 KV 共 576 维前 512 维NoPE 部分按1×128的 tile 粒度量化为 512 个float8_e4m3配 4 个float32缩放因子后 64 维RoPE 部分对精度敏感保持 bfloat16 不动。落到显存里就是 512 16 128 656 字节/token对比 bfloat16 的 1152 字节省约 43%。内核内先把 FP8 反量化回 bfloat16再全部用 bfloat16 做 MMA输出 fp32 累加。收益KV 显存直接决定能塞多少上下文这一步为长上下文推理腾出了空间代价是引入反量化开销——这正是设计 3 要还的债。设计 3反量化比 MMA 慢 47%用 Crossover 摊薄为什么需要H800 单个 SM 每周期能做 4096 次 MMA FLOP989 TFlops / 1830 MHz / 132 SM。一个 CTA 处理 64 个 query 头时每个 KV token 的 MMA 只需约34 周期而 H800 不能直接从float8_e4m3转 bfloat16反量化要 fp8→half→fp32→bf16 再乘缩放因子按 NVIDIA 指令吞吐表算下来每个 token 至少要50 周期。Tensor Core 等 CUDA Core这就是反量化受限dequantization-bound。怎么实现抓住 MQA 的一个事实——同一个 query token 的 128 个 query 头共享同一份 KV。于是把 CTA 以 cluster size 2 发射两个 CTA 各负责 64 个 query 头各只反量化一半 KV再用 Hopper 的分布式共享内存DSMcluster 内 CTA 可直读对方 shared memory把结果交换过去。交换用st.async异步写入对端 shared memory由 cluster transaction barrier 同步。对应代码在 csrc/sm90/decode/sparse_fp8/splitkv_mla.cuhst_async_128b写入sK_*_peer_base辅助原语在 csrc/sm90/decode/sparse_fp8/components/helpers.h。收益反量化工作量减半Tensor Core 利用率回升。官方给出的对比无 Crossover 的旧 FP8 稀疏解码内核 250 TFLOPS加上 Crossover 后410 TFLOPScompute-bound 配置bs128、128 头、s_q2、topk2048H800 SXM5。设计 4Seesaw 调度——只有一块输出矩阵怎么做 ping-pong为什么需要WGMMA 要求输出矩阵驻留寄存器。一个 64×512 的输出矩阵占32,768 个 32 位寄存器而 H800 每 SM 只有 65,536 个——塞得下 1 份塞不下 2 份。FlashAttention-3 的 ping-pong 调度靠两份输出矩阵交替占用 CUDA Core 和 Tensor Core在这里行不通。怎么实现把输出矩阵竖向劈成O_L、O_R各 64×256V同步劈成四块交给两个 warpgroup 各自持有半个输出。每个主循环取两个 KV 块K0, K1两个 warpgroup 交替推进一个算qK0ᵀ做 softmax 更新O_L另一个算qK1ᵀ更新O_R最后用两个缩放因子做交叉项p0·V0R进O_R、p1·V0L进O_L。数学上与 FlashAttention 的 online softmax 完全等价作者称之为 seesaw跷跷板调度。完整时序图见 docs/20250422-new-kernel-deep-dive.md 的附录图。附带好处某块 KV 数据用完后可以立刻发射下一块的 TMA 加载访存与计算天然重叠。收益单输出矩阵下依然让 CUDA Core 与 Tensor Core 全时段重叠官方报告达到降频理论峰值的80% Tensor Core 利用率。设计 5隐藏访存延迟——把 K 块切成 9 刀为什么需要即便算力受限数据没到位时 SM 也只能空转延迟仍是硬约束。怎么实现三个配合的小动作。一是细粒度 TMA 流水线一个 64×576 的 K 块拆成9 次 64×64 的 TMA 拷贝源码里对 shared memory 的 tile 划分正是(PAGE_BLOCK_SIZE, 64, 9)每块到齐就启动对应 GEMM不等整块。二是给 TMA 加EVICT_FIRST缓存提示实验证明能提升 L2 命中率csrc/sm90/decode/dense/splitkv_mla.cuh 中的launch_kv_tiles_copy_tma。三是稀疏解码路径用 128 位宽的__ldg加载量化 KV。另外splitkv_mla与combine两个内核通过 Programmatic Dependent Launch 重叠执行tile 调度器负责把请求/块均匀分给 SM。收益算力侧 80% 利用率与3 TB/s带宽同时达成H800 SXM5。实测数据值得记住的 4 个数字 ⚡660 TFLOPS密集 MLA 解码compute-bound 配置H800 SXM5 CUDA 12.83000 GB/s同内核在 memory-bound 配置下的带宽410 TFLOPSFP8 稀疏解码topk2048topk 调到 32768 可达460 TFLOPS640 TFLOPS稀疏注意力预填充前向H800 SXM5对照组不加 Crossover 的旧 FP8 稀疏解码内核是 250 TFLOPS。上手示例解码与稀疏预填充怎么调安装git clone https://gitcode.com/GitHub_Trending/fl/FlashMLA后git submodule update --init --recursive再pip install .。解码阶段先算一次 tile 调度元数据循环里每层调用内核from flash_mla import get_mla_metadata, flash_mla_with_kvcache sched, 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, sched, num_splits, is_causal, is_fp8_kvcache, indices)稀疏预填充内核flash_mla_sparse_fwd(q, kv, indices, sm_scale)的语义等价于下面这段 gather 以 2 为底的对数softmax注意它不支持 batch 维多批次要 reshape 模拟kv kv.squeeze(1) # [s_kv, d_qk]h_kv 必须为 1 indices indices.squeeze(1) # [s_q, topk] focused_kv kv[indices] # 按 indices gather 出 [s_q, topk, d_qk] P (Q focused_kv.transpose(-1, -2)) * sm_scale lse log2sumexp2(P) # 指数与对数均以 2 为底 S exp2(P - lse) out S focused_kv边界与权衡架构门槛只支持 SM90/SM100。Crossover 依赖 Hopper 的 CTA cluster 与 DSMSeesaw 依赖 WGMMA 的寄存器语义两者都无法移植到更老的架构。topk 小的代价稀疏内核的前置/收尾开销占比更大——topk2048 时 410 TFLOPS低于 bfloat16 密集解码 640 TFLOPS 的峰值topk32768 才到 460 TFLOPS。稀疏收益要序列足够长才体现。B200 尚未调优README 明确说 B200 上稀疏解码仅 350 TFLOPSnot really optimized yetSM100 侧目前只有 prefill 与 head64 解码路径sparse_fp8 的 Crossover 实现只在 SM90 目录里。量化是部分覆盖RoPE 那 64 维始终 bfloat16656 字节的布局是精度与显存的折中反量化路径fp8→half→fp32→bf16是 H800 硬件限制带来的额外开销换个架构周期账要重算。memory-bound 略有回退新版相对旧 ping-pong 版本在访存受限场景慢约 2%官方认为可接受。一句话收尾FlashMLA 给出的方法论比任何单个技巧都重要先用周期和字节做成本核算找出真瓶颈再用 DSM、TMA、PDL 这些硬件原语把每一拍空隙填平——660 TFLOPS 的解码性能就是这么算出来的。【免费下载链接】FlashMLAFlashMLA: Efficient Multi-head Latent Attention Kernels项目地址: https://gitcode.com/GitHub_Trending/fl/FlashMLA创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表