Skip to content

[M3] Improve indexer for long-context decode (sm100) - #48582

Merged
ywang96 merged 4 commits into
vllm-project:mainfrom
gau-nernst:m3_indexer_decode
Jul 17, 2026
Merged

ywang96 merged 4 commits into
vllm-project:mainfrom
gau-nernst:m3_indexer_decode

Conversation

@gau-nernst

@gau-nernst gau-nernst commented Jul 14, 2026 •

Copy link
Copy Markdown
Contributor

Purpose

This PR adds a CuteDSL indexer decode kernel for Minimax M3, optimized for long-context.

  • TMA + mma.sync design, instead of tcgen05 as I find saturating memory bandwidth with high occupancy (+ hardware scheduling) is easier than tcgen05 software pipelining, especially when there is context imbalance (different requests have different context lengths)
  • BF16 and FP8 Indexer cache support
  • Speculative decoding support. However, it's constrained to (1 + num_speculative_tokens) * num_idx_heads <= 32. Fallback to Triton otherwise

Microbenchmarks

#45743 reduces TARGET_GRID in Triton kernel from 4096 to 512. Even though this helps for uniform context lengths, it reduces load balancing effect, leading to worse performance for non-uniform context lengths. Hence, to make it a fair comparison, I'm comparing the CuteDSL kernel in this PR with both Triton-512 and Triton-4096 (i.e. using 512 or 4096 target CTAs)

Benchmark done on GB300

Microbenchmark script
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Benchmark MiniMax M3 Triton and CuteDSL decode index-score kernels.

The benchmark reports two workloads:

* A sweep over uniform batch sizes and context lengths.
* A batch-size sweep whose per-request context lengths are sampled from a
  reproducible log-uniform distribution.
