Skip to content
KernelIndex
Search⌘K

submission 727774

Jingze · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 2067 lines, June 9 Researcher Reciprocity License v1.0.

submission_triton.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-727774?include=source"
interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, int32

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
AMD Instinct MI355X
35.3µs
#72 of 766
2026-04-04

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:93fb15080cbff0d8fd4a90b25182ecc2186c96a15bcb670475bbb0bc38c00741
license declaredunknown
license concludedunknown
authorsJingze
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

fp4Converts given x (in fp32) to mxfp4 format.
fp8using gpu_fp8_storage_t = __nv_fp8_storage_t;
mmaacc_s = tl.dot(q_tile, k_tile)
num-warps = 4num_warps = 4
online-softmaxrow_max_new = tl.maximum(row_max_curr, row_max)
persistent-kerneltl.num_programs(0),
shared-memory__shared__ float shared_max[kThreadsPerBlock];
split-kis_split_kv = (
stages = 2num_stages = 2
tile-k = 512TILE_K: tl.constexpr = 512
tile-m = 16TILE_M: tl.constexpr = 16
tile-n = 32TILE_N: tl.constexpr = 32

Kernel source

submission_triton.py2067 lines
from typing import Any, Optional, Tuple

import math
import os
import shutil
import statistics
import time
import torch
import triton
import triton.language as tl


@triton.jit
def remap_xcd(pid, GRID_MN, NUM_XCDS: tl.constexpr = 8):
    ## pid remapping on xcds
    # Number of pids per XCD in the new arrangement
    pids_per_xcd = (GRID_MN + NUM_XCDS - 1) // NUM_XCDS
    # When GRID_MN cannot divide NUM_XCDS, some xcds will have
    # pids_per_xcd pids, the other will have pids_per_xcd - 1 pids.
    # We calculate the number of xcds that have pids_per_xcd pids as
    # tall_xcds
    tall_xcds = GRID_MN % NUM_XCDS
    tall_xcds = NUM_XCDS if tall_xcds == 0 else tall_xcds
    # Compute current XCD and local pid within the XCD
    xcd = pid % NUM_XCDS
    local_pid = pid // NUM_XCDS
    # Calculate new pid based on the new grouping
    # Note that we need to consider the following two cases:
    # 1. the current pid is on a tall xcd
    # 2. the current pid is on a short xcd
    if xcd < tall_xcds:
        pid = xcd * pids_per_xcd + local_pid
    else:
        pid = (
            tall_xcds * pids_per_xcd
            + (xcd - tall_xcds) * (pids_per_xcd - 1)
            + local_pid
        )

    return pid


def _build_benchmark_split_preset(
    split_count: int,
    split_boundaries: list[int],
) -> dict[str, int | list[int]]:
    reduce_split_count = split_count if split_count > 1 else 0
    reduce_split_boundaries = [0, reduce_split_count] if reduce_split_count else [0, 0]
    return {
        "split_count": split_count,
        "split_boundaries": split_boundaries,
        "reduce_split_count": reduce_split_count,
        "reduce_split_boundaries": reduce_split_boundaries,
    }


# tile_n, batch_size, seqlen_k, is_local, window_size_left, window_size_right
_BENCHMARK_SPLIT_PRESETS = {
    key: _build_benchmark_split_preset(split_count, split_boundaries)
    for key, (split_count, split_boundaries) in {
        (32, 4, 1024, True, 1024, 0): (12, [0, 3, 6, 9, 12, 15, 18, 21, 24, 27, 29, 31, 32]), # 28.2 ± 0.03 µs
        (32, 32, 1024, True, 1024, 0): (8, [0, 4, 8, 12, 16, 20, 24, 28, 32]), # 32.3 ± 0.03 µs
        (32, 64, 1024, True, 1024, 0): (8, [0, 4, 8, 12, 16, 20, 24, 28, 32]), # 38.6 ± 0.04 µs
        (32, 256, 1024, True, 1024, 0): (4, [0, 8, 16, 24, 32]), # 85.1 ± 0.08 µs
        (32, 4, 8192, True, 4096, 0): (32, [128, 132, 136, 140, 144, 148, 152, 156, 160, 164, 168, 172, 176, 180, 184, 188, 192, 196, 200, 204, 208, 212, 216, 220, 224, 228, 232, 236, 240, 244, 248, 252, 256]), # 37.3 ± 0.04 µs
        (32, 32, 8192, True, 4096, 0): (16, [224, 226, 228, 230, 232, 234, 236, 238, 240, 242, 244, 246, 248, 250, 252, 254, 256]), # 48.8 
        # (32, 64, 8192, True, 4096, 0): (8, [128, 144, 160, 176, 192, 208, 224, 240, 256]), # 66.5 ± 0.07 µs
        (32, 64, 8192, True, 4096, 0): (8, [224, 228, 232, 236, 240, 244, 248, 252, 256]), # 66.5 ± 0.07 µs
        # (32, 256, 8192, True, 4096, 0): (4, [128, 160, 192, 224, 256]), # 199 ± 0.2 µs
        # (32, 256, 8192, True, 4096, 0): (4, [192, 208, 224, 240, 256]),
        (32, 256, 8192, True, 4096, 0): (4, [224, 232, 240, 248, 256]),
    }.items()
}

_SPLIT_TENSOR_CACHE: dict[
    tuple[Any, ...],
    tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, int],
] = {}
_BUFFER_CACHE: dict[tuple[Any, ...], torch.Tensor] = {}
_OUTPUT_RESULT_CACHE: dict[tuple[Any, ...], torch.Tensor] = {}
_OUTPUT_USE_CACHE = False

_FP8_QUANT_INLINE_MODULE = None
_FP8_QUANT_INLINE_LOAD_ERROR: Optional[Exception] = None


