Skip to content
KernelIndex
Search⌘K

submission 600564

johnny.t.shi · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

v12_fp8_q_kv.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-600564?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
55.0µs
#185 of 766
2026-03-21

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:bebc30e0b45c9f207144aa8fcaa809c217a4c8aacc075910ffeb5fe06001902f
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 Decode v12 — FP8 Q + FP8 KV with persistent mode.

Kernel source

v12_fp8_q_kv.py117 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""MLA Decode v12 — FP8 Q + FP8 KV with persistent mode.

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
"""
from task import input_t, output_t
import torch
import aiter
from aiter import dtypes
from aiter.mla import mla_decode_fwd

_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)
    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,
    )

    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(
        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,
    )

    result = (wmd, wi, wis, ri, rfm, 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"]

    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]

    # 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)

    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)

    # 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,
        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,
        kv_scale=kv_fp8_scale,
    )

    return output
scrolls · 117 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 598454.

⋯ 19 unchanged lines
_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):
- """Compute and cache persistent mode metadata."""
+ """Cache metadata by shape key. Safe: metadata depends on shape, not data content."""
key = (batch_size, total_kv, num_heads, page_size)
cached = _meta_cache.get(key)
if cached is not None:
return cached
- max_seqlen_qo = 1 # decode
- max_split_per_batch = -1 # auto
-
- # Get metadata tensor sizes
(
- (work_meta_data_size, work_meta_data_type),
- (work_indptr_size, work_indptr_type),
- (work_info_set_size, work_info_set_type),
- (reduce_indptr_size, reduce_indptr_type),
- (reduce_final_map_size, reduce_final_map_type),
- (reduce_partial_map_size, reduce_partial_map_type),
+ (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,
- max_seqlen_qo,
- num_heads,
- dtypes.fp8, # q dtype
- dtypes.fp8, # kv dtype
- is_sparse=False,
- fast_mode=True,
- num_kv_splits=max_split_per_batch,
- intra_batch_mode=False,
+ batch_size, 1, num_heads, dtypes.fp8, dtypes.fp8,
+ is_sparse=False, fast_mode=True, num_kv_splits=-1, intra_batch_mode=False,
)
- # Pre-allocate metadata tensors
- work_meta_data = torch.empty(work_meta_data_size, dtype=work_meta_data_type, device=device)
- work_indptr = torch.empty(work_indptr_size, dtype=work_indptr_type, device=device)
- work_info_set = torch.empty(work_info_set_size, dtype=work_info_set_type, device=device)
- reduce_indptr = torch.empty(reduce_indptr_size, dtype=reduce_indptr_type, device=device)
- reduce_final_map = torch.empty(reduce_final_map_size, dtype=reduce_final_map_type, device=device)
- reduce_partial_map = torch.empty(reduce_partial_map_size, dtype=reduce_partial_map_type, device=device)
+ 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)
- # Populate metadata
aiter.get_mla_metadata_v1(
- qo_indptr,
- kv_indptr,
- kv_last_page_lens,
- num_heads // nhead_kv, # num_heads_per_head_k
- nhead_kv, # num_heads_k
- False, # is_causal
- work_meta_data,
- work_info_set,
- work_indptr,
- reduce_indptr,
- reduce_final_map,
- reduce_partial_map,
- page_size=page_size,
- kv_granularity=max(page_size, 16),
- max_seqlen_qo=max_seqlen_qo,
- uni_seqlen_qo=1, # decode
- fast_mode=True,
- max_split_per_batch=max_split_per_batch,
- intra_batch_mode=False,
- dtype_q=dtypes.fp8,
- dtype_kv=dtypes.fp8,
+ 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,
)
- result = (work_meta_data, work_indptr, work_info_set,
- reduce_indptr, reduce_final_map, reduce_partial_map)
+ result = (wmd, wi, wis, ri, rfm, rpm)
_meta_cache[key] = result
return result
⋯ 24 unchanged lines
kv_indptr_pages = kv_indptr // PAGE_SIZE
kv_last_page_lens = torch.full((batch_size,), PAGE_SIZE, device=q.device, dtype=torch.int32)
- # Get or compute persistent metadata
+ # 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,
scrolls · 102 diff lines total

Best evidence level for this revision: reported

JSON