Skip to content

vllm.models.minimax_m3.amd.ops.index_topk

Triton kernels for MiniMax M3 lightning-indexer block scoring + top-k.

Index queries score each 128-token block of index keys (max over the block), then the top-k blocks (plus forced init/local blocks) are selected per query token. Adapted to vLLM's paged KV cache: the KV page size is forced to equal the sparse block size (128), so one sparse block maps to exactly one page.

Index-K cache layout (vLLM): (num_blocks, 128, idx_head_dim) (one shared key vector per token).

Only the paths MiniMax M3 uses are implemented: score_type="max", index value disabled (score-only indexer), and shared index keys. Each local index-query head selects its own block ids for the block-sparse attention kernels in sparse_attn.

Functions:

_decode_score_split_launch_policy(num_reqs, head_dim, query_dtype, cache_dtype, *, is_gfx950)

Choose the generic split-K launch and high-batch specialization.

Source code in vllm/models/minimax_m3/amd/ops/index_topk.py
def _decode_score_split_launch_policy(
    num_reqs: int,
    head_dim: int,
    query_dtype: torch.dtype,
    cache_dtype: torch.dtype,
    *,
    is_gfx950: bool,
) -> tuple[int, bool]:
    """Choose the generic split-K launch and high-batch specialization."""
    use_high_batch_config = (
        is_gfx950
        and MIN_DECODE_SCORE_GFX950_HIGH_BATCH_REQUESTS
        <= num_reqs
        <= MAX_DECODE_SCORE_GFX950_HIGH_BATCH_REQUESTS
        and head_dim == 128
        and query_dtype == torch.bfloat16
        and cache_dtype == torch.bfloat16
    )
    target_grid = (
        DECODE_SCORE_GFX950_HIGH_BATCH_TARGET_GRID
        if use_high_batch_config
        else DECODE_SCORE_DEFAULT_TARGET_GRID
    )
    target = max(
        1,
        min(
            MAX_DECODE_SCORE_KV_CHUNKS,
            target_grid // max(1, num_reqs),
        ),
    )
    return 1 << (target.bit_length() - 1), use_high_batch_config

_decode_topk_launch_policy(max_block, total_q, num_idx_heads, topk, *, is_gfx950)

Choose the selector grid and compile-time launch configuration.

Source code in vllm/models/minimax_m3/amd/ops/index_topk.py
def _decode_topk_launch_policy(
    max_block: int,
    total_q: int,
    num_idx_heads: int,
    topk: int,
    *,
    is_gfx950: bool,
) -> tuple[int, int, int, int, bool, bool]:
    """Choose the selector grid and compile-time launch configuration."""
    if (
        is_gfx950
        and topk == 16
        and 0 < max_block <= MAX_DECODE_TOPK_FAST_CHUNKS * DECODE_TOPK_BLOCKS_PER_CHUNK
    ):
        if max_block <= DECODE_TOPK_SHORT_BLOCKS_PER_CHUNK:
            return 1, DECODE_TOPK_SHORT_BLOCKS_PER_CHUNK, 2, 1, True, True
        return (
            MAX_DECODE_TOPK_FAST_CHUNKS,
            DECODE_TOPK_BLOCKS_PER_CHUNK,
            4,
            2,
            True,
            True,
        )

    target = max(
        1,
        min(
            MAX_DECODE_TOPK_FAST_CHUNKS,
            DECODE_TOPK_TARGET_GRID // max(1, total_q * num_idx_heads),
        ),
    )
    return (
        1 << (target.bit_length() - 1),
        DECODE_TOPK_BLOCKS_PER_CHUNK,
        8,
        2,
        False,
        False,
    )

minimax_m3_index_decode(idx_q, index_kv_cache, block_table, seq_lens, max_seq_len, topk, init_blocks, local_blocks, num_kv_heads, decode_query_len, max_decode_query_len, out=None, *, attention_block_table=None, sparse_block_table_out=None, sparse_context_lens_out=None, block_page_stride=None, completion_counter=None)

Decode index block-score followed by fused adaptive top-k selection.

