Prhub

#48770 [2/N][Attention] Enable masked MHA for sparse MLA prefills

原始 PR 作者 MatthewBonanni 合并时间 2026-07-31 21:01 文件变更 11 提交数 28 评论 21 代码增减 +855 / -26

执行摘要

稀疏 MLA 预填充启用 masked MHA 加速

DeepSeek V3.2 论文指出短序列预填充下 masked MHA 模式能更高效地模拟 DSA。PR body 说明对于纯 prefill,当序列长度低于某阈值时 masked MHA 路径快于稀疏 MQA,并明确该 PR 构建在 #47327 实现的 dense MHA 捷径之上。

值得精读。重点关注位打包掩码的构建与复用(_build_topk_mask / _scatter_topk_kernel)、全局掩码 64 MiB 上限的设计权衡,以及 _use_masked_mha 阈值表的推导依据。若计划在非 Blackwell 或量化 KV 场景扩展该优化,可参考其门控与回落机制。

讨论亮点

LucasWilkinson 质疑同时保留全局掩码与分块掩码两条分支的必要性("do we really need both of these branches? i find this very confusing"),MatthewBonanni 回应全局掩码是性能优势,但其大小随 batch_size * query_len * seq_len 增长,因此设置 GLOBAL_TOPK_MASK_MAX_BYTES = 64 MiB 上限,超过后回落分块掩码。另有关于测试中 topk 索引 -1/0 语义、FA4 输出能否 inplace 复用以及若干删除注释/格式化的 nit,作者均已修复或澄清。

实现拆解

  1. 能力门控与 workspace 分配:在 sparse_mla_attention.py 新增 _is_masked_mha_available,限制设备能力 100(Blackwell)、DSV3 维度(128 heads / 512 kv_lora_rank / 128 qk_nope / 64 qk_rope / 128 v_head)、FA4 且非量化 KV;满足时分配 64 MiB 位打包掩码 workspace(topk_mask_workspace)。
  2. 掩码构建内核:新增两个 Triton kernel(_scatter_topk_kernel_scatter_topk_single_req_kernel)和 _build_topk_mask,将 topk 索引按 32 位一组写入位图,单请求与多请求分别走专用/通用路径,避免多余偏移计算。
  3. 路由启发式:在 mla_attention.py 新增 _use_masked_mha,按后端(FLASHMLA_SPARSE / FLASHINFER_MLA_SPARSE)和 TP 大小配置 seq_len 分桶的 query_len 阈值表 _DSV32_MASKED_MHA_THRESHOLDS,并在 forward_impl 中综合 use_dense_mhause_masked_mhasparse_mla_force_mqa 决策是否路由 MHA。
  4. FA4 接口扩展flash_attn_interface.py 透传 block_sparse_tensorsaux_tensor_leading_dims 到 FA4 _flash_attn_fwd,并断言 full_block_cnt / full_block_idx 已物化;sparse_mla_mask.py 新增 CUTLASS mask_moddense_mask_mod / offset_dense_mask_mod)用于按位提取掩码。
  5. 测试与基准:新增 test_sparse_mla_mask.py 验证单请求与通用路径一致性;扩展 test_sparse_mla_backends.py 覆盖 masked_mhamasked_mha_chunked_context;基准脚本增加 masked_mha variant 与 sparse_mla_masked_mha_max_seq_len 配置,并提供 mla_sparse_masked_mha_vs_mqa.yaml 扫描配置。
文件 模块 状态 重要度
vllm/model_executor/layers/attention/sparse_mla_attention.py 注意力层 modified 9.05
vllm/model_executor/layers/attention/sparse_mla_mask.py 注意力掩码 added 8.13
vllm/model_executor/layers/attention/mla_attention.py 注意力路由 modified 7.44
tests/v1/attention/test_sparse_mla_mask.py 掩码测试 added 6.05
vllm/vllm_flash_attn/flash_attn_interface.py 注意力接口 modified 5.99
benchmarks/attention_benchmarks/benchmark.py 基准脚本 modified 5.53

关键符号

_is_masked_mha_available _scatter_topk_kernel _scatter_topk_single_req_kernel _build_topk_mask _use_masked_mha dense_mask_mod offset_dense_mask_mod

关键源码片段

vllm/model_executor/layers/attention/sparse_mla_mask.py core-logic

新增 CUTLASS mask_mod 实现,供 FA4 按位消费 top-k 掩码,并支持分块上下文偏移。