"""

from __future__ import annotations

import argparse
import csv
import math
import random
import statistics
from dataclasses import dataclass
from pathlib import Path

import torch
from flashinfer.testing import bench_gpu_time_with_cupti

from vllm.models.minimax_m3.common.ops.index_topk import (
    SPARSE_BLOCK_SIZE,
    _decode_index_score_kernel,
)
from vllm.models.minimax_m3.nvidia.ops import (
    minimax_m3_index_decode_score_cutedsl,
)
from vllm.platforms import current_platform
from vllm.triton_utils import triton
from vllm.utils.math_utils import round_up

TOTAL_INDEX_HEADS = 4
INDEX_HEAD_DIM = 128
MAX_MODEL_LEN = 1_048_576
DEFAULT_UNIFORM_BATCH_SIZES = "1,4,16"
DEFAULT_BATCH_SIZES = "32,40,48,56,64"
DEFAULT_CONTEXT_LENS = "8192,32768,131072"


@dataclass
class Case:
    query: torch.Tensor
    key_cache: torch.Tensor
    block_table: torch.Tensor
    seq_lens: torch.Tensor
    scores: dict[str, torch.Tensor]
    seq_lens_cpu: list[int]


def _parse_ints(value: str) -> list[int]:
    return [int(item) for item in value.split(",") if item]


def _local_index_heads(tp: int) -> int:
    return max(1, TOTAL_INDEX_HEADS // tp)


def _sample_seq_lens(
    batch_size: int,
    min_context_len: int,
    max_context_len: int,
) -> list[int]:
    log_min = math.log(min_context_len)
    log_max = math.log(max_context_len)
    return [
        round(math.exp(random.uniform(log_min, log_max))) for _ in range(batch_size)
    ]


def _make_case(
    dtype: torch.dtype,
    seq_lens_cpu: list[int],
    num_heads: int,
    decode_query_len: int,
) -> Case:
    batch_size = len(seq_lens_cpu)
    blocks_per_request = [
        triton.cdiv(seq_len, SPARSE_BLOCK_SIZE) for seq_len in seq_lens_cpu
    ]
    num_pages = sum(blocks_per_request)
    total_queries = batch_size * decode_query_len
    static_blocks = triton.cdiv(MAX_MODEL_LEN, SPARSE_BLOCK_SIZE)

    query = torch.randn(
        total_queries,
        num_heads,
        INDEX_HEAD_DIM,
        device="cuda",
        dtype=torch.bfloat16,
    ).to(dtype)
    key_cache = torch.randn(
        num_pages,
        SPARSE_BLOCK_SIZE,
        INDEX_HEAD_DIM,
        device="cuda",
        dtype=torch.bfloat16,
    ).to(dtype)
    active_pages = torch.randperm(num_pages, device="cuda", dtype=torch.int32)
    block_table = torch.zeros(
        batch_size,
        static_blocks,
        device="cuda",
        dtype=torch.int32,
    )
    offset = 0
    for request_id, num_blocks in enumerate(blocks_per_request):
        block_table[request_id, :num_blocks] = active_pages[
            offset : offset + num_blocks
        ]
        offset += num_blocks

    seq_lens = torch.tensor(seq_lens_cpu, device="cuda", dtype=torch.int32)
    score_shape = (num_heads, total_queries, round_up(static_blocks, 16))
    scores = {
        provider: torch.full(
            score_shape,
            -float("inf"),
            device="cuda",
            dtype=torch.float32,
        )
        for provider in ("triton_512", "triton_4096", "cutedsl")
    }
    return Case(
        query=query,
        key_cache=key_cache,
        block_table=block_table,
        seq_lens=seq_lens,
        scores=scores,
        seq_lens_cpu=seq_lens_cpu,
    )


def _run_triton(
    case: Case,
    num_heads: int,
    decode_query_len: int,
    max_decode_query_len: int,
    target_grid: int,
) -> None:
    use_pdl = current_platform.is_arch_support_pdl()
    launch_kwargs: dict[str, bool | int] = {}
    if use_pdl:
        launch_kwargs["launch_pdl"] = True
    if num_heads > 1 and max_decode_query_len > 1:
        launch_kwargs.update(num_warps=4, num_stages=2)

    max_num_kv_chunks = 256
    target = max(
        1,
        min(
            max_num_kv_chunks,
            target_grid // max(1, case.seq_lens.shape[0]),
        ),
    )
    num_kv_chunks = 1 << (target.bit_length() - 1)
    score = case.scores[f"triton_{target_grid}"]
    _decode_index_score_kernel[(case.seq_lens.shape[0], num_kv_chunks)](
        case.query,
        case.key_cache,
        score,
        case.block_table,
        case.seq_lens,
        num_heads,
        INDEX_HEAD_DIM,
        0,
        0,
        decode_query_len,
        case.query.stride(0),
        case.query.stride(1),
        case.query.stride(2),
        case.key_cache.stride(0),
        case.key_cache.stride(1),
        case.key_cache.stride(2),
        score.stride(0),
        score.stride(1),
        score.stride(2),
        case.block_table.stride(0),
        BLOCK_SIZE_K=SPARSE_BLOCK_SIZE,
        BLOCK_SIZE_Q=triton.next_power_of_2(max_decode_query_len),
        num_kv_chunks=num_kv_chunks,
        USE_PDL=use_pdl,
        **launch_kwargs,
    )


def _make_calls(
    case: Case,
    num_heads: int,
    decode_query_len: int,
    max_decode_query_len: int,
) -> dict[str, object]:
    def run_triton_512() -> None:
        _run_triton(
            case,
            num_heads,
            decode_query_len,
            max_decode_query_len,
            512,
        )

    def run_triton_4096() -> None:
        _run_triton(
            case,
            num_heads,
            decode_query_len,
            max_decode_query_len,
            4096,
        )

    def run_cutedsl() -> None:
        minimax_m3_index_decode_score_cutedsl(
            case.query,
            case.key_cache,
            case.block_table,
            case.seq_lens,
            max_seq_len=MAX_MODEL_LEN,
            init_blocks=0,
            local_blocks=0,
            num_kv_heads=num_heads,
            decode_query_len=decode_query_len,
            max_decode_query_len=max_decode_query_len,
            score_out=case.scores["cutedsl"],
        )

    return {
        "triton_512": run_triton_512,
        "triton_4096": run_triton_4096,
        "cutedsl": run_cutedsl,
    }


def _check_correctness(case: Case, decode_query_len: int) -> None:
    expected = case.scores["triton_512"]
    for provider in ("triton_4096", "cutedsl"):
        actual = case.scores[provider]
        for request_id, seq_len in enumerate(case.seq_lens_cpu):
            num_blocks = triton.cdiv(seq_len, SPARSE_BLOCK_SIZE)
            query_start = request_id * decode_query_len
            query_end = query_start + decode_query_len
            torch.testing.assert_close(
                actual[:, query_start:query_end, :num_blocks],
                expected[:, query_start:query_end, :num_blocks],
            )


def _bench_us(fn: object, warmup_iters: int, repeat_iters: int) -> float:
    times_ms = bench_gpu_time_with_cupti(
        fn,
        dry_run_iters=warmup_iters,
        repeat_iters=repeat_iters,
        cold_l2_cache=True,
        use_cuda_graph=True,
    )
    return statistics.median(times_ms) * 1e3


def _logical_bytes(case: Case, decode_query_len: int) -> int:
    num_heads = case.query.shape[1]
    num_blocks = sum(
        triton.cdiv(seq_len, SPARSE_BLOCK_SIZE) for seq_len in case.seq_lens_cpu
    )
    return (
        case.query.numel() * case.query.element_size()
        + case.key_cache.numel() * case.key_cache.element_size()
        + num_blocks * num_heads * decode_query_len * torch.float32.itemsize
        + num_blocks * torch.int32.itemsize
        + case.seq_lens.numel() * case.seq_lens.element_size()
    )


def _run_case(
    args: argparse.Namespace,
    dtype: torch.dtype,
    seq_lens_cpu: list[int],
) -> dict[str, float | int | str]:
    num_heads = _local_index_heads(args.tp)
    case = _make_case(
        dtype,
        seq_lens_cpu,
        num_heads,
        args.decode_query_len,
    )
    calls = _make_calls(
        case,
        num_heads,
        args.decode_query_len,
        args.max_decode_query_len,
    )
    for call in calls.values():
        call()
    torch.cuda.synchronize()
    _check_correctness(case, args.decode_query_len)

    result: dict[str, float | int | str] = {
        "dtype": "bf16" if dtype is torch.bfloat16 else "fp8",
        "tp": args.tp,
        "batch": len(seq_lens_cpu),
        "dql": args.decode_query_len,
        "seq_min": min(seq_lens_cpu),
        "seq_p50": statistics.median(seq_lens_cpu),
        "seq_max": max(seq_lens_cpu),
    }
    logical_bytes = _logical_bytes(case, args.decode_query_len)
    for provider, call in calls.items():
        time_us = _bench_us(
            call,
            args.warmup_iters,
            args.repeat_iters,
        )
        result[f"{provider}_us"] = time_us
        result[f"{provider}_gbps"] = logical_bytes / time_us / 1e3
    del case
    torch.cuda.empty_cache()
    return result


def _print_rows(title: str, rows: list[dict[str, float | int | str]]) -> None:
    print(f"\n{title}")
    print(
        "dtype tp dql batch seq_min seq_p50 seq_max "
        "triton_512 triton_4096 cutedsl vs_512 vs_4096"
    )
    for row in rows:
        triton_512_us = float(row["triton_512_us"])
        triton_4096_us = float(row["triton_4096_us"])
        cutedsl_us = float(row["cutedsl_us"])
        print(
            f"{row['dtype']:>5} {row['tp']:2} {row['dql']:3} {row['batch']:5} "
            f"{row['seq_min']:7} {row['seq_p50']:7.0f} {row['seq_max']:7} "
            f"{triton_512_us:.1f} us / {row['triton_512_gbps']:.1f} GB/s  "
            f"{triton_4096_us:.1f} us / {row['triton_4096_gbps']:.1f} GB/s  "
            f"{cutedsl_us:.1f} us / {row['cutedsl_gbps']:.1f} GB/s  "
            f"{triton_512_us / cutedsl_us:.2f}x  "
            f"{triton_4096_us / cutedsl_us:.2f}x"
        )


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument(
        "--dtype",
        nargs="+",
        choices=("bf16", "fp8"),
        default=("bf16", "fp8"),
    )
    parser.add_argument("--tp", type=int, default=4)
    parser.add_argument("--decode-query-len", type=int, default=1)
    parser.add_argument("--max-decode-query-len", type=int)
    parser.add_argument(
        "--uniform-batch-sizes",
        default=DEFAULT_UNIFORM_BATCH_SIZES,
    )
    parser.add_argument("--batch-sizes", default=DEFAULT_BATCH_SIZES)
    parser.add_argument("--context-lens", default=DEFAULT_CONTEXT_LENS)
    parser.add_argument("--min-context-len", type=int, default=8_000)
    parser.add_argument("--max-context-len", type=int, default=250_000)
    parser.add_argument("--seed", type=int, default=0)
    parser.add_argument("--warmup-iters", type=int, default=20)
    parser.add_argument("--repeat-iters", type=int, default=100)
    parser.add_argument("--csv", type=Path)
    args = parser.parse_args()

    if not current_platform.is_device_capability_family(100):
        raise RuntimeError("The CuteDSL kernel requires a Blackwell GPU.")
    if TOTAL_INDEX_HEADS % args.tp != 0:
        raise ValueError("tp must divide the four total index heads")
    if args.max_decode_query_len is None:
        args.max_decode_query_len = args.decode_query_len
    if args.decode_query_len > args.max_decode_query_len:
        raise ValueError("decode-query-len must not exceed max-decode-query-len")
    if _local_index_heads(args.tp) * args.max_decode_query_len > 32:
        raise ValueError("local index heads * max decode query length must be <= 32")

    torch.manual_seed(args.seed)
    random.seed(args.seed)
    torch.cuda.set_device(0)
    dtypes = {
        "bf16": torch.bfloat16,
        "fp8": torch.float8_e4m3fn,
    }
    context_lens = _parse_ints(args.context_lens)
    uniform_batch_sizes = _parse_ints(args.uniform_batch_sizes)
    batch_sizes = _parse_ints(args.batch_sizes)
    sampled_context_lens = {
        batch_size: _sample_seq_lens(
            batch_size,
            args.min_context_len,
            args.max_context_len,
        )
        for batch_size in batch_sizes
    }

    uniform_rows = [
        _run_case(args, dtypes[dtype], [context_len] * batch_size)
        for dtype in args.dtype
        for context_len in context_lens
        for batch_size in uniform_batch_sizes
    ]
    random_rows = [
        _run_case(
            args,
            dtypes[dtype],
            sampled_context_lens[batch_size],
        )
        for dtype in args.dtype
        for batch_size in batch_sizes
    ]

    print(f"GPU: {current_platform.get_device_name()}")
    print(
        "random context distribution: "
        f"log-uniform [{args.min_context_len}, {args.max_context_len}], "
        f"seed={args.seed}"
    )
    _print_rows("Uniform context-length sweep", uniform_rows)
    _print_rows("Randomized context-length batch sweep", random_rows)

    if args.csv:
        rows = uniform_rows + random_rows
        args.csv.parent.mkdir(parents=True, exist_ok=True)
        with args.csv.open("w", newline="") as file:
            writer = csv.DictWriter(file, fieldnames=list(rows[0]))
            writer.writeheader()
            writer.writerows(rows)
        print(f"wrote {args.csv}")


if __name__ == "__main__":
    main()

Uniform context lengths

BF16, TP4, normal decode — DQL1

Context BS Triton-512 Triton-4096 CuteDSL Cute vs best
8K 1 4.7 µs / 443 GB/s 4.7 µs / 446 GB/s 4.2 µs / 497 GB/s 1.11×
8K 4 6.2 µs / 1,345 GB/s 6.3 µs / 1,338 GB/s 5.6 µs / 1,486 GB/s 1.10×
8K 16 10.6 µs / 3,169 GB/s 11.6 µs / 2,882 GB/s 11.9 µs / 2,812 GB/s 0.89×
32K 1 6.0 µs / 1,402 GB/s 6.0 µs / 1,395 GB/s 5.5 µs / 1,516 GB/s 1.08×
32K 4 10.7 µs / 3,131 GB/s 10.8 µs / 3,094 GB/s 10.4 µs / 3,217 GB/s 1.03×
32K 16 30.5 µs / 4,402 GB/s 29.8 µs / 4,502 GB/s 27.4 µs / 4,898 GB/s 1.09×
128K 1 12.0 µs / 2,804 GB/s 12.0 µs / 2,797 GB/s 9.9 µs / 3,389 GB/s 1.21×
128K 4 30.3 µs / 4,426 GB/s 32.3 µs / 4,158 GB/s 26.4 µs / 5,085 GB/s 1.15×
128K 16 95.2 µs / 5,640 GB/s 93.3 µs / 5,755 GB/s 85.7 µs / 6,264 GB/s 1.09×

FP8, TP4, normal decode — DQL1

Context BS Triton-512 Triton-4096 CuteDSL Cute vs best
8K 1 4.1 µs / 258 GB/s 4.1 µs / 258 GB/s 3.7 µs / 283 GB/s 1.09×
8K 4 4.9 µs / 852 GB/s 5.0 µs / 843 GB/s 4.5 µs / 930 GB/s 1.09×
8K 16 7.8 µs / 2,141 GB/s 7.8 µs / 2,150 GB/s 7.6 µs / 2,204 GB/s 1.03×
32K 1 5.0 µs / 841 GB/s 5.0 µs / 846 GB/s 4.6 µs / 917 GB/s 1.08×
32K 4 7.9 µs / 2,132 GB/s 7.9 µs / 2,115 GB/s 7.2 µs / 2,321 GB/s 1.09×
32K 16 19.8 µs / 3,387 GB/s 18.1 µs / 3,704 GB/s 15.7 µs / 4,278 GB/s 1.15×
128K 1 9.5 µs / 1,766 GB/s 9.5 µs / 1,760 GB/s 7.4 µs / 2,261 GB/s 1.28×
128K 4 19.6 µs / 3,423 GB/s 17.5 µs / 3,836 GB/s 16.4 µs / 4,102 GB/s 1.07×
128K 16 61.7 µs / 4,355 GB/s 55.3 µs / 4,860 GB/s 45.9 µs / 5,849 GB/s 1.20×

FP8, TP4, Eagle3-like decode — DQL4

Context BS Triton-512 Triton-4096 CuteDSL Cute vs best
8K 1 4.1 µs / 259 GB/s 4.1 µs / 259 GB/s 3.5 µs / 301 GB/s 1.17×
8K 4 5.1 µs / 821 GB/s 5.2 µs / 816 GB/s 4.6 µs / 912 GB/s 1.11×
8K 16 8.1 µs / 2,072 GB/s 9.8 µs / 1,716 GB/s 7.7 µs / 2,179 GB/s 1.05×
32K 1 5.1 µs / 820 GB/s 5.2 µs / 815 GB/s 4.8 µs / 884 GB/s 1.08×
32K 4 8.1 µs / 2,063 GB/s 8.7 µs / 1,923 GB/s 7.2 µs / 2,323 GB/s 1.13×
32K 16 20.4 µs / 3,297 GB/s 21.5 µs / 3,123 GB/s 16.1 µs / 4,183 GB/s 1.27×
128K 1 9.9 µs / 1,691 GB/s 9.9 µs / 1,699 GB/s 7.6 µs / 2,224 GB/s 1.31×
128K 4 20.5 µs / 3,276 GB/s 20.8 µs / 3,235 GB/s 16.8 µs / 3,996 GB/s 1.22×
128K 16 62.8 µs / 4,278 GB/s 58.3 µs / 4,611 GB/s 46.3 µs / 5,805 GB/s 1.26×

Non-uniform context lengths

BF16, TP4

DQL BS Context min/P50/max Triton-512 Triton-4096 CuteDSL Cute vs best
1 32 11K/66K/236K 191.4 µs / 3,939 GB/s 128.7 µs / 5,859 GB/s 116.8 µs / 6,457 GB/s 1.10×
1 48 9K/71K/247K 344.7 µs / 3,143 GB/s 181.4 µs / 5,973 GB/s 164.8 µs / 6,574 GB/s 1.10×
1 64 8K/39K/250K 358.4 µs / 3,036 GB/s 183.0 µs / 5,943 GB/s 165.8 µs / 6,561 GB/s 1.10×
4 32 11K/66K/236K 195.9 µs / 3,851 GB/s 128.7 µs / 5,862 GB/s 116.5 µs / 6,472 GB/s 1.10×
4 48 9K/71K/247K 356.6 µs / 3,039 GB/s 185.8 µs / 5,832 GB/s 164.9 µs / 6,572 GB/s 1.13×
4 56 8K/28K/240K 327.8 µs / 2,741 GB/s 157.7 µs / 5,697 GB/s 137.4 µs / 6,539 GB/s 1.15×
4 64 8K/39K/250K 363.7 µs / 2,992 GB/s 185.6 µs / 5,863 GB/s 165.0 µs / 6,594 GB/s 1.12×

FP8, TP4

DQL BS Context min/P50/max Triton-512 Triton-4096 CuteDSL Cute vs best
1 32 11K/66K/236K 148.2 µs / 2,544 GB/s 77.3 µs / 4,875 GB/s 62.3 µs / 6,055 GB/s 1.24×
1 48 9K/71K/247K 284.6 µs / 1,904 GB/s 110.7 µs / 4,893 GB/s 85.8 µs / 6,312 GB/s 1.29×
1 64 8K/39K/250K 287.5 µs / 1,892 GB/s 112.0 µs / 4,857 GB/s 86.5 µs / 6,287 GB/s 1.29×
4 32 11K/66K/236K 162.6 µs / 2,321 GB/s 81.5 µs / 4,628 GB/s 62.1 µs / 6,078 GB/s 1.31×
4 48 9K/71K/247K 309.7 µs / 1,751 GB/s 118.5 µs / 4,577 GB/s 86.4 µs / 6,272 GB/s 1.37×
4 64 8K/39K/250K 312.7 µs / 1,741 GB/s 120.3 µs / 4,527 GB/s 86.8 µs / 6,275 GB/s 1.39×

E2E accuracy

MiniMaxAI/MiniMax-M3-MXFP8, with FP8 KV cache and FP8 Indexer cache

AIME25 (4 epochs)

  • Exact match: 90.0%
  • Pass@4: 93.33%

E2E perf benchmarks

TP4, decode-only, 4xGB300, FP8 Indexer cache, FP8 KV cache, concurrency 48, input length 60k-120k

Eagle3? Main throughput Main P50 TPOT This PR throughput This PR P50 TPOT
No 1,667.58 tok/s 28.18 ms 1,749.01 tok/s (+4.88%) 26.82 ms (−4.83%)
Yes 2,432.05 tok/s 18.41 ms 2,604.86 tok/s (+7.11%) 17.33 ms (−5.87%)

Note: even though the kernel is tuned and validated on sm100 only, it should work well on SM90 and SM120 as well. For SM120, explicit FP8->FP16 dequantization is not necessary. This is left for future PRs.


Essential Elements of an Effective PR Description Checklist
  • The purpose of the PR, such as "Fix some issue (link existing issues this PR will resolve)".
  • The test plan, such as providing test command.
  • The test results, such as pasting the results comparison before and after, or e2e results
  • (Optional) The necessary documentation update, such as updating supported_models.md and examples for a new model.
Signed-off-by: Thien Tran <gau.nernst@yahoo.com.sg>
Signed-off-by: Thien Tran <gau.nernst@yahoo.com.sg>
Signed-off-by: Thien Tran <gau.nernst@yahoo.com.sg>

@claude claude Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Claude Code Review

This pull request is from a fork — automated review is disabled. A repository maintainer can comment @claude review to run a one-time review.

@zyongye zyongye added the ready ONLY add when PR is ready to merge/full CI is needed label Jul 16, 2026
@ywang96
ywang96 merged commit fe784ff into vllm-project:main Jul 17, 2026
94 of 96 checks passed
@gau-nernst
gau-nernst deleted the m3_indexer_decode branch July 17, 2026 01:12
brb-nv added a commit to brb-nv/TensorRT-LLM that referenced this pull request Aug 20, 2026
…wiring

The MiniMax-M3 decode path currently runs its generation rows through
fmha_sm100, which schedules a generation row like a context row. Three
kernels ported from vLLM replace that: a CuTe DSL indexer scoring kernel, a
Triton sparse block decode kernel and a trtllm-gen dense decode kernel.

This lands the kernels, their custom-op registration and their correctness
tests. Nothing dispatches to them yet: the MSA backend, indexer and cache
manager are untouched, so the decode path is byte-for-byte what it was and
the kernels are reachable only from the tests and the microbenchmark. The
dispatch is a follow-up, since it rests on the device-side length patching
and the fused per-layer cache writes that are still landing on
feat/m3_with_msa.

Split out of NVIDIA#17268 on feat/m3_with_msa, which carries the same kernels plus
that wiring. Two additions differ from it. msa_indexer gains only
cutedsl_score_runner and _cutedsl_score, the self-contained entry points the
scorer test drives, and not the run_indexer dispatch that calls them. The
tests reach fmha_sm100 through a local _flat_page_table helper, because
build_kv_page_indices does not take a block table until NVIDIA#16875; the helper
feeds it the slot map that block table implies, so the A/B comparisons still
run against the production page-table builder rather than a test-local copy.

The CuTe DSL indexer decode kernel and its tests were originally contributed
to vLLM by Thien Tran (vllm-project/vllm#48582), as were the CuTe utilities
(vllm-project/vllm#43273). The Triton sparse decode kernel and its tests
were originally contributed to vLLM by Kaichao You
(vllm-project/vllm#45381). Thanks to both.

No test-list change: l0_b300 already collects unittest/_torch/attention
wholesale, so the three new files are picked up there, and each skips itself
off SM100/SM103.

(cherry picked from commit 727c683)
Signed-off-by: Balaram Buddharaju <169953907+brb-nv@users.noreply.github.com>
brb-nv added a commit to brb-nv/TensorRT-LLM that referenced this pull request Aug 20, 2026
…wiring

The MiniMax-M3 decode path currently runs its generation rows through
fmha_sm100, which schedules a generation row like a context row. Three
kernels ported from vLLM replace that: a CuTe DSL indexer scoring kernel, a
Triton sparse block decode kernel and a trtllm-gen dense decode kernel.

This lands the kernels, their custom-op registration and their correctness
tests. Nothing dispatches to them yet: the MSA backend, indexer and cache
manager are untouched, so the decode path is byte-for-byte what it was and
the kernels are reachable only from the tests and the microbenchmark. The
dispatch is a follow-up, since it rests on the device-side length patching
and the fused per-layer cache writes that are still landing on
feat/m3_with_msa.

Split out of NVIDIA#17268 on feat/m3_with_msa, which carries the same kernels plus
that wiring. Two additions differ from it. msa_indexer gains only
cutedsl_score_runner and _cutedsl_score, the self-contained entry points the
scorer test drives, and not the run_indexer dispatch that calls them. The
tests reach fmha_sm100 through a local _flat_page_table helper, because
build_kv_page_indices does not take a block table until NVIDIA#16875; the helper
feeds it the slot map that block table implies, so the A/B comparisons still
run against the production page-table builder rather than a test-local copy.

The CuTe DSL indexer decode kernel and its tests were originally contributed
to vLLM by Thien Tran (vllm-project/vllm#48582), as were the CuTe utilities
(vllm-project/vllm#43273). The Triton sparse decode kernel and its tests
were originally contributed to vLLM by Kaichao You
(vllm-project/vllm#45381). Thanks to both.

No test-list change: l0_b300 already collects unittest/_torch/attention
wholesale, so the three new files are picked up there, and each skips itself
off SM100/SM103.

(cherry picked from commit 727c683)
Signed-off-by: Balaram Buddharaju <169953907+brb-nv@users.noreply.github.com>
brb-nv added a commit to brb-nv/TensorRT-LLM that referenced this pull request Aug 21, 2026
…wiring

The MiniMax-M3 decode path currently runs its generation rows through
fmha_sm100, which schedules a generation row like a context row. Three
kernels ported from vLLM replace that: a CuTe DSL indexer scoring kernel, a
Triton sparse block decode kernel and a trtllm-gen dense decode kernel.

This lands the kernels, their custom-op registration and their correctness
tests. Nothing dispatches to them yet: the MSA backend, indexer and cache
manager are untouched, so the decode path is byte-for-byte what it was and
the kernels are reachable only from the tests and the microbenchmark. The
dispatch is a follow-up, since it rests on the device-side length patching
and the fused per-layer cache writes that are still landing on
feat/m3_with_msa.

Split out of NVIDIA#17268 on feat/m3_with_msa, which carries the same kernels plus
that wiring. Two additions differ from it. msa_indexer gains only
cutedsl_score_runner and _cutedsl_score, the self-contained entry points the
scorer test drives, and not the run_indexer dispatch that calls them. The
tests reach fmha_sm100 through a local _flat_page_table helper, because
build_kv_page_indices does not take a block table until NVIDIA#16875; the helper
feeds it the slot map that block table implies, so the A/B comparisons still
run against the production page-table builder rather than a test-local copy.

The CuTe DSL indexer decode kernel and its tests were originally contributed
to vLLM by Thien Tran (vllm-project/vllm#48582), as were the CuTe utilities
(vllm-project/vllm#43273). The Triton sparse decode kernel and its tests
were originally contributed to vLLM by Kaichao You
(vllm-project/vllm#45381). Thanks to both.

No test-list change: l0_b300 already collects unittest/_torch/attention
wholesale, so the three new files are picked up there, and each skips itself
off SM100/SM103.

(cherry picked from commit 727c683)
Signed-off-by: Balaram Buddharaju <169953907+brb-nv@users.noreply.github.com>
brb-nv added a commit to brb-nv/TensorRT-LLM that referenced this pull request Aug 24, 2026
…wiring

The MiniMax-M3 decode path currently runs its generation rows through
fmha_sm100, which schedules a generation row like a context row. Three
kernels ported from vLLM replace that: a CuTe DSL indexer scoring kernel, a
Triton sparse block decode kernel and a trtllm-gen dense decode kernel.

This lands the kernels, their custom-op registration and their correctness
tests. Nothing dispatches to them yet: the MSA backend, indexer and cache
manager are untouched, so the decode path is byte-for-byte what it was and
the kernels are reachable only from the tests and the microbenchmark. The
dispatch is a follow-up, since it rests on the device-side length patching
and the fused per-layer cache writes that are still landing on
feat/m3_with_msa.

Split out of NVIDIA#17268 on feat/m3_with_msa, which carries the same kernels plus
that wiring. Two additions differ from it. msa_indexer gains only
cutedsl_score_runner and _cutedsl_score, the self-contained entry points the
scorer test drives, and not the run_indexer dispatch that calls them. The
tests reach fmha_sm100 through a local _flat_page_table helper, because
build_kv_page_indices does not take a block table until NVIDIA#16875; the helper
feeds it the slot map that block table implies, so the A/B comparisons still
run against the production page-table builder rather than a test-local copy.

The CuTe DSL indexer decode kernel and its tests were originally contributed
to vLLM by Thien Tran (vllm-project/vllm#48582), as were the CuTe utilities
(vllm-project/vllm#43273). The Triton sparse decode kernel and its tests
were originally contributed to vLLM by Kaichao You
(vllm-project/vllm#45381). Thanks to both.

No test-list change: l0_b300 already collects unittest/_torch/attention
wholesale, so the three new files are picked up there, and each skips itself
off SM100/SM103.

(cherry picked from commit 727c683)
Signed-off-by: Balaram Buddharaju <169953907+brb-nv@users.noreply.github.com>
brb-nv added a commit to brb-nv/TensorRT-LLM that referenced this pull request Aug 24, 2026
…wiring

The MiniMax-M3 decode path currently runs its generation rows through
fmha_sm100, which schedules a generation row like a context row. Three
kernels ported from vLLM replace that: a CuTe DSL indexer scoring kernel, a
Triton sparse block decode kernel and a trtllm-gen dense decode kernel.

This lands the kernels, their custom-op registration and their correctness
tests. Nothing dispatches to them yet: the MSA backend, indexer and cache
manager are untouched, so the decode path is byte-for-byte what it was and
the kernels are reachable only from the tests and the microbenchmark. The
dispatch is a follow-up, since it rests on the device-side length patching
and the fused per-layer cache writes that are still landing on
feat/m3_with_msa.

Split out of NVIDIA#17268 on feat/m3_with_msa, which carries the same kernels plus
that wiring. Two additions differ from it. msa_indexer gains only
cutedsl_score_runner and _cutedsl_score, the self-contained entry points the
scorer test drives, and not the run_indexer dispatch that calls them. The
tests reach fmha_sm100 through a local _flat_page_table helper, because
build_kv_page_indices does not take a block table until NVIDIA#16875; the helper
feeds it the slot map that block table implies, so the A/B comparisons still
run against the production page-table builder rather than a test-local copy.

The CuTe DSL indexer decode kernel and its tests were originally contributed
to vLLM by Thien Tran (vllm-project/vllm#48582), as were the CuTe utilities
(vllm-project/vllm#43273). The Triton sparse decode kernel and its tests
were originally contributed to vLLM by Kaichao You
(vllm-project/vllm#45381). Thanks to both.

No test-list change: l0_b300 already collects unittest/_torch/attention
wholesale, so the three new files are picked up there, and each skips itself
off SM100/SM103.

(cherry picked from commit 727c683)
Signed-off-by: Balaram Buddharaju <169953907+brb-nv@users.noreply.github.com>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

ready ONLY add when PR is ready to merge/full CI is needed

3 participants