# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project


import numpy as np
import pytest
import torch

from vllm.platforms import current_platform
from vllm.utils.torch_utils import set_random_seed

# Test parameters
NUM_ROWS = [1, 32, 2050]
TOP_K_VALUES = [2048, 3000]
BATCH_SIZE = [1, 2, 2048]
NEXT_N = [1, 8]
DATA_GENERATION = ["random", "10LSBits"]
RADIX_TOPK_WORKSPACE_SIZE = 1024 * 1024


def _has_device_capability(major: int) -> bool:
    return current_platform.is_cuda() and current_platform.has_device_capability(major)


def _on_gfx950() -> bool:
    if not current_platform.is_rocm():
        return False
    from vllm.platforms.rocm import on_gfx950

    return on_gfx950()


# DeepSelect is compiled for sm_100a/sm_103a only.
requires_sm100 = pytest.mark.skipif(
    not current_platform.is_device_capability_family(100),
    reason="DeepSelect requires SM100a/SM103a",
)
requires_gfx950 = pytest.mark.skipif(
    not _on_gfx950(), reason="This test exercises the gfx950 launch configuration"
)


COOPERATIVE_TOPK_BACKEND = pytest.param(
    "cooperative_topk",
    marks=pytest.mark.skipif(
        not _has_device_capability(90),
        reason="cooperative_topk requires SM90+",
    ),
)
WORKSPACE_TOPK_BACKENDS = ["persistent_topk", COOPERATIVE_TOPK_BACKEND]
TOPK_BACKENDS = ["top_k_per_row_decode", *WORKSPACE_TOPK_BACKENDS]


def _run_topk_backend(
    backend: str,
    logits: torch.Tensor,
    lengths: torch.Tensor,
    indices: torch.Tensor,
    top_k: int,
    max_seq_len: int,
    next_n: int = 1,
) -> None:
    if backend == "top_k_per_row_decode":
        torch.ops._C.top_k_per_row_decode(
            logits,
            next_n,
            lengths,
            indices,
            indices.shape[0],
            logits.stride(0),
            logits.stride(1),
            top_k,
        )
    elif backend == "persistent_topk":
        workspace = torch.empty(
            RADIX_TOPK_WORKSPACE_SIZE, dtype=torch.uint8, device="cuda"
        )
        torch.ops._C.persistent_topk(
            logits, lengths, indices, workspace, top_k, max_seq_len
        )
    elif backend == "cooperative_topk":
        if indices.shape[0] > 64:
            pytest.skip(
                "cooperative_topk supports <=64 rows; "
                "persistent_topk covers larger batches"
            )
        if logits.stride(0) % 4 != 0:
            pytest.skip(
                "cooperative_topk requires row stride divisible by 4; "
                "persistent_topk covers unaligned strides"
            )
        workspace = torch.empty(
            RADIX_TOPK_WORKSPACE_SIZE, dtype=torch.uint8, device="cuda"
        )
        torch.ops._C.cooperative_topk(
            logits, lengths, indices, workspace, top_k, max_seq_len
        )
    else:
        raise ValueError(f"Unknown top-k backend: {backend}")


def create_random_logits(
    row_starts: torch.Tensor,
    row_ends: torch.Tensor,
    dtype: torch.dtype,
    seed: int,
    clean_logits: bool,
    data_generation: str,
) -> torch.Tensor:
    """Create random logits tensor for testing."""
    torch.manual_seed(seed)
    np.random.seed(seed)
    # Generate logits with some structure to make testing more meaningful
    if data_generation == "random":
        logits = torch.randn(
            row_starts.shape[0], max(row_ends), dtype=dtype, device="cuda"
        )
    elif data_generation == "10LSBits":
        top_22_bits_mask = 0xFFFFFC00
        last_10_bits_mask = 0x000003FF
        fixed_top_22_bits = 0x3F900000
        # Generate random bits for the last 10 bits
        random_bottom_bits = torch.randint(
            0,
            2**10,
            (row_starts.shape[0], max(row_ends)),
            dtype=torch.int32,
            device="cuda",
        )
        # Combine: fixed top 22 bits with random last 10 bits
        logits_bits = (fixed_top_22_bits & top_22_bits_mask) | (
            random_bottom_bits & last_10_bits_mask
        )
        logits = logits_bits.view(dtype)

    if clean_logits:
        for i, end in enumerate(row_ends):
            logits[i, end:] = float("-inf")
    return logits


def create_row_boundaries(
    seq_len: int, vocab_size: int
) -> tuple[torch.Tensor, torch.Tensor]:
    """Create row start and end indices for testing."""
    row_starts = torch.zeros(seq_len, dtype=torch.int32, device="cuda")
    row_ends = torch.arange(1, seq_len + 1, device="cuda", dtype=torch.int32)
    return row_starts, row_ends


def compare_top_k_results(
    logits: torch.Tensor,
    cuda_indices: torch.Tensor,
    torch_indices: torch.Tensor,
    row_starts: torch.Tensor,
    row_ends: torch.Tensor,
    top_k: int,
    tolerance: float = 1e-5,
) -> bool:
    """Compare results from CUDA top_k_per_row with torch.topk.
    Both results should be sorted and contain the same top-k elements.
    """
    num_rows = cuda_indices.shape[0]

    for row_idx in range(num_rows):
        # Get valid elements using row boundaries
        row_start = row_starts[row_idx].item()
        row_end = row_ends[row_idx].item()
        row_length = row_end - row_start
        num_valid = min(top_k, row_length)
        cuda_row_indices = cuda_indices[row_idx][:num_valid].cpu()
        torch_row_indices = torch_indices[row_idx][:num_valid].cpu()

        # Compare the sets of indices first
        cuda_set = set(cuda_row_indices.tolist())
        torch_set = set(torch_row_indices.tolist())
        if cuda_set == torch_set:
            continue

        # Any difference in elements, compare the values
        logits_row = logits[row_idx]
        cuda_row_values = [logits_row[i] for i in cuda_row_indices]
        torch_row_values = [logits_row[i] for i in torch_row_indices]

        cuda_only_values, torch_only_values = [], []
        for idx in cuda_set - torch_set:
            cuda_pos = (cuda_row_indices == idx).nonzero(as_tuple=True)[0]
            cuda_only_values.append(cuda_row_values[cuda_pos[0]])

        for idx in torch_set - cuda_set:
            torch_pos = (torch_row_indices == idx).nonzero(as_tuple=True)[0]
            torch_only_values.append(torch_row_values[torch_pos[0]])

        if len(cuda_only_values) != len(torch_only_values):
            return False
        if not torch.allclose(
            torch.tensor(cuda_only_values),
            torch.tensor(torch_only_values),
            rtol=tolerance,
            atol=tolerance,
        ):
            return False

    return True


def validate_topk_against_reference(
    logits: torch.Tensor,
    cuda_indices: torch.Tensor,
    row_starts: torch.Tensor,
    row_ends: torch.Tensor,
    top_k: int,
    kernel_name: str,
) -> None:
    """Validate CUDA top-k results against PyTorch reference implementation.

    Args:
        logits: Input logits tensor
        cuda_indices: CUDA kernel output indices
        row_starts: Row start positions
        row_ends: Row end positions
        top_k: Number of top elements to select
        kernel_name: Name of the kernel being tested (for error messages)

    """
    num_rows = cuda_indices.shape[0]
    torch_indices = torch.empty((num_rows, top_k), dtype=torch.int32, device="cuda")

    for i in range(num_rows):
        row_end = int(row_ends[i])
        k_i = min(top_k, row_end)
        idx = logits[i, :row_end].topk(k_i, dim=-1)[1]
        torch_indices[i, :k_i] = idx

    assert compare_top_k_results(
        logits, cuda_indices, torch_indices, row_starts, row_ends, top_k
    ), f"{kernel_name} results don't match torch.topk"


@pytest.mark.parametrize("num_rows", NUM_ROWS)
@pytest.mark.parametrize("top_k", TOP_K_VALUES)
@pytest.mark.parametrize("clean_logits", [True, False])
@pytest.mark.skipif(not current_platform.is_cuda(), reason="This test requires CUDA")
@torch.inference_mode()
def test_top_k_per_row(
    num_rows: int,
    top_k: int,
    clean_logits: bool,
) -> None:
    """Test top_k_per_row."""
    set_random_seed(0)
    torch.set_default_device("cuda:0")

    # Create test data
    vocab_size = 20000
    row_starts, row_ends = create_row_boundaries(num_rows, vocab_size)
    logits = create_random_logits(
        row_starts, row_ends, torch.float32, 42, clean_logits, "random"
    )

    # Create output tensors
    indices = torch.empty((num_rows, top_k), dtype=torch.int32, device="cuda")

    # Run CUDA implementation
    torch.ops._C.top_k_per_row_prefill(
        logits,
        row_starts,
        row_ends,
        indices,
        num_rows,
        logits.stride(0),
        logits.stride(1),
        top_k,
    )

    # Run reference implementation
    torch_indices = torch.empty((num_rows, top_k), dtype=torch.int32, device="cuda")
    for i in range(num_rows):
        row_end = int(row_ends[i])
        k_i = min(top_k, row_end)
        idx = logits[i, :row_end].topk(k_i, dim=-1)[1]
        torch_indices[i, :k_i] = idx

    # Compare results
    assert compare_top_k_results(
        logits, indices, torch_indices, row_starts, row_ends, top_k
    ), "CUDA top_k_per_row_prefill results don't match torch.topk"