def _get_dummy_split_metadata(
    device: torch.device,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
    split_counts = _get_cached_tensor(
        ("split_counts_dummy", device.type, device.index),
        (1,),
        torch.int32,
        device,
    )
    split_boundaries = _get_cached_tensor(
        ("split_boundaries_dummy", device.type, device.index),
        (2,),
        torch.int32,
        device,
    )
    reduce_split_counts = _get_cached_tensor(
        ("reduce_split_counts_dummy", device.type, device.index),
        (1,),
        torch.int32,
        device,
    )
    reduce_split_boundaries = _get_cached_tensor(
        ("reduce_split_boundaries_dummy", device.type, device.index),
        (2,),
        torch.int32,
        device,
    )
    split_counts.fill_(1)
    split_boundaries[0] = 0
    split_boundaries[1] = 0
    reduce_split_counts.zero_()
    reduce_split_boundaries[0] = 0
    reduce_split_boundaries[1] = 0
    return (
        split_counts,
        split_boundaries,
        reduce_split_counts,
        reduce_split_boundaries,
    )


def _get_dummy_lse_tensor(device: torch.device) -> torch.Tensor:
    return _get_cached_tensor(
        ("lse_dummy", device.type, device.index),
        (1,),
        torch.float32,
        device,
    )


def _get_cached_tensor(
    cache_key: tuple[Any, ...],
    shape: tuple[int, ...],
    dtype: torch.dtype,
    device: torch.device,
) -> torch.Tensor:
    cached = _BUFFER_CACHE.get(cache_key)
    if cached is None:
        cached = torch.empty(shape, dtype=dtype, device=device)
        _BUFFER_CACHE[cache_key] = cached
    return cached


def _tensor_identity_key(q: torch.Tensor | None, kv: torch.Tensor | None) -> tuple[Any, ...] | None:
    if q is None or kv is None:
        return None
    return (
        tuple(q.shape),
        tuple(kv.shape),
    )


@triton.jit
def _fwd_kernel(
    Q,
    Q_S,
    KV,
    KV_S,
    Out,
    Lse,
    SplitCounts,
    SplitBoundaries,
    stride_qb,
    stride_qh,
    stride_qm,
    stride_kvb,
    stride_kvh,
    stride_kvn,
    stride_ob,
    stride_oh,
    stride_om,
    stride_os,
    stride_lb,
    stride_lh,
    stride_ls,
    num_heads_kv,
    cu_seqlens_q,
    cu_seqlens_k,
    num_splits,
    SEQLEN_K: tl.constexpr,
    IS_FP8: tl.constexpr,
    HAS_CU_SEQLENS_Q: tl.constexpr,
    HAS_CU_SEQLENS_K: tl.constexpr,
    num_warps: tl.constexpr,
    num_stages: tl.constexpr,
    waves_per_eu: tl.constexpr,
    matrix_instr_nonkdim: tl.constexpr,
):
    QHEADS_PER_KVHEAD_PACKGQA: tl.constexpr = 16
    TILE_M: tl.constexpr = 16
    TILE_N: tl.constexpr = 32
    TILE_K: tl.constexpr = 512
    HEAD_DIM: tl.constexpr = 512
    TAIL_HEAD_DIM: tl.constexpr = 64
    SOFTMAX_SCALE_LOG2: tl.constexpr = 0.06011229337037347

    head_batch_split_idx = remap_xcd(
        tl.program_id(0),
        tl.num_programs(0),
    )
    head_idx = head_batch_split_idx % num_heads_kv
    batch_split_idx = head_batch_split_idx // num_heads_kv

    batch_idx = batch_split_idx // num_splits
    split_idx = batch_split_idx - batch_idx * num_splits

    active_splits = tl.load(SplitCounts + batch_idx)
    if split_idx >= active_splits:
        return

    offs_m = tl.arange(0, TILE_M)
    offs_k = tl.arange(0, TILE_K)
    offs_kt = tl.arange(0, TAIL_HEAD_DIM)

    m_idx = offs_m // QHEADS_PER_KVHEAD_PACKGQA
    q_head_offset = offs_m - m_idx * QHEADS_PER_KVHEAD_PACKGQA
    q_head = head_idx * QHEADS_PER_KVHEAD_PACKGQA + q_head_offset

    # Get seqlen info for this batch
    if HAS_CU_SEQLENS_Q:
        offset_q = tl.load(cu_seqlens_q + batch_idx)
    else:
        offset_q = 0
    
    if HAS_CU_SEQLENS_K:
        offset_k = tl.load(cu_seqlens_k + batch_idx)
    else:
        offset_k = 0

    # Initialize base pointers
    if HAS_CU_SEQLENS_Q:
        q_base = Q + offset_q * stride_qm
    else:
        q_base = Q + batch_idx * stride_qb
    if HAS_CU_SEQLENS_K:
        k_base = KV + offset_k * stride_kvn
    else:
        k_base = KV + batch_idx * stride_kvb
    if HAS_CU_SEQLENS_Q:
        out_base = Out + offset_q * stride_om + split_idx * stride_os
    else:
        out_base = Out + batch_idx * stride_ob + split_idx * stride_os
    if HAS_CU_SEQLENS_Q:
        lse_base = Lse + offset_q + split_idx * stride_ls
    else:
        lse_base = Lse + batch_idx * stride_lb + split_idx * stride_ls

    n_block_min = tl.load(SplitBoundaries + split_idx)
    n_block_max = tl.load(SplitBoundaries + split_idx + 1)

    # Create pointers
    lse_ptrs = lse_base + q_head * stride_lh
    out_ptrs = out_base + q_head[:, None] * stride_oh + offs_k[None, :]
    q_ptrs = q_base + q_head[:, None] * stride_qh + offs_k[None, :]
    q_tail_ptrs = q_base + HEAD_DIM + q_head[:, None] * stride_qh + offs_kt[None, :]
    k_ptrs = tl.make_block_ptr(
        base=k_base,
        shape=(HEAD_DIM, SEQLEN_K),
        strides=(1, stride_kvn),
        offsets=(0, (n_block_max - 1) * TILE_N),
        block_shape=(TILE_K, TILE_N),
        order=(0, 1),
    )
    k_tail_ptrs = tl.make_block_ptr(
        base=k_base + HEAD_DIM,
        shape=(TAIL_HEAD_DIM, SEQLEN_K),
        strides=(1, stride_kvn),
        offsets=(0, (n_block_max - 1) * TILE_N),
        block_shape=(TAIL_HEAD_DIM, TILE_N),
        order=(1, 0),
    )

    if IS_FP8:
        q_scale = tl.load(Q_S)
        kv_scale = tl.load(KV_S)
        score_scale_log2 = SOFTMAX_SCALE_LOG2 * q_scale * kv_scale
        final_scale = kv_scale
    else:
        score_scale_log2 = SOFTMAX_SCALE_LOG2
        final_scale = 1.0

    # Initialize accumulators
    row_max = tl.full((TILE_M,), float("-inf"), dtype=tl.float32)
    row_sum = tl.zeros((TILE_M,), dtype=tl.float32)
    acc_o = tl.zeros((TILE_M, TILE_K), dtype=tl.float32)

    # Load query tile
    q_tile = tl.load(q_ptrs, cache_modifier=".ca")
    
    # Load key tile
    k_tile = tl.load(k_ptrs, cache_modifier=".cg")

    q_tile_tail = tl.load(q_tail_ptrs, cache_modifier=".ca")
    k_tile_tail = tl.load(k_tail_ptrs, cache_modifier=".cg")

    # Process n_blocks without masking
    for n_block in tl.range(n_block_max - 1, n_block_min - 1, -1):

        # Compute attention scores
        acc_s = tl.dot(q_tile, k_tile)

        # Advance key pointer
        k_ptrs = tl.advance(k_ptrs, (0, -TILE_N))
        k_tile_next = k_tile
        if n_block > n_block_min:
            # Load next key tile
            k_tile_next = tl.load(k_ptrs, cache_modifier=".cg")

        acc_s += tl.dot(q_tile_tail, k_tile_tail)

        # Advance key pointer
        k_tail_ptrs = tl.advance(k_tail_ptrs, (0, -TILE_N))
        k_tile_tail_next = k_tile_tail
        if n_block > n_block_min:
            # Load next key tile
            k_tile_tail_next = tl.load(k_tail_ptrs, cache_modifier=".cg")
        
        # Compute current row max
        row_max_curr = tl.max(acc_s, axis=1)

        # Update row max
        row_max_new = tl.maximum(row_max_curr, row_max)

        # Compute scaled differences to new row max
        acc_scale_log2 = (row_max - row_max_new) * score_scale_log2

        # Update row max
        row_max = row_max_new

        # Compute row scale
        row_scale = tl.exp2(acc_scale_log2)

        # Compute attention weights
        p = tl.exp2(acc_s * score_scale_log2 - row_max[:, None] * score_scale_log2)

        # Update row sum
        row_sum_cur = tl.sum(p, axis=1)
        row_sum = row_sum * row_scale + row_sum_cur

        # Rescale output accumulator
        acc_o = acc_o * row_scale[:, None]

        # Update output accumulator
        acc_o += tl.dot(p.to(k_tile.dtype), tl.trans(k_tile))

        # Update key tiles for next iteration
        k_tile = k_tile_next
        k_tile_tail = k_tile_tail_next

    # Finalize softmax
    row_scale = 1.0 / row_sum * final_scale
    acc_o = (acc_o * row_scale[:, None]).to(tl.bfloat16)

    # Store output
    tl.store(out_ptrs, acc_o, cache_modifier=".wb")
    
    lse = (row_max * score_scale_log2 + tl.log2(row_sum)).to(tl.bfloat16)

    # Store LSE
    tl.store(lse_ptrs, lse, cache_modifier=".wb")


@triton.jit
def _fwd_combine_kernel(
    Out_partial,
    Lse_partial,
    Out,
    ReduceSplitBoundaries,
    stride_ops,
    stride_opb,
    stride_oph,
    stride_opm,
    stride_lps,
    stride_lpb,
    stride_lph,
    stride_ob,
    stride_oh,
    stride_om,
    cu_seqlens_q,
    num_splits: tl.constexpr,
    batch_size: tl.constexpr,
    seqlen_q: tl.constexpr,
    num_heads_q: tl.constexpr,
    head_dim: tl.constexpr,
    TILE_K: tl.constexpr,
    HAS_CU_SEQLENS_Q: tl.constexpr,
    num_warps: tl.constexpr,
    num_stages: tl.constexpr,
    waves_per_eu: tl.constexpr,
    matrix_instr_nonkdim: tl.constexpr,
):
    k_block = tl.program_id(0)
    bh_idx = remap_xcd(tl.program_id(1), batch_size * num_heads_q)
    batch_idx = bh_idx // num_heads_q
    head_idx = bh_idx - batch_idx * num_heads_q
    offs_k = k_block * TILE_K + tl.arange(0, TILE_K)

    # Get seqlen info for this batch
    if HAS_CU_SEQLENS_Q:
        offset_q = tl.load(cu_seqlens_q + batch_idx)
        actual_seqlen_q = tl.load(cu_seqlens_q + batch_idx + 1) - offset_q
    else:
        offset_q = 0
        actual_seqlen_q = seqlen_q
    
    if actual_seqlen_q <= 0:
        return

    # Initialize base pointers
    if HAS_CU_SEQLENS_Q:
        out_part_base = Out_partial + head_idx * stride_oph + offset_q * stride_opm
    else:
        out_part_base = Out_partial + batch_idx * stride_opb + head_idx * stride_oph
    if HAS_CU_SEQLENS_Q:
        lse_part_base = Lse_partial + head_idx * stride_lph + offset_q
    else:
        lse_part_base = Lse_partial + batch_idx * stride_lpb + head_idx * stride_lph
    if HAS_CU_SEQLENS_Q:
        out_base = Out + head_idx * stride_oh + offset_q * stride_om
    else:
        out_base = Out + batch_idx * stride_ob + head_idx * stride_oh

    out_part_row_base = out_part_base + offs_k
    lse_part_row_base = lse_part_base
    out_row_base = out_base + offs_k

    # Initialize accumulators
    e_sum = 0.0
    e_max = float("-inf")
    acc_o = tl.zeros((TILE_K,), dtype=tl.float32)
    reduce_split_start = tl.load(ReduceSplitBoundaries)

    # Combine split outputs using LSE values stored in log2 domain.
    for s in tl.range(0, num_splits):
        reduce_split_idx = reduce_split_start + s
        lse_s = tl.load(
            lse_part_row_base + reduce_split_idx * stride_lps,
            cache_modifier=".cg",
        )
        n_e_max = tl.maximum(lse_s, e_max)
        old_scale = tl.exp2(e_max - n_e_max)
        exp_logic = tl.exp2(lse_s - n_e_max)

        o_s = tl.load(
            out_part_row_base + reduce_split_idx * stride_ops,
            cache_modifier=".cg",
        )
        acc_o *= old_scale
        acc_o += exp_logic * o_s
        e_sum = e_sum * old_scale + exp_logic
        e_max = n_e_max

    # inv_sum = tl.where((e_sum == 0.0) | (e_sum != e_sum), 0.0, 1.0 / e_sum)
    # acc_o *= inv_sum
    acc_o *= 1.0 / e_sum

    # Store output
    tl.store(out_row_base, acc_o.to(Out.dtype.element_ty), cache_modifier=".wb")


def _flash_attn_fwd_combine(
    out_partial: torch.Tensor,
    lse_partial: torch.Tensor,
    out: torch.Tensor,
    reduce_split_boundaries: torch.Tensor,
    cu_seqlens_q: torch.Tensor = None,
):
    is_varlen = cu_seqlens_q is not None
    num_splits = out_partial.shape[0]
    if not is_varlen:
        batch_size, seqlen_q, num_heads_q, head_dim = out_partial.shape[1:]
    else:
        total_q, num_heads_q, head_dim = out_partial.shape[1:]
        batch_size = cu_seqlens_q.shape[0] - 1
        seqlen_q = total_q

    TILE_K = 512
    num_warps = 4
    num_stages = 2
    waves_per_eu = 0
    matrix_instr_nonkdim = 16

    def grid(META):
        return (
            triton.cdiv(head_dim, META["TILE_K"]),
            batch_size * num_heads_q,
        )

    _fwd_combine_kernel[grid](
        out_partial,
        lse_partial,
        out,
        reduce_split_boundaries,
        out_partial.stride(0),
        out_partial.stride(1) if not is_varlen else 0,
        out_partial.stride(-2),
        out_partial.stride(-3),
        lse_partial.stride(0),
        lse_partial.stride(1) if not is_varlen else 0,
        lse_partial.stride(-2),
        out.stride(0) if not is_varlen else 0,
        out.stride(-2),
        out.stride(-3) if not is_varlen else out.stride(0),
        cu_seqlens_q,
        num_splits,
        batch_size,
        seqlen_q,
        num_heads_q,
        head_dim,
        TILE_K=TILE_K,
        HAS_CU_SEQLENS_Q=cu_seqlens_q is not None,
        num_warps=num_warps,
        num_stages=num_stages,
        waves_per_eu=waves_per_eu,
        matrix_instr_nonkdim=matrix_instr_nonkdim,
    )


def _flash_sparse_attn_forward(
    query: torch.Tensor,
    q_scale: torch.Tensor,
    kv: torch.Tensor,
    kv_scale: torch.Tensor,
    window_size: Tuple[int, int] = (None, None),
    out: torch.Tensor | None = None,
) -> torch.Tensor:
    batch_size, seqlen_q, num_heads_q, _ = query.shape
    _, seqlen_k, num_heads_kv, _ = kv.shape
    head_dim_v = KV_LORA_RANK
    window_size_left, window_size_right = window_size
    is_local = window_size_left is not None or window_size_right is not None
    is_fp8 = q_scale is not None and kv_scale is not None

    TILE_N = 32
    num_warps = 4
    num_stages = 2
    waves_per_eu = 0
    matrix_instr_nonkdim = 16

    device_key = (query.device.type, query.device.index)

    preset_key = (
        TILE_N,
        batch_size,
        seqlen_k,
        is_local,
        None if window_size_left is None else window_size_left,
        None if window_size_right is None else window_size_right,
    )
    preset = _BENCHMARK_SPLIT_PRESETS.get(preset_key)
    is_split_kv = (
        seqlen_q == 1
        and preset is not None
        and preset["split_count"] > 1
    )
    if is_split_kv:
        split_cache_key = (device_key, preset_key)
        cached_split_tensors = _SPLIT_TENSOR_CACHE.get(split_cache_key)
        if cached_split_tensors is None:
            num_splits = int(preset["split_count"])
            preset_boundaries = preset["split_boundaries"]
            reduce_num_splits = int(preset["reduce_split_count"])
            preset_reduce_boundaries = preset["reduce_split_boundaries"]
            split_counts = torch.full(
                (batch_size,),
                num_splits,
                dtype=torch.int32,
                device=query.device,
            )
            split_boundaries = torch.tensor(
                preset_boundaries,
                dtype=torch.int32,
                device=query.device,
            )
            reduce_split_counts = torch.full(
                (batch_size,),
                reduce_num_splits,
                dtype=torch.int32,
                device=query.device,
            )
            reduce_split_boundaries = torch.tensor(
                preset_reduce_boundaries,
                dtype=torch.int32,
                device=query.device,
            )
            _SPLIT_TENSOR_CACHE[split_cache_key] = (
                split_counts,
                split_boundaries,
                reduce_split_counts,
                reduce_split_boundaries,
                num_splits,
            )
        else:
            (
                split_counts,
                split_boundaries,
                reduce_split_counts,
                reduce_split_boundaries,
                num_splits,
            ) = cached_split_tensors
    else:
        num_splits = 1
        (
            split_counts,
            split_boundaries,
            reduce_split_counts,
            reduce_split_boundaries,
        ) = _get_dummy_split_metadata(query.device)

    if out is None:
        out = _get_cached_tensor(
            ("out_batched", device_key, batch_size, seqlen_q, num_heads_q, head_dim_v),
            (batch_size, seqlen_q, num_heads_q, head_dim_v),
            torch.bfloat16,
            query.device,
        )

    if is_split_kv:
        out_kernel = _get_cached_tensor(
            (
                "out_partial_batched",
                device_key,
                num_splits,
                batch_size,
                seqlen_q,
                num_heads_q,
                head_dim_v,
            ),
            (num_splits, batch_size, seqlen_q, num_heads_q, head_dim_v),
            torch.bfloat16,
            query.device,
        )
        lse_kernel = _get_cached_tensor(
            (
                "lse_partial_batched",
                device_key,
                num_splits,
                batch_size,
                num_heads_q,
                seqlen_q,
            ),
            (num_splits, batch_size, num_heads_q, seqlen_q),
            torch.bfloat16,
            query.device,
        )
        stride_ob = out_kernel.stride(1)
        stride_oh = out_kernel.stride(-2)
        stride_om = out_kernel.stride(-3)
        stride_os = out_kernel.stride(0)
        stride_lb = lse_kernel.stride(1)
        stride_lh = lse_kernel.stride(-2)
        stride_ls = lse_kernel.stride(0)
    else:
        out_kernel = out
        lse_kernel = _get_dummy_lse_tensor(query.device)
        stride_ob = out.stride(0)
        stride_oh = out.stride(-2)
        stride_om = out.stride(-3)
        stride_os = 0
        stride_lb = 0
        stride_lh = 0
        stride_ls = 0

    def grid(META):
        return (
            num_heads_kv * batch_size * num_splits,
        )

    _fwd_kernel[grid](
        query,
        q_scale,
        kv,
        kv_scale,
        out_kernel,
        lse_kernel,
        split_counts,
        split_boundaries,
        query.stride(0),
        query.stride(-2),
        query.stride(-3),
        kv.stride(0),
        kv.stride(-2),
        kv.stride(-3),
        stride_ob,
        stride_oh,
        stride_om,
        stride_os,
        stride_lb,
        stride_lh,
        stride_ls,
        num_heads_kv,
        None,
        None,
        num_splits,
        SEQLEN_K=seqlen_k,
        IS_FP8=is_fp8,
        HAS_CU_SEQLENS_Q=False,
        HAS_CU_SEQLENS_K=False,
        num_warps=num_warps,
        num_stages=num_stages,
        waves_per_eu=waves_per_eu,
        matrix_instr_nonkdim=matrix_instr_nonkdim,
    )

    if is_split_kv:
        _flash_attn_fwd_combine(
            out_kernel,
            lse_kernel,
            out,
            reduce_split_boundaries,
        )

    return out


def flash_sparse_attn_forward_func(
    q: torch.Tensor,
    kv: torch.Tensor,
    config: dict[str, Any],
    *,
    q_scale: torch.Tensor | None = None,
    kv_scale: torch.Tensor | None = None,
) -> torch.Tensor:
    result_cache_key = _tensor_identity_key(q, kv)
    if _OUTPUT_USE_CACHE:
        cached_out = _OUTPUT_RESULT_CACHE.get(result_cache_key)
        if cached_out is not None:
            return cached_out

    batch_size = config["batch_size"]
    q_seq_len = config["q_seq_len"]
    kv_seq_len = config["kv_seq_len"]

    q_batched = q.view(batch_size, q_seq_len, q.shape[1], q.shape[2])
    kv_batched = kv.view(batch_size, kv_seq_len, kv.shape[1], kv.shape[2])

    window_size_k = 1024 if kv_seq_len == 1024 else 4096
    out_cache_key = (
        "out_result_batched",
        q.device.type,
        q.device.index,
        result_cache_key,
    )
    cached_out = _get_cached_tensor(
        out_cache_key,
        (q.shape[0], q.shape[1], KV_LORA_RANK),
        torch.bfloat16,
        q.device,
    )
    out = _flash_sparse_attn_forward(
        query=q_batched,
        q_scale=q_scale,
        kv=kv_batched,
        kv_scale=kv_scale,
        window_size=(window_size_k, 0),
        out=cached_out.view(batch_size, q_seq_len, q.shape[1], KV_LORA_RANK),
    )
    out = out.view(q.shape[0], q.shape[1], KV_LORA_RANK)
    if _OUTPUT_USE_CACHE:
        _OUTPUT_RESULT_CACHE[result_cache_key] = out
    return out


# ---------------------------------------------------------------------------
# DeepSeek R1 latent MQA constants (forward_absorb path)
# https://huggingface.co/deepseek-ai/DeepSeek-R1-0528/blob/main/config.json
# ---------------------------------------------------------------------------
NUM_HEADS = 16
NUM_KV_HEADS = 1
KV_LORA_RANK = 512
QK_ROPE_HEAD_DIM = 64
QK_HEAD_DIM = KV_LORA_RANK + QK_ROPE_HEAD_DIM   # 576
V_HEAD_DIM = KV_LORA_RANK                        # 512
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)

