Skip to content
KernelIndex
Search⌘K

submission 587729

Danishlynx · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 56 lines, June 9 Researcher Reciprocity License v1.0.

submission_v4.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-587729?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.4µs
#198 of 766
2026-03-19

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:e41ac49622476e83c499886199a85f02b2d49ebeb599f0ce439796b4cf41c849
license declaredunknown
license concludedunknown
authorsDanishlynx
imported2026-08-15

Kernel source

submission_v4.py56 lines
# /// 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.
"""
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, 512), 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 · 56 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