def _run_top_k_per_row_decode_test(
    top_k: int,
    batch_size: int,
    next_n: int,
    vocab_size: int,
    clean_logits: bool,
    data_generation: str,
) -> None:
    """Helper function to run top_k_per_row_decode test with given parameters."""
    torch.set_default_device("cuda:0")

    # Create test data
    num_rows = batch_size * next_n
    seq_lens = torch.randint(
        low=next_n,
        high=vocab_size,
        size=(batch_size,),
        dtype=torch.int32,
        device="cuda",
    )
    row_starts = torch.zeros(num_rows, dtype=torch.int32, device="cuda")
    row_indices = torch.arange(num_rows, device="cuda") // next_n
    next_n_offset = torch.arange(num_rows, device="cuda") % next_n
    row_ends = seq_lens[row_indices] - next_n + next_n_offset + 1
    logits = create_random_logits(
        row_starts, row_ends, torch.float32, 42, clean_logits, data_generation
    )

    # Create output tensors
    indices = torch.empty((num_rows, top_k), dtype=torch.int32, device="cuda")

    # Run CUDA implementation
    torch.ops._C.top_k_per_row_decode(
        logits,
        next_n,
        seq_lens,
        indices,
        num_rows,
        logits.stride(0),
        logits.stride(1),
        top_k,
    )

    torch.accelerator.synchronize()

    # Run reference implementation
    torch_indices = torch.empty((num_rows, top_k), dtype=torch.int32, device="cuda")
    for i in range(num_rows):
        row_end = int(row_ends[i])
        k_i = min(top_k, row_end)
        idx = logits[i, :row_end].topk(k_i, dim=-1)[1]
        torch_indices[i, :k_i] = idx

    # Compare results
    assert compare_top_k_results(
        logits, indices, torch_indices, row_starts, row_ends, top_k
    ), "CUDA top_k_per_row_decode results don't match torch.topk"


@pytest.mark.parametrize("top_k", TOP_K_VALUES)
@pytest.mark.parametrize("batch_size", BATCH_SIZE)
@pytest.mark.parametrize("next_n", NEXT_N)
@pytest.mark.parametrize("clean_logits", [True, False])
@pytest.mark.parametrize("data_generation", DATA_GENERATION)
@pytest.mark.skipif(not current_platform.is_cuda(), reason="This test requires CUDA")
@torch.inference_mode()
def test_top_k_per_row_decode(
    top_k: int,
    batch_size: int,
    next_n: int,
    clean_logits: bool,
    data_generation: str,
) -> None:
    """Test top_k_per_row with seq_lens tensor."""
    set_random_seed(0)
    vocab_size = 20000
    _run_top_k_per_row_decode_test(
        top_k, batch_size, next_n, vocab_size, clean_logits, data_generation
    )


@pytest.mark.skipif(not current_platform.is_cuda(), reason="This test requires CUDA")
@pytest.mark.parametrize("clean_logits", [True, False])
@torch.inference_mode()
def test_top_k_per_row_decode_large_vocab_size(clean_logits: bool) -> None:
    """Test top_k_per_row_decode with large vocabulary size."""
    set_random_seed(0)
    top_k = 2048
    batch_size = 2
    next_n = 2
    vocab_size = 300000
    data_generation = "random"
    _run_top_k_per_row_decode_test(
        top_k, batch_size, next_n, vocab_size, clean_logits, data_generation
    )


@pytest.mark.parametrize(
    "batch_size,num_speculative_tokens,min_seq_len,max_seq_len",
    [
        pytest.param(1, 2, 125_000, 250_000, id="eight-splits"),
        pytest.param(7, 4, 125_000, 250_000, id="four-splits"),
        pytest.param(16, 4, 2_500, 2_501, id="dynamic-short-rows"),
        pytest.param(13, 4, 25_000, 25_001, id="dynamic-three-splits"),
        pytest.param(16, 4, 65_535, 65_536, id="dynamic-boundary-mixed"),
        pytest.param(16, 4, 125_000, 125_001, id="dynamic-two-splits"),
        pytest.param(64, 3, 125_000, 125_001, id="dynamic-two-splits-max-rows"),
        pytest.param(64, 4, 125_000, 125_001, id="baseline-fallback"),
    ],
)
@pytest.mark.skipif(not current_platform.is_rocm(), reason="This test requires ROCm")
@torch.inference_mode()
def test_top_k_per_row_decode_gfx950_long_c4a(
    batch_size: int,
    num_speculative_tokens: int,
    min_seq_len: int,
    max_seq_len: int,
) -> None:
    properties = torch.cuda.get_device_properties(0)
    if not properties.gcnArchName.startswith("gfx950"):
        pytest.skip("This test exercises the gfx950 launch configuration")

    query_tokens = num_speculative_tokens + 1
    num_rows = batch_size * query_tokens
    stride = 262_144
    top_k = 1024
    offsets = torch.arange(num_rows, dtype=torch.int32, device="cuda")
    seq_lens = min_seq_len + offsets * (max_seq_len - min_seq_len) // max(
        num_rows - 1, 1
    )
    seq_lens = seq_lens.reshape(batch_size, query_tokens)
    logits = torch.randn(num_rows, stride, dtype=torch.float32, device="cuda")
    indices = torch.empty((num_rows, top_k), dtype=torch.int32, device="cuda")

    torch.ops._C.top_k_per_row_decode(
        logits,
        query_tokens,
        seq_lens,
        indices,
        num_rows,
        logits.stride(0),
        logits.stride(1),
        top_k,
    )

    row_starts = torch.zeros(num_rows, dtype=torch.int32, device="cuda")
    validate_topk_against_reference(
        logits,
        indices,
        row_starts,
        seq_lens.reshape(-1),
        top_k,
        "gfx950 long C4A top-k",
    )


@pytest.mark.skipif(not current_platform.is_rocm(), reason="This test requires ROCm")
@torch.inference_mode()
def test_top_k_per_row_decode_gfx950_long_c4a_1d_seq_lens() -> None:
    properties = torch.cuda.get_device_properties(0)
    if not properties.gcnArchName.startswith("gfx950"):
        pytest.skip("This test exercises the gfx950 launch configuration")

    batch_size = 20
    query_tokens = 4
    num_rows = batch_size * query_tokens
    stride = 262_144
    top_k = 1024
    seq_lens = torch.linspace(
        125_003,
        250_000,
        batch_size,
        dtype=torch.int32,
        device="cuda",
    )
    row_offsets = torch.arange(query_tokens, dtype=torch.int32, device="cuda")
    row_ends = (seq_lens[:, None] - query_tokens + row_offsets + 1).reshape(-1)
    logits = torch.randn(num_rows, stride, dtype=torch.float32, device="cuda")
    indices = torch.empty((num_rows, top_k), dtype=torch.int32, device="cuda")

    torch.ops._C.top_k_per_row_decode(
        logits,
        query_tokens,
        seq_lens,
        indices,
        num_rows,
        logits.stride(0),
        logits.stride(1),
        top_k,
    )

    row_starts = torch.zeros(num_rows, dtype=torch.int32, device="cuda")
    validate_topk_against_reference(
        logits,
        indices,
        row_starts,
        row_ends,
        top_k,
        "gfx950 long C4A top-k with 1-D sequence lengths",
    )


def _assert_exact_topk(
    logits: torch.Tensor, indices: torch.Tensor, row_ends: torch.Tensor
) -> None:
    ends = row_ends.reshape(-1, 1).clamp_min(0)
    valid = (indices >= 0) & (indices < ends)
    assert torch.equal(valid.sum(1), ends.clamp_max(indices.shape[1]).flatten())
    assert torch.all(valid | (indices == -1))
    ordered = indices.sort(1).values
    assert not ((ordered[:, 1:] == ordered[:, :-1]) & (ordered[:, 1:] >= 0)).any()

    width = int(ends.max())
    visible = logits[:, :width].clone()
    columns = torch.arange(width, device=logits.device)
    visible.masked_fill_(columns[None, :] >= ends, float("-inf"))
    expected = visible.topk(min(indices.shape[1], width), dim=1).values
    selected = logits.gather(1, indices.clamp_min(0).long())
    selected.masked_fill_(~valid, float("-inf"))
    selected = selected.sort(1, descending=True).values
    assert torch.equal(selected[:, : expected.shape[1]], expected)


@requires_gfx950
@torch.inference_mode()
def test_top_k_per_row_decode_gfx950_k512_1d_seq_lens() -> None:
    """Honor causal offsets and exact compressed lengths, including short rows."""
    next_n, width, top_k = 6, 16_384, 512
    lengths = torch.tensor(
        [0, 1, 511, 512, 513, 1023, 1024, 16_000],
        dtype=torch.int32,
        device="cuda",
    )
    offsets = torch.arange(next_n, dtype=torch.int32, device="cuda")
    row_ends = (lengths[:, None] - next_n + offsets + 1).clamp_min(0).flatten()
    rows = row_ends.numel()
    storage = torch.full((rows, width + 8), 1e20, dtype=torch.float32, device="cuda")
    logits = storage[:, 1 : width + 1]
    logits[:, :16_000] = torch.arange(16_000, dtype=torch.float32, device="cuda")
    indices = torch.full((rows, top_k), -777, dtype=torch.int32, device="cuda")
    _run_topk_backend(
        "top_k_per_row_decode", logits, lengths, indices, top_k, width, next_n
    )
    _assert_exact_topk(logits, indices, row_ends)


