Skip to content
KernelIndex
Search⌘K

submission 742585

lgc0338 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_monkey.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-742585?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
52.6µs
#171 of 766
2026-04-06

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:d4c2f622c027ea2f6ad837007ea1565b35d7e8f9c9f83b0cc6fbb728c06c54cf
license declaredunknown
license concludedunknown
authorslgc0338
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

mmaqk += tl.dot(qq, tl.trans(kk))
num-warps = 4q.stride(0), q.stride(1), kv.stride(0), BN=64, NS=NS, num_warps=4)
online-softmaxm_ij = tl.max(qk, axis=1); m_new = tl.maximum(m_i, m_ij)
persistent-kernelNone, # non-persistent indptr (None = persistent mode)
tile-n = 64q.stride(0), q.stride(1), kv.stride(0), BN=64, NS=NS, num_warps=4)

Kernel source

submission_monkey.py260 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X

"""Monkey-patch: bypass mla_decode_fwd wrapper overhead.

Key findings from aiter source analysis:
1. Wrapper allocates splitData/splitLse EVERY call (~10μs)
2. Wrapper does module lookup every call
3. Wrapper has Python if/elif dispatch logic

Our patch: call stage1_asm + reduce directly with ALL buffers pre-cached.
Previous attempts said "3-6% slower" — but those may not have cached everything properly.

Combined with: Triton for small cases, patched ASM for large cases.
"""

import torch
import triton
import triton.language as tl
from task import input_t, output_t

from aiter import mla_decode_stage1_asm_fwd, mla_reduce_v1
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1

FP8 = aiter_dtypes.fp8
FI = torch.finfo(FP8)
QS_V = 5.0 / FI.max
IV = float(1.0 / QS_V)
SM_SCALE = tl.constexpr(1.0 / (576 ** 0.5))
LOG2E = tl.constexpr(1.44269504)
SM = 1.0 / (576 ** 0.5)
BF16 = torch.bfloat16

_c = {}
_asm = {}
_qs_ready = False


def _g(k, s, d):
    if k not in _c or _c[k].shape != s:
        _c[k] = torch.empty(s, dtype=d, device="cuda")
    return _c[k]


# ============================================================
# Triton kernels (same as final, proven fast for small cases)
# ============================================================
@triton.jit
def _fused_quant(O, I, iv, N: tl.constexpr, B: tl.constexpr):
    pid = tl.program_id(0)
    offs = pid * B + tl.arange(0, B)
    mask = offs < N
    x = tl.load(I + offs, mask=mask).to(tl.float32) * iv
    tl.store(O + offs, tl.clamp(x, -240.0, 240.0).to(O.dtype.element_ty), mask=mask)


@triton.jit
def _attn_fp8(
    Q, KV, PO, PLSE, kv_indptr, kv_scale_ptr,
    stride_qb, stride_qh, stride_kvt,
    BN: tl.constexpr, NS: tl.constexpr,
):
    bid = tl.program_id(0); sid = tl.program_id(1)
    ks = tl.load(kv_indptr + bid); ke = tl.load(kv_indptr + bid + 1)
    split_size = tl.cdiv(ke - ks, NS)
    my_start = ks + sid * split_size
    my_end = tl.minimum(my_start + split_size, ke)
    h = tl.arange(0, 16); v = tl.arange(0, 512)
    m_i = tl.full((16,), float("-inf"), dtype=tl.float32)
    l_i = tl.zeros((16,), dtype=tl.float32)
    acc = tl.zeros((16, 512), dtype=tl.float32)
    kvs = tl.load(kv_scale_ptr); eff_scale = SM_SCALE * kvs
    if my_start >= my_end:
        po = PO + (bid * NS + sid) * 16 * 512
        tl.store(po + h[:, None] * 512 + v[None, :], acc)
        tl.store(PLSE + (bid * NS + sid) * 16 + h, m_i); return
    for tile_start in range(my_start, my_end, BN):
        ti = tile_start + tl.arange(0, BN); tm = ti < my_end
        qk = tl.zeros((16, BN), dtype=tl.float32)
        for dk in range(0, 576, 64):
            do = dk + tl.arange(0, 64)
            qq = tl.load(Q + bid * stride_qb + h[:, None] * stride_qh + do[None, :])
            kk = tl.load(KV + ti[:, None] * stride_kvt + do[None, :],
                         mask=tm[:, None], other=0.0).to(tl.bfloat16)
            qk += tl.dot(qq, tl.trans(kk))
        qk = qk * eff_scale; qk = tl.where(tm[None, :], qk, float("-inf"))
        m_ij = tl.max(qk, axis=1); m_new = tl.maximum(m_i, m_ij)
        alpha = tl.math.exp2((m_i - m_new) * LOG2E)
        p = tl.math.exp2((qk - m_new[:, None]) * LOG2E)
        l_i = alpha * l_i + tl.sum(p, axis=1); acc = acc * alpha[:, None]
        vv = tl.load(KV + ti[:, None] * stride_kvt + v[None, :],
                     mask=tm[:, None], other=0.0).to(tl.bfloat16)
        acc += tl.dot(p.to(tl.bfloat16), vv); m_i = m_new
    acc = (acc * kvs) / l_i[:, None]; lse = m_i + tl.log(l_i)
    po = PO + (bid * NS + sid) * 16 * 512
    tl.store(po + h[:, None] * 512 + v[None, :], acc)
    tl.store(PLSE + (bid * NS + sid) * 16 + h, lse)


