Skip to content
KernelIndex
Search⌘K

submission 754086

PromptForcePrime · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

solution.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-754086?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
34.9µs
#70 of 766
2026-04-07

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:9461a2065034781f89c887ebeb7c6c13e1f1d2551088e8137a226af0db91ef45
license declaredunknown
license concludedunknown
authorsPromptForcePrime
imported2026-08-15

Kernel source

solution.py220 lines
"""
solution.py — v94: right-sized workspace (ws_ms=ns).

v93 had ws_ms=32/64 but actual splits (ns) often just 1.
This caused 32x oversized logits_buf/lse_buf (e.g. 256MB vs 8MB).
v94 sets ws_ms=ns so _meta_info sizes workspace exactly for the
actual split count, reducing allocation overhead and cache waste.
"""

import sys as _sys
import torch

_NUM_CUS = 304
_MAX_TOKS_PER_SPLIT = 2048
_MAX_KV_TOKENS = 256 * 8192

_mla_fwd = None
_meta_info = None
_meta_fn = None
_stage1 = None
_reduce = None


def _ensure_aiter():
    global _mla_fwd, _meta_info, _meta_fn, _stage1, _reduce
    if _mla_fwd is not None:
        return
    from aiter.mla import mla_decode_fwd
    from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
    _mla_fwd = mla_decode_fwd
    _meta_info = get_mla_metadata_info_v1
    _meta_fn = get_mla_metadata_v1
    try:
        import aiter as _a
        s1 = getattr(_a, 'mla_decode_stage1_asm_fwd', None)
        rd = getattr(_a, 'mla_reduce_v1', None)
        if callable(s1) and callable(rd):
            _stage1 = s1
            _reduce = rd
    except Exception:
        pass


_cache = {}
_kvi = {}
_qs = None
_direct_ok = True
_last_q_ptr = {}


