大家好,我最近在维护一个开源的大语言模型开发仓库,目标是让任何人都能针对各种架构(比如 GQA、MLA、Dense 等)训练任意规模的模型。
我在实现 MLA 时参考了 DeepSeek 的技术报告,发现他们用 FlashMLA 来做高速训练和推理。不过官方版本只支持 sm_100 和 sm_90,我在网上没找到有人把它编译到消费级 Blackwell 的 sm_120 上。于是我自己的仓库里就做了这个适配。
先说推理和服务的测试结果。在稀疏 FP8 解码(b=128、s_q=2、topk=2048)上,FlashMLA 耗时 0.809 毫秒,而 PyTorch SDPA 要 2.118 毫秒,快了约 2.62 倍;稀疏服务场景下更是快了约 5 倍;稀疏 prefill(s_q=512、s_kv=8192)也达到 2.61 倍。密集型解码用了 4K 缓存时,FlashMLA 跑出 1394 GB/s 的带宽,PyTorch 没有对应路径。使用 FP8 KV 缓存时,FlashMLA 延迟低 8%、显存占用仅为其 1/1.84。
训练方面,按模型真实注意力形状(192/128、H=22)测前向加反向:S=4096 时快 2.40 倍,S=8192 快 3.04 倍,S=1024 热身跑快 3.29 倍,稀疏 prefill 快 3.01 倍。
不过要注意,在完整模型层面,BF16 缓存解码其实和 PyTorch 基本持平,所以别把内核级数字直接当成端到端 3 倍加速。对长上下文训练和稀疏 prefill 这类注意力密集任务,提升才真正明显。另外我也欢迎有人提 PR,帮我补上错过的优化点。