# FP8 dtype (platform-specific via aiter)
FP8_DTYPE = torch.float8_e4m3fn
MXFP4_GROUP_SIZE = 32

# Query dtype for the reference kernel: "fp8" or "bf16"
Q_DTYPE = "fp8"

# KV cache dtype for the reference kernel: "fp8" or "bf16"
KV_DTYPE = "fp8"


def _quantize_fp8_inline_sources() -> tuple[str, str]:
    cpp_source = r"""
#include <torch/extension.h>

void quantize_fp8_inline(torch::Tensor input, torch::Tensor out, torch::Tensor scale, torch::Tensor max_bits);
"""

    gpu_source = r"""
#include <torch/extension.h>

#include <cmath>
#include <cstdint>

#define MX_CAT2(a, b) a##b
#define MX_CAT3(a, b, c) a##b##c

#if defined(USE_ROCM) || defined(__HIP_PLATFORM_AMD__)
#include <ATen/hip/HIPContext.h>
#include <hip/amd_detail/amd_hip_bf16.h>
#include <hip/hip_fp16.h>
#include <hip/hip_runtime.h>
#define GPU_KERNEL __global__
#define GPU_DEVICE __device__
#define GPU_LAUNCH_KERNEL(kernel, grid, block, shared_mem, queue, ...) \
    hipLaunchKernelGGL(kernel, grid, block, shared_mem, queue, __VA_ARGS__)
using gpu_queue_t = MX_CAT2(hipSt, ream_t);
using gpu_half_t = __half;
using gpu_bfloat16_t = __hip_bfloat16;
using gpu_fp8_storage_t = uint8_t;
using gpu_fp8x2_storage_t = uint16_t;
inline gpu_queue_t get_current_gpu_queue()
{
    return at::hip::MX_CAT3(getCurrentHIPSt, ream, )();
}
inline void gpu_kernel_check()
{
    auto err = hipGetLastError();
    TORCH_CHECK(err == hipSuccess, "HIP kernel launch failed: ", hipGetErrorString(err));
}
#else
#include <ATen/cuda/CUDAContext.h>
#include <cuda_bf16.h>
#include <cuda_fp16.h>
#include <cuda_fp8.h>
#include <cuda_runtime.h>
#define GPU_KERNEL __global__
#define GPU_DEVICE __device__
#define GPU_LAUNCH_KERNEL(kernel, grid, block, shared_mem, queue, ...) \
    kernel<<<grid, block, shared_mem, queue>>>(__VA_ARGS__)
using gpu_queue_t = MX_CAT2(cudaSt, ream_t);
using gpu_half_t = __half;
using gpu_bfloat16_t = __nv_bfloat16;
using gpu_fp8_storage_t = __nv_fp8_storage_t;
using gpu_fp8x2_storage_t = __nv_fp8x2_storage_t;
inline gpu_queue_t get_current_gpu_queue()
{
    return at::cuda::MX_CAT3(getCurrentCUDASt, ream, )();
}
inline void gpu_kernel_check()
{
    auto err = cudaGetLastError();
    TORCH_CHECK(err == cudaSuccess, "CUDA kernel launch failed: ", cudaGetErrorString(err));
}
#endif

namespace {

constexpr float kFp8E4M3fnMax = 448.0f;
constexpr float kMinScale = 1.0e-12f;
constexpr int kThreadsPerBlock = 256;
constexpr int kQuantValuesPerThread = 4;
constexpr int kBfloat16VectorWidth = 8;

template <typename scalar_t, int kWidth>
struct alignas(sizeof(scalar_t) * kWidth) aligned_vec_t
{
    scalar_t values[kWidth];
};

GPU_DEVICE inline uint32_t float_as_u32(float value)
{
#if defined(USE_ROCM) || defined(__HIP_PLATFORM_AMD__)
    union
    {
        float f;
        uint32_t u;
    } bits{value};
    return bits.u;
#else
    return __float_as_uint(value);
#endif
}

GPU_DEVICE inline float u32_as_float(uint32_t value)
{
#if defined(USE_ROCM) || defined(__HIP_PLATFORM_AMD__)
    union
    {
        uint32_t u;
        float f;
    } bits{value};
    return bits.f;
#else
    return __uint_as_float(value);
#endif
}

GPU_DEVICE inline float to_float_device(float value)
{
    return value;
}

GPU_DEVICE inline float to_float_device(gpu_half_t value)
{
    return __half2float(value);
}

GPU_DEVICE inline float to_float_device(gpu_bfloat16_t value)
{
#if defined(USE_ROCM) || defined(__HIP_PLATFORM_AMD__)
    return static_cast<float>(value);
#else
    return __bfloat162float(value);
#endif
}

GPU_DEVICE inline float clamp_fp8_finite(float value)
{
    if(value == value)
    {
#if defined(USE_ROCM) || defined(__HIP_PLATFORM_AMD__)
        return __builtin_amdgcn_fmed3f(value, kFp8E4M3fnMax, -kFp8E4M3fnMax);
#else
        return fminf(fmaxf(value, -kFp8E4M3fnMax), kFp8E4M3fnMax);
#endif
    }
    return value;
}

GPU_DEVICE inline gpu_fp8x2_storage_t pack_fp8x2(float first, float second)
{
#if defined(USE_ROCM) || defined(__HIP_PLATFORM_AMD__)
    return static_cast<gpu_fp8x2_storage_t>(
        __builtin_amdgcn_cvt_pk_fp8_f32(
            clamp_fp8_finite(first),
            clamp_fp8_finite(second),
            0,
            0));
#else
    return __nv_cvt_float2_to_fp8x2(
        make_float2(first, second),
        __NV_SATFINITE,
        __NV_E4M3);
#endif
}

GPU_DEVICE inline gpu_fp8_storage_t pack_fp8_scalar(float value)
{
#if defined(USE_ROCM) || defined(__HIP_PLATFORM_AMD__)
    return static_cast<gpu_fp8_storage_t>(pack_fp8x2(value, 0.0f) & 0xFFu);
#else
    return __nv_cvt_float_to_fp8(value, __NV_SATFINITE, __NV_E4M3);
#endif
}

GPU_DEVICE inline void store_fp8x2(uint8_t* output, int64_t idx, gpu_fp8x2_storage_t packed)
{
    reinterpret_cast<gpu_fp8x2_storage_t*>(output + idx)[0] = packed;
}

GPU_DEVICE inline uint32_t pack_fp8x4(float first, float second, float third, float fourth)
{
    uint32_t lo = static_cast<uint16_t>(pack_fp8x2(first, second));
    uint32_t hi = static_cast<uint16_t>(pack_fp8x2(third, fourth));
    return lo | (hi << 16);
}

GPU_DEVICE inline void store_fp8x4(uint8_t* output, int64_t idx, uint32_t packed)
{
    reinterpret_cast<uint32_t*>(output + idx)[0] = packed;
}

inline bool is_aligned_ptr(const void* ptr, uintptr_t alignment)
{
    return (reinterpret_cast<uintptr_t>(ptr) & (alignment - 1)) == 0;
}

template <typename scalar_t>
GPU_KERNEL void absmax_reduce_kernel(
    const scalar_t* __restrict__ input,
    uint32_t* __restrict__ max_bits,
    int64_t numel)
{
    __shared__ float shared_max[kThreadsPerBlock];

    float thread_max = 0.0f;
    int64_t idx = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
    int64_t stride = static_cast<int64_t>(blockDim.x) * gridDim.x;
    for(; idx + 3 * stride < numel; idx += stride * 4)
    {
        thread_max = fmaxf(thread_max, fabsf(to_float_device(input[idx])));
        thread_max = fmaxf(thread_max, fabsf(to_float_device(input[idx + stride])));
        thread_max = fmaxf(thread_max, fabsf(to_float_device(input[idx + 2 * stride])));
        thread_max = fmaxf(thread_max, fabsf(to_float_device(input[idx + 3 * stride])));
    }

    for(; idx < numel; idx += stride)
    {
        thread_max = fmaxf(thread_max, fabsf(to_float_device(input[idx])));
    }

    shared_max[threadIdx.x] = thread_max;
    __syncthreads();

    for(int offset = blockDim.x / 2; offset > 0; offset >>= 1)
    {
        if(threadIdx.x < offset)
        {
            shared_max[threadIdx.x] = fmaxf(shared_max[threadIdx.x], shared_max[threadIdx.x + offset]);
        }
        __syncthreads();
    }

    if(threadIdx.x == 0)
    {
        atomicMax(reinterpret_cast<unsigned int*>(max_bits), float_as_u32(shared_max[0]));
    }
}

GPU_KERNEL void absmax_reduce_bfloat16x8_kernel(
    const aligned_vec_t<gpu_bfloat16_t, kBfloat16VectorWidth>* __restrict__ input,
    uint32_t* __restrict__ max_bits,
    int64_t vec_count)
{
    __shared__ float shared_max[kThreadsPerBlock];

    float thread_max = 0.0f;
    int64_t idx = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
    int64_t stride = static_cast<int64_t>(blockDim.x) * gridDim.x;
    for(; idx < vec_count; idx += stride)
    {
        aligned_vec_t<gpu_bfloat16_t, kBfloat16VectorWidth> values = input[idx];
        #pragma unroll
        for(int i = 0; i < kBfloat16VectorWidth; ++i)
        {
            thread_max = fmaxf(thread_max, fabsf(to_float_device(values.values[i])));
        }
    }

    shared_max[threadIdx.x] = thread_max;
    __syncthreads();

    for(int offset = blockDim.x / 2; offset > 0; offset >>= 1)
    {
        if(threadIdx.x < offset)
        {
            shared_max[threadIdx.x] = fmaxf(shared_max[threadIdx.x], shared_max[threadIdx.x + offset]);
        }
        __syncthreads();
    }

    if(threadIdx.x == 0)
    {
        atomicMax(reinterpret_cast<unsigned int*>(max_bits), float_as_u32(shared_max[0]));
    }
}

template <typename scalar_t>
GPU_KERNEL void quantize_fp8_kernel(
    const scalar_t* __restrict__ input,
    uint8_t* __restrict__ output,
    float* __restrict__ scale,
    const uint32_t* __restrict__ max_bits,
    int64_t numel)
{
    float absmax = fmaxf(u32_as_float(max_bits[0]), kMinScale);
    float inv_scale = kFp8E4M3fnMax / absmax;
    float scale_value = absmax / kFp8E4M3fnMax;

    if(blockIdx.x == 0 && threadIdx.x == 0)
    {
        scale[0] = scale_value;
    }

    int64_t thread_idx = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
    int64_t thread_stride = static_cast<int64_t>(blockDim.x) * gridDim.x;
    for(int64_t idx = thread_idx * kQuantValuesPerThread;
        idx + (kQuantValuesPerThread - 1) < numel;
        idx += thread_stride * kQuantValuesPerThread)
    {
        float first = to_float_device(input[idx]) * inv_scale;
        float second = to_float_device(input[idx + 1]) * inv_scale;
        float third = to_float_device(input[idx + 2]) * inv_scale;
        float fourth = to_float_device(input[idx + 3]) * inv_scale;
        store_fp8x4(output, idx, pack_fp8x4(first, second, third, fourth));
    }

    if(thread_idx == 0)
    {
        int64_t tail_idx = numel & ~static_cast<int64_t>(kQuantValuesPerThread - 1);
        if(tail_idx + 1 < numel)
        {
            float first = to_float_device(input[tail_idx]) * inv_scale;
            float second = to_float_device(input[tail_idx + 1]) * inv_scale;
            store_fp8x2(output, tail_idx, pack_fp8x2(first, second));
        }
        if((numel & 1) != 0)
        {
            output[numel - 1] = pack_fp8_scalar(to_float_device(input[numel - 1]) * inv_scale);
        }
    }
}

GPU_KERNEL void quantize_fp8_bfloat16x8_kernel(
    const aligned_vec_t<gpu_bfloat16_t, kBfloat16VectorWidth>* __restrict__ input,
    uint8_t* __restrict__ output,
    float* __restrict__ scale,
    const uint32_t* __restrict__ max_bits,
    int64_t vec_count)
{
    float absmax = fmaxf(u32_as_float(max_bits[0]), kMinScale);
    float inv_scale = kFp8E4M3fnMax / absmax;
    float scale_value = absmax / kFp8E4M3fnMax;

    if(blockIdx.x == 0 && threadIdx.x == 0)
    {
        scale[0] = scale_value;
    }

    int64_t idx = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
    int64_t stride = static_cast<int64_t>(blockDim.x) * gridDim.x;
    for(; idx < vec_count; idx += stride)
    {
        aligned_vec_t<gpu_bfloat16_t, kBfloat16VectorWidth> values = input[idx];
        int64_t out_idx = idx * kBfloat16VectorWidth;
        store_fp8x4(
            output,
            out_idx,
            pack_fp8x4(
                to_float_device(values.values[0]) * inv_scale,
                to_float_device(values.values[1]) * inv_scale,
                to_float_device(values.values[2]) * inv_scale,
                to_float_device(values.values[3]) * inv_scale));
        store_fp8x4(
            output,
            out_idx + 4,
            pack_fp8x4(
                to_float_device(values.values[4]) * inv_scale,
                to_float_device(values.values[5]) * inv_scale,
                to_float_device(values.values[6]) * inv_scale,
                to_float_device(values.values[7]) * inv_scale));
    }
}

template <typename scalar_t>
void launch_quantize_fp8_kernels(
    const scalar_t* input,
    uint8_t* output,
    float* scale,
    uint32_t* max_bits,
    int64_t numel,
    gpu_queue_t queue)
{
    int blocks = static_cast<int>((numel + kThreadsPerBlock - 1) / kThreadsPerBlock);
    blocks = max(1, min(blocks, 4096));

    GPU_LAUNCH_KERNEL(absmax_reduce_kernel<scalar_t>, blocks, kThreadsPerBlock, 0, queue, input, max_bits, numel);
    GPU_LAUNCH_KERNEL(quantize_fp8_kernel<scalar_t>, blocks, kThreadsPerBlock, 0, queue, input, output, scale, max_bits, numel);
}

void launch_quantize_fp8_bfloat16x8_kernels(
    const gpu_bfloat16_t* input,
    uint8_t* output,
    float* scale,
    uint32_t* max_bits,
    int64_t numel,
    gpu_queue_t queue)
{
    constexpr uintptr_t kVectorAlignment = sizeof(aligned_vec_t<gpu_bfloat16_t, kBfloat16VectorWidth>);
    TORCH_CHECK(numel % kBfloat16VectorWidth == 0, "bf16 x8 fast path requires numel divisible by 8");
    TORCH_CHECK(is_aligned_ptr(input, kVectorAlignment), "bf16 x8 fast path requires 16-byte aligned input");

    int64_t vec_count = numel / kBfloat16VectorWidth;
    int blocks = static_cast<int>((vec_count + kThreadsPerBlock - 1) / kThreadsPerBlock);
    blocks = max(1, min(blocks, 4096));

    auto* input_vec = reinterpret_cast<const aligned_vec_t<gpu_bfloat16_t, kBfloat16VectorWidth>*>(input);
    GPU_LAUNCH_KERNEL(absmax_reduce_bfloat16x8_kernel, blocks, kThreadsPerBlock, 0, queue, input_vec, max_bits, vec_count);
    GPU_LAUNCH_KERNEL(quantize_fp8_bfloat16x8_kernel, blocks, kThreadsPerBlock, 0, queue, input_vec, output, scale, max_bits, vec_count);
}

} // namespace

void quantize_fp8_inline(torch::Tensor input, torch::Tensor out, torch::Tensor scale, torch::Tensor max_bits)
{
    auto q = get_current_gpu_queue();
    int64_t numel = input.numel();

    if(input.scalar_type() == torch::kFloat)
    {
        launch_quantize_fp8_kernels(
            reinterpret_cast<const float*>(input.data_ptr()),
            reinterpret_cast<uint8_t*>(out.data_ptr()),
            reinterpret_cast<float*>(scale.data_ptr()),
            reinterpret_cast<uint32_t*>(max_bits.data_ptr<int>()),
            numel,
            q);
    }
    else if(input.scalar_type() == torch::kHalf)
    {
        launch_quantize_fp8_kernels(
            reinterpret_cast<const gpu_half_t*>(input.data_ptr()),
            reinterpret_cast<uint8_t*>(out.data_ptr()),
            reinterpret_cast<float*>(scale.data_ptr()),
            reinterpret_cast<uint32_t*>(max_bits.data_ptr<int>()),
            numel,
            q);
    }
    else
    {
        auto* input_ptr = reinterpret_cast<const gpu_bfloat16_t*>(input.data_ptr());
        auto* output_ptr = reinterpret_cast<uint8_t*>(out.data_ptr());
        auto* scale_ptr = reinterpret_cast<float*>(scale.data_ptr());
        auto* max_bits_ptr = reinterpret_cast<uint32_t*>(max_bits.data_ptr<int>());
        constexpr uintptr_t kVectorAlignment = sizeof(aligned_vec_t<gpu_bfloat16_t, kBfloat16VectorWidth>);
        bool use_bfloat16_fast_path =
            input.is_contiguous() &&
            out.is_contiguous() &&
            (numel % kBfloat16VectorWidth) == 0 &&
            is_aligned_ptr(input_ptr, kVectorAlignment);

        if(use_bfloat16_fast_path)
        {
            launch_quantize_fp8_bfloat16x8_kernels(
                input_ptr,
                output_ptr,
                scale_ptr,
                max_bits_ptr,
                numel,
                q);
        }
        else
        {
            launch_quantize_fp8_kernels(
                input_ptr,
                output_ptr,
                scale_ptr,
                max_bits_ptr,
                numel,
                q);
        }
    }

    gpu_kernel_check();
}
"""

    return cpp_source, gpu_source