@triton.jit
def _attn_bf16(
    Q, KV, PO, PLSE, kv_indptr,
    stride_qb, stride_qh, stride_kvt,
    BN: tl.constexpr, NS: tl.constexpr,
):
    bid = tl.program_id(0); sid = tl.program_id(1)
    h = tl.arange(0, 16); v = tl.arange(0, 512)
    ks = tl.load(kv_indptr + bid); ke = tl.load(kv_indptr + bid + 1)
    split_size = tl.cdiv(ke - ks, NS)
    my_start = ks + sid * split_size
    my_end = tl.minimum(my_start + split_size, ke)
    m_i = tl.full((16,), float("-inf"), dtype=tl.float32)
    l_i = tl.zeros((16,), dtype=tl.float32)
    acc = tl.zeros((16, 512), dtype=tl.float32)
    if my_start >= my_end:
        po = PO + (bid * NS + sid) * 16 * 512
        tl.store(po + h[:, None] * 512 + v[None, :], acc)
        tl.store(PLSE + (bid * NS + sid) * 16 + h, m_i); return
    for tile_start in range(my_start, my_end, BN):
        ti = tile_start + tl.arange(0, BN); tm = ti < my_end
        qk = tl.zeros((16, BN), dtype=tl.float32)
        for dk in range(0, 576, 64):
            do = dk + tl.arange(0, 64)
            qq = tl.load(Q + bid * stride_qb + h[:, None] * stride_qh + do[None, :])
            kk = tl.load(KV + ti[:, None] * stride_kvt + do[None, :],
                         mask=tm[:, None], other=0.0)
            qk += tl.dot(qq, tl.trans(kk.to(qq.dtype)))
        qk = qk * SM_SCALE; qk = tl.where(tm[None, :], qk, float("-inf"))
        m_ij = tl.max(qk, axis=1); m_new = tl.maximum(m_i, m_ij)
        alpha = tl.math.exp2((m_i - m_new) * LOG2E)
        p = tl.math.exp2((qk - m_new[:, None]) * LOG2E)
        l_i = alpha * l_i + tl.sum(p, axis=1); acc = acc * alpha[:, None]
        vv = tl.load(KV + ti[:, None] * stride_kvt + v[None, :],
                     mask=tm[:, None], other=0.0)
        acc += tl.dot(p.to(vv.dtype), vv); m_i = m_new
    acc = acc / l_i[:, None]; lse = m_i + tl.log(l_i)
    po = PO + (bid * NS + sid) * 16 * 512
    tl.store(po + h[:, None] * 512 + v[None, :], acc)
    tl.store(PLSE + (bid * NS + sid) * 16 + h, lse)


@triton.jit
def _reduce(PO, PLSE, O, stride_ob, stride_oh, NS: tl.constexpr):
    b = tl.program_id(0); h = tl.program_id(1); v = tl.arange(0, 512)
    gm = tl.full((1,), float("-inf"), dtype=tl.float32)
    for s in tl.static_range(NS):
        gm = tl.maximum(gm, tl.load(PLSE + (b * NS + s) * 16 + h))
    acc = tl.zeros((512,), dtype=tl.float32); tw = tl.zeros((1,), dtype=tl.float32)
    for s in tl.static_range(NS):
        lse = tl.load(PLSE + (b * NS + s) * 16 + h)
        w = tl.math.exp2((lse - gm) * LOG2E)
        acc += w * tl.load(PO + (b * NS + s) * 16 * 512 + h * 512 + v); tw += w
    tl.store(O + b * stride_ob + h * stride_oh + v, (acc / tw).to(tl.bfloat16))


