submission 745253
augustus2024 · 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.
submission_v409.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-745253?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:c3c5a512c4a12bab0f97aee374f595990c39296a055cf3eba87159eca0342147
license declaredunknown
license concludedunknown
authorsaugustus2024
imported2026-08-15
Kernel source
submission_v409.py53 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 · 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 714518.
- """V254 — MI355X optimal: V188 kv=1024 + MI355X-validated kv=8192 nks.- Key: (4,8192) stays ps=4 (ps=2 slower on MI355X despite MI350X being faster).- Only change from V188: kv=8192 nks values from sweep."""+ #!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 osos.environ["HIP_FORCE_DEV_KERNARG"]="1";os.environ["HSA_ENABLE_SDMA"]="0"import torch⋯ 2 unchanged linesfrom aiter import dtypes as aiter_dtypesfrom aiter import get_mla_metadata_info_v1, get_mla_metadata_v1NH=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={}- # V188 kv=1024 (leaderboard proven) + kv=8192 sweep optimal nks- # (4,8192): ps=4 NOT ps=2 (MI355X: ps=4=24µs, ps=2=30µs)+ _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):(4,16), # V188 original- (32,1024):(2,2),(32,8192):(4,6), # nks=6 (was 8)- (64,1024):(2,4),(64,8192):(4,3), # nks=3 (was 8)- (256,1024):(2,2),(256,8192):(4,1), # nks=1 (was 4)+ (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//psqo=torch.arange(0,bs+1,dtype=torch.int32,device="cuda")kvi=torch.arange(0,bs+1,dtype=torch.int32,device="cuda")*nppklp=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 lineslg=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 _qsq,kv_data,_,_2,cfg=databs=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 · 69 diff lines total
Best evidence level for this revision: reported
JSON