Skip to content
KernelIndex
Search⌘K

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
AMD Instinct MI355X
55.7µs
#192 of 766
2026-04-03

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