Skip to content
KernelIndex
Search⌘K

submission 739598

johnny.t.shi · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

v61.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-739598?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
32.6µs
#31 of 766
2026-04-05

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:2c944642d6ae0e570be29d1eec9b04d5deeea38fc802e259ebd774cb925126a6
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"""a8w8 persistent: FP8 Q + FP8 KV — uses higher-throughput FP8×FP8 MFMA."""

Kernel source

v61.py148 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""MLA v61 — v60 + split tuning + a8w8 for bs=64/kv=1024.

Changes from v60:
1. bs=64/kv=1024: FP8 NP splits=1 → a8w8 ps=2 splits=16
2. bs=64/kv=8192: splits 4→8
3. bs=256/kv=8192: splits 4→8
"""
from task import input_t, output_t
import torch
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

_meta_cache = {}
_tensor_cache = {}


def _get_cached_tensors(key, create_fn):
    cached = _tensor_cache.get(key)
    if cached is None:
        cached = create_fn()
        _tensor_cache[key] = cached
    return cached


def _get_or_make_metadata(batch_size, total_kv, num_heads, nhead_kv, num_splits, page_size,
                          q_dtype, kv_dtype, qo_indptr, kv_indptr, kv_last_page_lens, device):
    key = (batch_size, total_kv, num_heads, num_splits, page_size, str(q_dtype), str(kv_dtype))
    cached = _meta_cache.get(key)
    if cached is not None:
        return cached

    info = get_mla_metadata_info_v1(
        batch_size, 1, num_heads, q_dtype, kv_dtype,
        is_sparse=False, fast_mode=True,
        num_kv_splits=num_splits, intra_batch_mode=False,
    )
    work = [torch.empty(s, dtype=t, device=device) for s, t in info]
    (wmd, wi, wis, ri, rfm, rpm) = work

    get_mla_metadata_v1(
        qo_indptr, kv_indptr, kv_last_page_lens,
        num_heads // nhead_kv, nhead_kv, False,
        wmd, wis, wi, ri, rfm, rpm,
        page_size=page_size, kv_granularity=max(page_size, 16),
        max_seqlen_qo=1, uni_seqlen_qo=1, fast_mode=True,
        max_split_per_batch=num_splits, intra_batch_mode=False,
        dtype_q=q_dtype, dtype_kv=kv_dtype,
    )

    result = dict(
        work_meta_data=wmd, work_indptr=wi, work_info_set=wis,
        reduce_indptr=ri, reduce_final_map=rfm, reduce_partial_map=rpm,
    )
    _meta_cache[key] = result
    return result


def _run_bf16(q, kv_bf16, output, qo_indptr, kv_indptr, config):
    batch_size = config['batch_size']
    total_kv = kv_bf16.shape[0]
    kv_buffer = kv_bf16.unsqueeze(1)

    tensors = _get_cached_tensors(
        ('bf16', batch_size, total_kv),
        lambda: {
            'kv_indices': torch.arange(total_kv, device=q.device, dtype=torch.int32),
            'kv_last_page_lens': torch.ones(batch_size, device=q.device, dtype=torch.int32),
        }
    )

    mla_decode_fwd(
        q=q, kv_buffer=kv_buffer, o=output,
        qo_indptr=qo_indptr, kv_indptr=kv_indptr,
        kv_indices=tensors['kv_indices'], kv_last_page_lens=tensors['kv_last_page_lens'],
        max_seqlen_q=1, page_size=1, nhead_kv=1, sm_scale=config['sm_scale'],
    )


def _run_a8w8(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config, page_size, num_splits=4):
    """a8w8 persistent: FP8 Q + FP8 KV — uses higher-throughput FP8×FP8 MFMA."""
    batch_size = config['batch_size']
    num_heads = config['num_heads']
    total_kv = kv_fp8_data.shape[0]

    q_fp8 = q.to(torch.float8_e4m3fn)
    num_pages = total_kv // page_size
    kv_buffer = kv_fp8_data.view(num_pages, page_size, 1, 576)

    tensors = _get_cached_tensors(
        ('a8w8', batch_size, total_kv, page_size),
        lambda: {
            'q_scale': torch.ones(1, dtype=torch.float32, device=q.device),
            'kv_indices': torch.arange(num_pages, device=q.device, dtype=torch.int32),
            'kv_indptr_pages': kv_indptr // page_size,
            'kv_last_page_lens': torch.full((batch_size,), page_size, device=q.device, dtype=torch.int32),
        }
    )

    meta = _get_or_make_metadata(
        batch_size, total_kv, num_heads, 1, num_splits, page_size,
        aiter_dtypes.fp8, aiter_dtypes.fp8,
        qo_indptr, tensors['kv_indptr_pages'], tensors['kv_last_page_lens'], q.device,
    )

    mla_decode_fwd(
        q_fp8, kv_buffer, output,
        qo_indptr, tensors['kv_indptr_pages'], tensors['kv_indices'], tensors['kv_last_page_lens'],
        1, page_size=page_size, nhead_kv=1, sm_scale=config['sm_scale'],
        logit_cap=0.0, num_kv_splits=num_splits,
        q_scale=tensors['q_scale'], kv_scale=kv_fp8_scale,
        intra_batch_mode=False, **meta,
    )


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data

    batch_size = config['batch_size']
    num_heads = config['num_heads']
    v_head_dim = config['v_head_dim']
    kv_seq_len = config['kv_seq_len']

    output = torch.empty((q.shape[0], num_heads, v_head_dim), dtype=q.dtype, device=q.device)

    if kv_seq_len <= 1024 and batch_size <= 4:
        # BF16 non-persistent — fastest for small batch, exact
        _run_bf16(q, kv_data["bf16"], output, qo_indptr, kv_indptr, config)
    elif kv_seq_len <= 1024:
        # a8w8 ps=2 for all kv=1024 shapes (including bs=64)
        kv_fp8_data, kv_fp8_scale = kv_data["fp8"]
        splits = 16 if batch_size <= 64 else 8
        _run_a8w8(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config, page_size=2, num_splits=splits)
    else:
        # a8w8 ps=8 for kv=8192 — increased splits for larger batches
        kv_fp8_data, kv_fp8_scale = kv_data["fp8"]
        if batch_size <= 4:
            splits = 16
        elif batch_size <= 64:
            splits = 8
        else:
            splits = 8  # was 4, now 8 for bs=256
        _run_a8w8(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config, page_size=8, num_splits=splits)

    return output
scrolls · 148 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 677054.

#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
- """MLA v90 — v89 with bs=64/kv=1024 reverted to FP8 NP (BF16p failed leaderboard).
+ """MLA v61 — v60 + split tuning + a8w8 for bs=64/kv=1024.
- 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
+ Changes from v60:
+ 1. bs=64/kv=1024: FP8 NP splits=1 → a8w8 ps=2 splits=16
+ 2. bs=64/kv=8192: splits 4→8
+ 3. bs=256/kv=8192: splits 4→8
"""
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 = {}
+ _meta_cache = {}
+ _tensor_cache = {}
- 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 _get_cached_tensors(key, create_fn):
+ cached = _tensor_cache.get(key)
+ if cached is None:
+ cached = create_fn()
+ _tensor_cache[key] = cached
+ return cached
- def _gi(key, fn):
- v = _ic.get(key)
- if v is None: v = fn(); _ic[key] = v
- return v
+ def _get_or_make_metadata(batch_size, total_kv, num_heads, nhead_kv, num_splits, page_size,
+ q_dtype, kv_dtype, qo_indptr, kv_indptr, kv_last_page_lens, device):
+ key = (batch_size, total_kv, num_heads, num_splits, page_size, str(q_dtype), str(kv_dtype))
+ cached = _meta_cache.get(key)
+ if cached is not None:
+ return cached
- 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]
+ info = get_mla_metadata_info_v1(
+ batch_size, 1, num_heads, q_dtype, kv_dtype,
+ is_sparse=False, fast_mode=True,
+ num_kv_splits=num_splits, intra_batch_mode=False,
+ )
+ work = [torch.empty(s, dtype=t, device=device) for s, t in info]
+ (wmd, wi, wis, ri, rfm, rpm) = work
- 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
+ get_mla_metadata_v1(
+ qo_indptr, kv_indptr, kv_last_page_lens,
+ num_heads // nhead_kv, nhead_kv, False,
+ wmd, wis, wi, ri, rfm, rpm,
+ page_size=page_size, kv_granularity=max(page_size, 16),
+ max_seqlen_qo=1, uni_seqlen_qo=1, fast_mode=True,
+ max_split_per_batch=num_splits, intra_batch_mode=False,
+ dtype_q=q_dtype, dtype_kv=kv_dtype,
+ )
- kd, ks = kv_data["fp8"]
- tot = kd.shape[0]
+ result = dict(
+ work_meta_data=wmd, work_indptr=wi, work_info_set=wis,
+ reduce_indptr=ri, reduce_final_map=rfm, reduce_partial_map=rpm,
+ )
+ _meta_cache[key] = result
+ return result
- # ---- 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
+ def _run_bf16(q, kv_bf16, output, qo_indptr, kv_indptr, config):
+ batch_size = config['batch_size']
+ total_kv = kv_bf16.shape[0]
+ kv_buffer = kv_bf16.unsqueeze(1)
- # ---- 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)
+ tensors = _get_cached_tensors(
+ ('bf16', batch_size, total_kv),
+ lambda: {
+ 'kv_indices': torch.arange(total_kv, device=q.device, dtype=torch.int32),
+ 'kv_last_page_lens': torch.ones(batch_size, device=q.device, dtype=torch.int32),
+ }
+ )
- 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)
+ mla_decode_fwd(
+ q=q, kv_buffer=kv_buffer, o=output,
+ qo_indptr=qo_indptr, kv_indptr=kv_indptr,
+ kv_indices=tensors['kv_indices'], kv_last_page_lens=tensors['kv_last_page_lens'],
+ max_seqlen_q=1, page_size=1, nhead_kv=1, sm_scale=config['sm_scale'],
+ )
- 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,
+ def _run_a8w8(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config, page_size, num_splits=4):
+ """a8w8 persistent: FP8 Q + FP8 KV — uses higher-throughput FP8×FP8 MFMA."""
+ batch_size = config['batch_size']
+ num_heads = config['num_heads']
+ total_kv = kv_fp8_data.shape[0]
+
+ q_fp8 = q.to(torch.float8_e4m3fn)
+ num_pages = total_kv // page_size
+ kv_buffer = kv_fp8_data.view(num_pages, page_size, 1, 576)
+
+ tensors = _get_cached_tensors(
+ ('a8w8', batch_size, total_kv, page_size),
+ lambda: {
+ 'q_scale': torch.ones(1, dtype=torch.float32, device=q.device),
+ 'kv_indices': torch.arange(num_pages, device=q.device, dtype=torch.int32),
+ 'kv_indptr_pages': kv_indptr // page_size,
+ 'kv_last_page_lens': torch.full((batch_size,), page_size, device=q.device, dtype=torch.int32),
+ }
)
- _aiter.mla_reduce_v1(
- inter['logits'], inter['attn_lse'],
- m['reduce_indptr'], m['reduce_final_map'], m['reduce_partial_map'],
- 1, o, None,
+ meta = _get_or_make_metadata(
+ batch_size, total_kv, num_heads, 1, num_splits, page_size,
+ aiter_dtypes.fp8, aiter_dtypes.fp8,
+ qo_indptr, tensors['kv_indptr_pages'], tensors['kv_last_page_lens'], q.device,
)
- return o
+
+ mla_decode_fwd(
+ q_fp8, kv_buffer, output,
+ qo_indptr, tensors['kv_indptr_pages'], tensors['kv_indices'], tensors['kv_last_page_lens'],
+ 1, page_size=page_size, nhead_kv=1, sm_scale=config['sm_scale'],
+ logit_cap=0.0, num_kv_splits=num_splits,
+ q_scale=tensors['q_scale'], kv_scale=kv_fp8_scale,
+ intra_batch_mode=False, **meta,
+ )
+
+
+ def custom_kernel(data: input_t) -> output_t:
+ q, kv_data, qo_indptr, kv_indptr, config = data
+
+ batch_size = config['batch_size']
+ num_heads = config['num_heads']
+ v_head_dim = config['v_head_dim']
+ kv_seq_len = config['kv_seq_len']
+
+ output = torch.empty((q.shape[0], num_heads, v_head_dim), dtype=q.dtype, device=q.device)
+
+ if kv_seq_len <= 1024 and batch_size <= 4:
+ # BF16 non-persistent — fastest for small batch, exact
+ _run_bf16(q, kv_data["bf16"], output, qo_indptr, kv_indptr, config)
+ elif kv_seq_len <= 1024:
+ # a8w8 ps=2 for all kv=1024 shapes (including bs=64)
+ kv_fp8_data, kv_fp8_scale = kv_data["fp8"]
+ splits = 16 if batch_size <= 64 else 8
+ _run_a8w8(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config, page_size=2, num_splits=splits)
+ else:
+ # a8w8 ps=8 for kv=8192 — increased splits for larger batches
+ kv_fp8_data, kv_fp8_scale = kv_data["fp8"]
+ if batch_size <= 4:
+ splits = 16
+ elif batch_size <= 64:
+ splits = 8
+ else:
+ splits = 8 # was 4, now 8 for bs=256
+ _run_a8w8(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config, page_size=8, num_splits=splits)
+
+ return output
scrolls · 257 diff lines total

Best evidence level for this revision: reported

JSON