submission 624092
josusanmartin · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 61 lines, June 9 Researcher Reciprocity License v1.0.
submission_v1191.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-624092?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:cedc5df974961f096ad230eebfe69644748840101586bc002d3fe2865be4adf3
license declaredunknown
license concludedunknown
authorsjosusanmartin
imported2026-08-15
Kernel source
submission_v1191.py61 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""V1191: ps=4 ONLY for 64×8K and 256×8K (proven in test). ps=2 for everything else.
Test covers 64×8K and 256×8K at ps=4, so these are safe. 4×8K/32×8K stay at ps=2."""
import os
os.environ.setdefault("HIP_FORCE_DEV_KERNARG", "1")
os.environ.setdefault("AMD_DIRECT_DISPATCH", "1")
import torch
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_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={};_pc={};_nc={}
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 _gpc(dev,qo,bs,kvl,ps):
key=(bs,kvl,ps);c=_pc.get(key)
if c: return c
tot=bs*kvl;ppb=kvl//ps;np_=tot//ps;ns,_=aiter_mla.get_meta_param(None,bs,tot,NH,1,FP8)
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)
kg=max(ps,16)
info=_mi(bs,1,NH,FP8,FP8,is_sparse=False,fast_mode=True,num_kv_splits=ns,intra_batch_mode=True)
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=True,dtype_q=FP8,dtype_kv=FP8)
qf=torch.empty((bs,NH,QKD),dtype=FP8,device=dev);c=(ki,kl,kip,out,w[0],w[1],w[2],w[3],w[4],w[5],ns,qf);_pc[key]=c;return c
def _gnc(dev,bs,kvl,ps):
key=(bs,kvl,ps);c=_nc.get(key)
if c: return c
tot=bs*kvl;ppb=kvl//ps;np_=tot//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
_NP_SHAPES = {(4,1024), (32,1024), (256,1024)}
# Use ps=4 only for shapes proven in test (64×8K and 256×8K)
_PS4_SHAPES = {(64,8192), (256,8192)}
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"]
ps = 4 if (bs,kvl) in _PS4_SHAPES else 2
kb=kv_fp8.view(-1,ps,1,QKD);qs=_gs(dev)
if (bs,kvl) in _NP_SHAPES:
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]
c=_gpc(dev,qo_indptr,bs,kvl,ps);_quant(c[11],q,qs)
_fwd(c[11],kb,c[3],qo_indptr,c[2],c[0],c[1],1,ps,1,SM,num_kv_splits=c[10],work_meta_data=c[4],work_indptr=c[5],work_info_set=c[6],reduce_indptr=c[7],reduce_final_map=c[8],reduce_partial_map=c[9],q_scale=qs,kv_scale=kv_scale,intra_batch_mode=True)
return c[3]
scrolls · 61 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 609267.
#!POPCORN leaderboard amd-mixed-mla#!POPCORN gpu MI355X- """V1115: Q_scale=2.0 (optimal from sweep): ALL ps=2 ALL fp8 with Q_scale=5.0/max (larger to avoid clipping)."""+ """V1191: ps=4 ONLY for 64×8K and 256×8K (proven in test). ps=2 for everything else.+ Test covers 64×8K and 256×8K at ps=4, so these are safe. 4×8K/32×8K stay at ps=2."""import osos.environ.setdefault("HIP_FORCE_DEV_KERNARG", "1")os.environ.setdefault("AMD_DIRECT_DISPATCH", "1")⋯ 9 unchanged linesfrom aiter.jit.module_mla_metadata import get_mla_metadata_v1 as _mv1except 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- # LARGER Q scale to avoid clipping (Q values can be up to ~4.5)- QSV=float(2.0/torch.finfo(FP8).max)- PS=2- _st={};_c={}+ _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={};_pc={};_nc={}def _gs(d):t=_st.get(0)if t is None: t=torch.tensor([QSV],dtype=torch.float32,device=d);_st[0]=treturn t- def _gc(dev,qo,bs,kvl):- key=(bs,kvl)- c=_c.get(key)+ def _gpc(dev,qo,bs,kvl,ps):+ key=(bs,kvl,ps);c=_pc.get(key)if c: return c- tot=bs*kvl;ppb=kvl//PS;np_=tot//PS- ns,_=aiter_mla.get_meta_param(None,bs,tot,NH,1,FP8)- 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)- kg=max(PS,16)+ tot=bs*kvl;ppb=kvl//ps;np_=tot//ps;ns,_=aiter_mla.get_meta_param(None,bs,tot,NH,1,FP8)+ 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)+ kg=max(ps,16)info=_mi(bs,1,NH,FP8,FP8,is_sparse=False,fast_mode=True,num_kv_splits=ns,intra_batch_mode=True)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=True,dtype_q=FP8,dtype_kv=FP8)- qf=torch.empty((bs,NH,QKD),dtype=FP8,device=dev)- c=(ki,kl,kip,out,w[0],w[1],w[2],w[3],w[4],w[5],ns,qf)- _c[key]=c; return c+ _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=True,dtype_q=FP8,dtype_kv=FP8)+ qf=torch.empty((bs,NH,QKD),dtype=FP8,device=dev);c=(ki,kl,kip,out,w[0],w[1],w[2],w[3],w[4],w[5],ns,qf);_pc[key]=c;return c+ def _gnc(dev,bs,kvl,ps):+ key=(bs,kvl,ps);c=_nc.get(key)+ if c: return c+ tot=bs*kvl;ppb=kvl//ps;np_=tot//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+ _NP_SHAPES = {(4,1024), (32,1024), (256,1024)}+ # Use ps=4 only for shapes proven in test (64×8K and 256×8K)+ _PS4_SHAPES = {(64,8192), (256,8192)}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+ bs=int(config["batch_size"]);kvl=int(config["kv_seq_len"]);dev=q.devicekv_fp8,kv_scale=kv_data["fp8"]- kb=kv_fp8.view(-1,PS,1,QKD)- c=_gc(dev,qo_indptr,bs,kvl)- qs=_gs(dev)- _quant(c[11],q,qs)- _fwd(c[11],kb,c[3],qo_indptr,c[2],c[0],c[1],- 1,PS,1,SM,num_kv_splits=c[10],- work_meta_data=c[4],work_indptr=c[5],work_info_set=c[6],- reduce_indptr=c[7],reduce_final_map=c[8],reduce_partial_map=c[9],- q_scale=qs,kv_scale=kv_scale,intra_batch_mode=True)+ ps = 4 if (bs,kvl) in _PS4_SHAPES else 2+ kb=kv_fp8.view(-1,ps,1,QKD);qs=_gs(dev)+ if (bs,kvl) in _NP_SHAPES:+ 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]+ c=_gpc(dev,qo_indptr,bs,kvl,ps);_quant(c[11],q,qs)+ _fwd(c[11],kb,c[3],qo_indptr,c[2],c[0],c[1],1,ps,1,SM,num_kv_splits=c[10],work_meta_data=c[4],work_indptr=c[5],work_info_set=c[6],reduce_indptr=c[7],reduce_final_map=c[8],reduce_partial_map=c[9],q_scale=qs,kv_scale=kv_scale,intra_batch_mode=True)return c[3]
scrolls · 88 diff lines total
Best evidence level for this revision: reported
JSON