Skip to content
KernelIndex
Search⌘K

submission 722487

Lewen-Cai · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-722487?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
90.7µs
#421 of 766
2026-04-04

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:cd176293f41cf3f23fcdd4c4d81210b1ef4ced16b32a825abb41512a652609a1
license declaredunknown
license concludedunknown
authorsLewen-Cai
imported2026-08-26

Techniques

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

persistent-kernel1. FP8 a8w8 kernel via aiter persistent mode

Kernel source

submission.py283 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X

"""
MLA decode kernel with aiter acceleration for AMD GPUs.

Optimizations over reference:
  1. FP8 a8w8 kernel via aiter persistent mode
  2. Safe metadata caching by batch geometry (scalar keys, not data_ptr)
  3. Pre-allocated output buffer and kv_indices reuse across calls
  4. kv_last_page_len reuse to avoid repeated tensor ops
"""

import torch
import torch.nn.functional as F

# ---------------------------------------------------------------------------
# Aiter availability — graceful fallback when not installed
# ---------------------------------------------------------------------------
_HAS_AITER = False
try:
    from aiter.mla import mla_decode_fwd
    from aiter import dtypes as aiter_dtypes
    from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
    _HAS_AITER = True
except ImportError:
    pass

# ---------------------------------------------------------------------------
# Constants
# ---------------------------------------------------------------------------
PAGE_SIZE = 1
NUM_KV_SPLITS = 32

if _HAS_AITER:
    _FP8_DTYPE = aiter_dtypes.fp8

# ---------------------------------------------------------------------------
# Geometry-based caches (safe: keyed by scalar batch params, not pointers)
# ---------------------------------------------------------------------------
_metadata_cache: dict = {}
_kv_indices_cache: dict = {}
_output_buf_cache: dict = {}
_kv_last_page_cache: dict = {}


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

def _quantize_fp8(tensor: torch.Tensor) -> tuple:
    finfo = torch.finfo(_FP8_DTYPE)
    amax = tensor.abs().amax().clamp(min=1e-12)
    scale = amax / finfo.max
    fp8 = (tensor / scale).clamp(min=finfo.min, max=finfo.max).to(_FP8_DTYPE)
    return fp8, scale.to(torch.float32).reshape(1)


# ---------------------------------------------------------------------------
# Cached helper allocations
# ---------------------------------------------------------------------------

def _get_kv_indices(total_kv: int) -> torch.Tensor:
    cached = _kv_indices_cache.get(total_kv)
    if cached is not None:
        return cached
    indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")
    _kv_indices_cache[total_kv] = indices
    return indices


def _get_output_buffer(total_q: int, H: int, Dv: int) -> torch.Tensor:
    key = (total_q, H, Dv)
    cached = _output_buf_cache.get(key)
    if cached is not None:
        return cached
    buf = torch.empty(key, dtype=torch.bfloat16, device="cuda")
    _output_buf_cache[key] = buf
    return buf


def _get_kv_last_page_len(kv_indptr: torch.Tensor, batch_size: int, kv_len: int) -> torch.Tensor:
    key = (batch_size, kv_len)
    cached = _kv_last_page_cache.get(key)
    if cached is not None:
        return cached
    kv_last = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
    _kv_last_page_cache[key] = kv_last
    return kv_last


# ---------------------------------------------------------------------------
# Safe metadata caching by batch geometry
# ---------------------------------------------------------------------------

def _build_persistent_metadata(
    batch_size, max_q_len, nhead, nhead_kv,
    q_dtype, kv_dtype, qo_indptr, kv_indptr, kv_last_page_len,
    kv_len,
):
    cache_key = (batch_size, max_q_len, kv_len, nhead, nhead_kv, q_dtype, kv_dtype)
    cached = _metadata_cache.get(cache_key)
    if cached is not None:
        return cached

    info = get_mla_metadata_info_v1(
        batch_size, max_q_len, nhead, q_dtype, kv_dtype,
        is_sparse=False, fast_mode=False,
        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,
        nhead // nhead_kv, nhead_kv, 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=False,
        max_split_per_batch=NUM_KV_SPLITS,
        intra_batch_mode=True,
        dtype_q=q_dtype,
        dtype_kv=kv_dtype,
    )

    meta = {
        "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,
    }
    _metadata_cache[cache_key] = meta
    return meta


# ---------------------------------------------------------------------------
# Aiter path: persistent-mode FP8 a8w8 MLA decode
# ---------------------------------------------------------------------------

