Skip to content
KernelIndex
Search⌘K

submission 616673

johnny.t.shi · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

v31.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-616673?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
53.1µs
#175 of 766
2026-03-23

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:e3bb08f19bc92bf53d432029f80aaf804b529aa3446aef70c9e431466c88692f
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"""MLA v29 — Optimal hybrid: BF16 non-persist + a16w8 (BF16 Q + FP8 KV) persistent.

Kernel source

v31.py125 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""MLA v29 — Optimal hybrid: BF16 non-persist + a16w8 (BF16 Q + FP8 KV) persistent.

Dispatch by best-of per shape:
  bs≤4, kv≤1024  → BF16 non-persistent (26µs, fastest for tiny shapes)
  everything else → a16w8 persistent (45-98µs, FP8 KV bandwidth + no Q quant overhead)
    ps=1 for kv≤1024, ps=8 for kv≥8192
"""
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 = {}


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):
    """BF16 non-persistent — fastest for tiny shapes."""
    batch_size = config['batch_size']
    total_kv = kv_bf16.shape[0]
    kv_buffer = kv_bf16.unsqueeze(1)
    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=kv_indices, kv_last_page_lens=kv_last_page_lens,
        max_seqlen_q=1, page_size=1, nhead_kv=1, sm_scale=config['sm_scale'],
    )


