submission 670399
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.
v027i.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-670399?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:28e962dcd950bc759603201798e4168c2a85c9ca12d8d0145230e8975f5b6c2f
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
64x8K was 33.9us persistent in v027h. NP ps=128 should be ~18us.Kernel source
v027i.py103 lines
#!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."""
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,"npi");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
_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)
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)
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 669133.
#!POPCORN leaderboard amd-mixed-mla#!POPCORN gpu MI355X- """V026o: SAFE version of v026l. 256x8K NP ns=1 ps=64 (fastest!).- 64x8K reverted to PERSISTENT (NP ns=2 failed ranked at seed 1360).- 32x8K: NP ns=3 + Triton reduce (frozen from v025h)."""+ """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."""import osos.environ.setdefault("HIP_FORCE_DEV_KERNARG", "1")os.environ.setdefault("AMD_DIRECT_DISPATCH", "1")⋯ 8 unchanged linesexcept Exception:from aiter.ops.quant import static_per_tensor_quant as _quanttry:- from aiter.jit.module_mla_metadata import get_mla_metadata_v1 as _mv1- except Exception:- _mv1 = aiter.get_mla_metadata_v1- try:from aiter.jit.module_mla_asm import mla_decode_stage1_asm_fwd as _s1except Exception:_s1 = aiter.mla_decode_stage1_asm_fwd⋯ 1 unchanged linesfrom aiter.jit.module_mla_reduce import mla_reduce_v1 as _rdexcept 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- @triton.jit- def _safe_reduce_branchless(- Mid_O, Mid_lse, O, qo_indptr, kv_indptr, num_kv_splits_indptr,- stride_mid_ob: tl.int64, stride_mid_oh: tl.int64, stride_mid_os: tl.int64,- stride_obs: tl.int64, stride_oh: tl.int64,- MAYBE_FINAL_OUT: tl.constexpr, BATCH_NUM: tl.constexpr,- BLOCK_DV: tl.constexpr, Lv: tl.constexpr, mgc: tl.constexpr,- ):- cur_batch = tl.program_id(0); cur_head = tl.program_id(1)- cur_qo_start = tl.load(qo_indptr + cur_batch); cur_qo_end = tl.load(qo_indptr + cur_batch + 1)- cur_split_start = tl.load(num_kv_splits_indptr + cur_batch); cur_split_end = tl.load(num_kv_splits_indptr + cur_batch + 1)- num_max_kv_splits = tl.load(num_kv_splits_indptr + BATCH_NUM)- cur_kv_seq_len = tl.load(kv_indptr + cur_batch + 1) - tl.load(kv_indptr + cur_batch)- offs_d = tl.arange(0, BLOCK_DV); mask_d = offs_d < Lv- offs_logic = cur_qo_start * stride_mid_ob + cur_head * stride_mid_oh- offs_v = offs_logic * Lv + offs_d- num_valid_kv_splits = tl.minimum(cur_split_end - cur_split_start, tl.cdiv(cur_kv_seq_len, mgc))- final_out = MAYBE_FINAL_OUT and num_max_kv_splits == BATCH_NUM- for cur_qo in range(cur_qo_start, cur_qo_end):- if final_out:- input_ptr = Mid_O.to(tl.pointer_type(O.type.element_ty))- out = tl.load(input_ptr + Lv * (cur_qo * stride_mid_os + cur_head * stride_mid_oh) + offs_d, mask=mask_d, other=0.0)- tl.store(O + cur_qo * stride_obs + cur_head * stride_oh + offs_d, out, mask=mask_d)- else:- e_sum = 0.0; e_max = -float("inf"); acc = tl.zeros((BLOCK_DV,), dtype=tl.float32)- for split_kv_id in range(0, num_valid_kv_splits):- tv = tl.load(Mid_O + offs_v + split_kv_id * stride_mid_os * Lv, mask=mask_d, other=0.0)- tlogic = tl.load(Mid_lse + offs_logic + split_kv_id * stride_mid_os)- tlogic = tl.where(tlogic == tlogic, tlogic, -1e30); tlogic = tl.minimum(tlogic, 1e30)- n_e_max = tl.maximum(tlogic, e_max); old_scale = tl.exp(e_max - n_e_max)- acc *= old_scale; exp_logic = tl.exp(tlogic - n_e_max)- acc += exp_logic * tv; e_sum = e_sum * old_scale + exp_logic; e_max = n_e_max- offs_logic += stride_mid_ob; offs_v += stride_mid_ob * Lv- tl.store(O + cur_qo * stride_obs + cur_head * stride_oh + offs_d, acc / e_sum, mask=mask_d)- aiter.mla._fwd_kernel_stage2_asm = _safe_reduce_branchless_fwd = aiter.mla.mla_decode_fwdNH=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);MGC=64;_st={};_pc={};_nc={};_s2c={}+ 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⋯ 1 unchanged linesdef _gnc(dev,bs,kvl,ps):key=(bs,kvl,ps);c=_nc.get(key)if c: return c- ppb=kvl//ps;np_=bs*kvl//ps- ki=torch.arange(np_,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+ 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 cdef _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;np_=tot//ps+ tot=bs*kvl;ppb=kvl//psns = split_override if split_override is not None else aiter_mla.get_meta_param(None,bs,tot,NH,1,FP8)[0]- ki=torch.arange(np_,dtype=torch.int32,device=dev);kl=torch.full((bs,),ps,dtype=torch.int32,device=dev)+ 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)⋯ 2 unchanged linesqf=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 _uniform_nsi(dev,bs,ns):- return torch.arange(bs+1,dtype=torch.int32,device=dev) * ns- def _g32s2(dev,bs,kvl,ps):- key=(bs,kvl,ps,"32split3");c=_s2c.get(key)+ def _g_np(dev,bs,kvl,ps):+ key=(bs,kvl,ps,"npi");c=_s2c.get(key)if c: return c- ppb=kvl//ps;np_=bs*kvl//ps;ns=3- ki=torch.arange(np_,dtype=torch.int32,device=dev)+ 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)*ppbout=torch.empty((bs,NH,VD),dtype=BF16,device=dev)qf=torch.empty((bs,NH,QKD),dtype=FP8,device=dev)- nsi=_uniform_nsi(dev,bs,ns)- logits=torch.empty((bs,ns,NH,VD),dtype=torch.float32,device=dev)- lse=torch.empty((bs,ns,NH,1),dtype=torch.float32,device=dev)- c=(ki,kl,kip,out,qf,ns,nsi,logits,lse);_s2c[key]=c;return c- def _g_np_direct(dev,bs,kvl,ps):- key=(bs,kvl,ps,"np_direct");c=_s2c.get(key)- if c: return c- ppb=kvl//ps;np_=bs*kvl//ps;ns=1- ki=torch.arange(np_,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=_uniform_nsi(dev,bs,ns)- logits=torch.empty((bs,ns,NH,VD),dtype=torch.float32,device=dev)- lse=torch.empty((bs,ns,NH,1),dtype=torch.float32,device=dev)- c=(ki,kl,kip,out,qf,ns,nsi,logits,lse);_s2c[key]=c;return c+ 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_CFG = {- (4,8192): (8, False, None, None),- (64,8192): (8, False, 3, None), # PERSISTENT — NP ns=2 failed ranked!(64,1024): (2, True, 2, None),(256,1024):(2, False, 1, 8),}⋯ 1 unchanged linesq,kv_data,qo_indptr,kv_indptr,config=databs=int(config["batch_size"]);kvl=int(config["kv_seq_len"]);dev=q.devicekv_fp8,kv_scale=kv_data["fp8"];qs=_gs(dev)+ # 1K bs<=32: NP via _fwd ps=2if 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]- # 32x8K: NP ns=3 + Triton reduce (from v025h)- if bs == 32 and kvl == 8192:- kb=kv_fp8.view(-1,8,1,QKD);c=_g32s2(dev,bs,kvl,8);_quant(c[4],q,qs)- _s1(c[4],kb,qo_indptr,c[2],c[0],c[1],c[6],None,None,None,1,8,1,SM,c[7],c[8],c[3],qs,kv_scale)- _safe_reduce_branchless[(bs,NH)](- c[7],c[8],c[3],qo_indptr,c[2],c[6],- c[8].stride(0),c[8].stride(2),c[8].stride(1),- c[3].stride(0),c[3].stride(1),- MAYBE_FINAL_OUT=False,BATCH_NUM=bs,BLOCK_DV=VD,Lv=VD,mgc=MGC,- num_warps=2,num_stages=1,waves_per_eu=2- )+ # 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)return c[3]- # 256x8K: NP ns=1 ps=64 direct output (FASTEST!)- if bs == 256 and kvl == 8192:- ps=64;kb=kv_fp8.view(-1,ps,1,QKD);c=_g_np_direct(dev,bs,kvl,ps);_quant(c[4],q,qs)- c[3].zero_()- c[8].fill_(-float("inf"))- _s1(c[4],kb,qo_indptr,c[2],c[0],c[1],c[6],None,None,None,1,ps,1,SM,c[7],c[8],c[3],qs,kv_scale)- return c[3]- # Everything else: persistent (4x8K, 64x1K, 64x8K, 256x1K)+ # 64x1K, 256x1K: persistentps,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)
scrolls · 185 diff lines total
Best evidence level for this revision: reported
JSON