def _aiter_decode(q, kv_data, qo_indptr, kv_indptr, config):
    B = config["batch_size"]
    H = config["num_heads"]
    HK = config["num_kv_heads"]
    D = config["qk_head_dim"]
    Dv = config["v_head_dim"]
    scale = config["sm_scale"]
    max_q = config["q_seq_len"]

    # FP8 quantize Q on-the-fly
    q_fp8, q_scale = _quantize_fp8(q)

    # Use pre-quantised FP8 KV; fall back to bf16 + on-the-fly quant
    if "fp8" in kv_data:
        kv_fp8, kv_scale = kv_data["fp8"]
    else:
        kv_fp8, kv_scale = _quantize_fp8(kv_data["bf16"])

    total_kv = kv_fp8.shape[0]
    kv_len = total_kv // B
    kv_4d = kv_fp8.view(total_kv, PAGE_SIZE, HK, D)

    # Cached allocations — avoid per-call torch.arange / torch.empty / subtraction
    kv_indices = _get_kv_indices(total_kv)
    kv_last_page_len = _get_kv_last_page_len(kv_indptr, B, kv_len)
    o = _get_output_buffer(q.shape[0], H, Dv)

    # Cached metadata — keyed by (batch_size, q_len, kv_len, heads, dtypes)
    meta = _build_persistent_metadata(
        B, max_q, H, HK,
        q_fp8.dtype, kv_fp8.dtype,
        qo_indptr, kv_indptr, kv_last_page_len,
        kv_len,
    )

    mla_decode_fwd(
        q_fp8.view(-1, H, D),
        kv_4d,
        o,
        qo_indptr,
        kv_indptr,
        kv_indices,
        kv_last_page_len,
        max_q,
        page_size=PAGE_SIZE,
        nhead_kv=HK,
        sm_scale=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


# ---------------------------------------------------------------------------
# SDPA fallback (no aiter / non-AMD)
# ---------------------------------------------------------------------------

def _sdpa_decode(q, kv_data, qo_indptr, kv_indptr, config):
    B = config["batch_size"]
    H = config["num_heads"]
    HK = config["num_kv_heads"]
    D = config["qk_head_dim"]
    Dv = config["v_head_dim"]
    scale = config["sm_scale"]

    kv_buf = kv_data["bf16"]
    K = kv_buf
    V = kv_buf[..., :Dv].contiguous()

    kv_lens = kv_indptr[1:] - kv_indptr[:-1]
    q_lens = qo_indptr[1:] - qo_indptr[:-1]

    all_equal_kv = (kv_lens == kv_lens[0]).all().item()
    all_equal_q = (q_lens == q_lens[0]).all().item()

    if all_equal_kv and all_equal_q:
        return _sdpa_equal(q, K, V, B, H, HK, D, Dv, scale,
                           q_lens[0].item(), kv_lens[0].item())
    return _sdpa_variable(q, K, V, qo_indptr, kv_indptr,
                          B, H, HK, D, Dv, scale, q_lens, kv_lens)


def _sdpa_equal(q, K, V, B, H, HK, D, Dv, scale, Lq, Lk):
    Qs = q.view(B, Lq, H, D).permute(0, 2, 1, 3).contiguous()
    Ks = K.view(B, Lk, HK, D).permute(0, 2, 1, 3).contiguous()
    Vs = V.view(B, Lk, HK, Dv).permute(0, 2, 1, 3).contiguous()
    O = F.scaled_dot_product_attention(Qs, Ks, Vs, scale=scale)
    return O.permute(0, 2, 1, 3).reshape(-1, H, Dv).contiguous()


def _sdpa_variable(q, K, V, qo_indptr, kv_indptr,
                   B, H, HK, D, Dv, scale, q_lens, kv_lens):
    max_kv = kv_lens.max().item()
    max_q = q_lens.max().item()

    Q_pad = q.new_zeros(B, max_q, H, D)
    K_pad = K.new_zeros(B, max_kv, HK, D)
    V_pad = K.new_zeros(B, max_kv, HK, Dv)
    kv_mask = torch.ones(B, max_kv, dtype=torch.bool, device=q.device)

    for i in range(B):
        qs, qe = int(qo_indptr[i]), int(qo_indptr[i + 1])
        ks, ke = int(kv_indptr[i]), int(kv_indptr[i + 1])
        ql, kl = qe - qs, ke - ks
        Q_pad[i, :ql] = q[qs:qe].view(ql, H, D)
        K_pad[i, :kl] = K[ks:ke].view(kl, HK, D)
        V_pad[i, :kl] = V[ks:ke].view(kl, HK, Dv)
        kv_mask[i, kl:] = False

    Qs = Q_pad.permute(0, 2, 1, 3)
    Ks = K_pad.permute(0, 2, 1, 3)
    Vs = V_pad.permute(0, 2, 1, 3)
    attn_mask = kv_mask.unsqueeze(1).unsqueeze(2)
    O = F.scaled_dot_product_attention(Qs, Ks, Vs, attn_mask=attn_mask, scale=scale)

    results = []
    for i in range(B):
        ql = q_lens[i].item()
        results.append(O[i, :, :ql].permute(1, 0, 2))
    return torch.cat(results, dim=0).contiguous()


# ---------------------------------------------------------------------------
# Entry point
# ---------------------------------------------------------------------------

def custom_kernel(data):
    q, kv_data, qo_indptr, kv_indptr, config = data

    if _HAS_AITER and torch.cuda.is_available():
        return _aiter_decode(q, kv_data, qo_indptr, kv_indptr, config)
    return _sdpa_decode(q, kv_data, qo_indptr, kv_indptr, config)
scrolls · 283 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