FlashMLA sm_120 内核构建,性能比 SDPA 提升 2-3 倍
摘要
作者为消费级 Blackwell sm_120 构建了 FlashMLA,在长上下文训练和稀疏预填充等注意力密集型工作负载中,性能比 PyTorch SDPA 提升了 2-3 倍。
大家好,我一直在开发一个开源的大语言模型开发仓库。任何人都可以针对任何架构(GQA、MLA、Dense...)训练任意参数量。当我尝试实现 MLA 时,我阅读了 DeepSeek 实现的技術研究论文,发现他们使用了 FlashMLA 来加速训练和推理。问题是,它只针对 sm_100 和 sm_90 编译,我没有找到任何人尝试为消费级 Blackwell sm_120 构建它。https://github.com/IISuperluminaLII/FlashMLA_Windows_Linux_sm120 *我不擅长起名字
推理 FlashMLA 与 PyTorch SDPA 基准测试
推理/服务工作负载 FlashMLA SDPA 加速比
稀疏 FP8 解码 — b=128, s_q=2, topk=2048 0.809 ms 2.118 ms (gather + math) 2.62×
稀疏服务 — b=4, s_q=1 (CFG=4, warm) 0.050 ms 0.257 ms (CFG=1 作为代理) ~5×
稀疏预填充前向 — s_q=512, s_kv=8192 1.240 ms 3.232 ms 2.61×
密集解码 — H=22, s_q=1, 4K 缓存 (CFG=4) 0.440 ms / 1394 GB/s 无等效 PyTorch 路径 —
模型级别 BF16 缓存解码步骤 ~持平 ~持平 ~1.0×
FP8 KV 缓存,启用 FlashMLA 延迟降低 8.0%,缓存内存降低 1.84× 延迟提高 1.8% —
训练
对于我的用例——可能很快会有更多人——使用模型的实际注意力形状:192/128, H=22 进行前向和反向传播。
工作负载 FlashMLA SDPA 加速比
密集 S=4096 3.630 ms 8.696 ms 2.40×
密集 S=8192 9.911 ms 30.105 ms 3.04×
密集 S=1024 (warm-clock 运行) 0.306 ms 1.007 ms 3.29×
稀疏预填充 — s_q=512, topk=2048 6.651 ms 20.054 ms 3.01×
在全模型级别,BF16 缓存解码基本持平,因此我不会将内核级别的数字解释为自动的端到端 3× 模型加速。但对于注意力密集型工作负载——特别是长上下文训练和稀疏预填充——差异是显著的。
相似文章
FlashPrefill V2:面向长上下文LLM服务的块稀疏预填充注意力
FlashPrefill V2通过均值校正的稀疏注意力和优化的GPU算子改善了长上下文LLM服务,相比FlashAttention-2和稠密基线提供了显著的加速。
Flash-MSA: 利用稀疏注意力内核加速百万token训练
介绍Flash-MSA,首个针对MiniMax稀疏注意力在Hopper和Blackwell GPU上的高性能开源训练内核,实现高效的百万token训练。
Blackwell 与 PDL 性能提升
Llama.cpp 现已支持适用于 Blackwell GPU 的 Nvidia 程序化依赖启动 (PDL),在 Token 生成时可带来 5-10% 的性能提升。该功能默认未启用,需通过编译标志开启。
@modal: 我们与 @lmsysorg 和 http://z-lab.ai 合作,将 DFlash 规范集成到 @sgl_project,并通过重叠加速……
Modal 与 LMSys 和 Z Lab 合作,将 DFlash 推测解码集成到 SGLang,在大型语言模型上实现了相比基准最高 4.3 倍的吞吐量提升,比原生多 token 预测提升 1.5 倍。
@pupposandro: https://x.com/pupposandro/status/2054241934164492328
该文章宣布了 llama.cpp 对 AMD Strix Halo 集成 GPU (iGPU) 上的 DFlash 和 PFlash 投机解码的支持,并展示了使用 ROCm 时推理性能的显著提升。