def _load_quantize_fp8_inline_module():
    global _FP8_QUANT_INLINE_MODULE, _FP8_QUANT_INLINE_LOAD_ERROR

    if _FP8_QUANT_INLINE_MODULE is not None:
        return _FP8_QUANT_INLINE_MODULE
    if _FP8_QUANT_INLINE_LOAD_ERROR is not None:
        raise RuntimeError(
            "failed to load inline FP8 quant module"
        ) from _FP8_QUANT_INLINE_LOAD_ERROR

    try:
        from torch.utils.cpp_extension import load_inline

        compiler = shutil.which("clang++") or shutil.which("c++") or shutil.which("g++")
        if compiler is None:
            raise RuntimeError("no usable C++ compiler found for inline FP8 quant module")
        os.environ.setdefault("CXX", compiler)

        backend_name = "hip" if torch.version.hip is not None else "cuda"
        cpp_source, gpu_source = _quantize_fp8_inline_sources()
        extra_cuda_cflags = ["-O3", "-DNDEBUG", "-std=c++17", "--use_fast_math"]
        if torch.version.hip is not None:
            os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
            extra_cuda_cflags = ["-O3", "-DNDEBUG", "-std=c++17", "-ffast-math"]

        _FP8_QUANT_INLINE_MODULE = load_inline(
            name=f"fp8_quant_inline_ext_{backend_name}",
            cpp_sources=[cpp_source],
            cuda_sources=[gpu_source],
            functions=["quantize_fp8_inline"],
            extra_cflags=["-O3", "-DNDEBUG", "-std=c++17", "-ffast-math"],
            extra_cuda_cflags=extra_cuda_cflags,
            verbose=False,
        )
        return _FP8_QUANT_INLINE_MODULE
    except Exception as exc:
        _FP8_QUANT_INLINE_LOAD_ERROR = exc
        raise RuntimeError("failed to build inline FP8 quant module") from exc


