submission 715629
bill_97933 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 54 lines, June 9 Researcher Reciprocity License v1.0.
m347_s1s2_merged.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-715629?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:fddc5d6a09cd914531f07e8bb19988a5fd3e8c794b59e39b3b8d9a5f1cd0422e
license declaredunknown
license concludedunknown
authorsbill_97933
imported2026-08-15
Kernel source
m347_s1s2_merged.py54 lines
"""
MLA m347: Merged S1+S2 optimizations.
- S1: sp 8→16 (-1.7µs from m330)
- S2: sp 16→24 (-2µs confirmed from m336/m340)
All other shapes: m312 baseline.
"""
import torch, aiter
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
from task import input_t, output_t
NUM_HEADS=16;NUM_KV_HEADS=1;QK_HEAD_DIM=576;V_HEAD_DIM=512
SM_SCALE=1.0/(QK_HEAD_DIM**0.5);PAGE_SIZE=1;FP8_DTYPE=aiter_dtypes.fp8
SHAPE_CONFIG={
(4,1024):16, # S1: CHANGED 8→16 (m330: -1.7µs)
(4,8192):24, # S2: CHANGED 16→24 (confirmed -2µs)
(32,1024):8, # S3: baseline
(32,8192):8, # S4: baseline
(64,1024):4, # S5: baseline
(64,8192):4, # S6: baseline (sp=8 fails LB)
(256,1024):1, # S7: no reduce
(256,8192):1, # S8: no reduce
}
NO_REDUCE_SHAPES={(256,1024),(256,8192)}
_cache={};_one=None
def _get_cached(bs,kvl,qsl,tq,qoi,kvi):
global _one
k=(bs,kvl)
if k in _cache:return _cache[k]
sp=SHAPE_CONFIG.get(k,8);tkv=bs*kvl
ki=torch.arange(tkv,dtype=torch.int32,device="cuda")
klp=(kvi[1:]-kvi[:-1]).to(torch.int32);qd=FP8_DTYPE
if _one is None:_one=torch.ones(1,dtype=torch.float32,device="cuda")
info=get_mla_metadata_info_v1(bs,qsl,NUM_HEADS,qd,FP8_DTYPE,is_sparse=False,fast_mode=True,num_kv_splits=sp,intra_batch_mode=True)
w=[torch.empty(s,dtype=t,device="cuda") for s,t in info]
get_mla_metadata_v1(qoi,kvi,klp,NUM_HEADS//NUM_KV_HEADS,NUM_KV_HEADS,True,w[0],w[2],w[1],w[3],w[4],w[5],page_size=PAGE_SIZE,kv_granularity=16,max_seqlen_qo=qsl,uni_seqlen_qo=qsl,fast_mode=True,max_split_per_batch=sp,intra_batch_mode=True,dtype_q=qd,dtype_kv=FP8_DTYPE)
n_partial=w[5].size(0)*qsl
logits=torch.empty((n_partial,1,NUM_HEADS,V_HEAD_DIM),dtype=torch.float32,device="cuda")
attn_lse=torch.empty((n_partial,1,NUM_HEADS,1),dtype=torch.float32,device="cuda")
ob=torch.empty((tq,NUM_HEADS,V_HEAD_DIM),dtype=torch.bfloat16,device="cuda")
qi=torch.empty((tq,NUM_HEADS,QK_HEAD_DIM),dtype=FP8_DTYPE,device="cuda")
_cache[k]=(ki,klp,w,sp,logits,attn_lse,ob,qi)
return _cache[k]
def custom_kernel(data:input_t)->output_t:
q,kv_data,qo_indptr,kv_indptr,config=data
bs=config["batch_size"];nkv=config["num_kv_heads"];qsl=config["q_seq_len"];kvl=config["kv_seq_len"];tq=q.shape[0]
kvf,kvs=kv_data["fp8"]
ki,klp,w,sp,logits,attn_lse,o,qi=_get_cached(bs,kvl,qsl,tq,qo_indptr,kv_indptr)
qi.copy_(q.view(tq,NUM_HEADS,QK_HEAD_DIM))
kv4=kvf.view(kvf.shape[0],PAGE_SIZE,nkv,QK_HEAD_DIM)
aiter.mla_decode_stage1_asm_fwd(qi,kv4,qo_indptr,kv_indptr,ki,klp,None,w[0],w[1],w[2],qsl,PAGE_SIZE,nkv,SM_SCALE,logits,attn_lse,o,_one,kvs)
if(bs,kvl) not in NO_REDUCE_SHAPES:
aiter.mla_reduce_v1(logits,attn_lse,w[3],w[4],w[5],qsl,o,None)
return o
scrolls · 54 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Best evidence level for this revision: reported
JSON