submission 586726
josusanmartin · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 184 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-586726?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:7c851eb4da17ed4e2e14896932e74b855333c3b83bc67d4a868c2d8af5a85c57
license declaredunknown
license concludedunknown
authorsjosusanmartin
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
persistent-kernel
"""v707 - bf16 for small shapes, 1-split for (64,1K)+(256,1K), persistent for 8K."""Kernel source
submission.py184 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""v707 - bf16 for small shapes, 1-split for (64,1K)+(256,1K), persistent for 8K."""
import os
os.environ.setdefault("HIP_FORCE_DEV_KERNARG", "1")
os.environ.setdefault("AMD_DIRECT_DISPATCH", "1")
import torch
from task import input_t, output_t
import aiter
from aiter import dtypes as aiter_dtypes
from aiter import mla as aiter_mla
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
try:
from aiter.jit.module_quant import static_per_tensor_quant
except Exception:
try:
from aiter.ops.quant import static_per_tensor_quant
except Exception:
static_per_tensor_quant = None
NUM_HEADS = 16
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
FP8_DTYPE = aiter_dtypes.fp8
_FP8_FINFO = torch.finfo(FP8_DTYPE)
_cache = {}
def _quantize_q_fp8(q, q_fp8_buf):
amax = q.abs().amax().clamp(min=1e-12)
scale = (amax / _FP8_FINFO.max).reshape(1).to(torch.float32)
if static_per_tensor_quant is not None:
static_per_tensor_quant(q_fp8_buf, q, scale)
else:
q_fp8_buf.copy_(
(q / scale).clamp(min=_FP8_FINFO.min, max=_FP8_FINFO.max).to(FP8_DTYPE)
)
return scale
def _build_nonpersist_1split(dev, bs, kvlen, kv_indptr):
"""Non-persistent mode with 1 split: output written directly, no reduce needed."""
total_kv = bs * kvlen
_, num_splits_indptr = aiter_mla.get_meta_param(1, bs, total_kv, NUM_HEADS, 1, FP8_DTYPE)
kv_indices = torch.arange(total_kv, dtype=torch.int32, device=dev)
kv_lpl = torch.full((bs,), kvlen, dtype=torch.int32, device=dev)
out = torch.empty((bs, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=dev)
# With 1 split, logits shares memory with out
logits = out.view(bs, 1, NUM_HEADS, V_HEAD_DIM)
attn_lse = torch.empty((bs, 1, NUM_HEADS, 1), dtype=torch.float32, device=dev)
q_fp8 = torch.empty((bs, NUM_HEADS, QK_HEAD_DIM), dtype=FP8_DTYPE, device=dev)
return (num_splits_indptr, kv_indices, kv_lpl, out, logits, attn_lse, q_fp8)
def _build_persistent(dev, qo_indptr, kv_indptr, bs, kvlen, kv_gran, intra):
"""Persistent mode with metadata — uses mla_reduce_v1."""
total_kv = bs * kvlen
num_splits, _ = aiter_mla.get_meta_param(None, bs, total_kv, NUM_HEADS, 1, FP8_DTYPE)
kv_indices = torch.arange(total_kv, dtype=torch.int32, device=dev)
kv_lpl = torch.full((bs,), kvlen, dtype=torch.int32, device=dev)
out = torch.empty((bs, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=dev)
info = get_mla_metadata_info_v1(
bs, 1, NUM_HEADS, FP8_DTYPE, FP8_DTYPE,
is_sparse=False, fast_mode=True, num_kv_splits=num_splits, intra_batch_mode=intra,
)
bufs = [torch.empty(s, dtype=t, device=dev) for s, t in info]
wmd, wi, wis, ri, rfm, rpm = bufs
get_mla_metadata_v1(
qo_indptr, kv_indptr, kv_lpl, 16, 1, False,
wmd, wis, wi, ri, rfm, rpm,
page_size=1, kv_granularity=kv_gran, max_seqlen_qo=1, uni_seqlen_qo=1,
fast_mode=True, max_split_per_batch=num_splits, intra_batch_mode=intra,
dtype_q=FP8_DTYPE, dtype_kv=FP8_DTYPE,
)
pt = int(rpm.numel())
po = torch.empty((pt, 1, NUM_HEADS, V_HEAD_DIM), dtype=torch.float32, device=dev)
pl = torch.empty((pt, 1, NUM_HEADS, 1), dtype=torch.float32, device=dev)
q_fp8 = torch.empty((bs, NUM_HEADS, QK_HEAD_DIM), dtype=FP8_DTYPE, device=dev)
return (kv_indices, kv_lpl, out, wmd, wi, wis, ri, rfm, rpm, po, pl, q_fp8)
def _build_bf16_persistent(dev, qo_indptr, kv_indptr, bs, kvlen, kv_gran, intra):
"""Persistent mode with bf16 Q + bf16 KV — no quantization needed."""
total_kv = bs * kvlen
num_splits, _ = aiter_mla.get_meta_param(None, bs, total_kv, NUM_HEADS, 1, torch.bfloat16)
kv_indices = torch.arange(total_kv, dtype=torch.int32, device=dev)
kv_lpl = torch.full((bs,), kvlen, dtype=torch.int32, device=dev)
out = torch.empty((bs, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=dev)
info = get_mla_metadata_info_v1(
bs, 1, NUM_HEADS, torch.bfloat16, torch.bfloat16,
is_sparse=False, fast_mode=True, num_kv_splits=num_splits, intra_batch_mode=intra,
)
bufs = [torch.empty(s, dtype=t, device=dev) for s, t in info]
wmd, wi, wis, ri, rfm, rpm = bufs
get_mla_metadata_v1(
qo_indptr, kv_indptr, kv_lpl, 16, 1, False,
wmd, wis, wi, ri, rfm, rpm,
page_size=1, kv_granularity=kv_gran, max_seqlen_qo=1, uni_seqlen_qo=1,
fast_mode=True, max_split_per_batch=num_splits, intra_batch_mode=intra,
dtype_q=torch.bfloat16, dtype_kv=torch.bfloat16,
)
pt = int(rpm.numel())
po = torch.empty((pt, 1, NUM_HEADS, V_HEAD_DIM), dtype=torch.float32, device=dev)
pl = torch.empty((pt, 1, NUM_HEADS, 1), dtype=torch.float32, device=dev)
return (kv_indices, kv_lpl, out, wmd, wi, wis, ri, rfm, rpm, po, pl)
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
bs = int(config["batch_size"])
kvlen = int(config["kv_seq_len"])
sm_scale = float(config["sm_scale"])
dev = q.device
kv_fp8, kv_scale = kv_data["fp8"]
kv_buf = kv_fp8.view(-1, 1, 1, QK_HEAD_DIM)
key = (dev.index, bs, kvlen)
# --- bf16 path for small batches: skip Q quantization entirely ---
if bs <= 32 and kvlen == 1024:
bkey = ("bf16", *key)
if bkey not in _cache:
_cache[bkey] = _build_bf16_persistent(dev, qo_indptr, kv_indptr, bs, kvlen, 16, True)
c = _cache[bkey]
kv_bf16 = kv_data["bf16"].view(-1, 1, 1, QK_HEAD_DIM)
aiter.mla_decode_stage1_asm_fwd(
q, kv_bf16, qo_indptr, kv_indptr, c[0], c[1], None,
c[3], c[4], c[5], 1, 1, 1, sm_scale, c[9], c[10], c[2], None, None,
)
aiter.mla_reduce_v1(c[9], c[10], c[6], c[7], c[8], 1, c[2], None)
return c[2]
if bs == 4 and kvlen == 8192:
bkey = ("bf16", *key)
if bkey not in _cache:
_cache[bkey] = _build_bf16_persistent(dev, qo_indptr, kv_indptr, bs, kvlen, 64, True)
c = _cache[bkey]
kv_bf16 = kv_data["bf16"].view(-1, 1, 1, QK_HEAD_DIM)
aiter.mla_decode_stage1_asm_fwd(
q, kv_bf16, qo_indptr, kv_indptr, c[0], c[1], None,
c[3], c[4], c[5], 1, 1, 1, sm_scale, c[9], c[10], c[2], None, None,
)
aiter.mla_reduce_v1(c[9], c[10], c[6], c[7], c[8], 1, c[2], None)
return c[2]
# --- 1-split non-persistent for select shapes: skip reduce ---
if (bs, kvlen) in ((64, 1024), (256, 1024)):
nkey = ("np1", *key)
if nkey not in _cache:
_cache[nkey] = _build_nonpersist_1split(dev, bs, kvlen, kv_indptr)
c = _cache[nkey]
q_scale = _quantize_q_fp8(q, c[6])
aiter.mla_decode_stage1_asm_fwd(
c[6], kv_buf, qo_indptr, kv_indptr, c[1], c[2], c[0],
None, None, None, 1, 1, 1, sm_scale, c[4], c[5], c[3], q_scale, kv_scale,
)
return c[3]
# --- Persistent mode for remaining shapes ---
pkey = ("persist", *key)
if pkey not in _cache:
kv_gran = 64 if kvlen == 8192 else 8
intra = True
_cache[pkey] = _build_persistent(dev, qo_indptr, kv_indptr, bs, kvlen, kv_gran, intra)
c = _cache[pkey]
q_scale = _quantize_q_fp8(q, c[11])
aiter.mla_decode_stage1_asm_fwd(
c[11], kv_buf, qo_indptr, kv_indptr, c[0], c[1], None,
c[3], c[4], c[5], 1, 1, 1, sm_scale, c[9], c[10], c[2], q_scale, kv_scale,
)
aiter.mla_reduce_v1(c[9], c[10], c[6], c[7], c[8], 1, c[2], None)
return c[2]
scrolls · 184 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 586016.
#!POPCORN leaderboard amd-mixed-mla#!POPCORN gpu MI355X- """v697 - Direct stage1+reduce with PAGE_SIZE=1, pre-allocated split buffers."""+ """v707 - bf16 for small shapes, 1-split for (64,1K)+(256,1K), persistent for 8K."""import os⋯ 5 unchanged linesimport aiterfrom aiter import dtypes as aiter_dtypes+ from aiter import mla as aiter_mlafrom aiter import get_mla_metadata_info_v1, get_mla_metadata_v1try:⋯ 7 unchanged linesNUM_HEADS = 16QK_HEAD_DIM = 576V_HEAD_DIM = 512- NUM_KV_SPLITS = 32FP8_DTYPE = aiter_dtypes.fp8_FP8_FINFO = torch.finfo(FP8_DTYPE)⋯ 13 unchanged linesreturn scale- def _get_cache(dev, qo_indptr, kv_indptr, bs, kvlen):- key = (dev.index, bs, kvlen)- c = _cache.get(key)- if c is not None:- return c-+ def _build_nonpersist_1split(dev, bs, kvlen, kv_indptr):+ """Non-persistent mode with 1 split: output written directly, no reduce needed."""total_kv = bs * kvlen+ _, num_splits_indptr = aiter_mla.get_meta_param(1, bs, total_kv, NUM_HEADS, 1, FP8_DTYPE)kv_indices = torch.arange(total_kv, dtype=torch.int32, device=dev)- kv_lpl = torch.ones(bs, dtype=torch.int32, device=dev)+ kv_lpl = torch.full((bs,), kvlen, dtype=torch.int32, device=dev)out = torch.empty((bs, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=dev)+ # With 1 split, logits shares memory with out+ logits = out.view(bs, 1, NUM_HEADS, V_HEAD_DIM)+ attn_lse = torch.empty((bs, 1, NUM_HEADS, 1), dtype=torch.float32, device=dev)q_fp8 = torch.empty((bs, NUM_HEADS, QK_HEAD_DIM), dtype=FP8_DTYPE, device=dev)+ return (num_splits_indptr, kv_indices, kv_lpl, out, logits, attn_lse, q_fp8)- # Metadata matching reference: PAGE_SIZE=1, is_causal=True, fast_mode=False++ def _build_persistent(dev, qo_indptr, kv_indptr, bs, kvlen, kv_gran, intra):+ """Persistent mode with metadata — uses mla_reduce_v1."""+ total_kv = bs * kvlen+ num_splits, _ = aiter_mla.get_meta_param(None, bs, total_kv, NUM_HEADS, 1, FP8_DTYPE)+ kv_indices = torch.arange(total_kv, dtype=torch.int32, device=dev)+ kv_lpl = torch.full((bs,), kvlen, dtype=torch.int32, device=dev)+ out = torch.empty((bs, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=dev)+info = get_mla_metadata_info_v1(bs, 1, NUM_HEADS, FP8_DTYPE, FP8_DTYPE,- is_sparse=False, fast_mode=False,- num_kv_splits=NUM_KV_SPLITS, intra_batch_mode=True,+ is_sparse=False, fast_mode=True, num_kv_splits=num_splits, intra_batch_mode=intra,)bufs = [torch.empty(s, dtype=t, device=dev) for s, t in info]wmd, wi, wis, ri, rfm, rpm = bufsget_mla_metadata_v1(- qo_indptr, kv_indptr, kv_lpl,- 16, 1, True,+ qo_indptr, kv_indptr, kv_lpl, 16, 1, False,wmd, wis, wi, ri, rfm, rpm,- page_size=1, kv_granularity=16,- max_seqlen_qo=1, uni_seqlen_qo=1,- fast_mode=False, max_split_per_batch=NUM_KV_SPLITS,- intra_batch_mode=True,+ page_size=1, kv_granularity=kv_gran, max_seqlen_qo=1, uni_seqlen_qo=1,+ fast_mode=True, max_split_per_batch=num_splits, intra_batch_mode=intra,dtype_q=FP8_DTYPE, dtype_kv=FP8_DTYPE,)-- # Pre-allocate split buffers (size from reduce_partial_map)pt = int(rpm.numel())- split_out = torch.empty((pt, 1, NUM_HEADS, V_HEAD_DIM), dtype=torch.float32, device=dev)- split_lse = torch.empty((pt, 1, NUM_HEADS, 1), dtype=torch.float32, device=dev)+ po = torch.empty((pt, 1, NUM_HEADS, V_HEAD_DIM), dtype=torch.float32, device=dev)+ pl = torch.empty((pt, 1, NUM_HEADS, 1), dtype=torch.float32, device=dev)+ q_fp8 = torch.empty((bs, NUM_HEADS, QK_HEAD_DIM), dtype=FP8_DTYPE, device=dev)+ return (kv_indices, kv_lpl, out, wmd, wi, wis, ri, rfm, rpm, po, pl, q_fp8)- c = (kv_indices, kv_lpl, out, q_fp8, wmd, wi, wis, ri, rfm, rpm, split_out, split_lse)- _cache[key] = c- return c+ def _build_bf16_persistent(dev, qo_indptr, kv_indptr, bs, kvlen, kv_gran, intra):+ """Persistent mode with bf16 Q + bf16 KV — no quantization needed."""+ total_kv = bs * kvlen+ num_splits, _ = aiter_mla.get_meta_param(None, bs, total_kv, NUM_HEADS, 1, torch.bfloat16)+ kv_indices = torch.arange(total_kv, dtype=torch.int32, device=dev)+ kv_lpl = torch.full((bs,), kvlen, dtype=torch.int32, device=dev)+ out = torch.empty((bs, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=dev)+ info = get_mla_metadata_info_v1(+ bs, 1, NUM_HEADS, torch.bfloat16, torch.bfloat16,+ is_sparse=False, fast_mode=True, num_kv_splits=num_splits, intra_batch_mode=intra,+ )+ bufs = [torch.empty(s, dtype=t, device=dev) for s, t in info]+ wmd, wi, wis, ri, rfm, rpm = bufs+ get_mla_metadata_v1(+ qo_indptr, kv_indptr, kv_lpl, 16, 1, False,+ wmd, wis, wi, ri, rfm, rpm,+ page_size=1, kv_granularity=kv_gran, max_seqlen_qo=1, uni_seqlen_qo=1,+ fast_mode=True, max_split_per_batch=num_splits, intra_batch_mode=intra,+ dtype_q=torch.bfloat16, dtype_kv=torch.bfloat16,+ )+ pt = int(rpm.numel())+ po = torch.empty((pt, 1, NUM_HEADS, V_HEAD_DIM), dtype=torch.float32, device=dev)+ pl = torch.empty((pt, 1, NUM_HEADS, 1), dtype=torch.float32, device=dev)+ return (kv_indices, kv_lpl, out, wmd, wi, wis, ri, rfm, rpm, po, pl)++def custom_kernel(data: input_t) -> output_t:q, kv_data, qo_indptr, kv_indptr, config = databs = int(config["batch_size"])kvlen = int(config["kv_seq_len"])sm_scale = float(config["sm_scale"])+ dev = q.device- cache = _get_cache(q.device, qo_indptr, kv_indptr, bs, kvlen)- kv_indices, kv_lpl, out, q_fp8, wmd, wi, wis, ri, rfm, rpm, split_out, split_lse = cache-- q_scale = _quantize_q_fp8(q, q_fp8)kv_fp8, kv_scale = kv_data["fp8"]kv_buf = kv_fp8.view(-1, 1, 1, QK_HEAD_DIM)+ key = (dev.index, bs, kvlen)++ # --- bf16 path for small batches: skip Q quantization entirely ---+ if bs <= 32 and kvlen == 1024:+ bkey = ("bf16", *key)+ if bkey not in _cache:+ _cache[bkey] = _build_bf16_persistent(dev, qo_indptr, kv_indptr, bs, kvlen, 16, True)+ c = _cache[bkey]+ kv_bf16 = kv_data["bf16"].view(-1, 1, 1, QK_HEAD_DIM)+ aiter.mla_decode_stage1_asm_fwd(+ q, kv_bf16, qo_indptr, kv_indptr, c[0], c[1], None,+ c[3], c[4], c[5], 1, 1, 1, sm_scale, c[9], c[10], c[2], None, None,+ )+ aiter.mla_reduce_v1(c[9], c[10], c[6], c[7], c[8], 1, c[2], None)+ return c[2]++ if bs == 4 and kvlen == 8192:+ bkey = ("bf16", *key)+ if bkey not in _cache:+ _cache[bkey] = _build_bf16_persistent(dev, qo_indptr, kv_indptr, bs, kvlen, 64, True)+ c = _cache[bkey]+ kv_bf16 = kv_data["bf16"].view(-1, 1, 1, QK_HEAD_DIM)+ aiter.mla_decode_stage1_asm_fwd(+ q, kv_bf16, qo_indptr, kv_indptr, c[0], c[1], None,+ c[3], c[4], c[5], 1, 1, 1, sm_scale, c[9], c[10], c[2], None, None,+ )+ aiter.mla_reduce_v1(c[9], c[10], c[6], c[7], c[8], 1, c[2], None)+ return c[2]++ # --- 1-split non-persistent for select shapes: skip reduce ---+ if (bs, kvlen) in ((64, 1024), (256, 1024)):+ nkey = ("np1", *key)+ if nkey not in _cache:+ _cache[nkey] = _build_nonpersist_1split(dev, bs, kvlen, kv_indptr)+ c = _cache[nkey]+ q_scale = _quantize_q_fp8(q, c[6])+ aiter.mla_decode_stage1_asm_fwd(+ c[6], kv_buf, qo_indptr, kv_indptr, c[1], c[2], c[0],+ None, None, None, 1, 1, 1, sm_scale, c[4], c[5], c[3], q_scale, kv_scale,+ )+ return c[3]++ # --- Persistent mode for remaining shapes ---+ pkey = ("persist", *key)+ if pkey not in _cache:+ kv_gran = 64 if kvlen == 8192 else 8+ intra = True+ _cache[pkey] = _build_persistent(dev, qo_indptr, kv_indptr, bs, kvlen, kv_gran, intra)+ c = _cache[pkey]+ q_scale = _quantize_q_fp8(q, c[11])aiter.mla_decode_stage1_asm_fwd(- q_fp8, kv_buf, qo_indptr, kv_indptr, kv_indices, kv_lpl, None,- wmd, wi, wis, 1, 1, 1, sm_scale, split_out, split_lse, out, q_scale, kv_scale,+ c[11], kv_buf, qo_indptr, kv_indptr, c[0], c[1], None,+ c[3], c[4], c[5], 1, 1, 1, sm_scale, c[9], c[10], c[2], q_scale, kv_scale,)- aiter.mla_reduce_v1(split_out, split_lse, ri, rfm, rpm, 1, out, None)- return out+ aiter.mla_reduce_v1(c[9], c[10], c[6], c[7], c[8], 1, c[2], None)+ return c[2]
scrolls · 194 diff lines total
Best evidence level for this revision: reported
JSON