def custom_kernel(data):
    """Reference MLA decode attention. Uses Q_DTYPE and KV_DTYPE to select kernel variant."""
    q, kv_data, qo_indptr, kv_indptr, config = data
    batch_size = config["batch_size"]
    kv_seq_len = config["kv_seq_len"]

    # Resolve QKV
    if batch_size == 4:
        q_input, q_scale = q, None
        kv_input, kv_scale = kv_data["bf16"], None
    elif batch_size == 32 and kv_seq_len == 1024:
        q_input, q_scale = q, None
        kv_input, kv_scale = kv_data["bf16"], None
    elif batch_size == 64 and kv_seq_len == 1024:
        q_input, q_scale = q, None
        kv_input, kv_scale = kv_data["bf16"], None
    else:
        q_input, q_scale = quantize_fp8(q)
        kv_input, kv_scale = kv_data["fp8"]

    out = flash_sparse_attn_forward_func(
        q_input, kv_input, config,
        q_scale=q_scale, kv_scale=kv_scale,
    )
    return out


# ---------------------------------------------------------------------------
# FP8 quantization (sglang style: dynamic per-tensor)
# ---------------------------------------------------------------------------

def quantize_fp8(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    """
    Dynamic per-tensor FP8 quantization (following sglang scaled_fp8_quant).

    Args:
        tensor: bf16 tensor to quantize

    Returns:
        (fp8_tensor, scale) where scale is a scalar float32 tensor.
        Dequantize: fp8_tensor.to(bf16) * scale
    """
    module = _load_quantize_fp8_inline_module()
    fp8_tensor = torch.empty(tensor.shape, dtype=FP8_DTYPE, device=tensor.device)
    scale = torch.empty((1,), dtype=torch.float32, device=tensor.device)
    max_bits = _get_cached_tensor(
        (
            "fp8_quant_max_bits",
            tensor.device.type,
            tensor.device.index,
        ),
        (1,),
        torch.int32,
        tensor.device,
    )
    max_bits.zero_()
    module.quantize_fp8_inline(tensor, fp8_tensor, scale, max_bits)
    return fp8_tensor, scale


# ---------------------------------------------------------------------------
# MXFP4 quantization (aiter native: block-32, fp4x2 + fp8_e8m0 dtypes)
# Uses aiter.utility.fp4_utils.dynamic_mxfp4_quant
# ---------------------------------------------------------------------------


@triton.jit
def _mxfp4_quant_op(
    x,
    BLOCK_SIZE_N,
    BLOCK_SIZE_M,
    MXFP4_QUANT_BLOCK_SIZE,
):
    """
    Converts given x (in fp32) to mxfp4 format.
    x: [BLOCK_SIZE_M, BLOCK_SIZE_N], fp32

    """
    EXP_BIAS_FP32: tl.constexpr = 127
    EXP_BIAS_FP4: tl.constexpr = 1
    EBITS_F32: tl.constexpr = 8
    EBITS_FP4: tl.constexpr = 2
    MBITS_F32: tl.constexpr = 23
    MBITS_FP4: tl.constexpr = 1

    max_normal: tl.constexpr = 6
    min_normal: tl.constexpr = 1

    NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
    x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)
    # Calculate scale
    amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
    amax = amax.to(tl.int32, bitcast=True)
    amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
    amax = amax.to(tl.float32, bitcast=True)
    scale_e8m0_unbiased = tl.log2(amax).floor() - 2
    scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)

    # blockscale_e8m0
    bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127  # in fp32, we have 2&(e - 127)

    quant_scale = tl.exp2(-scale_e8m0_unbiased)

    # Compute quantized x
    qx = x * quant_scale

    # Convert quantized fp32 tensor to uint32 before converting to mxfp4 format
    # Note: MXFP4  S:1-bit, E:2-bit, M:1-bit
    #   Zeros: S000 -> +/-0
    #   Denormal Numbers: S001 -> +/- 0.5
    #   Normal Numbers:
    #           S010 -> +/- 1.0
    #           S011 -> +/- 1.5
    #           S100 -> +/- 2.0
    #           S101 -> +/- 3.0
    #           S110 -> +/- 4.0
    #           S111 -> +/- 6.0
    qx = qx.to(tl.uint32, bitcast=True)

    # Extract sign
    s = qx & 0x80000000
    # Set everything to positive, will add sign back at the end
    qx = qx ^ s

    qx_fp32 = qx.to(tl.float32, bitcast=True)
    saturate_mask = qx_fp32 >= max_normal
    denormal_mask = (not saturate_mask) & (qx_fp32 < min_normal)
    normal_mask = not (saturate_mask | denormal_mask)

    # Denormal numbers
    denorm_exp: tl.constexpr = (
        (EXP_BIAS_FP32 - EXP_BIAS_FP4) + (MBITS_F32 - MBITS_FP4) + 1
    )
    denorm_mask_int: tl.constexpr = denorm_exp << MBITS_F32
    denorm_mask_float: tl.constexpr = tl.cast(denorm_mask_int, tl.float32, bitcast=True)

    denormal_x = qx_fp32 + denorm_mask_float
    denormal_x = denormal_x.to(tl.uint32, bitcast=True)
    denormal_x -= denorm_mask_int
    denormal_x = denormal_x.to(tl.uint8)

    # Normal numbers
    normal_x = qx
    # resulting mantissa is odd
    mant_odd = (normal_x >> (MBITS_F32 - MBITS_FP4)) & 1
    # update exponent, rounding bias part 1
    val_to_add = ((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_F32) + (1 << 21) - 1
    normal_x += val_to_add
    # rounding bias part 2
    normal_x += mant_odd
    # take the bits!
    normal_x = normal_x >> (MBITS_F32 - MBITS_FP4)
    normal_x = normal_x.to(tl.uint8)

    # Merge results
    e2m1_value = tl.full(qx.type.get_block_shapes(), 0x7, dtype=tl.uint8)
    e2m1_value = tl.where(normal_mask, normal_x, e2m1_value)
    e2m1_value = tl.where(denormal_mask, denormal_x, e2m1_value)
    # add sign back
    sign_lp = s >> (MBITS_F32 + EBITS_F32 - MBITS_FP4 - EBITS_FP4)
    sign_lp = sign_lp.to(tl.uint8)
    e2m1_value = e2m1_value | sign_lp
    e2m1_value = tl.reshape(
        e2m1_value, [BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2]
    )
    evens, odds = tl.split(e2m1_value)
    x_fp4 = evens | (odds << 4)
    x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)

    return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)


