Skip to content
KernelIndex
Search⌘K

submission 720383

div22 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

solution_373.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-720383?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.3µs
#74 of 766
2026-04-04

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:c1742990b53892ca2e57649afb82f211ccde6f1a2192eeba4c4d644f48a97dc0
license declaredunknown
license concludedunknown
authorsdiv22
imported2026-08-15

Techniques

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

mmascores = tl.dot(Q1, tl.trans(K1)) + tl.dot(Q2, tl.trans(K2))
online-softmaxm_ij = tl.max(scores, axis=1); m_new = tl.maximum(m_i, m_ij)
persistent-kerneldef _aiter_persistent(q, kv_data, qo_indptr, config, ps, nks=None, kg=16, fast_mode=False, ibm_override=None):
tile-n = 32BM = triton.next_power_of_2(nh); BN = 32; sq = nh * 576; sh = 576

Kernel source

solution_373.py240 lines
"""
S370: S369 + Triton nw=4/nst=2 for bs>=32/kv=1K (s365 config).
s365 used nw=4/nst=2 for bs=32 and bs=64 → 24.6/37.2 vs default 26.7/40.7us.
"""
import math
import torch
import triton
import triton.language as tl
from task import input_t, output_t

from aiter.mla import mla_decode_fwd
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1

NUM_KV_HEADS = 1
SM_SCALE = 1.0 / (576 ** 0.5)
KV_GRANULARITY = 16
CU_NUM = 256
FP8_DTYPE = torch.float8_e4m3fn
FP8_MAX = torch.finfo(FP8_DTYPE).max


def _auto_nks(bs, kv_seq_len):
    avg_kv = kv_seq_len
    best_eff = -1.0
    best_nks = 1
    for i in range(1, 17):
        waves = math.ceil(bs * i / CU_NUM)
        eff = (bs * i / waves) * CU_NUM * avg_kv / (avg_kv + 84.1 * i)
        if eff > best_eff:
            best_eff = eff
            best_nks = i
    return best_nks


_meta_cache = {}
_out_cache = {}
_indices_cache = {}
_kv_lpl_cache = {}
_kv_indptr_page_cache = {}


def _ci(n):
    if n not in _indices_cache:
        _indices_cache[n] = torch.arange(n, dtype=torch.int32, device="cuda")
    return _indices_cache[n]


def _cl(bs, ps):
    k = (bs, ps)
    if k not in _kv_lpl_cache:
        _kv_lpl_cache[k] = torch.full((bs,), ps, dtype=torch.int32, device="cuda")
    return _kv_lpl_cache[k]


def _cp(bs, kvsl, ps):
    k = (bs, kvsl, ps)
    if k not in _kv_indptr_page_cache:
        ppb = kvsl // ps
        _kv_indptr_page_cache[k] = torch.arange(0, (bs + 1) * ppb, ppb, dtype=torch.int32, device="cuda")
    return _kv_indptr_page_cache[k]


def _co(tq, nh, tag=""):
    k = (tq, nh, tag)
    if k not in _out_cache:
        _out_cache[k] = torch.empty(tq, nh, 512, dtype=torch.bfloat16, device="cuda")
    return _out_cache[k]


def _get_cached_meta(batch_size, q_seq, nhead, q_dtype, kv_dtype, page_size,
                     qo_indptr, kv_indptr_page, kv_lpl, nks, kv_seq_len, kg=16, fast_mode=False, ibm_override=None):
    ibm = ibm_override if ibm_override is not None else (not fast_mode)
    key = (batch_size, q_seq, nhead, q_dtype, kv_dtype, page_size, nks, kv_seq_len, kg, fast_mode, ibm)
    if key not in _meta_cache:
        info = get_mla_metadata_info_v1(
            batch_size, q_seq, nhead, q_dtype, kv_dtype,
            is_sparse=False, fast_mode=fast_mode,
            num_kv_splits=nks, intra_batch_mode=ibm,
        )
        work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
        wm, wi, wis, ri, rfm, rpm = work
        get_mla_metadata_v1(
            qo_indptr, kv_indptr_page, kv_lpl,
            nhead // NUM_KV_HEADS, NUM_KV_HEADS, True,
            wm, wis, wi, ri, rfm, rpm,
            page_size=page_size, kv_granularity=max(page_size, kg),
            max_seqlen_qo=q_seq, uni_seqlen_qo=q_seq,
            fast_mode=fast_mode, max_split_per_batch=nks,
            intra_batch_mode=ibm, dtype_q=q_dtype, dtype_kv=kv_dtype,
        )
        _meta_cache[key] = work
    return _meta_cache[key]


