submission 707711
augustus2024 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 51 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-707711?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:58257288b92c637ebbee1b1e31ca845681cbbc33e218bd6ab1689fb563546439
license declaredunknown
license concludedunknown
authorsaugustus2024
imported2026-08-15
Kernel source
submission.py51 lines
"""V188 — V185 splits + ps=4 for kv=8192 + safer (64,1024) nks=4.
kv=8192 ps=4 gives ~20-47% speedup. kv=1024 uses V79's proven ps."""
import os
os.environ["HIP_FORCE_DEV_KERNARG"]="1";os.environ["HSA_ENABLE_SDMA"]="0"
import torch
from task import input_t, output_t
import aiter
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
NH=16;NKH=1;QKD=576;VD=512;SM=1.0/(QKD**0.5);FP8=aiter_dtypes.fp8
_stage1=aiter.mla_decode_stage1_asm_fwd
_reduce=aiter.mla_reduce_v1
_qs=None;_cache={}
# (ps, nks) per config — ps=4 for kv=8192, V79 ps for kv=1024
_CFG={
(4,1024):(1,16),(4,8192):(4,16),
(32,1024):(2,2),(32,8192):(4,8),
(64,1024):(2,4),(64,8192):(4,8), # nks=4 for safety (V79 value)
(256,1024):(2,2),(256,8192):(4,4),
}
def _build(bs,kvs,qd,kd,ps,nks):
npp=kvs//ps
qo=torch.arange(0,bs+1,dtype=torch.int32,device="cuda")
kvi=torch.arange(0,bs+1,dtype=torch.int32,device="cuda")*npp
klp=torch.full((bs,),ps,dtype=torch.int32,device="cuda")
kvidx=torch.arange(bs*npp,dtype=torch.int32,device="cuda")
info=get_mla_metadata_info_v1(bs,1,NH,qd,kd,is_sparse=False,fast_mode=False,num_kv_splits=nks,intra_batch_mode=True)
w=[torch.empty(s,dtype=t,device="cuda") for s,t in info]
wm,wi,wis,ri,rfm,rpm=w
get_mla_metadata_v1(qo,kvi,klp,NH,NKH,True,wm,wis,wi,ri,rfm,rpm,
page_size=ps,kv_granularity=max(ps,16),max_seqlen_qo=1,uni_seqlen_qo=1,
fast_mode=False,max_split_per_batch=nks,intra_batch_mode=True,dtype_q=qd,dtype_kv=kd)
lg=torch.empty((rpm.size(0),1,NH,VD),dtype=torch.float32,device="cuda")
al=torch.empty((rpm.size(0),1,NH,1),dtype=torch.float32,device="cuda")
o=torch.empty((bs,NH,VD),dtype=torch.bfloat16,device="cuda")
return (wm,wi,wis,ri,rfm,rpm,kvidx,klp,qo,kvi,lg,al,o,ps)
def custom_kernel(data:input_t)->output_t:
global _qs
q,kv_data,_,_2,cfg=data
bs=cfg["batch_size"];kvs=cfg["kv_seq_len"]
kf,ks=kv_data["fp8"];qf=q.to(FP8)
if _qs is None:_qs=torch.ones(1,dtype=torch.float32,device="cuda")
ps,nks=_CFG.get((bs,kvs),(2,32))
key=(bs,kvs)
if key not in _cache:_cache[key]=_build(bs,kvs,qf.dtype,kf.dtype,ps,nks)
wm,wi,wis,ri,rfm,rpm,kvidx,klp,qo,kvi,lg,al,o,ps_=_cache[key]
np_=kf.shape[0]//ps_;kv4d=kf.view(np_,ps_,NKH,QKD)
_stage1(qf.view(-1,NH,QKD),kv4d,qo,kvi,kvidx,klp,None,wm,wi,wis,1,ps_,NKH,SM,lg,al,o,_qs,ks)
_reduce(lg,al,ri,rfm,rpm,1,o,None)
return o
scrolls · 51 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 691798.
- """- MLA decode V29 — Surgical hybrid: pg1 ONLY for bs=4 kv≤1024, pg2 for ALL else.-- Ranked test fails only on bs=4 kv=1024 with pg2+skip_quant.- All other configs pass easily. So use pg1 only for that one case.-- bs=4 kv=8K with pg2: 26.2µs (vs pg1: 32.6µs) — 20% win, keeps pg2 here.- """-+ """V188 — V185 splits + ps=4 for kv=8192 + safer (64,1024) nks=4.+ kv=8192 ps=4 gives ~20-47% speedup. kv=1024 uses V79's proven ps."""import os- os.environ["HIP_FORCE_DEV_KERNARG"] = "1"- os.environ["HSA_ENABLE_SDMA"] = "0"-+ os.environ["HIP_FORCE_DEV_KERNARG"]="1";os.environ["HSA_ENABLE_SDMA"]="0"import torchfrom task import input_t, output_t-import aiterfrom aiter import dtypes as aiter_dtypesfrom aiter import get_mla_metadata_info_v1, get_mla_metadata_v1-- NUM_HEADS = 16- NUM_KV_HEADS = 1- QK_HEAD_DIM = 576- V_HEAD_DIM = 512- SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)- FP8_DTYPE = aiter_dtypes.fp8-- _cache = {}- _q_scale_one = None--- def _build_cache(batch_size, q_seq_len, kv_seq_len, q_dtype, kv_dtype, page_size):- total_q = batch_size * q_seq_len- num_kv_splits = 32- num_pages_per_batch = kv_seq_len // page_size-- qo_indptr = torch.arange(0, batch_size + 1, dtype=torch.int32, device="cuda") * q_seq_len- kv_indptr = torch.arange(0, batch_size + 1, dtype=torch.int32, device="cuda") * num_pages_per_batch- kv_last_page_len = torch.full((batch_size,), page_size, dtype=torch.int32, device="cuda")- total_pages = batch_size * num_pages_per_batch- kv_indices = torch.arange(total_pages, dtype=torch.int32, device="cuda")-- info = get_mla_metadata_info_v1(- batch_size, q_seq_len, NUM_HEADS, q_dtype, kv_dtype,- is_sparse=False, fast_mode=False,- num_kv_splits=num_kv_splits, intra_batch_mode=True,- )- work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]- (work_metadata, work_indptr, work_info_set,- reduce_indptr, reduce_final_map, reduce_partial_map) = work-- get_mla_metadata_v1(- qo_indptr, kv_indptr, kv_last_page_len,- NUM_HEADS // NUM_KV_HEADS, NUM_KV_HEADS, True,- work_metadata, work_info_set, work_indptr,- reduce_indptr, reduce_final_map, reduce_partial_map,- page_size=page_size, kv_granularity=max(page_size, 16),- max_seqlen_qo=q_seq_len, uni_seqlen_qo=q_seq_len,- fast_mode=False, max_split_per_batch=num_kv_splits,- intra_batch_mode=True, dtype_q=q_dtype, dtype_kv=kv_dtype,- )-- logits = torch.empty(- (reduce_partial_map.size(0) * q_seq_len, 1, NUM_HEADS, V_HEAD_DIM),- dtype=torch.float32, device="cuda",- )- attn_lse = torch.empty(- (reduce_partial_map.size(0) * q_seq_len, 1, NUM_HEADS, 1),- dtype=torch.float32, device="cuda",- )- o = torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")-- return {- "meta": {- "work_meta_data": work_metadata, "work_indptr": work_indptr,- "work_info_set": work_info_set, "reduce_indptr": reduce_indptr,- "reduce_final_map": reduce_final_map, "reduce_partial_map": reduce_partial_map,- },- "kv_indices": kv_indices, "kv_last_page_len": kv_last_page_len,- "qo_indptr": qo_indptr, "kv_indptr": kv_indptr,- "logits": logits, "attn_lse": attn_lse, "o": o,- "page_size": page_size,- }--- def custom_kernel(data: input_t) -> output_t:- global _q_scale_one- q, kv_data, qo_indptr, kv_indptr, config = data- batch_size = config["batch_size"]- q_seq_len = config["q_seq_len"]- kv_seq_len = config["kv_seq_len"]- kv_buffer_fp8, kv_scale = kv_data["fp8"]-- q_fp8 = q.to(FP8_DTYPE)- if _q_scale_one is None:- _q_scale_one = torch.ones(1, dtype=torch.float32, device="cuda")-- # Surgical hybrid: pg1 ONLY for the problematic case- page_size = 1 if (batch_size <= 4 and kv_seq_len <= 1024) else 2-- key = (batch_size, q_seq_len, kv_seq_len, page_size)- if key not in _cache:- _cache[key] = _build_cache(batch_size, q_seq_len, kv_seq_len, q_fp8.dtype, kv_buffer_fp8.dtype, page_size)- c = _cache[key]- ps = c["page_size"]-- num_pages = kv_buffer_fp8.shape[0] // ps- kv_buffer_4d = kv_buffer_fp8.view(num_pages, ps, NUM_KV_HEADS, QK_HEAD_DIM)-- aiter.mla_decode_stage1_asm_fwd(- q_fp8.view(-1, NUM_HEADS, QK_HEAD_DIM), kv_buffer_4d,- c["qo_indptr"], c["kv_indptr"], c["kv_indices"], c["kv_last_page_len"],- None, c["meta"]["work_meta_data"], c["meta"]["work_indptr"], c["meta"]["work_info_set"],- q_seq_len, ps, NUM_KV_HEADS, SM_SCALE,- c["logits"], c["attn_lse"], c["o"],- _q_scale_one, kv_scale,- )-- aiter.mla_reduce_v1(- c["logits"], c["attn_lse"],- c["meta"]["reduce_indptr"], c["meta"]["reduce_final_map"], c["meta"]["reduce_partial_map"],- q_seq_len, c["o"], None,- )- return c["o"]+ NH=16;NKH=1;QKD=576;VD=512;SM=1.0/(QKD**0.5);FP8=aiter_dtypes.fp8+ _stage1=aiter.mla_decode_stage1_asm_fwd+ _reduce=aiter.mla_reduce_v1+ _qs=None;_cache={}+ # (ps, nks) per config — ps=4 for kv=8192, V79 ps for kv=1024+ _CFG={+ (4,1024):(1,16),(4,8192):(4,16),+ (32,1024):(2,2),(32,8192):(4,8),+ (64,1024):(2,4),(64,8192):(4,8), # nks=4 for safety (V79 value)+ (256,1024):(2,2),(256,8192):(4,4),+ }+ def _build(bs,kvs,qd,kd,ps,nks):+ npp=kvs//ps+ qo=torch.arange(0,bs+1,dtype=torch.int32,device="cuda")+ kvi=torch.arange(0,bs+1,dtype=torch.int32,device="cuda")*npp+ klp=torch.full((bs,),ps,dtype=torch.int32,device="cuda")+ kvidx=torch.arange(bs*npp,dtype=torch.int32,device="cuda")+ info=get_mla_metadata_info_v1(bs,1,NH,qd,kd,is_sparse=False,fast_mode=False,num_kv_splits=nks,intra_batch_mode=True)+ w=[torch.empty(s,dtype=t,device="cuda") for s,t in info]+ wm,wi,wis,ri,rfm,rpm=w+ get_mla_metadata_v1(qo,kvi,klp,NH,NKH,True,wm,wis,wi,ri,rfm,rpm,+ page_size=ps,kv_granularity=max(ps,16),max_seqlen_qo=1,uni_seqlen_qo=1,+ fast_mode=False,max_split_per_batch=nks,intra_batch_mode=True,dtype_q=qd,dtype_kv=kd)+ lg=torch.empty((rpm.size(0),1,NH,VD),dtype=torch.float32,device="cuda")+ al=torch.empty((rpm.size(0),1,NH,1),dtype=torch.float32,device="cuda")+ o=torch.empty((bs,NH,VD),dtype=torch.bfloat16,device="cuda")+ return (wm,wi,wis,ri,rfm,rpm,kvidx,klp,qo,kvi,lg,al,o,ps)+ def custom_kernel(data:input_t)->output_t:+ global _qs+ q,kv_data,_,_2,cfg=data+ bs=cfg["batch_size"];kvs=cfg["kv_seq_len"]+ kf,ks=kv_data["fp8"];qf=q.to(FP8)+ if _qs is None:_qs=torch.ones(1,dtype=torch.float32,device="cuda")+ ps,nks=_CFG.get((bs,kvs),(2,32))+ key=(bs,kvs)+ if key not in _cache:_cache[key]=_build(bs,kvs,qf.dtype,kf.dtype,ps,nks)+ wm,wi,wis,ri,rfm,rpm,kvidx,klp,qo,kvi,lg,al,o,ps_=_cache[key]+ np_=kf.shape[0]//ps_;kv4d=kf.view(np_,ps_,NKH,QKD)+ _stage1(qf.view(-1,NH,QKD),kv4d,qo,kvi,kvidx,klp,None,wm,wi,wis,1,ps_,NKH,SM,lg,al,o,_qs,ks)+ _reduce(lg,al,ri,rfm,rpm,1,o,None)+ return o
scrolls · 168 diff lines total
Best evidence level for this revision: reported
JSON