Skip to content
KernelIndex
Search⌘K

submission 733177

Shellmia0 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_20260406_v249a_kv8k.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-733177?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
35.4µs
#76 of 766
2026-04-05

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:65b7e55c8fb5bdec9d0e4f08976aae0f5213898cc81b5d0e326c655de1f36b1a
license declaredunknown
license concludedunknown
authorsShellmia0
imported2026-08-15

Techniques

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

mmas = tl.dot(q0, tl.trans(kv0)) + tl.dot(q1, tl.trans(kv1))
num-warps = 4KVLEN=kvlen, BKV=bkv, sm_scale=SM_SCALE, NS=ns, num_warps=4,

Kernel source

submission_20260406_v249a_kv8k.py189 lines
"""
v143: Optimal hybrid routing — Triton fp8 where it wins, aiter elsewhere.
Routing table (benchmark-validated):
- bs<=4, kv=1024: Triton NS=16 BKV=32 → 16μs
- bs<=4, kv=8192: Triton NS=16 BKV=64 → 34μs
- bs=32, kv=1024: Triton NS=16 BKV=32 → 27μs
- bs=32, kv=8192: aiter page8 gran=32 → 32μs
- bs>=64: aiter page1/page8 → 38-89μs
"""
import os as _os
import sys as _sys
_devnull_fd = _os.open(_os.devnull, _os.O_WRONLY)
_os.dup2(_devnull_fd, 2)
_sys.stderr = open(_os.devnull, 'w')

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

NUM_HEADS = 16
V_HEAD_DIM = 512
QK_HEAD_DIM = 576
SM_SCALE = QK_HEAD_DIM ** -0.5

from aiter.mla import mla_decode_fwd
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1

FP8_DTYPE = aiter_dtypes.fp8
_cache_triton = {}
_cache_aiter = {}


@triton.jit
def _fd_s1(
    Q, KV, KVsc, Part, Lse,
    stride_qb, stride_qh, stride_kvt,
    KVLEN: tl.constexpr, BKV: tl.constexpr, sm_scale,
    NS: tl.constexpr,
):
    pid = tl.program_id(0)
    sid = pid % NS
    bid = pid // NS
    kps = KVLEN // NS
    ks = sid * kps

    h = tl.arange(0, 16)
    d0 = tl.arange(0, 512)
    d1 = tl.arange(0, 64)

    q0 = tl.load(Q + bid * stride_qb + h[:, None] * stride_qh + d0[None, :]).to(tl.bfloat16)
    q1 = tl.load(Q + bid * stride_qb + h[:, None] * stride_qh + (512 + d1[None, :])).to(tl.bfloat16)

    sc = tl.load(KVsc).to(tl.float32)
    kvbase = bid * KVLEN * stride_kvt

    m = tl.full([16], value=-float('inf'), dtype=tl.float32)
    l = tl.zeros([16], dtype=tl.float32)
    a = tl.zeros([16, 512], dtype=tl.float32)

    rows = tl.arange(0, BKV)
    n_iters = kps // BKV

    for i in range(n_iters):
        o = ks + i * BKV
        kv0 = tl.load(KV + kvbase + (o + rows[:, None]) * stride_kvt + d0[None, :]).to(tl.bfloat16)
        kv1 = tl.load(KV + kvbase + (o + rows[:, None]) * stride_kvt + (512 + d1[None, :])).to(tl.bfloat16)
        s = tl.dot(q0, tl.trans(kv0)) + tl.dot(q1, tl.trans(kv1))
        s = s.to(tl.float32) * sc * sm_scale
        bm = tl.max(s, axis=1)
        mn = tl.maximum(m, bm)
        al = tl.exp(m - mn)
        p = tl.exp(s - mn[:, None])
        l = l * al + tl.sum(p, axis=1)
        m = mn
        a = a * al[:, None] + tl.dot(p.to(tl.bfloat16), kv0).to(tl.float32) * sc

    a = a / l[:, None]
    lse = m + tl.log(l)
    base = (bid * NS * 16 + sid * 16) * 512
    tl.store(Part + base + h[:, None] * 512 + d0[None, :], a.to(tl.bfloat16))
    tl.store(Lse + bid * NS * 16 + sid * 16 + h, lse)


@triton.jit
def _fd_red(Part, Lse, Out, NS: tl.constexpr):
    pid = tl.program_id(0)
    hid = pid % 16
    bid = pid // 16
    acc = tl.zeros([512], dtype=tl.float32)
    mg = -float('inf')
    lg = 0.0
    for s in range(NS):
        lse = tl.load(Lse + bid * NS * 16 + s * 16 + hid)
        part = tl.load(Part + (bid * NS * 16 + s * 16 + hid) * 512 + tl.arange(0, 512)).to(tl.float32)
        mn = tl.maximum(mg, lse)
        a = tl.exp(mg - mn)
        b = tl.exp(lse - mn)
        lg = lg * a + b
        mg = mn
        acc = acc * a + b * part
    acc = acc / lg
    tl.store(Out + bid * 16 * 512 + hid * 512 + tl.arange(0, 512), acc.to(tl.bfloat16))


