submission 663004
Maxwell Cipher · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 53 lines, June 9 Researcher Reciprocity License v1.0.
mla_v41.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-663004?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:759f799b1c657e53ddf306c20cdda5df46f7abbaaf2a1c53aab588214a5429c8
license declaredunknown
license concludedunknown
authorsMaxwell Cipher
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
persistent-kernel
"""v41: Full bf16 non-persistent MLA decode.Kernel source
mla_v41.py53 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""v41: Full bf16 non-persistent MLA decode.
Skip fp8 quantization entirely — single kernel launch with bf16 Q and KV.
Overhead savings (no quant, no metadata, no reduce) outweigh 2x bandwidth cost."""
import torch
from task import input_t, output_t
from aiter.mla import mla_decode_fwd
_NH = 16
_NKV = 1
_QK_DIM = 576
_V_DIM = 512
_SM_SC = 1.0 / (_QK_DIM ** 0.5)
_shape_bufs = {}
def _get_shape_buffers(bs, seq_len, dev):
k = (bs, seq_len)
if k not in _shape_bufs:
n = bs * seq_len
page_ids = torch.arange(n, dtype=torch.int32, device=dev)
seq_lens = torch.full((bs,), seq_len, dtype=torch.int32, device=dev)
_shape_bufs[k] = (page_ids, seq_lens)
return _shape_bufs[k]
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
bs = int(config["batch_size"])
seq_len = int(config["kv_seq_len"])
n_tokens = q.shape[0]
q_reshaped = q.view(n_tokens, _NH, _QK_DIM)
kv_raw = kv_data["bf16"]
kv_paged = kv_raw.view(-1, 1, _NKV, kv_raw.shape[-1])
page_ids, seq_lens = _get_shape_buffers(bs, seq_len, q.device)
out = torch.empty((n_tokens, _NH, _V_DIM), dtype=torch.bfloat16, device=q.device)
mla_decode_fwd(
q_reshaped, kv_paged, out,
qo_indptr, kv_indptr,
page_ids, seq_lens, 1,
page_size=1, nhead_kv=_NKV, sm_scale=_SM_SC,
intra_batch_mode=False,
)
return out
scrolls · 53 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 660382.
#!POPCORN leaderboard amd-mixed-mla#!POPCORN gpu MI355X- """v41: Pure bf16 non-persistent — the dgavriloff insight.- Eliminates ALL overhead: no Q quantization, no metadata, no reduce.- Only 1 kernel launch per call. Trades 2x bandwidth for zero overhead."""+ """v41: Full bf16 non-persistent MLA decode.+ Skip fp8 quantization entirely — single kernel launch with bf16 Q and KV.+ Overhead savings (no quant, no metadata, no reduce) outweigh 2x bandwidth cost."""import torchfrom task import input_t, output_tfrom aiter.mla import mla_decode_fwd- NUM_HEADS = 16- NUM_KV_HEADS = 1- QK_HEAD_DIM = 576- V_HEAD_DIM = 512- SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)+ _NH = 16+ _NKV = 1+ _QK_DIM = 576+ _V_DIM = 512+ _SM_SC = 1.0 / (_QK_DIM ** 0.5)- _cache = {}+ _shape_bufs = {}+ def _get_shape_buffers(bs, seq_len, dev):+ k = (bs, seq_len)+ if k not in _shape_bufs:+ n = bs * seq_len+ page_ids = torch.arange(n, dtype=torch.int32, device=dev)+ seq_lens = torch.full((bs,), seq_len, dtype=torch.int32, device=dev)+ _shape_bufs[k] = (page_ids, seq_lens)+ return _shape_bufs[k]++def custom_kernel(data: input_t) -> output_t:q, kv_data, qo_indptr, kv_indptr, config = data- batch_size = int(config["batch_size"])- kv_seq_len = int(config["kv_seq_len"])- q_total = q.shape[0]+ bs = int(config["batch_size"])+ seq_len = int(config["kv_seq_len"])+ n_tokens = q.shape[0]- kv_bf16 = kv_data["bf16"]- q_bf16 = q.view(-1, NUM_HEADS, QK_HEAD_DIM)- kv_4d = kv_bf16.view(-1, 1, NUM_KV_HEADS, kv_bf16.shape[-1])+ q_reshaped = q.view(n_tokens, _NH, _QK_DIM)+ kv_raw = kv_data["bf16"]+ kv_paged = kv_raw.view(-1, 1, _NKV, kv_raw.shape[-1])- key = (batch_size, kv_seq_len)- if key not in _cache:- total_kv = batch_size * kv_seq_len- _cache[key] = (- torch.arange(total_kv, dtype=torch.int32, device="cuda"),- torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device="cuda"),- )+ page_ids, seq_lens = _get_shape_buffers(bs, seq_len, q.device)+ out = torch.empty((n_tokens, _NH, _V_DIM), dtype=torch.bfloat16, device=q.device)- kv_indices, kv_last_page_len = _cache[key]- output = torch.empty((q_total, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")-mla_decode_fwd(- q_bf16, kv_4d, output,+ q_reshaped, kv_paged, out,qo_indptr, kv_indptr,- kv_indices, kv_last_page_len,- 1,- page_size=1, nhead_kv=NUM_KV_HEADS, sm_scale=SM_SCALE,+ page_ids, seq_lens, 1,+ page_size=1, nhead_kv=_NKV, sm_scale=_SM_SC,intra_batch_mode=False,)- return output+ return out
scrolls · 82 diff lines total
Best evidence level for this revision: reported
JSON