Returns topk_idx [num_kv_heads, total_q, topk] (0-indexed block ids, -1 pad). When out ([num_kv_heads, >=total_q, topk]) is given, writes into out[:, :total_q, :] (stable address for cudagraph) instead of allocating. The optional sparse-table arguments fuse current-layer table construction into the selector. They must be provided together. completion_counter provides stable per-query synchronization storage for CUDA graphs. It must be zero before its first launch and must not be shared by overlapping selector invocations; every completed launch resets its active entries.

Source code in vllm/models/minimax_m3/amd/ops/index_topk.py
1296
1297
1298
1299
1300
1301
1302
1303
1304
1305
1306
1307
1308
1309
1310
1311
1312
1313
1314
1315
1316
1317
1318
1319
1320
1321
1322
1323
1324
1325
1326
1327
1328
1329
1330
1331
1332
1333
1334
1335
1336
1337
1338
1339
1340
1341
1342
1343
1344
1345
1346
1347
1348
1349
1350
1351
1352
1353
1354
1355
1356
1357
1358
1359
1360
1361
1362
1363
1364
1365
1366
1367
1368
1369
1370
1371
1372
1373
1374
1375
1376
1377
1378
1379
1380
1381
1382
1383
1384
1385
1386
1387
1388
1389
1390
1391
1392
1393
1394
1395
1396
1397
1398
1399
1400
1401
1402
1403
1404
1405
1406
1407
1408
1409
1410
1411
1412
1413
1414
1415
1416
1417
1418
1419
1420
1421
1422
1423
1424
1425
1426
1427
1428
1429
1430
1431
1432
1433
1434
1435
1436
1437
1438
1439
1440
1441
1442
1443
1444
1445
1446
1447
1448
1449
1450
1451
1452
1453
1454
1455
1456
1457
1458
1459
1460
1461
1462
1463
1464
1465
1466
1467
1468
1469
1470
1471
1472
1473
1474
1475
1476
1477
1478
1479
1480
1481
1482
1483
1484
1485
1486
1487
1488
1489
1490
1491
1492
1493
1494
1495
1496
1497
1498
1499
1500
1501
1502
1503
1504
1505
1506
1507
1508
1509
1510
1511
1512
1513
1514
1515
1516
1517
1518
1519
1520
1521
1522
1523
1524
1525
1526
1527
1528
1529
1530
1531
1532
1533
1534
1535
1536
1537
1538
1539
1540
1541
1542
1543
1544
1545
1546
1547
1548
1549
1550
1551
1552
1553
1554
1555
1556
1557
1558
1559
1560
1561
1562
1563
1564
1565
1566
1567
1568
1569
1570
1571
1572
1573
1574
1575
1576
1577
1578
1579
1580
1581
1582
1583
1584
1585
1586
1587
1588
1589
1590
1591
1592
1593
1594
1595
1596
1597
1598
1599
1600
1601
1602
1603
1604
1605
1606
1607
1608
1609
1610
1611
1612
1613
1614
1615
1616
1617
1618
1619
1620
1621
1622
1623
1624
1625
1626
1627
1628
1629
1630
1631
1632
1633
1634
1635
1636
1637
1638
1639
@torch.no_grad()
def minimax_m3_index_decode(
    idx_q: torch.Tensor,  # [total_q, num_idx_heads, head_dim]
    index_kv_cache: torch.Tensor,  # [num_blocks, 128, head_dim]
    block_table: torch.Tensor,  # [num_reqs, max_blocks]
    seq_lens: torch.Tensor,  # [num_reqs] int32
    max_seq_len: int,
    topk: int,
    init_blocks: int,
    local_blocks: int,
    num_kv_heads: int,
    decode_query_len: int,
    max_decode_query_len: int,
    out: torch.Tensor | None = None,
    *,
    attention_block_table: torch.Tensor | None = None,
    sparse_block_table_out: torch.Tensor | None = None,
    sparse_context_lens_out: torch.Tensor | None = None,
    block_page_stride: int | None = None,
    completion_counter: torch.Tensor | None = None,
) -> torch.Tensor:
    """Decode index block-score followed by fused adaptive top-k selection.

    Returns topk_idx [num_kv_heads, total_q, topk] (0-indexed block ids, -1 pad).
    When ``out`` ([num_kv_heads, >=total_q, topk]) is given, writes into
    ``out[:, :total_q, :]`` (stable address for cudagraph) instead of allocating.
    The optional sparse-table arguments fuse current-layer table construction
    into the selector. They must be provided together. ``completion_counter``
    provides stable per-query synchronization storage for CUDA graphs. It must
    be zero before its first launch and must not be shared by overlapping
    selector invocations; every completed launch resets its active entries.
    """
    total_q, num_idx_heads, head_dim = idx_q.shape
    assert num_idx_heads == num_kv_heads, (
        "M3 expects num_idx_heads == num_kv_heads (no topk index reduce)"
    )
    assert 0 < decode_query_len <= max_decode_query_len
    assert total_q == seq_lens.shape[0] * decode_query_len
    batch = total_q
    emit_sparse_table = attention_block_table is not None
    if emit_sparse_table and (
        sparse_block_table_out is None
        or sparse_context_lens_out is None
        or block_page_stride is None
    ):
        raise ValueError(
            "MiniMax-M3 fused decode sparse-table arguments must be provided together"
        )
    if not emit_sparse_table and (
        sparse_block_table_out is not None
        or sparse_context_lens_out is not None
        or block_page_stride is not None
    ):
        raise ValueError(
            "MiniMax-M3 fused decode sparse-table arguments must be provided together"
        )
    if emit_sparse_table:
        assert attention_block_table is not None
        assert sparse_block_table_out is not None
        assert sparse_context_lens_out is not None
        assert block_page_stride is not None
        if num_idx_heads != 1:
            raise ValueError(
                "MiniMax-M3 fused decode sparse-table construction requires "
                f"one index head, got {num_idx_heads}"
            )
        if block_page_stride not in (
            PAGES_PER_SPARSE_BLOCK,
            2 * PAGES_PER_SPARSE_BLOCK,
        ):
            raise ValueError(
                "MiniMax-M3 fused decode sparse-table page stride must be "
                f"{PAGES_PER_SPARSE_BLOCK} or {2 * PAGES_PER_SPARSE_BLOCK}, "
                f"got {block_page_stride}"
            )
        expected_sparse_shape = (total_q, topk * PAGES_PER_SPARSE_BLOCK)
        if sparse_block_table_out.shape != expected_sparse_shape:
            raise ValueError(
                "MiniMax-M3 fused decode sparse block table has shape "
                f"{tuple(sparse_block_table_out.shape)}, expected "
                f"{expected_sparse_shape}"
            )
        if sparse_context_lens_out.shape != (total_q,):
            raise ValueError(
                "MiniMax-M3 fused decode sparse context lengths have shape "
                f"{tuple(sparse_context_lens_out.shape)}, expected {(total_q,)}"
            )
        if (
            attention_block_table.dim() != 2
            or attention_block_table.shape[0] != seq_lens.shape[0]
        ):
            raise ValueError(
                "MiniMax-M3 fused decode attention block table must be "
                "[num_requests, max_blocks]"
            )
        int32_tensors = (
            attention_block_table,
            sparse_block_table_out,
            sparse_context_lens_out,
        )
        if any(tensor.dtype != torch.int32 for tensor in int32_tensors):
            raise ValueError("MiniMax-M3 fused decode sparse tables require int32")
        if any(tensor.device != idx_q.device for tensor in int32_tensors):
            raise ValueError(
                "MiniMax-M3 fused decode sparse tables must share the query device"
            )
        if (
            attention_block_table.stride(1) != 1
            or sparse_block_table_out.stride(1) != 1
            or sparse_context_lens_out.stride(0) != 1
        ):
            raise ValueError(
                "MiniMax-M3 fused decode sparse tables require contiguous rows"
            )
    max_block = triton.cdiv(max_seq_len, SPARSE_BLOCK_SIZE)
    if max_block >= 0xFFFF:
        raise ValueError(
            "MiniMax-M3 decode top-k supports fewer than 65535 sparse blocks"
        )
    if emit_sparse_table:
        assert attention_block_table is not None
        if attention_block_table.shape[1] < max_block:
            raise ValueError(
                "MiniMax-M3 fused decode attention block table is shorter "
                f"than the required {max_block} sparse blocks"
            )
    use_pdl = current_platform.is_arch_support_pdl()
    # `launch_pdl` is a Triton runtime kwarg only some backends accept (CUDA
    # SM9+); this ROCm Triton rejects it even when False ("Keyword argument
    # launch_pdl was specified but unrecognised"). Only pass it when PDL is
    # actually supported -- on ROCm use_pdl is always False, so it's omitted.
    pdl_kwargs: dict[str, bool | int] = {}
    if use_pdl:
        pdl_kwargs.update({"launch_pdl": True})
    is_gfx950 = False
    if current_platform.is_rocm():
        from vllm.platforms.rocm import on_gfx950

        is_gfx950 = on_gfx950()
    # Multi-head spec decode scores a wider head-position tile per K block;
    # reduce stages to ease memory/register pressure on the fallback path.
    score_kwargs = pdl_kwargs.copy()
    if num_idx_heads > 1 and max_decode_query_len > 1:
        score_kwargs.update({"num_warps": 4, "num_stages": 2})

    # Keep score strides 16-divisible to avoid Triton recompiles.
    score_block_stride = round_up(max_block, 16)
    score = torch.empty(
        (num_idx_heads, total_q, score_block_stride),
        dtype=torch.float32,
        device=idx_q.device,
    )
    # Use the configured max decode length to avoid Triton recompiles when
    # switching between qlen=1 and spec-decode verification batches.
    BLOCK_SIZE_Q = triton.next_power_of_2(max_decode_query_len)
    num_reqs = seq_lens.shape[0]
    score_program_budget = _decode_score_program_budget(
        num_reqs,
        head_dim,
        idx_q.dtype,
        index_kv_cache.dtype,
        is_gfx950=is_gfx950,
    )
    grid_score: tuple[int, ...]
    if score_program_budget is not None:
        grid_score = (score_program_budget + num_reqs - 1,)
        _decode_index_score_balanced_kernel[grid_score](
            idx_q,
            index_kv_cache,
            score,
            block_table,
            seq_lens,
            num_idx_heads,
            head_dim,
            init_blocks,
            local_blocks,
            num_reqs,
            score_program_budget,
            decode_query_len,
            idx_q.stride(0),
            idx_q.stride(1),
            idx_q.stride(2),
            index_kv_cache.stride(0),
            index_kv_cache.stride(1),
            index_kv_cache.stride(2),
            score.stride(0),
            score.stride(1),
            score.stride(2),
            block_table.stride(0),
            BLOCK_SIZE_K=SPARSE_BLOCK_SIZE,
            BLOCK_SIZE_Q=BLOCK_SIZE_Q,
            num_warps=2,
            num_stages=1,
        )
    else:
        # Increase independent work for the measured high-batch gfx950 decode
        # range while preserving the deployed split for every other shape.
        num_kv_chunks, use_high_batch_config = _decode_score_split_launch_policy(
            num_reqs,
            head_dim,
            idx_q.dtype,
            index_kv_cache.dtype,
            is_gfx950=is_gfx950,
        )
        if use_high_batch_config and not (
            num_idx_heads > 1 and max_decode_query_len > 1
        ):
            score_kwargs.update({"num_warps": 2, "num_stages": 1})
        grid_score = (num_reqs, num_kv_chunks)
        _decode_index_score_kernel[grid_score](
            idx_q,
            index_kv_cache,
            score,
            block_table,
            seq_lens,
            num_idx_heads,
            head_dim,
            init_blocks,
            local_blocks,
            decode_query_len,
            idx_q.stride(0),
            idx_q.stride(1),
            idx_q.stride(2),
            index_kv_cache.stride(0),
            index_kv_cache.stride(1),
            index_kv_cache.stride(2),
            score.stride(0),
            score.stride(1),
            score.stride(2),
            block_table.stride(0),
            BLOCK_SIZE_K=SPARSE_BLOCK_SIZE,
            BLOCK_SIZE_Q=BLOCK_SIZE_Q,
            num_kv_chunks=num_kv_chunks,
            USE_PDL=use_pdl,
            **score_kwargs,
        )

    if out is not None:
        topk_idx = out[:, :total_q, :]
    else:
        topk_idx = torch.empty(
            (num_idx_heads, total_q, topk),
            dtype=torch.int32,
            device=idx_q.device,
        )
    # The launch grid remains shape-constant for CUDA graphs. Each query uses
    # only the number of chunks needed for its live context.
    (
        num_topk_chunks,
        topk_block_size,
        selector_num_warps,
        selector_num_stages,
        single_tile_guaranteed,
        adaptive_final_merge,
    ) = _decode_topk_launch_policy(
        max_block,
        batch,
        num_idx_heads,
        topk,
        is_gfx950=is_gfx950,
    )
    block_size_t = triton.next_power_of_2(topk)
    topk_partial = torch.empty(
        num_topk_chunks,
        num_idx_heads,
        batch,
        block_size_t,
        dtype=torch.int64,
        device=idx_q.device,
    )
    if completion_counter is None:
        active_counter = torch.zeros(
            (num_idx_heads, batch),
            dtype=torch.int32,
            device=idx_q.device,
        )
    else:
        if (
            completion_counter.dim() != 2
            or completion_counter.shape[0] != num_idx_heads
            or completion_counter.shape[1] < batch
            or completion_counter.dtype != torch.int32
            or completion_counter.device != idx_q.device
            or not completion_counter.is_contiguous()
        ):
            raise ValueError(
                "MiniMax-M3 completion counter must be contiguous int32 "
                "[num_idx_heads, >=total_q] on the query device"
            )
        active_counter = completion_counter[:, :batch]

    selector_attention_block_table = block_table
    selector_sparse_block_table = topk_idx
    selector_sparse_context_lens = seq_lens
    selector_block_page_stride = PAGES_PER_SPARSE_BLOCK
    selector_sparse_block_stride = topk_idx.stride(1)
    if emit_sparse_table:
        assert attention_block_table is not None
        assert sparse_block_table_out is not None
        assert sparse_context_lens_out is not None
        assert block_page_stride is not None
        selector_attention_block_table = attention_block_table
        selector_sparse_block_table = sparse_block_table_out
        selector_sparse_context_lens = sparse_context_lens_out
        selector_block_page_stride = block_page_stride
        selector_sparse_block_stride = sparse_block_table_out.stride(0)
    _decode_topk_fused_kernel[(batch, num_idx_heads, num_topk_chunks)](
        score,
        topk_partial,
        active_counter,
        topk_idx,
        seq_lens,
        selector_attention_block_table,
        selector_sparse_block_table,
        selector_sparse_context_lens,
        decode_query_len,
        score.stride(0),
        score.stride(1),
        score.stride(2),
        topk_partial.stride(0),
        topk_partial.stride(1),
        topk_partial.stride(2),
        topk_partial.stride(3),
        active_counter.stride(0),
        active_counter.stride(1),
        topk_idx.stride(0),
        topk_idx.stride(1),
        topk_idx.stride(2),
        selector_attention_block_table.stride(0),
        selector_sparse_block_stride,
        topk=topk,
        block_size=SPARSE_BLOCK_SIZE,
        pages_per_sparse_block=PAGES_PER_SPARSE_BLOCK,
        block_page_stride=selector_block_page_stride,
        NUM_TOPK_CHUNKS=num_topk_chunks,
        BLOCK_SIZE_K=topk_block_size,
        BLOCK_SIZE_T=block_size_t,
        EMIT_SPARSE_TABLE=emit_sparse_table,
        SINGLE_TILE_GUARANTEED=single_tile_guaranteed,
        ADAPTIVE_FINAL_MERGE=adaptive_final_merge,
        num_warps=selector_num_warps,
        num_stages=selector_num_stages,
    )
    return topk_idx

