submission 594605
rikashi · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 103 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-594605?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:63756fe375e51761f6107c53003fc002e7bcc940e90ff4ca84d5d0ee73fd9b98
license declaredunknown
license concludedunknown
authorsrikashi
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
sc=tl.dot(qt,tl.trans(kt),sc,out_dtype=tl.float32)num-warps = 1
qd=dq,vd=dv,BH=16,TK=64,TD=128,NS=ns,num_warps=1,num_stages=1)stages = 1
qd=dq,vd=dv,BH=16,TK=64,TD=128,NS=ns,num_warps=1,num_stages=1)Kernel source
submission.py103 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
import torch, triton, triton.language as tl
from task import input_t, output_t
from aiter.mla import mla_decode_fwd
from aiter import dtypes as aiter_dtypes, get_mla_metadata_info_v1, get_mla_metadata_v1
FP8_DTYPE = aiter_dtypes.fp8
NUM_KV_SPLITS = 16
_cache = {}
@triton.jit
def _decode_split(q_ptr, kv_ptr, pp, lp, qip, kvip, sm,
qd: tl.constexpr, vd: tl.constexpr, BH: tl.constexpr, TK: tl.constexpr,
TD: tl.constexpr, NS: tl.constexpr):
bi=tl.program_id(0); si=tl.program_id(1)
qs=tl.load(qip+bi).to(tl.int32); ks=tl.load(kvip+bi).to(tl.int32)
ke=tl.load(kvip+bi+1).to(tl.int32); kl=ke-ks
ch=(kl+NS-1)//NS; ms=ks+si*ch; me=tl.minimum(ks+(si+1)*ch,ke); ml=me-ms
oh=tl.arange(0,BH); ok=tl.arange(0,TK); ov=tl.arange(0,vd)
mi=tl.full([BH],float("-inf"),dtype=tl.float32)
li=tl.zeros([BH],dtype=tl.float32); acc=tl.zeros([BH,vd],dtype=tl.float32)
for ko in range(0,ml,TK):
tm=(ko+ok)<ml; kb=ms+ko
sc=tl.zeros([BH,TK],dtype=tl.float32)
for do in range(0,qd,TD):
di=do+tl.arange(0,TD)
qt=tl.load(q_ptr+qs*BH*qd+oh[:,None]*qd+di[None,:],mask=di[None,:]<qd,other=0.0).to(tl.float16)
kt=tl.load(kv_ptr+(kb+ok[:,None])*qd+di[None,:],mask=tm[:,None]&(di[None,:]<qd),other=0.0).to(tl.float16)
sc=tl.dot(qt,tl.trans(kt),sc,out_dtype=tl.float32)
sc=sc*sm; sc=tl.where(tm[None,:],sc,float("-inf"))
mn=tl.maximum(mi,tl.max(sc,axis=1))
al=tl.math.exp2((mi-mn)*1.44269504); p=tl.math.exp2((sc-mn[:,None])*1.44269504)
vt=tl.load(kv_ptr+(kb+ok[:,None])*qd+ov[None,:],mask=tm[:,None]&(ov[None,:]<vd),other=0.0).to(tl.float16)
acc=acc*al[:,None]; acc=tl.dot(p.to(tl.float16),vt,acc,out_dtype=tl.float32)
li=li*al+tl.sum(p,axis=1); mi=mn
out=acc/(li[:,None]+1e-12); base=(bi*NS+si)*BH
tl.store(pp+(base+oh[:,None])*vd+ov[None,:],out)
tl.store(lp+base+oh,mi+tl.log(li+1e-12))
@triton.jit
def _reduce(pp,lp,op,NS: tl.constexpr,BH: tl.constexpr,vd: tl.constexpr):
bi=tl.program_id(0); hi=tl.program_id(1); ov=tl.arange(0,vd)
gm=float("-inf")
for s in range(NS): gm=tl.maximum(gm,tl.load(lp+(bi*NS+s)*BH+hi))
acc=tl.zeros([vd],dtype=tl.float32); tw=0.0
for s in range(NS):
lse=tl.load(lp+(bi*NS+s)*BH+hi); w=tl.math.exp2((lse-gm)*1.44269504); tw+=w
acc+=w*tl.load(pp+((bi*NS+s)*BH+hi)*vd+ov)
tl.store(op+(bi*BH+hi)*vd+ov,(acc/tw).to(tl.bfloat16))
def custom_kernel(data: input_t) -> output_t:
q,kv_data,qo_indptr,kv_indptr,config=data
bs=config["batch_size"]; nq=config["num_heads"]; nkv=config["num_kv_heads"]
dq=config["qk_head_dim"]; dv=config["v_head_dim"]
kvl=config["kv_seq_len"]; sm=config["sm_scale"]; tkv=bs*kvl
ps=2 if kvl%2==0 else 1
if bs<=4 and kvl<=2048:
ns=16; kv=kv_data["bf16"].view(tkv,dq).contiguous()
key=("t",bs,kvl)
if key not in _cache:
_cache[key]={"o":torch.empty((bs,nq,dv),dtype=torch.bfloat16,device="cuda"),
"p":torch.empty((bs*ns*nq,dv),dtype=torch.float32,device="cuda"),
"l":torch.empty((bs*ns*nq,),dtype=torch.float32,device="cuda")}
c=_cache[key]
_decode_split[(bs,ns)](q.view(bs,nq*dq),kv,c["p"],c["l"],qo_indptr,kv_indptr,sm,
qd=dq,vd=dv,BH=16,TK=64,TD=128,NS=ns,num_warps=1,num_stages=1)
_reduce[(bs,nq)](c["p"],c["l"],c["o"].view(bs*nq,dv),NS=ns,BH=16,vd=dv)
return c["o"]
ua=(bs>=64 and kvl<=2048); uf=(not ua)and(kvl>=8192)
npg=tkv//ps
kvip=kv_indptr//ps if ps>1 else kv_indptr
if ua:
kb,ks=kv_data["fp8"]; k4=kb.view(npg,ps,nkv,dq)
qu=q.view(-1,nq,dq); qsc=None; kvsc=ks; qdt=q.dtype; kvdt=FP8_DTYPE
elif uf:
kb,ks=kv_data["fp8"]; k4=kb.view(npg,ps,nkv,dq)
qu=q.to(FP8_DTYPE).view(-1,nq,dq); qsc=torch.ones(1,dtype=torch.float32,device="cuda")
kvsc=ks; qdt=FP8_DTYPE; kvdt=FP8_DTYPE
else:
k4=kv_data["bf16"].view(npg,ps,nkv,dq)
qu=q.view(-1,nq,dq); qsc=None; kvsc=None; qdt=q.dtype; kvdt=kv_data["bf16"].dtype
key=("a",bs,kvl,ua,uf,ps)
if key not in _cache:
ki=torch.arange(npg,dtype=torch.int32,device="cuda")
kl_t=kv_indptr[1:]-kv_indptr[:-1]; kl=((kl_t-1)%ps+1).to(torch.int32) if ps>1 else kl_t.to(torch.int32)
info=get_mla_metadata_info_v1(bs,1,nq,qdt,kvdt,is_sparse=False,fast_mode=False,
num_kv_splits=NUM_KV_SPLITS,intra_batch_mode=True)
w=[torch.empty(s,dtype=t,device="cuda") for s,t in info]
get_mla_metadata_v1(qo_indptr,kvip,kl,nq//nkv,nkv,True,
w[0],w[2],w[1],w[3],w[4],w[5],page_size=ps,
kv_granularity=max(ps,16),max_seqlen_qo=1,uni_seqlen_qo=1,
fast_mode=False,max_split_per_batch=NUM_KV_SPLITS,intra_batch_mode=True,
dtype_q=qdt,dtype_kv=kvdt)
_cache[key]={"i":ki,"l":kl,"w":w,"kvip":kvip}
sc=_cache[key]
o=torch.empty((q.shape[0],nq,dv),dtype=torch.bfloat16,device="cuda")
mla_decode_fwd(qu,k4,o,qo_indptr,sc["kvip"],sc["i"],sc["l"],1,
page_size=ps,nhead_kv=nkv,sm_scale=sm,logit_cap=0.0,
num_kv_splits=NUM_KV_SPLITS,q_scale=qsc,kv_scale=kvsc,intra_batch_mode=True,
work_meta_data=sc["w"][0],work_indptr=sc["w"][1],work_info_set=sc["w"][2],
reduce_indptr=sc["w"][3],reduce_final_map=sc["w"][4],reduce_partial_map=sc["w"][5])
return o
scrolls · 103 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Best evidence level for this revision: reported
JSON