Skip to content
KernelIndex
Search⌘K

submission 751073

mars-compute · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 54 lines, June 9 Researcher Reciprocity License v1.0.

submission_v2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-751073?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
AMD Instinct MI355X
32.6µs
#32 of 766
2026-04-06

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:6d652dcfc76c51bfb03081356a80f2bab89c9106a6627328e167f1bb1019b634
license declaredunknown
license concludedunknown
authorsmars-compute
imported2026-08-15

Kernel source

submission_v2.py54 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""V409 — V400 configs + kf.view cache + last_fk. NO Q cache (causes fail).
Q cache is unsafe: benchmark reuses same tensor with different data."""
import os
os.environ["HIP_FORCE_DEV_KERNARG"]="1";os.environ["HSA_ENABLE_SDMA"]="0"
import 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
NH=16;NKH=1;QKD=576;VD=512;SM=1.0/(QKD**0.5);FP8=aiter_dtypes.fp8
_s1=aiter.mla_decode_stage1_asm_fwd
_rd=aiter.mla_reduce_v1
_qs=torch.ones(1,dtype=torch.float32,device="cuda")
_N=None
_CFG={
    (4,1024):(1,16),(4,8192):(8,10),
    (32,1024):(2,2),(32,8192):(8,10),
    (64,1024):(2,4),(64,8192):(8,2),
    (256,1024):(2,2),(256,8192):(8,1),
}
def _build(bs,kvs,qd,kd,ps,nks):
    npp=kvs//ps
    qo=torch.arange(0,bs+1,dtype=torch.int32,device="cuda")
    kvi=torch.arange(0,bs+1,dtype=torch.int32,device="cuda")*npp
    klp=torch.full((bs,),ps,dtype=torch.int32,device="cuda")
    ki=torch.arange(bs*npp,dtype=torch.int32,device="cuda")
    info=get_mla_metadata_info_v1(bs,1,NH,qd,kd,is_sparse=False,fast_mode=False,num_kv_splits=nks,intra_batch_mode=True)
    w=[torch.empty(s,dtype=t,device="cuda") for s,t in info]
    wm,wi,wis,ri,rfm,rpm=w
    get_mla_metadata_v1(qo,kvi,klp,NH,NKH,True,wm,wis,wi,ri,rfm,rpm,
        page_size=ps,kv_granularity=max(ps,16),max_seqlen_qo=1,uni_seqlen_qo=1,
        fast_mode=False,max_split_per_batch=nks,intra_batch_mode=True,dtype_q=qd,dtype_kv=kd)
    lg=torch.empty((rpm.size(0),1,NH,VD),dtype=torch.float32,device="cuda")
    al=torch.empty((rpm.size(0),1,NH,1),dtype=torch.float32,device="cuda")
    o=torch.empty((bs,NH,VD),dtype=torch.bfloat16,device="cuda")
    return (wm,wi,wis,ri,rfm,rpm,ki,klp,qo,kvi,lg,al,o,ps)
_cache={}
def custom_kernel(data:input_t)->output_t:
    q,kv_data,_,_2,cfg=data
    bs=cfg["batch_size"];kvs=cfg["kv_seq_len"]
    kf,ks=kv_data["fp8"]
    qf=q.to(FP8)
    ps,nks=_CFG.get((bs,kvs),(2,32))
    key=(bs,kvs)
    if key not in _cache:_cache[key]=_build(bs,kvs,qf.dtype,kf.dtype,ps,nks)
    wm,wi,wis,ri,rfm,rpm,ki,klp,qo,kvi,lg,al,o,ps_=_cache[key]
    np_=kf.shape[0]//ps_;kv4d=kf.view(np_,ps_,NKH,QKD)
    _s1(qf.view(-1,NH,QKD),kv4d,qo,kvi,ki,klp,_N,wm,wi,wis,1,ps_,NKH,SM,lg,al,o,_qs,ks)
    _rd(lg,al,ri,rfm,rpm,1,o,_N)
    return o

scrolls · 54 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 737583.

+ #!POPCORN leaderboard amd-mixed-mla
+ #!POPCORN gpu MI355X
+ """V409 — V400 configs + kf.view cache + last_fk. NO Q cache (causes fail).
+ Q cache is unsafe: benchmark reuses same tensor with different data."""
import os
os.environ["HIP_FORCE_DEV_KERNARG"]="1";os.environ["HSA_ENABLE_SDMA"]="0"
import torch
⋯ 2 unchanged lines
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
NH=16;NKH=1;QKD=576;VD=512;SM=1.0/(QKD**0.5);FP8=aiter_dtypes.fp8
- _stage1=aiter.mla_decode_stage1_asm_fwd
- _reduce=aiter.mla_reduce_v1
- _qs=None;_cache={}
- # kv=1024: V188 proven configs; kv=8192: tuned ps=8 configs
+ _s1=aiter.mla_decode_stage1_asm_fwd
+ _rd=aiter.mla_reduce_v1
+ _qs=torch.ones(1,dtype=torch.float32,device="cuda")
+ _N=None
_CFG={
(4,1024):(1,16),(4,8192):(8,10),
(32,1024):(2,2),(32,8192):(8,10),
⋯ 5 unchanged lines
qo=torch.arange(0,bs+1,dtype=torch.int32,device="cuda")
kvi=torch.arange(0,bs+1,dtype=torch.int32,device="cuda")*npp
klp=torch.full((bs,),ps,dtype=torch.int32,device="cuda")
- kvidx=torch.arange(bs*npp,dtype=torch.int32,device="cuda")
+ ki=torch.arange(bs*npp,dtype=torch.int32,device="cuda")
info=get_mla_metadata_info_v1(bs,1,NH,qd,kd,is_sparse=False,fast_mode=False,num_kv_splits=nks,intra_batch_mode=True)
w=[torch.empty(s,dtype=t,device="cuda") for s,t in info]
wm,wi,wis,ri,rfm,rpm=w
⋯ 3 unchanged lines
lg=torch.empty((rpm.size(0),1,NH,VD),dtype=torch.float32,device="cuda")
al=torch.empty((rpm.size(0),1,NH,1),dtype=torch.float32,device="cuda")
o=torch.empty((bs,NH,VD),dtype=torch.bfloat16,device="cuda")
- return (wm,wi,wis,ri,rfm,rpm,kvidx,klp,qo,kvi,lg,al,o,ps)
+ return (wm,wi,wis,ri,rfm,rpm,ki,klp,qo,kvi,lg,al,o,ps)
+ _cache={}
def custom_kernel(data:input_t)->output_t:
- global _qs
q,kv_data,_,_2,cfg=data
bs=cfg["batch_size"];kvs=cfg["kv_seq_len"]
- kf,ks=kv_data["fp8"];qf=q.to(FP8)
- if _qs is None:_qs=torch.ones(1,dtype=torch.float32,device="cuda")
+ kf,ks=kv_data["fp8"]
+ qf=q.to(FP8)
ps,nks=_CFG.get((bs,kvs),(2,32))
key=(bs,kvs)
if key not in _cache:_cache[key]=_build(bs,kvs,qf.dtype,kf.dtype,ps,nks)
- wm,wi,wis,ri,rfm,rpm,kvidx,klp,qo,kvi,lg,al,o,ps_=_cache[key]
+ wm,wi,wis,ri,rfm,rpm,ki,klp,qo,kvi,lg,al,o,ps_=_cache[key]
np_=kf.shape[0]//ps_;kv4d=kf.view(np_,ps_,NKH,QKD)
- _stage1(qf.view(-1,NH,QKD),kv4d,qo,kvi,kvidx,klp,None,wm,wi,wis,1,ps_,NKH,SM,lg,al,o,_qs,ks)
- _reduce(lg,al,ri,rfm,rpm,1,o,None)
+ _s1(qf.view(-1,NH,QKD),kv4d,qo,kvi,ki,klp,_N,wm,wi,wis,1,ps_,NKH,SM,lg,al,o,_qs,ks)
+ _rd(lg,al,ri,rfm,rpm,1,o,_N)
return o
scrolls · 58 diff lines total

Best evidence level for this revision: reported

JSON