minimax_m3_index_score(idx_q, index_kv_cache, block_table, cu_seqlens_q, seq_lens, prefix_lens, max_query_len, max_seq_len, num_kv_heads)

Compute per-token index scores for each visible sparse block.

Returns score [num_kv_heads, total_q, max_block], where each score is the max over a 128-token index-K block. M3 has num_idx_heads == num_kv_heads.

Source code in vllm/models/minimax_m3/amd/ops/index_topk.py
@torch.no_grad()
def minimax_m3_index_score(
    idx_q: torch.Tensor,  # [total_q, num_idx_heads, head_dim]
    index_kv_cache: torch.Tensor,  # [num_blocks, 128, head_dim]
    block_table: torch.Tensor,  # [batch, max_blocks]
    cu_seqlens_q: torch.Tensor,  # [batch+1] int32
    seq_lens: torch.Tensor,  # [batch] int32
    prefix_lens: torch.Tensor,  # [batch] int32
    max_query_len: int,
    max_seq_len: int,
    num_kv_heads: int,
) -> torch.Tensor:
    """Compute per-token index scores for each visible sparse block.

    Returns score [num_kv_heads, total_q, max_block], where each score is the
    max over a 128-token index-K block. M3 has num_idx_heads == num_kv_heads.
    """
    total_q, num_idx_heads, head_dim = idx_q.shape
    assert num_idx_heads == num_kv_heads, (
        "M3 expects num_idx_heads == num_kv_heads (no topk index reduce)"
    )
    batch = cu_seqlens_q.shape[0] - 1
    max_block = triton.cdiv(max_seq_len, SPARSE_BLOCK_SIZE)

    # Keep score strides 16-divisible to avoid Triton recompiles.
    score_block_stride = round_up(max_block, 16)
    score = torch.empty(
        (num_idx_heads, total_q, score_block_stride),
        dtype=torch.float32,
        device=idx_q.device,
    )
    BLOCK_SIZE_Q = 64
    grid_score = (triton.cdiv(max_query_len, BLOCK_SIZE_Q), batch * num_idx_heads)
    _index_block_score_kernel[grid_score](
        idx_q,
        index_kv_cache,
        score,
        block_table,
        cu_seqlens_q,
        seq_lens,
        prefix_lens,
        num_idx_heads,
        head_dim,
        idx_q.stride(0),
        idx_q.stride(1),
        idx_q.stride(2),
        index_kv_cache.stride(0),
        index_kv_cache.stride(1),
        index_kv_cache.stride(2),
        score.stride(0),
        score.stride(1),
        score.stride(2),
        block_table.stride(0),
        BLOCK_SIZE_Q=BLOCK_SIZE_Q,
        BLOCK_SIZE_K=SPARSE_BLOCK_SIZE,
    )
    return score

