Skip to content
KernelIndex
Search⌘K

submission 609267

josusanmartin · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-609267?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
43.6µs
#131 of 766
2026-03-22

Reported · How evidence levels are derived →

Source and license

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

Kernel source

submission.py65 lines
#!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)."""
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
# 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={}
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 _gc(dev,qo,bs,kvl):
    key=(bs,kvl)
    c=_c.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)
    _c[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"]
    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)
    return c[3]
scrolls · 65 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 605758.

#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
- """V1092: ALL ps=2 ALL fp8 with Q_scale=5.0/max (larger to avoid clipping)."""
+ """V1115: Q_scale=2.0 (optimal from sweep): ALL ps=2 ALL fp8 with Q_scale=5.0/max (larger to avoid clipping)."""
import os
os.environ.setdefault("HIP_FORCE_DEV_KERNARG", "1")
os.environ.setdefault("AMD_DIRECT_DISPATCH", "1")
⋯ 14 unchanged lines
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(5.0/torch.finfo(FP8).max)
+ QSV=float(2.0/torch.finfo(FP8).max)
PS=2
_st={};_c={}
def _gs(d):

Best evidence level for this revision: reported

JSON