FlashMLA sm_120 内核构建,性能比 SDPA 提升 2-3 倍

Reddit r/LocalLLaMA 工具

摘要

作者为消费级 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× 模型加速。但对于注意力密集型工作负载——特别是长上下文训练和稀疏预填充——差异是显著的。
查看原文

相似文章

Blackwell 与 PDL 性能提升

Reddit r/LocalLLaMA

Llama.cpp 现已支持适用于 Blackwell GPU 的 Nvidia 程序化依赖启动 (PDL),在 Token 生成时可带来 5-10% 的性能提升。该功能默认未启用,需通过编译标志开启。