@pytest.mark.parametrize(
    ("rows", "row_length"),
    [
        pytest.param(128, 65_536, id="single-before-64k"),
        pytest.param(128, 65_537, id="four-splits-after-64k"),
        pytest.param(129, 262_144, id="single-before-three-split-length"),
        pytest.param(129, 262_145, id="three-splits-after-256k"),
        pytest.param(160, 262_145, id="three-splits-upper-row-bound"),
        pytest.param(161, 262_144, id="single-before-two-split-length"),
        pytest.param(161, 262_145, id="two-splits-after-256k"),
        pytest.param(256, 262_145, id="two-splits-upper-row-bound"),
        pytest.param(257, 262_145, id="large-row-single-block"),
        pytest.param(384, 262_145, id="optimized-row-limit"),
        pytest.param(385, 262_145, id="native-fallback"),
    ],
)
@requires_gfx950
@torch.inference_mode()
def test_top_k_per_row_decode_gfx950_k512_split_policy_boundaries(
    rows: int, row_length: int
) -> None:
    width, top_k = 262_145, 512
    row = torch.arange(width, dtype=torch.float32, device="cuda")
    logits = row.expand(rows, -1)
    lengths = torch.full((rows,), row_length, dtype=torch.int32, device="cuda")
    indices = torch.empty((rows, top_k), dtype=torch.int32, device="cuda")

    _run_topk_backend("top_k_per_row_decode", logits, lengths, indices, top_k, width)

    expected = torch.arange(
        row_length - top_k, row_length, dtype=torch.int32, device="cuda"
    ).expand_as(indices)
    assert torch.equal(indices.sort(1).values, expected)


