# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Unit tests for the horizontally-fused deepseek_v32 (NVIDIA SM100) Triton
kernels used by the specialized DSA model:

  fused_norm_rope
    - q  : q_lora RMSNorm
    - kv : kv_lora RMSNorm + (interleaved) RoPE on k_pe + MLA cache insert
           (bf16 or per-tensor fp8)
    - idx: indexer-K LayerNorm + RoPE (interleaved or NeoX) + UE8M0 fp8 quant +
           packed indexer cache insert; plus the top-k buffer (-1) fill
  fused_q
    - mqa: ql_nope + (interleaved) RoPE'd q_pe, concat-quantized to the fp8 MQA
           query
    - idx: indexer-Q RoPE (interleaved or NeoX) + UE8M0 fp8 quant + folded
           index weights
  fused_eh_norm (MTP): zero-at-pos-0 + enorm RMSNorm(embeds) + hnorm
           RMSNorm(prev), concatenated side-by-side

Each kernel is compared against a PyTorch reference. The kernel keeps the whole
pipeline in fp32 and rounds once, so it can land on the opposite side of a
round-to-nearest tie from the reference for a few elements: deterministic fp8
outputs are checked within 1 representable-step (ULP); bf16 norm/RoPE outputs use
rtol/atol=1e-2 (the tolerance the sibling deepseek_v4 fused-kernel test uses).
"""

from functools import partial
from typing import NamedTuple

import pytest
import torch

from vllm.models.deepseek_v32.common import kernels as K
from vllm.platforms import current_platform

FP8 = torch.float8_e4m3fn
FP8_MAX = 448.0

EPS = 1e-6


class ModelConfig(NamedTuple):
    """Per-model attention dims (num_heads is the TP8 local count)."""

    q_lora: int
    kv_lora: int
    rope_dim: int
    num_heads: int
    index_heads: int
    index_head_dim: int
    hidden: int


MODEL_CONFIGS = [
    ModelConfig(2048, 512, 64, 8, 32, 128, 6144),  # GLM-5.2
    ModelConfig(1536, 512, 64, 16, 64, 128, 7168),  # DeepSeek-V3.2
]
MODEL_IDS = ["glm5.2", "dsv3.2"]

pytestmark = pytest.mark.skipif(
    not current_platform.supports_fp8(),
    reason="deepseek_v32 fused kernels require fp8 support",
)


# ── reference helpers ────────────────────────────────────────────────────────


def make_cos_sin(max_pos: int, rot_dim: int, device) -> torch.Tensor:
    """cos||sin cache: row[pos] = [cos(theta)(rot/2), sin(theta)(rot/2)]."""
    half = rot_dim // 2
    inv_freq = 1.0 / (
        10000.0 ** (torch.arange(0, half, dtype=torch.float32, device=device) / half)
    )
    t = torch.arange(max_pos, dtype=torch.float32, device=device)
    freqs = torch.einsum("i,j->ij", t, inv_freq)
    return torch.cat([freqs.cos(), freqs.sin()], dim=-1)


def rms_norm(x: torch.Tensor, w: torch.Tensor) -> torch.Tensor:
    """RMSNorm matching kernels._rms_norm (fp32, eps inside rsqrt). Returns fp32."""
    xf = x.float()
    ms = xf.pow(2).mean(dim=-1, keepdim=True)
    return xf * torch.rsqrt(ms + EPS) * w.float()


def layer_norm(x: torch.Tensor, w: torch.Tensor, b: torch.Tensor) -> torch.Tensor:
    xf = x.float()
    mean = xf.mean(dim=-1, keepdim=True)
    var = (xf - mean).pow(2).mean(dim=-1, keepdim=True)
    return (xf - mean) * torch.rsqrt(var + EPS) * w.float() + b.float()


def rope(
    x: torch.Tensor, pos: torch.Tensor, cos_sin: torch.Tensor, interleave: bool
) -> torch.Tensor:
    """Apply RoPE to the first ``rot_dim`` elements of x's last dim.

    x: [..., head_dim] fp32. ``cos_sin`` is [max_pos, rot_dim]. ``interleave``
    selects adjacent-pair (GLM) vs split-half NeoX (DeepSeek-V3.2) layout.
    """
    rot = cos_sin.shape[-1]
    half = rot // 2
    cs = cos_sin[pos.long()]
    cos, sin = cs[..., :half], cs[..., half:]
    out = x.float().clone()
    r = out[..., :rot]
    if interleave:
        x1, x2 = r[..., 0::2].clone(), r[..., 1::2].clone()
        r[..., 0::2] = x1 * cos - x2 * sin
        r[..., 1::2] = x2 * cos + x1 * sin
    else:
        x1, x2 = r[..., :half].clone(), r[..., half:].clone()
        r[..., :half] = x1 * cos - x2 * sin
        r[..., half:] = x2 * cos + x1 * sin
    return out


def ue8m0_quant(vals: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    """Per-row (last dim) UE8M0 fp8 quant matching kernels._fp8_ue8m0_quantize."""
    amax = vals.float().abs().amax(dim=-1, keepdim=True)
    scale = torch.clamp(amax, min=1e-4) / FP8_MAX
    scale = torch.exp2(torch.ceil(torch.log2(scale)))
    q = (vals.float() / scale).to(FP8)
    return q, scale.squeeze(-1)


def _bf16_ulp(a: torch.Tensor, b: torch.Tensor) -> int:
    def key(t):
        u = t.contiguous().view(torch.int16).to(torch.int64) & 0xFFFF
        return torch.where(u >= 0x8000, 0xFFFF - u, u + 0x8000)

    return int((key(a) - key(b)).abs().max().item())


def _fp8_ulp(a: torch.Tensor, b: torch.Tensor) -> int:
    def key(t):
        u = t.contiguous().view(torch.uint8).to(torch.int64)
        return torch.where(u >= 0x80, 0xFF - u, u + 0x80)

    return int((key(a) - key(b)).abs().max().item())


def assert_bf16(got: torch.Tensor, ref_fp32: torch.Tensor, msg: str):
    # Kernel keeps RMSNorm/RoPE in fp32 and rounds to bf16 once; the fp32
    # reduction/FMA order differs from torch, so a few elements land on the
    # opposite side of a round-to-nearest tie. Use the same tolerance the
    # sibling deepseek_v4 fused-kernel test uses for this bf16 norm+rope class.
    torch.testing.assert_close(
        got.float(), ref_fp32.float(), rtol=1e-2, atol=1e-2, msg=lambda m: f"{msg}: {m}"
    )


def assert_fp8(got: torch.Tensor, ref: torch.Tensor, msg: str):
    assert _fp8_ulp(got, ref) <= 1, f"{msg}: >1 fp8 ULP"


# ── fused_norm_rope ──────────────────────────────────────────────────────────


@pytest.mark.parametrize("cfg", MODEL_CONFIGS, ids=MODEL_IDS)
@pytest.mark.parametrize("num_tokens", [1, 4, 17, 512, 4096])
@pytest.mark.parametrize("index_interleave", [True, False])
@pytest.mark.parametrize("mla_dtype", ["auto", "bfloat16", "fp8"])
def test_fused_norm_rope(
    num_tokens: int, index_interleave: bool, mla_dtype: str, cfg: ModelConfig
):
    torch.manual_seed(0)
    dev = "cuda"
    max_pos = 8192
    pos = torch.arange(num_tokens, device=dev, dtype=torch.int64) % max_pos

    q_c = torch.randn(num_tokens, cfg.q_lora, device=dev, dtype=torch.bfloat16)
    kv_c = torch.randn(num_tokens, cfg.kv_lora, device=dev, dtype=torch.bfloat16)
    k_pe = torch.randn(num_tokens, cfg.rope_dim, device=dev, dtype=torch.bfloat16)
    qw = torch.randn(cfg.q_lora, device=dev, dtype=torch.bfloat16)
    kvw = torch.randn(cfg.kv_lora, device=dev, dtype=torch.bfloat16)
    ik = torch.randn(num_tokens, cfg.index_head_dim, device=dev, dtype=torch.bfloat16)
    ikw = torch.randn(cfg.index_head_dim, device=dev, dtype=torch.float32)
    ikb = torch.randn(cfg.index_head_dim, device=dev, dtype=torch.float32)

    mla_cos_sin = make_cos_sin(max_pos, cfg.rope_dim, dev)  # MLA k_pe: interleaved
    idx_cos_sin = make_cos_sin(max_pos, cfg.rope_dim, dev)

    bs = max_pos  # single block covering all tokens
    mla_dim = cfg.kv_lora + cfg.rope_dim
    mla_fp8 = mla_dtype == "fp8"
    if mla_fp8:
        mla_cache = torch.zeros(1, bs, mla_dim, device=dev, dtype=torch.uint8)
        mla_k_scale = torch.tensor([0.3], device=dev, dtype=torch.float32)
    else:
        mla_cache = torch.zeros(1, bs, mla_dim, device=dev, dtype=torch.bfloat16)
        mla_k_scale = None
    idx_row = cfg.index_head_dim + cfg.index_head_dim // 128 * 4  # 132
    idx_cache = torch.zeros(1, bs, idx_row, device=dev, dtype=torch.uint8)
    slot = torch.arange(num_tokens, device=dev, dtype=torch.int64)
    topk = torch.full((num_tokens, 2048), 7, device=dev, dtype=torch.int32)

    q_out = K.fused_norm_rope(
        pos,
        q_c,
        qw,
        EPS,
        kv_c,
        kvw,
        EPS,
        k_pe,
        mla_cos_sin,
        ik,
        ikw,
        ikb,
        EPS,
        idx_cos_sin,
        topk,
        slot_mapping=slot,
        indexer_k_cache=idx_cache,
        mla_kv_cache=mla_cache,
        mla_kv_cache_dtype=mla_dtype,
        mla_k_scale=mla_k_scale,
        has_indexer=True,
        index_rope_interleave=index_interleave,
    )

    # q_lora RMSNorm
    assert_bf16(q_out, rms_norm(q_c, qw), "q_c rmsnorm")

    # MLA cache: [kv_c_normed | k_pe_roped(interleaved)]
    kv_ref = rms_norm(kv_c, kvw)
    kpe_ref = rope(k_pe.float(), pos, mla_cos_sin, interleave=True)
    if mla_fp8:
        cache = mla_cache.view(FP8)[0, :num_tokens]
        s = mla_k_scale.item()
        assert_fp8(cache[:, : cfg.kv_lora], (kv_ref / s).to(FP8), "MLA kv fp8")
        assert_fp8(cache[:, cfg.kv_lora :], (kpe_ref / s).to(FP8), "MLA k_pe fp8")
    else:
        cache = mla_cache[0, :num_tokens]
        assert_bf16(cache[:, : cfg.kv_lora], kv_ref, "MLA kv bf16")
        assert_bf16(cache[:, cfg.kv_lora :], kpe_ref, "MLA k_pe bf16")

    # Indexer-K cache (packed [bs*head_dim fp8 | bs*4 fp32 scale]).
    ik_ref = layer_norm(ik, ikw, ikb)
    ik_ref = rope(ik_ref, pos, idx_cos_sin, interleave=index_interleave)
    q_ref, s_ref = ue8m0_quant(ik_ref)
    flat = idx_cache[0].reshape(-1)
    vals = flat[: bs * cfg.index_head_dim].view(FP8).reshape(bs, cfg.index_head_dim)
    scales = flat[bs * cfg.index_head_dim :].view(torch.float32)
    assert_fp8(vals[:num_tokens], q_ref, "indexer-K fp8")
    torch.testing.assert_close(scales[:num_tokens], s_ref, rtol=0, atol=0)

    # Top-k buffer cleared to -1 on indexer layers.
    assert (topk == -1).all(), "topk buffer not cleared on indexer layer"


@pytest.mark.parametrize("cfg", MODEL_CONFIGS, ids=MODEL_IDS)
def test_fused_norm_rope_packed_indexer_block_stride(cfg: ModelConfig):
    """Indexer writes use their HMA block stride and independent slot map."""
    torch.manual_seed(7)
    dev = "cuda"
    num_tokens = block_size = 2
    pos = torch.arange(num_tokens, device=dev, dtype=torch.int64)
    q_c = torch.randn(num_tokens, cfg.q_lora, device=dev, dtype=torch.bfloat16)
    kv_c = torch.randn(num_tokens, cfg.kv_lora, device=dev, dtype=torch.bfloat16)
    k_pe = torch.randn(num_tokens, cfg.rope_dim, device=dev, dtype=torch.bfloat16)
    qw = torch.randn(cfg.q_lora, device=dev, dtype=torch.bfloat16)
    kvw = torch.randn(cfg.kv_lora, device=dev, dtype=torch.bfloat16)
    ik = torch.randn(num_tokens, cfg.index_head_dim, device=dev, dtype=torch.bfloat16)
    ikw = torch.randn(cfg.index_head_dim, device=dev, dtype=torch.float32)
    ikb = torch.randn(cfg.index_head_dim, device=dev, dtype=torch.float32)
    cos_sin = make_cos_sin(32, cfg.rope_dim, dev)

    idx_row = cfg.index_head_dim + cfg.index_head_dim // 128 * 4
    packed_block_stride = block_size * idx_row + 64
    backing = torch.zeros(2 * packed_block_stride, device=dev, dtype=torch.uint8)
    idx_cache = torch.as_strided(
        backing,
        (2, block_size, idx_row),
        (packed_block_stride, idx_row, 1),
    )
    mla_cache = torch.zeros(
        1,
        block_size,
        cfg.kv_lora + cfg.rope_dim,
        device=dev,
        dtype=torch.bfloat16,
    )
    topk = torch.zeros(num_tokens, 8, device=dev, dtype=torch.int32)

    K.fused_norm_rope(
        pos,
        q_c,
        qw,
        EPS,
        kv_c,
        kvw,
        EPS,
        k_pe,
        cos_sin,
        ik,
        ikw,
        ikb,
        EPS,
        cos_sin,
        topk,
        slot_mapping=torch.arange(num_tokens, device=dev, dtype=torch.int64),
        indexer_slot_mapping=torch.arange(
            block_size, block_size + num_tokens, device=dev, dtype=torch.int64
        ),
        indexer_k_cache=idx_cache,
        mla_kv_cache=mla_cache,
        has_indexer=True,
    )

    ik_ref = rope(layer_norm(ik, ikw, ikb), pos, cos_sin, interleave=False)
    q_ref, s_ref = ue8m0_quant(ik_ref)
    packed = idx_cache[1].reshape(-1)
    values = (
        packed[: block_size * cfg.index_head_dim]
        .view(FP8)
        .view(block_size, cfg.index_head_dim)
    )
    scales = packed[block_size * cfg.index_head_dim :].view(torch.float32)
    assert_fp8(values, q_ref, "packed indexer-K fp8")
    torch.testing.assert_close(scales, s_ref, rtol=0, atol=0)
    assert (backing[block_size * idx_row : packed_block_stride] == 0).all()


@pytest.mark.parametrize("cfg", MODEL_CONFIGS, ids=MODEL_IDS)
def test_fused_norm_rope_materializes_qk_without_local_cache_slots(
    cfg: ModelConfig,
):
    """DCP non-owner ranks still need valid query shards for query AllGather."""
    torch.manual_seed(7)
    dev = "cuda"
    num_tokens = 4
    max_pos = 16
    pos = torch.arange(num_tokens, device=dev, dtype=torch.int64)

    q_c = torch.randn(num_tokens, cfg.q_lora, device=dev, dtype=torch.bfloat16)
    kv_c = torch.randn(num_tokens, cfg.kv_lora, device=dev, dtype=torch.bfloat16)
    k_pe = torch.randn(num_tokens, cfg.rope_dim, device=dev, dtype=torch.bfloat16)
    qw = torch.randn(cfg.q_lora, device=dev, dtype=torch.bfloat16)
    kvw = torch.randn(cfg.kv_lora, device=dev, dtype=torch.bfloat16)
    ik = torch.randn(num_tokens, cfg.index_head_dim, device=dev, dtype=torch.bfloat16)
    ikw = torch.randn(cfg.index_head_dim, device=dev, dtype=torch.float32)
    ikb = torch.randn(cfg.index_head_dim, device=dev, dtype=torch.float32)
    cos_sin = make_cos_sin(max_pos, cfg.rope_dim, dev)

    mla_cache = torch.zeros(
        1, max_pos, cfg.kv_lora + cfg.rope_dim, device=dev, dtype=torch.bfloat16
    )
    idx_row = cfg.index_head_dim + cfg.index_head_dim // 128 * 4
    idx_cache = torch.zeros(1, max_pos, idx_row, device=dev, dtype=torch.uint8)
    no_local_slots = torch.full((num_tokens,), -1, device=dev, dtype=torch.int64)
    topk = torch.full((num_tokens, 2048), 7, device=dev, dtype=torch.int32)
    kv_out = torch.empty_like(kv_c)
    kpe_out = torch.empty_like(k_pe)

    q_out = K.fused_norm_rope(
        pos,
        q_c,
        qw,
        EPS,
        kv_c,
        kvw,
        EPS,
        k_pe,
        cos_sin,
        ik,
        ikw,
        ikb,
        EPS,
        cos_sin,
        topk,
        slot_mapping=no_local_slots,
        indexer_k_cache=idx_cache,
        mla_kv_cache=mla_cache,
        mla_kv_cache_dtype="auto",
        mla_k_scale=None,
        has_indexer=True,
        index_rope_interleave=True,
        kv_c_out=kv_out,
        k_pe_out=kpe_out,
    )

    assert_bf16(q_out, rms_norm(q_c, qw), "q_c rmsnorm without local cache slots")
    assert_bf16(kv_out, rms_norm(kv_c, kvw), "kv_c rmsnorm without local slots")
    assert_bf16(
        kpe_out,
        rope(k_pe.float(), pos, cos_sin, interleave=True),
        "k_pe RoPE without local slots",
    )
    assert not mla_cache.any(), "non-owner rank wrote the MLA KV cache"
    assert not idx_cache.any(), "non-owner rank wrote the indexer KV cache"
    assert (topk == -1).all(), "topk buffer not cleared on non-owner rank"


@pytest.mark.parametrize("cfg", MODEL_CONFIGS, ids=MODEL_IDS)
@pytest.mark.parametrize("num_tokens", [1, 17, 512])
def test_fused_norm_rope_no_indexer(num_tokens: int, cfg: ModelConfig):
    """Shared (no-indexer) layer: q + kv/MLA only; top-k buffer untouched."""
    torch.manual_seed(1)
    dev = "cuda"
    max_pos = 8192
    pos = torch.arange(num_tokens, device=dev, dtype=torch.int64)

    q_c = torch.randn(num_tokens, cfg.q_lora, device=dev, dtype=torch.bfloat16)
    kv_c = torch.randn(num_tokens, cfg.kv_lora, device=dev, dtype=torch.bfloat16)
    k_pe = torch.randn(num_tokens, cfg.rope_dim, device=dev, dtype=torch.bfloat16)
    qw = torch.randn(cfg.q_lora, device=dev, dtype=torch.bfloat16)
    kvw = torch.randn(cfg.kv_lora, device=dev, dtype=torch.bfloat16)
    mla_cos_sin = make_cos_sin(max_pos, cfg.rope_dim, dev)

    bs = max_pos
    mla_cache = torch.zeros(
        1, bs, cfg.kv_lora + cfg.rope_dim, device=dev, dtype=torch.bfloat16
    )
    slot = torch.arange(num_tokens, device=dev, dtype=torch.int64)
    topk = torch.full((num_tokens, 2048), 7, device=dev, dtype=torch.int32)

    q_out = K.fused_norm_rope(
        pos,
        q_c,
        qw,
        EPS,
        kv_c,
        kvw,
        EPS,
        k_pe,
        mla_cos_sin,
        None,
        None,
        None,
        EPS,
        None,
        topk,
        slot_mapping=slot,
        indexer_k_cache=None,
        mla_kv_cache=mla_cache,
        mla_kv_cache_dtype="auto",
        mla_k_scale=None,
        has_indexer=False,
        index_rope_interleave=False,
    )

    assert_bf16(q_out, rms_norm(q_c, qw), "q_c rmsnorm (no-indexer)")
    cache = mla_cache[0, :num_tokens]
    assert_bf16(cache[:, : cfg.kv_lora], rms_norm(kv_c, kvw), "MLA kv (no-indexer)")
    assert_bf16(
        cache[:, cfg.kv_lora :],
        rope(k_pe.float(), pos, mla_cos_sin, interleave=True),
        "MLA k_pe (no-indexer)",
    )
    # Shared layers reuse the previous indexer's top-k: buffer must be untouched.
    assert (topk == 7).all(), "topk buffer should be untouched on shared layer"


@pytest.mark.parametrize("cfg", MODEL_CONFIGS, ids=MODEL_IDS)
@pytest.mark.parametrize("has_indexer", [False, True])
def test_fused_norm_rope_materializes_pcp_cache_inputs(
    has_indexer: bool, cfg: ModelConfig
):
    """PCP gets local normalized/rotated K rows without direct cache writes."""
    torch.manual_seed(6)
    dev = "cuda"
    num_tokens = 17
    max_pos = 8192
    pos = torch.arange(num_tokens, device=dev, dtype=torch.int64)
    q_c = torch.randn(num_tokens, cfg.q_lora, device=dev, dtype=torch.bfloat16)
    kv_c = torch.randn(num_tokens, cfg.kv_lora, device=dev, dtype=torch.bfloat16)
    k_pe = torch.randn(num_tokens, cfg.rope_dim, device=dev, dtype=torch.bfloat16)
    qw = torch.randn(cfg.q_lora, device=dev, dtype=torch.bfloat16)
    kvw = torch.randn(cfg.kv_lora, device=dev, dtype=torch.bfloat16)
    ik = (
        torch.randn(num_tokens, cfg.index_head_dim, device=dev, dtype=torch.bfloat16)
        if has_indexer
        else None
    )
    ikw = (
        torch.randn(cfg.index_head_dim, device=dev, dtype=torch.float32)
        if has_indexer
        else None
    )
    ikb = (
        torch.randn(cfg.index_head_dim, device=dev, dtype=torch.float32)
        if has_indexer
        else None
    )
    mla_cos_sin = make_cos_sin(max_pos, cfg.rope_dim, dev)
    idx_cos_sin = make_cos_sin(max_pos, cfg.rope_dim, dev) if has_indexer else None
    q_out = torch.empty_like(q_c)
    kv_out = torch.empty_like(kv_c)
    kpe_out = torch.empty_like(k_pe)
    ik_out = torch.empty_like(ik) if ik is not None else None
    topk = torch.full((num_tokens, 2048), 7, device=dev, dtype=torch.int32)

    actual_q = K.fused_norm_rope(
        pos,
        q_c,
        qw,
        EPS,
        kv_c,
        kvw,
        EPS,
        k_pe,
        mla_cos_sin,
        ik,
        ikw,
        ikb,
        EPS,
        idx_cos_sin,
        topk,
        has_indexer=has_indexer,
        index_rope_interleave=True,
        q_c_out=q_out,
        kv_c_out=kv_out,
        k_pe_out=kpe_out,
        index_k_out=ik_out,
    )

    assert actual_q.data_ptr() == q_out.data_ptr()
    assert_bf16(actual_q, rms_norm(q_c, qw), "PCP q norm")
    assert_bf16(kv_out, rms_norm(kv_c, kvw), "PCP kv norm")
    assert_bf16(
        kpe_out,
        rope(k_pe.float(), pos, mla_cos_sin, interleave=True),
        "PCP k_pe RoPE",
    )
    if has_indexer:
        assert ik is not None and ikw is not None and ikb is not None
        assert ik_out is not None and idx_cos_sin is not None
        ik_ref = rope(layer_norm(ik, ikw, ikb), pos, idx_cos_sin, interleave=True)
        assert_bf16(ik_out, ik_ref, "PCP indexer-K")


@pytest.mark.parametrize("cfg", MODEL_CONFIGS, ids=MODEL_IDS)
@pytest.mark.parametrize("num_tokens", [1, 4, 17, 512])
def test_fused_norm_rope_ds_mla(num_tokens: int, cfg: ModelConfig):
    """fp8_ds_mla MLA cache layout (FlashMLA sparse, bf16-query path; SM90/SM100).

    Per-token 656-byte entry: 512 fp8 NoPE (4 per-128 tiles, power-of-two scales
    stored as float32) | 4 float32 scales | 64 bf16 (unquantized) RoPE.
    """
    torch.manual_seed(5)
    dev = "cuda"
    max_pos = 8192
    pos = torch.arange(num_tokens, device=dev, dtype=torch.int64) % max_pos

    q_c = torch.randn(num_tokens, cfg.q_lora, device=dev, dtype=torch.bfloat16)
    kv_c = torch.randn(num_tokens, cfg.kv_lora, device=dev, dtype=torch.bfloat16)
    k_pe = torch.randn(num_tokens, cfg.rope_dim, device=dev, dtype=torch.bfloat16)
    qw = torch.randn(cfg.q_lora, device=dev, dtype=torch.bfloat16)
    kvw = torch.randn(cfg.kv_lora, device=dev, dtype=torch.bfloat16)
    mla_cos_sin = make_cos_sin(max_pos, cfg.rope_dim, dev)

    bs = max_pos
    mla_cache = torch.zeros(1, bs, 656, device=dev, dtype=torch.uint8)
    slot = torch.arange(num_tokens, device=dev, dtype=torch.int64)
    topk = torch.full((num_tokens, 2048), 7, device=dev, dtype=torch.int32)

    q_out = K.fused_norm_rope(
        pos,
        q_c,
        qw,
        EPS,
        kv_c,
        kvw,
        EPS,
        k_pe,
        mla_cos_sin,
        None,
        None,
        None,
        EPS,
        None,
        topk,
        slot_mapping=slot,
        indexer_k_cache=None,
        mla_kv_cache=mla_cache,
        mla_kv_cache_dtype="fp8_ds_mla",
        mla_k_scale=None,
        has_indexer=False,
        index_rope_interleave=False,
    )

    assert_bf16(q_out, rms_norm(q_c, qw), "q_c rmsnorm (ds_mla)")

    kv_ref = rms_norm(kv_c, kvw)  # [N, 512] fp32
    kpe_ref = rope(k_pe.float(), pos, mla_cos_sin, interleave=True)  # [N, 64]
    tiles = kv_ref.view(num_tokens, 4, 128)
    ref_scale = torch.clamp(tiles.abs().amax(dim=-1) / FP8_MAX, min=1e-4)
    ref_scale = torch.exp2(torch.ceil(torch.log2(ref_scale)))
    ref_nope = (tiles / ref_scale[..., None]).reshape(num_tokens, cfg.kv_lora).to(FP8)

    cache = mla_cache[0, :num_tokens]  # [N, 656] uint8
    nope = cache[:, : cfg.kv_lora].view(FP8)
    scales = cache.view(torch.float32)[:, cfg.kv_lora // 4 : cfg.kv_lora // 4 + 4]
    rope_off = cfg.kv_lora // 2 + 8
    rope_vals = cache.view(torch.bfloat16)[:, rope_off : rope_off + cfg.rope_dim]

    torch.testing.assert_close(scales, ref_scale, rtol=0, atol=0)
    assert_fp8(nope, ref_nope, "ds_mla NoPE fp8")
    assert_bf16(rope_vals, kpe_ref, "ds_mla RoPE bf16")
    # No indexer on this call: top-k buffer must be untouched.
    assert (topk == 7).all(), "topk buffer should be untouched (no indexer)"


E2M1_MAGNITUDES = [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0]


def quantize_to_e2m1(x: torch.Tensor) -> torch.Tensor:
    """Round to nearest e2m1, saturating to +-6 like cvt.rn.satfinite.e2m1x2.f32."""
    mags = torch.tensor(E2M1_MAGNITUDES, dtype=torch.float32, device=x.device)
    a = x.float().abs().clamp_max(6.0)
    mids = (mags[:-1] + mags[1:]) / 2
    code = torch.bucketize(a, mids, right=True)
    on_tie = (a.unsqueeze(-1) == mids).any(dim=-1)
    tie_code = torch.bucketize(a, mids, right=False)
    code = torch.where(on_tie, tie_code + (tie_code & 1), code)
    return (torch.signbit(x).to(torch.uint8) << 3) | code.to(torch.uint8)


def nvfp4_sf_byte(s: torch.Tensor) -> torch.Tensor:
    """Scale-factor byte permutation shared with FlashMLA: 8*(s&3) + (s>>2)."""
    return 8 * (s % 4) + (s // 4)


@pytest.mark.skipif(
    not current_platform.is_device_capability(100),
    reason="nvfp4_ds_mla requires SM100 (Blackwell)",
)
@pytest.mark.parametrize("cfg", MODEL_CONFIGS, ids=MODEL_IDS)
@pytest.mark.parametrize("num_tokens", [1, 4, 17, 512])
def test_fused_norm_rope_nvfp4_ds_mla(num_tokens: int, cfg: ModelConfig):
    """nvfp4_ds_mla MLA cache layout (FlashMLA sparse, SM100 only).

    Per-token 352-byte entry: 256 B of 512 e2m1 NoPE packed 2/byte (low nibble
    = even element) | 64 B unscaled e4m3 RoPE | 32 B byte-permuted e4m3 tile
    scales, one per 16 NoPE elements.
    """
    torch.manual_seed(5)
    dev = "cuda"
    max_pos = 8192
    pos = torch.arange(num_tokens, device=dev, dtype=torch.int64) % max_pos

    q_c = torch.randn(num_tokens, cfg.q_lora, device=dev, dtype=torch.bfloat16)
    kv_c = torch.randn(num_tokens, cfg.kv_lora, device=dev, dtype=torch.bfloat16)
    # Tile 0 is deliberately low-magnitude: amax/6 lands between two e4m3
    # subnormals (spaced a flat 2^-9 there), so round-to-nearest picks a scale
    # BELOW amax/6 and the tile's peak saturates at +-6. Asserted below, so a
    # future change to the scale rule cannot silently stop covering this.
    # amax must land in (6, 9) * 2^-9 after rms_norm; 7.5 sits mid-band, with
    # margin for the per-token RMS to vary.
    kv_c[:, :16] = (
        7.5 * (2.0**-9) * torch.linspace(0.2, 1.0, 16, device=dev, dtype=torch.bfloat16)
    )
    k_pe = torch.randn(num_tokens, cfg.rope_dim, device=dev, dtype=torch.bfloat16)
    qw = torch.randn(cfg.q_lora, device=dev, dtype=torch.bfloat16)
    kvw = torch.ones(cfg.kv_lora, device=dev, dtype=torch.bfloat16)
    mla_cos_sin = make_cos_sin(max_pos, cfg.rope_dim, dev)

    bs = max_pos
    mla_cache = torch.zeros(1, bs, 352, device=dev, dtype=torch.uint8)
    slot = torch.arange(num_tokens, device=dev, dtype=torch.int64)
    topk = torch.full((num_tokens, 2048), 7, device=dev, dtype=torch.int32)

    q_out = K.fused_norm_rope(
        pos,
        q_c,
        qw,
        EPS,
        kv_c,
        kvw,
        EPS,
        k_pe,
        mla_cos_sin,
        None,
        None,
        None,
        EPS,
        None,
        topk,
        slot_mapping=slot,
        indexer_k_cache=None,
        mla_kv_cache=mla_cache,
        mla_kv_cache_dtype="nvfp4_ds_mla",
        mla_k_scale=None,
        has_indexer=False,
        index_rope_interleave=False,
    )

    assert_bf16(q_out, rms_norm(q_c, qw), "q_c rmsnorm (nvfp4_ds_mla)")

    kv_ref = rms_norm(kv_c, kvw)  # [N, 512] fp32
    kpe_ref = rope(k_pe.float(), pos, mla_cos_sin, interleave=True)  # [N, 64]
    tiles = kv_ref.view(num_tokens, 32, 16)
    amax = tiles.abs().amax(dim=-1)
    scale_target = torch.clamp_min(amax / 6.0, 2.0**-9)
    ref_scale = scale_target.to(FP8)  # round-to-nearest

    # The crafted tile must actually exercise the saturating path.
    assert (ref_scale[:, 0].float() < scale_target[:, 0]).all(), (
        "tile 0 should round its scale DOWN; adjust the crafted magnitude"
    )
    assert (amax[:, 0] / ref_scale[:, 0].float() > 6.0).all(), (
        "tile 0 should saturate e2m1; adjust the crafted magnitude"
    )

    ref_codes = quantize_to_e2m1(tiles / ref_scale.float().unsqueeze(-1))
    ref_codes = ref_codes.reshape(num_tokens, cfg.kv_lora)

    cache = mla_cache[0, :num_tokens]  # [N, 352] uint8
    packed = cache[:, : cfg.kv_lora // 2]
    got_codes = torch.stack([packed & 0xF, packed >> 4], dim=-1)
    got_codes = got_codes.reshape(num_tokens, cfg.kv_lora)  # low nibble = even elem
    torch.testing.assert_close(got_codes, ref_codes, rtol=0, atol=0)

    got_rope = cache[:, cfg.kv_lora // 2 : cfg.kv_lora // 2 + cfg.rope_dim].view(FP8)
    assert_fp8(got_rope, kpe_ref.to(FP8), "nvfp4_ds_mla RoPE e4m3")

    perm = nvfp4_sf_byte(torch.arange(32, device=dev))
    got_scale = cache[:, cfg.kv_lora // 2 + cfg.rope_dim :].view(FP8)[:, perm]
    torch.testing.assert_close(got_scale.float(), ref_scale.float(), rtol=0, atol=0)

    # No indexer on this call: top-k buffer must be untouched.
    assert (topk == 7).all(), "topk buffer should be untouched (no indexer)"


def test_fused_norm_rope_supports_large_token_count():
    """Keep the token count off CUDA grid-y at its 65,536-block boundary."""
    num_tokens = 65536
    dev = "cuda"
    dtype = torch.bfloat16
    positions = torch.zeros(num_tokens, device=dev, dtype=torch.int64)
    q_c = torch.ones((num_tokens, 1), device=dev, dtype=dtype)
    kv_c = torch.ones((num_tokens, 1), device=dev, dtype=dtype)
    k_pe = torch.ones((num_tokens, 2), device=dev, dtype=dtype)
    norm_w = torch.ones(1, device=dev, dtype=dtype)
    cos_sin = torch.tensor([[1.0, 0.0]], device=dev, dtype=torch.float32)
    topk = torch.empty((num_tokens, 1), device=dev, dtype=torch.int32)
    slot_mapping = torch.arange(num_tokens, device=dev, dtype=torch.int64)
    mla_cache = torch.empty((1, num_tokens, 3), device=dev, dtype=dtype)

    q_out = K.fused_norm_rope(
        positions,
        q_c,
        norm_w,
        EPS,
        kv_c,
        norm_w,
        EPS,
        k_pe,
        cos_sin,
        None,
        None,
        None,
        EPS,
        None,
        topk,
        slot_mapping=slot_mapping,
        mla_kv_cache=mla_cache,
        has_indexer=False,
    )

    rows = torch.tensor([0, num_tokens - 1], device=dev)
    assert_bf16(q_out[rows], rms_norm(q_c[rows], norm_w), "large-token q norm")


# ── fused_q ──────────────────────────────────────────────────────────────────


@pytest.mark.parametrize("cfg", MODEL_CONFIGS, ids=MODEL_IDS)
@pytest.mark.parametrize("num_tokens", [1, 4, 17, 512, 4096])
@pytest.mark.parametrize("index_interleave", [True, False])
def test_fused_q(num_tokens: int, index_interleave: bool, cfg: ModelConfig):
    torch.manual_seed(2)
    dev = "cuda"
    max_pos = 8192
    pos = torch.arange(num_tokens, device=dev, dtype=torch.int64) % max_pos

    q_pe = torch.randn(
        num_tokens, cfg.num_heads, cfg.rope_dim, device=dev, dtype=torch.bfloat16
    )
    ql_nope = torch.randn(
        num_tokens, cfg.num_heads, cfg.kv_lora, device=dev, dtype=torch.bfloat16
    )
    index_q = torch.randn(
        num_tokens,
        cfg.index_heads,
        cfg.index_head_dim,
        device=dev,
        dtype=torch.bfloat16,
    )
    index_w = torch.randn(num_tokens, cfg.index_heads, device=dev, dtype=torch.float32)
    q_scale = torch.tensor([0.37], device=dev, dtype=torch.float32)
    softmax_scale = cfg.index_head_dim**-0.5
    head_scale = cfg.index_heads**-0.5
    q_cos_sin = make_cos_sin(max_pos, cfg.rope_dim, dev)  # q_pe: interleaved
    idx_cos_sin = make_cos_sin(max_pos, cfg.rope_dim, dev)

    iq_fp8, iw_out, mqa = K.fused_q(
        pos,
        q_pe,
        q_cos_sin,
        index_q,
        idx_cos_sin,
        ql_nope,
        q_scale,
        index_w,
        softmax_scale,
        head_scale,
        has_indexer=True,
        index_rope_interleave=index_interleave,
    )

    s = q_scale.item()
    # MQA query: [ql_nope | q_pe RoPE'd (interleaved)], per-tensor fp8.
    mqa_nope_ref = (ql_nope.float() / s).to(FP8)
    qpe_ref = rope(
        q_pe.float(),
        pos.unsqueeze(-1).expand(num_tokens, cfg.num_heads),
        q_cos_sin,
        interleave=True,
    )
    mqa_pe_ref = (qpe_ref / s).to(FP8)
    assert_fp8(mqa[:, :, : cfg.kv_lora], mqa_nope_ref, "mqa ql_nope")
    assert_fp8(mqa[:, :, cfg.kv_lora :], mqa_pe_ref, "mqa q_pe")

    # Indexer-Q: RoPE + UE8M0 fp8 quant; index weights fold in q-scale.
    iq_ref = rope(
        index_q.float(),
        pos.unsqueeze(-1).expand(num_tokens, cfg.index_heads),
        idx_cos_sin,
        interleave=index_interleave,
    )
    q_ref, scale_ref = ue8m0_quant(iq_ref)
    assert_fp8(iq_fp8, q_ref, "indexer-Q fp8")
    iw_ref = index_w * scale_ref * softmax_scale * head_scale
    torch.testing.assert_close(iw_out, iw_ref, rtol=1e-3, atol=1e-3)


@pytest.mark.parametrize("cfg", MODEL_CONFIGS, ids=MODEL_IDS)
@pytest.mark.parametrize("num_tokens", [1, 17, 512])
def test_fused_q_no_indexer(num_tokens: int, cfg: ModelConfig):
    torch.manual_seed(3)
    dev = "cuda"
    max_pos = 8192
    pos = torch.arange(num_tokens, device=dev, dtype=torch.int64)
    q_pe = torch.randn(
        num_tokens, cfg.num_heads, cfg.rope_dim, device=dev, dtype=torch.bfloat16
    )
    ql_nope = torch.randn(
        num_tokens, cfg.num_heads, cfg.kv_lora, device=dev, dtype=torch.bfloat16
    )
    q_scale = torch.tensor([0.5], device=dev, dtype=torch.float32)
    q_cos_sin = make_cos_sin(max_pos, cfg.rope_dim, dev)
    _, _, mqa = K.fused_q(
        pos,
        q_pe,
        q_cos_sin,
        None,
        None,
        ql_nope,
        q_scale,
        None,
        0.0,
        0.0,
        has_indexer=False,
        index_rope_interleave=False,
    )
    s = q_scale.item()
    assert_fp8(mqa[:, :, : cfg.kv_lora], (ql_nope.float() / s).to(FP8), "mqa ql_nope")
    qpe_ref = rope(
        q_pe.float(),
        pos.unsqueeze(-1).expand(num_tokens, cfg.num_heads),
        q_cos_sin,
        interleave=True,
    )
    assert_fp8(mqa[:, :, cfg.kv_lora :], (qpe_ref / s).to(FP8), "mqa q_pe")


@pytest.mark.parametrize("cfg", MODEL_CONFIGS, ids=MODEL_IDS)
@pytest.mark.parametrize("num_tokens", [1, 17, 512])
@pytest.mark.parametrize("has_indexer", [True, False])
@pytest.mark.parametrize("use_mha", [True, False])
@pytest.mark.parametrize(
    "index_interleave,cache_dtype", [(False, torch.float32), (True, torch.bfloat16)]
)
@pytest.mark.parametrize("capture", [False, True])
def test_fused_q_bf16_query(
    num_tokens: int,
    has_indexer: bool,
    use_mha: bool,
    index_interleave: bool,
    cache_dtype: torch.dtype,
    capture: bool,
    cfg: ModelConfig,
):
    """bf16-query path (FlashMLA sparse, SM90/SM100): only the RoPE'd q_pe is
    produced; MHA does not need an absorbed ql_nope projection."""
    torch.manual_seed(6)
    dev = "cuda"
    max_pos = 8192
    pos = torch.arange(num_tokens, device=dev, dtype=torch.int64) % max_pos

    # Attention takes q_pe as a strided slice of the full Q projection.
    q_pe = torch.randn(
        num_tokens, cfg.num_heads, 4 * cfg.rope_dim, device=dev, dtype=torch.bfloat16
    )[..., -cfg.rope_dim :]
    ql_nope = torch.randn(
        num_tokens, cfg.num_heads, cfg.kv_lora, device=dev, dtype=torch.bfloat16
    )
    q_scale = torch.tensor([0.37], device=dev, dtype=torch.float32)
    q_cos_sin = make_cos_sin(max_pos, cfg.rope_dim, dev).to(cache_dtype)

    index_q = index_w = idx_cos_sin = None
    if has_indexer:
        index_q = torch.randn(
            num_tokens,
            cfg.index_heads,
            cfg.index_head_dim,
            device=dev,
            dtype=torch.bfloat16,
        )
        index_w = torch.randn(
            num_tokens, cfg.index_heads, device=dev, dtype=torch.float32
        )
        idx_cos_sin = make_cos_sin(max_pos, cfg.rope_dim, dev).to(cache_dtype)

    run = partial(
        K.fused_q,
        pos,
        q_pe,
        q_cos_sin,
        index_q,
        idx_cos_sin,
        q_pe if use_mha else ql_nope,
        q_scale,
        index_w,
        cfg.index_head_dim**-0.5,
        cfg.index_heads**-0.5,
        has_indexer=has_indexer,
        index_rope_interleave=index_interleave,
        quantize_mqa=False,
    )

    iq_fp8, iw_out, q_pe_out = run()
    if capture:
        graph = torch.cuda.CUDAGraph()
        with torch.cuda.graph(graph):
            iq_fp8, iw_out, q_pe_out = run()
        q_pe.mul_(0.5)
        if index_q is not None:
            index_q.mul_(1.5)
        graph.replay()

    # MQA query: only the RoPE'd q_pe, bf16, unquantized.
    assert q_pe_out.dtype == torch.bfloat16
    assert q_pe_out.shape == (num_tokens, cfg.num_heads, cfg.rope_dim)
    qpe_ref = rope(
        q_pe.float(),
        pos.unsqueeze(-1).expand(num_tokens, cfg.num_heads),
        q_cos_sin,
        interleave=True,
    )
    assert_bf16(q_pe_out, qpe_ref, "bf16 q_pe RoPE")

    # Indexer-Q is unchanged on this path (still UE8M0 fp8 + folded weights).
    if has_indexer:
        assert index_q is not None
        iq_ref = rope(
            index_q.float(),
            pos.unsqueeze(-1).expand(num_tokens, cfg.index_heads),
            idx_cos_sin,
            interleave=index_interleave,
        )
        q_ref, scale_ref = ue8m0_quant(iq_ref)
        assert_fp8(iq_fp8, q_ref, "indexer-Q fp8 (bf16-query path)")
        iw_ref = (
            index_w * scale_ref * (cfg.index_head_dim**-0.5) * (cfg.index_heads**-0.5)
        )
        torch.testing.assert_close(iw_out, iw_ref, rtol=1e-3, atol=1e-3)


def test_fused_q_triton_supports_large_token_count():
    """Keep the token count off CUDA grid-y in the Triton fallback.

    The minimal dimensions also bypass CuTeDSL on SM100.
    """
    num_tokens = 65536
    dev = "cuda"
    dtype = torch.bfloat16
    positions = torch.zeros(num_tokens, device=dev, dtype=torch.int64)
    q_pe = torch.ones((num_tokens, 1, 2), device=dev, dtype=dtype)
    ql_nope = torch.ones((num_tokens, 1, 1), device=dev, dtype=dtype)
    cos_sin = torch.tensor([[1.0, 0.0]], device=dev, dtype=torch.float32)
    q_scale = torch.ones(1, device=dev, dtype=torch.float32)

    _, _, mqa_q = K.fused_q(
        positions,
        q_pe,
        cos_sin,
        None,
        None,
        ql_nope,
        q_scale,
        None,
        0.0,
        0.0,
        has_indexer=False,
    )

    rows = torch.tensor([0, num_tokens - 1], device=dev)
    ref = torch.cat([ql_nope, q_pe], dim=-1).to(FP8)
    assert_fp8(mqa_q[rows], ref[rows], "large-token fused Q")


# ── fused_eh_norm (MTP) ──────────────────────────────────────────────────────


@pytest.mark.parametrize("cfg", MODEL_CONFIGS, ids=MODEL_IDS)
@pytest.mark.parametrize("num_tokens", [1, 4, 17, 512])
def test_fused_eh_norm(num_tokens: int, cfg: ModelConfig):
    torch.manual_seed(4)
    dev = "cuda"
    # Mix in a position-0 token to exercise the embeds-zeroing branch.
    pos = torch.arange(num_tokens, device=dev, dtype=torch.int64)
    pos[0] = 0
    embeds = torch.randn(num_tokens, cfg.hidden, device=dev, dtype=torch.bfloat16)
    prev = torch.randn(num_tokens, cfg.hidden, device=dev, dtype=torch.bfloat16)
    ew = torch.randn(cfg.hidden, device=dev, dtype=torch.bfloat16)
    hw = torch.randn(cfg.hidden, device=dev, dtype=torch.bfloat16)

    out = K.fused_eh_norm(pos, embeds, prev, ew, hw, EPS)

    masked = torch.where(pos.unsqueeze(-1) == 0, torch.zeros_like(embeds), embeds)
    ref = torch.cat([rms_norm(masked, ew), rms_norm(prev, hw)], dim=-1)
    assert out.shape == (num_tokens, 2 * cfg.hidden)
    assert_bf16(out, ref, "eh_norm")
