Skip to content
KernelIndex
Search⌘K

submission 677054

johnny.t.shi · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

v90.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-677054?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
33.2µs
#42 of 766
2026-03-31

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:a642fa6770dbf638e697a5a5b53e598b2c2962d443ddb41a107dd0b70b7a8f32
license declaredunknown
license concludedunknown
authorsjohnny.t.shi
imported2026-08-15

Techniques

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

persistent-kernel- bs=4/kv=1024: BF16 NP → BF16 persistent (-1.2µs) ← KEPT

Kernel source

v90.py137 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""MLA v90 — v89 with bs=64/kv=1024 reverted to FP8 NP (BF16p failed leaderboard).

Changes from v86:
- bs=4/kv=1024: BF16 NP → BF16 persistent (-1.2µs) ← KEPT
- bs=64/kv=1024: stays FP8 NP splits=1 (BF16p fails secret runner) ← REVERTED
"""
from task import input_t, output_t
import torch
import aiter as _aiter
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

_mc = {}
_tc = {}
_ic = {}
_oc = {}

def _gt(key, fn):
    v = _tc.get(key)
    if v is None: v = fn(); _tc[key] = v
    return v

def _gm(bs, tot, nh, ns, ps, qi, ki, kl, dev, qd, kdd):
    key = (bs, tot, nh, ns, ps, str(qd), str(kdd))
    v = _mc.get(key)
    if v is not None: return v
    info = get_mla_metadata_info_v1(bs, 1, nh, qd, kdd, is_sparse=False, fast_mode=True, num_kv_splits=ns, intra_batch_mode=False)
    work = [torch.empty(s, dtype=t, device=dev) for s, t in info]
    wmd, wi, wis, ri, rfm, rpm = work
    get_mla_metadata_v1(qi, ki, kl, nh, 1, False, wmd, wis, wi, ri, rfm, rpm,
        page_size=ps, kv_granularity=max(ps,16), max_seqlen_qo=1, uni_seqlen_qo=1,
        fast_mode=True, max_split_per_batch=ns, intra_batch_mode=False, dtype_q=qd, dtype_kv=kdd)
    v = dict(work_meta_data=wmd, work_indptr=wi, work_info_set=wis, reduce_indptr=ri, reduce_final_map=rfm, reduce_partial_map=rpm)
    _mc[key] = v; return v

def _gi(key, fn):
    v = _ic.get(key)
    if v is None: v = fn(); _ic[key] = v
    return v


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    bs = config['batch_size']
    nh = config['num_heads']
    vd = config['v_head_dim']
    kvl = config['kv_seq_len']
    sms = config['sm_scale']
    dev = q.device
    tq = q.shape[0]

    ok = (tq, nh, vd)
    o = _oc.get(ok)
    if o is None:
        o = torch.empty((tq, nh, vd), dtype=q.dtype, device=dev)
        _oc[ok] = o

    kd, ks = kv_data["fp8"]
    tot = kd.shape[0]

    # ---- BF16 persistent for bs<=4/kv<=1024 only (bs=64 FAILS leaderboard) ----
    if kvl <= 1024 and bs <= 4:
        kv_bf16 = kv_data["bf16"]
        ps, ns = 2, 16
        np = tot // ps
        t = _gt(('bfp', bs, tot, ps), lambda: {
            'ki': torch.arange(np, device=dev, dtype=torch.int32),
            'kip': kv_indptr // ps,
            'kl': torch.full((bs,), ps, device=dev, dtype=torch.int32),
        })
        m = _gm(bs, tot, nh, ns, ps, qo_indptr, t['kip'], t['kl'], dev, torch.bfloat16, torch.bfloat16)
        mla_decode_fwd(q, kv_bf16.view(np, ps, 1, 576), o,
            qo_indptr, t['kip'], t['ki'], t['kl'],
            1, page_size=ps, nhead_kv=1, sm_scale=sms,
            logit_cap=0.0, num_kv_splits=ns,
            intra_batch_mode=False, **m)
        return o

    # ---- FP8 NP splits=1 for bs=64/kv=1024 (BF16p fails leaderboard) ----
    if kvl <= 1024 and bs == 64:
        q8 = q.to(torch.float8_e4m3fn)
        t = _gt(('fn', bs, tot), lambda: {
            'qs': torch.ones(1, dtype=torch.float32, device=dev),
            'ki': torch.arange(tot, device=dev, dtype=torch.int32),
            'kl': torch.ones(bs, device=dev, dtype=torch.int32),
            'si': torch.arange(bs+1, dtype=torch.int32, device=dev),
        })
        mla_decode_fwd(q=q8, kv_buffer=kd.unsqueeze(1), o=o,
            qo_indptr=qo_indptr, kv_indptr=kv_indptr,
            kv_indices=t['ki'], kv_last_page_lens=t['kl'],
            max_seqlen_q=1, page_size=1, nhead_kv=1, sm_scale=sms,
            num_kv_splits=1, num_kv_splits_indptr=t['si'],
            q_scale=t['qs'], kv_scale=ks)
        return o

    # ---- a8w8 persistent (direct calls) for all other shapes ----
    q8 = q.to(torch.float8_e4m3fn)
    if kvl <= 1024:
        ps, ns = 2, (16 if bs <= 32 else 8)
    else:
        ps, ns = 8, (16 if bs <= 4 else 8)

    np = tot // ps
    t = _gt(('a8', bs, tot, ps), lambda: {
        'qs': torch.ones(1, dtype=torch.float32, device=dev),
        'ki': torch.arange(np, device=dev, dtype=torch.int32),
        'kip': kv_indptr // ps,
        'kl': torch.full((bs,), ps, device=dev, dtype=torch.int32),
    })
    m = _gm(bs, tot, nh, ns, ps, qo_indptr, t['kip'], t['kl'], dev, aiter_dtypes.fp8, aiter_dtypes.fp8)

    rpm_sz = m['reduce_partial_map'].size(0)
    inter = _gi(('a8_i', rpm_sz, nh, vd), lambda: {
        'logits': torch.empty((rpm_sz, 1, nh, vd), dtype=torch.float32, device=dev),
        'attn_lse': torch.empty((rpm_sz, 1, nh, 1), dtype=torch.float32, device=dev),
    })

    _aiter.mla_decode_stage1_asm_fwd(
        q8, kd.view(np, ps, 1, 576), qo_indptr, t['kip'],
        t['ki'], t['kl'],
        None,
        m['work_meta_data'], m['work_indptr'], m['work_info_set'],
        1, ps, 1, sms,
        inter['logits'], inter['attn_lse'], o,
        t['qs'], ks,
    )

    _aiter.mla_reduce_v1(
        inter['logits'], inter['attn_lse'],
        m['reduce_indptr'], m['reduce_final_map'], m['reduce_partial_map'],
        1, o, None,
    )
    return o
scrolls · 137 lines total

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

Changes from previous submission

Against this author's previous submission submission 651367.

#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
- """MLA v108 — Full bypass: ALL paths use direct stage1_asm (+ reduce where needed).
+ """MLA v90 — v89 with bs=64/kv=1024 reverted to FP8 NP (BF16p failed leaderboard).
- v90 bypassed mla_decode_fwd only for a8w8 persistent. v108 also bypasses for:
- - BF16 persistent (bs<=4/kv<=1024): saves ~3µs from torch.empty elimination
- - FP8 NP splits=1 (bs=64/kv<=1024): saves ~2µs, MAYBE_FINAL_OUT=True (no stage2)
+ Changes from v86:
+ - bs=4/kv=1024: BF16 NP → BF16 persistent (-1.2µs) ← KEPT
+ - bs=64/kv=1024: stays FP8 NP splits=1 (BF16p fails secret runner) ← REVERTED
"""
from task import input_t, output_t
import torch
import aiter as _aiter
+ 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
⋯ 45 unchanged lines
kd, ks = kv_data["fp8"]
tot = kd.shape[0]
- # ---- BF16 persistent DIRECT (bs<=4/kv<=1024) ----
+ # ---- BF16 persistent for bs<=4/kv<=1024 only (bs=64 FAILS leaderboard) ----
if kvl <= 1024 and bs <= 4:
kv_bf16 = kv_data["bf16"]
ps, ns = 2, 16
⋯ 4 unchanged lines
'kl': torch.full((bs,), ps, device=dev, dtype=torch.int32),
})
m = _gm(bs, tot, nh, ns, ps, qo_indptr, t['kip'], t['kl'], dev, torch.bfloat16, torch.bfloat16)
- rpm_sz = m['reduce_partial_map'].size(0)
- inter = _gi(('bfp_i', rpm_sz, nh, vd), lambda: {
- 'logits': torch.empty((rpm_sz, 1, nh, vd), dtype=torch.float32, device=dev),
- 'attn_lse': torch.empty((rpm_sz, 1, nh, 1), dtype=torch.float32, device=dev),
- })
- _aiter.mla_decode_stage1_asm_fwd(
- q, kv_bf16.view(np, ps, 1, 576), qo_indptr, t['kip'],
- t['ki'], t['kl'], None,
- m['work_meta_data'], m['work_indptr'], m['work_info_set'],
- 1, ps, 1, sms,
- inter['logits'], inter['attn_lse'], o,
- )
- _aiter.mla_reduce_v1(
- inter['logits'], inter['attn_lse'],
- m['reduce_indptr'], m['reduce_final_map'], m['reduce_partial_map'],
- 1, o, None,
- )
+ mla_decode_fwd(q, kv_bf16.view(np, ps, 1, 576), o,
+ qo_indptr, t['kip'], t['ki'], t['kl'],
+ 1, page_size=ps, nhead_kv=1, sm_scale=sms,
+ logit_cap=0.0, num_kv_splits=ns,
+ intra_batch_mode=False, **m)
return o
- # ---- FP8 NP splits=1 DIRECT (bs=64/kv<=1024) ----
- # MAYBE_FINAL_OUT=True (v_dim=512<=512, mgc=0): stage1 writes directly to o
+ # ---- FP8 NP splits=1 for bs=64/kv=1024 (BF16p fails leaderboard) ----
if kvl <= 1024 and bs == 64:
q8 = q.to(torch.float8_e4m3fn)
t = _gt(('fn', bs, tot), lambda: {
⋯ 2 unchanged lines
'kl': torch.ones(bs, device=dev, dtype=torch.int32),
'si': torch.arange(bs+1, dtype=torch.int32, device=dev),
})
- inter = _gi(('fn_i', tq, nh), lambda: {
- 'attn_lse': torch.empty((tq, 1, nh, 1), dtype=torch.float32, device=dev),
- })
- logits = o.view(tq, 1, nh, vd) # View of output — stage1 writes here directly
- _aiter.mla_decode_stage1_asm_fwd(
- q8, kd.unsqueeze(1), qo_indptr, kv_indptr,
- t['ki'], t['kl'], t['si'],
- None, None, None,
- 1, 1, 1, sms,
- logits, inter['attn_lse'], o,
- t['qs'], ks,
- )
- # NO stage2 reduce needed — MAYBE_FINAL_OUT=True
+ mla_decode_fwd(q=q8, kv_buffer=kd.unsqueeze(1), o=o,
+ qo_indptr=qo_indptr, kv_indptr=kv_indptr,
+ kv_indices=t['ki'], kv_last_page_lens=t['kl'],
+ max_seqlen_q=1, page_size=1, nhead_kv=1, sm_scale=sms,
+ num_kv_splits=1, num_kv_splits_indptr=t['si'],
+ q_scale=t['qs'], kv_scale=ks)
return o
- # ---- a8w8 persistent DIRECT (all other shapes) ----
+ # ---- a8w8 persistent (direct calls) for all other shapes ----
q8 = q.to(torch.float8_e4m3fn)
if kvl <= 1024:
ps, ns = 2, (16 if bs <= 32 else 8)
⋯ 17 unchanged lines
_aiter.mla_decode_stage1_asm_fwd(
q8, kd.view(np, ps, 1, 576), qo_indptr, t['kip'],
- t['ki'], t['kl'], None,
+ t['ki'], t['kl'],
+ None,
m['work_meta_data'], m['work_indptr'], m['work_info_set'],
1, ps, 1, sms,
inter['logits'], inter['attn_lse'], o,
scrolls · 102 diff lines total

Best evidence level for this revision: reported

JSON