Prhub

#48642 [Bugfix] Sparse MLA: enable fp8_ds_mla dense prefill

原始 PR 作者 MatthewBonanni 合并时间 2026-07-17 06:44 文件变更 13 提交数 4 评论 3 代码增减 +516 / -85

执行摘要

修复 fp8_ds_mla dense prefill 的越界和 gather 问题

PR #47327添加了dense-MHA prefill路由到sparse MLA,暴露了两个fp8_ds_mla混合batch bug:1)req_id_per_token可能超过top-k张量长度导致越界写;2)dense MHA无法收集packed 656字节FP8 cache,且FlashMLA假定其metadata覆盖整个batch。本PR解决这些问题以支持fp8_ds_mla下的dense prefill。

建议精读flashmla_sparse.py中metadata重构和mla_attention.py中prefill路径切换的设计,理解如何将decode-only metadata与dense prefill解耦。对有FP8 sparse attention需求的团队,此PR提供了可参考的修复模式。

讨论亮点

审核人 mgoin 认为设计合理,依赖CI验证。drakosha在相关Issue的评论中确认了该分支在4×H200上通过DCP前缀缓存测试,decode吞吐与封闭构建持平,且gather正确性无问题。

实现拆解

  1. 边界检查sparse_utils.py):在triton_convert_req_index_to_global_index中增加长度断言,拒绝req_idtoken_indices形状不匹配的调用,防止kernel grid越界。
  2. FP8 gather kernel扩展csrc/cache.hcsrc/libtorch_stable/ops.h):cp_gather_and_upconvert_fp8_kv_cache新增seq_starts参数,支持从任意缓存位置开始gather;同时移除seq_lens参数简化接口。所有调用点(_custom_ops.py、benchmark、测试)同步更新。
  3. Metadata重构flashmla_sparse.py):在FlashMLASparseMetadataBuilder中引入require_uniform_decodes = True标志。_build_fp8_separate_prefill_decode现在直接使用metadata中已解析的decode/prefill计数,而非再次调用split_decodes_and_prefills,使得当dense MHA处理prefill时,decode-only MQA metadata能正确构建。
  4. Prefill路径适配mla_attention.py):当cache_dtype == "fp8_ds_mla"时,chunked_prefill_workspace dtype固定为bfloat16(而非query dtype),因为FP8 cache上转换后为BF16。在_compute_prefill_context_context_parallel_compute_prefill_context中,对fp8_ds_mla调用cp_gather_and_upconvert_fp8_kv_cache进行gather+上转换,替代原有gather_and_maybe_dequant_cache
  5. 测试覆盖test_sparse_mla_backends.pytest_cp_gather_fp8.pytest_mla_backends.py):新增越界拒绝、decode slice正确性、seq_starts gather、cache dtype映射等测试。
文件 模块 状态 重要度
vllm/v1/attention/backends/mla/flashmla_sparse.py 稀疏注意后端 modified 6.9
vllm/model_executor/layers/attention/mla_attention.py MLA 预填充 modified 6.73
tests/v1/attention/test_sparse_mla_backends.py 测试覆盖 modified 7.82
tests/kernels/test_cp_gather_fp8.py 测试覆盖 modified 6.21
vllm/v1/attention/backends/mla/sparse_utils.py 工具函数 modified 4.94

关键符号

test_triton_convert_rejects_req_id_longer_than_token_indices test_flashmla_forward_bf16_kv_slices_req_id_to_mqa_tokens test_cp_gather_fp8_with_sequence_starts test_mla_kv_cache_spec_uses_layer_cache_dtype _compute_prefill_context _build_fp8_separate_prefill_decode triton_convert_req_index_to_global_index

关键源码片段

vllm/v1/attention/backends/mla/flashmla_sparse.py core-logic

核心逻辑变更:重构 metadata 构建,支持 decode-only MQA metadata 与 dense MHA prefill 分离,引入 require_uniform_decodes 标志,修改 _build_fp8_separate_prefill_decode 接口。

