Skip to content
KernelIndex
Search⌘K

submission 617326

johnny.t.shi · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:9bb44675ed25c64d1d9c25612a5005bd585838ad1926cad32dab483ca8666d83
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 v41 — Optimal safe hybrid: BF16 + FP8 NP + a16w8 persistent.

Kernel source

v41.py146 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""MLA v41 — Optimal safe hybrid: BF16 + FP8 NP + a16w8 persistent.

Per-shape optimized dispatch based on full benchmark sweep (v32/v33/v35/v39):
  bs≤32, kv≤1024  → BF16 non-persist (26-37µs, exact, faster than FP8 NP for small batch)
  bs≥64, kv≤1024  → FP8 non-persist splits=1 (40-78µs, no reduce step)
  kv≥8192         → a16w8 persist ps=8 (37-100µs, persistent wins by 3-4x for large kv)
Target geomean: ~47µs
"""
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 — exact correctness."""
    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_fp8_nonpersist(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config):
    """FP8+FP8 non-persistent splits=1 — single kernel, no reduce."""
    batch_size = config['batch_size']
    total_kv = kv_fp8_data.shape[0]

    q_fp8 = q.to(torch.float8_e4m3fn)
    q_scale = torch.ones(1, dtype=torch.float32, device=q.device)

    kv_buffer = kv_fp8_data.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)

    num_kv_splits = 1
    num_kv_splits_indptr = torch.arange(batch_size + 1, dtype=torch.int32, device=q.device)

    mla_decode_fwd(
        q=q_fp8, 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'],
        num_kv_splits=num_kv_splits, num_kv_splits_indptr=num_kv_splits_indptr,
        q_scale=q_scale, kv_scale=kv_fp8_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."""
    batch_size = config['batch_size']
    num_heads = config['num_heads']
    total_kv = kv_fp8_data.shape[0]

    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_splits, page_size,
        torch.bfloat16, aiter_dtypes.fp8,
        qo_indptr, kv_indptr_pages, kv_last_page_lens, q.device,
    )

    mla_decode_fwd(
        q, kv_buffer, 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_splits,
        q_scale=None, 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 <= 32:
        # BF16 non-persistent: exact, fastest for small batch + short kv
        _run_bf16(q, kv_data["bf16"], output, qo_indptr, kv_indptr, config)
    elif kv_seq_len <= 1024:
        # FP8 non-persistent splits=1: faster than BF16/a16w8 for bs≥64/kv≤1024
        kv_fp8_data, kv_fp8_scale = kv_data["fp8"]
        _run_fp8_nonpersist(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config)
    else:
        # a16w8 persistent ps=8 for kv=8192: persistent wins by 3-4x for large kv
        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 · 146 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 616673.

#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
- """MLA v29 — Optimal hybrid: BF16 non-persist + a16w8 (BF16 Q + FP8 KV) persistent.
+ """MLA v41 — Optimal safe hybrid: BF16 + FP8 NP + a16w8 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
+ Per-shape optimized dispatch based on full benchmark sweep (v32/v33/v35/v39):
+ bs≤32, kv≤1024 → BF16 non-persist (26-37µs, exact, faster than FP8 NP for small batch)
+ bs≥64, kv≤1024 → FP8 non-persist splits=1 (40-78µs, no reduce step)
+ kv≥8192 → a16w8 persist ps=8 (37-100µs, persistent wins by 3-4x for large kv)
+ Target geomean: ~47µs
"""
from task import input_t, output_t
import torch
⋯ 38 unchanged lines
def _run_bf16(q, kv_bf16, output, qo_indptr, kv_indptr, config):
- """BF16 non-persistent — fastest for tiny shapes."""
+ """BF16 non-persistent — exact correctness."""
batch_size = config['batch_size']
total_kv = kv_bf16.shape[0]
kv_buffer = kv_bf16.unsqueeze(1)
⋯ 8 unchanged lines
)
+ def _run_fp8_nonpersist(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config):
+ """FP8+FP8 non-persistent splits=1 — single kernel, no reduce."""
+ batch_size = config['batch_size']
+ total_kv = kv_fp8_data.shape[0]
+
+ q_fp8 = q.to(torch.float8_e4m3fn)
+ q_scale = torch.ones(1, dtype=torch.float32, device=q.device)
+
+ kv_buffer = kv_fp8_data.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)
+
+ num_kv_splits = 1
+ num_kv_splits_indptr = torch.arange(batch_size + 1, dtype=torch.int32, device=q.device)
+
+ mla_decode_fwd(
+ q=q_fp8, 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'],
+ num_kv_splits=num_kv_splits, num_kv_splits_indptr=num_kv_splits_indptr,
+ q_scale=q_scale, kv_scale=kv_fp8_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."""
+ """a16w8: BF16 Q + FP8 KV persistent."""
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)
⋯ 2 unchanged lines
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,
+ batch_size, total_kv, num_heads, 1, num_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,
+ q, kv_buffer, 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,
+ logit_cap=0.0, num_kv_splits=num_splits,
+ q_scale=None, kv_scale=kv_fp8_scale,
intra_batch_mode=False, **meta,
)
⋯ 8 unchanged lines
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)
+ if kv_seq_len <= 1024 and batch_size <= 32:
+ # BF16 non-persistent: exact, fastest for small batch + short kv
_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)
+ # FP8 non-persistent splits=1: faster than BF16/a16w8 for bs≥64/kv≤1024
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)
+ _run_fp8_nonpersist(q, kv_fp8_data, kv_fp8_scale, output, qo_indptr, kv_indptr, config)
else:
- # a16w8 persistent ps=8: adaptive splits
+ # a16w8 persistent ps=8 for kv=8192: persistent wins by 3-4x for large kv
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)
scrolls · 112 diff lines total

Best evidence level for this revision: reported

JSON