minimax_m3_index_topk(score, cu_seqlens_q, prefix_lens, max_query_len, topk, init_blocks, local_blocks, out=None)

Select index top-k from a precomputed score tensor.

When out is provided (a [num_idx_heads, >=total_q, topk] buffer), the result is written into out[:, :total_q, :] instead of a fresh tensor -- used to keep the top-k output at a stable address for cudagraph capture.

Source code in vllm/models/minimax_m3/amd/ops/index_topk.py
@torch.no_grad()
def minimax_m3_index_topk(
    score: torch.Tensor,  # [num_idx_heads, total_q, max_block]
    cu_seqlens_q: torch.Tensor,  # [batch+1] int32
    prefix_lens: torch.Tensor,  # [batch] int32
    max_query_len: int,
    topk: int,
    init_blocks: int,
    local_blocks: int,
    out: torch.Tensor | None = None,
) -> torch.Tensor:
    """Select index top-k from a precomputed score tensor.

    When ``out`` is provided (a ``[num_idx_heads, >=total_q, topk]`` buffer), the
    result is written into ``out[:, :total_q, :]`` instead of a fresh tensor --
    used to keep the top-k output at a stable address for cudagraph capture.
    """
    num_idx_heads = score.shape[0]
    batch = cu_seqlens_q.shape[0] - 1
    total_q = score.shape[1]
    if out is not None:
        topk_idx = out[:, :total_q, :]
    else:
        topk_idx = torch.empty(
            (num_idx_heads, total_q, topk),
            dtype=torch.int32,
            device=score.device,
        )
    # block_size_q == 1 -> query blocks coincide with query tokens.
    grid_topk = (max_query_len, batch, num_idx_heads)
    _topk_index_kernel[grid_topk](
        score,
        topk_idx,
        1,  # sample_interval (block_size_q)
        SPARSE_BLOCK_SIZE,
        cu_seqlens_q,
        cu_seqlens_q,  # cu_seqblocks_q == cu_seqlens_q when block_size_q == 1
        prefix_lens,
        topk,
        init_blocks,
        local_blocks,
        score.stride(0),
        score.stride(1),
        score.stride(2),
        topk_idx.stride(0),
        topk_idx.stride(1),
        topk_idx.stride(2),
        MASK_INIT=False,
        MASK_LOCAL=False,
    )
    return topk_idx