# FlashMLASparseMetadataBuilder 中 require_uniform_decodes 标志的引入class FlashMLASparseMetadataBuilder(SparseMLACommonMetadataBuilder[FlashMLASparseMetadata]):
    _cudagraph_support: ClassVar[AttentionCGSupport] = AttentionCGSupport.UNIFORM_BATCH
    # 当 require_uniform_decodes=True 时,metadata 构建器会跳过对 decode 的 split_decodes_and_prefills
    # 调用,直接使用已解析的 metadata 字段,从而支持 decode-only MQA kernel metadata
    require_uniform_decodes: ClassVar[bool] = True
​
    def _build_fp8_separate_prefill_decode(
        self,
        common_attn_metadata: CommonAttentionMetadata,
        metadata: FlashMLASparseMetadata, # 传入已构建的 metadata,避免重复 split
    ) -> "FlashMLASparseMetadata.FP8SeparatePrefillDecode":
        num_tokens = common_attn_metadata.num_actual_tokens
        # 直接从 metadata 获取,而非再次调用 split_decodes_and_prefills
        num_decodes = metadata.num_decodes
        num_prefills = metadata.num_prefills
        num_decode_tokens = metadata.num_decode_tokens
        num_prefill_tokens = num_tokens - num_decode_tokens
        # ... 构建 decode 和 prefill 子结构 ...
vllm/model_executor/layers/attention/mla_attention.py data-contract

prefill 路径调整:在 fp8_ds_mla 时使用 cp_gather_and_upconvert kernel 替换原有 gather,workspace dtype 固定为 bfloat16

# _compute_prefill_context 中 fp8_ds_mla 分支def _compute_prefill_context(self, q, kv_c_and_k_pe_cache, attn_metadata, k_scale):
    # ... 前置代码 ...
    for i in range(iters):
        toks = prefill_metadata.chunked_context.seq_tot[i]
        if self.kv_cache_dtype == 'fp8_ds_mla':
            # 使用支持任意 seq_starts 的 gather+upconvert kernel
            ops.cp_gather_and_upconvert_fp8_kv_cache(
                src_cache=kv_c_and_k_pe_cache,
                dst=workspace[:toks],
                block_table=prefill_metadata.block_table,
                workspace_starts=prefill_metadata.chunked_context.cu_seq_lens[i],
                batch_size=attn_metadata.num_prefills,
                seq_starts=prefill_metadata.chunked_context.starts[i], # 新增:每行的起始偏移
            )
        elif not use_fp8_prefill:
            ops.gather_and_maybe_dequant_cache(...) # 原有 BF16 路径
        else:
            # fp8 query prefill 路径不变
            ...

评论区精华

DCP 前缀缓存验证 测试

drakosha 在 Issue 评论中测试了本分支在 4×H200 上的 DCP 路径,证实 dense-MHA seq_starts gather 正确,prefix cache 命中无误,decode 吞吐与封闭构建持平。

结论:DCP 路径兼容性得到确认,无回归。 · 已解决

风险与影响

主要风险包括:

1) cp_gather_and_upconvert_fp8_kv_cache移除了seq_lens参数,改为seq_starts,若其他未发现的调用点未更新将引发编译或运行时错误;
2) metadata重构改变了_build_fp8_separate_prefill_decode的输入依赖,可能在某些未测试的batch组合下产生不一致;
3) workspace dtype从q_data_type改为固定bfloat16,在非fp8_ds_mla模式下可能引入额外精度损失或性能变化;
4) 所有变更仅在CUDA SM90+(Blackwell)下测试,其他平台可能无法正常工作。但测试覆盖较全面,基准测试也同步更新,风险可控。

直接影响:使用fp8_ds_mla和sparse MLA的用户(如DeepSeek模型)的dense prefill路由不再崩溃,混合batch准确率与sparse decode一致。间接影响:kernel接口变更可能影响内部其他分支或未来集成(如DCP #46514)。团队需关注后续merge的DCP变更是否与此处seq_starts逻辑兼容。

越界风险修复 kernel 接口变更 metadata 重构 仅 CUDA SM90+ 测试

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论