def _aiter_persistent(q, kv_data, qo_indptr, config, ps, nks=None, kg=16, fast_mode=False, ibm_override=None):
    """Persistent mode bf16 Q + fp8 KV."""
    bs, nh, qsl, kvsl = config["batch_size"], config["num_heads"], config["q_seq_len"], config["kv_seq_len"]
    tq, tkv = bs * qsl, bs * kvsl
    if nks is None:
        nks = _auto_nks(bs, kvsl)
    np_ = tkv // ps
    kv_fp8, kv_s = kv_data["fp8"]
    kv_4d = kv_fp8.view(np_, ps, NUM_KV_HEADS, 576)
    output = _co(tq, nh, "p")
    ibm = ibm_override if ibm_override is not None else (not fast_mode)
    wm, wi, wis, ri, rfm, rpm = _get_cached_meta(
        bs, qsl, nh, q.dtype, kv_fp8.dtype, ps,
        qo_indptr, _cp(bs, kvsl, ps), _cl(bs, ps), nks, kvsl, kg, fast_mode, ibm_override)
    mla_decode_fwd(
        q.view(-1, nh, 576), kv_4d, output,
        qo_indptr, _cp(bs, kvsl, ps), _ci(np_), _cl(bs, ps), qsl,
        page_size=ps, nhead_kv=NUM_KV_HEADS, sm_scale=SM_SCALE,
        num_kv_splits=nks, q_scale=None, kv_scale=kv_s, intra_batch_mode=ibm,
        work_meta_data=wm, work_indptr=wi, work_info_set=wis,
        reduce_indptr=ri, reduce_final_map=rfm, reduce_partial_map=rpm,
    )
    return output


# ============ TRITON (gfx950: BLOCK_N=32, nw=8, ns=3) ============

@triton.jit
def _mla_stage1(
    Q_ptr, KV_ptr, PO_ptr, PM_ptr, PS_ptr, output_ptr, kv_indptr_ptr,
    stride_q_seq, stride_q_head, total_q, num_heads: tl.constexpr,
    q_seq_len, num_splits, sm_scale,
    IS_FUSED: tl.constexpr, IS_FP8: tl.constexpr,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
):
    pid_split = tl.program_id(0); qi = tl.program_id(1); bi = qi // q_seq_len
    kv_s = tl.load(kv_indptr_ptr + bi); kv_e = tl.load(kv_indptr_ptr + bi + 1)
    if IS_FUSED: cs = kv_s; ce = kv_e
    else:
        kv_len = kv_e - kv_s; chunk = tl.cdiv(kv_len, num_splits)
        cs = kv_s + pid_split * chunk; ce = tl.minimum(cs + chunk, kv_e)
    heads = tl.arange(0, BLOCK_M); d512 = tl.arange(0, 512); d64 = tl.arange(0, 64)
    mask_h = heads < num_heads; q_base = qi * stride_q_seq
    Q1 = tl.load(Q_ptr + q_base + heads[:, None] * stride_q_head + d512[None, :], mask=mask_h[:, None], other=0.0)
    Q2 = tl.load(Q_ptr + q_base + heads[:, None] * stride_q_head + 512 + d64[None, :], mask=mask_h[:, None], other=0.0)
    m_i = tl.full([BLOCK_M], float("-inf"), dtype=tl.float32)
    l_i = tl.zeros([BLOCK_M], dtype=tl.float32); acc = tl.zeros([BLOCK_M, 512], dtype=tl.float32)
    for start_n in range(cs, ce, BLOCK_N):
        n_offs = start_n + tl.arange(0, BLOCK_N); mask_n = n_offs < ce
        if IS_FP8:
            K1 = tl.load(KV_ptr + n_offs[:, None] * 576 + d512[None, :], mask=mask_n[:, None], other=0.0).to(tl.bfloat16)
            K2 = tl.load(KV_ptr + n_offs[:, None] * 576 + 512 + d64[None, :], mask=mask_n[:, None], other=0.0).to(tl.bfloat16)
        else:
            K1 = tl.load(KV_ptr + n_offs[:, None] * 576 + d512[None, :], mask=mask_n[:, None], other=0.0)
            K2 = tl.load(KV_ptr + n_offs[:, None] * 576 + 512 + d64[None, :], mask=mask_n[:, None], other=0.0)
        scores = tl.dot(Q1, tl.trans(K1)) + tl.dot(Q2, tl.trans(K2))
        scores = scores * sm_scale; scores = tl.where(mask_n[None, :], scores, float("-inf"))
        m_ij = tl.max(scores, axis=1); m_new = tl.maximum(m_i, m_ij)
        alpha = tl.math.exp2((m_i - m_new) * 1.44269504)
        p = tl.math.exp2((scores - m_new[:, None]) * 1.44269504)
        l_i = l_i * alpha + tl.sum(p, axis=1)
        acc = acc * alpha[:, None] + tl.dot(p.to(tl.bfloat16), K1); m_i = m_new
    if IS_FUSED:
        inv = 1.0 / tl.maximum(l_i, 1e-12); result = (acc * inv[:, None]).to(tl.bfloat16)
        out_base = qi * num_heads
        tl.store(output_ptr + (out_base + heads[:, None]) * 512 + d512[None, :], result, mask=mask_h[:, None])
    else:
        total_qh = total_q * num_heads; out_base = pid_split * total_qh + qi * num_heads
        tl.store(PM_ptr + out_base + heads, m_i, mask=mask_h)
        tl.store(PS_ptr + out_base + heads, l_i, mask=mask_h)
        tl.store(PO_ptr + (out_base + heads[:, None]) * 512 + d512[None, :], acc, mask=mask_h[:, None])


