Skip to content
KernelIndex
Search⌘K

submission 754992

josusanmartin · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 127 lines, June 9 Researcher Reciprocity License v1.0.

v033b.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-754992?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
21.2µs
#3 of 766
2026-04-07

Reported · How evidence levels are derived →

Source and license

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

Kernel source

v033b.py127 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""V033b: v029n + Triton fused init kernel for 8K shapes.
Replace zero_() + fill_(-inf) with one Triton kernel that does both.
Saves ~1.5us dispatch overhead per 8K shape (4 ops -> 3 ops)."""
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
import triton
import triton.language as tl
from task import input_t, output_t
import aiter
from aiter import dtypes as _ad, mla as _am
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=_ad.fp8;BF16=torch.bfloat16
QSV=float(2.0/torch.finfo(FP8).max);_st={};_nc={};_pc={};_s2c={}
_NEG_INF=float("-inf")

@triton.jit
def _init_kernel(out_ptr, lse_ptr, out_numel, lse_numel, BLOCK: tl.constexpr):
    pid = tl.program_id(0)
    out_blocks = tl.cdiv(out_numel, BLOCK)
    if pid < out_blocks:
        offs = pid * BLOCK + tl.arange(0, BLOCK)
        mask = offs < out_numel
        tl.store(out_ptr + offs, tl.zeros([BLOCK], dtype=tl.bfloat16), mask=mask)
    else:
        lse_pid = pid - out_blocks
        offs = lse_pid * BLOCK + tl.arange(0, BLOCK)
        mask = offs < lse_numel
        neg_inf = tl.full([BLOCK], -float("inf"), dtype=tl.float32)
        tl.store(lse_ptr + offs, neg_inf, mask=mask)

def _fused_init(out, lse):
    out_n = out.numel()
    lse_n = lse.numel()
    BLOCK = 1024
    out_blocks = (out_n + BLOCK - 1) // BLOCK
    lse_blocks = (lse_n + BLOCK - 1) // BLOCK
    grid = (out_blocks + lse_blocks,)
    _init_kernel[grid](out, lse, out_n, lse_n, BLOCK=BLOCK)

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,ns,kg):
    key=(bs,kvl,ps,intra,ns,kg);c=_pc.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)
    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,"npn2");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
_8K_PS={4:2048,32:1024,64:512,256:2048}
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,kvs=kv_data["fp8"];qs=_gs(dev)
    if kvl==1024:
        if bs<=32:
            c=_gnc(dev,bs,kvl,2);_quant(c[4],q,qs)
            _fwd(c[4],kv.view(-1,2,1,QKD),c[3],qo_indptr,c[2],c[0],c[1],1,2,1,SM,q_scale=qs,kv_scale=kvs,intra_batch_mode=True)
            return c[3]
        elif bs==64:
            c=_gpc(dev,qo_indptr,64,1024,2,True,1,4);_quant(c[11],q,qs)
            _s1(c[11],kv.view(-1,2,1,QKD),qo_indptr,c[2],c[0],c[1],None,c[4],c[5],c[6],1,2,1,SM,c[12],c[13],c[3],qs,kvs)
            if c[10]>1:_rd(c[12],c[13],c[7],c[8],c[9],1,c[3],None)
            return c[3]
        else:
            c=_gpc(dev,qo_indptr,256,1024,2,True,1,2);_quant(c[11],q,qs)
            _s1(c[11],kv.view(-1,2,1,QKD),qo_indptr,c[2],c[0],c[1],None,c[4],c[5],c[6],1,2,1,SM,c[12],c[13],c[3],qs,kvs)
            if c[10]>1:_rd(c[12],c[13],c[7],c[8],c[9],1,c[3],None)
            return c[3]
    # 8K: NP ns=1 with fused Triton init (1 launch instead of 2)
    ps=_8K_PS[bs]
    c=_g_np(dev,bs,kvl,ps);_quant(c[4],q,qs)
    _fused_init(c[3], c[7])
    _s1(c[4],kv.view(-1,ps,1,QKD),qo_indptr,c[2],c[0],c[1],c[5],None,None,None,1,ps,1,SM,c[6],c[7],c[3],qs,kvs)
    return c[3]
scrolls · 127 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 736348.

#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
- """V029n: Push ALL 8K shapes to maximum page sizes.
- - 4x8K: ps=2048 (8192/2048=4 pages, 4*4=16 total)
- - 32x8K: ps=1024 (8192/1024=8 pages, 32*8=256 total)
- - 64x8K: ps=512 (8192/512=16 pages, 64*16=1024 total)
- - 256x8K: ps=2048 (proven in v029b leaderboard)
- 1K shapes: best proven configs from v027m."""
+ """V033b: v029n + Triton fused init kernel for 8K shapes.
+ Replace zero_() + fill_(-inf) with one Triton kernel that does both.
+ Saves ~1.5us dispatch overhead per 8K shape (4 ops -> 3 ops)."""
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
+ import triton
+ import triton.language as tl
from task import input_t, output_t
import aiter
from aiter import dtypes as _ad, mla as _am
⋯ 16 unchanged lines
_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=_ad.fp8;BF16=torch.bfloat16
QSV=float(2.0/torch.finfo(FP8).max);_st={};_nc={};_pc={};_s2c={}
+ _NEG_INF=float("-inf")
+
+ @triton.jit
+ def _init_kernel(out_ptr, lse_ptr, out_numel, lse_numel, BLOCK: tl.constexpr):
+ pid = tl.program_id(0)
+ out_blocks = tl.cdiv(out_numel, BLOCK)
+ if pid < out_blocks:
+ offs = pid * BLOCK + tl.arange(0, BLOCK)
+ mask = offs < out_numel
+ tl.store(out_ptr + offs, tl.zeros([BLOCK], dtype=tl.bfloat16), mask=mask)
+ else:
+ lse_pid = pid - out_blocks
+ offs = lse_pid * BLOCK + tl.arange(0, BLOCK)
+ mask = offs < lse_numel
+ neg_inf = tl.full([BLOCK], -float("inf"), dtype=tl.float32)
+ tl.store(lse_ptr + offs, neg_inf, mask=mask)
+
+ def _fused_init(out, lse):
+ out_n = out.numel()
+ lse_n = lse.numel()
+ BLOCK = 1024
+ out_blocks = (out_n + BLOCK - 1) // BLOCK
+ lse_blocks = (lse_n + BLOCK - 1) // BLOCK
+ grid = (out_blocks + lse_blocks,)
+ _init_kernel[grid](out, lse, out_n, lse_n, BLOCK=BLOCK)
+
def _gs(d):
t=_st.get(0)
if t is None: t=torch.tensor([QSV],dtype=torch.float32,device=d);_st[0]=t
⋯ 8 unchanged lines
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)
+ def _gpc(dev,qo,bs,kvl,ps,intra,ns,kg):
+ key=(bs,kvl,ps,intra,ns,kg);c=_pc.get(key)
if c: return c
- tot=bs*kvl;ppb=kvl//ps
- ns=split_override if split_override else _am.get_meta_param(None,bs,tot,NH,1,FP8)[0]
+ 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)
- kg=kg_override if kg_override 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)
⋯ 13 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
- # Per-shape optimal page sizes
- _8K_PS = {4: 2048, 32: 1024, 64: 512, 256: 2048}
+ _8K_PS={4:2048,32:1024,64:512,256:2048}
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
⋯ 13 unchanged lines
_s1(c[11],kv.view(-1,2,1,QKD),qo_indptr,c[2],c[0],c[1],None,c[4],c[5],c[6],1,2,1,SM,c[12],c[13],c[3],qs,kvs)
if c[10]>1:_rd(c[12],c[13],c[7],c[8],c[9],1,c[3],None)
return c[3]
- # 8K: NP ns=1 with per-shape optimal large ps
+ # 8K: NP ns=1 with fused Triton init (1 launch instead of 2)
ps=_8K_PS[bs]
c=_g_np(dev,bs,kvl,ps);_quant(c[4],q,qs)
- c[3].zero_();c[7].fill_(float("-inf"))
+ _fused_init(c[3], c[7])
_s1(c[4],kv.view(-1,ps,1,QKD),qo_indptr,c[2],c[0],c[1],c[5],None,None,None,1,ps,1,SM,c[6],c[7],c[3],qs,kvs)
return c[3]
scrolls · 95 diff lines total

Best evidence level for this revision: reported

JSON