def _run_a16w8(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config, page_size, num_splits=4):
    """a16w8: BF16 Q + FP8 KV persistent — no Q quant, FP8 bandwidth."""
    batch_size = config['batch_size']
    num_heads = config['num_heads']
    total_kv = kv_fp8_data.shape[0]
    NUM_KV_SPLITS = num_splits

    num_pages = total_kv // page_size
    kv_buffer = kv_fp8_data.view(num_pages, page_size, 1, 576)
    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_KV_SPLITS, page_size,
        torch.bfloat16, aiter_dtypes.fp8,
        qo_indptr, kv_indptr_pages, kv_last_page_lens, q.device,
    )

    mla_decode_fwd(
        q,              # BF16 Q — no quantization!
        kv_buffer,      # FP8 KV
        output,
        qo_indptr, kv_indptr_pages, kv_indices, 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_KV_SPLITS,
        q_scale=None,   # BF16 Q doesn't need 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 batch_size <= 4 and kv_seq_len <= 1024:
        # BF16 non-persistent: fastest for tiny shapes (26µs)
        _run_bf16(q, kv_data["bf16"], output, qo_indptr, kv_indptr, config)
    elif kv_seq_len <= 1024:
        # a16w8 persistent ps=1: adaptive splits (more for small batch, fewer for large)
        kv_fp8_data, kv_fp8_scale = kv_data["fp8"]
        splits = 8 if batch_size <= 32 else 4
        _run_a16w8(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config, page_size=1, num_splits=splits)
    else:
        # a16w8 persistent ps=8: adaptive splits
        kv_fp8_data, kv_fp8_scale = kv_data["fp8"]
        splits = 8 if batch_size <= 32 else 4
        _run_a16w8(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config, page_size=8, num_splits=splits)

    return output
scrolls · 125 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 600564.

#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
- """MLA Decode v12 — FP8 Q + FP8 KV with persistent mode.
+ """MLA v29 — Optimal hybrid: BF16 non-persist + a16w8 (BF16 Q + FP8 KV) persistent.
- Key insight from recon: competitors at ranks 6,9 use 'a16w8_ps2':
- - Cast Q from BF16 → FP8 (a16 = original precision, w8 = KV in FP8)
- - Use persistent mode (ps=1) with pre-computed metadata
- - aiter only supports fp8+fp8 in persistent mode
-
- Steps:
- 1. Cast Q to float8_e4m3fn
- 2. Pre-compute metadata via get_mla_metadata_info_v1 + get_mla_metadata_v1
- 3. Call mla_decode_fwd in persistent mode
+ Dispatch by best-of per shape:
+ bs≤4, kv≤1024 → BF16 non-persistent (26µs, fastest for tiny shapes)
+ everything else → a16w8 persistent (45-98µs, FP8 KV bandwidth + no Q quant overhead)
+ ps=1 for kv≤1024, ps=8 for kv≥8192
"""
from task import input_t, output_t
import torch
- import aiter
- from aiter import dtypes
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 = {}
- def _get_or_compute_metadata(batch_size, total_kv, num_heads, nhead_kv, page_size,
- qo_indptr, kv_indptr, kv_last_page_lens, device):
- """Cache metadata by shape key. Safe: metadata depends on shape, not data content."""
- key = (batch_size, total_kv, num_heads, page_size)
+
+ 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
- (
- (wmd_sz, wmd_ty), (wi_sz, wi_ty), (wis_sz, wis_ty),
- (ri_sz, ri_ty), (rfm_sz, rfm_ty), (rpm_sz, rpm_ty),
- ) = aiter.get_mla_metadata_info_v1(
- batch_size, 1, num_heads, dtypes.fp8, dtypes.fp8,
- is_sparse=False, fast_mode=True, num_kv_splits=-1, intra_batch_mode=False,
+ 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
- wmd = torch.empty(wmd_sz, dtype=wmd_ty, device=device)
- wi = torch.empty(wi_sz, dtype=wi_ty, device=device)
- wis = torch.empty(wis_sz, dtype=wis_ty, device=device)
- ri = torch.empty(ri_sz, dtype=ri_ty, device=device)
- rfm = torch.empty(rfm_sz, dtype=rfm_ty, device=device)
- rpm = torch.empty(rpm_sz, dtype=rpm_ty, device=device)
-
- aiter.get_mla_metadata_v1(
+ 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=-1, intra_batch_mode=False,
- dtype_q=dtypes.fp8, dtype_kv=dtypes.fp8,
+ max_split_per_batch=num_splits, intra_batch_mode=False,
+ dtype_q=q_dtype, dtype_kv=kv_dtype,
)
- result = (wmd, wi, wis, ri, rfm, rpm)
+ 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 custom_kernel(data: input_t) -> output_t:
- q, kv_data, qo_indptr, kv_indptr, config = data
- kv_fp8_data, kv_fp8_scale = kv_data["fp8"]
+ def _run_bf16(q, kv_bf16, output, qo_indptr, kv_indptr, config):
+ """BF16 non-persistent — fastest for tiny shapes."""
+ batch_size = config['batch_size']
+ total_kv = kv_bf16.shape[0]
+ kv_buffer = kv_bf16.unsqueeze(1)
+ 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=kv_indices, kv_last_page_lens=kv_last_page_lens,
+ max_seqlen_q=1, page_size=1, nhead_kv=1, sm_scale=config['sm_scale'],
+ )
+
+
+ def _run_a16w8(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config, page_size, num_splits=4):
+ """a16w8: BF16 Q + FP8 KV persistent — no Q quant, FP8 bandwidth."""
batch_size = config['batch_size']
num_heads = config['num_heads']
- v_head_dim = config['v_head_dim']
- sm_scale = config['sm_scale']
-
total_kv = kv_fp8_data.shape[0]
+ NUM_KV_SPLITS = num_splits
- # Cast Q to FP8
- q_fp8 = q.to(torch.float8_e4m3fn)
- q_scale = torch.ones([1], dtype=torch.float32, device=q.device)
-
- output = torch.empty((q.shape[0], num_heads, v_head_dim), dtype=q.dtype, device=q.device)
-
- # page_size=2
- PAGE_SIZE = 2
- num_pages = total_kv // PAGE_SIZE
- kv_buffer = kv_fp8_data.view(num_pages, PAGE_SIZE, 1, 576)
-
+ num_pages = total_kv // page_size
+ kv_buffer = kv_fp8_data.view(num_pages, page_size, 1, 576)
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)
+ kv_indptr_pages = kv_indptr // page_size
+ kv_last_page_lens = torch.full((batch_size,), page_size, device=q.device, dtype=torch.int32)
- # Cached metadata: depends on shape only, not data content. Safe for leaderboard.
- (work_meta_data, work_indptr, work_info_set,
- reduce_indptr, reduce_final_map, reduce_partial_map) = _get_or_compute_metadata(
- batch_size, total_kv, num_heads, 1, PAGE_SIZE,
+ meta = _get_or_make_metadata(
+ batch_size, total_kv, num_heads, 1, NUM_KV_SPLITS, page_size,
+ torch.bfloat16, aiter_dtypes.fp8,
qo_indptr, kv_indptr_pages, kv_last_page_lens, q.device,
)
mla_decode_fwd(
- q=q_fp8,
- kv_buffer=kv_buffer,
- o=output,
- qo_indptr=qo_indptr,
- kv_indptr=kv_indptr_pages,
- kv_indices=kv_indices,
- kv_last_page_lens=kv_last_page_lens,
- max_seqlen_q=1,
- page_size=PAGE_SIZE,
- nhead_kv=1,
- sm_scale=sm_scale,
- work_meta_data=work_meta_data,
- 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,
- q_scale=q_scale,
+ q, # BF16 Q — no quantization!
+ kv_buffer, # FP8 KV
+ output,
+ qo_indptr, kv_indptr_pages, kv_indices, 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_KV_SPLITS,
+ q_scale=None, # BF16 Q doesn't need 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 batch_size <= 4 and kv_seq_len <= 1024:
+ # BF16 non-persistent: fastest for tiny shapes (26µs)
+ _run_bf16(q, kv_data["bf16"], output, qo_indptr, kv_indptr, config)
+ elif kv_seq_len <= 1024:
+ # a16w8 persistent ps=1: adaptive splits (more for small batch, fewer for large)
+ kv_fp8_data, kv_fp8_scale = kv_data["fp8"]
+ splits = 8 if batch_size <= 32 else 4
+ _run_a16w8(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config, page_size=1, num_splits=splits)
+ else:
+ # a16w8 persistent ps=8: adaptive splits
+ kv_fp8_data, kv_fp8_scale = kv_data["fp8"]
+ splits = 8 if batch_size <= 32 else 4
+ _run_a16w8(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config, page_size=8, num_splits=splits)
+
return output
scrolls · 198 diff lines total

Best evidence level for this revision: reported

JSON