def _run_triton(q, kv_fp8, kv_scale, bs, kvlen, dev, ns, bkv):
    key = ('tri', bs, kvlen, ns, bkv)
    if key not in _cache_triton:
        part = torch.empty(bs * ns * NUM_HEADS * V_HEAD_DIM, dtype=torch.bfloat16, device=dev)
        lse = torch.empty(bs * ns * NUM_HEADS, dtype=torch.float32, device=dev)
        out = torch.empty(bs, NUM_HEADS, V_HEAD_DIM, dtype=torch.bfloat16, device=dev)
        _cache_triton[key] = (part, lse, out)
    part, lse, out = _cache_triton[key]
    kv_flat = kv_fp8.view(bs * kvlen, QK_HEAD_DIM)
    _fd_s1[(bs * ns,)](
        q, kv_flat, kv_scale, part, lse,
        q.stride(0), q.stride(1), kv_flat.stride(0),
        KVLEN=kvlen, BKV=bkv, sm_scale=SM_SCALE, NS=ns, num_warps=4,
    )
    _fd_red[(bs * NUM_HEADS,)](part, lse, out, NS=ns)
    return out


def _build_aiter(bs, kvlen, dev, ps, gran, fm, nsplits):
    pc = kvlen // ps if ps > 1 else kvlen
    np_ = bs * pc if ps > 1 else bs * kvlen
    lpl = ps if ps > 1 else kvlen
    qoi = torch.arange(bs + 1, dtype=torch.int32, device=dev)
    kvi = torch.arange(bs + 1, dtype=torch.int32, device=dev) * pc
    klp = torch.full((bs,), lpl, dtype=torch.int32, device=dev)
    ki = torch.arange(np_, dtype=torch.int32, device=dev)
    out = torch.empty((bs, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=dev)
    info = get_mla_metadata_info_v1(bs, 1, NUM_HEADS, torch.bfloat16, FP8_DTYPE,
        is_sparse=False, fast_mode=fm, num_kv_splits=nsplits, intra_batch_mode=False)
    w = [torch.empty(s, dtype=d, device=dev) for s, d in info]
    get_mla_metadata_v1(qoi, kvi, klp, 16, 1, True, w[0], w[2], w[1], w[3], w[4], w[5],
        page_size=ps, kv_granularity=gran, max_seqlen_qo=1, uni_seqlen_qo=1,
        fast_mode=fm, max_split_per_batch=nsplits, intra_batch_mode=False,
        dtype_q=torch.bfloat16, dtype_kv=FP8_DTYPE)
    return qoi, kvi, klp, ki, out, w, np_, ps


def _run_aiter(q, kv_fp8, kv_scale, bs, kvlen, dev, ps, gran, fm, nsplits):
    key = ('ait', bs, kvlen, ps, gran, fm, nsplits)
    if key not in _cache_aiter:
        _cache_aiter[key] = _build_aiter(bs, kvlen, dev, ps, gran, fm, nsplits)
    qoi, kvi, klp, ki, out, w, np_, ps = _cache_aiter[key]
    kv4d = kv_fp8.view(np_, ps, 1, QK_HEAD_DIM)
    mla_decode_fwd(q, kv4d, out, qoi, kvi, ki, klp, 1,
        page_size=ps, nhead_kv=1, sm_scale=SM_SCALE, logit_cap=0.0,
        num_kv_splits=nsplits, q_scale=None, kv_scale=kv_scale,
        intra_batch_mode=False, work_meta_data=w[0], work_indptr=w[1],
        work_info_set=w[2], reduce_indptr=w[3], reduce_final_map=w[4],
        reduce_partial_map=w[5])
    return out


def custom_kernel(data):
    q, kv_data, qo_indptr, _, _ = data
    kv_fp8, kv_scale = kv_data["fp8"]
    bs = qo_indptr.numel() - 1
    kvlen = kv_fp8.shape[0] // bs
    dev = q.device

    # Triton fp8: wins for small batch + short kv
    if bs <= 4:
        if kvlen <= 1024:
            return _run_triton(q, kv_fp8, kv_scale, bs, kvlen, dev, ns=16, bkv=32)
        else:
            return _run_triton(q, kv_fp8, kv_scale, bs, kvlen, dev, ns=16, bkv=64)

    if bs <= 4 and kvlen <= 1024:
        return _run_triton(q, kv_fp8, kv_scale, bs, kvlen, dev, ns=16, bkv=32)
    if bs <= 32 and kvlen <= 1024:
        return _run_triton(q, kv_fp8, kv_scale, bs, kvlen, dev, ns=4, bkv=32)

    # v247: Triton for bs=64, pg2 for bs>=128
    if bs >= 128 and kvlen <= 1024 and kvlen % 2 == 0:
        ps, gran = 2, 16
        return _run_aiter(q, kv_fp8, kv_scale, bs, kvlen, dev, ps, gran, fm=False, nsplits=32)
    if bs >= 64 and kvlen <= 1024:
        return _run_triton(q, kv_fp8, kv_scale, bs, kvlen, dev, ns=4, bkv=32)
    if kvlen >= 8192:
        return _run_aiter(q, kv_fp8, kv_scale, bs, kvlen, dev, ps=8, gran=16, fm=False, nsplits=64)
    ps, gran = 1, 16
    fm = bs <= 32
    return _run_aiter(q, kv_fp8, kv_scale, bs, kvlen, dev, ps, gran, fm, 32)
scrolls · 189 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