Skip to content
KernelIndex
Search⌘K

submission 673003

josusanmartin · 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.

v027m.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-673003?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
23.5µs
#7 of 766
2026-03-30

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:3afcdd8ea0713e14035a72630520dc42e54fbe5934f540e51f1c96888f20b55f
license declaredunknown
license concludedunknown
authorsjosusanmartin
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

persistent-kernel"""V027m: v027l but push 256x8K to ps=128 and try 256x1K persistent intra=True kg=2.

Kernel source

v027m.py103 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""V027m: v027l but push 256x8K to ps=128 and try 256x1K persistent intra=True kg=2.
256x8K: ps=128 (was ps=64 at 26.4us, should drop to ~18us like other 8K)
256x1K: try persistent intra=True split=1 kg=2 (smaller granularity)"""
import os
os.environ.setdefault("HIP_FORCE_DEV_KERNARG", "1")
os.environ.setdefault("AMD_DIRECT_DISPATCH", "1")
os.environ.setdefault("HIPBLASLT_ALLOW_FLUSH_DENORM", "1")
os.environ.setdefault("GPU_MAX_HW_QUEUES", "2")
import torch, triton, triton.language as tl
from task import input_t, output_t
import aiter
from aiter import dtypes as aiter_dtypes, mla as aiter_mla
try:
    from aiter.jit.module_quant import static_per_tensor_quant as _quant
except Exception:
    from aiter.ops.quant import static_per_tensor_quant as _quant
try:
    from aiter.jit.module_mla_asm import mla_decode_stage1_asm_fwd as _s1
except Exception:
    _s1 = aiter.mla_decode_stage1_asm_fwd
try:
    from aiter.jit.module_mla_reduce import mla_reduce_v1 as _rd
except Exception:
    _rd = aiter.mla_reduce_v1
try:
    from aiter.jit.module_mla_metadata import get_mla_metadata_v1 as _mv1
except Exception:
    _mv1 = aiter.get_mla_metadata_v1
_mi = aiter.get_mla_metadata_info_v1
_fwd = aiter.mla.mla_decode_fwd
NH=16;QKD=576;VD=512;SM=float(1.0/(QKD**0.5));FP8=aiter_dtypes.fp8;BF16=torch.bfloat16
QSV=float(2.0/torch.finfo(FP8).max);_st={};_nc={};_pc={};_s2c={}
def _gs(d):
    t=_st.get(0)
    if t is None: t=torch.tensor([QSV],dtype=torch.float32,device=d);_st[0]=t
    return t
def _gnc(dev,bs,kvl,ps):
    key=(bs,kvl,ps);c=_nc.get(key)
    if c: return c
    ppb=kvl//ps
    ki=torch.arange(bs*ppb,dtype=torch.int32,device=dev)
    kl=torch.full((bs,),ps,dtype=torch.int32,device=dev)
    kip=torch.arange(bs+1,dtype=torch.int32,device=dev)*ppb
    out=torch.empty((bs,NH,VD),dtype=BF16,device=dev)
    qf=torch.empty((bs,NH,QKD),dtype=FP8,device=dev)
    c=(ki,kl,kip,out,qf);_nc[key]=c;return c
def _gpc(dev,qo,bs,kvl,ps,intra=True,split_override=None,kg_override=None):
    key=(bs,kvl,ps,intra,split_override,kg_override);c=_pc.get(key)
    if c: return c
    tot=bs*kvl;ppb=kvl//ps
    ns = split_override if split_override is not None else aiter_mla.get_meta_param(None,bs,tot,NH,1,FP8)[0]
    ki=torch.arange(bs*ppb,dtype=torch.int32,device=dev);kl=torch.full((bs,),ps,dtype=torch.int32,device=dev)
    kip=torch.arange(bs+1,dtype=torch.int32,device=dev)*ppb;out=torch.empty((bs,NH,VD),dtype=BF16,device=dev)
    kg=kg_override if kg_override is not None else max(ps,16)
    info=_mi(bs,1,NH,FP8,FP8,is_sparse=False,fast_mode=True,num_kv_splits=ns,intra_batch_mode=intra)
    w=[torch.empty(s,dtype=t,device=dev) for s,t in info]
    _mv1(qo,kip,kl,16,1,False,w[0],w[2],w[1],w[3],w[4],w[5],page_size=ps,kv_granularity=kg,max_seqlen_qo=1,uni_seqlen_qo=1,fast_mode=True,max_split_per_batch=ns,intra_batch_mode=intra,dtype_q=FP8,dtype_kv=FP8)
    qf=torch.empty((bs,NH,QKD),dtype=FP8,device=dev)
    pt=int(w[5].numel());po=torch.empty((pt,1,NH,VD),dtype=torch.float32,device=dev);pl=torch.empty((pt,1,NH,1),dtype=torch.float32,device=dev)
    c=(ki,kl,kip,out,w[0],w[1],w[2],w[3],w[4],w[5],ns,qf,po,pl);_pc[key]=c;return c
