# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Round-trip tests for compressor → FP8 quant + KV cache insert → gather + dequant.

These tests cover:
  A) DeepseekV4 Attention: head_dim=512 (448 FP8 nope + 64 bf16 rope), quant_block=64
  B) Fused dequant+gather K cache
  C) Indexer:       head_dim=128 (all FP8), quant_block=128
  D) DeepseekV4 Attention magnitude range: correctness across small/large values
  E) Indexer fused Triton kernel: compress+norm+rope+quant+insert
  F) Indexer fused two-stage Triton kernel: head=512 cr>=128 (no-overlap)
"""

import math
from types import SimpleNamespace

import pytest
import torch

from vllm import _custom_ops as ops
from vllm.models.deepseek_v4.common.ops import (
    dequantize_and_gather_k_cache,
    quantize_and_insert_k_cache,
)
from vllm.models.deepseek_v4.common.ops.fused_compress_quant_cache import (
    _fused_kv_compress_norm_rope_insert_indexer_attn,
    _fused_kv_compress_norm_rope_insert_indexer_mxfp4_attn,
    _launch_two_stage_sparse_attn_compressor,
    compress_norm_rope_store_triton,
)
from vllm.models.deepseek_v4.compressor import _get_c128_boundary
from vllm.platforms import current_platform
from vllm.v1.attention.backends.mla.compressor_utils import (
    get_dspark_swa_index_width,
)
from vllm.v1.attention.ops.rocm_aiter_mla_sparse import (
    cp_gather_indexer_k_quant_cache_triton,
    indexer_k_quant_and_cache_triton,
)

from .test_fused_indexer_q_rope_quant import quantize_to_mxfp4


@pytest.mark.skipif(not current_platform.is_cuda(), reason="CUDA only")
@pytest.mark.parametrize(
    "cache_dtype,kv_mxfp8",
    [
        (torch.uint8, False),
        (torch.uint8, True),
        (torch.bfloat16, False),
        (torch.float8_e4m3fn, False),
    ],
)
@pytest.mark.parametrize("num_tokens", [1, 17, 1023, 1024])
@pytest.mark.parametrize("use_graph", [False, True])
def test_dspark_context_kv_matches_query_insert(
    cache_dtype, kv_mxfp8, num_tokens, use_graph
):
    """KV-only insertion must preserve every cache byte, including graph replay."""
    from vllm.models.deepseek_v41.nvidia.dspark import _insert_context_kv

    torch.manual_seed(42)
    block_size = 256
    num_blocks = math.ceil((num_tokens + 7) / block_size)
    row_size = (528 if kv_mxfp8 else 584) if cache_dtype == torch.uint8 else 512
    page_stride = math.ceil(block_size * row_size / 576) * 576 + 576
    backing = torch.full((num_blocks, page_stride), 3, device="cuda").to(cache_dtype)
    cache = backing.as_strided(
        (num_blocks, block_size, row_size), (page_stride, row_size, 1)
    )
    reference_backing = backing.clone()
    reference = reference_backing.as_strided(cache.shape, cache.stride())
    kv = torch.randn(num_tokens + 3, 512, device="cuda", dtype=torch.bfloat16)
    positions = torch.arange(num_tokens + 3, device="cuda") + 7
    angles = torch.randn(num_tokens + 16, 32, device="cuda")
    cos_sin = torch.cat((angles.cos(), angles.sin()), dim=-1)
    slots = torch.randperm(num_blocks * block_size, device="cuda")[:num_tokens]
    slots[::5] = -1
    scale = torch.tensor([0.7], device="cuda")
    # Only what _insert_context_kv actually reads: the record width comes off
    # the cache tensor, so no per-record width attribute is stubbed here.
    attn = SimpleNamespace(
        swa_cache_layer=SimpleNamespace(kv_cache=cache, block_size=block_size),
        kv_mxfp8=kv_mxfp8,
        rotary_emb=SimpleNamespace(cos_sin_cache=cos_sin),
        _flashinfer_fp8_kv_scale=scale,
    )

    def legacy_insert():
        q = torch.zeros(kv.shape[0], 8, 512, dtype=kv.dtype, device="cuda")
        if cache_dtype == torch.uint8:
            torch.ops._C.fused_deepseek_v4_qnorm_rope_kv_rope_quant_insert(
                q,
                kv,
                reference.view(num_blocks, -1),
                slots,
                positions,
                cos_sin,
                8,
                1e-20,
                block_size,
                True,
                kv_mxfp8,
            )
        elif cache_dtype == torch.bfloat16:
            torch.ops._C.fused_deepseek_v4_qnorm_rope_kv_rope_full_cache_bf16_insert(
                q,
                kv,
                reference,
                slots,
                positions,
                cos_sin,
                1e-20,
                block_size,
            )
        else:
            q_fp8 = torch.empty_like(q, dtype=cache_dtype)
            torch.ops._C.fused_deepseek_v4_qnorm_rope_kv_rope_full_cache_fp8_insert(
                q,
                kv,
                q_fp8,
                reference,
                slots,
                positions,
                cos_sin,
                scale,
                scale,
                1e-20,
                block_size,
            )

    def insert():
        _insert_context_kv(attn, kv, positions, slots)

    insert()
    graph = None
    if use_graph:
        graph = torch.cuda.CUDAGraph()
        with torch.cuda.graph(graph):
            insert()
    for iteration in range(3):
        kv.normal_()
        positions.add_(1)
        if iteration == 2:
            slots.fill_(-1)
        backing.view(torch.uint8).fill_(165)
        reference_backing.copy_(backing)
        before = kv.clone()
        legacy_insert()
        insert() if graph is None else graph.replay()
        torch.testing.assert_close(
            backing.view(torch.uint8),
            reference_backing.view(torch.uint8),
            rtol=0,
            atol=0,
        )
        torch.testing.assert_close(kv, before, rtol=0, atol=0)


@pytest.mark.skipif(
    not current_platform.is_cuda_alike(), reason="graph capture coverage"
)
@pytest.mark.parametrize("compress_ratio", [1, 2])
@pytest.mark.parametrize("use_graph", [False, True])
@pytest.mark.parametrize(
    "lengths",
    [(1, 5, 13), (1, 1, 1), (1, 2, 0), (7, 8, 8), (2, 3, 3)],
    ids=["prefill", "decode", "empty", "wrap_threshold", "packed_boundary"],
)
def test_v41_fused_save_compress_and_insert(
    compress_ratio: int, use_graph: bool, lengths: tuple[int, ...]
):
    """Ring states: a group boundary whose previous token is outside the chunk
    pools the ring row that chunk left behind, even when this chunk overwrites
    that row; only the last ``capacity`` tokens of a chunk are stored.

    Exercise historical group starts, adjacent requests, a chunk longer than
    the ring, single-token decode pairs, empty requests, wrap thresholds at
    odd and even starts, padded graph rows, and padded physical pages against
    a Torch reference with BF16 rounding before RoPE. Replays change both
    inputs and history so stale scratch/cache values cannot satisfy the test.
    """
    from vllm.models.deepseek_v41.common.ops.fused_compress_quant_cache import (
        fused_save_compress_norm,
        rope_quant_insert,
    )

    torch.manual_seed(42)
    device = "cuda"
    capacity = 8
    starts = [7, 8, 13]
    actual = sum(lengths)
    num_tokens = actual + 5
    positions = torch.zeros(num_tokens, dtype=torch.int64, device=device)
    req_ids = torch.zeros(actual + 2, dtype=torch.int32, device=device)
    state_slots = torch.full((actual + 2,), -1, dtype=torch.int64, device=device)
    cache_slots = torch.full_like(state_slots, -1)
    query_start_loc = torch.tensor(
        [0, *torch.tensor(lengths).cumsum(0).tolist()], dtype=torch.int32
    ).to(device)
    ring_blocks = torch.randperm(6, device=device)[:3]
    raw = torch.randn(num_tokens, 512 * compress_ratio, device=device)
    norm = torch.randn(512, dtype=torch.bfloat16, device=device)
    latent = torch.empty(num_tokens, 512, dtype=torch.bfloat16, device=device)
    angles = torch.randn(32, 32, device=device)
    cos_sin = torch.cat((angles.cos(), angles.sin()), dim=-1)
    # capacity x 1024 floats rounded up to 576-byte alignment.
    state_backing = torch.empty(6, 8208, device=device)
    state = state_backing.as_strided((6, capacity, 1024), (8208, 1024, 1))
    cache_block = 256 // compress_ratio
    cache_stride = math.ceil(cache_block * 584 / 576) * 576
    cache_backing = torch.empty(3, cache_stride, dtype=torch.uint8, device=device)
    cache = cache_backing.as_strided((3, cache_block, 584), (cache_stride, 584, 1))
    cursor = 0
    for req, (start, length) in enumerate(zip(starts, lengths)):
        pos = torch.arange(start, start + length, device=device)
        rows = slice(cursor, cursor + length)
        positions[rows] = pos
        req_ids[rows] = req
        state_slots[rows] = ring_blocks[req] * capacity + pos % capacity
        cache_slots[rows] = (2 - req) * cache_block + pos // compress_ratio
        cursor += length

    has_ring = compress_ratio == 2

    def run():
        fused_save_compress_norm(
            raw,
            positions,
            state if has_ring else None,
            state_slots,
            query_start_loc if has_ring else None,
            req_ids if has_ring else None,
            norm,
            1e-20,
            compress_ratio,
            latent,
        )
        rope_quant_insert(
            latent, positions, cos_sin, cache, cache_slots, compress_ratio
        )

    state_backing.normal_()
    run()
    graph = None
    if use_graph:
        stream = torch.cuda.Stream()
        stream.wait_stream(torch.cuda.current_stream())
        with torch.cuda.stream(stream):
            run()
        torch.cuda.current_stream().wait_stream(stream)
        graph = torch.cuda.CUDAGraph()
        with torch.cuda.graph(graph):
            run()

    for _ in range(3):
        raw.normal_()
        state_backing.normal_()
        history = state.clone()
        expected_state = state_backing.clone()
        expected_rows = expected_state.as_strided((6, capacity, 1024), (8208, 1024, 1))
        cache_backing.fill_(165)
        expected_cache = cache_backing.clone()
        latent.fill_(float("nan"))
        # Ratio 1 has no ring; ratio 2 stores each chunk's last `capacity` rows.
        for t in range(actual if has_ring else 0):
            req = req_ids[t].item()
            if query_start_loc[req + 1].item() - t > capacity:
                continue
            slot = state_slots[t].item()
            expected_rows[slot // capacity, slot % capacity, :512] = raw[t, :512]
            expected_rows[slot // capacity, slot % capacity, 512:] = raw[t, 512:]

        if graph is None:
            run()
        else:
            graph.replay()

        torch.testing.assert_close(state_backing, expected_state, rtol=0, atol=0)
        for t in range(num_tokens):
            pos = positions[t].item()
            if t >= actual or (pos + 1) % compress_ratio:
                assert torch.isnan(latent[t]).all()
                continue
            req = req_ids[t].item()
            group_rows = []
            for k in range(compress_ratio - 1, -1, -1):
                if t - k >= query_start_loc[req].item():
                    kv_score = raw[t - k]
                    if compress_ratio == 1:
                        kv_score = torch.cat((kv_score, torch.zeros_like(kv_score)))
                    group_rows.append(kv_score)
                else:
                    ring = ring_blocks[req].item()
                    group_rows.append(history[ring, (pos - k) % capacity])
            group = torch.stack(group_rows)
            pooled = (group[:, :512] * group[:, 512:].softmax(0)).sum(0)
            normed = pooled * torch.rsqrt(pooled.square().mean() + 1e-20) * norm
            torch.testing.assert_close(latent[t].float(), normed, rtol=0.004, atol=1e-6)

            slot = cache_slots[t].item()
            page, row = divmod(slot, cache_block)
            values = expected_cache[page, row * 576 : (row + 1) * 576]
            quantized, scales = _ue8m0_reference(latent[t, :448], 64, 448.0)
            values[:448] = quantized.view(torch.uint8)
            c, s = cos_sin[pos // compress_ratio * compress_ratio].chunk(2)
            rope_input = latent[t, 448:].float()
            rotated = (
                torch.stack(
                    (
                        rope_input[0::2] * c - rope_input[1::2] * s,
                        rope_input[1::2] * c + rope_input[0::2] * s,
                    ),
                    dim=-1,
                )
                .flatten()
                .to(torch.bfloat16)
            )
            actual_rope = cache_backing[page, row * 576 + 448 : (row + 1) * 576].view(
                torch.bfloat16
            )
            torch.testing.assert_close(actual_rope, rotated, rtol=0.008, atol=1e-6)
            values[448:] = actual_rope.view(torch.uint8)
            scale_offset = cache_block * 576 + row * 8
            expected_cache[page, scale_offset : scale_offset + 7] = (
                scales.log2() + 127
            ).to(torch.uint8)
            expected_cache[page, scale_offset + 7] = 0
        # Includes unwritten rows, page padding, quantized values and scales.
        torch.testing.assert_close(cache_backing, expected_cache, rtol=0, atol=0)


def _rotate_rope_tail(
    latent_row: torch.Tensor, cos_sin_row: torch.Tensor
) -> torch.Tensor:
    """GPT-J RoPE of one latent row's last 64 dims, in fp32."""
    row = latent_row.float()
    c, s = cos_sin_row.chunk(2)
    rotated = row.clone()
    even, odd = row[448::2], row[449::2]
    rotated[448::2] = even * c - odd * s
    rotated[449::2] = odd * c + even * s
    return rotated


