Prhub

#45743 [M3] Tune Triton indexer score decode for spec-decode

原始 PR 作者 gau-nernst 合并时间 2026-06-17 12:07 文件变更 3 提交数 3 评论 0 代码增减 +72 / -49

执行摘要

MiniMax-M3 Triton indexer decode 按 request 批量处理 spec-decode tokens

之前为支持 spec-decode,只是简单地为每个 decode token 启动额外的 CTA,这样 kernel 保持简单但效率不佳,因为没有将多个 query tokens 批量在一起进行 QK MMA。本 PR 通过每个 request 一个 CTA 并 batch 多个 query tokens 来提升计算效率。

值得精读,尤其关注 Triton 内核优化模式:利用 constexpr 元数据避免重编译、通过虚拟 batch 批量处理 query tokens、移除冗余计算。设计决策中对无 spec-decode 场景保持兼容的做法也值得借鉴。

讨论亮点

本次 PR 无公开 review 评论,但作者提供了详尽的微基准和 E2E 测试结果,并且维护者已 Approved。

实现拆解

  1. 元数据扩展:在 MiniMaxM3IndexerDecodeMetadata 中添加 max_decode_query_len 字段,在 builder 的 init 中根据 reorder_batch_threshold 初始化。该值作为 constexpr 传递给 kernel,避免因实际 decode_query_len 变化而重编译。
  2. 内核重构:修改 _decode_index_score_kernel,将启动维度从每个 token 一个 CTA 改为每个 request 一个 CTA,并在 BLOCK_SIZE_Q 维度上并行处理多个 query tokens(头维度合并)。移除了 sm_scale,因为 decode 阶段只需要 block 顺序,不需要 scale。调整 TARGET_GRID 从 4096 降到 512。
  3. 测试适配:在 test_minimax_m3.py 中移除 sm_scale 参数,添加 max_decode_query_len 参数并扩展 parametrize 覆盖多种 (decode_query_len, max_decode_query_len) 组合,包括 pad 场景。
文件 模块 状态 重要度
vllm/models/minimax_m3/common/indexer.py MiniMax-M3 modified 5.84
vllm/models/minimax_m3/common/ops/index_topk.py Kernel modified 5.74
tests/kernels/attention/test_minimax_m3.py 测试 modified 4.46

关键符号

_decode_index_score_kernel MiniMaxM3IndexerDecodeMetadata MiniMaxM3IndexerMetadataBuilder.__init__ MiniMaxM3IndexerTritonMetadataBuilder.build test_decode_index_topk_correctness

关键源码片段

vllm/models/minimax_m3/common/indexer.py data-contract

添加 max_decode_query_len 字段并传递到 decode 内核,是元数据契约变更的关键文件。

@dataclass
class MiniMaxM3IndexerDecodeMetadata:
    """Per-decode state (cudagraph-safe). ``decode_query_len`` is the uniform
    per-request query length (1, or 1 + num_speculative_tokens).    Added ``max_decode_query_len`` field (NEW) to serve as a constexpr for the
    Triton kernel, avoiding recompilation when the actual decode query length
    varies across steps. It is equal to ``reorder_batch_threshold`` (1 or
    1 + num_speculative_tokens) and stays constant for the worker lifetime.
    """
    seq_lens: torch.Tensor # [num_decodes] int32
    block_table: torch.Tensor
    max_seq_len: int
    decode_query_len: int
    max_decode_query_len: int # NEW: maximum possible decode_query_len
​
​
class MiniMaxM3IndexerMetadataBuilder(...):
    def __init__(self, ...):
        ...
        self._init_reorder_batch_threshold(1, supports_spec_as_decode=True)
        assert self.reorder_batch_threshold is not None
        # NEW: persist the threshold as max for later use in forward()
        self.max_decode_query_len = self.reorder_batch_threshold
​
    def build(self, ...) -> MiniMaxM3IndexerMetadata:
        ...
        if num_decodes > 0:
            decode_query_len = ... # computed from common_attn_metadata
            decode_metadata = MiniMaxM3IndexerDecodeMetadata(
                ...,
                decode_query_len=decode_query_len,
                max_decode_query_len=self.max_decode_query_len, # NEW field
            )
        ...
vllm/models/minimax_m3/common/ops/index_topk.py core-logic

核心内核优化:改 CTA 粒度、移除 sm_scale、调整 BLOCK_SIZE_Q,是性能提升的关键。

@triton.jit(do_not_specialize=["num_kv_chunks", "decode_query_len"])
def _decode_index_score_kernel(
    # ... (various pointers)
    head_dim: tl.constexpr,
    init_blocks, local_blocks,
    # REMOVED: sm_scale – decode only needs block ordering, scaling is unnecessary.
    decode_query_len, # actual query length for this request (can vary)
    stride_q_n, stride_q_h, stride_q_d,
    ...
    BLOCK_SIZE_K: tl.constexpr,
    BLOCK_SIZE_Q: tl.constexpr, # NEW: number of query tokens processed per inner block
    num_kv_chunks,
    USE_PDL: tl.constexpr,
):
    # NEW: flatten heads and Q dimensions into one grid dimension (per request)
    BLOCK_SIZE_HQ: tl.constexpr = num_idx_heads * BLOCK_SIZE_Q
    hq_offsets = tl.arange(0, BLOCK_SIZE_HQ)
    q_offsets = hq_offsets % BLOCK_SIZE_Q
    q_mask = q_offsets < decode_query_len
​
    pid_r = tl.program_id(0) # request id (instead of per-token id)
    q_ids = pid_r * decode_query_len + q_offsets # global token indices
​
    # Load KV block addresses and compute QK dot product (no scaling)
    ...
    qk = tl.dot(q, k) # sm_scale removed
​
    # Causal mask & chunk scanning unchanged...

评论区精华

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

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

风险与影响

核心风险是正确性:新增 max_decode_query_len 字段必须在所有使用处一致传递,否则可能导致 kernel 执行错误;测试覆盖了多种组合但需关注边界。性能风险:无 spec-decode 场景微基准有 ±2% 波动;E2E benchmark 显示无退化。可移植性:Triton 内核依赖 CUDA,非 NVIDIA GPU 可能不适用。

用户影响:MiniMax-M3 模型使用 spec-decode 时生成速度显著提升(长上下文下 TPOT 降低可达 12.9%);无 spec-decode 场景无影响。影响范围仅限 MiniMax-M3 模型的 Triton indexer,不涉及其他模型或系统组件。

新字段传播风险 Triton 内核兼容性 BLOCK_SIZE_Q 限制

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论