ARTICLE DETAIL

资讯详情

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

FlashMLA 注意力内核源码走读:656 字节 KV 缓存背后的完整链路

FlashMLA 注意力内核源码走读:656 字节 KV 缓存背后的完整链路 FlashMLA 注意力内核源码走读656 字节 KV 缓存背后的完整链路【免费下载链接】FlashMLAFlashMLA: Efficient Multi-head Latent Attention Kernels项目地址: https://gitcode.com/GitHub_Trending/fl/FlashMLAFlashMLA 注意力内核库是 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_e4m3576 维 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 × (576512) × 2 / 4096 ≈34 个周期。50 34内核处于反量化受限状态Tensor Core 被迫等数据。两个 CTA 分摊一份 KV关键事实MQA 模式下同一 query token 的 128 个 query 头读的是同一份 K/V。每个 CTACUDA 线程块被调度到 SM 上并行运行的一组线程只负责 64 个头正好可以分工。用 Hopper 的 CTA clusterCTA 集群一组可互相直接访问对方共享内存的 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算完的每一段各自留下部分输出与 lsecombine内核用 lse 把它们归一成最终结果。两个内核通过 Programmatic Dependent Launch程序化依赖启动让后一个内核在前一个收尾阶段就开始准备的启动机制重叠执行省掉一次完整的内核切换间隙。tile scheduler 与 PDL 组合的效果SM 之间不抢活、内核之间不空转长上下文下尤其明显。性能账本H800 SXM5 除非另注场景指标数值一句话说明密集解码计算受限TFLOPS660新版内核较旧版提升 5%~15%密集解码访存受限带宽3000 GB/s逼近 H800 约 3.35 TB/s 的理论上限密集解码新内核Tensor Core 利用率 / 带宽~80% / 3 TB/s利用率相对降频后理论峰值FP8 稀疏解码topk2048TFLOPS410batch128、128 头的计算受限配置FP8 稀疏解码topk32768TFLOPS460topk 更大前后处理占比下降FP8 稀疏解码无 CrossoverTFLOPS250对照基线量化 Crossover 的收益稀疏预填充H800 / B200TFLOPS640 / 1450驱动 DeepSeek-V3.2 的稀疏注意力稠密 MHA 预填充B200TFLOPS前向 1460 / 反向 1000NVIDIA 报告值回头看这条链路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),仅供参考
返回列表