def _g_np(dev,bs,kvl,ps):
    key=(bs,kvl,ps,"npm");c=_s2c.get(key)
    if c: return c
    ppb=kvl//ps
    ki=torch.arange(bs*ppb,dtype=torch.int32,device=dev)
    kl=torch.full((bs,),ps,dtype=torch.int32,device=dev)
    kip=torch.arange(bs+1,dtype=torch.int32,device=dev)*ppb
    out=torch.empty((bs,NH,VD),dtype=BF16,device=dev)
    qf=torch.empty((bs,NH,QKD),dtype=FP8,device=dev)
    nsi=torch.arange(bs+1,dtype=torch.int32,device=dev)
    sd=torch.empty((bs,1,NH,VD),dtype=torch.float32,device=dev)
    sl=torch.empty((bs,1,NH,1),dtype=torch.float32,device=dev)
    c=(ki,kl,kip,out,qf,nsi,sd,sl);_s2c[key]=c;return c
def custom_kernel(data: input_t) -> output_t:
    q,kv_data,qo_indptr,kv_indptr,config=data
    bs=int(config["batch_size"]);kvl=int(config["kv_seq_len"]);dev=q.device
    kv_fp8,kv_scale=kv_data["fp8"];qs=_gs(dev)
    if kvl == 1024 and bs <= 32:
        ps=2;kb=kv_fp8.view(-1,ps,1,QKD);c=_gnc(dev,bs,kvl,ps);_quant(c[4],q,qs)
        _fwd(c[4],kb,c[3],qo_indptr,c[2],c[0],c[1],1,ps,1,SM,q_scale=qs,kv_scale=kv_scale,intra_batch_mode=True)
        return c[3]
    if bs == 64 and kvl == 1024:
        ps=2;kb=kv_fp8.view(-1,ps,1,QKD)
        c=_gpc(dev,qo_indptr,bs,kvl,ps,intra=True,split_override=1,kg_override=4)
        _quant(c[11],q,qs)
        _s1(c[11],kb,qo_indptr,c[2],c[0],c[1],None,c[4],c[5],c[6],1,ps,1,SM,c[12],c[13],c[3],qs,kv_scale)
        if c[10]>1: _rd(c[12],c[13],c[7],c[8],c[9],1,c[3],None)
        return c[3]
    if bs == 256 and kvl == 1024:
        ps=2;kb=kv_fp8.view(-1,ps,1,QKD)
        c=_gpc(dev,qo_indptr,bs,kvl,ps,intra=True,split_override=1,kg_override=2)
        _quant(c[11],q,qs)
        _s1(c[11],kb,qo_indptr,c[2],c[0],c[1],None,c[4],c[5],c[6],1,ps,1,SM,c[12],c[13],c[3],qs,kv_scale)
        if c[10]>1: _rd(c[12],c[13],c[7],c[8],c[9],1,c[3],None)
        return c[3]
    # ALL 8K: NP ns=1 ps=128 for all shapes
    ps=128;kb=kv_fp8.view(-1,ps,1,QKD);c=_g_np(dev,bs,kvl,ps);_quant(c[4],q,qs)
    c[3].zero_(); c[7].fill_(-float("inf"))
    _s1(c[4],kb,qo_indptr,c[2],c[0],c[1],c[5],None,None,None,1,ps,1,SM,c[6],c[7],c[3],qs,kv_scale)
    return c[3]
scrolls · 103 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 670399.

