submission 635670
josusanmartin · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 142 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-635670?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:9e761a6a2145759bf1aaad411dd62537cde1c2f84c30ff3d102e95be68f721b1
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
Persistent wrapper for 64×1K, 256×1K (from v012g — avoids NP for large batch 1K).split-k
for split_kv_id in range(0, num_valid_kv_splits):Kernel source
submission.py142 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""V1261: Best of v012g safety + v1260 ps=8 speed.
NP only for 4×1K, 32×1K (safest, small batch, safe reduce patched).
Persistent wrapper for 64×1K, 256×1K (from v012g — avoids NP for large batch 1K).
Direct _s1+_rd ps=8 intra=False for ALL 8K (from v1260 — ps=8 is 10µs faster than ps=4).
Target ~31.5µs."""
import os
os.environ.setdefault("HIP_FORCE_DEV_KERNARG", "1")
os.environ.setdefault("AMD_DIRECT_DISPATCH", "1")
import torch
import triton
import 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_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 _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
_mi = aiter.get_mla_metadata_info_v1
_orig_fwd = aiter.mla.mla_decode_fwd
@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_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 _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
def _gpc(dev,qo,bs,kvl,ps,intra=True):
key=(bs,kvl,ps,intra);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=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,intra,po,pl);_pc[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)
# NP: only 4×1K, 32×1K (small batch, safe reduce patched)
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: direct _s1+_rd ps=8 intra=False
if kvl == 8192:
ps=8;kb=kv_fp8.view(-1,ps,1,QKD)
c=_gpc(dev,qo_indptr,bs,kvl,ps,intra=False);_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[13],c[14],c[3],qs,kv_scale)
if c[10]>1: _rd(c[13],c[14],c[7],c[8],c[9],1,c[3],None)
return c[3]
# 64×1K, 256×1K: persistent wrapper (safe, ASM reduce internally)
ps=2;kb=kv_fp8.view(-1,ps,1,QKD)
c=_gpc(dev,qo_indptr,bs,kvl,ps);_quant(c[11],q,qs)
_orig_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=c[12])
return c[3]
scrolls · 142 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 631834.
#!POPCORN leaderboard amd-mixed-mla#!POPCORN gpu MI355X- """V010i: v010d plus intra_batch_mode=False for 256x1k."""+ """V1261: Best of v012g safety + v1260 ps=8 speed.+ NP only for 4×1K, 32×1K (safest, small batch, safe reduce patched).+ Persistent wrapper for 64×1K, 256×1K (from v012g — avoids NP for large batch 1K).+ Direct _s1+_rd ps=8 intra=False for ALL 8K (from v1260 — ps=8 is 10µs faster than ps=4).+ Target ~31.5µs."""import osos.environ.setdefault("HIP_FORCE_DEV_KERNARG", "1")os.environ.setdefault("AMD_DIRECT_DISPATCH", "1")import torch+ import triton+ import triton.language as tlfrom task import input_t, output_timport aiterfrom aiter import dtypes as aiter_dtypes, mla as aiter_mla⋯ 5 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+ 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+ _mi = aiter.get_mla_metadata_info_v1+ _orig_fwd = aiter.mla.mla_decode_fwd++ @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_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={}+ 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 _gpc(dev,qo,bs,kvl,ps,split_override=None,intra=True):- key=(bs,kvl,ps,split_override,intra);c=_pc.get(key)+ def _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+ def _gpc(dev,qo,bs,kvl,ps,intra=True):+ key=(bs,kvl,ps,intra);c=_pc.get(key)+ if c: return ctot=bs*kvl;ppb=kvl//ps;np_=tot//ps- ns = split_override if split_override is not None else aiter_mla.get_meta_param(None,bs,tot,NH,1,FP8)[0]+ 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=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);c=(ki,kl,kip,out,w[0],w[1],w[2],w[3],w[4],w[5],ns,qf,intra);_pc[key]=c;return c- _SPLIT_CFG = {- (4,8192): 6, # was 4, try 6 (more K-parallelism)- (32,8192): 6, # proven- (64,8192): 4, # was 2, try 4 (more K-parallelism)- }+ 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,intra,po,pl);_pc[key]=c;return cdef custom_kernel(data: input_t) -> output_t:q,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"]- ps = 4 if kvl == 8192 else 2- kb=kv_fp8.view(-1,ps,1,QKD);qs=_gs(dev)- split_ov = _SPLIT_CFG.get((bs,kvl))- intra = False if (bs,kvl) in {(32,8192), (256,1024)} else True- c=_gpc(dev,qo_indptr,bs,kvl,ps,split_ov,intra);_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=c[12])+ qs=_gs(dev)+ # NP: only 4×1K, 32×1K (small batch, safe reduce patched)+ 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: direct _s1+_rd ps=8 intra=False+ if kvl == 8192:+ ps=8;kb=kv_fp8.view(-1,ps,1,QKD)+ c=_gpc(dev,qo_indptr,bs,kvl,ps,intra=False);_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[13],c[14],c[3],qs,kv_scale)+ if c[10]>1: _rd(c[13],c[14],c[7],c[8],c[9],1,c[3],None)+ return c[3]+ # 64×1K, 256×1K: persistent wrapper (safe, ASM reduce internally)+ ps=2;kb=kv_fp8.view(-1,ps,1,QKD)+ c=_gpc(dev,qo_indptr,bs,kvl,ps);_quant(c[11],q,qs)+ _orig_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=c[12])return c[3]
scrolls · 155 diff lines total
Best evidence level for this revision: reported
JSON