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