Skip to content
KernelIndex
Search⌘K

submission 598454

johnny.t.shi · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:5f0c604718b62e345bc3390cd5eb651b69f50872b8400b81cb74310d8d39264d
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.py150 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):
    """Compute and cache persistent mode metadata."""
    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),
    ) = 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,
    )

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

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

    result = (work_meta_data, work_indptr, work_info_set,
              reduce_indptr, reduce_final_map, reduce_partial_map)
    _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)

    # Get or compute persistent metadata
    (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 · 150 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 592885.

#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
- """MLA Decode v1 — use aiter's ASM-optimized MLA decode."""
+ """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):
+ """Compute and cache persistent mode metadata."""
+ 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),
+ ) = 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,
+ )
+
+ # 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)
+
+ # 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,
+ )
+
+ result = (work_meta_data, work_indptr, work_info_set,
+ reduce_indptr, reduce_final_map, reduce_partial_map)
+ _meta_cache[key] = result
+ return result
+
+
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
- kv_bf16 = kv_data["bf16"]
+ kv_fp8_data, kv_fp8_scale = kv_data["fp8"]
batch_size = config['batch_size']
num_heads = config['num_heads']
- qk_head_dim = config['qk_head_dim']
v_head_dim = config['v_head_dim']
sm_scale = config['sm_scale']
- total_kv = kv_bf16.shape[0]
+ 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)
- # Treat contiguous KV as paged with page_size=1
- # kv_bf16 shape: [total_kv, nhead_kv, head_dim] → reshape to [total_kv, 1, nhead_kv, head_dim]
- if kv_bf16.dim() == 3:
- kv_buffer = kv_bf16.unsqueeze(1) # [total_kv, 1, nhead_kv, head_dim]
- elif kv_bf16.dim() == 2:
- kv_buffer = kv_bf16.unsqueeze(1).unsqueeze(2) # [total_kv, 1, 1, head_dim]
- else:
- kv_buffer = kv_bf16
+ # 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(total_kv, device=q.device, dtype=torch.int32)
- kv_last_page_lens = torch.ones(batch_size, device=q.device, dtype=torch.int32)
+ 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)
+ # Get or compute persistent metadata
+ (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,
+ q=q_fp8,
kv_buffer=kv_buffer,
o=output,
qo_indptr=qo_indptr,
- kv_indptr=kv_indptr,
+ kv_indptr=kv_indptr_pages,
kv_indices=kv_indices,
kv_last_page_lens=kv_last_page_lens,
max_seqlen_q=1,
- page_size=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 · 166 diff lines total

Best evidence level for this revision: reported

JSON