def _get_config(bs, kvl):
    if bs <= 8:
        ps, ibm = 1, True
    elif kvl >= 4096:
        ps, ibm = 8, False
    else:
        ps, ibm = 2, False

    if 32 <= bs < 256:
        ns = 1
    elif bs >= 256:
        if kvl >= 4096:
            ns = max(1, (kvl + _MAX_TOKS_PER_SPLIT - 1) // _MAX_TOKS_PER_SPLIT)
        else:
            ns = 1
    else:
        ms = 64 if (bs <= 8 and kvl >= 4096) else 32
        tok_ceil = kvl // 64
        if kvl >= 4096:
            ns = min(tok_ceil, ms)
        else:
            one_wave = (_NUM_CUS + bs - 1) // bs
            half_wave = max(1, one_wave // 2)
            ns = min(tok_ceil, ms)
            for s in (1, 2, 4, 8, 16):
                if s >= half_wave and s <= tok_ceil and s <= ms:
                    ns = s
                    break

    gran = max(1, (64 if kvl >= 4096 else 16) // ps)
    return ps, ibm, ns, gran


def _aiter_mla_decode(q, kv_fp8, kv_scale, qo_indptr, kv_indptr,
                       bs, nh, nkv, dq, dv, sm, qsl):
    global _qs, _direct_ok
    _ensure_aiter()

    fp8d = kv_fp8.dtype
    tkv = kv_fp8.shape[0]
    kvl = tkv // bs
    ck = (bs, tkv)

    if ck not in _cache:
        ps, ibm, ns, gran = _get_config(bs, kvl)

        if ps not in _kvi:
            _kvi[ps] = torch.arange(_MAX_KV_TOKENS // ps,
                                     dtype=torch.int32, device=q.device)

        if _qs is None:
            _qs = torch.ones(1, dtype=torch.float32, device="cuda")

        info = _meta_info(
            bs, qsl, nh, fp8d, fp8d,
            is_sparse=False, fast_mode=False,
            num_kv_splits=ns, intra_batch_mode=ibm,
        )
        wt = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
        wm, wi, wis, ri, rfm, rpm = wt
        tq = q.shape[0]
        ot = torch.empty((tq, nh, dv), dtype=torch.bfloat16, device="cuda")
        sl = kv_indptr[1:] - kv_indptr[:-1]
        lpl = ((sl - 1) % ps + 1).to(torch.int32)
        ip = (kv_indptr // ps).to(torch.int32)

        _meta_fn(
            qo_indptr, ip, lpl,
            nh // nkv, nkv, False,
            wm, wis, wi, ri, rfm, rpm,
            page_size=ps, kv_granularity=gran,
            max_seqlen_qo=qsl, uni_seqlen_qo=qsl,
            fast_mode=False, max_split_per_batch=ns,
            intra_batch_mode=ibm,
            dtype_q=fp8d, dtype_kv=fp8d,
        )

        qv_buf = torch.empty((tq, nh, dq), dtype=fp8d, device="cuda")
        rpm_rows = rpm.size(0)
        logits_buf = torch.empty((rpm_rows, 1, nh, dv),
                                  dtype=torch.float32, device="cuda")
        lse_buf = torch.empty((rpm_rows, 1, nh, 1),
                               dtype=torch.float32, device="cuda")

        _cache[ck] = (ps, ns, _kvi[ps], wm, wi, wis, ri, rfm, rpm,
                      ot, lpl, ip, ibm, fp8d, qv_buf, logits_buf, lse_buf)

    (ps, ns, kvi, wm, wi, wis, ri, rfm, rpm,
     o, lpl, ip, ibm, fp8d, qv_buf, logits_buf, lse_buf) = _cache[ck]

    kv4d = kv_fp8.view(tkv // ps, ps, nkv, dq)

    q_ptr = q.data_ptr()
    if _last_q_ptr.get(ck) != q_ptr:
        qv_buf.copy_(q.view(-1, nh, dq))
        _last_q_ptr[ck] = q_ptr

    # --- Direct ops (primary path) ---
    if _stage1 is not None and _direct_ok:
        try:
            _stage1(
                qv_buf, kv4d,
                qo_indptr, ip, kvi, lpl,
                None, wm, wi, wis,
                qsl, ps, nkv, sm,
                logits_buf, lse_buf, o,
                _qs, kv_scale,
            )
            _reduce(
                logits_buf, lse_buf,
                ri, rfm, rpm,
                qsl, o, None,
            )
            return o
        except Exception as e:
            _direct_ok = False
            print(f"[v92] direct ops failed: {type(e).__name__}: {e}",
                  file=_sys.stderr)

    # --- Wrapper fallback ---
    _mla_fwd(
        qv_buf, kv4d, o,
        qo_indptr, ip, kvi, lpl,
        qsl,
        page_size=ps, nhead_kv=nkv,
        sm_scale=sm, logit_cap=0.0,
        num_kv_splits=ns,
        q_scale=_qs, kv_scale=kv_scale,
        intra_batch_mode=ibm,
        work_meta_data=wm, work_indptr=wi, work_info_set=wis,
        reduce_indptr=ri, reduce_final_map=rfm, reduce_partial_map=rpm,
    )
    return o


def _naive_fallback(q, kv_data, qo_indptr, kv_indptr, config):
    kv = kv_data["bf16"]
    bs = config["batch_size"]
    sm = config["sm_scale"]
    dv = config["v_head_dim"]
    out = []
    for b in range(bs):
        qs, qe = int(qo_indptr[b].item()), int(qo_indptr[b + 1].item())
        ks, ke = int(kv_indptr[b].item()), int(kv_indptr[b + 1].item())
        qb = q[qs:qe].float()
        kb = kv[ks:ke, 0, :].float()
        sc = torch.einsum("qhd,kd->qhk", qb, kb) * sm
        at = torch.softmax(sc, dim=-1)
        out.append(torch.einsum("qhk,kd->qhd", at, kb[:, :dv]))
    return torch.cat(out, dim=0).to(torch.bfloat16)


def custom_kernel(data):
    q, kv_data, qo_indptr, kv_indptr, config = data
    if q.is_cuda and "fp8" in kv_data:
        fp8_pair = kv_data["fp8"]
        if isinstance(fp8_pair, (tuple, list)) and len(fp8_pair) == 2:
            try:
                return _aiter_mla_decode(
                    q, fp8_pair[0], fp8_pair[1],
                    qo_indptr, kv_indptr,
                    config["batch_size"], config["num_heads"],
                    config["num_kv_heads"], config["qk_head_dim"],
                    config["v_head_dim"], config["sm_scale"],
                    config["q_seq_len"],
                )
            except Exception as e:
                print(f"[aiter fail] {type(e).__name__}: {e}", file=_sys.stderr)
    return _naive_fallback(q, kv_data, qo_indptr, kv_indptr, config)
scrolls · 220 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