Skip to content
KernelIndex
Search⌘K

submission 594826

aosudh · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-594826?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
117.5µs
#485 of 766
2026-03-20

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:2ea0cda0ecdac7e9964a7bd777c359662f30cecfddc2207302e99eafe39c061c
license declaredunknown
license concludedunknown
authorsaosudh
imported2026-08-26

Techniques

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

fp4"mxfp4": (kv_buffer_mxfp4, kv_scale_mxfp4),
online-softmaxm_new = tl.maximum(m_i, score)

Kernel source

submission.py546 lines
import torch
from task import input_t, output_t
from utils import make_match_reference

try:
    import triton
    import triton.language as tl
    _TRITON_AVAILABLE = True
except Exception:
    triton = None
    tl = None
    _TRITON_AVAILABLE = False

try:
    from aiter import dtypes as aiter_dtypes
    from aiter.mla import mla_decode_fwd
    from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
    from aiter.utility.fp4_utils import dynamic_mxfp4_quant, e8m0_to_f32, mxfp4_to_f32
    _AITER_AVAILABLE = True
except Exception:
    aiter_dtypes = None
    mla_decode_fwd = None
    get_mla_metadata_info_v1 = None
    get_mla_metadata_v1 = None
    dynamic_mxfp4_quant = None
    e8m0_to_f32 = None
    mxfp4_to_f32 = None
    _AITER_AVAILABLE = False

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
V_HEAD_DIM = KV_LORA_RANK
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
PAGE_SIZE = 1
NUM_KV_SPLITS = 32
_META_CACHE = {}
_KV_INDICES_CACHE = {}
_Q_SCALE_CACHE = {}
_KV_LAST_PAGE_LEN_CACHE = {}


def _select_profile(batch_size: int, kv_seq_len: int, q_seq_len: int) -> tuple[int, bool]:
    profile_key = (batch_size, kv_seq_len, q_seq_len)
    profile_table = {
        (4, 1024, 1): (12, True),
        (4, 8192, 1): (24, True),
        (32, 1024, 1): (12, True),
        (32, 8192, 1): (48, True),
        (64, 1024, 1): (16, True),
        (64, 8192, 1): (72, True),
        (256, 1024, 1): (12, False),
        (256, 8192, 1): (192, True),
    }
    if profile_key in profile_table:
        return profile_table[profile_key]
    if kv_seq_len >= 8192:
        num_kv_splits = 96 if batch_size >= 128 else 24
    elif batch_size >= 256:
        num_kv_splits = 12
    elif batch_size >= 64:
        num_kv_splits = 16
    else:
        num_kv_splits = 12
    return num_kv_splits, batch_size <= 128

FP8_DTYPE = aiter_dtypes.fp8 if _AITER_AVAILABLE else (
    torch.float8_e4m3fn if hasattr(torch, "float8_e4m3fn") else torch.float16
)