# ============================================================
# Direct ASM call with ALL buffers pre-cached (bypass wrapper)
# ============================================================
def _setup_direct_asm(bs, kv_len, total_kv, ns=32):
    key = (bs, kv_len)
    if key in _asm:
        return _asm[key]

    ki = torch.arange(total_kv, dtype=torch.int32, device="cuda")
    lpl = torch.full((bs,), kv_len, dtype=torch.int32, device="cuda")

    # Metadata
    info = get_mla_metadata_info_v1(bs, 1, 16, FP8, FP8,
        is_sparse=False, fast_mode=False, num_kv_splits=ns, intra_batch_mode=True)
    bufs = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
    wm, wi, wis, ri, rfm, rpm = bufs

    dqo = torch.arange(0, bs + 1, dtype=torch.int32, device="cuda")
    dkv = dqo * kv_len
    get_mla_metadata_v1(dqo, dkv, lpl, 16, 1, True, wm, wis, wi, ri, rfm, rpm,
        page_size=1, kv_granularity=16, max_seqlen_qo=1, uni_seqlen_qo=1,
        fast_mode=False, max_split_per_batch=ns, intra_batch_mode=True,
        dtype_q=FP8, dtype_kv=FP8)

    # Pre-allocate split buffers (wrapper allocates these every call!)
    sd = torch.empty((ns * bs, 16, 512), dtype=torch.float32, device="cuda")
    sl = torch.empty((ns * bs, 16, 1), dtype=torch.float32, device="cuda")

    _asm[key] = {
        'ki': ki, 'lpl': lpl,
        'wm': wm, 'wi': wi, 'wis': wis,
        'ri': ri, 'rfm': rfm, 'rpm': rpm,
        'sd': sd, 'sl': sl,
    }
    return _asm[key]


# ============================================================
# Dispatch
# ============================================================
@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
    global _qs_ready
    q, kv_data, qo_indptr, kv_indptr, config = data
    bs = config["batch_size"]
    kv_len = config["kv_seq_len"]
    total_kv = bs * kv_len
    o = _g("o", (bs, 16, 512), BF16)

    # === Triton bf16: bs<=4/kv<=1024 ===
    if bs <= 4 and kv_len <= 1024:
        kv = kv_data["bf16"].view(-1, 576); NS = 16
        po = _g("tpo", (bs, NS, 16, 512), torch.float32)
        pl = _g("tpl", (bs, NS, 16), torch.float32)
        _attn_bf16[(bs, NS)](q, kv, po, pl, kv_indptr,
            q.stride(0), q.stride(1), kv.stride(0), BN=64, NS=NS, num_warps=4)
        _reduce[(bs, 16)](po, pl, o, o.stride(0), o.stride(1), NS=NS, num_warps=4)
        return o

    kv8, kvs = kv_data["fp8"]
    kv_flat = kv8.view(-1, 576)

    # === Triton FP8 for small/medium ===
    if bs <= 4:
        NS = 64; nw = 4
    elif bs <= 32 and kv_len <= 1024:
        NS = 8; nw = 4
    elif bs <= 64 and kv_len <= 1024:
        NS = 4; nw = 8
    else:
        # === Direct ASM (bypass wrapper) for large cases ===
        q8 = _g("q8", q.shape, FP8)
        N = q.numel()
        _fused_quant[(N + 1023) // 1024,](q8, q, IV, N=N, B=1024, num_warps=4)
        qs = _g("qs", (1,), torch.float32)
        if not _qs_ready:
            qs.fill_(QS_V)
            _qs_ready = True

        kv4d = kv8.view(total_kv, 1, 1, -1)
        sc = _setup_direct_asm(bs, kv_len, total_kv)

        # Direct stage1 + reduce (bypass mla_decode_fwd wrapper)
        mla_decode_stage1_asm_fwd(
            q8.view(-1, 16, 576), kv4d,
            qo_indptr, kv_indptr, sc['ki'], sc['lpl'],
            None,  # non-persistent indptr (None = persistent mode)
            sc['wm'], sc['wi'], sc['wis'],
            1, 1, 1, SM,
            sc['sd'], sc['sl'], o,
            q_scale=qs, kv_scale=kvs,
        )
        mla_reduce_v1(sc['sd'], sc['sl'], sc['ri'], sc['rfm'], sc['rpm'], 1, o, None)
        return o

    # Triton FP8 path
    po = _g(f"po_{bs}_{NS}", (bs, NS, 16, 512), torch.float32)
    pl = _g(f"pl_{bs}_{NS}", (bs, NS, 16), torch.float32)
    _attn_fp8[(bs, NS)](q, kv_flat, po, pl, kv_indptr, kvs,
        q.stride(0), q.stride(1), kv_flat.stride(0),
        BN=128, NS=NS, num_warps=nw)
    _reduce[(bs, 16)](po, pl, o, o.stride(0), o.stride(1), NS=NS, num_warps=4)
    return o
scrolls · 260 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