Prhub

#45720 [Bugfix][ROCm] Fix MiniMax-M3 FP8 KV cache dtype

原始 PR 作者 cquil11 合并时间 2026-06-17 11:14 文件变更 3 提交数 6 评论 3 代码增减 +60 / -5

执行摘要

修复 ROCm 上 MiniMax-M3 FP8 KV 缓存 dtype 错误

关联Issue #45562 报告在ROCm MI300X上,使用--kv-cache-dtype fp8时MiniMax-M3模型生成结果严重错误,GSM8K strict accuracy从0.9666降至0.0099。根因是MiniMax-M3的sparse attention后端硬编码 torch.float8_e4m3fn 来重新解释字节型FP8 KV缓存,但ROCm平台使用 torch.float8_e4m3fnuz,两种编码不同导致KV值被错误读取。

本 PR 值得关注,尤其是跨平台 dtype 决策的提炼方式:引入 current_platform.fp8_dtype()is_fp8_fnuz() 避免了平台硬编码,是类似问题的标准处理模式。同时,新增的测试覆盖了平台参数化,可以作为 dtype 选择场景的参考模板。建议开发者在涉及平台类型差异时采用类似策略。

讨论亮点

PR 没有公开的 review 评论线程,但 reviewer @hongxiayang 参与了 co-authored commit(3e9faf5),在 E5M2 分支中增加了对 float8_e5m2fnuz 的判断。该修改使得 ROCm 的 E5M2 也能正确选择 FNUZ 变体,完善了原实现的平台覆盖。

实现拆解

  1. sparse_attention.pyMiniMaxM3SparseImpl.__init__ 中,将硬编码的 dtype 选择改为:若 kv_cache_dtype 包含 'e5m2',则根据 current_platform.is_fp8_fnuz() 选择 torch.float8_e5m2fnuztorch.float8_e5m2;否则使用 current_platform.fp8_dtype()(对ROCm返回 float8_e4m3fnuz,对CUDA返回 float8_e4m3fn)。
  2. sparse_attn.py 中定义了 _FP8_DTYPES 元组,包含全部四种FP8类型(e4m3fn、e4m3fnuz、e5m2、e5m2fnuz),并将内核中判断是否使用FP8路径的条件从 (torch.float8_e4m3fn, torch.float8_e5m2) 改为 in _FP8_DTYPES
  3. 新增两个测试函数:test_sparse_impl_uses_platform_fp8_dtype 参数化验证不同 kv_cache_dtype 字符串下构造的 impl 对象能正确返回平台期望的 dtype;test_sparse_kernels_recognize_fp8_dtypes 验证所有四种FP8类型都被 _FP8_DTYPES 识别。
文件 模块 状态 重要度
vllm/models/minimax_m3/common/sparse_attention.py 注意力 modified 6.92
vllm/models/minimax_m3/common/ops/sparse_attn.py 算子 modified 4.81
tests/kernels/attention/test_minimax_m3.py 测试 modified 5.79

关键符号

MiniMaxM3SparseImpl.__init__ minimax_m3_sparse_attn minimax_m3_sparse_attn_decode test_sparse_impl_uses_platform_fp8_dtype test_sparse_kernels_recognize_fp8_dtypes

关键源码片段

vllm/models/minimax_m3/common/sparse_attention.py core-logic

核心修复,在 MiniMaxM3SparseImpl.__init__ 中使用 current_platform.fp8_dtype() 替换硬编码,并处理 E5M2FNUZ 分支,直接影响 KV 缓存的 dtype 选择。

class MiniMaxM3SparseImpl(AttentionImplBase):
    def __init__(self, ..., kv_cache_dtype: str = "auto", ...):
        self.kv_cache_dtype = kv_cache_dtype
        self.use_fp8_kv = is_quantized_kv_cache(kv_cache_dtype)
        # 原来是硬编码:
        # torch.float8_e5m2 if "e5m2" in kv_cache_dtype else torch.float8_e4m3fn
        # 现在改为动态平台感知:
        if "e5m2" in kv_cache_dtype:
            self.kv_cache_fp8_dtype = (
                torch.float8_e5m2fnuz
                if current_platform.is_fp8_fnuz()
                else torch.float8_e5m2
            )
        else:
            # 对 ROCm 返回 float8_e4m3fnuz,对 CUDA 返回 float8_e4m3fn
            self.kv_cache_fp8_dtype = current_platform.fp8_dtype()
        self.topk_blocks = topk_blocks
        self.block_size = sparse_block_size
vllm/models/minimax_m3/common/ops/sparse_attn.py infrastructure

定义了 _FP8_DTYPES 元组供内核使用,并修改了 prefill 和 decode 两个函数的 dtype 检查条件,确保所有 FP8 变体都能触发 FP8 转换路径。

# 定义全部 FP8 类型集合,供内核检查使用
_FP8_DTYPES = (
    torch.float8_e4m3fn,
    torch.float8_e4m3fnuz,
    torch.float8_e5m2,
    torch.float8_e5m2fnuz,
)def minimax_m3_sparse_attn(..., kv_cache, ...):
    # 之前只检查 (torch.float8_e4m3fn, torch.float8_e5m2)
    # 现在改为检查全集
    use_fp8 = kv_cache.dtype in _FP8_DTYPES
    ...def minimax_m3_sparse_attn_decode(..., kv_cache, ...):
    use_fp8 = kv_cache.dtype in _FP8_DTYPES
    ...

评论区精华

没有提炼出高价值讨论线程

当前评论区没有形成足够清晰的争议点或结论,后续有更多讨论时会体现在这里。

风险与影响

本次修改仅影响 MiniMax-M3 模型,且只在 ROCm 平台上原有行为有差异。核心风险在于:如果未来新增其他 FP8 变体(如 e4m3fnuz 的另一种形式),需要同步更新 _FP8_DTYPESMiniMaxM3SparseImpl。目前测试覆盖了所有已知的 FP8 dtype,降低了回归风险。性能上无影响,仅改变了 dtype 选择逻辑。

修复了 ROCm 平台(主要是 MI300X、MI325X)使用 MiniMax-M3 模型并开启 FP8 KV 缓存时的极端精度问题(GSM8K 从0.01恢复至0.956)。用户无需修改配置即可获得正常精度。对 CUDA 平台无影响,因为 current_platform.fp8_dtype() 返回 float8_e4m3fn 与之前硬编码一致。系统其他组件不受影响。

跨平台 dtype 兼容性 核心路径变更 仅影响 ROCm

关联 Issue

#45562 [Bug]: ROCm MI300X FP8 KV cache MiniMax-M3-MXFP8 accuracy issues

完整报告

参与讨论