Skip to content
KernelIndex
Search⌘K

submission 692590

Jingze · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_sparse.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-692590?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
69.4µs
#312 of 766
2026-04-01

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:8be5ecbe9d1953713e447e45d1d70574345b9259f6bb30034e2db5682b887fc2
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.
fp8"e4m3",
num-warps = 4num_warps = 4
online-softmax:return row_max_new: Updated maximum values per row of shape [BLOCK_M].
shared-memory__shared__ float shared_max[kThreadsPerBlock];
split-kIS_SPLIT_KV: tl.constexpr,
stages = 2num_stages = 2
tile-k = 1TILE_K=1,
tile-m = 16TILE_M = 16
tile-n = 16TILE_N = 16

Kernel source

submission_sparse.py2700 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


@triton.jit
def pid_grid(pid: int, num_pid_m: int, num_pid_n: int, GROUP_SIZE_M: tl.constexpr = 1):
    """
    Maps 1D pid to 2D grid coords (pid_m, pid_n).

    Args:
        - pid: 1D pid
        - num_pid_m: grid m size
        - num_pid_n: grid n size
        - GROUP_SIZE_M: tl.constexpr: default is 1
    """
    if GROUP_SIZE_M == 1:
        pid_m = pid // num_pid_n
        pid_n = pid % num_pid_n
    else:
        num_pid_in_group = GROUP_SIZE_M * num_pid_n
        group_id = pid // num_pid_in_group
        first_pid_m = group_id * GROUP_SIZE_M
        group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M)
        tl.assume(group_size_m >= 0)
        pid_m = first_pid_m + (pid % group_size_m)
        pid_n = (pid % num_pid_in_group) // group_size_m

    return pid_m, pid_n