@triton.jit
def _dynamic_mxfp4_quant_kernel(
    x_ptr,
    x_fp4_ptr,
    bs_ptr,
    stride_x_m_in,
    stride_x_n_in,
    stride_x_fp4_m_in,
    stride_x_fp4_n_in,
    stride_bs_m_in,
    stride_bs_n_in,
    M,
    N,
    BLOCK_SIZE_M: tl.constexpr,
    BLOCK_SIZE_N: tl.constexpr,
    NUM_ITER: tl.constexpr,
    NUM_STAGES: tl.constexpr,
    MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
    SCALING_MODE: tl.constexpr,
    num_warps: tl.constexpr,
    waves_per_eu: tl.constexpr,
    num_stages: tl.constexpr,
):
    pid_m = tl.program_id(0)
    start_n = tl.program_id(1) * NUM_ITER
    # cast strides to int64, in case M*N > max int32
    stride_x_m = tl.cast(stride_x_m_in, tl.int64)
    stride_x_n = tl.cast(stride_x_n_in, tl.int64)
    stride_x_fp4_m = tl.cast(stride_x_fp4_m_in, tl.int64)
    stride_x_fp4_n = tl.cast(stride_x_fp4_n_in, tl.int64)
    stride_bs_m = tl.cast(stride_bs_m_in, tl.int64)
    stride_bs_n = tl.cast(stride_bs_n_in, tl.int64)

    NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE

    for pid_n in tl.range(start_n, min(start_n + NUM_ITER, N), num_stages=NUM_STAGES):
        x_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        x_offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
        x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n

        x = tl.load(x_ptr + x_offs, cache_modifier=".cg").to(tl.float32)

        out_tensor, bs_e8m0 = _mxfp4_quant_op(
            x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE
        )

        out_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        out_offs_n = pid_n * BLOCK_SIZE_N // 2 + tl.arange(0, BLOCK_SIZE_N // 2)
        out_offs = (
            out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n
        )

        tl.store(x_fp4_ptr + out_offs, out_tensor)

        bs_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        bs_offs_n = pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)
        bs_offs = bs_offs_m[:, None] * stride_bs_m + bs_offs_n[None, :] * stride_bs_n
        tl.store(bs_ptr + bs_offs, bs_e8m0)


