Skip to content
KernelIndex
Search⌘K

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
AMD Instinct MI355X
54.3µs
#182 of 766
2026-03-20

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.

mmasc=tl.dot(qt,tl.trans(kt),sc,out_dtype=tl.float32)
num-warps = 1qd=dq,vd=dv,BH=16,TK=64,TD=128,NS=ns,num_warps=1,num_stages=1)
stages = 1qd=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