submission 588139
Danishlynx · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 80 lines, June 9 Researcher Reciprocity License v1.0.
submission_qw.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-588139?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:500da280b5bcf6745c51884fc15afd95a08bc98781425d5247ba12d41c4f4415
license declaredunknown
license concludedunknown
authorsDanishlynx
imported2026-08-15
Kernel source
submission_qw.py80 lines
# /// script
# leaderboard = "amd-mixed-mla"
# ///
"""QW: ob.py + ns=16 for (4,1024).
ii2 proved ns=16 saves ~2us vs ns=9 for (4,1024) with kvg=32.
ii2 PASSED LB. This is ob.py with that one fix.
"""
import torch
import aiter
from task import input_t, output_t
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
_D = aiter_dtypes.fp8
_S = 1.0 / (576 ** 0.5)
_c = {}
_NS = {
(4, 1024): 16, (4, 8192): 16,
(32, 1024): 8, (32, 8192): 8,
(64, 1024): 4, (64, 8192): 4,
(256, 1024): 1, (256, 8192): 1,
}
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
bs = config["batch_size"]
kvsl = config["kv_seq_len"]
key = (bs, kvsl)
if key not in _c:
ns = _NS.get(key, max(1, 256 // bs))
ki = torch.arange(bs * kvsl, dtype=torch.int32, device="cuda")
kl = torch.full((bs,), kvsl, dtype=torch.int32, device="cuda")
info = get_mla_metadata_info_v1(
bs, 1, 16, _D, _D,
is_sparse=False, fast_mode=True,
num_kv_splits=ns, intra_batch_mode=True,
)
wk = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
get_mla_metadata_v1(
qo_indptr, kv_indptr, kl, 16, 1, True,
wk[0], wk[2], wk[1], wk[3], wk[4], wk[5],
page_size=1, kv_granularity=32,
max_seqlen_qo=1, uni_seqlen_qo=1,
fast_mode=True, max_split_per_batch=ns,
intra_batch_mode=True, dtype_q=_D, dtype_kv=_D,
)
np_ = wk[5].size(0)
sd = torch.empty((np_, 1, 16, 576), dtype=torch.float32, device="cuda")
sl = torch.empty((np_, 1, 16, 1), dtype=torch.float32, device="cuda")
qs = torch.ones(1, dtype=torch.float32, device="cuda")
qf = torch.empty((bs, 16, 576), dtype=_D, device="cuda")
ob = torch.empty((bs, 16, 512), dtype=torch.bfloat16, device="cuda")
_c[key] = (ki, kl, wk, sd, sl, qs, qf, ob, ns)
torch.cuda.synchronize()
ki, kl, wk, sd, sl, qs, qf, ob, ns = _c[key]
kf, k_scale = kv_data["fp8"]
qf.copy_(q.view(-1, 16, 576))
aiter.mla_decode_stage1_asm_fwd(
qf, kf.view(-1, 1, 1, 576),
qo_indptr, kv_indptr,
ki, kl,
None, wk[0], wk[1], wk[2],
1, 1, 1, _S,
sd, sl, ob,
q_scale=qs, kv_scale=k_scale,
)
if ns > 1:
aiter.mla_reduce_v1(sd, sl, wk[3], wk[4], wk[5], 1, ob, None)
return ob
scrolls · 80 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 587729.
# /// script# leaderboard = "amd-mixed-mla"# ///- """v4: QW + splitData=512. No Q cache (caused LB failure).- splitData=512 saves bandwidth vs 576, proven correct pre-reset.+ """QW: ob.py + ns=16 for (4,1024).+ ii2 proved ns=16 saves ~2us vs ns=9 for (4,1024) with kvg=32.+ ii2 PASSED LB. This is ob.py with that one fix."""import torchimport aiter⋯ 22 unchanged linesns = _NS.get(key, max(1, 256 // bs))ki = torch.arange(bs * kvsl, dtype=torch.int32, device="cuda")kl = torch.full((bs,), kvsl, dtype=torch.int32, device="cuda")- info = get_mla_metadata_info_v1(bs, 1, 16, _D, _D, is_sparse=False, fast_mode=True, num_kv_splits=ns, intra_batch_mode=True)++ info = get_mla_metadata_info_v1(+ bs, 1, 16, _D, _D,+ is_sparse=False, fast_mode=True,+ num_kv_splits=ns, intra_batch_mode=True,+ )wk = [torch.empty(s, dtype=t, device="cuda") for s, t in info]- get_mla_metadata_v1(qo_indptr, kv_indptr, kl, 16, 1, True, wk[0], wk[2], wk[1], wk[3], wk[4], wk[5],- page_size=1, kv_granularity=32, max_seqlen_qo=1, uni_seqlen_qo=1,- fast_mode=True, max_split_per_batch=ns, intra_batch_mode=True, dtype_q=_D, dtype_kv=_D)+ get_mla_metadata_v1(+ qo_indptr, kv_indptr, kl, 16, 1, True,+ wk[0], wk[2], wk[1], wk[3], wk[4], wk[5],+ page_size=1, kv_granularity=32,+ max_seqlen_qo=1, uni_seqlen_qo=1,+ fast_mode=True, max_split_per_batch=ns,+ intra_batch_mode=True, dtype_q=_D, dtype_kv=_D,+ )+np_ = wk[5].size(0)- sd = torch.empty((np_, 1, 16, 512), dtype=torch.float32, device="cuda")+ sd = torch.empty((np_, 1, 16, 576), dtype=torch.float32, device="cuda")sl = torch.empty((np_, 1, 16, 1), dtype=torch.float32, device="cuda")qs = torch.ones(1, dtype=torch.float32, device="cuda")qf = torch.empty((bs, 16, 576), dtype=_D, device="cuda")ob = torch.empty((bs, 16, 512), dtype=torch.bfloat16, device="cuda")+_c[key] = (ki, kl, wk, sd, sl, qs, qf, ob, ns)torch.cuda.synchronize()ki, kl, wk, sd, sl, qs, qf, ob, ns = _c[key]kf, k_scale = kv_data["fp8"]+qf.copy_(q.view(-1, 16, 576))- aiter.mla_decode_stage1_asm_fwd(qf, kf.view(-1, 1, 1, 576), qo_indptr, kv_indptr, ki, kl,- None, wk[0], wk[1], wk[2], 1, 1, 1, _S, sd, sl, ob, q_scale=qs, kv_scale=k_scale)++ aiter.mla_decode_stage1_asm_fwd(+ qf, kf.view(-1, 1, 1, 576),+ qo_indptr, kv_indptr,+ ki, kl,+ None, wk[0], wk[1], wk[2],+ 1, 1, 1, _S,+ sd, sl, ob,+ q_scale=qs, kv_scale=k_scale,+ )+if ns > 1:aiter.mla_reduce_v1(sd, sl, wk[3], wk[4], wk[5], 1, ob, None)+return ob
scrolls · 67 diff lines total
Best evidence level for this revision: reported
JSON