@requires_gfx950
@torch.inference_mode()
def test_top_k_per_row_decode_gfx950_k512_masked_graph_replay() -> None:
    """Preserve visible top-k scores when masks and lengths change on replay."""
    rows, next_n, width, top_k = 24, 6, 262_145, 512
    logits = torch.full((rows, width), 1e20, dtype=torch.float32, device="cuda")
    lengths = torch.full(
        (rows // next_n, next_n), 70_000, dtype=torch.int32, device="cuda"
    )
    indices = torch.full((rows, top_k), -777, dtype=torch.int32, device="cuda")
    logits[:, :70_000] = float("-inf")
    bin_sizes = [128, 129, 1024, 1025]
    for row in range(rows):
        pattern = row % 9
        if pattern < 4:
            count = [0, 1, 511, 513][pattern]
            logits[row, :count] = torch.arange(
                count, dtype=torch.float32, device="cuda"
            )
        elif pattern == 4:
            logits[row, :511] = 0
            logits[row, 511] = torch.finfo(torch.float32).min
        else:
            count = bin_sizes[pattern - 5]
            higher = max(0, top_k - count // 2)
            logits[row, :higher] = 2
            logits[row, higher : higher + count] = 1
            logits[row, higher] = 1 + 2**-23

    def run() -> None:
        _run_topk_backend(
            "top_k_per_row_decode", logits, lengths, indices, top_k, width, next_n
        )

    run()
    _assert_exact_topk(logits, indices, lengths)
    torch.accelerator.synchronize()
    graph = torch.cuda.CUDAGraph()
    with torch.cuda.graph(graph):
        run()
    indices.fill_(-777)
    graph.replay()
    _assert_exact_topk(logits, indices, lengths)

    bounds = [
        0,
        511,
        512,
        513,
        16_384,
        16_385,
        131_072,
        131_073,
        262_144,
        262_145,
    ]
    row_bounds = torch.tensor(bounds, dtype=torch.int32, device="cuda")
    repeats = (rows + len(bounds) - 1) // len(bounds)
    lengths.flatten().copy_(row_bounds.repeat(repeats)[:rows])
    logits[:] = torch.arange(width, dtype=torch.float32, device="cuda")
    indices.fill_(-777)
    graph.replay()
    _assert_exact_topk(logits, indices, lengths)


@pytest.mark.skipif(not current_platform.is_rocm(), reason="This test requires ROCm")
@torch.inference_mode()
def test_aiter_c4a_prefill_topk_returns_sequence_local_indices() -> None:
    from vllm._aiter_ops import rocm_aiter_ops
    from vllm.v1.attention.ops.rocm_aiter_mla_sparse import (
        _launch_aiter_top_k_per_row_prefill,
    )

    row_starts = torch.tensor([1, 6, 20], dtype=torch.int32, device="cuda")
    row_ends = torch.tensor([3, 18, 23], dtype=torch.int32, device="cuda")
    logits = torch.arange(96, dtype=torch.float32, device="cuda").reshape(3, 32)
    top_k = 4
    indices = torch.empty((3, top_k), dtype=torch.int32, device="cuda")

    if not rocm_aiter_ops.is_indexer_top_k_supported(
        is_prefill=True,
        compress_ratio=4,
        num_rows=logits.shape[0],
    ):
        pytest.skip("AITER top-k is unavailable")
    assert (
        _launch_aiter_top_k_per_row_prefill(
            logits,
            row_starts,
            row_ends,
            indices,
            top_k,
        )
        is None
    )

    for row_idx, (row_start, row_end) in enumerate(
        zip(row_starts.tolist(), row_ends.tolist())
    ):
        num_valid = min(top_k, row_end - row_start)
        expected = logits[row_idx, row_start:row_end].topk(num_valid).indices
        actual = indices[row_idx, :num_valid]
        assert set(actual.tolist()) == set(expected.tolist())
        assert torch.all(indices[row_idx, num_valid:] == -1)


@pytest.mark.parametrize("num_speculative_tokens", [2, 3, 4])
@pytest.mark.skipif(not current_platform.is_rocm(), reason="This test requires ROCm")
@torch.inference_mode()
def test_aiter_c4a_decode_topk_uses_exact_mtp_lengths(
    num_speculative_tokens: int,
) -> None:
    from vllm._aiter_ops import rocm_aiter_ops

    query_tokens = num_speculative_tokens + 1
    offsets = torch.arange(query_tokens, dtype=torch.int32, device="cuda")
    final_uncompressed_lens = torch.tensor([43, 59], dtype=torch.int32, device="cuda")
    seq_lens = (final_uncompressed_lens[:, None] - query_tokens + offsets + 1) // 4
    seq_lens[0, 0] = 0

    num_rows = seq_lens.numel()
    logits = torch.arange(num_rows * 16, dtype=torch.float32, device="cuda").reshape(
        num_rows, 16
    )
    top_k = 4
    indices = torch.empty((num_rows, top_k), dtype=torch.int32, device="cuda")

    if not rocm_aiter_ops.is_indexer_top_k_supported(
        is_prefill=False,
        compress_ratio=4,
        num_rows=num_rows,
        max_valid_seq_len=seq_lens.max().item(),
    ):
        pytest.skip("AITER top-k is unavailable")
    assert (
        rocm_aiter_ops.indexer_top_k_decode(
            logits,
            1,
            seq_lens.reshape(-1),
            indices,
            top_k,
        )
        is None
    )

    for row_idx, row_end in enumerate(seq_lens.reshape(-1).tolist()):
        num_valid = min(top_k, row_end)
        expected = logits[row_idx, :row_end].topk(num_valid).indices
        actual = indices[row_idx, :num_valid]
        assert set(actual.tolist()) == set(expected.tolist())
        assert torch.all(indices[row_idx, num_valid:] == -1)


@requires_gfx950
@torch.inference_mode()
def test_aiter_decode_topk_matches_per_row_under_mtp() -> None:
    """The aiter backend expands (B, 1) seq_lens into per-row ends before
    launching, so next_n > 1 must agree with the in-tree kernel, including
    the rows that clamp to a zero-length prefix.
    """
    from vllm.model_executor.layers.indexer_topk import get_indexer_topk

    next_n = 4
    top_k = 512
    num_columns = 1024
    seq_lens = torch.tensor([[2], [513], [700]], dtype=torch.int32, device="cuda")
    num_rows = seq_lens.shape[0] * next_n

    torch.manual_seed(0)
    logits = torch.randn(num_rows, num_columns, dtype=torch.float32, device="cuda")

    outputs = {}
    for backend in ("aiter", "per_row"):
        indices = torch.empty((num_rows, top_k), dtype=torch.int32, device="cuda")
        get_indexer_topk(backend)(logits, seq_lens, next_n, indices, top_k, num_columns)
        outputs[backend] = indices

    row_ends = (
        (seq_lens - next_n + 1 + torch.arange(next_n, device="cuda"))
        .clamp_(min=0)
        .reshape(-1)
        .tolist()
    )
    assert row_ends[:next_n] == [0, 0, 1, 2]

    for row_idx, row_end in enumerate(row_ends):
        num_valid = min(top_k, row_end)
        aiter_row = outputs["aiter"][row_idx]
        assert set(aiter_row[:num_valid].tolist()) == set(
            outputs["per_row"][row_idx][:num_valid].tolist()
        )
        assert torch.all(aiter_row[num_valid:] == -1)


def _enable_aiter_topk(monkeypatch, enabled: bool = True) -> None:
    from vllm._aiter_ops import rocm_aiter_ops

    monkeypatch.setattr(
        rocm_aiter_ops, "is_indexer_top_k_enabled", lambda: enabled, raising=True
    )


def test_aiter_c4a_topk_kernel_selection(monkeypatch) -> None:
    from vllm._aiter_ops import rocm_aiter_ops

    _enable_aiter_topk(monkeypatch)
    is_supported = rocm_aiter_ops.is_indexer_top_k_supported

    eligible_cases = [
        dict(is_prefill=True, compress_ratio=4, num_rows=1),
        dict(is_prefill=False, compress_ratio=4, num_rows=1, max_valid_seq_len=65_536),
        dict(
            is_prefill=False,
            compress_ratio=4,
            num_rows=257,
            max_valid_seq_len=250_000,
        ),
        dict(
            is_prefill=False,
            compress_ratio=2,
            num_rows=384,
            max_valid_seq_len=1_048_576,
        ),
    ]
    for kwargs in eligible_cases:
        assert is_supported(**kwargs) is True

    ineligible_cases = [
        dict(is_prefill=True, compress_ratio=1, num_rows=1),
        dict(is_prefill=False, compress_ratio=4, num_rows=1, max_valid_seq_len=65_537),
        dict(
            is_prefill=False,
            compress_ratio=4,
            num_rows=256,
            max_valid_seq_len=125_000,
        ),
    ]
    for kwargs in ineligible_cases:
        assert is_supported(**kwargs) is False

    _enable_aiter_topk(monkeypatch, enabled=False)
    assert is_supported(is_prefill=True, compress_ratio=4, num_rows=1) is False


def test_dsv4_native_top_k_window() -> None:
    """The in-tree decode kernel's measured advantage band over AITER, which
    applies to the DSV4 indexer only."""
    from vllm._aiter_ops import rocm_aiter_ops

    prefers_native = rocm_aiter_ops.dsv4_indexer_prefers_native_top_k
    assert prefers_native(num_rows=128, num_columns=524_288, topk_tokens=512) is True
    assert prefers_native(num_rows=385, num_columns=524_288, topk_tokens=512) is False
    assert prefers_native(num_rows=128, num_columns=1_048_577, topk_tokens=512) is False
    assert prefers_native(num_rows=128, num_columns=524_288, topk_tokens=1024) is False


@pytest.mark.skipif(not current_platform.is_cuda(), reason="This test requires CUDA")
@pytest.mark.parametrize(
    "seq_len_range,test_id",
    [
        pytest.param((4000, 8000), "short_sequences", id="short"),
        pytest.param((8000, 32000), "medium_sequences", id="medium"),
        pytest.param((32000, 163840), "long_sequences", id="long"),
    ],
)
@pytest.mark.parametrize("clean_logits", [True, False])
@pytest.mark.parametrize("top_k", [2048])
@pytest.mark.parametrize("next_n", [1, 4])
@pytest.mark.parametrize("backend", WORKSPACE_TOPK_BACKENDS)
@torch.inference_mode()
def test_deepseek_workspace_topk(
    seq_len_range: tuple[int, int],
    test_id: str,
    clean_logits: bool,
    top_k: int,
    next_n: int,
    backend: str,
) -> None:
    """Test workspace top-k backends with varying sequence lengths and speculative
    decoding.
    Supports speculative decoding with next_n > 1.
    """
    set_random_seed(42 if test_id == "short_sequences" else 43)
    torch.set_default_device("cuda:0")

    batch_size = 4
    num_rows = batch_size * next_n

    seq_lens = torch.randint(
        seq_len_range[0],
        seq_len_range[1],
        (batch_size,),
        dtype=torch.int32,
        device="cuda",
    )
    seq_lens = (seq_lens + 3) & ~3  # align to 4 for TMA

    # Compute row boundaries for speculative decoding
    row_starts = torch.zeros(num_rows, dtype=torch.int32, device="cuda")
    row_indices = torch.arange(num_rows, device="cuda") // next_n
    next_n_offset = torch.arange(num_rows, device="cuda") % next_n
    row_ends = seq_lens[row_indices] - next_n + next_n_offset + 1

    logits = create_random_logits(
        row_starts, row_ends, torch.float32, 42, clean_logits, "random"
    )

    indices = torch.empty((num_rows, top_k), dtype=torch.int32, device="cuda")

    if next_n == 1:
        lengths = seq_lens
    else:
        offsets = torch.arange(next_n, device=logits.device, dtype=torch.int32)
        lengths = (seq_lens.unsqueeze(1) - next_n + 1 + offsets).flatten()

    max_seq_len = int(seq_lens.max().item())
    _run_topk_backend(backend, logits, lengths, indices, top_k, max_seq_len, next_n)

    validate_topk_against_reference(
        logits, indices, row_starts, row_ends, top_k, f"{backend} ({test_id})"
    )


@pytest.mark.skipif(not _has_device_capability(80), reason="Requires SM80+")
@pytest.mark.parametrize("rows", [65, 128])
@pytest.mark.parametrize("top_k", [512, 1024, 2048])
@pytest.mark.parametrize(
    "distribution", ["random", "10LSBits", "ascending", "constant", "sampled_peaks"]
)
@torch.inference_mode()
def test_persistent_topk_sampled_graph(
    rows: int, top_k: int, distribution: str
) -> None:
    """Sampling and both fallbacks preserve exact values across graph replays."""
    if torch.cuda.get_device_properties(0).shared_memory_per_block_optin < 144 * 1024:
        pytest.skip("Sampled top-k requires at least 144 KiB of shared memory")
    set_random_seed(42)
    min_sampled_length = 98304 if top_k == 512 else 65536
    width = min_sampled_length + 3
    lengths = torch.full((rows,), width, dtype=torch.int32, device="cuda")
    logits = create_random_logits(
        torch.zeros_like(lengths),
        lengths,
        torch.float32,
        42,
        False,
        "10LSBits" if distribution == "10LSBits" else "random",
    )
    if distribution == "ascending":
        logits.copy_(torch.arange(width, device="cuda", dtype=torch.float32))
    elif distribution == "constant":
        logits.fill_(1.0)
    elif distribution == "sampled_peaks":
        # A biased sample leaves fewer than k survivors and must fall back.
        logits.zero_()
        sample = torch.arange(4096, device="cuda")
        positions = (sample // 32) * (width // 128) + sample % 32
        logits[:, positions] = 1 + sample.float() / 4096
    original = logits.clone()
    indices = torch.empty((rows, top_k), dtype=torch.int32, device="cuda")
    workspace = torch.empty(RADIX_TOPK_WORKSPACE_SIZE, dtype=torch.uint8, device="cuda")

    def run() -> None:
        torch.ops._C.persistent_topk(
            logits,
            lengths.view(-1, 4 if rows == 128 else 1),
            indices,
            workspace,
            top_k,
            width,
        )

    run()
    graph = torch.cuda.CUDAGraph()
    with torch.cuda.graph(graph):
        run()
    positions = torch.arange(width, device="cuda")
    slots = torch.arange(top_k, device="cuda")
    for step in range(3):
        bounds = torch.tensor(
            [
                -1,
                0,
                1,
                top_k - 1,
                top_k,
                32769,
                min_sampled_length - 1,
                min_sampled_length,
                width,
                width + 17,
            ],
            dtype=torch.int32,
            device="cuda",
        )
        lengths.copy_(
            bounds[(torch.arange(rows, device="cuda") + step) % bounds.numel()]
        )
        logits.copy_(original if step != 1 else original.flip(1))
        logits.masked_fill_(positions[None] >= lengths[:, None], float("nan"))
        indices.fill_(-2)
        graph.replay()
        valid = slots[None] < lengths.clamp(max=top_k)[:, None]
        assert torch.all(indices[~valid] == -1)
        assert torch.all(((indices >= 0) & (indices < lengths[:, None])) == valid)
        ordered_indices = indices.sort(dim=1).values
        assert torch.all(
            (ordered_indices[:, 1:] != ordered_indices[:, :-1])
            | (ordered_indices[:, 1:] == -1)
        )
        selected = logits.gather(1, indices.clamp_min(0).long())
        selected.masked_fill_(~valid, -float("inf"))
        expected = logits.masked_fill(
            positions[None] >= lengths[:, None], -float("inf")
        )
        torch.testing.assert_close(
            selected.sort(dim=1, descending=True).values,
            expected.topk(top_k, dim=1).values,
            atol=0,
            rtol=0,
        )


def run_large_context_topk_test(
    batch_size: int,
    seq_lens: list[int],
    top_k: int,
    data_type: str = "random",
    seed: int = 42,
    backend: str = "cooperative_topk",
) -> None:
    """Helper to run a top-k backend test with given parameters.

    Args:
        batch_size: Number of rows/sequences
        seq_lens: List of sequence lengths (one per row)
        top_k: Number of top elements to select
        data_type: Type of test data to generate
        seed: Random seed for reproducibility
        backend: Top-k backend to test

    """
    torch.set_default_device("cuda:0")
    set_random_seed(seed)

    # Create test data
    num_rows = batch_size
    max_len = max(seq_lens)
    lengths = torch.tensor(seq_lens, dtype=torch.int32, device="cuda")

    if data_type == "random":
        logits = torch.randn(num_rows, max_len, dtype=torch.float32, device="cuda")
    elif data_type == "sorted_asc":
        # Each row gets its own ascending sequence based on its length
        logits = torch.empty(num_rows, max_len, dtype=torch.float32, device="cuda")
        for i, length in enumerate(seq_lens):
            logits[i, :length] = torch.arange(
                length, dtype=torch.float32, device="cuda"
            )
            if length < max_len:
                logits[i, length:] = float("-inf")
    elif data_type == "sorted_desc":
        # Each row gets its own descending sequence based on its length
        logits = torch.empty(num_rows, max_len, dtype=torch.float32, device="cuda")
        for i, length in enumerate(seq_lens):
            logits[i, :length] = torch.arange(
                length, 0, -1, dtype=torch.float32, device="cuda"
            )
            if length < max_len:
                logits[i, length:] = float("-inf")
    elif data_type == "all_same":
        logits = torch.ones(num_rows, max_len, dtype=torch.float32, device="cuda")
        for i, length in enumerate(seq_lens):
            if length < max_len:
                logits[i, length:] = float("-inf")
    elif data_type == "many_ties":
        # Only 10 unique values, many duplicates
        logits = torch.randint(0, 10, (num_rows, max_len), device="cuda").float() / 10.0
        for i, length in enumerate(seq_lens):
            if length < max_len:
                logits[i, length:] = float("-inf")
    elif data_type == "small_differences":
        # Very small differences to test float precision
        base = torch.randn(num_rows, max_len, dtype=torch.float32, device="cuda")
        noise = (
            torch.randn(num_rows, max_len, dtype=torch.float32, device="cuda") * 1e-6
        )
        logits = base + noise
        for i, length in enumerate(seq_lens):
            if length < max_len:
                logits[i, length:] = float("-inf")
    else:
        raise ValueError(f"Unknown data_type: {data_type}")

    # Create output tensor
    indices = torch.empty((num_rows, top_k), dtype=torch.int32, device="cuda")

    max_seq_len = max(seq_lens)
    _run_topk_backend(backend, logits, lengths, indices, top_k, max_seq_len)

    torch.accelerator.synchronize()

    torch_indices = torch.empty((num_rows, top_k), dtype=torch.int32, device="cuda")
    for i in range(num_rows):
        length = seq_lens[i]
        k_i = min(top_k, length)
        if k_i > 0:
            idx = logits[i, :length].topk(k_i, dim=-1)[1]
            torch_indices[i, :k_i] = idx
            if k_i < top_k:
                torch_indices[i, k_i:] = -1
        else:
            torch_indices[i, :] = -1

    # Compare results
    for i in range(num_rows):
        length = seq_lens[i]
        k_i = min(top_k, length)

        if k_i == 0:
            continue

        cuda_row = indices[i, :k_i].cpu()
        torch_row = torch_indices[i, :k_i].cpu()

        # Filter out -1 padding values from cuda_row
        valid_mask = cuda_row >= 0
        cuda_row = cuda_row[valid_mask]

        # Compare sets (order may differ for ties)
        cuda_set = set(cuda_row.tolist())
        torch_set = set(torch_row.tolist())

        if cuda_set == torch_set:
            continue

        # If sets differ, check if it's due to equal values (ties)
        cuda_vals = logits[i, cuda_row].cpu()
        torch_vals = logits[i, torch_row].cpu()

        # Check that min CUDA value >= max of values NOT in top-k
        if k_i < length:
            non_topk_indices = torch.tensor(
                list(set(range(length)) - cuda_set), dtype=torch.int32
            )
            if len(non_topk_indices) > 0:
                non_topk_vals = logits[i, non_topk_indices].cpu()
                min_cuda_val = cuda_vals.min()
                max_non_topk = non_topk_vals.max()

                # Allow small tolerance for floating point errors
                assert min_cuda_val >= max_non_topk - 1e-4, (
                    f"Row {i}: CUDA top-k contains values smaller than non-top-k. "
                    f"Min CUDA: {min_cuda_val}, Max non-top-k: {max_non_topk}, "
                    f"Length: {length}, k: {k_i}, CUDA indices: {sorted(cuda_set)[:10]}..., "  # noqa: E501
                    f"Expected indices: {sorted(torch_set)[:10]}..."
                )

        # For ties, verify the values are close
        assert torch.allclose(
            cuda_vals.sort(descending=True)[0],
            torch_vals.sort(descending=True)[0],
            rtol=1e-4,
            atol=1e-4,
        ), f"""Row {i}: Top-k values don't match.
            CUDA: {cuda_vals.sort(descending=True)[0][:10]},
            Torch: {torch_vals.sort(descending=True)[0][:10]}"""


@pytest.mark.skipif(not current_platform.is_cuda(), reason="This test requires CUDA")
@pytest.mark.parametrize("top_k", [512, 1024, 2048])
def test_cooperative_topk_cs2(top_k: int) -> None:
    """The 64-row dispatch uses the two-CTA cooperative kernel."""
    run_large_context_topk_test(
        batch_size=64,
        seq_lens=[65536] * 64,
        top_k=top_k,
        backend="cooperative_topk",
    )


@pytest.mark.skipif(not current_platform.is_cuda(), reason="This test requires CUDA")
@pytest.mark.parametrize(
    "test_config",
    [
        # ==================== CATEGORY: Sequence Length Edge Cases ====================
        pytest.param(
            {"seq_lens": [1, 10, 100, 2048], "top_k": 2048, "data_type": "random"},
            id="seq_len_edge_very_small_to_medium",
        ),
        pytest.param(
            {
                "seq_lens": [2049, 2100, 2500, 3000],
                "top_k": 2048,
                "data_type": "random",
            },
            id="seq_len_edge_above_k",
        ),
        pytest.param(
            {"seq_lens": [8000, 16384, 20000], "top_k": 2048, "data_type": "random"},
            id="algo_transition_filtered_radix",
        ),
        # ==================== CATEGORY: Data Distributions ====================
        pytest.param(
            {"seq_lens": [5000, 10000], "top_k": 2048, "data_type": "sorted_asc"},
            id="data_sorted_ascending",
        ),
        pytest.param(
            {"seq_lens": [5000, 10000], "top_k": 2048, "data_type": "sorted_desc"},
            id="data_sorted_descending",
        ),
        pytest.param(
            {"seq_lens": [5000, 10000], "top_k": 2048, "data_type": "all_same"},
            id="data_all_same",
        ),
        pytest.param(
            {"seq_lens": [5000, 10000], "top_k": 2048, "data_type": "many_ties"},
            id="data_many_ties",
        ),
        pytest.param(
            {
                "seq_lens": [5000, 10000],
                "top_k": 2048,
                "data_type": "small_differences",
            },
            id="data_float_precision",
        ),
        # ==================== CATEGORY: Alignment / Vectorization ====================
        pytest.param(
            {
                "seq_lens": [2055, 2056, 2057, 2063],
                "top_k": 2048,
                "data_type": "random",
            },
            id="align_vec_boundaries_low",
        ),
        pytest.param(
            {
                "seq_lens": [4095, 4096, 4097, 4102],
                "top_k": 2048,
                "data_type": "random",
            },
            id="align_4k_boundary",
        ),
        pytest.param(
            {
                "seq_lens": [8191, 8192, 8193, 8198],
                "top_k": 2048,
                "data_type": "random",
            },
            id="align_8k_boundary",
        ),
        pytest.param(
            {
                "seq_lens": [16383, 16384, 16385, 16390],
                "top_k": 2048,
                "data_type": "random",
            },
            id="align_16k_boundary",
        ),
    ],
)
@pytest.mark.parametrize("backend", WORKSPACE_TOPK_BACKENDS)
@torch.inference_mode()
def test_workspace_topk_correctness(test_config: dict, backend: str) -> None:
    """Comprehensive correctness tests covering:
    - Sequence length edge cases (trivial, boundary, varied)
    - Very small sequences (< 100 elements)
    - Mixed sequence lengths in same batch
    - Data distributions (sorted, ties, precision)
    - Memory alignment / vectorization boundaries
    """
    run_large_context_topk_test(
        batch_size=len(test_config["seq_lens"]),
        seq_lens=test_config["seq_lens"],
        top_k=test_config["top_k"],
        data_type=test_config.get("data_type", "random"),
        backend=backend,
    )


@pytest.mark.skipif(not current_platform.is_cuda(), reason="This test requires CUDA")
@pytest.mark.parametrize(
    "test_config",
    [
        # ==================== CATEGORY: Batch Size Scalability ====================
        pytest.param(
            {"batch_size": 1, "seq_len": 5000, "top_k": 2048},
            id="batch_1",
        ),
        pytest.param(
            {"batch_size": 4, "seq_len": 5000, "top_k": 2048},
            id="batch_4",
        ),
        pytest.param(
            {"batch_size": 32, "seq_len": 5000, "top_k": 2048},
            id="batch_32",
        ),
        pytest.param(
            {"batch_size": 256, "seq_len": 5000, "top_k": 2048},
            id="batch_256",
        ),
        # ==================== CATEGORY: Single-CTA vs Multi-CTA ====================
        pytest.param(
            {"batch_size": 2, "seq_len": 4096, "top_k": 2048},
            id="single_cta_4k",
        ),
        pytest.param(
            {"batch_size": 2, "seq_len": 8192, "top_k": 2048},
            id="single_cta_8k",
        ),
        pytest.param(
            {"batch_size": 2, "seq_len": 163840, "top_k": 2048},
            id="multi_cta_163840_dsv3_max",
        ),
        # ==================== CATEGORY: Extreme Cases ====================
        pytest.param(
            {"batch_size": 512, "seq_len": 5000, "top_k": 2048},
            id="extreme_large_batch",
        ),
        pytest.param(
            {"batch_size": 2, "seq_len": 163840, "top_k": 2048},
            id="extreme_dsv3_max_context",
        ),
    ],
)
@pytest.mark.parametrize("backend", WORKSPACE_TOPK_BACKENDS)
@torch.inference_mode()
def test_workspace_topk_algorithm_paths(test_config: dict, backend: str) -> None:
    """Test different algorithm execution paths (capped at 163840 for DeepSeek V3.2):
    - Batch size scalability (1, 4, 32, 256)
    - Single-CTA vs Multi-CTA execution
    - Extreme configurations (large batch, max context length)
    """
    run_large_context_topk_test(
        batch_size=test_config["batch_size"],
        seq_lens=[test_config["seq_len"]] * test_config["batch_size"],
        top_k=test_config["top_k"],
        backend=backend,
    )


@pytest.mark.skipif(not current_platform.is_cuda(), reason="This test requires CUDA")
@pytest.mark.parametrize("backend", WORKSPACE_TOPK_BACKENDS)
@torch.inference_mode()
def test_workspace_topk_stress(backend: str) -> None:
    """Stress test with random configurations to catch edge cases.
    Capped at 163840 (DeepSeek V3.2 max context) for realistic testing.
    """
    torch.set_default_device("cuda:0")
    top_k = 2048

    for seed in range(3):
        set_random_seed(seed)

        # Random batch size (limited for speed)
        batch_size = torch.randint(1, 32, (1,)).item()

        # Random sequence lengths capped at DeepSeek V3.2 max context
        seq_lens_tensor = torch.randint(100, 163840, (batch_size,))
        if backend == "cooperative_topk":
            seq_lens = ((seq_lens_tensor + 3) & ~3).tolist()
        else:
            seq_lens = seq_lens_tensor.tolist()

        run_large_context_topk_test(
            batch_size=batch_size,
            seq_lens=seq_lens,
            top_k=top_k,
            seed=seed,
            backend=backend,
        )


@pytest.mark.skipif(not current_platform.is_cuda(), reason="This test requires CUDA")
@pytest.mark.parametrize("backend", TOPK_BACKENDS)
@pytest.mark.parametrize("top_k", [512, 1024, 2048])
@torch.inference_mode()
def test_deepseek_topk_backends_no_error_and_reference(
    backend: str,
    top_k: int,
) -> None:
    """Exercise every production top-k backend on the same inputs."""
    run_large_context_topk_test(
        batch_size=4,
        seq_lens=[2049, 4097, 8191, 12000],
        top_k=top_k,
        data_type="random",
        seed=123,
        backend=backend,
    )


@pytest.mark.skipif(not _has_device_capability(90), reason="This test requires SM90+")
@torch.inference_mode()
def test_cooperative_topk_512_tie_workspace_is_per_row() -> None:
    """Regression test for TopK=512 tie workspace row overlap."""
    torch.set_default_device("cuda:0")

    top_k = 512
    num_rows = 2
    stride = 65536
    lengths = torch.tensor([40960, 65536], dtype=torch.int32, device="cuda")
    logits = torch.full(
        (num_rows, stride), float("-inf"), dtype=torch.float32, device="cuda"
    )

    # Row 0 must never select these low indices: many better row-0 ties exist.
    logits[0, :2048] = -10.0
    logits[0, 2048 : lengths[0]] = 1.0
    # Row 1 has higher exact tie scores. With the old row * TopK tie_ws stride,
    # these row-1 ties could overwrite row 0's TopK=512 refinement workspace.
    logits[1, : lengths[1]] = 2.0

    indices = torch.empty((num_rows, top_k), dtype=torch.int32, device="cuda")
    workspace = torch.empty(RADIX_TOPK_WORKSPACE_SIZE, dtype=torch.uint8, device="cuda")
    torch.ops._C.cooperative_topk(logits, lengths, indices, workspace, top_k, stride)
    torch.accelerator.synchronize()

    row0 = indices[0].cpu()
    assert torch.all(row0 >= 2048), (
        "cooperative_topk TopK=512 selected row-0 low-score indices, likely "
        "from overlapping tie_ws rows"
    )


@pytest.mark.skipif(not current_platform.is_cuda(), reason="This test requires CUDA")
@pytest.mark.parametrize(
    "test_config",
    [
        # Mixed batch: rows spanning all four paths (trivial, decode, medium, large)
        pytest.param(
            {
                "seq_lens": [2000, 6000, 30000, 80000],
                "data_type": "random",
            },
            id="mixed_all_paths",
        ),
        # All decode/medium rows (typical decode scenario)
        pytest.param(
            {
                "seq_lens": [2048, 4096, 8192, 16000],
                "data_type": "random",
            },
            id="all_decode_medium",
        ),
        # All large rows
        pytest.param(
            {
                "seq_lens": [70000, 100000, 163840],
                "data_type": "random",
            },
            id="all_large",
        ),
        # Boundary around LARGE_THRESHOLD (32K)
        pytest.param(
            {
                "seq_lens": [32767, 32768, 32769, 32772],
                "data_type": "random",
            },
            id="large_threshold_boundary",
        ),
        # Single row medium
        pytest.param(
            {
                "seq_lens": [5000],
                "data_type": "random",
            },
            id="single_row_medium",
        ),
        # Single row large
        pytest.param(
            {
                "seq_lens": [100000],
                "top_k": 2048,
                "data_type": "random",
            },
            id="single_row_large",
        ),
        # Trivial rows mixed with medium and large
        pytest.param(
            {
                "seq_lens": [100, 2048, 10000, 80000],
                "data_type": "random",
            },
            id="trivial_medium_large_mix",
        ),
    ],
)
@pytest.mark.parametrize("top_k", [512, 2048])
@pytest.mark.parametrize("backend", WORKSPACE_TOPK_BACKENDS)
@torch.inference_mode()
def test_workspace_topk(test_config: dict, top_k: int, backend: str) -> None:
    """Tests specific to workspace top-k backends:
    - Mixed medium/large rows in the same batch (dynamic per-row dispatch)
    - Boundary around LARGE_THRESHOLD (32K)
    - Trivial + medium + large rows in a single batch
    """
    run_large_context_topk_test(
        batch_size=len(test_config["seq_lens"]),
        seq_lens=test_config["seq_lens"],
        top_k=top_k,
        data_type=test_config.get("data_type", "random"),
        backend=backend,
    )


@pytest.mark.skipif(not current_platform.is_cuda(), reason="This test requires CUDA")
@torch.inference_mode()
def test_persistent_topk_reused_group_after_short_row() -> None:
    """A short row must not advance a group's radix histogram ring."""
    torch.set_default_device("cuda:0")
    set_random_seed(0)

    top_k = 2048
    long_seq_len = 32769
    radix = 256
    fixed_smem = ((radix + radix + 5) * 4 + 15) & ~15
    props = torch.cuda.get_device_properties(0)
    max_smem = props.shared_memory_per_block_optin
    if max_smem <= props.shared_memory_per_multiprocessor // 2:
        pytest.skip("Cannot force one persistent_topk CTA per SM")

    max_chunk = ((max_smem - fixed_smem) // 4 // 4) * 4
    ctas_per_group = max(
        (props.multi_processor_count - 1 + 9) // 10,
        (long_seq_len + max_chunk - 1) // max_chunk,
    )
    if ctas_per_group >= props.multi_processor_count:
        pytest.skip("Not enough SMs to construct a reused CTA group")

    stride = ctas_per_group * max_chunk
    num_groups = max(1, (props.multi_processor_count - 1) // ctas_per_group)
    num_rows = 3 * num_groups
    lengths = torch.full((num_rows,), top_k, dtype=torch.int32, device="cuda")
    target_row = 2 * num_groups
    lengths[0] = long_seq_len
    lengths[num_groups] = long_seq_len - 1
    lengths[target_row] = long_seq_len

    logits = torch.randn(num_rows, stride, dtype=torch.float32, device="cuda")
    indices = torch.empty((num_rows, top_k), dtype=torch.int32, device="cuda")
    _run_topk_backend("persistent_topk", logits, lengths, indices, top_k, stride)
    torch.accelerator.synchronize()

    expected = logits[target_row, :long_seq_len].topk(top_k).indices
    assert set(indices[target_row].cpu().tolist()) == set(expected.cpu().tolist())


@pytest.mark.skipif(not current_platform.is_cuda(), reason="This test requires CUDA")
@pytest.mark.parametrize("top_k", [512, 2048])
@pytest.mark.parametrize("backend", WORKSPACE_TOPK_BACKENDS)
@torch.inference_mode()
def test_workspace_topk_padded_stride(top_k: int, backend: str) -> None:
    """Test workspace top-k backends with padded logits (large stride, small seq_len)
    to simulate the e2e CUDAGraph scenario where fp8_paged_mqa_logits
    returns [B, max_model_len] with max_model_len=163840.
    """
    set_random_seed(42)
    torch.set_default_device("cuda:0")

    batch_size = 4
    padded_stride = 163840  # DeepSeek-V3.2 max_model_len
    actual_seq_lens = [3000, 5000, 8000, 12000]

    # Create padded logits tensor (like fp8_paged_mqa_logits output)
    logits = torch.full(
        (batch_size, padded_stride),
        float("-inf"),
        dtype=torch.float32,
        device="cuda",
    )
    for i, sl in enumerate(actual_seq_lens):
        logits[i, :sl] = torch.randn(sl, dtype=torch.float32, device="cuda")

    lengths = torch.tensor(actual_seq_lens, dtype=torch.int32, device="cuda")
    indices = torch.empty((batch_size, top_k), dtype=torch.int32, device="cuda")
    _run_topk_backend(backend, logits, lengths, indices, top_k, max(actual_seq_lens))
    torch.accelerator.synchronize()

    # Validate against torch.topk
    for i in range(batch_size):
        sl = actual_seq_lens[i]
        k_i = min(top_k, sl)
        expected = logits[i, :sl].topk(k_i, dim=-1)[1].cpu()
        actual = indices[i, :k_i].cpu()

        expected_set = set(expected.tolist())
        actual_set = set(actual.tolist())

        if expected_set != actual_set:
            # Allow ties
            expected_vals = logits[i, expected].cpu().sort(descending=True)[0]
            actual_vals = logits[i, actual].cpu().sort(descending=True)[0]
            assert torch.allclose(expected_vals, actual_vals, rtol=1e-4, atol=1e-4), (
                f"Row {i}: {backend} with padded stride doesn't match. "
                f"seq_len={sl}, stride={padded_stride}"
            )


@requires_sm100
@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16])
@pytest.mark.parametrize("batch_size", [1, 8, 64])
@pytest.mark.parametrize("vocab_size", [4096, 131072])
@pytest.mark.parametrize("top_k", [512, 2048])
@torch.inference_mode()
def test_deep_select_topk(
    dtype: torch.dtype,
    batch_size: int,
    vocab_size: int,
    top_k: int,
) -> None:
    """DeepSelect wrapper vs torch.topk with variable per-row ends."""
    from vllm.model_executor.layers import indexer_topk

    set_random_seed(0)
    torch.set_default_device("cuda:0")

    row_ends = torch.randint(
        1, vocab_size + 1, (batch_size,), dtype=torch.int32, device="cuda"
    )
    # Force one short row to exercise the -1 out-of-bounds fill.
    row_ends[0] = min(top_k - 1, vocab_size)
    logits = torch.randn(batch_size, vocab_size, dtype=dtype, device="cuda")
    # Values beyond each row's end must never be read by the kernel: fill
    # them with NaN (production tails are uninitialized; DeepSelect traps on
    # NaN with abort_when_nan_found=True, so a wrong read fails loudly).
    col_idx = torch.arange(vocab_size, device="cuda")
    logits[col_idx[None, :] >= row_ends[:, None]] = float("nan")

    indices = indexer_topk.deep_select_topk(logits, top_k, end=row_ends)
    torch.accelerator.synchronize()

    assert indices.shape == (batch_size, top_k)
    row_starts = torch.zeros(batch_size, dtype=torch.int32, device="cuda")
    torch_indices = torch.empty((batch_size, top_k), dtype=torch.int32, device="cuda")
    for i in range(batch_size):
        row_end = int(row_ends[i])
        k_i = min(top_k, row_end)
        idx = logits[i, :row_end].topk(k_i, dim=-1)[1]
        torch_indices[i, :k_i] = idx
        assert torch.all(indices[i, k_i:] == -1), (
            f"Row {i}: expected -1 fill after {k_i} valid indices"
        )
        # All selected indices must lie inside [0, row_end).
        assert torch.all(indices[i, :k_i] < row_end)

    assert compare_top_k_results(
        logits, indices, torch_indices, row_starts, row_ends, top_k
    ), "DeepSelect topk results don't match torch.topk"


@requires_sm100
@torch.inference_mode()
def test_deep_select_topk_preallocated_output() -> None:
    """DeepSelect writes into a preallocated aligned buffer (no `end`)."""
    from vllm.model_executor.layers import indexer_topk

    set_random_seed(0)
    torch.set_default_device("cuda:0")

    batch_size = 8
    vocab_size = 131072
    top_k = 2048
    logits = torch.randn(batch_size, vocab_size, dtype=torch.float32, device="cuda")

    # Mimic vLLM's topk_indices_buffer: slice of a wider contiguous buffer,
    # whose stride(0) is 32B-aligned.
    buffer = torch.full((batch_size, 4096), -2, dtype=torch.int32, device="cuda")
    output_idx = buffer[:, :top_k]
    result = indexer_topk.deep_select_topk(logits, top_k, output_idx=output_idx)
    torch.accelerator.synchronize()

    assert result.data_ptr() == output_idx.data_ptr()
    assert torch.all(buffer[:, top_k:] == -2), "kernel wrote out of bounds"

    row_ends = torch.full((batch_size,), vocab_size, dtype=torch.int32, device="cuda")
    row_starts = torch.zeros(batch_size, dtype=torch.int32, device="cuda")
    torch_indices = torch.empty((batch_size, top_k), dtype=torch.int32, device="cuda")
    for i in range(batch_size):
        torch_indices[i] = logits[i].topk(top_k, dim=-1)[1]

    assert compare_top_k_results(
        logits, output_idx, torch_indices, row_starts, row_ends, top_k
    ), "DeepSelect topk (preallocated output) results don't match torch.topk"


def _has_flashinfer_topk() -> bool:
    if not current_platform.is_cuda():
        return False
    # find_spec only: importing flashinfer initializes CUDA at import time.
    import importlib.util

    return importlib.util.find_spec("flashinfer") is not None


def _has_cooperative_topk() -> bool:
    return _has_device_capability(
        90
    ) and not current_platform.is_device_capability_family(120)


SPARSE_INDEXER_EXPLICIT_BACKENDS = [
    pytest.param(
        "per_row",
        marks=pytest.mark.skipif(
            not current_platform.is_cuda(), reason="requires CUDA"
        ),
    ),
    "torch",
    pytest.param(
        "persistent",
        marks=pytest.mark.skipif(
            not current_platform.is_cuda(), reason="requires CUDA"
        ),
    ),
    pytest.param(
        "cooperative",
        marks=pytest.mark.skipif(
            not _has_cooperative_topk(),
            reason="cooperative_topk requires SM90+ (non-SM12x)",
        ),
    ),
    pytest.param(
        "flashinfer",
        marks=pytest.mark.skipif(
            not _has_flashinfer_topk(), reason="requires flashinfer top-k"
        ),
    ),
    pytest.param("deep_select", marks=requires_sm100),
]


@pytest.mark.parametrize("next_n", [1, 4])
@pytest.mark.parametrize("backend", SPARSE_INDEXER_EXPLICIT_BACKENDS)
@torch.inference_mode()
def test_sparse_indexer_decode_topk_explicit_backends(
    backend: str, next_n: int, workspace_init
) -> None:
    """Every explicit sparse-indexer decode top-k backend must restrict each
    row to its seq_len (dirty data past the end must never be selected) and
    match torch.topk on the valid region."""
    from vllm.config import VllmConfig, set_current_vllm_config
    from vllm.model_executor.layers.indexer_topk import SparseIndexerTopk

    set_random_seed(0)
    torch.set_default_device("cuda:0")

    top_k = 2048
    batch_size = 4
    num_rows = batch_size * next_n

    # Per-row effective lens in the (B, next_n) "native spec decode" form;
    # cooperative/persistent_topk require lengths.numel() == num_rows.
    row_ends = torch.randint(4000, 16000, (num_rows,), dtype=torch.int32, device="cuda")
    seq_lens = row_ends.reshape(batch_size, next_n)
    max_seq_len = int(row_ends.max())
    # Pad the stride: multiple of 4 for cooperative_topk (TMA) and
    # 1024B-aligned (256 fp32) for DeepSelect.
    stride = (max_seq_len + 255) // 256 * 256

    row_starts = torch.zeros(num_rows, dtype=torch.int32, device="cuda")

    logits = torch.randn(num_rows, stride, dtype=torch.float32, device="cuda")
    col_idx = torch.arange(stride, device="cuda")
    logits[col_idx[None, :] >= row_ends[:, None]] = 1e30  # dirty tail

    indices = torch.full((num_rows, top_k), -2, dtype=torch.int32, device="cuda")
    cfg = VllmConfig(kernel_config={"sparse_indexer_topk_backend": backend})
    with set_current_vllm_config(cfg):
        SparseIndexerTopk()(logits, seq_lens, next_n, indices, top_k, max_seq_len)
    torch.accelerator.synchronize()

    # k_i == top_k for every row here (row_ends >= 3997 > top_k).
    assert torch.all(indices >= 0)
    torch_indices = torch.empty((num_rows, top_k), dtype=torch.int32, device="cuda")
    for i in range(num_rows):
        row_end = int(row_ends[i])
        torch_indices[i] = logits[i, :row_end].topk(top_k, dim=-1)[1]
        assert torch.all(indices[i] < row_end), (
            f"{backend}: row {i} selected indices past its seq_len"
        )

    assert compare_top_k_results(
        logits, indices, torch_indices, row_starts, row_ends, top_k
    ), f"{backend} results don't match torch.topk"


@pytest.mark.parametrize(
    "backend",
    [
        "torch",
        pytest.param(
            "per_row",
            marks=pytest.mark.skipif(
                not current_platform.is_cuda(), reason="requires CUDA"
            ),
        ),
        pytest.param(
            "flashinfer",
            marks=pytest.mark.skipif(
                not _has_flashinfer_topk(), reason="requires flashinfer top-k"
            ),
        ),
        pytest.param("deep_select", marks=requires_sm100),
    ],
)
@torch.inference_mode()
def test_sparse_indexer_decode_topk_short_seq_lens(
    backend: str, workspace_init
) -> None:
    """seq_len < next_n: derived per-row ends must clamp at 0 (the reference
    kernels do max(0, ...)); negative ends are OOB for the ragged kernels."""
    from vllm.config import VllmConfig, set_current_vllm_config
    from vllm.model_executor.layers.indexer_topk import SparseIndexerTopk

    set_random_seed(0)
    next_n = 4
    top_k = 512
    vocab = 8192  # 1024B-aligned rows for DeepSelect
    # 1D (B,) seq_lens: rows of a request share its len and each kernel
    # derives per-row ends (with clamping) itself.
    seq_lens = torch.tensor([2, 6], dtype=torch.int32, device="cuda")
    num_rows = 2 * next_n
    logits = torch.randn(num_rows, vocab, dtype=torch.float32, device="cuda")
    indices = torch.full((num_rows, top_k), -2, dtype=torch.int32, device="cuda")

    cfg = VllmConfig(kernel_config={"sparse_indexer_topk_backend": backend})
    with set_current_vllm_config(cfg):
        SparseIndexerTopk()(logits, seq_lens, next_n, indices, top_k, vocab)
    torch.accelerator.synchronize()

    # Derived ends: [2,6] - next_n + 1 + arange(next_n) -> [-1,0,1,2] and
    # [3,4,5,6], clamped to [0,0,1,2] and [3,4,5,6].
    for row, end in enumerate([0, 0, 1, 2, 3, 4, 5, 6]):
        k_i = min(top_k, end)
        assert torch.all(indices[row, k_i:] == -1), f"{backend}: row {row}"
        if k_i:
            ref = logits[row, :end].topk(k_i).indices.sort().values
            got = indices[row, :k_i].sort().values
            assert torch.equal(got, ref), f"{backend}: row {row}"


@pytest.mark.skipif(not current_platform.is_cuda(), reason="requires CUDA")
@torch.inference_mode()
def test_sparse_indexer_topk_backend_resolution() -> None:
    """Auto heuristic chain and fail-fast validation of explicit backends."""
    from vllm.config import VllmConfig, set_current_vllm_config
    from vllm.model_executor.layers.indexer_topk import SparseIndexerTopk

    # 1024B-aligned stride: 131072 * 4 bytes.
    logits = torch.randn(64, 131072, dtype=torch.float32, device="cuda")
    # 4-divisible (cooperative TMA ok) but not 1024B-aligned (DeepSelect no):
    # 131076 * 4 = 524304 bytes.
    unaligned_logits = torch.randn(128, 131076, dtype=torch.float32, device="cuda")[
        :, :131072
    ]

    def resolve(
        backend: str,
        t: torch.Tensor = logits,
        k: int = 2048,
        num_rows: int = 64,
    ) -> str:
        cfg = VllmConfig(kernel_config={"sparse_indexer_topk_backend": backend})
        with set_current_vllm_config(cfg):
            return SparseIndexerTopk().resolve_backend(t, k, num_rows)

    is_sm100 = current_platform.is_device_capability_family(100)
    has_coop = _has_cooperative_topk()
    has_fi = _has_flashinfer_topk()

    # "auto" is the chain cooperative -> persistent -> per_row, with the
    # cooperative/persistent switch at cooperative_topk's cluster-wave row
    # limit. It must never select the opt-in backends by itself.
    if has_coop:
        from vllm.model_executor.layers.indexer_topk import (
            AUTO_COOPERATIVE_MAX_ROWS,
        )

        # cooperative for every batch within its cluster-wave row limit...
        assert (
            resolve("auto", k=512, num_rows=AUTO_COOPERATIVE_MAX_ROWS) == "cooperative"
        )
        assert resolve("auto", k=512, num_rows=8) == "cooperative"
        # ...persistent past the row limit, even where DeepSelect would be
        # applicable.
        assert (
            resolve("auto", k=512, num_rows=AUTO_COOPERATIVE_MAX_ROWS + 1)
            == "persistent"
        )
        assert resolve("auto", k=512, num_rows=128) == "persistent"
        # 16B-aligned (cooperative TMA ok) but not DeepSelect-aligned: the
        # same row-limit logic applies.
        assert (
            resolve("auto", unaligned_logits, num_rows=AUTO_COOPERATIVE_MAX_ROWS)
            == "cooperative"
        )
    # ...and persistent past the row limit.
    assert resolve("auto", unaligned_logits, num_rows=128) == "persistent"
    # Unsupported topk (outside {512, 1024, 2048} for the workspace kernels)
    # -> per_row fallback.
    assert resolve("auto", k=5000) == "per_row"

    # Explicit values are returned as-is when constraints are met.
    assert resolve("per_row") == "per_row"
    assert resolve("torch") == "torch"
    assert resolve("persistent") == "persistent"
    if has_coop:
        assert resolve("cooperative") == "cooperative"
    if has_fi:
        assert resolve("flashinfer") == "flashinfer"
    if is_sm100:
        # Explicit deep_select has no row-count requirement.
        assert resolve("deep_select", num_rows=1) == "deep_select"
        with pytest.raises(RuntimeError, match="constraints"):
            resolve("deep_select", unaligned_logits)

    # Fail-fast on unmet constraints.
    with pytest.raises(RuntimeError, match="num_rows must be <= 64"):
        resolve("cooperative", num_rows=128)
    with pytest.raises(RuntimeError, match="topk_tokens must be in"):
        resolve("persistent", k=3000)
    with pytest.raises(RuntimeError, match="topk_tokens must be in"):
        resolve("cooperative", k=3000)
    if not is_sm100:
        with pytest.raises(RuntimeError, match="SM100a/SM103a"):
            resolve("deep_select")


def _patch_aiter_topk(monkeypatch, decode_kernel, enabled: bool = True) -> None:
    from vllm._aiter_ops import rocm_aiter_ops

    _enable_aiter_topk(monkeypatch, enabled=enabled)
    monkeypatch.setattr(rocm_aiter_ops, "is_enabled", lambda: enabled, raising=True)
    monkeypatch.setattr(rocm_aiter_ops, "indexer_top_k_decode", decode_kernel)


def _make_topk(backend: str):
    from vllm.config import VllmConfig, set_current_vllm_config
    from vllm.model_executor.layers.indexer_topk import SparseIndexerTopk

    cfg = VllmConfig(kernel_config={"sparse_indexer_topk_backend": backend})
    with set_current_vllm_config(cfg):
        return SparseIndexerTopk()


@pytest.mark.skipif(not current_platform.is_cuda(), reason="requires CUDA")
def test_sparse_indexer_topk_aiter_requires_rocm() -> None:
    """Explicitly requesting the ROCm-only backend must fail fast elsewhere."""
    logits = torch.randn(8, 4096, dtype=torch.float32, device="cuda")
    with pytest.raises(RuntimeError, match="ROCm"):
        _make_topk("aiter").resolve_backend(logits, 2048, 8)


@pytest.mark.skipif(not current_platform.is_rocm(), reason="requires ROCm")
def test_sparse_indexer_topk_backend_resolution_aiter(monkeypatch) -> None:
    """Aiter is opt-in only: "auto" never selects it (the shape-dependent perf
    gate lives at the call site), and an explicit request validates hard
    capability only."""

    def decode_kernel(*args, **kwargs) -> None:
        pass

    _patch_aiter_topk(monkeypatch, decode_kernel)
    logits = torch.randn(8, 4096, dtype=torch.float32, device="cuda")

    assert _make_topk("auto").resolve_backend(logits, 2048, 8) == "per_row"
    assert _make_topk("aiter").resolve_backend(logits, 2048, 8) == "aiter"
    # Shapes the call-site perf gate would decline still resolve.
    assert _make_topk("aiter").resolve_backend(logits, 512, 8) == "aiter"
    with pytest.raises(RuntimeError, match="topk_tokens must be in"):
        _make_topk("aiter").resolve_backend(logits, 3000, 8)

    _enable_aiter_topk(monkeypatch, enabled=False)
    with pytest.raises(RuntimeError, match="no tuned indexer top-k kernels"):
        _make_topk("aiter").resolve_backend(logits, 2048, 8)

    _patch_aiter_topk(monkeypatch, decode_kernel, enabled=False)
    assert _make_topk("auto").resolve_backend(logits, 2048, 8) == "per_row"
    with pytest.raises(RuntimeError, match="VLLM_ROCM_USE_AITER"):
        _make_topk("aiter").resolve_backend(logits, 2048, 8)


@pytest.mark.skipif(not current_platform.is_rocm(), reason="requires ROCm")
@torch.inference_mode()
def test_sparse_indexer_topk_aiter_applies_no_shape_gate(monkeypatch) -> None:
    """Shape-dependent perf gating belongs to the call site: once "aiter" is
    requested the dispatcher must run it for every supported k, including the
    k=512 band the in-tree kernel wins on DSV4's compressed-KV logits."""
    from vllm.model_executor.layers import indexer_topk

    calls: list[tuple] = []
    _patch_aiter_topk(monkeypatch, lambda *args: calls.append(args))
    monkeypatch.setattr(
        indexer_topk.ops,
        "top_k_per_row_decode",
        lambda *args: pytest.fail("an explicit aiter request must not fall back"),
    )

    num_rows = 8
    logits = torch.randn(num_rows, 4096, dtype=torch.float32, device="cuda")
    seq_lens = torch.full((num_rows, 1), 4096, dtype=torch.int32, device="cuda")
    indices = torch.empty((num_rows, 512), dtype=torch.int32, device="cuda")
    op = _make_topk("aiter")

    op(logits, seq_lens, 1, indices, 512, 4096)
    # Past the compressed-context boundary the call-site gate declines on.
    op(logits, seq_lens, 1, indices, 512, 4 * (64 * 1024 + 1))
    assert len(calls) == 2


@pytest.mark.skipif(not current_platform.is_rocm(), reason="requires ROCm")
@torch.inference_mode()
def test_sparse_indexer_topk_auto_stays_on_per_row(monkeypatch) -> None:
    """The "auto" chain must not reach AITER on its own: the decode top-k runs
    through top_k_per_row_decode unless the caller asks for the aiter backend."""
    from vllm.model_executor.layers import indexer_topk

    def decode_kernel(*args, **kwargs) -> None:
        pytest.fail("AITER is opt-in only")

    _patch_aiter_topk(monkeypatch, decode_kernel)
    calls: list[tuple] = []
    monkeypatch.setattr(
        indexer_topk.ops, "top_k_per_row_decode", lambda *args: calls.append(args)
    )

    num_rows = 8
    logits = torch.randn(num_rows, 4096, dtype=torch.float32, device="cuda")
    seq_lens = torch.full((num_rows, 1), 4096, dtype=torch.int32, device="cuda")
    indices = torch.empty((num_rows, 1024), dtype=torch.int32, device="cuda")

    _make_topk("auto")(logits, seq_lens, 1, indices, 1024, 4096)
    assert len(calls) == 1


@pytest.mark.skipif(not current_platform.is_rocm(), reason="requires ROCm")
@torch.inference_mode()
def test_sparse_indexer_topk_aiter_expands_mtp_row_ends(monkeypatch) -> None:
    """AITER's decode kernel only implements the next_n == 1 form, so spec
    decode batches must be handed the equivalent expanded per-row ends."""
    from vllm.model_executor.layers import indexer_topk

    recorded: dict = {}

    def decode_kernel(logits, next_n, seq_lens, *args, **kwargs) -> None:
        recorded["next_n"] = next_n
        recorded["row_ends"] = seq_lens.tolist()

    _patch_aiter_topk(monkeypatch, decode_kernel)
    monkeypatch.setattr(
        indexer_topk.ops,
        "top_k_per_row_decode",
        lambda *args: pytest.fail("an explicit aiter request must not fall back"),
    )

    next_n = 2
    logits = torch.randn(4, 4096, dtype=torch.float32, device="cuda")
    seq_lens = torch.tensor([[100], [200]], dtype=torch.int32, device="cuda")
    indices = torch.empty((4, 1024), dtype=torch.int32, device="cuda")

    _make_topk("aiter")(logits, seq_lens, next_n, indices, 1024, 800)

    assert recorded["next_n"] == 1
    assert recorded["row_ends"] == [99, 100, 199, 200]