@triton.jit
def _mla_reduce(PO_ptr, PM_ptr, PS_ptr, output_ptr, total_qh, out_scale, num_splits: tl.constexpr):
    qh = tl.program_id(0); d512 = tl.arange(0, 512); gmax = float("-inf")
    for s in range(num_splits): gmax = tl.maximum(gmax, tl.load(PM_ptr + s * total_qh + qh))
    acc = tl.zeros([512], dtype=tl.float32); total_sum = 0.0
    for s in range(num_splits):
        idx = s * total_qh + qh
        rs = tl.math.exp2((tl.load(PM_ptr + idx) - gmax) * 1.44269504)
        total_sum += tl.load(PS_ptr + idx) * rs
        acc += tl.load(PO_ptr + idx * 512 + d512) * rs
    inv = 1.0 / tl.maximum(total_sum, 1e-12)
    tl.store(output_ptr + qh * 512 + d512, (acc * inv * out_scale).to(tl.bfloat16))


_ws_cache = {}


def _triton_kernel(q, kv_data, qo_indptr, kv_indptr, config, nw=8, nst=3):
    nh, qsl, kvsl, sm, bs = config["num_heads"], config["q_seq_len"], config["kv_seq_len"], config["sm_scale"], config["batch_size"]
    tq = bs * qsl; tqh = tq * nh; tkv = bs * kvsl; dev = q.device
    BM = triton.next_power_of_2(nh); BN = 32; sq = nh * 576; sh = 576
    ns = max(1, min(32, kvsl // 64, -(-CU_NUM // tq)))
    use_fp8 = tkv > 100000
    if use_fp8:
        kf, ks = kv_data["fp8"]; kv = kf.view(-1, 576); kvs = ks.item(); es = sm * kvs; os_val = kvs
    else:
        kv = kv_data["bf16"].view(-1, 576); es = sm; os_val = 1.0
    output = _co(tq, nh, "t")
    if ns == 1:
        _mla_stage1[(1, tq)](q, kv, None, None, None, output, kv_indptr,
            sq, sh, tq, nh, qsl, 1, es, True, use_fp8, BM, BN, num_warps=nw, num_stages=nst)
    else:
        wk = (ns, tqh)
        if wk not in _ws_cache:
            n = ns * tqh
            _ws_cache[wk] = (torch.empty(n, 512, dtype=torch.float32, device=dev),
                             torch.empty(n, dtype=torch.float32, device=dev),
                             torch.empty(n, dtype=torch.float32, device=dev))
        po, pm, ps = _ws_cache[wk]
        _mla_stage1[(ns, tq)](q, kv, po, pm, ps, None, kv_indptr,
            sq, sh, tq, nh, qsl, ns, es, False, use_fp8, BM, BN, num_warps=nw, num_stages=nst)
        _mla_reduce[(tqh,)](po, pm, ps, output, tqh, os_val, ns)
    return output


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    bs = config["batch_size"]
    kv_seq = config["kv_seq_len"]

    # bs=4/kv≤1K: Triton nw=8 nst=3 (default, proven best for small batch)
    if bs <= 4 and kv_seq <= 1024:
        return _triton_kernel(q, kv_data, qo_indptr, kv_indptr, config)

    # bs=32,64/kv≤1K: Triton nw=4 nst=2 (s365: 24.6/37.2 vs default 26.7/40.7)
    if bs <= 64 and kv_seq <= 1024:
        return _triton_kernel(q, kv_data, qo_indptr, kv_indptr, config, nw=4, nst=2)

    # bs=256/kv=1K: persistent ps=2 nks=1
    if bs >= 256 and kv_seq <= 1024:
        return _aiter_persistent(q, kv_data, qo_indptr, config, ps=2, nks=1)

    # bs=4/kv=8K: fast_mode=True (S289: -6us improvement)
    if bs <= 4:
        return _aiter_persistent(q, kv_data, qo_indptr, config, ps=8, fast_mode=True)

    # bs=32/kv=8K: ibm=False (s351: -2.8us improvement)
    if bs <= 32:
        return _aiter_persistent(q, kv_data, qo_indptr, config, ps=8, ibm_override=False)

    # kv>=8K: persistent ps=8 (ibm=True default - ibm=False HURTS bs>=64)
    return _aiter_persistent(q, kv_data, qo_indptr, config, ps=8)
scrolls · 240 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