Skip to content
KernelIndex
Search⌘K

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
AMD Instinct MI355X
56.0µs
#193 of 766
2026-03-19

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 torch
import aiter
⋯ 22 unchanged lines
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)
+
+ 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