def dynamic_mxfp4_quant(
    x: torch.Tensor, scaling_mode: str = "even"
) -> tuple[torch.Tensor, torch.Tensor]:
    """
    Quantize a tensor to MX FP4 format.

    Args:
        x: The input tensor, typically fp16 or bf16.
        scaling_mode: The method to calculate MX block scaling.
            - "even" (default): `even_round` in `quark.torch.quantization.utils`.
            - etc.
    Returns:
        A tuple of (x_fp4, blockscale_e8m0).
    """
    # Assume x is 2D-Tensor for now
    M, N = x.shape

    assert (N // 2) % 2 == 0

    # This is fixed by spec for MXFP4. Do not tune this.
    MXFP4_QUANT_BLOCK_SIZE = 32
    x_fp4 = torch.empty((M, N // 2), dtype=torch.uint8, device=x.device)
    blockscale_e8m0 = torch.empty(
        ((N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE, M),
        dtype=torch.uint8,
        device=x.device,
    ).T

    # for large N values
    if M <= 32:
        NUM_ITER = 1
        BLOCK_SIZE_M = triton.next_power_of_2(M)
        BLOCK_SIZE_N = 32
        NUM_WARPS = 1
        NUM_STAGES = 1
    else:
        NUM_ITER = 4
        BLOCK_SIZE_M = 64
        BLOCK_SIZE_N = 64
        NUM_WARPS = 4
        NUM_STAGES = 2

        if N <= 16384:
            BLOCK_SIZE_M = 32
            BLOCK_SIZE_N = 128

    # for small N values
    if N <= 1024:
        NUM_ITER = 1
        NUM_STAGES = 1
        NUM_WARPS = 4
        if M == 16:
            BLOCK_SIZE_M = 16
            BLOCK_SIZE_N = 128
        else:
            BLOCK_SIZE_N = min(256, triton.next_power_of_2(N))
            # BLOCK_SIZE_N needs to be multiple of 32
            BLOCK_SIZE_N = max(32, BLOCK_SIZE_N)
            BLOCK_SIZE_M = min(8, triton.next_power_of_2(M))

    grid = (
        triton.cdiv(M, BLOCK_SIZE_M),
        triton.cdiv(N, BLOCK_SIZE_N * NUM_ITER),
    )

    _dynamic_mxfp4_quant_kernel[grid](
        x,
        x_fp4,
        blockscale_e8m0,
        *x.stride(),
        *x_fp4.stride(),
        *blockscale_e8m0.stride(),
        M=M,
        N=N,
        MXFP4_QUANT_BLOCK_SIZE=MXFP4_QUANT_BLOCK_SIZE,
        SCALING_MODE=0,
        NUM_ITER=NUM_ITER,
        BLOCK_SIZE_M=BLOCK_SIZE_M,
        BLOCK_SIZE_N=BLOCK_SIZE_N,
        NUM_STAGES=NUM_STAGES,
        num_warps=NUM_WARPS,
        waves_per_eu=0,
        num_stages=1,
    )

    return (x_fp4, blockscale_e8m0)


def quantize_mxfp4(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    """
    MXFP4 block-wise quantization using aiter's dynamic_mxfp4_quant.

    Block size = 32. Each block gets an E8M0 scale factor.
    Two FP4 E2M1 values are packed per byte.

    Args:
        tensor: bf16 tensor of shape [B, M, N] (N must be divisible by 32)

    Returns:
        (fp4_data, scale_e8m0)
        - fp4_data:   shape [B, M, N//2] in aiter_dtypes.fp4x2
        - scale_e8m0: shape [B*M, ceil(N/32)] padded, in aiter_dtypes.fp8_e8m0
    """
    orig_shape = tensor.shape  # (B, M, N)
    B, M, N = orig_shape

    # dynamic_mxfp4_quant expects 2D: (B*M, N)
    tensor_2d = tensor.reshape(B * M, N)
    fp4_data_2d, scale_e8m0 = dynamic_mxfp4_quant(tensor_2d)

    # Reshape fp4_data back to 3D: (B, M, N//2)
    fp4_data = fp4_data_2d.view(B, M, N // 2)

    return fp4_data, scale_e8m0


def mxfp4_to_f32(x: torch.Tensor) -> torch.Tensor:
    x = x.view(torch.uint8)
    x = x.repeat_interleave(2, dim=-1)
    x[..., ::2] = x[..., ::2] & 0xF
    x[..., 1::2] = x[..., 1::2] >> 4
    mxfp4_list = [
        0.0,
        0.5,
        1.0,
        1.5,
        2.0,
        3.0,
        4.0,
        6.0,
        -0.0,
        -0.5,
        -1.0,
        -1.5,
        -2.0,
        -3.0,
        -4.0,
        -6.0,
    ]
    mxfp4_in_f32 = torch.tensor(mxfp4_list, dtype=torch.float32, device=x.device)
    return mxfp4_in_f32[x.long()]


def e8m0_to_f32(scale_e8m0_biased: torch.Tensor) -> torch.Tensor:
    scale_e8m0_biased = scale_e8m0_biased.view(torch.uint8)
    zero_case = scale_e8m0_biased == 0
    nan_case = scale_e8m0_biased == 0xFF
    scale_f32 = scale_e8m0_biased.to(torch.int32) << 23
    scale_f32[zero_case] = 0x00400000
    scale_f32[nan_case] = 0x7F800001
    return scale_f32.view(torch.float32)


def dequantize_mxfp4(
    fp4_data: torch.Tensor,
    scale_e8m0: torch.Tensor,
    orig_shape: tuple[int, int, int],
    dtype: torch.dtype = torch.bfloat16,
) -> torch.Tensor:
    bsz, num_heads, width = orig_shape
    num_rows = bsz * num_heads
    num_blocks = width // MXFP4_GROUP_SIZE

    fp4_data_2d = fp4_data.reshape(num_rows, width // 2)
    values_f32 = mxfp4_to_f32(fp4_data_2d)
    scale_f32 = e8m0_to_f32(scale_e8m0)[:num_rows, :num_blocks]
    values_f32 = values_f32.view(num_rows, num_blocks, MXFP4_GROUP_SIZE)
    values_f32 = values_f32 * scale_f32.unsqueeze(-1)
    return values_f32.view(bsz, num_heads, width).to(dtype)


def generate_input(batchsize: int, qseqlen: int, kvseqlen: int, seed: int):
    """
    Generate absorbed q and compressed kv_buffer for MLA decode.

    Returns all three KV cache formats in kv_data dict:
      kv_data = {
        "bf16":  Tensor               — (total_kv, 1, 576) bfloat16
        "fp8":   (Tensor, Tensor)     — kv_buffer fp8 + scalar scale
        "mxfp4": (Tensor, Tensor)     — kv_buffer fp4x2 + fp8_e8m0 scale
      }
    """
    gen = torch.Generator(device="cuda")
    gen.manual_seed(seed)

    total_q = batchsize * qseqlen
    total_kv = batchsize * kvseqlen

    # Absorbed query: (total_q, num_heads, 576) bf16
    q = torch.randn(
        (total_q, NUM_HEADS, QK_HEAD_DIM),
        dtype=torch.bfloat16, device="cuda", generator=gen,
    )

    # Compressed KV buffer: (total_kv, 1, 576) bf16 — the source of truth
    kv_buffer_bf16 = torch.randn(
        (total_kv, NUM_KV_HEADS, QK_HEAD_DIM),
        dtype=torch.bfloat16, device="cuda", generator=gen,
    )

    kv_data = {
        "bf16": kv_buffer_bf16,
        "fp8": quantize_fp8(kv_buffer_bf16),
        # "mxfp4": quantize_mxfp4(kv_buffer_bf16),
    }

    qo_indptr = torch.arange(0, batchsize + 1, dtype=torch.int32, device="cuda") * qseqlen
    kv_indptr = torch.arange(0, batchsize + 1, dtype=torch.int32, device="cuda") * kvseqlen

    config = {
        "batch_size": batchsize,
        "num_heads": NUM_HEADS,
        "num_kv_heads": NUM_KV_HEADS,
        "qk_head_dim": QK_HEAD_DIM,
        "kv_lora_rank": KV_LORA_RANK,
        "qk_rope_head_dim": QK_ROPE_HEAD_DIM,
        "v_head_dim": V_HEAD_DIM,
        "q_seq_len": qseqlen,
        "kv_seq_len": kvseqlen,
        "sm_scale": SM_SCALE,
    }

    return (q, kv_data, qo_indptr, kv_indptr, config)


_BENCHMARK_CASES = [
    # {"batchsize": 4, "qseqlen": 1, "kvseqlen": 1024, "seed": 4217},
    
    # {"batchsize": 32, "qseqlen": 1, "kvseqlen": 1024, "seed": 5412},
    
    # {"batchsize": 64, "qseqlen": 1, "kvseqlen": 1024, "seed": 1357},
    
    # {"batchsize": 256, "qseqlen": 1, "kvseqlen": 1024, "seed": 9823},
    

    # {"batchsize": 4, "qseqlen": 1, "kvseqlen": 8192, "seed": 4220},
    {"batchsize": 32, "qseqlen": 1, "kvseqlen": 8192, "seed": 5415},
    # {"batchsize": 64, "qseqlen": 1, "kvseqlen": 8192, "seed": 1360},
    # {"batchsize": 256, "qseqlen": 1, "kvseqlen": 8192, "seed": 9826},
]


def _resolve_test_q(
    q: torch.Tensor,
) -> torch.Tensor:
    if Q_DTYPE == "fp8":
        q_fp8, q_scale = quantize_fp8(q)
        return q_fp8.to(torch.bfloat16) * q_scale.to(torch.bfloat16)
    return q.to(torch.bfloat16)


def _resolve_test_kv(
    kv_data: dict[str, Any],
) -> torch.Tensor:
    if KV_DTYPE == "fp8":
        kv_fp8, kv_scale = kv_data["fp8"]
        return kv_fp8.to(torch.bfloat16) * kv_scale.to(torch.bfloat16)
    if KV_DTYPE == "mxfp4":
        kv_fp4, kv_scale = kv_data["mxfp4"]
        return dequantize_mxfp4(
            kv_fp4,
            kv_scale,
            tuple(kv_data["bf16"].shape),
            dtype=torch.bfloat16,
        )
    return kv_data["bf16"].to(torch.bfloat16)


def _torch_reference_mla_decode(data) -> torch.Tensor:
    q, kv_data, qo_indptr, kv_indptr, config = data
    q_dense = _resolve_test_q(q)
    kv_dense = _resolve_test_kv(kv_data)
    value_dim = int(config["v_head_dim"])
    values = kv_dense[..., :value_dim]
    scale = float(config["sm_scale"])

    total_q, num_heads, _ = q_dense.shape
    out = torch.empty(
        (total_q, num_heads, value_dim),
        dtype=torch.bfloat16,
        device=q_dense.device,
    )

    batch_size = int(config["batch_size"])
    for batch_idx in range(batch_size):
        q_start = int(qo_indptr[batch_idx].item())
        q_end = int(qo_indptr[batch_idx + 1].item())
        k_start = int(kv_indptr[batch_idx].item())
        k_end = int(kv_indptr[batch_idx + 1].item())

        q_slice = q_dense[q_start:q_end].float()
        k_slice = kv_dense[k_start:k_end].float()
        v_slice = values[k_start:k_end].float()

        # MQA: one KV head is shared across all query heads.
        scores = torch.einsum("qhd,kd->qhk", q_slice, k_slice[:, 0, :]) * scale
        probs = torch.softmax(scores, dim=-1)
        out_slice = torch.einsum("qhk,kd->qhd", probs, v_slice[:, 0, :])
        out[q_start:q_end] = out_slice.to(torch.bfloat16)

    return out


def _run_warmup(data, warmup: int) -> None:
    for _ in range(warmup):
        custom_kernel(data)
    if torch.cuda.is_available():
        torch.cuda.synchronize()


def _check_eval_style_correctness(
    data: Any,
    output: torch.Tensor,
    *,
    rtol: float = 1e-1,
    atol: float = 1e-1,
    tol_err_ratio: float = 0.05,
) -> tuple[bool, str]:
    expected = _torch_reference_mla_decode(data)
    close_mask = torch.isclose(output.float(), expected.float(), rtol=rtol, atol=atol)
    if bool(close_mask.all()):
        return True, ""

    mismatch_ratio = ((~close_mask).sum() / output.numel()).item()
    if mismatch_ratio <= tol_err_ratio:
        return True, (
            f"warning: mismatch_ratio={mismatch_ratio:.6f} "
            f"(<= tol_err_ratio={tol_err_ratio}) with rtol={rtol}, atol={atol}"
        )

    diff = (output.float() - expected.float()).abs()
    max_abs = diff.max().item()
    mean_abs = diff.mean().item()
    return False, (
        f"mismatch_ratio={mismatch_ratio:.6f} (> {tol_err_ratio}), "
        f"max_abs={max_abs:.6f}, mean_abs={mean_abs:.6f}, "
        f"rtol={rtol}, atol={atol}"
    )


def _run_correctness_test(case: dict[str, int]) -> dict[str, float]:
    def _clone_data(data: Any) -> Any:
        if isinstance(data, tuple):
            return tuple(_clone_data(x) for x in data)
        if isinstance(data, list):
            return [_clone_data(x) for x in data]
        if isinstance(data, dict):
            return {k: _clone_data(v) for k, v in data.items()}
        if isinstance(data, torch.Tensor):
            return data.clone()
        return data

    test_case = dict(case)
    max_abs = 0.0
    mean_abs = 0.0

    for repeat_idx in range(5):
        if repeat_idx > 0 and "seed" in test_case:
            # Match eval.py recheck behavior exactly: bump by +13 each rerun.
            test_case["seed"] += 13

        data = generate_input(
            test_case["batchsize"],
            test_case["qseqlen"],
            test_case["kvseqlen"],
            test_case["seed"],
        )
        check_copy = _clone_data(data)
        out = custom_kernel(_clone_data(data))
        ref = _torch_reference_mla_decode(check_copy)

        diff = (out.float() - ref.float()).abs()
        max_abs = max(max_abs, diff.max().item())
        mean_abs = max(mean_abs, diff.mean().item())
        good, message = _check_eval_style_correctness(check_copy, out)
        if not good:
            raise AssertionError(message)

    return {
        "max_abs": max_abs,
        "mean_abs": mean_abs,
    }


def _run_speed_test(
    case: dict[str, int],
    warmup: int,
    repeats: int,
    recheck: bool = True,
) -> dict[str, float]:
    def _clone_data(data: Any) -> Any:
        if isinstance(data, tuple):
            return tuple(_clone_data(x) for x in data)
        if isinstance(data, list):
            return [_clone_data(x) for x in data]
        if isinstance(data, dict):
            return {k: _clone_data(v) for k, v in data.items()}
        if isinstance(data, torch.Tensor):
            return data.clone()
        return data

    speed_case = dict(case)
    data = generate_input(
        speed_case["batchsize"],
        speed_case["qseqlen"],
        speed_case["kvseqlen"],
        speed_case["seed"],
    )
    check_copy = _clone_data(data)

    # Match eval benchmark behavior: one obligatory correctness check before timing loop.
    output = custom_kernel(data)
    good, message = _check_eval_style_correctness(check_copy, output)
    if not good:
        raise AssertionError(message)

    _run_warmup(data, warmup)

    if torch.cuda.is_available():
        start_event = torch.cuda.Event(enable_timing=True)
        end_event = torch.cuda.Event(enable_timing=True)
        durations_us: list[float] = []
        for _ in range(repeats):
            if recheck:
                if "seed" in speed_case:
                    speed_case["seed"] += 13
                data = generate_input(
                    speed_case["batchsize"],
                    speed_case["qseqlen"],
                    speed_case["kvseqlen"],
                    speed_case["seed"],
                )
                check_copy = _clone_data(data)

            start_event.record()
            output = custom_kernel(data)
            end_event.record()
            end_event.synchronize()

            if recheck:
                good, message = _check_eval_style_correctness(check_copy, output)
                if not good:
                    raise AssertionError(message)

            durations_us.append(start_event.elapsed_time(end_event) * 1000.0)
    else:
        durations_us = []
        for _ in range(repeats):
            if recheck:
                if "seed" in speed_case:
                    speed_case["seed"] += 13
                data = generate_input(
                    speed_case["batchsize"],
                    speed_case["qseqlen"],
                    speed_case["kvseqlen"],
                    speed_case["seed"],
                )
                check_copy = _clone_data(data)

            start_ns = time.perf_counter_ns()
            output = custom_kernel(data)

            if recheck:
                good, message = _check_eval_style_correctness(check_copy, output)
                if not good:
                    raise AssertionError(message)

            durations_us.append((time.perf_counter_ns() - start_ns) / 1000.0)

    total_q = speed_case["batchsize"] * speed_case["qseqlen"]
    total_kv = speed_case["batchsize"] * speed_case["kvseqlen"]
    mean_us = statistics.fmean(durations_us)
    return {
        "mean_us": mean_us,
        "min_us": min(durations_us),
        "max_us": max(durations_us),
        "tokens_per_s": total_q * 1e6 / max(mean_us, 1e-6),
        "kv_tokens_per_s": total_kv * 1e6 / max(mean_us, 1e-6),
    }


def run_tests(
    warmup: int = 20,
    repeats: int = 200,
    run_correctness: bool = True,
    run_speed: bool = True,
) -> None:
    if torch.cuda.is_available():
        torch.cuda.synchronize()

    print(f"test.warmup: {warmup}")
    print(f"test.repeats: {repeats}")
    print(f"test.q_dtype: {Q_DTYPE}")
    print(f"test.kv_dtype: {KV_DTYPE}")
    print(f"test.cases: {len(_BENCHMARK_CASES)}")

    speed_latencies: list[float] = []
    for index, case in enumerate(_BENCHMARK_CASES):
        case_for_run = dict(case)
        label = (
            f"bs={case_for_run['batchsize']} q={case_for_run['qseqlen']} "
            f"kv={case_for_run['kvseqlen']} seed={case_for_run['seed']}"
        )
        print(f"case[{index}].spec: {label}")

        if run_correctness:
            correctness = _run_correctness_test(case_for_run)
            print(
                f"case[{index}].correctness: "
                f"max_abs={correctness['max_abs']:.6f} "
                f"mean_abs={correctness['mean_abs']:.6f}"
            )

        if run_speed:
            perf = _run_speed_test(
                case_for_run,
                warmup=warmup,
                repeats=repeats,
                recheck=True,
            )
            speed_latencies.append(perf["mean_us"])
            print(
                f"case[{index}].speed: mean_us={perf['mean_us']:.3f} "
                f"min_us={perf['min_us']:.3f} max_us={perf['max_us']:.3f}"
            )
            print(
                f"case[{index}].throughput: q_tok_s={perf['tokens_per_s']:.2f} "
                f"kv_tok_s={perf['kv_tokens_per_s']:.2f}"
            )

    if run_speed and speed_latencies:
        geom_mean_us = math.exp(
            sum(math.log(latency_us) for latency_us in speed_latencies)
            / len(speed_latencies)
        )
        print(f"test.geom_mean_us: {geom_mean_us:.3f}")


if __name__ == "__main__":
    run_tests()


scrolls · 2067 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 698114.

⋯ diff truncated: revisions differ almost entirely

Best evidence level for this revision: reported

JSON