Skip to content
KernelIndex
Search⌘K

submission 689325

nanbeilvdougao · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_20260401_v34_pg8only.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-689325?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
45.2µs
#141 of 766
2026-04-01

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:6d72079950aabc506d7d389c1bd8a74c125292c23a0029b1e4c13ff5303e1c86
license declaredunknown
license concludedunknown
authorsnanbeilvdougao
imported2026-08-15

Kernel source

submission_20260401_v34_pg8only.py126 lines
"""
v34: Same as v30 (page1 only, a16w8, minimal) but with page2 ONLY for kv=8192.
kv=1024 stays page1 (precision safe).
This should give better perf than v30 (page2 helps for kv=8192 large batch)
while keeping kv=1024 precision clean.
"""
import os as _os
import sys as _sys

_devnull_fd = _os.open(_os.devnull, _os.O_WRONLY)
_orig_stderr_fd = _os.dup(2)
_os.dup2(_devnull_fd, 2)
_sys.stderr = open(_os.devnull, 'w')

import torch
from task import input_t, output_t

NUM_HEADS = 16
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
SM_SCALE = QK_HEAD_DIM ** -0.5
PAGE_SIZE = 1
PAGE_SIZE_LONG = 8
KV_GRANULARITY = 16
KV_GRANULARITY_LONG = 32
NUM_KV_SPLITS = 32

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

FP8_DTYPE = aiter_dtypes.fp8

_cache = {}


def _build_meta(batch_size, device, q_dtype, kv_dtype, qo_indptr, kv_indptr, kv_last_page_len, page_size, kv_gran):
    info = get_mla_metadata_info_v1(
        batch_size, 1, NUM_HEADS, q_dtype, kv_dtype,
        is_sparse=False, fast_mode=False,
        num_kv_splits=NUM_KV_SPLITS, intra_batch_mode=True,
    )
    work = [torch.empty(shape, dtype=dtype, device=device) for shape, dtype in info]
    get_mla_metadata_v1(
        qo_indptr, kv_indptr, kv_last_page_len,
        NUM_HEADS // NUM_KV_HEADS, NUM_KV_HEADS, True,
        work[0], work[2], work[1], work[3], work[4], work[5],
        page_size=page_size, kv_granularity=kv_gran,
        max_seqlen_qo=1, uni_seqlen_qo=1,
        fast_mode=False, max_split_per_batch=NUM_KV_SPLITS,
        intra_batch_mode=True, dtype_q=q_dtype, dtype_kv=kv_dtype,
    )
    return work


def _build_state(batch_size, kv_seq_len, device):
    total_q = batch_size
    total_kv = batch_size * kv_seq_len

    qo_indptr = torch.arange(batch_size + 1, dtype=torch.int32, device=device)
    out = torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=device)

    # Decide page size based on kv_seq_len
    if kv_seq_len >= 8192 and kv_seq_len % PAGE_SIZE_LONG == 0 and total_kv % PAGE_SIZE_LONG == 0:
        ps = PAGE_SIZE_LONG
        kv_gran = KV_GRANULARITY_LONG
    else:
        ps = PAGE_SIZE
        kv_gran = KV_GRANULARITY

    if ps == 1:
        page_count = kv_seq_len
        kv_indptr = torch.arange(batch_size + 1, dtype=torch.int32, device=device) * page_count
        kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device=device)
        kv_indices = torch.arange(total_kv, dtype=torch.int32, device=device)
        num_pages = total_kv
    else:
        page_count = kv_seq_len // ps
        last_pl = kv_seq_len % ps
        if last_pl == 0:
            last_pl = ps
        kv_indptr = torch.arange(batch_size + 1, dtype=torch.int32, device=device) * page_count
        kv_last_page_len = torch.full((batch_size,), last_pl, dtype=torch.int32, device=device)
        kv_indices = torch.arange(total_kv // ps, dtype=torch.int32, device=device)
        num_pages = batch_size * page_count

    work = _build_meta(batch_size, device, torch.bfloat16, FP8_DTYPE,
                        qo_indptr, kv_indptr, kv_last_page_len, ps, kv_gran)

    return {
        'qo': qo_indptr, 'kv': kv_indptr, 'klp': kv_last_page_len,
        'ki': kv_indices, 'out': out, 'work': work,
        'total_kv': total_kv, 'ps': ps, 'num_pages': num_pages,
    }


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, _kv_indptr, _config = data
    kv_fp8, kv_scale = kv_data["fp8"]

    batch_size = qo_indptr.numel() - 1
    kv_seq_len = kv_fp8.shape[0] // batch_size

    key = (batch_size, kv_seq_len)
    if key not in _cache:
        _cache[key] = _build_state(batch_size, kv_seq_len, q.device)

    s = _cache[key]
    ps = s['ps']
    kv_4d = kv_fp8.view(s['num_pages'], ps, NUM_KV_HEADS, QK_HEAD_DIM)
    w = s['work']

    mla_decode_fwd(
        q, kv_4d, s['out'],
        s['qo'], s['kv'], s['ki'], s['klp'],
        1,
        page_size=ps, nhead_kv=NUM_KV_HEADS,
        sm_scale=SM_SCALE, logit_cap=0.0, num_kv_splits=NUM_KV_SPLITS,
        q_scale=None, kv_scale=kv_scale,
        intra_batch_mode=True,
        work_meta_data=w[0], work_indptr=w[1], work_info_set=w[2],
        reduce_indptr=w[3], reduce_final_map=w[4], reduce_partial_map=w[5],
    )
    return s['out']
scrolls · 126 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Best evidence level for this revision: reported

JSON