# vllm/model_executor/layers/attention/sparse_mla_mask.pyimport cutlass
import cutlass.cute as cutefrom vllm.vllm_flash_attn.cute import utils
​
​
@cute.jit
def dense_mask_mod(
    batch: cute.TensorSSA,
    head: cute.TensorSSA,
    q_idx: cute.TensorSSA,
    kv_idx: cute.TensorSSA,
    seqlen_info,
    aux_tensors: list,
) -> cute.TensorSSA:
    # 从 aux_tensors[0] 读取位打包掩码,按 (batch, q_idx, word_idx) 定位
    dense_mask = aux_tensors[0]
    batch_idx = utils.ssa_to_scalar(batch)
    q_idx = utils.ssa_to_scalar(q_idx)
    kv_idx = utils.ssa_to_scalar(kv_idx)
    word_idx = kv_idx >> 5
    bit_idx = cutlass.Uint32(kv_idx & 31)
    word = dense_mask[batch_idx, q_idx, word_idx]
    result = cute.make_rmem_tensor(1, dtype=cutlass.Uint32)
    result[0] = utils.shr_u32(cutlass.Uint32(word), bit_idx)
    return result.load()
​
​
dense_mask_mod.__vec_size__ = 32
​
​
@cute.jit
def offset_dense_mask_mod(
    batch: cute.TensorSSA,
    head: cute.TensorSSA,
    q_idx: cute.TensorSSA,
    kv_idx: cute.TensorSSA,
    seqlen_info,
    aux_tensors: list,
) -> cute.TensorSSA:
    # 分块上下文变体:掩码最后一维的最后一个元素保存 key_start 偏移
    dense_mask = aux_tensors[0]
    batch_idx = utils.ssa_to_scalar(batch)
    q_idx = utils.ssa_to_scalar(q_idx)
    key_start = dense_mask[batch_idx, 0, dense_mask.shape[2] - 1]
    kv_idx = utils.ssa_to_scalar(kv_idx) + key_start
    word_idx = kv_idx >> 5
    bit_idx = cutlass.Uint32(kv_idx & 31)
    word = dense_mask[batch_idx, q_idx, word_idx]
    result = cute.make_rmem_tensor(1, dtype=cutlass.Uint32)
    result[0] = utils.shr_u32(cutlass.Uint32(word), bit_idx)
    return result.load()
​
​
offset_dense_mask_mod.__vec_size__ = 32

评论区精华

全局掩码 vs 分块掩码双分支的必要性 设计

LucasWilkinson 质疑同时保留全局掩码和分块掩码两条路径容易混淆,问是否有强理由不统一为一种。

结论:MatthewBonanni 解释全局掩码是性能优势,但尺寸随 batch_size * query_len * seq_len 增长,因此设置 GLOBAL_TOPK_MASK_MAX_BYTES = 64 MiB 上限,超出后回落分块掩码。 · 已解决

测试中 topk 索引 -1/0 填充语义 question

LucasWilkinson 问测试里 masked 分支用 -1 填充而原分支用 0,两者作用是否相同,能否统一。

结论:作者在后续 commit cfc29572 中补充注释或调整,回复 done。 · 已解决

FA4 输出能否 inplace 复用 性能

LucasWilkinson 建议 output / output_lse 使用 inplace 指向同一 buffer,减少分配。

结论:作者在 commit c68f1a1 清理,标记为 fixed。 · 已解决

删除注释与格式化调整的 nits style

LucasWilkinson 对多处删除注释、重排返回类型注解提出 why remove? / why reformat?。

结论:作者逐条回复已修复,集中在 commit cfc29572 与 7a0c45a。 · 已解决

风险与影响

  1. 核心路径变更forward_impl 与预填充元数据结构改动影响所有 MLA 后端,即便 unmasked 配置也走新检查逻辑,存在回归风险。
  2. 依赖外部 FA 版本:路径仅在 FA4 且非量化 KV 下启用,依赖 flash-attention PR #155 的 block_sparse_tensors 接口,外部版本不满足时自动禁用,但需确保构建期版本一致。
  3. 硬编码阈值_DSV32_MASKED_MHA_THRESHOLDS 按 DeepSeek V3 维度硬编码,未来模型变体或新后端可能不适用,需人工维护。
  4. 仅限 Blackwell_is_masked_mha_available 要求设备能力 100,其他硬件不受影响但无法获得优化。
  5. 掩码 workspace 容量:64 MiB 上限若被突破则切换分块路径,路径切换正确性依赖新增测试(masked_mha_chunked_context),覆盖率仍有限。

对使用 DeepSeek V3.x 稀疏 MLA 且后端为 FLASHMLA_SPARSE / FLASHINFER_MLA_SPARSE 的用户,短序列纯 prefill 吞吐将显著提升;代码影响 mla_attention.pysparse_mla_attention.py、FA4 接口等核心注意力模块,并扩展了 attention 基准工具链;团队后续维护需关注阈值表和掩码上限的调优。

核心路径变更 依赖外部 FA 版本 硬编码阈值 仅限 Blackwell

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论