submission 660305
yuzhou_lithos · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 140 lines, June 9 Researcher Reciprocity License v1.0.
submission_v96_persist_split1.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-660305?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:68481b60e0fad5b80a16725003a84a9c50b4eab14382fb6d0367e6cd3aff427a
license declaredunknown
license concludedunknown
authorsyuzhou_lithos
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
persistent-kernel
"""MLA decode v96 — persistent with max_split_per_batch=1 for bs=256 — Wave 248.Kernel source
submission_v96_persist_split1.py140 lines
"""MLA decode v96 — persistent with max_split_per_batch=1 for bs=256 — Wave 248.
v94's non-persistent 1-split was unreliable (correctness failures in leaderboard).
This version stays FULLY persistent but uses max_split_per_batch=1 for bs=256,
minimizing the reduce work while keeping the proven .co kernel.
For bs=256: persistent with max_split=1 → 256 CUs, 1 batch/CU, minimal reduce.
For other bs: persistent with max_split=32 (v82 approach, proven correct).
"""
import math, torch
from task import input_t, output_t
import aiter
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
SM_SCALE = 1.0 / math.sqrt(576)
FP8_DTYPE = aiter_dtypes.fp8
PAGE_SIZE = 1
H = 16
NKV = 1
DK = 576
DV = 512
_cache = {}
def _get_splits(bs):
"""Choose max_split_per_batch based on batch size."""
if bs >= 256:
return 1 # 256 batches fully fill 256 CUs, no KV splitting needed
return 32 # default: 32 max splits for CU utilization
def _ensure_cached(bs, kvl, device):
key = (bs, kvl)
if key in _cache:
return _cache[key]
total_kv = bs * kvl
max_splits = _get_splits(bs)
qo_indptr = torch.arange(0, bs + 1, dtype=torch.int32, device=device)
kv_indptr = torch.arange(0, bs + 1, dtype=torch.int32, device=device) * kvl
kv_last_page_len = torch.full((bs,), kvl, dtype=torch.int32, device=device)
kv_indices = torch.arange(total_kv, dtype=torch.int32, device=device)
out = torch.empty((bs, H, DV), dtype=torch.bfloat16, device=device)
q_fp8 = torch.empty((bs * H, DK), dtype=FP8_DTYPE, device=device)
q_fp8_3d = q_fp8.view(bs, H, DK)
q_scale = torch.ones(1, dtype=torch.float32, device=device)
info = get_mla_metadata_info_v1(
bs, 1, H, FP8_DTYPE, FP8_DTYPE,
is_sparse=False, fast_mode=False,
num_kv_splits=max_splits, intra_batch_mode=True,
)
work = [torch.empty(s, dtype=t, device=device) for s, t in info]
(work_metadata, work_indptr, work_info_set,
reduce_indptr, reduce_final_map, reduce_partial_map) = work
get_mla_metadata_v1(
qo_indptr, kv_indptr, kv_last_page_len,
H // NKV, NKV, True,
work_metadata, 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=1,
uni_seqlen_qo=1,
fast_mode=False,
max_split_per_batch=max_splits,
intra_batch_mode=True,
dtype_q=FP8_DTYPE,
dtype_kv=FP8_DTYPE,
)
rp_size = reduce_partial_map.size(0)
logits = torch.empty((rp_size, 1, H, DV), dtype=torch.float32, device=device)
attn_lse = torch.empty((rp_size, 1, H, 1), dtype=torch.float32, device=device)
_cache[key] = {
"qo_indptr": qo_indptr,
"kv_indptr": kv_indptr,
"kv_indices": kv_indices,
"kv_last_page_len": kv_last_page_len,
"work_meta_data": work_metadata,
"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,
"output": out,
"q_fp8": q_fp8,
"q_fp8_3d": q_fp8_3d,
"q_scale": q_scale,
"logits": logits,
"attn_lse": attn_lse,
}
return _cache[key]
@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr_in, kv_indptr_in, cfg = data
bs = cfg["batch_size"]
kvl = cfg["kv_seq_len"]
c = _ensure_cached(bs, kvl, q.device)
q_fp8 = c["q_fp8"]
q_fp8_3d = c["q_fp8_3d"]
logits = c["logits"]
attn_lse = c["attn_lse"]
output = c["output"]
q_scale = c["q_scale"]
q_2d = q.view(-1, DK)
q_fp8.copy_(q_2d)
kv_fp8, kv_scale = kv_data["fp8"]
tkv = bs * kvl
kv_4d = kv_fp8[:tkv].view(tkv, PAGE_SIZE, NKV, DK)
aiter.mla_decode_stage1_asm_fwd(
q_fp8_3d, kv_4d,
c["qo_indptr"], c["kv_indptr"], c["kv_indices"], c["kv_last_page_len"],
None, c["work_meta_data"], c["work_indptr"], c["work_info_set"],
1, PAGE_SIZE, NKV, SM_SCALE,
logits, attn_lse, output, q_scale, kv_scale,
)
aiter.mla_reduce_v1(
logits, attn_lse,
c["reduce_indptr"], c["reduce_final_map"], c["reduce_partial_map"],
1, output, None,
)
return output
scrolls · 140 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