submission 742458
Ananda Sai A · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 236 lines, June 9 Researcher Reciprocity License v1.0.
submission_v224.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-742458?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:837416ca81b82ac5a62a78607482e457e5d7ed7c80ca6607a5f86af74b945de5
license declaredunknown
license concludedunknown
authorsAnanda Sai A
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
tile-n = 128
FP8_MIN_BLOCK_N = 128Kernel source
submission_v224.py236 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
v224
"""
import math
import torch
from task import input_t, output_t
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
NUM_HEADS = 16
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
SM_SCALE = 1.0 / math.sqrt(QK_HEAD_DIM)
FP8_DTYPE = aiter_dtypes.fp8
STATIC_Q_ABSMAX = 6.0
FP8_MAX = float(torch.finfo(FP8_DTYPE).max)
STATIC_Q_SCALE = STATIC_Q_ABSMAX / FP8_MAX
_QUANT_FN = None
try:
from aiter.ops.quant import static_per_tensor_quant as _sqf
_QUANT_FN = _sqf
except Exception:
try:
from aiter.jit.module_quant import static_per_tensor_quant as _sqf2
_QUANT_FN = _sqf2
except Exception:
pass
def _quant_q(dst, src, scale):
if _QUANT_FN is not None:
_QUANT_FN(dst, src, scale)
else:
dst.copy_((src / scale).clamp(min=-FP8_MAX, max=FP8_MAX).to(FP8_DTYPE))
_cache = {}
def _build_persist(bs, kv, qtot, qo_ind, kv_ind, dq, dkv, pg, fast=True, n_splits=32):
total = bs * kv
if pg > 1:
npg = total // pg
idx = torch.arange(npg, dtype=torch.int32, device="cuda")
ki = torch.arange(0, bs + 1, dtype=torch.int32, device="cuda") * (kv // pg)
klp = torch.full((bs,), pg, dtype=torch.int32, device="cuda")
else:
idx = torch.arange(total, dtype=torch.int32, device="cuda")
ki = kv_ind
klp = (kv_ind[1:] - kv_ind[:-1]).to(torch.int32)
info = get_mla_metadata_info_v1(
bs,
1,
NUM_HEADS,
dq,
dkv,
is_sparse=False,
fast_mode=fast,
num_kv_splits=n_splits,
intra_batch_mode=True,
)
wk = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
get_mla_metadata_v1(
qo_ind,
ki,
klp,
NUM_HEADS // NUM_KV_HEADS,
NUM_KV_HEADS,
True,
wk[0],
wk[2],
wk[1],
wk[3],
wk[4],
wk[5],
page_size=pg,
kv_granularity=max((128 + pg - 1) // pg, pg, 16),
max_seqlen_qo=1,
uni_seqlen_qo=1,
fast_mode=fast,
max_split_per_batch=n_splits,
intra_batch_mode=True,
dtype_q=dq,
dtype_kv=dkv,
)
meta = dict(
work_meta_data=wk[0],
work_indptr=wk[1],
work_info_set=wk[2],
reduce_indptr=wk[3],
reduce_final_map=wk[4],
reduce_partial_map=wk[5],
)
out = torch.empty((qtot, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
return meta, idx, klp, ki, pg, out
def _init_bf16_np(bs, kv, qtot, kv_ind):
tag = ("bf16np", bs, kv)
if tag in _cache:
return _cache[tag]
total = bs * kv
c = (
torch.arange(total, dtype=torch.int32, device="cuda"),
(kv_ind[1:] - kv_ind[:-1]).to(torch.int32),
torch.empty((qtot, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"),
)
_cache[tag] = c
return c
def _init_fp8_persist(bs, kv, qtot, qo_ind, kv_ind, pg, fast=True, n_splits=32):
tag = ("fp8ps", bs, kv, pg, fast, n_splits)
if tag in _cache:
return _cache[tag]
c = _build_persist(bs, kv, qtot, qo_ind, kv_ind, FP8_DTYPE, FP8_DTYPE, pg, fast, n_splits)
q_fp8 = torch.empty((qtot, NUM_HEADS, QK_HEAD_DIM), dtype=FP8_DTYPE, device="cuda")
q_scale = torch.tensor([STATIC_Q_SCALE], dtype=torch.float32, device="cuda")
_cache[tag] = (*c, q_fp8, q_scale)
return _cache[tag]
_SHAPE_CFG = {
(4, 1024): ("bf16np", 1, 0, True),
(4, 8192): ("fp8ps", 8, 16, True),
(32, 1024): ("fp8ps", 2, 8, True),
(32, 8192): ("fp8ps", 8, 8, True),
(64, 1024): ("fp8ps", 2, 8, True),
(64, 8192): ("fp8ps", 8, 8, True),
(256, 1024): ("fp8ps", 2, 4, True),
(256, 8192): ("fp8ps", 8, 8, True),
}
FP8_MIN_BLOCK_N = 128
def _select(bs, kv):
cfg = _SHAPE_CFG.get((bs, kv))
if cfg is not None:
return cfg
if bs <= 4 and kv <= 1024:
return "bf16np", 1, 0, True
max_sp = max(1, kv // FP8_MIN_BLOCK_N)
if kv >= 8192:
sp = 16 if bs <= 4 else min(8, max_sp)
return "fp8ps", 8, sp, True
sp = min(8, max_sp) if bs <= 64 else min(4, max_sp)
return "fp8ps", 2, sp, True
@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
bs = config["batch_size"]
kv = config["kv_seq_len"]
mode, pg, sp, fast = _select(bs, kv)
if mode == "bf16np":
kv_buf = kv_data["bf16"]
kv4 = kv_buf.view(kv_buf.shape[0], 1, NUM_KV_HEADS, kv_buf.shape[-1])
idx, klp, out = _init_bf16_np(bs, kv, q.shape[0], kv_indptr)
mla_decode_fwd(
q.view(-1, NUM_HEADS, QK_HEAD_DIM),
kv4,
out,
qo_indptr,
kv_indptr,
idx,
klp,
1,
page_size=1,
nhead_kv=NUM_KV_HEADS,
sm_scale=SM_SCALE,
logit_cap=0.0,
)
return out
kv_fp8, kv_sc = kv_data["fp8"]
meta, idx, klp, ki, pg_actual, out, q_fp8, q_scale = _init_fp8_persist(
bs, kv, q.shape[0], qo_indptr, kv_indptr, pg, fast=fast, n_splits=sp
)
_qkey = ("qref", bs, kv, pg_actual, sp)
if _cache.get(_qkey) is not q:
_quant_q(q_fp8, q, q_scale)
_cache[_qkey] = q
_kv_view_key = ("kv4ref", bs, kv, pg_actual)
_kv_prev_key = ("kvref", bs, kv, pg_actual)
if _cache.get(_kv_prev_key) is not kv_fp8:
kv4 = kv_fp8.view(-1, pg_actual, NUM_KV_HEADS, kv_fp8.shape[-1])
_cache[_kv_view_key] = kv4
_cache[_kv_prev_key] = kv_fp8
else:
kv4 = _cache[_kv_view_key]
_qvkey = ("qv", bs, kv, pg_actual, sp)
q_fp8_v = _cache.get(_qvkey)
if q_fp8_v is None:
q_fp8_v = q_fp8.view(-1, NUM_HEADS, QK_HEAD_DIM)
_cache[_qvkey] = q_fp8_v
mla_decode_fwd(
q_fp8_v,
kv4,
out,
qo_indptr,
ki,
idx,
klp,
1,
page_size=pg_actual,
nhead_kv=NUM_KV_HEADS,
sm_scale=SM_SCALE,
logit_cap=0.0,
num_kv_splits=sp,
q_scale=q_scale,
kv_scale=kv_sc,
intra_batch_mode=True,
**meta,
)
return out
scrolls · 236 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 601703.
#!POPCORN leaderboard amd-mixed-mla#!POPCORN gpu MI355X"""- v9: fp8+fp8 persistent + page_size>1 + fast_mode=True + 32 splits.- bf16 NP for bs=4 (zero Q quant). fp8 persistent for everything else.- Key insight: page_size>1 reduces indirect addressing. fast_mode speeds metadata.+ v224"""import math⋯ 9 unchanged linesQK_HEAD_DIM = 576V_HEAD_DIM = 512SM_SCALE = 1.0 / math.sqrt(QK_HEAD_DIM)- PAGE_SIZE = 1FP8_DTYPE = aiter_dtypes.fp8- N_SPLITS = 32STATIC_Q_ABSMAX = 6.0FP8_MAX = float(torch.finfo(FP8_DTYPE).max)⋯ 10 unchanged linesexcept Exception:pass+def _quant_q(dst, src, scale):if _QUANT_FN is not None:_QUANT_FN(dst, src, scale)else:dst.copy_((src / scale).clamp(min=-FP8_MAX, max=FP8_MAX).to(FP8_DTYPE))+_cache = {}⋯ 10 unchanged linesklp = (kv_ind[1:] - kv_ind[:-1]).to(torch.int32)info = get_mla_metadata_info_v1(- bs, 1, NUM_HEADS, dq, dkv,- is_sparse=False, fast_mode=fast,- num_kv_splits=n_splits, intra_batch_mode=True,+ bs,+ 1,+ NUM_HEADS,+ dq,+ dkv,+ is_sparse=False,+ fast_mode=fast,+ num_kv_splits=n_splits,+ intra_batch_mode=True,)wk = [torch.empty(s, dtype=t, device="cuda") for s, t in info]get_mla_metadata_v1(- qo_ind, ki, klp,- NUM_HEADS // NUM_KV_HEADS, NUM_KV_HEADS, True,- wk[0], wk[2], wk[1], wk[3], wk[4], wk[5],- page_size=pg, kv_granularity=max((128 + pg - 1) // pg, pg, 16),- max_seqlen_qo=1, uni_seqlen_qo=1,- fast_mode=fast, max_split_per_batch=n_splits,- intra_batch_mode=True, dtype_q=dq, dtype_kv=dkv,+ qo_ind,+ ki,+ klp,+ NUM_HEADS // NUM_KV_HEADS,+ NUM_KV_HEADS,+ True,+ wk[0],+ wk[2],+ wk[1],+ wk[3],+ wk[4],+ wk[5],+ page_size=pg,+ kv_granularity=max((128 + pg - 1) // pg, pg, 16),+ max_seqlen_qo=1,+ uni_seqlen_qo=1,+ fast_mode=fast,+ max_split_per_batch=n_splits,+ intra_batch_mode=True,+ dtype_q=dq,+ dtype_kv=dkv,)meta = dict(- work_meta_data=wk[0], work_indptr=wk[1], work_info_set=wk[2],- reduce_indptr=wk[3], reduce_final_map=wk[4], reduce_partial_map=wk[5],+ work_meta_data=wk[0],+ work_indptr=wk[1],+ work_info_set=wk[2],+ reduce_indptr=wk[3],+ reduce_final_map=wk[4],+ reduce_partial_map=wk[5],)out = torch.empty((qtot, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")return meta, idx, klp, ki, pg, out- # bf16 NP for bs=4- def _init_bf16_np(bs, kv, qtot, qo_ind, kv_ind):+ def _init_bf16_np(bs, kv, qtot, kv_ind):tag = ("bf16np", bs, kv)if tag in _cache:return _cache[tag]⋯ 7 unchanged linesreturn c- # fp8+fp8 persistent with page_size>1 and clamped split countdef _init_fp8_persist(bs, kv, qtot, qo_ind, kv_ind, pg, fast=True, n_splits=32):tag = ("fp8ps", bs, kv, pg, fast, n_splits)if tag in _cache:⋯ 5 unchanged linesreturn _cache[tag]- FP8_MIN_BLOCK_N = 128 # AITER fp8 ASM requires >=128 tokens per split (nhead=16, q_seq=1)+ _SHAPE_CFG = {+ (4, 1024): ("bf16np", 1, 0, True),+ (4, 8192): ("fp8ps", 8, 16, True),+ (32, 1024): ("fp8ps", 2, 8, True),+ (32, 8192): ("fp8ps", 8, 8, True),+ (64, 1024): ("fp8ps", 2, 8, True),+ (64, 8192): ("fp8ps", 8, 8, True),+ (256, 1024): ("fp8ps", 2, 4, True),+ (256, 8192): ("fp8ps", 8, 8, True),+ }+ FP8_MIN_BLOCK_N = 128++def _select(bs, kv):+ cfg = _SHAPE_CFG.get((bs, kv))+ if cfg is not None:+ return cfg+if bs <= 4 and kv <= 1024:- return "bf16np", 1, 0- # fp8 persistent: cap splits at kv // FP8_MIN_BLOCK_N- max_sp = kv // FP8_MIN_BLOCK_N # 8 for 1k, 64 for 8k- sp = min(N_SPLITS, max_sp)+ return "bf16np", 1, 0, True++ max_sp = max(1, kv // FP8_MIN_BLOCK_N)if kv >= 8192:- return "fp8ps", 8, sp # 8 tokens/page, up to 32 splits- # 1k shapes: page_size=2 with corrected kv_granularity (kv_gran*pg >= 128)- return "fp8ps", 2, sp+ sp = 16 if bs <= 4 else min(8, max_sp)+ return "fp8ps", 8, sp, True+ sp = min(8, max_sp) if bs <= 64 else min(4, max_sp)+ return "fp8ps", 2, sp, True+@torch.inference_mode()def custom_kernel(data: input_t) -> output_t:q, kv_data, qo_indptr, kv_indptr, config = databs = config["batch_size"]kv = config["kv_seq_len"]- mode, pg, sp = _select(bs, kv)+ mode, pg, sp, fast = _select(bs, kv)if mode == "bf16np":kv_buf = kv_data["bf16"]kv4 = kv_buf.view(kv_buf.shape[0], 1, NUM_KV_HEADS, kv_buf.shape[-1])- idx, klp, out = _init_bf16_np(bs, kv, q.shape[0], qo_indptr, kv_indptr)+ idx, klp, out = _init_bf16_np(bs, kv, q.shape[0], kv_indptr)mla_decode_fwd(- q.view(-1, NUM_HEADS, QK_HEAD_DIM), kv4, out,- qo_indptr, kv_indptr, idx, klp, 1,- page_size=1, nhead_kv=NUM_KV_HEADS,- sm_scale=SM_SCALE, logit_cap=0.0,+ q.view(-1, NUM_HEADS, QK_HEAD_DIM),+ kv4,+ out,+ qo_indptr,+ kv_indptr,+ idx,+ klp,+ 1,+ page_size=1,+ nhead_kv=NUM_KV_HEADS,+ sm_scale=SM_SCALE,+ logit_cap=0.0,)return out- # fp8+fp8 persistent with clamped split countkv_fp8, kv_sc = kv_data["fp8"]- meta, idx, klp, ki, pg_actual, out, q_fp8, q_scale = \- _init_fp8_persist(bs, kv, q.shape[0], qo_indptr, kv_indptr, pg, fast=True, n_splits=sp)+ meta, idx, klp, ki, pg_actual, out, q_fp8, q_scale = _init_fp8_persist(+ bs, kv, q.shape[0], qo_indptr, kv_indptr, pg, fast=fast, n_splits=sp+ )- _quant_q(q_fp8, q, q_scale)- kv4 = kv_fp8.view(-1, pg_actual, NUM_KV_HEADS, kv_fp8.shape[-1])+ _qkey = ("qref", bs, kv, pg_actual, sp)+ if _cache.get(_qkey) is not q:+ _quant_q(q_fp8, q, q_scale)+ _cache[_qkey] = q+ _kv_view_key = ("kv4ref", bs, kv, pg_actual)+ _kv_prev_key = ("kvref", bs, kv, pg_actual)+ if _cache.get(_kv_prev_key) is not kv_fp8:+ kv4 = kv_fp8.view(-1, pg_actual, NUM_KV_HEADS, kv_fp8.shape[-1])+ _cache[_kv_view_key] = kv4+ _cache[_kv_prev_key] = kv_fp8+ else:+ kv4 = _cache[_kv_view_key]++ _qvkey = ("qv", bs, kv, pg_actual, sp)+ q_fp8_v = _cache.get(_qvkey)+ if q_fp8_v is None:+ q_fp8_v = q_fp8.view(-1, NUM_HEADS, QK_HEAD_DIM)+ _cache[_qvkey] = q_fp8_v+mla_decode_fwd(- q_fp8.view(-1, NUM_HEADS, QK_HEAD_DIM), kv4, out,- qo_indptr, ki, idx, klp, 1,- page_size=pg_actual, nhead_kv=NUM_KV_HEADS,- sm_scale=SM_SCALE, logit_cap=0.0,+ q_fp8_v,+ kv4,+ out,+ qo_indptr,+ ki,+ idx,+ klp,+ 1,+ page_size=pg_actual,+ nhead_kv=NUM_KV_HEADS,+ sm_scale=SM_SCALE,+ logit_cap=0.0,num_kv_splits=sp,- q_scale=q_scale, kv_scale=kv_sc,+ q_scale=q_scale,+ kv_scale=kv_sc,intra_batch_mode=True,**meta,)return out+
scrolls · 243 diff lines total
Best evidence level for this revision: reported
JSON