def _mxfp8_record_reference(
    latent_row: torch.Tensor, cos_sin_row: torch.Tensor
) -> tuple[torch.Tensor, torch.Tensor]:
    """Reference V4.1 record for one token: 512 fp8 bytes + 16 UE8M0 scales.

    Mirrors FlashMLA's ``KVCacheLayout.V41_FP8Sparse`` quantizer: rotate the
    RoPE tail first, then scale every 32-dim tile, RoPE tiles included.
    """
    rotated = _rotate_rope_tail(latent_row, cos_sin_row)
    quantized, scales = _ue8m0_reference(rotated, 32, 448.0)
    return quantized.view(torch.uint8), (scales.log2() + 127).to(torch.uint8)


@pytest.mark.skipif(not current_platform.is_cuda(), reason="CUDA only")
@pytest.mark.parametrize("compress_ratio", [1, 2])
def test_v41_rope_insert_mxfp8_record(compress_ratio: int):
    """The 528-byte V4.1 record quantizes the RoPE dims instead of keeping them
    in bf16, so a page is [block_size x 512 fp8][block_size x 16 UE8M0].

    Byte-exact against the reference quantizer, including the rows and page
    padding the kernel must leave alone.
    """
    from vllm.models.deepseek_v41.common.ops.fused_compress_quant_cache import (
        rope_quant_insert,
    )

    torch.manual_seed(7)
    device = "cuda"
    cache_block = 64
    cache_stride = math.ceil(cache_block * 528 / 512) * 512
    num_tokens = 12

    positions = torch.arange(num_tokens, dtype=torch.int64, device=device)
    latent = torch.randn(num_tokens, 512, dtype=torch.bfloat16, device=device)
    angles = torch.randn(64, 32, device=device)
    cos_sin = torch.cat((angles.cos(), angles.sin()), dim=-1)

    cache_backing = torch.full((2, cache_stride), 165, dtype=torch.uint8, device=device)
    cache = cache_backing.as_strided((2, cache_block, 528), (cache_stride, 528, 1))
    # Leave one token unmapped so the kernel's negative-slot skip is covered.
    slots = torch.arange(num_tokens, dtype=torch.int64, device=device)
    slots[3] = -1

    rope_quant_insert(latent, positions, cos_sin, cache, slots, compress_ratio)

    expected = torch.full_like(cache_backing, 165)
    for t in range(num_tokens):
        slot = slots[t].item()
        pos = positions[t].item()
        if slot < 0 or (pos + 1) % compress_ratio:
            continue
        page, row = divmod(slot, cache_block)
        values, scales = _mxfp8_record_reference(
            latent[t], cos_sin[pos // compress_ratio * compress_ratio]
        )
        expected[page, row * 512 : (row + 1) * 512] = values
        scale_offset = cache_block * 512 + row * 16
        expected[page, scale_offset : scale_offset + 16] = scales

    torch.testing.assert_close(cache_backing, expected, rtol=0, atol=0)


@pytest.mark.skipif(not current_platform.is_cuda(), reason="CUDA only")
def test_v41_mxfp8_cache_round_trip():
    """quantize_and_insert -> dequantize_and_gather recovers the V4.1 record.

    The RoPE dims now go through fp8, so they carry quantization error too --
    the check is that every dim survives within one MXFP8 tile step, and that
    the bytes on the way match the reference quantizer.
    """
    from vllm.models.deepseek_v41.common.ops import (
        dequantize_and_gather_k_cache,
        quantize_and_insert_k_cache,
    )

    torch.manual_seed(11)
    device = "cuda"
    block_size = 64
    num_tokens = 70
    num_blocks = 4
    page_bytes = math.ceil(block_size * 528 / 512) * 512

    k = torch.randn(num_tokens, 512, dtype=torch.bfloat16, device=device)
    cache = torch.zeros(num_blocks, page_bytes, dtype=torch.uint8, device=device)
    slot_mapping = torch.arange(num_tokens, dtype=torch.int64, device=device)
    quantize_and_insert_k_cache(
        k, cache, slot_mapping, block_size=block_size, bytes_per_token=528
    )

    for t in (0, 1, block_size, num_tokens - 1):
        values, scales = _ue8m0_reference(k[t].float(), 32, 448.0)
        page, row = divmod(t, block_size)
        torch.testing.assert_close(
            cache[page, row * 512 : (row + 1) * 512],
            values.view(torch.uint8),
            rtol=0,
            atol=0,
        )
        scale_offset = block_size * 512 + row * 16
        torch.testing.assert_close(
            cache[page, scale_offset : scale_offset + 16],
            (scales.log2() + 127).to(torch.uint8),
            rtol=0,
            atol=0,
        )

    out = torch.zeros(1, num_tokens, 512, dtype=torch.bfloat16, device=device)
    dequantize_and_gather_k_cache(
        out,
        cache.view(num_blocks, block_size, 528),
        seq_lens=torch.tensor([num_tokens], dtype=torch.int32, device=device),
        gather_lens=None,
        block_table=torch.arange(num_blocks, dtype=torch.int32, device=device).view(
            1, -1
        ),
        block_size=block_size,
        offset=0,
    )
    # Half an e4m3 ULP at the top of the range (16 units) bounds the error.
    tile_amax = k.float().abs().view(num_tokens, 16, 32).amax(-1).clamp(min=1e-4)
    tolerance = 16.0 * (tile_amax / 448.0).log2().ceil().exp2()
    error = (out[0].float() - k.float()).abs()
    assert (error <= tolerance.repeat_interleave(32, dim=-1)).all()


# e2m1 magnitudes indexed by the low 3 bits of a code; bit 3 is the sign.
_E2M1_MAGNITUDES = [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0]


def _decode_nvfp4_row(packed: torch.Tensor, scales: torch.Tensor) -> torch.Tensor:
    """Unpack e2m1 rows and apply their e4m3 tile scales."""
    mags = torch.tensor(_E2M1_MAGNITUDES, device=packed.device)
    codes = torch.empty(
        (*packed.shape[:-1], 512), dtype=torch.uint8, device=packed.device
    )
    codes[..., 0::2] = packed & 0xF  # even element in the low nibble
    codes[..., 1::2] = packed >> 4
    vals = mags[(codes & 7).long()] * torch.where(codes >= 8, -1.0, 1.0)
    return (
        vals.reshape(*packed.shape[:-1], 32, 16)
        * scales.view(torch.float8_e4m3fn).float().reshape(*packed.shape[:-1], 32, 1)
    ).flatten(-2)


@pytest.mark.skipif(not current_platform.is_cuda(), reason="CUDA only")
@pytest.mark.parametrize("compress_ratio", [1, 2])
def test_v41_rope_insert_nvfp4_record(compress_ratio: int):
    """The 288-byte V4.1 NVFP4 record: 256 B of e2m1 pairs then 32 e4m3 scales.

    Scales are byte-exact against the reference (``amax / 6`` clamped to the
    e4m3 range); values are checked as round-to-nearest on the e2m1 grid, whose
    coarsest step is 2 between 4 and 6, so half a step is ``1.0 * scale``.
    """
    from vllm.models.deepseek_v41.common.ops.fused_compress_quant_cache import (
        rope_quant_insert,
    )

    torch.manual_seed(13)
    device = "cuda"
    cache_block = 64
    cache_stride = math.ceil(cache_block * 288 / 512) * 512
    num_tokens = 12

    positions = torch.arange(num_tokens, dtype=torch.int64, device=device)
    latent = torch.randn(num_tokens, 512, dtype=torch.bfloat16, device=device)
    angles = torch.randn(64, 32, device=device)
    cos_sin = torch.cat((angles.cos(), angles.sin()), dim=-1)

    cache_backing = torch.full((2, cache_stride), 165, dtype=torch.uint8, device=device)
    cache = cache_backing.as_strided((2, cache_block, 288), (cache_stride, 288, 1))
    slots = torch.arange(num_tokens, dtype=torch.int64, device=device)
    slots[3] = -1  # covers the negative-slot skip

    rope_quant_insert(latent, positions, cos_sin, cache, slots, compress_ratio)

    untouched = torch.full_like(cache_backing, 165)
    for t in range(num_tokens):
        slot = slots[t].item()
        pos = positions[t].item()
        if slot < 0 or (pos + 1) % compress_ratio:
            continue
        page, row = divmod(slot, cache_block)
        cs = cos_sin[pos // compress_ratio * compress_ratio]
        rotated = _rotate_rope_tail(latent[t], cs)
        amax = rotated.view(32, 16).abs().amax(-1)
        scale = (amax / 6.0).clamp(2.0**-9, 448.0).to(torch.float8_e4m3fn)

        got_scales = cache_backing[page, cache_block * 256 + row * 32 :][:32]
        torch.testing.assert_close(got_scales, scale.view(torch.uint8), rtol=0, atol=0)
        decoded = _decode_nvfp4_row(
            cache_backing[page, row * 256 : (row + 1) * 256], got_scales
        )
        step = scale.float().repeat_interleave(16)
        assert ((decoded - rotated).abs() <= step).all()
        untouched[page, row * 256 : (row + 1) * 256] = cache_backing[
            page, row * 256 : (row + 1) * 256
        ]
        untouched[page, cache_block * 256 + row * 32 :][:32] = got_scales
    # Rows the kernel must not have touched, page padding included.
    torch.testing.assert_close(cache_backing, untouched, rtol=0, atol=0)


@pytest.mark.skipif(not current_platform.is_cuda(), reason="CUDA only")
@pytest.mark.parametrize("num_reqs", [1, 8, 32])
@pytest.mark.parametrize("num_tokens", [70, 128, 129, 2051])
@pytest.mark.parametrize("partial", [False, True])
def test_v41_nvfp4_gather_matches_insert(num_reqs, num_tokens, partial):
    """The NVFP4 gather dequantizes exactly what the insert kernel wrote."""
    from vllm.models.deepseek_v41.common.ops import dequantize_and_gather_k_cache
    from vllm.models.deepseek_v41.common.ops.fused_compress_quant_cache import (
        rope_quant_insert,
    )

    torch.manual_seed(17)
    device = "cuda"
    block_size = 64
    num_blocks = math.ceil(num_tokens / block_size)
    page_bytes = math.ceil(block_size * 288 / 512) * 512

    positions = torch.arange(num_tokens, dtype=torch.int64, device=device)
    latent = torch.randn(num_tokens, 512, dtype=torch.bfloat16, device=device)
    angles = torch.randn(num_tokens, 32, device=device)
    cos_sin = torch.cat((angles.cos(), angles.sin()), dim=-1)
    backing = torch.zeros(num_blocks, page_bytes, dtype=torch.uint8, device=device)
    cache = backing.as_strided((num_blocks, block_size, 288), (page_bytes, 288, 1))
    slots = torch.arange(num_tokens, dtype=torch.int64, device=device)
    rope_quant_insert(latent, positions, cos_sin, cache, slots, 1)

    offset = 7 if partial else 0
    seq_lens = num_tokens - torch.arange(num_reqs, dtype=torch.int32, device=device)
    gather_lens = seq_lens - 3 if partial else None
    storage = torch.full(
        (num_reqs, num_tokens + offset + 16, 512),
        -123.0,
        dtype=torch.bfloat16,
        device=device,
    )
    out = storage[:, : num_tokens + offset]
    block_table = torch.arange(num_blocks, dtype=torch.int32, device=device).repeat(
        num_reqs, 1
    )

    def gather():
        dequantize_and_gather_k_cache(
            out,
            cache,
            seq_lens=seq_lens,
            gather_lens=gather_lens,
            block_table=block_table,
            block_size=block_size,
            offset=offset,
        )

    expected = _decode_nvfp4_row(
        backing[:, : block_size * 256].reshape(-1, 256),
        backing[:, block_size * 256 : block_size * 288].reshape(-1, 32),
    )

    def check(shorter_by):
        assert (storage[:, num_tokens + offset :] == -123).all()
        for req in range(num_reqs):
            start = 3 if partial else 0
            length = num_tokens - req - start - shorter_by
            assert (out[req, :offset] == -123).all()
            assert (out[req, offset + length :] == -123).all()
            torch.testing.assert_close(
                out[req, offset : offset + length].float(),
                expected[start : start + length],
                rtol=0,
                atol=0,
            )

    gather()
    check(0)
    graph = torch.cuda.CUDAGraph()
    with torch.cuda.graph(graph):
        gather()
    storage.fill_(-123)
    seq_lens.sub_(1)
    if gather_lens is not None:
        gather_lens.sub_(1)
    graph.replay()
    check(1)


@pytest.mark.skipif(not current_platform.is_cuda(), reason="CUDA only")
@pytest.mark.parametrize("compress_ratio", [1, 2])
@pytest.mark.parametrize("store_fp8", [False, True])
def test_v41_rope_insert_plain_row(compress_ratio: int, store_fp8: bool):
    """The FlashInfer plain-row compressed cache gets [448 NoPE | 64 RoPE] rows.

    bf16 caches store the latent verbatim with the RoPE tail rotated; per-tensor
    fp8 caches scale the bf16-rounded row by 1/fp8_scale and clamp to e4m3.
    Non-boundary positions, negative slots and untouched rows must keep their
    prior contents.
    """
    from vllm.models.deepseek_v41.common.ops.fused_compress_quant_cache import (
        rope_quant_insert,
    )

    torch.manual_seed(3)
    device = "cuda"
    num_tokens = 9
    cache_block = 4
    positions = torch.arange(5, 5 + num_tokens, dtype=torch.int64, device=device)
    cache_slots = torch.randperm(3 * cache_block, device=device)[:num_tokens]
    cache_slots[2] = -1
    latent = torch.randn(num_tokens, 512, device=device).to(torch.bfloat16)
    angles = torch.randn(32, 32, device=device)
    cos_sin = torch.cat((angles.cos(), angles.sin()), dim=-1)
    dtype = torch.float8_e4m3fn if store_fp8 else torch.bfloat16
    # Pad the page stride so the kernel must honour strides, not shape.
    cache_backing = torch.full((3, (cache_block + 1) * 512), 3.0, device=device)
    cache_backing = cache_backing.to(dtype)
    cache = cache_backing.as_strided(
        (3, cache_block, 512), ((cache_block + 1) * 512, 512, 1)
    )
    fp8_scale = torch.tensor([0.5], dtype=torch.float32, device=device)
    expected = cache_backing.clone().float()

    rope_quant_insert(
        latent,
        positions,
        cos_sin,
        cache,
        cache_slots,
        compress_ratio,
        fp8_scale=fp8_scale if store_fp8 else None,
    )

    for t in range(num_tokens):
        pos = positions[t].item()
        slot = cache_slots[t].item()
        if slot < 0 or (pos + 1) % compress_ratio:
            continue
        c, s = cos_sin[pos // compress_ratio * compress_ratio].chunk(2)
        rope_input = latent[t, 448:].float()
        rotated = torch.stack(
            (
                rope_input[0::2] * c - rope_input[1::2] * s,
                rope_input[1::2] * c + rope_input[0::2] * s,
            ),
            dim=-1,
        ).flatten()
        row = torch.cat((latent[t, :448], rotated.to(torch.bfloat16))).float()
        if store_fp8:
            row = (row * (1.0 / fp8_scale)).clamp(-448.0, 448.0)
            row = row.to(torch.float8_e4m3fn).float()
        page, idx = divmod(slot, cache_block)
        expected[page, idx * 512 : (idx + 1) * 512] = row

    actual = cache_backing.float()
    if store_fp8:
        # One e4m3 ulp of slack for FMA-vs-separate rounding in the RoPE tail.
        torch.testing.assert_close(actual, expected, rtol=0.13, atol=1e-2)
    else:
        nope = torch.arange(512, device=device) < 448
        nope = nope.repeat(cache_block + 1)
        torch.testing.assert_close(actual[:, nope], expected[:, nope], rtol=0, atol=0)
        torch.testing.assert_close(
            actual[:, ~nope], expected[:, ~nope], rtol=0.008, atol=1e-6
        )


@pytest.mark.skipif(not current_platform.is_cuda(), reason="Triton kernel")
def test_v41_compressor_metadata_maps_tokens_to_their_ring():
    """The ring group's generic slot mapping is disabled (all PAD), so the
    builder must map every real token to ``ring_block * capacity + pos %
    capacity`` and keep padding tokens at PAD."""
    from unittest.mock import MagicMock

    from vllm.models.deepseek_v41.compressor import CompressorMetadataBuilder
    from vllm.v1.attention.backend import CommonAttentionMetadata
    from vllm.v1.kv_cache_interface import CircularBufferSpec

    capacity = 8
    vllm_config = MagicMock()
    vllm_config.scheduler_config.max_num_batched_tokens = 16
    spec = CircularBufferSpec(
        block_size=capacity,
        num_kv_heads=1,
        head_size=1024,
        head_size_v=0,
        dtype=torch.float32,
    )
    device = torch.device("cuda")
    builder = CompressorMetadataBuilder(spec, ["state"], vllm_config, device)

    # Two requests: 3 tokens at positions 13..15 on ring block 5, then 2
    # tokens at positions 7..8 on ring block 2; three padding tokens.
    query_start_loc = torch.tensor([0, 3, 5], dtype=torch.int32, device=device)
    positions = torch.tensor([13, 14, 15, 7, 8, 0, 0, 0], device=device)
    block_table = torch.tensor([[5], [2]], dtype=torch.int32, device=device)
    common = CommonAttentionMetadata(
        query_start_loc=query_start_loc,
        query_start_loc_cpu=query_start_loc.cpu(),
        seq_lens=torch.tensor([16, 9], dtype=torch.int32, device=device),
        num_reqs=2,
        num_actual_tokens=5,
        max_query_len=3,
        max_seq_len=16,
        block_table_tensor=block_table,
        slot_mapping=torch.full((8,), -1, dtype=torch.int64, device=device),
        positions=positions,
    )
    metadata = builder.build(0, common)

    expected = [5 * 8 + 5, 5 * 8 + 6, 5 * 8 + 7, 2 * 8 + 7, 2 * 8 + 0, -1, -1, -1]
    assert metadata.slot_mapping.tolist() == expected
    assert metadata.query_start_loc is query_start_loc
    assert metadata.token_to_req_indices.tolist() == [0, 0, 0, 1, 1]


@pytest.mark.skipif(not current_platform.is_cuda(), reason="CUDA stream coverage")
@pytest.mark.parametrize(
    "use_aux,use_graph", [(False, False), (True, False), (True, True)]
)
@torch.inference_mode()
def test_v41_attention_joins_cache_writes_before_consumption(use_aux, use_graph):
    """Both reused-event joins must publish this forward's states and cache rows."""
    from vllm.forward_context import ForwardContext, override_forward_context
    from vllm.models.deepseek_v41.attention import (
        DeepseekV4Attention,
        DeepseekV4Indexer,
    )
    from vllm.models.deepseek_v41.compressor import DeepseekCompressor

    torch.manual_seed(43)
    raw = torch.randn(19, 1024, device="cuda")
    q = torch.zeros(19, 512, dtype=torch.bfloat16, device="cuda")
    positions = torch.arange(7, 26, device="cuda")
    state_slots = positions.clone()
    state_slots[-2:] = -1
    cache_slots = torch.where(state_slots >= 0, positions // 2, -1)
    state = torch.empty(4, 8, 1024, device="cuda")
    main = torch.empty(1, 128, 584, dtype=torch.uint8, device="cuda")
    index = torch.empty(1, 128, 132, dtype=torch.uint8, device="cuda")
    caches = (state, main, index)
    observed = tuple(torch.empty_like(cache) for cache in caches)
    rotary = SimpleNamespace(cos_sin_cache=torch.randn(32, 64, device="cuda"))
    metadata = {
        "state": SimpleNamespace(
            slot_mapping=state_slots,
            query_start_loc=torch.tensor([0, 19], dtype=torch.int32, device="cuda"),
            token_to_req_indices=torch.zeros(19, dtype=torch.int32, device="cuda"),
        ),
        "main": SimpleNamespace(slot_mapping=cache_slots),
        "index": SimpleNamespace(slot_mapping=cache_slots),
    }
    context = ForwardContext({}, metadata, {})
    compressor = DeepseekCompressor.__new__(DeepseekCompressor)
    torch.nn.Module.__init__(compressor)
    compressor.head_dim, compressor.rope_head_dim, compressor.compress_ratio = (
        512,
        64,
        2,
    )
    compressor.rms_norm_eps = 1e-20
    compressor.norm = SimpleNamespace(
        weight=torch.ones(512, dtype=torch.bfloat16, device="cuda")
    )
    compressor.state_cache = SimpleNamespace(prefix="state", kv_cache=state)
    compressor.k_cache_prefix = "main"
    compressor._static_forward_context = {"main": SimpleNamespace(kv_cache=main)}
    indexer_weight = torch.randn(128, 512, dtype=torch.bfloat16, device="cuda")
    indexer = SimpleNamespace(
        owns_k=True,
        wk=lambda latent: (torch.nn.functional.linear(latent, indexer_weight), None),
        k_norm=SimpleNamespace(
            weight=torch.ones(128, dtype=torch.bfloat16, device="cuda"),
            variance_epsilon=1e-20,
        ),
        k_cache=SimpleNamespace(prefix="index", kv_cache=index),
        compress_ratio=2,
        use_fp4_kv=False,
    )

    def prepare_indexer(qr, latent, weights, positions, rotary, qr_scale):
        DeepseekV4Indexer._produce_k(indexer, latent, positions, rotary)
        return None, None, None

    def observe(*args):
        for output, cache in zip(observed, caches):
            output.copy_(cache)

    auxiliary = [torch.cuda.Stream()] if use_aux else None
    attention = SimpleNamespace(
        compressor=compressor,
        indexer=prepare_indexer,
        aux_stream_list=auxiliary,
        ln_events=[torch.cuda.Event(), torch.cuda.Event()],
        rotary_emb=rotary,
        indexer_rotary_emb=rotary,
        n_local_heads=1,
        head_dim=512,
        _wq_b_proj=lambda qr, scale: qr.clone(),
        _fused_qnorm_rope_kv_insert=lambda q, kv, pos, meta: q,
        _sparse_indexer_and_attn=observe,
    )

    def run():
        with override_forward_context(context):
            DeepseekV4Attention._prepare_and_attn(
                attention, q, q, q, None, raw, q, positions, q
            )

    def reset():
        state.zero_()
        main.fill_(165)
        index.fill_(165)

    reset()
    run()
    graph = None
    if use_graph:
        stream = torch.cuda.Stream()
        stream.wait_stream(torch.cuda.current_stream())
        with torch.cuda.stream(stream):
            run()
        torch.cuda.current_stream().wait_stream(stream)
        graph = torch.cuda.CUDAGraph()
        with torch.cuda.graph(graph):
            run()

    previous = None
    for _ in range(3):
        raw.normal_()
        reset()
        attention.aux_stream_list = None
        run()
        expected = tuple(output.clone() for output in observed)
        reset()
        attention.aux_stream_list = auxiliary
        if graph is None:
            run()
        else:
            graph.replay()
        for actual, reference in zip(observed, expected):
            torch.testing.assert_close(actual, reference, rtol=0, atol=0)
        if previous is not None:
            assert all(not torch.equal(a, b) for a, b in zip(expected, previous))
        previous = expected


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

        return _ON_GFX950
    except Exception:
        return False


@pytest.mark.skipif(not current_platform.is_rocm(), reason="ROCm-only dispatch")
def test_cp_gather_despecialized_kernel_is_gfx950_only(monkeypatch):
    from vllm.v1.attention.ops import rocm_aiter_mla_sparse as mod

    class FakeKernel:
        def __init__(self):
            self.calls = []

        def __getitem__(self, grid):
            def launch(*args):
                self.calls.append((grid, args))

            return launch

    legacy_kernel = FakeKernel()
    gfx950_kernel = FakeKernel()
    monkeypatch.setattr(mod, "_cp_gather_indexer_quant_cache_kernel", legacy_kernel)
    monkeypatch.setattr(
        mod,
        "_cp_gather_indexer_quant_cache_gfx950_kernel",
        gfx950_kernel,
    )

    k_cache = torch.zeros((4, 1, 132), dtype=torch.uint8)
    k_fp8 = torch.empty((5, 128), dtype=current_platform.fp8_dtype())
    k_scale = torch.empty((5, 4), dtype=torch.uint8)
    block_table = torch.zeros((2, 7), dtype=torch.int32)
    cu_seqlen = torch.tensor([0, 2, 5], dtype=torch.int32)
    token_to_seq = torch.tensor([0, 0, 1, 1, 1], dtype=torch.int32)
    args = (k_cache, k_fp8, k_scale, block_table, cu_seqlen, token_to_seq)

    monkeypatch.setattr(mod, "_ON_GFX950", True)
    mod.cp_gather_indexer_k_quant_cache_triton(*args)
    assert len(gfx950_kernel.calls) == 1
    assert not legacy_kernel.calls
    gfx950_grid, gfx950_args = gfx950_kernel.calls[0]
    assert gfx950_grid == (5,)
    assert len(gfx950_args) == 18
    assert gfx950_args[-3:] == (2, 7, 4)

    monkeypatch.setattr(mod, "_ON_GFX950", False)
    mod.cp_gather_indexer_k_quant_cache_triton(*args)
    assert len(legacy_kernel.calls) == 1
    legacy_grid, legacy_args = legacy_kernel.calls[0]
    assert legacy_grid == (5,)
    assert len(legacy_args) == 19
    assert legacy_args[-4:] == (5, 2, 7, 4)


@pytest.mark.parametrize(
    ("window_size", "num_speculative_tokens", "expected"),
    [(128, 5, 192), (512, 5, 576), (1024, 0, 1024)],
)
def test_get_dspark_swa_index_width(
    window_size: int, num_speculative_tokens: int, expected: int
):
    assert get_dspark_swa_index_width(window_size, num_speculative_tokens) == expected


def _ue8m0_reference(x: torch.Tensor, block_size: int, fp8_max: float):
    """PyTorch reference for UE8M0 FP8 quantization (per-block, power-of-2 scale).

    Returns (x_fp8, scales) where x_fp8 is float8_e4m3fn and scales are float32.
    """
    assert x.dim() == 1
    n = x.numel()
    n_blocks = math.ceil(n / block_size)
    x_fp8 = torch.zeros(n, dtype=torch.float8_e4m3fn, device=x.device)
    scales = torch.zeros(n_blocks, dtype=torch.float32, device=x.device)

    for i in range(n_blocks):
        start = i * block_size
        end = min(start + block_size, n)
        block = x[start:end].float()
        amax = block.abs().max().clamp(min=1e-4)
        raw_scale = amax / fp8_max
        exponent = math.ceil(math.log2(raw_scale.item()))
        scale = 2.0**exponent
        scales[i] = scale
        quantized = (block / scale).clamp(-fp8_max, fp8_max)
        x_fp8[start:end] = quantized.to(torch.float8_e4m3fn)

    return x_fp8, scales


def _decode_dsv4_cache_row(
    cache: torch.Tensor, block_size: int, scrub_nan: bool
) -> torch.Tensor:
    flat = cache.flatten()
    nope = flat[:448].view(torch.float8_e4m3fn).to(torch.bfloat16)
    encoded = flat[block_size * 576 : block_size * 576 + 7]
    scales = torch.exp2(encoded.to(torch.float32) - 127.0).to(torch.bfloat16)
    nope = nope * scales.repeat_interleave(64)
    rope = flat[448:576].view(torch.bfloat16)
    decoded = torch.cat((nope, rope))
    if scrub_nan:
        decoded = torch.where(decoded == decoded, decoded, 0.0)
    return decoded


def _assert_nan_free_cache_matches_legacy_scrub(
    cache: torch.Tensor, block_size: int
) -> None:
    flat = cache.flatten()
    scale_base = block_size * 576
    scale_codes = flat[scale_base : scale_base + 8]
    assert scale_codes[0].item() == 254
    assert scale_codes[1].item() == 247
    assert scale_codes[:7].max().item() <= 254
    nope_bytes = flat[:448]
    assert not ((nope_bytes == 0x7F) | (nope_bytes == 0xFF)).any()

    rope = flat[448:576].view(torch.bfloat16)
    assert not torch.isnan(rope).any()
    assert torch.isposinf(rope[0])
    assert torch.equal(rope[1:4], torch.zeros_like(rope[1:4]))

    legacy_cache = cache.clone()
    legacy_flat = legacy_cache.flatten()
    legacy_flat[scale_base] = 255
    legacy_rope = legacy_flat[448:576].view(torch.bfloat16)
    legacy_rope[1:4] = float("nan")
    canonical = _decode_dsv4_cache_row(cache, block_size, scrub_nan=False)
    legacy = _decode_dsv4_cache_row(legacy_cache, block_size, scrub_nan=True)
    torch.testing.assert_close(canonical, legacy, rtol=0, atol=0)
    assert torch.isinf(canonical[0])
    assert torch.isposinf(canonical[64])


@pytest.mark.skipif(
    not _on_gfx950(),
    reason="NaN-free fp8_ds_mla compressed-cache contract is gfx950-only",
)
@pytest.mark.parametrize("writer", ["single_pass", "two_stage_finalizer"])
def test_gfx950_compressed_cache_canonicalizes_nonfinite(writer: str) -> None:
    head_dim = 512
    rope_dim = 64
    block_size = 4
    device = "cuda"

    positions = torch.zeros(1, dtype=torch.int64, device=device)
    slot_mapping = torch.zeros(1, dtype=torch.int64, device=device)
    rms_weight = torch.ones(head_dim, dtype=torch.bfloat16, device=device)
    rms_weight[0] = float("inf")
    rms_weight[64] = torch.finfo(torch.bfloat16).max
    rms_weight[448] = float("inf")
    rms_weight[450] = float("nan")
    cos_sin_cache = torch.zeros(1, rope_dim, dtype=torch.float32, device=device)
    cos_sin_cache[:, : rope_dim // 2] = 1.0
    cache = torch.zeros(1, block_size, 584, dtype=torch.uint8, device=device)

    state_cache = torch.zeros(1, 1, 2 * head_dim, dtype=torch.float32, device=device)
    state_cache[..., :head_dim] = 1.0
    token_to_req = torch.zeros(1, dtype=torch.int32, device=device)
    block_table = torch.zeros(1, 1, dtype=torch.int32, device=device)

    if writer == "single_pass":
        compress_norm_rope_store_triton(
            state_cache=state_cache,
            num_actual=1,
            token_to_req_indices=token_to_req,
            positions=positions,
            slot_mapping=slot_mapping,
            block_table=block_table,
            block_size=1,
            state_width=head_dim,
            cos_sin_cache=cos_sin_cache,
            kv_cache=cache,
            k_cache_metadata=SimpleNamespace(slot_mapping=slot_mapping),
            pdl_kwargs={},
            head_dim=head_dim,
            rope_head_dim=rope_dim,
            compress_ratio=1,
            overlap=False,
            use_fp4_cache=False,
            rms_norm_weight=rms_weight,
            rms_norm_eps=1e-6,
            quant_block=64,
            token_stride=576,
            scale_dim=8,
        )
    else:
        _launch_two_stage_sparse_attn_compressor(
            state_cache,
            token_to_req,
            positions,
            slot_mapping,
            block_table,
            1,
            head_dim,
            1,
            cos_sin_cache,
            cache,
            slot_mapping,
            rms_weight,
            1e-6,
            64,
            576,
            8,
            head_dim,
            rope_dim,
            1,
            torch.empty(1, head_dim, dtype=torch.float32, device=device),
        )

    _assert_nan_free_cache_matches_legacy_scrub(cache, block_size)


@pytest.mark.parametrize(
    ("starts", "query_start_loc", "expected"),
    [
        ([0], [0, 127], False),
        ([0], [0, 128], True),
        ([127], [0, 1], True),
        ([128], [0, 127], False),
        ([1, 255], [0, 1, 2], True),
        (None, [0, 1], None),
    ],
)
def test_get_c128_boundary(starts, query_start_loc, expected):
    query_start_loc_tensor = torch.tensor(query_start_loc)
    query_lens = query_start_loc_tensor[1:] - query_start_loc_tensor[:-1]
    metadata = SimpleNamespace(
        seq_lens_cpu_upper_bound=(
            None if starts is None else torch.tensor(starts) + query_lens
        ),
        query_start_loc_cpu=query_start_loc_tensor,
    )
    assert _get_c128_boundary(metadata) is expected


# ── Test A: DeepseekV4 Attention path ──────────────────────────────────────────────


@pytest.mark.parametrize("num_tokens", [1, 4, 8, 17])
@pytest.mark.parametrize("block_size", [16, 64])
def test_deepseek_v4_attention_quant_cache_roundtrip(num_tokens: int, block_size: int):
    """compressed_kv → quantize_and_insert_k_cache → dequantize_and_gather_k_cache
    → compare against original."""
    HEAD_DIM = 512
    NOPE_DIM = 448
    HEAD_BYTES = 584  # 448 fp8 + 128 bf16 + 8 uint8 scale
    FP8_MAX = 448.0
    QUANT_BLOCK = 64

    num_blocks = (num_tokens + block_size - 1) // block_size + 1
    device = "cuda"

    # Random compressed_kv (simulates compressor output)
    compressed_kv = torch.randn(
        num_tokens, HEAD_DIM, dtype=torch.bfloat16, device=device
    )

    # ── Quant + insert ──────────────────────────────────────────────────
    k_cache = torch.zeros(
        num_blocks, block_size, HEAD_BYTES, dtype=torch.uint8, device=device
    )
    k_cache_2d = k_cache.view(num_blocks, -1)

    # Sequential slot mapping: token i → slot i
    slot_mapping = torch.arange(num_tokens, dtype=torch.int64, device=device)

    quantize_and_insert_k_cache(
        compressed_kv, k_cache_2d, slot_mapping, block_size=block_size
    )

    # ── Gather + dequant ────────────────────────────────────────────────
    num_reqs = 1
    max_blocks_per_seq = num_blocks
    out = torch.zeros(
        num_reqs, num_tokens, HEAD_DIM, dtype=torch.bfloat16, device=device
    )
    seq_lens = torch.tensor([num_tokens], dtype=torch.int32, device=device)
    # block_table: request 0 uses physical blocks 0, 1, ...
    block_table = torch.arange(
        max_blocks_per_seq, dtype=torch.int32, device=device
    ).unsqueeze(0)

    dequantize_and_gather_k_cache(
        out, k_cache, seq_lens, None, block_table, block_size, offset=0
    )

    recovered = out[0, :num_tokens]

    # ── NoPE portion (first 448): FP8 quantized, expect UE8M0 error ──
    nope_orig = compressed_kv[:, :NOPE_DIM].float()
    nope_recv = recovered[:, :NOPE_DIM].float()
    nope_diff = (nope_recv - nope_orig).abs()

    # Per-token check: FP8 e4m3 (3-bit mantissa) worst-case error is
    # half-ULP at the largest representable value.  At y ≈ 448 (max),
    # ULP = 2^(8-3) = 32, so error ≤ 16 * scale.
    for t in range(num_tokens):
        _, scales = _ue8m0_reference(
            compressed_kv[t, :NOPE_DIM].float(), QUANT_BLOCK, FP8_MAX
        )
        max_allowed = 16.0 * scales.max().item()
        token_diff = nope_diff[t].max().item()
        assert token_diff <= max_allowed, (
            f"Token {t} nope diff {token_diff} exceeds max_allowed "
            f"{max_allowed} (scale={scales.max().item()})"
        )

    # ── RoPE portion (last 64): stored as bf16, should be exact ─────
    rope_diff = (recovered[:, NOPE_DIM:] - compressed_kv[:, NOPE_DIM:]).abs()
    assert rope_diff.max().item() == 0.0, (
        f"RoPE portion should be exact but got max diff {rope_diff.max().item()}"
    )


# ── Test B: Fused dequant+gather K cache ────────────────────────────────────


def _dequantize_and_gather_k_cache_reference(
    out: torch.Tensor,
    k_cache: torch.Tensor,
    seq_lens: torch.Tensor,
    gather_lens: torch.Tensor | None,
    block_table: torch.Tensor,
    block_size: int,
    offset: int,
) -> None:
    fp8_dim = 448
    bf16_dim = 64
    scale_dim = 8
    quant_block = 64
    token_data_size = fp8_dim + bf16_dim * 2

    for req_id in range(seq_lens.shape[0]):
        seq_len = seq_lens[req_id].item()
        gather_len = gather_lens[req_id].item() if gather_lens is not None else seq_len
        start_pos = seq_len - gather_len

        for i in range(gather_len):
            pos = start_pos + i
            pos_in_block = pos % block_size
            block_idx = block_table[req_id, pos // block_size].item()
            cache_block = k_cache[block_idx].view(-1)

            token_data_start = pos_in_block * token_data_size
            fp8_bytes = cache_block[token_data_start : token_data_start + fp8_dim]
            fp8_vals = fp8_bytes.view(torch.float8_e4m3fn).float()

            scale_start = block_size * token_data_size + pos_in_block * scale_dim
            encoded_scales = cache_block[scale_start : scale_start + scale_dim]
            scales = torch.exp2(encoded_scales[:7].float() - 127.0)
            dequant = fp8_vals * scales.repeat_interleave(quant_block)

            bf16_start = token_data_start + fp8_dim
            bf16_bytes = cache_block[bf16_start : bf16_start + bf16_dim * 2]
            bf16_tail = bf16_bytes.view(torch.bfloat16)

            out[req_id, offset + i, :fp8_dim] = dequant
            out[req_id, offset + i, fp8_dim:] = bf16_tail


@pytest.mark.parametrize(
    ("seq_lens_host", "gather_lens_host", "offset"),
    [
        ([9, 23, 7], None, 0),
        ([19, 8, 257], [6, 8, 129], 5),
    ],
)
def test_dequantize_and_gather_k_cache(
    seq_lens_host: list[int],
    gather_lens_host: list[int] | None,
    offset: int,
):
    block_size = 64
    head_dim = 512
    nope_dim = 448
    scale_dim = 8
    head_bytes = nope_dim + (head_dim - nope_dim) * 2 + scale_dim
    device = "cuda"
    num_reqs = len(seq_lens_host)
    num_tokens = sum(seq_lens_host)
    max_gather_len = max(gather_lens_host or seq_lens_host)
    max_blocks_per_seq = math.ceil(max(seq_lens_host) / block_size)
    num_blocks = sum(math.ceil(seq_len / block_size) for seq_len in seq_lens_host)

    compressed_kv = torch.randn(
        num_tokens, head_dim, dtype=torch.bfloat16, device=device
    )

    # Randomize physical pages so the test covers block-table translation.
    # Keep padded block-table entries invalid to catch accidental reads.
    physical_blocks = torch.randperm(num_blocks, device=device)
    block_table = torch.full(
        (num_reqs, max_blocks_per_seq), int(-1e6), dtype=torch.int32, device=device
    )
    start = 0
    for req_id, seq_len in enumerate(seq_lens_host):
        num_req_blocks = math.ceil(seq_len / block_size)
        req_blocks = physical_blocks[start : start + num_req_blocks]
        block_table[req_id, :num_req_blocks] = req_blocks
        start += num_req_blocks

    # Build slot_mapping for quantize_and_insert_k_cache.
    slot_mapping = torch.empty(num_tokens, dtype=torch.int64, device=device)
    start = 0
    for req_id, seq_len in enumerate(seq_lens_host):
        logical_pos = torch.arange(seq_len, dtype=torch.int64, device=device)
        block_idx = block_table[req_id, logical_pos // block_size].to(torch.int64)
        token_slots = block_idx * block_size + logical_pos % block_size
        slot_mapping[start : start + seq_len] = token_slots
        start += seq_len

    # Insert compressed K into the paged cache layout used by the gather op.
    k_cache = torch.empty(
        num_blocks, block_size, head_bytes, dtype=torch.uint8, device=device
    )
    k_cache_2d = k_cache.view(num_blocks, -1)
    quantize_and_insert_k_cache(compressed_kv, k_cache_2d, slot_mapping, block_size)

    out_shape = (num_reqs, offset + max_gather_len + 3, head_dim)
    ref_out = torch.empty(out_shape, dtype=torch.bfloat16, device=device)
    actual_out = torch.empty_like(ref_out)
    seq_lens = torch.tensor(seq_lens_host, dtype=torch.int32, device=device)
    gather_lens = (
        torch.tensor(gather_lens_host, dtype=torch.int32, device=device)
        if gather_lens_host is not None
        else None
    )

    # Compare production gather against a PyTorch reference for valid output rows.
    _dequantize_and_gather_k_cache_reference(
        ref_out, k_cache, seq_lens, gather_lens, block_table, block_size, offset
    )
    dequantize_and_gather_k_cache(
        actual_out, k_cache, seq_lens, gather_lens, block_table, block_size, offset
    )
    torch.accelerator.synchronize()

    # only check non-padded content
    for req_id, seq_len in enumerate(seq_lens_host):
        gather_len = (
            gather_lens_host[req_id] if gather_lens_host is not None else seq_len
        )
        actual = actual_out[req_id, offset : offset + gather_len]
        expected = ref_out[req_id, offset : offset + gather_len]
        torch.testing.assert_close(actual, expected, rtol=0, atol=0)


# ── Test C: Indexer path ────────────────────────────────────────────────────


@pytest.mark.parametrize("num_tokens", [1, 4, 8, 17])
@pytest.mark.parametrize("block_size", [16, 64])
def test_indexer_quant_cache_roundtrip(num_tokens: int, block_size: int):
    """K → indexer_k_quant_and_cache → cp_gather_indexer_k_quant_cache
    → manual dequant → compare against original."""
    HEAD_DIM = 128
    QUANT_BLOCK_SIZE = 128
    # cache_stride = head_dim + (head_dim * 4 / quant_block_size) = 128 + 4 = 132
    CACHE_STRIDE = HEAD_DIM + HEAD_DIM * 4 // QUANT_BLOCK_SIZE

    num_blocks = (num_tokens + block_size - 1) // block_size + 1
    device = "cuda"

    # Random K (simulates compressor output for indexer)
    k = torch.randn(num_tokens, HEAD_DIM, dtype=torch.bfloat16, device=device)

    # ── Quant + insert ──────────────────────────────────────────────────
    kv_cache = torch.zeros(
        num_blocks, block_size, CACHE_STRIDE, dtype=torch.uint8, device=device
    )
    slot_mapping = torch.arange(num_tokens, dtype=torch.int64, device=device)

    ops.indexer_k_quant_and_cache(k, kv_cache, slot_mapping, QUANT_BLOCK_SIZE, "ue8m0")

    # ── Gather ──────────────────────────────────────────────────────────
    max_blocks_per_seq = num_blocks
    block_table = torch.arange(
        max_blocks_per_seq, dtype=torch.int32, device=device
    ).unsqueeze(0)
    cu_seq_lens = torch.tensor([0, num_tokens], dtype=torch.int32, device=device)

    # dst_k: [total_seq_len, head_dim] as uint8 (raw FP8 bytes)
    dst_k = torch.zeros(num_tokens, HEAD_DIM, dtype=torch.uint8, device=device)
    # dst_scale: [total_seq_len, head_dim/quant_block*4] as uint8 (raw float32 bytes)
    num_scale_bytes = HEAD_DIM * 4 // QUANT_BLOCK_SIZE  # 4
    dst_scale = torch.zeros(
        num_tokens, num_scale_bytes, dtype=torch.uint8, device=device
    )

    ops.cp_gather_indexer_k_quant_cache(
        kv_cache, dst_k, dst_scale, block_table, cu_seq_lens
    )

    # ── Manual dequant ──────────────────────────────────────────────────
    k_fp8 = dst_k.view(torch.float8_e4m3fn).float()  # [num_tokens, 128]
    scale = dst_scale.view(torch.float32)  # [num_tokens, 1]
    k_recovered = k_fp8 * scale  # [num_tokens, 128]

    # ── Compare ─────────────────────────────────────────────────────────
    diff = (k_recovered - k.float()).abs()
    k_abs = k.float().abs()

    for t in range(num_tokens):
        amax = k_abs[t].max().clamp(min=1e-4).item()
        # UE8M0: scale = 2^ceil(log2(amax / 448))
        exponent = math.ceil(math.log2(amax / 448.0))
        ue8m0_scale = 2.0**exponent
        # FP8 e4m3 (3-bit mantissa): worst-case error = 16 * scale
        max_allowed = 16.0 * ue8m0_scale
        token_diff = diff[t].max().item()
        assert token_diff <= max_allowed, (
            f"Token {t} diff {token_diff} exceeds max_allowed "
            f"{max_allowed} (scale={ue8m0_scale})"
        )


def test_indexer_gather_accepts_upper_bound_output():
    """Gather only exact cu_seq_lens even when dst is over-allocated."""
    head_dim = 128
    quant_block_size = 128
    cache_stride = head_dim + head_dim * 4 // quant_block_size
    valid_tokens = 9
    upper_bound_tokens = 13
    block_size = 16
    num_seqs = 3
    num_blocks = num_seqs
    sentinel = 123
    device = "cuda"

    k = torch.randn(valid_tokens, head_dim, dtype=torch.bfloat16, device=device)
    kv_cache = torch.zeros(
        num_blocks, block_size, cache_stride, dtype=torch.uint8, device=device
    )
    slot_mapping = torch.tensor(
        [0, 1, 2, 16, 17, 18, 32, 33, 34], dtype=torch.int64, device=device
    )
    ops.indexer_k_quant_and_cache(k, kv_cache, slot_mapping, quant_block_size, "ue8m0")

    block_table = torch.arange(num_blocks, dtype=torch.int32, device=device).unsqueeze(
        1
    )
    cu_seq_lens = torch.tensor([0, 3, 6, 9], dtype=torch.int32, device=device)
    dst_k = torch.full(
        (upper_bound_tokens, head_dim), sentinel, dtype=torch.uint8, device=device
    )
    num_scale_bytes = head_dim * 4 // quant_block_size
    dst_scale = torch.full(
        (upper_bound_tokens, num_scale_bytes),
        sentinel,
        dtype=torch.uint8,
        device=device,
    )

    ops.cp_gather_indexer_k_quant_cache(
        kv_cache, dst_k, dst_scale, block_table, cu_seq_lens
    )

    if current_platform.is_rocm():
        triton_kv_cache = torch.zeros_like(kv_cache)
        indexer_k_quant_and_cache_triton(
            k,
            triton_kv_cache,
            slot_mapping,
            quant_block_size,
            "ue8m0",
        )
        triton_dst_k = torch.full_like(dst_k, sentinel)
        triton_dst_scale = torch.full_like(dst_scale, sentinel)
        token_to_seq = torch.cat(
            (
                torch.repeat_interleave(
                    torch.arange(num_seqs, dtype=torch.int32, device=device), 3
                ),
                torch.full(
                    (upper_bound_tokens - valid_tokens,),
                    -1,
                    dtype=torch.int32,
                    device=device,
                ),
            )
        )
        cp_gather_indexer_k_quant_cache_triton(
            triton_kv_cache,
            triton_dst_k.view(current_platform.fp8_dtype()),
            triton_dst_scale,
            block_table,
            cu_seq_lens,
            token_to_seq,
        )
    torch.accelerator.synchronize()

    if current_platform.is_rocm():
        triton_recovered = triton_dst_k[:valid_tokens].view(
            current_platform.fp8_dtype()
        ).float() * triton_dst_scale[:valid_tokens].view(torch.float32)
        triton_error = (triton_recovered - k.float()).abs().amax(dim=1)
        max_triton_error = (
            16.0 * triton_dst_scale[:valid_tokens].view(torch.float32).flatten()
        )
        assert torch.all(triton_error <= max_triton_error)
        assert torch.all(triton_dst_k[valid_tokens:] == sentinel)
        assert torch.all(triton_dst_scale[valid_tokens:] == sentinel)
    k_recovered = dst_k[:valid_tokens].view(torch.float8_e4m3fn).float() * dst_scale[
        :valid_tokens
    ].view(torch.float32)
    diff = (k_recovered - k.float()).abs()
    max_allowed = (16.0 * dst_scale[:valid_tokens].view(torch.float32).max()).item()
    assert diff.max().item() <= max_allowed
    assert torch.all(dst_k[valid_tokens:] == sentinel)
    assert torch.all(dst_scale[valid_tokens:] == sentinel)


# ── Test D: DeepseekV4 attention with values at different magnitudes ───────────


def test_deepseek_v4_quant_magnitude_range():
    """Test that quantization handles a range of magnitudes correctly."""
    HEAD_DIM = 512
    NOPE_DIM = 448
    HEAD_BYTES = 584
    block_size = 16
    num_tokens = 4
    num_blocks = 2
    device = "cuda"

    # Create inputs with varying magnitudes: small, medium, large
    compressed_kv = torch.zeros(
        num_tokens, HEAD_DIM, dtype=torch.bfloat16, device=device
    )
    compressed_kv[0] = 0.001  # very small
    compressed_kv[1] = 1.0  # unit scale
    compressed_kv[2] = 100.0  # large
    compressed_kv[3] = torch.randn(HEAD_DIM, dtype=torch.bfloat16, device=device)

    k_cache = torch.zeros(
        num_blocks, block_size, HEAD_BYTES, dtype=torch.uint8, device=device
    )
    slot_mapping = torch.arange(num_tokens, dtype=torch.int64, device=device)

    quantize_and_insert_k_cache(
        compressed_kv, k_cache.view(num_blocks, -1), slot_mapping, block_size
    )

    out = torch.zeros(1, num_tokens, HEAD_DIM, dtype=torch.bfloat16, device=device)
    seq_lens = torch.tensor([num_tokens], dtype=torch.int32, device=device)
    block_table = torch.arange(num_blocks, dtype=torch.int32, device=device).unsqueeze(
        0
    )

    dequantize_and_gather_k_cache(
        out, k_cache, seq_lens, None, block_table, block_size, offset=0
    )

    recovered = out[0, :num_tokens]

    # RoPE portion must be exact
    rope_diff = (recovered[:, NOPE_DIM:] - compressed_kv[:, NOPE_DIM:]).abs().max()
    assert rope_diff.item() == 0.0, f"RoPE diff {rope_diff.item()}"

    # NoPE: relative error should be reasonable
    for t in range(num_tokens):
        orig = compressed_kv[t, :NOPE_DIM].float()
        recv = recovered[t, :NOPE_DIM].float()
        abs_diff = (recv - orig).abs().max().item()
        magnitude = orig.abs().max().item()
        if magnitude > 0.01:
            rel_err = abs_diff / magnitude
            assert rel_err < 0.15, (
                f"Token {t}: rel_err={rel_err:.4f}, abs_diff={abs_diff:.6f}, "
                f"magnitude={magnitude:.4f}"
            )


# ── Test E: Indexer fused K-cache insert (Triton kernels) ────────────────────
#
# Both kernels share the same Triton signature; use_fp4 selects between them.
# Full pipeline: state-cache gather → softmax-weighted compress → RMSNorm →
#   GPT-J RoPE → quant (MXFP4 or FP8) → paged cache insert.


def _reference_kv_compress_norm_rope(
    state_cache: torch.Tensor,
    block_table: torch.Tensor,
    positions: torch.Tensor,
    rms_weight: torch.Tensor,
    cos_sin_cache: torch.Tensor,
    compress_ratio: int = 1,
    overlap: int = 0,
    use_fp4: bool = False,
    rms_eps: float = 1e-6,
    fp8_max: float = 448.0,
    return_full_cache: bool = False,
):
    """Compress → RMSNorm → GPT-J RoPE → quantize.

    Gathers (1+overlap)*compress_ratio state entries per output token, applies
    per-element softmax over the scores, and computes the weighted kv sum.
    Returns (quantized_values, scale) matching the kernel's output layout.
    """
    device = state_cache.device
    head_dim = rms_weight.shape[0]
    rope_dim = cos_sin_cache.shape[-1]
    state_block_size = state_cache.shape[1]
    state_width = state_cache.shape[-1] // 2
    nope_dim = head_dim - rope_dim
    total = (1 + overlap) * compress_ratio
    results = []
    for pos in positions.tolist():
        src = torch.arange(pos - total + 1, pos + 1, dtype=torch.int64, device=device)
        valid = src >= 0
        idx = src.clamp(min=0)
        pages = block_table[0, idx // state_block_size]
        offsets = idx % state_block_size
        raw = state_cache[pages, offsets].float()  # [total, state_dim]

        # Group 0 (tokens 0..cr-1):   kv[:H],   score[SW:SW+H]
        # Group 1 (tokens cr..2cr-1): kv[H:2H], score[SW+H:SW+2H]
        if overlap:
            sw = state_width
            g0_kv = raw[:compress_ratio, :head_dim]
            g1_kv = raw[compress_ratio:, head_dim : 2 * head_dim]
            g0_scores = raw[:compress_ratio, sw : sw + head_dim]
            g1_scores = raw[compress_ratio:, sw + head_dim : sw + 2 * head_dim]
            kv = torch.cat([g0_kv, g1_kv])
            scores = torch.cat([g0_scores, g1_scores])
        else:
            kv = raw[:, :head_dim]
            scores = raw[:, state_width : state_width + head_dim]

        scores[~valid] = float("-inf")
        kv[~valid] = 0.0
        weights = torch.softmax(scores, dim=0)
        compressed = (kv * weights).sum(dim=0)  # [H]
        var = (compressed * compressed).mean()
        normed = compressed * torch.rsqrt(var + rms_eps) * rms_weight.float()
        compressed_pos = (pos // compress_ratio) * compress_ratio
        cos, sin = cos_sin_cache[compressed_pos].float().chunk(2)
        nope, rope = normed.split([nope_dim, rope_dim])
        rope = torch.stack(
            [rope[0::2] * cos - rope[1::2] * sin, rope[1::2] * cos + rope[0::2] * sin],
            dim=-1,
        ).reshape(rope_dim)
        results.append(torch.cat([nope, rope]).to(state_cache.dtype))
    result = torch.stack(results)

    if return_full_cache:
        # Contiguous 512-wide bf16 row (nope unrotated + rope rotated), matching
        # the FlashInfer full-cache layout before any per-tensor fp8 quant. The
        # kernel rounds the fp32 result to bf16 once at the store.
        return result.to(torch.bfloat16)

    if use_fp4:
        return quantize_to_mxfp4(result)
    else:
        pairs = [
            _ue8m0_reference(result[t], head_dim, fp8_max) for t in range(len(result))
        ]
        quants, scales = zip(*pairs)
        return torch.stack(quants), torch.cat(scales)


@pytest.mark.parametrize("num_tokens", [1, 7, 32])
@pytest.mark.parametrize("kv_block_size", [16, 32])
@pytest.mark.parametrize(
    "use_fp4",
    [
        False,
        pytest.param(
            True,
            marks=pytest.mark.skipif(
                not (
                    current_platform.is_cuda()
                    and current_platform.is_device_capability_family(100)
                ),
                reason="MXFP4 indexer cache requires an SM100-family GPU",
            ),
        ),
    ],
)
def test_fused_kv_insert_indexer(num_tokens: int, kv_block_size: int, use_fp4: bool):
    """Fused K compress+norm+rope+quant+insert for the indexer KV cache."""
    HEAD_DIM = 128
    ROPE_DIM = 64
    BLOCK_SIZE = 16
    RMS_EPS = 1e-6
    FP8_MAX = 448.0

    device = "cuda"
    torch.manual_seed(42)
    compress_ratio = 4

    if use_fp4:
        TOKEN_STRIDE = HEAD_DIM // 2  # packed nibbles: 64 bytes
        SCALE_DIM = HEAD_DIM // 32  # ue8m0 bytes: 4
        QUANT_BLOCK = 32
        kernel = _fused_kv_compress_norm_rope_insert_indexer_mxfp4_attn
    else:
        TOKEN_STRIDE = HEAD_DIM  # FP8 bytes: 128
        SCALE_DIM = 4  # 1 float32: 4 bytes
        QUANT_BLOCK = HEAD_DIM
        kernel = _fused_kv_compress_norm_rope_insert_indexer_attn

    # overlap=1 whenever compress_ratio==4, matching DeepseekCompressor logic.
    overlap = 1 if compress_ratio == 4 else 0
    coff = 1 + overlap  # multiplier for state_dim per entry

    num_pages = (compress_ratio * num_tokens - 1) // BLOCK_SIZE + 2
    state_cache = torch.randn(
        num_pages,
        BLOCK_SIZE,
        2 * coff * HEAD_DIM,  # kv_state + score_state, each coff*HEAD_DIM wide
        dtype=torch.bfloat16,
        device=device,
    )
    block_table = torch.arange(num_pages, dtype=torch.int32, device=device).unsqueeze(0)
    token_to_req = torch.zeros(num_tokens, dtype=torch.int32, device=device)
    slot_mapping = torch.arange(num_tokens, dtype=torch.int64, device=device)
    positions = torch.arange(
        compress_ratio - 1,
        compress_ratio * num_tokens,
        compress_ratio,
        dtype=torch.int64,
        device=device,
    )
    rms_weight = torch.randn(HEAD_DIM, dtype=torch.bfloat16, device=device)
    cos_sin_cache = torch.randn(compress_ratio * num_tokens, ROPE_DIM, device=device)

    kv_n_blocks = (num_tokens + kv_block_size - 1) // kv_block_size + 1
    kv_cache = torch.zeros(
        kv_n_blocks,
        kv_block_size * (TOKEN_STRIDE + SCALE_DIM),
        dtype=torch.uint8,
        device=device,
    )

    kernel[(num_tokens,)](
        state_cache,
        state_cache.stride(0),
        state_cache.stride(1),
        token_to_req,
        positions,
        slot_mapping,
        block_table,
        block_table.stride(0),
        BLOCK_SIZE,
        rms_weight,
        RMS_EPS,
        cos_sin_cache,
        cos_sin_cache.stride(0),
        kv_cache,
        slot_mapping,
        kv_block_size,
        HEAD_SIZE=HEAD_DIM,
        TRITON_BLOCK_SIZE=HEAD_DIM,
        STATE_WIDTH=coff * HEAD_DIM,
        COMPRESS_RATIO=compress_ratio,
        OVERLAP=overlap,
        ROPE_HEAD_DIM=ROPE_DIM,
        FP8_MAX=FP8_MAX,
        QUANT_BLOCK=QUANT_BLOCK,
        TOKEN_STRIDE=TOKEN_STRIDE,
        SCALE_DIM=SCALE_DIM,
        KV_BLOCK_STRIDE=kv_cache.stride(0),
        num_warps=1,
    )

    k_quant, scale = _reference_kv_compress_norm_rope(
        state_cache,
        block_table,
        positions,
        rms_weight,
        cos_sin_cache,
        compress_ratio,
        overlap,
        use_fp4,
        rms_eps=RMS_EPS,
        fp8_max=FP8_MAX,
    )

    if use_fp4:
        for i in range(num_tokens):
            blk, pos = i // kv_block_size, i % kv_block_size
            val_off = pos * TOKEN_STRIDE
            fp4_actual = kv_cache[blk, val_off : val_off + TOKEN_STRIDE]
            assert torch.equal(k_quant[i], fp4_actual), (
                f"token {i}: packed nibbles differ, "
                f"{(k_quant[i] != fp4_actual).sum()} "
                f"/ {TOKEN_STRIDE}"
            )

            scale_off = kv_block_size * TOKEN_STRIDE + pos * SCALE_DIM
            scale_actual = kv_cache[blk, scale_off : scale_off + SCALE_DIM]
            assert torch.equal(scale_actual, scale[i]), (
                f"token {i}: ue8m0 {scale_actual.tolist()} != {scale[i].tolist()}"
            )

    else:
        k_quant = k_quant.view(torch.uint8)
        for i in range(num_tokens):
            blk, pos = i // kv_block_size, i % kv_block_size
            val_off = pos * TOKEN_STRIDE
            assert torch.equal(
                k_quant[i], kv_cache[blk, val_off : val_off + TOKEN_STRIDE]
            ), f"token {i}: FP8 bytes differ"

            scale_off = kv_block_size * TOKEN_STRIDE + pos * SCALE_DIM
            actual_scale = kv_cache[blk, scale_off : scale_off + SCALE_DIM].view(
                torch.float32
            )
            assert torch.equal(actual_scale, scale[i : i + 1]), (
                f"token {i}: scale {actual_scale.item()} != {scale[i].item()}"
            )


@pytest.mark.parametrize("compress_ratio", [4, 128])
@pytest.mark.parametrize("store_fp8", [False, True])
def test_cutedsl_full_cache_store(compress_ratio: int, store_fp8: bool):
    """CuTeDSL compressor full-cache (FlashInfer) store parity for head=512.

    Exercises the contiguous bf16 / per-tensor fp8 store branch of both the C4
    fused kernel and the C128 split kernel against the PyTorch reference.
    """
    cutedsl = pytest.importorskip("cutlass")  # noqa: F841
    from vllm.models.deepseek_v4.nvidia.ops.sparse_attn_compress_cutedsl import (
        fused_kv_compress_norm_rope_insert_sparse_attn_cutedsl,
        split_kv_compress_norm_rope_insert_sparse_attn_cutedsl,
    )

    HEAD_DIM = 512
    ROPE_DIM = 64
    RMS_EPS = 1e-6
    FP8_MAX = 448.0
    # C128 compress (Block8 kernel) requires state-cache block_size=8; C4 uses 16.
    BLOCK_SIZE = 8 if compress_ratio == 128 else 16
    KV_BLOCK_SIZE = 64
    device = "cuda"
    torch.manual_seed(7)

    overlap = 1 if compress_ratio == 4 else 0
    coff = 1 + overlap
    num_tokens = 8

    num_pages = (compress_ratio * num_tokens - 1) // BLOCK_SIZE + 2
    # The production CompressorStateCache is fp32.
    state_cache = torch.randn(
        num_pages, BLOCK_SIZE, 2 * coff * HEAD_DIM, dtype=torch.float32, device=device
    )
    block_table = torch.arange(num_pages, dtype=torch.int32, device=device).unsqueeze(0)
    token_to_req = torch.zeros(num_tokens, dtype=torch.int32, device=device)
    slot_mapping = torch.arange(num_tokens, dtype=torch.int64, device=device)
    positions = torch.arange(
        compress_ratio - 1,
        compress_ratio * num_tokens,
        compress_ratio,
        dtype=torch.int64,
        device=device,
    )
    rms_weight = torch.randn(HEAD_DIM, dtype=torch.bfloat16, device=device)
    cos_sin_cache = torch.randn(
        compress_ratio * num_tokens, ROPE_DIM, dtype=torch.float32, device=device
    )

    dtype = torch.float8_e4m3fn if store_fp8 else torch.bfloat16
    kv_n_blocks = (num_tokens + KV_BLOCK_SIZE - 1) // KV_BLOCK_SIZE + 1
    k_cache = torch.zeros(
        kv_n_blocks, KV_BLOCK_SIZE, HEAD_DIM, dtype=dtype, device=device
    )
    fp8_scale = torch.tensor(
        [0.5 if store_fp8 else 1.0], dtype=torch.float32, device=device
    )

    if compress_ratio == 4:
        fused_kv_compress_norm_rope_insert_sparse_attn_cutedsl(
            state_cache,
            token_to_req,
            positions,
            slot_mapping,
            block_table,
            BLOCK_SIZE,
            rms_weight,
            RMS_EPS,
            cos_sin_cache,
            k_cache,
            slot_mapping,
            KV_BLOCK_SIZE,
            k_cache.stride(0),
            head_size=HEAD_DIM,
            state_width=coff * HEAD_DIM,
            rope_head_dim=ROPE_DIM,
            fp8_max=FP8_MAX,
            quant_block=64,
            token_stride=576,
            scale_dim=8,
            compress_ratio=compress_ratio,
            overlap=True,
            store_full_kv=True,
            store_full_fp8=store_fp8,
            fp8_scale=fp8_scale,
        )
    else:
        compressed_kv = torch.empty(
            (num_tokens, HEAD_DIM), dtype=torch.float32, device=device
        )
        split_kv_compress_norm_rope_insert_sparse_attn_cutedsl(
            state_cache,
            token_to_req,
            positions,
            slot_mapping,
            block_table,
            BLOCK_SIZE,
            compressed_kv,
            rms_weight,
            RMS_EPS,
            cos_sin_cache,
            k_cache,
            slot_mapping,
            KV_BLOCK_SIZE,
            k_cache.stride(0),
            head_size=HEAD_DIM,
            state_width=coff * HEAD_DIM,
            rope_head_dim=ROPE_DIM,
            fp8_max=FP8_MAX,
            quant_block=64,
            token_stride=576,
            scale_dim=8,
            compress_ratio=compress_ratio,
            overlap=bool(overlap),
            store_full_kv=True,
            store_full_fp8=store_fp8,
            fp8_scale=fp8_scale,
        )

    ref = _reference_kv_compress_norm_rope(
        state_cache,
        block_table,
        positions,
        rms_weight,
        cos_sin_cache,
        compress_ratio,
        overlap,
        rms_eps=RMS_EPS,
        return_full_cache=True,
    )  # [num_tokens, HEAD_DIM] bf16

    actual = torch.stack(
        [k_cache[i // KV_BLOCK_SIZE, i % KV_BLOCK_SIZE] for i in range(num_tokens)]
    )
    if store_fp8:
        ref_fp8 = torch.clamp(ref.float() / fp8_scale, -FP8_MAX, FP8_MAX).to(
            torch.float8_e4m3fn
        )
        torch.testing.assert_close(actual.float(), ref_fp8.float(), rtol=0.0, atol=0.3)
    else:
        torch.testing.assert_close(actual.float(), ref.float(), rtol=3e-2, atol=3e-2)


# ── Test F: DeepseekV4 Attention two-stage split compressor (Triton) ─────────
#
# Same full pipeline as Test E (state-cache gather -> softmax-weighted compress
# -> RMSNorm -> GPT-J RoPE -> quant -> paged insert), but for the head=512
# fp8_ds_mla layout via the two-stage split


@pytest.mark.skipif(
    not current_platform.is_rocm(),
    reason="two-stage split compressor is only enabled for ROCm at the moment",
)
@pytest.mark.parametrize("num_tokens", [1, 4, 8, 17])
@pytest.mark.parametrize("kv_block_size", [16, 64])
def test_fused_kv_insert_split(num_tokens: int, kv_block_size: int):
    """Two-stage split compress+norm+rope+quant+insert for the head=512 KV cache."""
    HEAD_DIM = 512
    NOPE_DIM = 448
    ROPE_DIM = 64
    HEAD_BYTES = 584  # 448 fp8 + 128 bf16 + 8 uint8 scale
    RMS_EPS = 1e-6
    FP8_MAX = 448.0
    QUANT_BLOCK = 64
    TOKEN_STRIDE = 576
    SCALE_DIM = 8
    STATE_BLOCK_SIZE = 8  # CompressorStateCache block_size for cr=128

    device = "cuda"
    torch.manual_seed(42)
    compress_ratio = 128
    overlap = 0  # no overlap for cr=128
    coff = 1 + overlap

    num_pages = (compress_ratio * num_tokens - 1) // STATE_BLOCK_SIZE + 2
    state_cache = torch.randn(
        num_pages,
        STATE_BLOCK_SIZE,
        2 * coff * HEAD_DIM,  # kv_state + score_state
        dtype=torch.float32,
        device=device,
    )
    block_table = torch.arange(num_pages, dtype=torch.int32, device=device).unsqueeze(0)
    token_to_req = torch.zeros(num_tokens, dtype=torch.int32, device=device)
    slot_mapping = torch.arange(num_tokens, dtype=torch.int64, device=device)
    positions = torch.arange(
        compress_ratio - 1,
        compress_ratio * num_tokens,
        compress_ratio,
        dtype=torch.int64,
        device=device,
    )
    rms_weight = torch.randn(HEAD_DIM, dtype=torch.bfloat16, device=device)
    cos_sin_cache = torch.randn(
        compress_ratio * num_tokens, ROPE_DIM, dtype=torch.float32, device=device
    )

    kv_n_blocks = (num_tokens + kv_block_size - 1) // kv_block_size + 1
    kv_cache = torch.zeros(
        kv_n_blocks, kv_block_size, HEAD_BYTES, dtype=torch.uint8, device=device
    )
    compress_scratch = torch.empty(
        num_tokens, HEAD_DIM, dtype=torch.float32, device=device
    )

    _launch_two_stage_sparse_attn_compressor(
        state_cache,
        token_to_req,
        positions,
        slot_mapping,
        block_table,
        STATE_BLOCK_SIZE,
        coff * HEAD_DIM,
        compress_ratio,
        cos_sin_cache,
        kv_cache,
        slot_mapping,
        rms_weight,
        RMS_EPS,
        QUANT_BLOCK,
        TOKEN_STRIDE,
        SCALE_DIM,
        HEAD_DIM,
        ROPE_DIM,
        num_tokens,
        compress_scratch,
    )

    # PyTorch reference: compress -> RMSNorm -> GPT-J RoPE (pre-quant bf16 row).
    ref = _reference_kv_compress_norm_rope(
        state_cache,
        block_table,
        positions,
        rms_weight,
        cos_sin_cache,
        compress_ratio,
        overlap,
        rms_eps=RMS_EPS,
        fp8_max=FP8_MAX,
        return_full_cache=True,
    )  # [num_tokens, HEAD_DIM] bf16

    # Dequant + gather the fp8_ds_mla cache back to bf16 (Test B op).
    out = torch.zeros(1, num_tokens, HEAD_DIM, dtype=torch.bfloat16, device=device)
    seq_lens = torch.tensor([num_tokens], dtype=torch.int32, device=device)
    gather_block_table = torch.arange(
        kv_n_blocks, dtype=torch.int32, device=device
    ).unsqueeze(0)
    dequantize_and_gather_k_cache(
        out, kv_cache, seq_lens, None, gather_block_table, kv_block_size, offset=0
    )
    recovered = out[0, :num_tokens]

    # NoPE (first 448): FP8 quantized, expect UE8M0 error (same bound as Test A).
    nope_diff = (recovered[:, :NOPE_DIM].float() - ref[:, :NOPE_DIM].float()).abs()
    for t in range(num_tokens):
        _, scales = _ue8m0_reference(ref[t, :NOPE_DIM].float(), QUANT_BLOCK, FP8_MAX)
        max_allowed = 16.0 * scales.max().item()
        token_diff = nope_diff[t].max().item()
        assert token_diff <= max_allowed, (
            f"Token {t} nope diff {token_diff} exceeds max_allowed "
            f"{max_allowed} (scale={scales.max().item()})"
        )

    # RoPE (last 64): stored as bf16. The kernel recomputes the rotation, so it
    # is bf16-close to the reference rather than bit-exact (cf. test_cutedsl).
    torch.testing.assert_close(recovered[:, NOPE_DIM:], ref[:, NOPE_DIM:])


@pytest.mark.skipif(not current_platform.is_rocm(), reason="ROCm stream coverage")
def test_v41_rocm_csa2_full_pipeline():
    """_forward_csa2_full: one fork, then K write and q-side after the join.

    The default chain must run first on the current stream, the compressor
    chain on aux 0 and the indexer weights projection on aux 1; the join must
    publish both cache writes before _produce_k / forward_q run.
    """
    from vllm.forward_context import ForwardContext, override_forward_context
    from vllm.models.deepseek_v41.amd.rocm import DeepseekV41ROCMAiterMLAAttention
    from vllm.models.deepseek_v41.compressor import DeepseekCompressor

    torch.manual_seed(43)
    num_tokens = 19
    hidden_states = torch.randn(num_tokens, 1024, device="cuda")
    positions = torch.arange(7, 26, device="cuda")
    qr_kv = torch.zeros(num_tokens, 1536, device="cuda")
    qr = torch.zeros(num_tokens, 4, device="cuda")
    kv = torch.zeros(num_tokens, 512, device="cuda")
    q = torch.zeros(num_tokens, 512, device="cuda")
    latent = torch.ones(num_tokens, 512, dtype=torch.bfloat16, device="cuda")
    weights = torch.randn(num_tokens, 32, device="cuda")
    index_q = torch.full((num_tokens, 32, 128), 2.0, device="cuda")
    index_q_scale = torch.full((num_tokens, 32), 3.0, device="cuda")
    index_weights_out = torch.full((num_tokens, 32), 4.0, device="cuda")
    main_cache = torch.full((1, 128, 584), 165, dtype=torch.uint8, device="cuda")
    index_cache = torch.full((1, 128, 132), 165, dtype=torch.uint8, device="cuda")

    order: list[str] = []
    streams: dict[str, torch.cuda.Stream] = {}

    compressor = DeepseekCompressor.__new__(DeepseekCompressor)
    torch.nn.Module.__init__(compressor)
    compressor.fused_wkv_wgate = SimpleNamespace(
        weight=torch.randn(1024, 1024, device="cuda")
    )

    def compressor_forward(kv_score, positions_):
        order.append("compressor")
        streams["compressor"] = torch.cuda.current_stream()
        return latent

    def compressor_insert(latent_, positions_, rotary_emb):
        order.append("insert_cache")
        streams["insert_cache"] = torch.cuda.current_stream()
        main_cache.fill_(5)

    compressor.forward = compressor_forward
    compressor.insert_cache = compressor_insert

    def weights_proj(hidden_states_):
        order.append("weights_proj")
        streams["weights_proj"] = torch.cuda.current_stream()
        return weights, None

    def produce_k(latent_, positions_, rotary_emb):
        order.append("produce_k")
        streams["produce_k"] = torch.cuda.current_stream()
        index_cache.fill_(7)

    def forward_q(qr_, qr_scale, indexer_weights, positions_, rotary_emb):
        order.append("forward_q")
        streams["forward_q"] = torch.cuda.current_stream()
        return index_q, index_q_scale, index_weights_out

    indexer = SimpleNamespace(
        weights_proj=weights_proj,
        _produce_k=produce_k,
        forward_q=forward_q,
    )
    rotary = SimpleNamespace(cos_sin_cache=torch.randn(32, 64, device="cuda"))
    aux_streams = [torch.cuda.Stream(), torch.cuda.Stream()]

    def wq_b(qr_, qr_scale):
        order.append("wq_b")
        streams["wq_b"] = torch.cuda.current_stream()
        return q

    def qnorm_insert(q_, kv_, positions_, attn_metadata):
        order.append("qnorm_insert")
        streams["qnorm_insert"] = torch.cuda.current_stream()
        return q_

    attention = SimpleNamespace(
        compressor=compressor,
        indexer=indexer,
        aux_stream_list=aux_streams,
        ln_events=[torch.cuda.Event() for _ in range(4)],
        rotary_emb=rotary,
        indexer_rotary_emb=rotary,
        n_local_heads=1,
        head_dim=512,
        _fused_wqa_wkv_gemm=lambda hidden_states_: qr_kv,
        _split_qkv_and_norm=lambda qr_kv_: (qr, None, kv),
        _wq_b_proj=wq_b,
        _fused_qnorm_rope_kv_insert=qnorm_insert,
    )

    context = ForwardContext({}, {}, {})
    with override_forward_context(context):
        result = DeepseekV41ROCMAiterMLAAttention._forward_csa2_full(
            attention, hidden_states, positions
        )

    q_out, kv_out, index_q_out, index_q_scale_out, index_weights_out_out = result
    assert torch.equal(
        q_out, q.view(num_tokens, attention.n_local_heads, attention.head_dim)
    )
    assert torch.equal(kv_out, kv)
    assert torch.equal(index_q_out, index_q)
    assert torch.equal(index_q_scale_out, index_q_scale)
    assert torch.equal(index_weights_out_out, index_weights_out)
    # Join publishes the aux-0 cache insert and orders the post-join K write.
    assert torch.equal(main_cache, torch.full_like(main_cache, 5))
    assert torch.equal(index_cache, torch.full_like(index_cache, 7))

    assert order == [
        "wq_b",
        "qnorm_insert",
        "compressor",
        "insert_cache",
        "weights_proj",
        "produce_k",
        "forward_q",
    ]
    current = torch.cuda.current_stream()
    for name in ("wq_b", "qnorm_insert", "produce_k", "forward_q"):
        assert streams[name] == current, name
    for name in ("compressor", "insert_cache"):
        assert streams[name] == aux_streams[0], name
    assert streams["weights_proj"] == aux_streams[1]


@pytest.mark.skipif(not current_platform.is_rocm(), reason="ROCm stream coverage")
def test_v41_rocm_csa2_reindex_pipeline():
    """_forward_csa2_reindex: serial projections, then SWA q path vs q-side."""
    from vllm.forward_context import ForwardContext, override_forward_context
    from vllm.models.deepseek_v41.amd.rocm import DeepseekV41ROCMAiterMLAAttention

    torch.manual_seed(45)
    num_tokens = 19
    hidden_states = torch.randn(num_tokens, 1024, device="cuda")
    positions = torch.arange(7, 26, device="cuda")
    qr_kv = torch.zeros(num_tokens, 1536, device="cuda")
    qr = torch.zeros(num_tokens, 4, device="cuda")
    kv = torch.zeros(num_tokens, 512, device="cuda")
    q = torch.zeros(num_tokens, 512, device="cuda")
    indexer_weights = torch.randn(num_tokens, 32, device="cuda")
    index_q = torch.full((num_tokens, 32, 128), 2.0, device="cuda")
    index_q_scale = torch.full((num_tokens, 32), 3.0, device="cuda")
    index_weights_out = torch.full((num_tokens, 32), 4.0, device="cuda")

    order: list[str] = []
    streams: dict[str, torch.cuda.Stream] = {}

    def forward_q(qr_, qr_scale, indexer_weights_, positions_, rotary_emb):
        order.append("forward_q")
        streams["forward_q"] = torch.cuda.current_stream()
        return index_q, index_q_scale, index_weights_out

    indexer = SimpleNamespace(forward_q=forward_q)
    rotary = SimpleNamespace(cos_sin_cache=torch.randn(32, 64, device="cuda"))
    aux0 = torch.cuda.Stream()

    def wq_b(qr_, qr_scale):
        order.append("wq_b")
        streams["wq_b"] = torch.cuda.current_stream()
        return q

    def qnorm_insert(q_, kv_, positions_, attn_metadata):
        order.append("qnorm_insert")
        streams["qnorm_insert"] = torch.cuda.current_stream()
        return q_

    attention = SimpleNamespace(
        indexer=indexer,
        aux_stream_list=[aux0],
        ln_events=[torch.cuda.Event() for _ in range(2)],
        indexer_rotary_emb=rotary,
        n_local_heads=1,
        head_dim=512,
        _run_parallel_input_projections=lambda hidden_states_: (
            qr_kv,
            None,
            indexer_weights,
        ),
        _split_qkv_and_norm=lambda qr_kv_: (qr, None, kv),
        _wq_b_proj=wq_b,
        _fused_qnorm_rope_kv_insert=qnorm_insert,
    )

    context = ForwardContext({}, {}, {})
    with override_forward_context(context):
        result = DeepseekV41ROCMAiterMLAAttention._forward_csa2_reindex(
            attention, hidden_states, positions
        )

    q_out, kv_out, index_q_out, index_q_scale_out, index_weights_out_out = result
    assert torch.equal(
        q_out, q.view(num_tokens, attention.n_local_heads, attention.head_dim)
    )
    assert torch.equal(kv_out, kv)
    assert torch.equal(index_q_out, index_q)
    assert torch.equal(index_q_scale_out, index_q_scale)
    assert torch.equal(index_weights_out_out, index_weights_out)

    assert order == ["wq_b", "qnorm_insert", "forward_q"]
    current = torch.cuda.current_stream()
    for name in ("wq_b", "qnorm_insert"):
        assert streams[name] == current, name
    assert streams["forward_q"] == aux0


@pytest.mark.parametrize("quant_tuple", [False, True])
def test_v41_indexer_forward_q(quant_tuple):
    """forward_q runs wq_b then the fused RoPE/quant and unpacks both shapes."""
    from unittest.mock import patch

    from vllm.models.deepseek_v41.attention import DeepseekV4Indexer

    num_tokens, n_head, head_dim = 5, 32, 128
    qr = torch.randn(num_tokens, 1280)
    q = torch.randn(num_tokens, n_head * head_dim)
    positions = torch.arange(num_tokens)
    indexer_weights = torch.randn(num_tokens, n_head)
    rotary = SimpleNamespace(cos_sin_cache=torch.randn(32, 64))
    fake = SimpleNamespace(
        n_head=n_head,
        head_dim=head_dim,
        softmax_scale=head_dim**-0.5,
        use_fp4_kv=False,
        indexer_weights_dtype=torch.float32,
        _wq_b_proj=lambda qr_, qr_scale: q,
    )
    q_quant = torch.randn(num_tokens, n_head, head_dim)
    q_scale = torch.randn(num_tokens, n_head)
    weights_out = torch.randn(num_tokens, n_head)
    if quant_tuple:
        fused_return = ((q_quant, q_scale), weights_out)
        expected = (q_quant, q_scale, weights_out)
    else:
        fused_return = (q_quant, weights_out)
        expected = (q_quant, None, weights_out)

    with patch(
        "vllm.models.deepseek_v41.attention.fused_indexer_q_rope_quant",
        return_value=fused_return,
    ) as mock_fused:
        result = DeepseekV4Indexer.forward_q(
            fake, qr, None, indexer_weights, positions, rotary
        )

    assert torch.equal(result[0], expected[0])
    assert (result[1] is None) == (expected[1] is None)
    if expected[1] is not None:
        assert torch.equal(result[1], expected[1])
    assert torch.equal(result[2], expected[2])
    args = mock_fused.call_args.args
    assert args[0] is positions
    torch.testing.assert_close(args[1], q.view(-1, n_head, head_dim))
    assert args[2] is rotary.cos_sin_cache
    assert args[3] is indexer_weights
    assert args[4] == head_dim**-0.5
    assert args[5] == n_head**-0.5
    assert mock_fused.call_args.kwargs == {
        "use_fp4": False,
        "weights_out_dtype": torch.float32,
    }
