社区开发者将 FlashMLA 移植至消费级 Blackwell 显卡
r/LocalLLaMA 用户为 sm_120 架构编译 FlashMLA 内核,推理与训练环节较 PyTorch SD…
一名开源社区开发者在 Reddit r/LocalLLaMA 发布了一项针对消费级 NVIDIA Blackwell 显卡(sm_120)的 FlashMLA 内核移植工作。该项目在其个人 LLM 训练仓库中完成,覆盖 GQA、MLA、Dense 等常见架构,主要解决了 FlashMLA 官方版本仅支持 sm_100 与 sm_90 编译目标的限制,使新一批消费级 Blackwell 设备也能用上这套源于 DeepSeek 的高效注意力实现。
推理性能:稀疏场景显著领先
在 FP8 解码与稀疏服务等典型场景下,自编译的 FlashMLA 内核相对 PyTorch SDPA 表现出明显加速:
- Sparse FP8 decode(b=128,s_q=2,topk=2048):0.809 ms vs 2.118 ms,加速约 2.62×
- Sparse serving(b=4,s_q=1,CFG=4,warm):0.050 ms vs 0.257 ms(以 CFG=1 作参照),约 5×
- Sparse prefill forward(s_q=512,s_kv=8192):1.240 ms vs 3.232 ms,约 2.61×
- Dense decode(H=22,s_q=1,4K cache):0.440 ms / 1394 GB/s,PyTorch 无对应路径
在 FP8 KV cache 路径下,FlashMLA 实现了延迟降低约 8.0% 与缓存内存降至 1/1.84 倍的效果;但在 BF16 缓存的整模型解码环节,两者基本持平(约 1.0×),说明内核层面的加速并不必然转化为端到端的 3× 模型级加速。
训练性能:长上下文与稀疏 Prefill 受益最大
开发者按模型实际注意力形状(192/128,H=22)进行前向 + 反向测试,结果如下:
- Dense S=4096:3.630 ms vs 8.696 ms,加速约 2.40×
- Dense S=8192:9.911 ms vs 30.105 ms,加速约 3.04×
- Dense S=1024(warm-clock):0.306 ms vs 1.007 ms,加速约 3.29×
- Sparse prefill(s_q=512,topk=2048):6.651 ms vs 20.054 ms,加速约 3.01×
数据表明,注意力占比越高的训练任务——例如长上下文训练与稀疏 prefill——获得的收益越大。
适用范围与局限
整体来看,该移植在「注意力密集型」场景下可带来 2–3 倍加速,适合本地训练 MLA 类模型或运行长上下文推理;但在 BF16 缓存的常规解码路径上则与 PyTorch 持平,开发者也提醒读者不要直接把这些内核级数字当作整模型 3× 提速的承诺。代码已在 GitHub 公开(仓库:IISuperluminaLII/FlashMLA_Windows_Linux_sm120),同时支持 Windows 与 Linux。