def num_splits_heuristic(
    seqlen_q: int,
    seqlen_k: int,
    num_SMs: int,
    TILE_M: int,
    TILE_N: int,
) -> int:
    """
    Determine the number of KV splits for FlashDecoding.

    Splits only when there are enough KV blocks to benefit from parallelism,
    and targets full SM occupancy by over-subscribing the M-block count.

    :param seqlen_q: Sequence length of queries.
    :param seqlen_k: Sequence length of keys.
    :param num_SMs: Number of SM on the device.
    :param TILE_M: Tile size for M dimension.
    :param TILE_N: Tile size for N dimension.

    :return: Number of splits.
    """
    total_mblocks = triton.cdiv(seqlen_q, TILE_M)
    num_n_blocks = triton.cdiv(seqlen_k, TILE_N)
    max_splits = triton.next_power_of_2(num_SMs)
    if num_n_blocks <= 4:
        # 1 means no splitting
        return 1
    return min(num_SMs // max(total_mblocks, 1), max_splits, num_n_blocks)

# tile_n, batch_size, seqlen_k, is_local, window_size_left, window_size_right
_BENCHMARK_SPLIT_PRESETS = {
    (16, 4, 1024, True, 1024, 0): (8, [0, 9, 18, 27, 35, 43, 50, 57, 64]),
    (16, 32, 1024, True, 1024, 0): (8, [0, 9, 17, 25, 33, 41, 49, 57, 64]),
    (16, 64, 1024, True, 1024, 0): (8, [0, 9, 17, 25, 33, 41, 49, 57, 64]),
    (16, 256, 1024, True, 1024, 0): (4, [0, 17, 33, 49, 64]),
    (16, 4, 8192, True, 4096, 0): (16, [255, 273, 291, 309, 327, 344, 361, 378, 394, 410, 426, 441, 456, 470, 484, 498, 512]),
    (16, 32, 8192, True, 4096, 0): (16, [255, 273, 290, 307, 324, 341, 358, 374, 390, 406, 422, 437, 452, 467, 482, 497, 512]),
    (16, 64, 8192, True, 4096, 0): (16, [255, 273, 290, 307, 324, 341, 358, 374, 390, 406, 422, 437, 452, 467, 482, 497, 512]),
    (16, 256, 8192, True, 4096, 0): (16, [255, 272, 289, 306, 322, 338, 354, 370, 386, 402, 418, 434, 450, 466, 481, 497, 512]),
    (32, 4, 1024, True, 1024, 0): (4, [0, 9, 17, 25, 32]),
    (32, 32, 1024, True, 1024, 0): (4, [0, 9, 17, 25, 32]),
    (32, 64, 1024, True, 1024, 0): (4, [0, 9, 17, 25, 32]),
    (32, 256, 1024, True, 1024, 0): (4, [0, 8, 16, 24, 32]),
    (32, 4, 8192, True, 4096, 0): (8, [127, 146, 164, 181, 197, 213, 228, 242, 256]),
    (32, 32, 8192, True, 4096, 0): (8, [127, 145, 162, 179, 195, 211, 226, 241, 256]),
    (32, 64, 8192, True, 4096, 0): (8, [127, 145, 162, 179, 195, 211, 226, 241, 256]),
    (32, 256, 8192, True, 4096, 0): (8, [127, 144, 161, 177, 193, 209, 225, 241, 256]),
}



_SPLIT_TENSOR_CACHE: dict[tuple[Any, ...], tuple[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_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 get_seqlen_info(
    batch_idx,
    seqlen_static,
    cu_seqlens,
    seqused,
    HAS_CU_SEQLENS: tl.constexpr,
):
    """
    Get offset and seqlen for a given batch index.

    :param batch_idx: Index of the batch.
    :param seqlen_static: Static sequence length if cu_seqlens is not provided.
    :param cu_seqlens: Cumulative sequence lengths tensor.
    :param seqused: Actual sequence lengths tensor.
    :param HAS_CU_SEQLENS: Boolean flag indicating if cu_seqlens is provided.

    :return offset: Offset for the given batch index.
    :return seqlen: Sequence length for the given batch index.
    """
    if HAS_CU_SEQLENS:
        offset = tl.load(cu_seqlens + batch_idx)
        seqlen = tl.load(cu_seqlens + batch_idx + 1) - offset
    else:
        offset = 0
        seqlen = seqlen_static
    return offset, seqlen


@triton.jit
def get_seqlen_info_qk(
    batch_idx,
    seqlen_q_static,
    seqlen_k_static,
    cu_seqlens_q,
    cu_seqlens_k,
    HAS_CU_SEQLENS_Q: tl.constexpr,
    HAS_CU_SEQLENS_K: tl.constexpr,
):
    """
    Get offset, padded_offset, and seqlen for both Q and K.

    :param batch_idx: Index of the batch.
    :param seqlen_q_static: Static sequence length for Q if cu_seqlens_q is not provided.
    :param seqlen_k_static: Static sequence length for K if cu_seqlens_k is not provided.
    :param cu_seqlens_q: Cumulative sequence lengths tensor for Q.
    :param cu_seqlens_k: Cumulative sequence lengths tensor for K.
    :param seqused_q: Actual sequence lengths tensor for Q.
    :param seqused_k: Actual sequence lengths tensor for K.
    :param HAS_CU_SEQLENS_Q: Boolean flag indicating if cu_seqlens_q is provided.
    :param HAS_CU_SEQLENS_K: Boolean flag indicating if cu_seqlens_k is provided.

    :return offset_q: Offset for Q for the given batch index.
    :return offset_k: Offset for K for the given batch index.
    :return padded_offset_q: Padded offset for Q aligned to TILE_M.
    :return padded_offset_k: Padded offset for K aligned to TILE_N.
    :return seqlen_q: Sequence length for Q for the given batch index.
    :return seqlen_k: Sequence length for K for the given batch index.
    """

    # Q offset and seqlen
    if HAS_CU_SEQLENS_Q:
        offset_q = tl.load(cu_seqlens_q + batch_idx)
        seqlen_q = tl.load(cu_seqlens_q + batch_idx + 1) - offset_q
    else:
        offset_q = 0
        seqlen_q = seqlen_q_static

    # K offset and seqlen
    if HAS_CU_SEQLENS_K:
        offset_k = tl.load(cu_seqlens_k + batch_idx)
        seqlen_k = tl.load(cu_seqlens_k + batch_idx + 1) - offset_k
    else:
        offset_k = 0
        seqlen_k = seqlen_k_static

    return offset_q, offset_k, seqlen_q, seqlen_k


@triton.jit
def get_n_block_min_max(
    seqlen_q,
    seqlen_k,
    m_block,
    split_idx,
    num_splits,
    TILE_N: tl.constexpr,
    TILE_M: tl.constexpr,
    IS_CAUSAL: tl.constexpr,
    IS_LOCAL: tl.constexpr,
    IS_SPLIT_KV: tl.constexpr,
    WINDOW_SIZE_LEFT: tl.constexpr,
    WINDOW_SIZE_RIGHT: tl.constexpr,
    QHEAD_PER_KVHEAD_PACKGQA: tl.constexpr,
):
    n_block_max = tl.cdiv(seqlen_k, TILE_N)
    if IS_CAUSAL or (IS_LOCAL and WINDOW_SIZE_RIGHT is not None):
        m_idx_max = (m_block + 1) * TILE_M
        if QHEAD_PER_KVHEAD_PACKGQA > 1:
            m_idx_max = tl.cdiv(m_idx_max, QHEAD_PER_KVHEAD_PACKGQA)
        n_idx = m_idx_max + seqlen_k - seqlen_q
        n_idx_right = n_idx if IS_CAUSAL else n_idx + WINDOW_SIZE_RIGHT
        n_block_max = tl.minimum(n_block_max, tl.cdiv(n_idx_right, TILE_N))
    n_block_min = 0
    if IS_LOCAL and WINDOW_SIZE_LEFT is not None:
        m_idx_min = m_block * TILE_M
        if QHEAD_PER_KVHEAD_PACKGQA > 1:
            m_idx_min = m_idx_min // QHEAD_PER_KVHEAD_PACKGQA
        n_idx = m_idx_min + seqlen_k - seqlen_q
        n_idx_left = n_idx - WINDOW_SIZE_LEFT
        n_block_min = tl.maximum(n_idx_left // TILE_N, 0)
    if IS_SPLIT_KV:
        num_n_blocks_per_split = (
            0
            if n_block_max <= n_block_min
            else (n_block_max - n_block_min + num_splits - 1) // num_splits
        )
        n_block_min = n_block_min + split_idx * num_n_blocks_per_split
        n_block_max = tl.minimum(n_block_min + num_n_blocks_per_split, n_block_max)
    return n_block_min, n_block_max


@triton.jit
def get_m_block_min_max(
    seqlen_q,
    seqlen_k,
    n_block,
    TILE_N: tl.constexpr,
    TILE_M: tl.constexpr,
    IS_CAUSAL: tl.constexpr,
    IS_LOCAL: tl.constexpr,
    WINDOW_SIZE_LEFT: tl.constexpr,
    WINDOW_SIZE_RIGHT: tl.constexpr,
):
    m_block_max = tl.cdiv(seqlen_q, TILE_M)
    m_block_min = 0
    if IS_CAUSAL or (IS_LOCAL and WINDOW_SIZE_RIGHT is not None):
        n_idx_min = n_block * TILE_N
        m_idx = n_idx_min + seqlen_q - seqlen_k
        m_idx_right = m_idx if IS_CAUSAL else m_idx - WINDOW_SIZE_RIGHT
        m_block_min = tl.maximum(m_block_min, m_idx_right // TILE_M)
    if IS_LOCAL and WINDOW_SIZE_LEFT is not None:
        n_idx_max = (n_block + 1) * TILE_N
        m_idx = n_idx_max + seqlen_q - seqlen_k
        m_idx_left = m_idx + WINDOW_SIZE_LEFT
        m_block_max = tl.minimum(m_block_max, tl.cdiv(m_idx_left, TILE_M))
    return m_block_min, m_block_max


@triton.jit
def get_n_block_min_causal_local_mask(
    seqlen_q,
    seqlen_k,
    m_block,
    n_block_min,
    TILE_N: tl.constexpr,
    TILE_M: tl.constexpr,
    IS_LOCAL: tl.constexpr,
    WINDOW_SIZE_RIGHT: tl.constexpr,
    QHEAD_PER_KVHEAD_PACKGQA: tl.constexpr,
):
    m_idx_min = m_block * TILE_M
    if QHEAD_PER_KVHEAD_PACKGQA > 1:
        m_idx_min = m_idx_min // QHEAD_PER_KVHEAD_PACKGQA
    n_idx = m_idx_min + seqlen_k - seqlen_q
    n_idx_right = (
        n_idx
        if (not IS_LOCAL or WINDOW_SIZE_RIGHT is None)
        else n_idx + WINDOW_SIZE_RIGHT
    )
    return tl.maximum(n_block_min, n_idx_right // TILE_N)


@triton.jit
def get_n_block_min_before_local_mask(
    seqlen_q,
    seqlen_k,
    m_block,
    n_block_min,
    TILE_N: tl.constexpr,
    TILE_M: tl.constexpr,
    IS_LOCAL: tl.constexpr,
    WINDOW_SIZE_LEFT: tl.constexpr,
    QHEAD_PER_KVHEAD_PACKGQA: tl.constexpr,
):
    if not IS_LOCAL or WINDOW_SIZE_LEFT is None:
        return n_block_min
    else:
        m_idx_max = (m_block + 1) * TILE_M
        if QHEAD_PER_KVHEAD_PACKGQA > 1:
            m_idx_max = tl.cdiv(m_idx_max, QHEAD_PER_KVHEAD_PACKGQA)
        n_idx = m_idx_max + seqlen_k - seqlen_q
        n_idx_left = n_idx - WINDOW_SIZE_LEFT
        return tl.maximum(n_block_min, tl.cdiv(n_idx_left, TILE_N))


@triton.jit
def get_m_block_min_causal_local_mask(
    seqlen_q,
    seqlen_k,
    n_block,
    m_block_min,
    TILE_N: tl.constexpr,
    TILE_M: tl.constexpr,
    IS_CAUSAL: tl.constexpr,
    IS_LOCAL: tl.constexpr,
    WINDOW_SIZE_RIGHT: tl.constexpr,
):
    if not IS_CAUSAL and (not IS_LOCAL or WINDOW_SIZE_RIGHT is None):
        return m_block_min
    else:
        n_idx_max = (n_block + 1) * TILE_N
        m_idx = n_idx_max + seqlen_q - seqlen_k
        m_idx_right = m_idx if IS_CAUSAL else m_idx - WINDOW_SIZE_RIGHT
        return tl.maximum(m_block_min, tl.cdiv(m_idx_right, TILE_M))


@triton.jit
def get_m_block_max_before_local_mask(
    seqlen_q,
    seqlen_k,
    n_block,
    m_block_max,
    TILE_N: tl.constexpr,
    TILE_M: tl.constexpr,
    IS_LOCAL: tl.constexpr,
    WINDOW_SIZE_LEFT: tl.constexpr,
):
    if not IS_LOCAL or WINDOW_SIZE_LEFT is None:
        return m_block_max
    else:
        n_idx_min = n_block * TILE_N
        m_idx = n_idx_min + seqlen_q - seqlen_k
        m_idx_left = m_idx + WINDOW_SIZE_LEFT
        return tl.minimum(m_block_max, m_idx_left // TILE_M)


@triton.jit
def offset_batch_Q(
    base_ptr,
    batch_idx,
    offset,
    stride_batch,
    stride_seq,
    HAS_CU_SEQLENS: tl.constexpr,
):
    if HAS_CU_SEQLENS:
        return base_ptr + offset * stride_seq
    else:
        return base_ptr + batch_idx * stride_batch


@triton.jit
def offset_batch_K(
    base_ptr,
    batch_idx,
    offset,
    stride_batch,
    stride_seq,
    HAS_CU_SEQLENS: tl.constexpr,
):
    if HAS_CU_SEQLENS:
        return base_ptr + offset * stride_seq
    else:
        return base_ptr + batch_idx * stride_batch


@triton.jit
def make_ptrs(
    base_ptrs,
    mn_block,
    stride_seq,
    TILE_MN: tl.constexpr,
    TILE_K: tl.constexpr,
    SWAP_AB: tl.constexpr,
):
    offs_mn = mn_block * TILE_MN + tl.arange(0, TILE_MN)
    if TILE_K > 1:
        offs_k = tl.arange(0, TILE_K)
        if SWAP_AB:
            ptrs = base_ptrs + offs_mn[None, :] * stride_seq + offs_k[:, None]
        else:
            ptrs = base_ptrs + offs_mn[:, None] * stride_seq + offs_k[None, :]
    else:
        ptrs = base_ptrs + offs_mn
    return ptrs


@triton.jit
def make_pack_gqa_ptrs(
    base_ptrs,
    m_block,
    head_idx,
    stride_head,
    stride_seq,
    TILE_M: tl.constexpr,
    TILE_K: tl.constexpr,
    QHEADS_PER_KVHEAD_PACKGQA: tl.constexpr,
):
    offs_m = m_block * TILE_M + tl.arange(0, TILE_M)
    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
    if TILE_K > 1:
        offs_k = tl.arange(0, TILE_K)
        ptrs = (
            base_ptrs
            + m_idx[:, None] * stride_seq
            + q_head[:, None] * stride_head
            + offs_k[None, :]
        )
    else:
        ptrs = base_ptrs + m_idx * stride_seq + q_head * stride_head
    return ptrs


@triton.jit
def check_inf(x):
    return tl.where(x == float("-inf"), 0.0, x)


@triton.jit
def online_softmax(
    acc_s,
    row_max,
    row_sum,
    scale_log2,
    CHECK_INF: tl.constexpr,
):
    """
    Apply online softmax to acc_s, and update block_max, row_max and row_sum.

    :param acc_s: Attention scores tensor of shape [BLOCK_M, BLOCK_N].
    :param block_max: Running block-wise maximum scalar, init to -inf.
    :param row_max: Current maximum values per row of shape [BLOCK_M], init to -inf.
    :param row_sum: Current sum values per row of shape [BLOCK_M], init to 0.
    :param scale_log2: Log2 of the scaling factor to be applied to acc_s.
    :param CHECK_INF: Boolean flag indicating if -inf row_max should be clamped to 0.
    :param RESCALE_THRESHOLD: Threshold for rescaling to avoid underflow. If <= 0, rescaling is disabled.

    :return p: Softmax probabilities tensor of shape [BLOCK_M, BLOCK_N].
    :return block_max_new: Updated block-wise maximum scalar.
    :return row_max_new: Updated maximum values per row of shape [BLOCK_M].
    :return row_sum_new: Updated sum values per row of shape [BLOCK_M].
    :return row_scale: Scaling factors per row of shape [BLOCK_M].
    :return skip_softmax: Boolean indicating whether this block was skipped.
    """

    # 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)

    # Avoid exp(-inf - (-inf)) = nan by clamping -inf to 0
    if CHECK_INF:
        row_max_new = check_inf(row_max_new)

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

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

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

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

    return p, row_max_new, row_sum_new, row_scale


@triton.jit
def finalize(
    row_max,
    row_sum,
    scale_log2,
):
    """
    Finalize online softmax by computing output scale and logsumexp.

    :param row_max: Final maximum values per row of shape [BLOCK_M].
    :param row_sum: Final sum values per row of shape [BLOCK_M].

    :return row_scale: Final scaling factors per row of shape [BLOCK_M].
    :return lse: Logsumexp values per row of shape [BLOCK_M].
    """
    # # if row_sum is zero or nan, set it to 1 to avoid division by zero
    # acc_o_is_zero_or_nan = (row_sum == 0.0) | (row_sum != row_sum)
    # row_scale = tl.where(acc_o_is_zero_or_nan, 1.0, 1.0 / row_sum)
    # # ln2 = math.log(2.0)
    # ln2 = 0.6931471805599453
    # lse = tl.where(
    #     acc_o_is_zero_or_nan,
    #     float("-inf"),
    #     (row_max * scale_log2 + tl.log2(row_sum)) * ln2,
    # )

    row_scale = 1.0 / row_sum
    ln2 = 0.6931471805599453
    lse = (row_max * scale_log2 + tl.log2(row_sum)) * ln2
    return row_scale, lse


@triton.jit
def rescale_o(
    acc_o,
    row_scale,
):
    """
    Rescale output accumulator by row_scale.

    :param acc_o: Output accumulator tensor of shape [BLOCK_M, BLOCK_N].
    :param row_scale: Scaling factors per row of shape [BLOCK_M].

    :return: Rescaled output accumulator tensor of shape [BLOCK_M, BLOCK_N].
    """
    acc_o = acc_o * row_scale[:, None]
    return acc_o


@triton.jit
def apply_mask(
    acc_s,
    m_block,
    n_block,
    seqlen_q,
    seqlen_k,
    MASK_SEQLEN: tl.constexpr,
    MASK_CAUSAL: tl.constexpr,
    MASK_LOCAL: tl.constexpr,
    TILE_M: tl.constexpr,
    TILE_N: tl.constexpr,
    WINDOW_SIZE_LEFT: tl.constexpr,
    WINDOW_SIZE_RIGHT: tl.constexpr,
    QHEADS_PER_KVHEAD_PACKGQA: tl.constexpr,
    SWAP_AB: tl.constexpr,
):
    """
    Apply seqlen, causal, and local masks to the attention scores.

    :param acc_s: Attention scores tensor of shape [BLOCK_M, BLOCK_N].
    :param m_block: Current block index along the M dimension.
    :param n_block: Current block index along the N dimension.
    :param seqlen_q: The sequence length of the query.
    :param seqlen_k: The sequence length of the key.
    :param MASK_SEQLEN: Boolean flag indicating if seqlen masking should be applied.
    :param MASK_CAUSAL: Boolean flag indicating if causal masking should be applied.
    :param MASK_LOCAL: Boolean flag indicating if local masking should be applied.
    :param TILE_M: Tile size along the M dimension.
    :param TILE_N: Tile size along the N dimension.
    :param WINDOW_SIZE_LEFT: Left window size for local masking.
    :param WINDOW_SIZE_RIGHT: Right window size for local masking.
    :param QHEADS_PER_KVHEAD_PACKGQA: Ratio of query heads to key/value heads for packed GQA.
    :param SWAP_AB: Boolean flag indicating if query and key dimensions are swapped.

    :return acc_s: Masked attention scores tensor of shape [BLOCK_M, BLOCK_N].
    """
    tl.static_assert(
        not (MASK_CAUSAL and MASK_LOCAL),
        "MASK_CAUSAL and MASK_LOCAL cannot be both True",
    )
    offs_m = m_block * TILE_M + tl.arange(0, TILE_M)
    offs_n = n_block * TILE_N + tl.arange(0, TILE_N)

    if SWAP_AB:
        tl.static_assert(
            QHEADS_PER_KVHEAD_PACKGQA == 1, "SWAP_AB with PACKGQA > 1 not supported"
        )
        q_idx = offs_m[None, :]
        k_idx = offs_n[:, None]
    else:
        q_idx = offs_m[:, None]
        k_idx = offs_n[None, :]
        if QHEADS_PER_KVHEAD_PACKGQA > 1:
            q_idx = q_idx // QHEADS_PER_KVHEAD_PACKGQA

    if MASK_SEQLEN:
        acc_s = tl.where(
            (k_idx < seqlen_k) & (q_idx < seqlen_q),
            acc_s,
            float("-inf"),
        )

    if MASK_CAUSAL or MASK_LOCAL:
        causal_offset = seqlen_k - seqlen_q

        if MASK_CAUSAL:
            acc_s = tl.where(
                q_idx + causal_offset >= k_idx,
                acc_s,
                float("-inf"),
            )
        else:
            if WINDOW_SIZE_RIGHT is not None:
                acc_s = tl.where(
                    q_idx + causal_offset + WINDOW_SIZE_RIGHT >= k_idx,
                    acc_s,
                    float("-inf"),
                )
            if WINDOW_SIZE_LEFT is not None:
                acc_s = tl.where(
                    q_idx + causal_offset - WINDOW_SIZE_LEFT <= k_idx,
                    acc_s,
                    float("-inf"),
                )

    return acc_s


@triton.jit
def _fwd_inner_sparse_fp8_kernel(
    q_tile,
    q_tile_tail,
    q_scale,
    kv_scale,
    k_tile,
    k_tile_tail,
    k_ptrs,
    k_tail_ptrs,
    acc_o,
    row_max,
    row_sum,
    softmax_scale_log2: tl.constexpr,
    n_block,
    n_block_min,
    TILE_M: tl.constexpr,
    TILE_N: tl.constexpr,
    CHECK_INF: tl.constexpr,
):
    k_tile_next = k_tile
    # Compute attention scores
    acc_s = tl.zeros((TILE_M, TILE_N), dtype=tl.float32)
    acc_s = tl.dot_scaled(
        q_tile,
        None,
        "e4m3",
        k_tile,
        None,
        "e4m3",
        acc=acc_s,
        fast_math=True,
    )
    acc_s = tl.dot_scaled(
        q_tile_tail,
        None,
        "e4m3",
        k_tile_tail,
        None,
        "e4m3",
        acc=acc_s,
        fast_math=True,
    )
    acc_s *= q_scale * kv_scale

    # Advance key pointer
    k_ptrs = tl.advance(k_ptrs, (0, -TILE_N))
    k_tail_ptrs = tl.advance(k_tail_ptrs, (0, -TILE_N))
    if n_block > n_block_min:
        # Load next key tile
        k_tile_next = tl.load(k_ptrs, boundary_check=(0, 1))
        k_tile_tail = tl.load(k_tail_ptrs, boundary_check=(0, 1))

    # Apply online softmax
    p, row_max, row_sum, row_scale = (
        online_softmax(
            acc_s=acc_s,
            row_max=row_max,
            row_sum=row_sum,
            scale_log2=softmax_scale_log2,
            CHECK_INF=CHECK_INF,
        )
    )

    v_tile = tl.trans(k_tile)

    # Rescale output accumulator
    acc_o = rescale_o(acc_o, row_scale)

    # Update output accumulator
    acc_o += tl.dot_scaled(
        p.to(v_tile.dtype),
        None,
        "e4m3",
        v_tile,
        None,
        "e4m3",
        fast_math=True,
    ) * kv_scale

    k_tile = k_tile_next

    return (
        k_tile,
        k_tile_tail,
        k_ptrs,
        k_tail_ptrs,
        acc_o,
        row_max,
        row_sum,
    )


@triton.jit
def _fwd_base_sparse_kernel(
    Q,
    Q_S,
    KV,
    KV_S,
    Out,
    Lse,
    SplitCounts,
    SplitBoundaries,
    softmax_scale_log2: tl.constexpr,
    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,
    cu_seqlens_q,
    cu_seqlens_k,
    num_splits,
    seqlen_q: tl.constexpr,
    seqlen_k: tl.constexpr,
    head_dim_qk: tl.constexpr,
    head_dim_v: tl.constexpr,
    QHEADS_PER_KVHEAD_PACKGQA: tl.constexpr,
    TILE_M: tl.constexpr,
    TILE_N: tl.constexpr,
    TILE_K: tl.constexpr,
    IS_CAUSAL: tl.constexpr,
    IS_LOCAL: tl.constexpr,
    IS_SPLIT_KV: tl.constexpr,
    WINDOW_SIZE_LEFT: tl.constexpr,
    WINDOW_SIZE_RIGHT: 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,
):
    m_block = tl.program_id(0)
    head_idx = tl.program_id(1)
    batch_split_idx = tl.program_id(2)

    QKV_HEAD_DIM: tl.constexpr = 512
    QK_ROPE_HEAD_DIM: tl.constexpr = 64

    if IS_SPLIT_KV:
        batch_idx = batch_split_idx // num_splits
        split_idx = batch_split_idx - batch_idx * num_splits
    else:
        batch_idx = batch_split_idx
        split_idx = 0

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

    # Get seqlen info for this batch
    (
        offset_q,
        offset_k,
        actual_seqlen_q,
        actual_seqlen_k,
    ) = get_seqlen_info_qk(
        batch_idx=batch_idx,
        seqlen_q_static=seqlen_q,
        seqlen_k_static=seqlen_k,
        cu_seqlens_q=cu_seqlens_q,
        cu_seqlens_k=cu_seqlens_k,
        HAS_CU_SEQLENS_Q=HAS_CU_SEQLENS_Q,
        HAS_CU_SEQLENS_K=HAS_CU_SEQLENS_K,
    )

    # Initialize base pointers
    q_base = offset_batch_Q(
        Q,
        batch_idx,
        offset_q,
        stride_qb,
        stride_qm,
        HAS_CU_SEQLENS_Q,
    )
    k_base = offset_batch_K(
        KV + head_idx * stride_kvh,
        batch_idx,
        offset_k,
        stride_kvb,
        stride_kvn,
        HAS_CU_SEQLENS_K,
    )
    out_base = offset_batch_Q(
        Out,
        batch_idx,
        offset_q,
        stride_ob,
        stride_om,
        HAS_CU_SEQLENS_Q,
    )
    lse_base = offset_batch_Q(
        Lse,
        batch_idx,
        offset_q,
        stride_lb,
        1,
        HAS_CU_SEQLENS_Q,
    )

    # For split KV, offset output and LSE base pointers by split_idx
    if IS_SPLIT_KV:
        out_base += split_idx * stride_os
        lse_base += split_idx * stride_ls

    # Compute n_block range for this m_block
    n_block_min, n_block_max = get_n_block_min_max(
        seqlen_q=actual_seqlen_q,
        seqlen_k=actual_seqlen_k,
        m_block=m_block,
        split_idx=0,
        num_splits=1,
        TILE_N=TILE_N,
        TILE_M=TILE_M,
        IS_CAUSAL=IS_CAUSAL,
        IS_LOCAL=IS_LOCAL,
        IS_SPLIT_KV=IS_SPLIT_KV,
        WINDOW_SIZE_LEFT=WINDOW_SIZE_LEFT,
        WINDOW_SIZE_RIGHT=WINDOW_SIZE_RIGHT,
        QHEAD_PER_KVHEAD_PACKGQA=QHEADS_PER_KVHEAD_PACKGQA,
    )
    if IS_SPLIT_KV:
        split_boundaries = SplitBoundaries
        split_n_block_min = tl.load(split_boundaries + split_idx)
        split_n_block_max = tl.load(split_boundaries + split_idx + 1)
        n_block_min = tl.maximum(n_block_min, split_n_block_min)
        n_block_max = tl.minimum(n_block_max, split_n_block_max)

    n_block_min_no_mask = get_n_block_min_before_local_mask(
        seqlen_q=actual_seqlen_q,
        seqlen_k=actual_seqlen_k,
        m_block=m_block,
        n_block_min=n_block_min,
        TILE_N=TILE_N,
        TILE_M=TILE_M,
        IS_LOCAL=IS_LOCAL,
        WINDOW_SIZE_LEFT=WINDOW_SIZE_LEFT,
        QHEAD_PER_KVHEAD_PACKGQA=QHEADS_PER_KVHEAD_PACKGQA,
    )
    n_block_max_no_mask = get_n_block_min_causal_local_mask(
        seqlen_q=actual_seqlen_q,
        seqlen_k=actual_seqlen_k,
        m_block=m_block,
        n_block_min=n_block_min,
        TILE_N=TILE_N,
        TILE_M=TILE_M,
        IS_LOCAL=IS_LOCAL,
        WINDOW_SIZE_RIGHT=WINDOW_SIZE_RIGHT,
        QHEAD_PER_KVHEAD_PACKGQA=QHEADS_PER_KVHEAD_PACKGQA,
    )

    # Clamp to split's range so the no-mask loop stays within bounds
    if IS_SPLIT_KV:
        n_block_min_no_mask = tl.maximum(n_block_min_no_mask, n_block_min)
        n_block_max_no_mask = tl.minimum(n_block_max_no_mask, n_block_max)

    # Create pointers
    lse_ptrs = make_pack_gqa_ptrs(
        lse_base,
        m_block,
        head_idx,
        stride_lh,
        1,
        TILE_M=TILE_M,
        TILE_K=1,
        QHEADS_PER_KVHEAD_PACKGQA=QHEADS_PER_KVHEAD_PACKGQA,
    )
    out_ptrs = make_pack_gqa_ptrs(
        out_base,
        m_block,
        head_idx,
        stride_oh,
        stride_om,
        TILE_M=TILE_M,
        TILE_K=TILE_K,
        QHEADS_PER_KVHEAD_PACKGQA=QHEADS_PER_KVHEAD_PACKGQA,
    )

    q_ptrs = make_pack_gqa_ptrs(
        q_base,
        m_block,
        head_idx,
        stride_qh,
        stride_qm,
        TILE_M=TILE_M,
        TILE_K=TILE_K,
        QHEADS_PER_KVHEAD_PACKGQA=QHEADS_PER_KVHEAD_PACKGQA,
    )
    q_tail_ptrs = make_pack_gqa_ptrs(
        q_base + QKV_HEAD_DIM,
        m_block,
        head_idx,
        stride_qh,
        stride_qm,
        TILE_M=TILE_M,
        TILE_K=QK_ROPE_HEAD_DIM,
        QHEADS_PER_KVHEAD_PACKGQA=QHEADS_PER_KVHEAD_PACKGQA,
    )
    k_ptrs = tl.make_block_ptr(
        base=k_base,
        shape=(head_dim_v, actual_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_v,
        shape=(head_dim_qk - head_dim_v, actual_seqlen_k),
        strides=(1, stride_kvn),
        offsets=(0, (n_block_max - 1) * TILE_N),
        block_shape=(QK_ROPE_HEAD_DIM, TILE_N),
        order=(1, 0),
    )

    q_scale = tl.load(Q_S)
    kv_scale = tl.load(KV_S)

    # Load query tile
    q_tile = tl.load(
        q_ptrs,
        cache_modifier=".ca",
    )
    q_tile_tail = tl.load(
        q_tail_ptrs,
        cache_modifier=".ca",
    )

    # 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 key tile
    k_tile = tl.load(k_ptrs)
    k_tile_tail = tl.load(k_tail_ptrs)

    n_block_max_no_mask = n_block_max
    n_block_min_no_mask = tl.minimum(n_block_min_no_mask, n_block_max_no_mask)

    # Process n_blocks without masking
    if n_block_max_no_mask > n_block_min_no_mask:
        for n_block in tl.range(n_block_max_no_mask - 1, n_block_min_no_mask - 1, -1):
            (
                k_tile,
                k_tile_tail,
                k_ptrs,
                k_tail_ptrs,
                acc_o,
                row_max,
                row_sum,
            ) = _fwd_inner_sparse_fp8_kernel(
                q_tile=q_tile,
                q_tile_tail=q_tile_tail,
                q_scale=q_scale,
                kv_scale=kv_scale,
                k_tile=k_tile,
                k_tile_tail=k_tile_tail,
                k_ptrs=k_ptrs,
                k_tail_ptrs=k_tail_ptrs,
                acc_o=acc_o,
                row_max=row_max,
                row_sum=row_sum,
                softmax_scale_log2=softmax_scale_log2,
                n_block=n_block,
                n_block_min=n_block_min_no_mask,
                TILE_M=TILE_M,
                TILE_N=TILE_N,
                CHECK_INF=False,
            )

    # Process n_blocks with masking
    if IS_LOCAL and n_block_min_no_mask > n_block_min:
        k_ptrs = tl.make_block_ptr(
            base=k_base,
            shape=(head_dim_v, actual_seqlen_k),
            strides=(1, stride_kvn),
            offsets=(0, (n_block_min_no_mask - 1) * TILE_N),
            block_shape=(TILE_K, TILE_N),
            order=(0, 1),
        )
        k_tail_ptrs = tl.make_block_ptr(
            base=k_base + head_dim_v,
            shape=(head_dim_qk - head_dim_v, actual_seqlen_k),
            strides=(1, stride_kvn),
            offsets=(0, (n_block_min_no_mask - 1) * TILE_N),
            block_shape=(QK_ROPE_HEAD_DIM, TILE_N),
            order=(0, 1),
        )

        k_tile = tl.load(k_ptrs)
        k_tile_tail = tl.load(k_tail_ptrs)
        for n_block in tl.range(n_block_min_no_mask - 1, n_block_min - 1, -1):
            (
                k_tile,
                k_tile_tail,
                k_ptrs,
                k_tail_ptrs,
                acc_o,
                row_max,
                row_sum,
            ) = _fwd_inner_sparse_fp8_kernel(
                q_tile=q_tile,
                q_tile_tail=q_tile_tail,
                q_scale=q_scale,
                kv_scale=kv_scale,
                k_tile=k_tile,
                k_tile_tail=k_tile_tail,
                k_ptrs=k_ptrs,
                k_tail_ptrs=k_tail_ptrs,
                acc_o=acc_o,
                row_max=row_max,
                row_sum=row_sum,
                softmax_scale_log2=softmax_scale_log2,
                n_block=n_block,
                n_block_min=n_block_min,
                TILE_M=TILE_M,
                TILE_N=TILE_N,
                CHECK_INF=False,
            )

    # Finalize softmax
    row_scale, lse_tile = finalize(
        row_max=row_max,
        row_sum=row_sum,
        scale_log2=softmax_scale_log2,
    )
    acc_o = rescale_o(acc_o, row_scale)

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

    # Store output
    # When IS_SPLIT_KV, store float32 partial results.
    # Otherwise, convert back to input dtype.
    if not IS_SPLIT_KV:
        acc_o = acc_o.to(Out.dtype.element_ty)

    tl.store(
        out_ptrs,
        acc_o,
        cache_modifier=".wb",
    )


@triton.jit
def _fwd_combine_kernel(
    Out_partial,
    Lse_partial,
    Out,
    stride_ops,
    stride_opb,
    stride_oph,
    stride_opm,
    stride_lps,
    stride_lpb,
    stride_lph,
    stride_ob,
    stride_oh,
    stride_om,
    cu_seqlens_q,
    seqused_q,
    num_splits,
    batch_size,
    seqlen_q,
    num_heads_q,
    head_dim,
    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
    offset_q, actual_seqlen_q = get_seqlen_info(
        batch_idx=batch_idx,
        seqlen_static=seqlen_q,
        cu_seqlens=cu_seqlens_q,
        seqused=seqused_q,
        HAS_CU_SEQLENS=HAS_CU_SEQLENS_Q,
    )
    if actual_seqlen_q <= 0:
        return

    # Initialize base pointers
    out_part_base = offset_batch_Q(
        Out_partial + head_idx * stride_oph,
        batch_idx,
        offset_q,
        stride_opb,
        stride_opm,
        HAS_CU_SEQLENS_Q,
    )
    lse_part_base = offset_batch_Q(
        Lse_partial + head_idx * stride_lph,
        batch_idx,
        offset_q,
        stride_lpb,
        1,
        HAS_CU_SEQLENS_Q,
    )
    out_base = offset_batch_Q(
        Out + head_idx * stride_oh,
        batch_idx,
        offset_q,
        stride_ob,
        stride_om,
        HAS_CU_SEQLENS_Q,
    )

    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)

    # Compute max across splits
    for s in tl.range(0, num_splits):
        lse_s = tl.load(
            lse_part_row_base + s * stride_lps,
            cache_modifier=".cg",
        )
        o_s = tl.load(
            out_part_row_base + s * stride_ops,
            cache_modifier=".cg",
        )

        n_e_max = tl.maximum(lse_s, e_max)
        old_scale = tl.exp(e_max - n_e_max)
        exp_logic = tl.exp(lse_s - n_e_max)

        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

    # 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,
    cu_seqlens_q: torch.Tensor = None,
    seqused_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,
        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,
        seqused_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_varlen_forward(
    query: torch.Tensor,
    q_scale: torch.Tensor,
    kv: torch.Tensor,
    kv_scale: torch.Tensor,
    cu_seqlens_q: torch.Tensor,
    cu_seqlens_k: torch.Tensor,
    max_seqlen_q: int,
    max_seqlen_k: int,
    is_causal: bool = False,
    softmax_scale: float = None,
    window_size: Tuple[int, int] = (None, None),
    out: torch.Tensor | None = None,
) -> Tuple[torch.Tensor, torch.Tensor, float, float]:
    total_seqlen_q, num_heads_q, head_dim_qk = query.shape
    _, num_heads_kv, _ = kv.shape
    head_dim_v = KV_LORA_RANK
    batch_size = cu_seqlens_q.shape[0] - 1
    seqlen_q = max_seqlen_q
    seqlen_k = max_seqlen_k
    window_size_left, window_size_right = window_size
    is_local = window_size_left is not None or window_size_right is not None
    softmax_scale = (
        1.0 / (head_dim_qk**0.5) if softmax_scale is None else softmax_scale
    )
    softmax_scale_log2 = softmax_scale * math.log2(math.e)

    qheads_per_kvhead_packgqa = num_heads_q // num_heads_kv

    TILE_K = 512
    TILE_M = 16
    TILE_N = 16
    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 int(window_size_left),
        None if window_size_right is None else int(window_size_right),
    )
    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, preset_boundaries = _BENCHMARK_SPLIT_PRESETS[preset_key]
        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,
        )
        _SPLIT_TENSOR_CACHE[split_cache_key] = (
            split_counts,
            split_boundaries,
            num_splits,
        )
    else:
        split_counts, split_boundaries, num_splits = cached_split_tensors

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

    out_partial = _get_cached_tensor(
        (
            "out_partial",
            device_key,
            num_splits,
            total_seqlen_q,
            num_heads_q,
            head_dim_v,
        ),
        (num_splits, total_seqlen_q, num_heads_q, head_dim_v),
        torch.float32,
        query.device,
    )
    lse_partial = _get_cached_tensor(
        ("lse_partial", device_key, num_splits, num_heads_q, total_seqlen_q),
        (num_splits, num_heads_q, total_seqlen_q),
        torch.float32,
        query.device,
    )

    def grid(META):
        return (
            triton.cdiv(
                seqlen_q * (num_heads_q // num_heads_kv), META["TILE_M"]
            ),
            num_heads_kv,
            batch_size * num_splits,
        )

    _fwd_base_sparse_kernel[grid](
        query,
        q_scale,
        kv,
        kv_scale,
        out_partial,
        lse_partial,
        split_counts,
        split_boundaries,
        softmax_scale_log2,
        0,
        query.stride(-2),
        query.stride(0),
        0,
        kv.stride(-2),
        kv.stride(0),
        0,
        out_partial.stride(-2),
        out_partial.stride(-3),
        out_partial.stride(0),
        0,
        lse_partial.stride(-2),
        lse_partial.stride(0),
        cu_seqlens_q,
        cu_seqlens_k,
        num_splits,
        seqlen_q,
        seqlen_k,
        head_dim_qk,
        head_dim_v,
        QHEADS_PER_KVHEAD_PACKGQA=qheads_per_kvhead_packgqa,
        TILE_M=TILE_M,
        TILE_N=TILE_N,
        TILE_K=TILE_K,
        IS_CAUSAL=is_causal,
        IS_LOCAL=is_local,
        IS_SPLIT_KV=True,
        WINDOW_SIZE_LEFT=window_size_left,
        WINDOW_SIZE_RIGHT=window_size_right,
        HAS_CU_SEQLENS_Q=True,
        HAS_CU_SEQLENS_K=True,
        num_warps=num_warps,
        num_stages=num_stages,
        waves_per_eu=waves_per_eu,
        matrix_instr_nonkdim=matrix_instr_nonkdim,
    )

    _flash_attn_fwd_combine(
        out_partial,
        lse_partial,
        out,
        cu_seqlens_q=cu_seqlens_q,
    )

    return out


def flash_sparse_attn_varlen_forward_func(
    q: torch.Tensor,
    kv: torch.Tensor,
    qo_indptr: torch.Tensor,
    kv_indptr: 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

    window_size_k = 1024 if config["kv_seq_len"] == 1024 else 4096
    out_cache_key = (
        "out_result",
        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_varlen_forward(
        query=q,
        q_scale=q_scale,
        kv=kv,
        kv_scale=kv_scale,
        cu_seqlens_q=qo_indptr,
        cu_seqlens_k=kv_indptr,
        max_seqlen_q=config["q_seq_len"],
        max_seqlen_k=config["kv_seq_len"],
        is_causal=False,
        softmax_scale=config["sm_scale"],
        window_size=(window_size_k, 0),
        out=cached_out,
    )
    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

    # Resolve Q
    if Q_DTYPE == "fp8":
        q_input, q_scale = quantize_fp8(q)
    else:
        q_input, q_scale = q, None

    # Resolve KV
    if KV_DTYPE == "fp8":
        kv_input, kv_scale = kv_data["fp8"]
    else:
        kv_input, kv_scale = kv_data["bf16"], None

    return flash_sparse_attn_varlen_forward_func(
        q_input, kv_input, qo_indptr, kv_indptr, config,
        q_scale=q_scale, kv_scale=kv_scale,
    )


# ---------------------------------------------------------------------------
# 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,
#     }

#     if KV_DTYPE == "fp8":
#         kv_data["fp8"] = quantize_fp8(kv_buffer_bf16)
#     elif KV_DTYPE == "mxfp4":
#         kv_data["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 _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:
#             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())
#         torch.testing.assert_close(out.float(), ref.float(), rtol=1e-1, atol=1e-1)

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


# def _run_speed_test(
#     case: dict[str, int],
#     warmup: int,
#     repeats: int,
# ) -> dict[str, float]:
#     data = generate_input(
#         case["batchsize"],
#         case["qseqlen"],
#         case["kvseqlen"],
#         case["seed"],
#     )
#     _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):
#             start_event.record()
#             custom_kernel(data)
#             end_event.record()
#             end_event.synchronize()
#             durations_us.append(start_event.elapsed_time(end_event) * 1000.0)
#     else:
#         durations_us = []
#         for _ in range(repeats):
#             start_ns = time.perf_counter_ns()
#             custom_kernel(data)
#             durations_us.append((time.perf_counter_ns() - start_ns) / 1000.0)

#     total_q = case["batchsize"] * case["qseqlen"]
#     total_kv = case["batchsize"] * 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):
#         label = (
#             f"bs={case['batchsize']} q={case['qseqlen']} "
#             f"kv={case['kvseqlen']} seed={case['seed']}"
#         )
#         print(f"case[{index}].spec: {label}")

#         if run_correctness:
#             correctness = _run_correctness_test(case)
#             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, warmup=warmup, repeats=repeats)
#             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 · 2700 lines total

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

Best evidence level for this revision: reported

JSON