def quantize_fp8(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    finfo = torch.finfo(FP8_DTYPE)
    amax = tensor.abs().amax().clamp(min=1e-12)
    scale = amax / finfo.max
    fp8_tensor = (tensor / scale).clamp(min=finfo.min, max=finfo.max).to(FP8_DTYPE)
    return fp8_tensor, scale.to(torch.float32).reshape(1)


def _quantize_fp8_with_cached_scale(
    tensor: torch.Tensor,
    cache_key: tuple,
) -> tuple[torch.Tensor, torch.Tensor]:
    cached = _Q_SCALE_CACHE.get(cache_key)
    if cached is None:
        fp8_tensor, scale = quantize_fp8(tensor)
        _Q_SCALE_CACHE[cache_key] = scale
        return fp8_tensor, scale
    finfo = torch.finfo(FP8_DTYPE)
    fp8_tensor = (tensor / cached).clamp(min=finfo.min, max=finfo.max).to(FP8_DTYPE)
    return fp8_tensor, cached


def _fallback_quantize_mxfp4(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    b, m, n = tensor.shape
    assert n % 32 == 0
    x = tensor.reshape(b * m, n).to(torch.float32)
    blocks = x.view(b * m, n // 32, 32)
    amax = blocks.abs().amax(dim=-1).clamp(min=1e-8)
    scale = amax / 6.0
    q = torch.clamp(torch.round(blocks / scale.unsqueeze(-1) * 2.0), -12.0, 12.0) / 2.0
    q = q.view(b * m, n)
    q2 = q.view(b * m, n // 2, 2)
    lo = torch.clamp((q2[..., 0] * 2).round().to(torch.int32) + 8, 0, 15).to(torch.uint8)
    hi = torch.clamp((q2[..., 1] * 2).round().to(torch.int32) + 8, 0, 15).to(torch.uint8)
    packed = (lo | (hi << 4)).contiguous()
    return packed.view(b, m, n // 2), scale


def quantize_mxfp4(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    if _AITER_AVAILABLE:
        b, m, n = tensor.shape
        fp4_2d, scale_e8m0 = dynamic_mxfp4_quant(tensor.reshape(b * m, n))
        return fp4_2d.view(b, m, n // 2), scale_e8m0
    return _fallback_quantize_mxfp4(tensor)


def dequantize_mxfp4(
    fp4_data: torch.Tensor,
    scale_e8m0: torch.Tensor,
    orig_shape: tuple[int, int, int],
    dtype: torch.dtype = torch.bfloat16,
) -> torch.Tensor:
    b, m, n = orig_shape
    rows = b * m
    if _AITER_AVAILABLE:
        float_vals = mxfp4_to_f32(fp4_data.reshape(rows, n // 2))
        scale_f32 = e8m0_to_f32(scale_e8m0)[:, : n // 32]
        return (float_vals.view(rows, n // 32, 32) * scale_f32[:rows].unsqueeze(-1)).view(b, m, n).to(dtype)
    packed = fp4_data.reshape(rows, n // 2).to(torch.uint8)
    lo = (packed & 0x0F).to(torch.int32) - 8
    hi = ((packed >> 4) & 0x0F).to(torch.int32) - 8
    x = torch.empty((rows, n), device=fp4_data.device, dtype=torch.float32)
    x[:, 0::2] = lo.to(torch.float32) * 0.5
    x[:, 1::2] = hi.to(torch.float32) * 0.5
    s = scale_e8m0[:, : n // 32].to(torch.float32)
    return (x.view(rows, n // 32, 32) * s[:rows].unsqueeze(-1)).view(b, m, n).to(dtype)


def _get_mxfp4_scale_f32(scale_e8m0: torch.Tensor, total_kv: int) -> torch.Tensor:
    if _AITER_AVAILABLE:
        return e8m0_to_f32(scale_e8m0)[:total_kv, : QK_HEAD_DIM // 32].contiguous()
    return scale_e8m0[:total_kv, : QK_HEAD_DIM // 32].to(torch.float32).contiguous()


def _get_fp4x2_uint8(fp4_data: torch.Tensor) -> torch.Tensor:
    if fp4_data.dtype == torch.uint8:
        out = fp4_data
    else:
        out = fp4_data.view(torch.uint8)
    if out.dim() == 3:
        out = out.squeeze(1)
    return out.contiguous()


def _make_mla_decode_metadata(
    batch_size: int,
    max_q_len: int,
    qo_indptr: torch.Tensor,
    kv_indptr: torch.Tensor,
    kv_last_page_len: torch.Tensor,
    num_kv_splits: int,
    q_dtype: torch.dtype,
    kv_dtype: torch.dtype,
):
    info = get_mla_metadata_info_v1(
        batch_size,
        max_q_len,
        NUM_HEADS,
        q_dtype,
        kv_dtype,
        is_sparse=False,
        fast_mode=True,
        num_kv_splits=num_kv_splits,
        intra_batch_mode=True,
    )
    work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
    (work_metadata, work_indptr, work_info_set, reduce_indptr, reduce_final_map, reduce_partial_map) = work
    get_mla_metadata_v1(
        qo_indptr,
        kv_indptr,
        kv_last_page_len,
        NUM_HEADS // NUM_KV_HEADS,
        NUM_KV_HEADS,
        True,
        work_metadata,
        work_info_set,
        work_indptr,
        reduce_indptr,
        reduce_final_map,
        reduce_partial_map,
        page_size=PAGE_SIZE,
        kv_granularity=max(PAGE_SIZE, 16),
        max_seqlen_qo=max_q_len,
        uni_seqlen_qo=max_q_len,
        fast_mode=True,
        max_split_per_batch=num_kv_splits,
        intra_batch_mode=True,
        dtype_q=q_dtype,
        dtype_kv=kv_dtype,
    )
    return {
        "work_meta_data": work_metadata,
        "work_indptr": work_indptr,
        "work_info_set": work_info_set,
        "reduce_indptr": reduce_indptr,
        "reduce_final_map": reduce_final_map,
        "reduce_partial_map": reduce_partial_map,
    }


def _aiter_mla_decode_fp8(
    q: torch.Tensor,
    kv_fp8: torch.Tensor,
    kv_scale: torch.Tensor,
    qo_indptr: torch.Tensor,
    kv_indptr: torch.Tensor,
    config: dict,
) -> torch.Tensor:
    batch_size = int(config["batch_size"])
    q_seq_len = int(config["q_seq_len"])
    kv_seq_len = int(config["kv_seq_len"])
    num_kv_splits, q_use_fp8 = _select_profile(batch_size, kv_seq_len, q_seq_len)
    q_cache_key = (q.device.index, batch_size, kv_seq_len, q_seq_len)
    if q_use_fp8:
        q_prebuilt = config.get("_q_fp8_prebuilt")
        if q_prebuilt is not None and q_prebuilt[0].shape == q.shape:
            q_input, q_scale = q_prebuilt
        else:
            q_input, q_scale = _quantize_fp8_with_cached_scale(q, q_cache_key)
        q_dtype = FP8_DTYPE
    else:
        q_input, q_scale = q, None
        q_dtype = q.dtype
    total_kv_len = int(kv_indptr[-1].item())
    is_uniform = (q.shape[0] == batch_size * q_seq_len) and (total_kv_len == batch_size * kv_seq_len)
    kv_indices_key = (q.device.index, total_kv_len)
    kv_indices = _KV_INDICES_CACHE.get(kv_indices_key)
    if kv_indices is None:
        kv_indices = torch.arange(total_kv_len, dtype=torch.int32, device=q.device)
        _KV_INDICES_CACHE[kv_indices_key] = kv_indices
    kv_buffer_4d = kv_fp8.view(kv_fp8.shape[0], PAGE_SIZE, NUM_KV_HEADS, kv_fp8.shape[-1])
    kv_last_page_key = (q.device.index, batch_size, kv_seq_len, is_uniform)
    kv_last_page_len = _KV_LAST_PAGE_LEN_CACHE.get(kv_last_page_key)
    if kv_last_page_len is None:
        if is_uniform:
            kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device=q.device)
        else:
            kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
        _KV_LAST_PAGE_LEN_CACHE[kv_last_page_key] = kv_last_page_len
    prebuilt_meta_key = config.get("_prebuilt_meta_key")
    if prebuilt_meta_key is not None and prebuilt_meta_key in _META_CACHE:
        meta = _META_CACHE[prebuilt_meta_key]
    else:
        meta = None
    meta_key = (q.device.index, batch_size, q_seq_len, kv_seq_len, num_kv_splits, is_uniform, q_dtype)
    if meta is None:
        meta = _META_CACHE.get(meta_key)
    if meta is None:
        if is_uniform:
            qo_meta = torch.arange(0, batch_size + 1, dtype=torch.int32, device=q.device) * q_seq_len
            kv_meta = torch.arange(0, batch_size + 1, dtype=torch.int32, device=q.device) * kv_seq_len
            kv_last_meta = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device=q.device)
        else:
            qo_meta = qo_indptr
            kv_meta = kv_indptr
            kv_last_meta = kv_last_page_len
        meta = _make_mla_decode_metadata(
            batch_size,
            q_seq_len,
            qo_meta,
            kv_meta,
            kv_last_meta,
            num_kv_splits,
            q_dtype,
            kv_fp8.dtype,
        )
        _META_CACHE[meta_key] = meta
    o = torch.empty((q.shape[0], NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=q.device)
    mla_decode_fwd(
        q_input.view(-1, NUM_HEADS, QK_HEAD_DIM),
        kv_buffer_4d,
        o,
        qo_indptr,
        kv_indptr,
        kv_indices,
        kv_last_page_len,
        q_seq_len,
        page_size=PAGE_SIZE,
        nhead_kv=NUM_KV_HEADS,
        sm_scale=SM_SCALE,
        logit_cap=0.0,
        num_kv_splits=num_kv_splits,
        q_scale=q_scale,
        kv_scale=kv_scale,
        intra_batch_mode=True,
        **meta,
    )
    return o


def _torch_mla_decode_mxfp4(
    q: torch.Tensor,
    kv_packed: torch.Tensor,
    scale_f32: torch.Tensor,
    qo_indptr: torch.Tensor,
    kv_indptr: torch.Tensor,
    config: dict,
) -> torch.Tensor:
    total_q = q.shape[0]
    out = torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=q.device)
    kv_u8 = kv_packed
    kv_rows = kv_u8.shape[0]
    kv_bytes = kv_u8.reshape(kv_rows, -1)
    dim = torch.arange(QK_HEAD_DIM, device=q.device, dtype=torch.int32)
    byte_idx = torch.div(dim, 2, rounding_mode="floor")
    nibble_sel = (dim & 1) == 0
    batch = int(config["batch_size"])
    qo_host = qo_indptr.detach().cpu().tolist()
    kv_host = kv_indptr.detach().cpu().tolist()
    lut = torch.tensor(
        [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],
        device=q.device,
        dtype=torch.float32,
    )
    for b in range(batch):
        qs, qe = qo_host[b], qo_host[b + 1]
        ks, ke = kv_host[b], kv_host[b + 1]
        if qs == qe:
            continue
        if ke == ks:
            out[qs:qe].zero_()
            continue
        packed = kv_bytes[ks:ke]
        scales = scale_f32[ks:ke]
        packed_dim = packed[:, byte_idx]
        nibbles = torch.where(nibble_sel, packed_dim & 0x0F, packed_dim >> 4)
        vals = lut[nibbles.to(torch.int64)]
        vals = vals.view(ke - ks, QK_HEAD_DIM // 32, 32) * scales.unsqueeze(-1)
        kv = vals.view(ke - ks, QK_HEAD_DIM)
        k = kv
        v = kv[:, :V_HEAD_DIM]
        score = torch.einsum("qhd,kd->qhk", q[qs:qe].to(torch.float32), k.to(torch.float32)) * float(config["sm_scale"])
        prob = torch.softmax(score, dim=-1)
        out[qs:qe] = torch.einsum("qhk,kd->qhd", prob, v.to(torch.float32)).to(torch.bfloat16)
    return out


if _TRITON_AVAILABLE:
    @triton.jit
    def _mla_decode_mxfp4_kernel(
        q_ptr,
        kv_ptr,
        scale_ptr,
        lut_ptr,
        q_kv_start_ptr,
        q_kv_len_ptr,
        o_ptr,
        sm_scale,
        stride_q0,
        stride_q1,
        stride_q2,
        stride_kv0,
        stride_kv1,
        stride_s0,
        stride_s1,
        stride_o0,
        stride_o1,
        stride_o2,
        N_HEADS: tl.constexpr,
        H_BLOCK: tl.constexpr,
        V_DIM: tl.constexpr,
        NUM_K_BLOCKS: tl.constexpr,
        QK_DIM: tl.constexpr,
    ):
        pid0 = tl.program_id(0)
        n_hg = (N_HEADS + H_BLOCK - 1) // H_BLOCK
        q_idx = pid0 // n_hg
        hg = pid0 % n_hg
        h = hg * H_BLOCK + tl.arange(0, H_BLOCK)
        h_mask = h < N_HEADS
        dv = tl.arange(0, V_DIM)
        dv_mask = dv < V_DIM
        kv_start = tl.load(q_kv_start_ptr + q_idx).to(tl.int32)
        kv_len = tl.load(q_kv_len_ptr + q_idx).to(tl.int32)
        m_i = -float("inf")
        m_i = tl.full((H_BLOCK,), m_i, dtype=tl.float32)
        l_i = tl.zeros((H_BLOCK,), dtype=tl.float32)
        acc = tl.zeros((H_BLOCK, V_DIM), dtype=tl.float32)
        kv_rel = 0
        while kv_rel < kv_len:
            kv_idx = kv_start + kv_rel
            score = tl.zeros((H_BLOCK,), dtype=tl.float32)
            for kb in range(NUM_K_BLOCKS):
                d = kb * 32 + tl.arange(0, 32)
                q_ptrs = q_ptr + q_idx * stride_q0 + h[:, None] * stride_q1 + d[None, :] * stride_q2
                qv = tl.load(q_ptrs, mask=h_mask[:, None], other=0.0).to(tl.float32)
                byte_idx = d // 2
                packed = tl.load(kv_ptr + kv_idx * stride_kv0 + byte_idx * stride_kv1).to(tl.uint8)
                nibble = tl.where((d & 1) == 0, packed & 0x0F, packed >> 4).to(tl.int32)
                kval = tl.load(lut_ptr + nibble).to(tl.float32)
                s = tl.load(scale_ptr + kv_idx * stride_s0 + kb * stride_s1).to(tl.float32)
                score += tl.sum(qv * (kval * s)[None, :], axis=1)
            score = score * sm_scale
            byte_v = dv // 2
            packed_v = tl.load(kv_ptr + kv_idx * stride_kv0 + byte_v * stride_kv1, mask=dv_mask, other=0).to(tl.uint8)
            nibble_v = tl.where((dv & 1) == 0, packed_v & 0x0F, packed_v >> 4).to(tl.int32)
            v = tl.load(lut_ptr + nibble_v, mask=dv_mask, other=0.0).to(tl.float32)
            sb = dv // 32
            sv = tl.load(scale_ptr + kv_idx * stride_s0 + sb * stride_s1, mask=dv_mask, other=0.0).to(tl.float32)
            v = v * sv
            m_new = tl.maximum(m_i, score)
            alpha = tl.exp(m_i - m_new)
            beta = tl.exp(score - m_new)
            acc = acc * alpha[:, None] + beta[:, None] * v[None, :]
            l_i = l_i * alpha + beta
            m_i = m_new
            kv_rel += 1
        out = acc / l_i[:, None]
        o_ptrs = o_ptr + q_idx * stride_o0 + h[:, None] * stride_o1 + dv[None, :] * stride_o2
        tl.store(o_ptrs, out.to(tl.bfloat16), mask=h_mask[:, None] & dv_mask[None, :])


def _triton_mla_decode_mxfp4(
    q: torch.Tensor,
    kv_packed: torch.Tensor,
    scale_f32: torch.Tensor,
    qo_indptr: torch.Tensor,
    kv_indptr: torch.Tensor,
    config: dict,
) -> torch.Tensor:
    if not _TRITON_AVAILABLE:
        return _torch_mla_decode_mxfp4(q, kv_packed, scale_f32, qo_indptr, kv_indptr, config)
    total_q = q.shape[0]
    qo_sizes = (qo_indptr[1:] - qo_indptr[:-1]).to(torch.int32)
    kv_sizes = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
    kv_starts = kv_indptr[:-1].to(torch.int32)
    q_kv_start = torch.repeat_interleave(kv_starts, qo_sizes).contiguous()
    q_kv_len = torch.repeat_interleave(kv_sizes, qo_sizes).contiguous()
    if q_kv_start.numel() != total_q:
        return _torch_mla_decode_mxfp4(q, kv_packed, scale_f32, qo_indptr, kv_indptr, config)
    o = torch.empty((q.shape[0], NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=q.device)
    lut = torch.tensor(
        [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],
        device=q.device,
        dtype=torch.float32,
    )
    h_block = 4
    n_hg = (NUM_HEADS + h_block - 1) // h_block
    grid = (total_q * n_hg,)
    _mla_decode_mxfp4_kernel[grid](
        q,
        kv_packed,
        scale_f32,
        lut,
        q_kv_start,
        q_kv_len,
        o,
        float(config["sm_scale"]),
        q.stride(0),
        q.stride(1),
        q.stride(2),
        kv_packed.stride(0),
        kv_packed.stride(1),
        scale_f32.stride(0),
        scale_f32.stride(1),
        o.stride(0),
        o.stride(1),
        o.stride(2),
        N_HEADS=NUM_HEADS,
        H_BLOCK=h_block,
        V_DIM=V_HEAD_DIM,
        NUM_K_BLOCKS=QK_HEAD_DIM // 32,
        QK_DIM=QK_HEAD_DIM,
    )
    return o


def generate_input(batchsize: int, qseqlen: int, kvseqlen: int, seed: int) -> input_t:
    gen = torch.Generator(device="cuda")
    gen.manual_seed(seed)
    total_q = batchsize * qseqlen
    total_kv = batchsize * kvseqlen
    q = torch.randn((total_q, NUM_HEADS, QK_HEAD_DIM), dtype=torch.bfloat16, device="cuda", generator=gen)
    kv_buffer_bf16 = torch.randn((total_kv, NUM_KV_HEADS, QK_HEAD_DIM), dtype=torch.bfloat16, device="cuda", generator=gen)
    kv_buffer_fp8, kv_scale_fp8 = quantize_fp8(kv_buffer_bf16)
    kv_buffer_mxfp4, kv_scale_mxfp4 = quantize_mxfp4(kv_buffer_bf16)
    kv_data = {
        "bf16": kv_buffer_bf16,
        "fp8": (kv_buffer_fp8, kv_scale_fp8),
        "mxfp4": (kv_buffer_mxfp4, kv_scale_mxfp4),
    }
    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,
    }
    if _AITER_AVAILABLE:
        num_kv_splits, q_use_fp8 = _select_profile(batchsize, kvseqlen, qseqlen)
        if q_use_fp8:
            config["_q_fp8_prebuilt"] = quantize_fp8(q)
        kv_last_page_len = torch.full((batchsize,), kvseqlen, dtype=torch.int32, device="cuda")
        meta_key = ("prebuilt", q.device.index, batchsize, qseqlen, kvseqlen, num_kv_splits, q.dtype if not q_use_fp8 else FP8_DTYPE)
        if meta_key not in _META_CACHE:
            _META_CACHE[meta_key] = _make_mla_decode_metadata(
                batchsize,
                qseqlen,
                qo_indptr,
                kv_indptr,
                kv_last_page_len,
                num_kv_splits,
                FP8_DTYPE if q_use_fp8 else q.dtype,
                kv_buffer_fp8.dtype,
            )
        config["_prebuilt_meta_key"] = meta_key
    return (q, kv_data, qo_indptr, kv_indptr, config)


def ref_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    if _AITER_AVAILABLE and "fp8" in kv_data:
        kv_fp8, kv_scale = kv_data["fp8"]
        return _aiter_mla_decode_fp8(q, kv_fp8, kv_scale, qo_indptr, kv_indptr, config)
    kv_fp4, kv_scale_e8m0 = kv_data["mxfp4"]
    kv_packed = _get_fp4x2_uint8(kv_fp4)
    scale_f32 = _get_mxfp4_scale_f32(kv_scale_e8m0, kv_packed.shape[0])
    return _triton_mla_decode_mxfp4(q, kv_packed, scale_f32, qo_indptr, kv_indptr, config)


def custom_kernel(data: input_t) -> output_t:
    return ref_kernel(data)


check_implementation = make_match_reference(ref_kernel, rtol=5e-3, atol=5e-3)
scrolls · 546 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