#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
- """V027i: v027h + add 64x8K NP ns=1 ps=128.
- v027h passed ranked with 4+32x8K NP ps=128. Now add 64x8K too.
- 64x8K was 33.9us persistent in v027h. NP ps=128 should be ~18us.
- v027b (all 8K NP) failed ranked — isolating whether 64x8K is the culprit."""
+ """V027m: v027l but push 256x8K to ps=128 and try 256x1K persistent intra=True kg=2.
+ 256x8K: ps=128 (was ps=64 at 26.4us, should drop to ~18us like other 8K)
+ 256x1K: try persistent intra=True split=1 kg=2 (smaller granularity)"""
import os
os.environ.setdefault("HIP_FORCE_DEV_KERNARG", "1")
os.environ.setdefault("AMD_DIRECT_DISPATCH", "1")
⋯ 52 unchanged lines
pt=int(w[5].numel());po=torch.empty((pt,1,NH,VD),dtype=torch.float32,device=dev);pl=torch.empty((pt,1,NH,1),dtype=torch.float32,device=dev)
c=(ki,kl,kip,out,w[0],w[1],w[2],w[3],w[4],w[5],ns,qf,po,pl);_pc[key]=c;return c
def _g_np(dev,bs,kvl,ps):
- key=(bs,kvl,ps,"npi");c=_s2c.get(key)
+ key=(bs,kvl,ps,"npm");c=_s2c.get(key)
if c: return c
ppb=kvl//ps
ki=torch.arange(bs*ppb,dtype=torch.int32,device=dev)
⋯ 5 unchanged lines
sd=torch.empty((bs,1,NH,VD),dtype=torch.float32,device=dev)
sl=torch.empty((bs,1,NH,1),dtype=torch.float32,device=dev)
c=(ki,kl,kip,out,qf,nsi,sd,sl);_s2c[key]=c;return c
- _CFG = {
- (64,1024): (2, True, 2, None),
- (256,1024):(2, False, 1, 8),
- }
def custom_kernel(data: input_t) -> output_t:
q,kv_data,qo_indptr,kv_indptr,config=data
bs=int(config["batch_size"]);kvl=int(config["kv_seq_len"]);dev=q.device
kv_fp8,kv_scale=kv_data["fp8"];qs=_gs(dev)
- # 1K bs<=32: NP via _fwd ps=2
if kvl == 1024 and bs <= 32:
ps=2;kb=kv_fp8.view(-1,ps,1,QKD);c=_gnc(dev,bs,kvl,ps);_quant(c[4],q,qs)
_fwd(c[4],kb,c[3],qo_indptr,c[2],c[0],c[1],1,ps,1,SM,q_scale=qs,kv_scale=kv_scale,intra_batch_mode=True)
return c[3]
- # ALL 8K shapes: NP ns=1 large ps
- if kvl == 8192:
- ps = 64 if bs == 256 else 128
- kb=kv_fp8.view(-1,ps,1,QKD);c=_g_np(dev,bs,kvl,ps);_quant(c[4],q,qs)
- c[3].zero_(); c[7].fill_(-float("inf"))
- _s1(c[4],kb,qo_indptr,c[2],c[0],c[1],c[5],None,None,None,1,ps,1,SM,c[6],c[7],c[3],qs,kv_scale)
+ if bs == 64 and kvl == 1024:
+ ps=2;kb=kv_fp8.view(-1,ps,1,QKD)
+ c=_gpc(dev,qo_indptr,bs,kvl,ps,intra=True,split_override=1,kg_override=4)
+ _quant(c[11],q,qs)
+ _s1(c[11],kb,qo_indptr,c[2],c[0],c[1],None,c[4],c[5],c[6],1,ps,1,SM,c[12],c[13],c[3],qs,kv_scale)
+ if c[10]>1: _rd(c[12],c[13],c[7],c[8],c[9],1,c[3],None)
return c[3]
- # 64x1K, 256x1K: persistent
- ps,intra,split_ov,kg_ov = _CFG.get((bs,kvl), (2, True, None, None))
- kb=kv_fp8.view(-1,ps,1,QKD);c=_gpc(dev,qo_indptr,bs,kvl,ps,intra=intra,split_override=split_ov,kg_override=kg_ov);_quant(c[11],q,qs)
- _s1(c[11],kb,qo_indptr,c[2],c[0],c[1],None,c[4],c[5],c[6],1,ps,1,SM,c[12],c[13],c[3],qs,kv_scale)
- if c[10]>1: _rd(c[12],c[13],c[7],c[8],c[9],1,c[3],None)
+ if bs == 256 and kvl == 1024:
+ ps=2;kb=kv_fp8.view(-1,ps,1,QKD)
+ c=_gpc(dev,qo_indptr,bs,kvl,ps,intra=True,split_override=1,kg_override=2)
+ _quant(c[11],q,qs)
+ _s1(c[11],kb,qo_indptr,c[2],c[0],c[1],None,c[4],c[5],c[6],1,ps,1,SM,c[12],c[13],c[3],qs,kv_scale)
+ if c[10]>1: _rd(c[12],c[13],c[7],c[8],c[9],1,c[3],None)
+ return c[3]
+ # ALL 8K: NP ns=1 ps=128 for all shapes
+ ps=128;kb=kv_fp8.view(-1,ps,1,QKD);c=_g_np(dev,bs,kvl,ps);_quant(c[4],q,qs)
+ c[3].zero_(); c[7].fill_(-float("inf"))
+ _s1(c[4],kb,qo_indptr,c[2],c[0],c[1],c[5],None,None,None,1,ps,1,SM,c[6],c[7],c[3],qs,kv_scale)
return c[3]
scrolls · 68 diff lines total

Best evidence level for this revision: reported

JSON