Skip to content
KernelIndex
Search⌘K

submission 695691

dark4scope · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-695691?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
66.1µs
#291 of 766
2026-04-02

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:763dd83f6a79d0b6f0362fd6b7f7f65aecc18ca0c2f706b4429af5f0105688a7
license declaredunknown
license concludedunknown
authorsdark4scope
imported2026-08-26

Techniques

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

autotune@torch.compile(mode="max-autotune-no-cudagraphs")
mmaqk = tl.dot(q_c, tl.trans(kv_c))
split-k2. Triton fp8 split-K: bs>4, kv<=2048 — fused attention beats AITER for short kv
tile-n = 32BLOCK_N = 32

Kernel source

submission.py463 lines
"""
V3.2 Mixed MLA decode: Optimal hybrid routing.

Based on benchmark data, each path wins for specific (bs, kv_len) ranges:
1. bmm (torch.compile): bs<=4 — fastest for tiny batches (16-18us)
2. Triton fp8 split-K: bs>4, kv<=2048 — fused attention beats AITER for short kv
3. AITER fp8 ASM: bs>4, kv>2048 — hand-tuned ASM wins for bandwidth-heavy cases
"""

import math
import os
import sys
from pathlib import Path

import torch
from task import input_t, output_t

try:
    import triton
    import triton.language as tl
    _HAS_TRITON = True
except Exception:
    _HAS_TRITON = False

# ── MLA constants ──────────────────────────────────────────────────────────
NUM_HEADS = 16
NUM_KV_HEADS = 1
KV_LORA_RANK = 512
QK_ROPE_DIM = 64
QK_DIM = KV_LORA_RANK + QK_ROPE_DIM
V_DIM = KV_LORA_RANK
SM_SCALE = 1.0 / math.sqrt(QK_DIM)
SM_SCALE_LOG2 = SM_SCALE * math.log2(math.e)
PAGE_SIZE = 1
CU_NUM = 304
BLOCK_N = 32
MAX_SPLITS = 32
MIN_TOK_PER_SPLIT = 64

_FP8_DTYPES = tuple(
    d for d in (getattr(torch, "float8_e4m3fnuz", None),
                getattr(torch, "float8_e4m3fn", None)) if d is not None
)
FP8_DTYPE = _FP8_DTYPES[0] if _FP8_DTYPES else None
_ROOT = Path(__file__).resolve().parents[2]
_VENDORED_AITER = _ROOT / "refs" / "aiter"

# ── AITER lazy imports ─────────────────────────────────────────────────────
mla_decode_fwd = None
get_mla_metadata_info_v1 = None
get_mla_metadata_v1 = None
_HAS_AITER: bool | None = None

# ── Caches ─────────────────────────────────────────────────────────────────
_buf_cache: dict = {}
_split_cache: dict = {}
_meta_cache: dict = {}
_kv_indices_cache: dict = {}
_kv_last_page_cache: dict = {}
_triton_fp8_ok: bool | None = None
_triton_bf16_ok: bool | None = None


# ── AITER import ───────────────────────────────────────────────────────────
def _import_aiter() -> bool:
    global mla_decode_fwd, get_mla_metadata_info_v1, get_mla_metadata_v1
    global _HAS_AITER, FP8_DTYPE

    if _HAS_AITER is not None:
        return _HAS_AITER

    def _load():
        try:
            from aiter.mla import mla_decode_fwd as _f
            from aiter import dtypes as _d
            from aiter import get_mla_metadata_info_v1 as _i, get_mla_metadata_v1 as _v
            globals().update(mla_decode_fwd=_f, get_mla_metadata_info_v1=_i,
                             get_mla_metadata_v1=_v, FP8_DTYPE=_d.fp8)
            return True
        except Exception:
            return False

    if _load():
        _HAS_AITER = True
        return True
    if _VENDORED_AITER.exists():
        v = str(_VENDORED_AITER)
        if v not in sys.path:
            sys.path.insert(0, v)
        if _load():
            _HAS_AITER = True
            return True
    _HAS_AITER = False
    return False


# ── Split heuristics ───────────────────────────────────────────────────────
def _pick_triton_splits(bs: int, kv_len: int) -> int:
    key = ("t", bs, kv_len)
    if key in _split_cache:
        return _split_cache[key]
    target = max(CU_NUM, 256)
    raw = max(1, target // bs)
    s = 1
    while s < raw:
        s <<= 1
    s = min(s, MAX_SPLITS, max(1, kv_len // MIN_TOK_PER_SPLIT))
    s = max(1, s)
    _split_cache[key] = s
    return s


def _pick_aiter_splits(bs: int, kv_len: int) -> int:
    key = ("a", bs, kv_len)
    if key in _split_cache:
        return _split_cache[key]
    overhead = 84.1
    best_score, best = 0.0, 1
    for n in range(1, 33):
        util = bs * n / (math.ceil(bs * n / CU_NUM) * CU_NUM)
        eff = kv_len / (kv_len + overhead * n)
        score = util * eff
        if score > best_score:
            best_score, best = score, n
    best = min(best, max(1, kv_len // 128))
    best = max(1, best)
    _split_cache[key] = best
    return best


def _get_bufs(bs, ns, device):
    key = (bs, ns)
    if key not in _buf_cache:
        _buf_cache[key] = (
            torch.empty((bs, NUM_HEADS, ns, V_DIM), dtype=torch.bfloat16, device=device),
            torch.empty((bs, NUM_HEADS, ns), dtype=torch.float32, device=device),
        )
    return _buf_cache[key]


# ── Triton kernels ─────────────────────────────────────────────────────────
if _HAS_TRITON:

    @triton.jit
    def _xcd_remap(pid, total, NUM_XCDS: tl.constexpr = 8):
        per = (total + NUM_XCDS - 1) // NUM_XCDS
        tall = total % NUM_XCDS
        tall = tl.where(tall == 0, NUM_XCDS, tall)
        xcd = pid % NUM_XCDS
        loc = pid // NUM_XCDS
        return tl.where(
            xcd < tall,
            xcd * per + loc,
            tall * per + (xcd - tall) * (per - 1) + loc,
        )

    @triton.jit
    def _stage1(
        Q, KV, KV_SCALE, INDPTR, MO, ML,
        sq0, sq1, skv0,
        smo0, smo1, smo2,
        sml0, sml1, sml2,
        scale_log2,
        BS: tl.constexpr, BN: tl.constexpr, NS: tl.constexpr,
        NH: tl.constexpr, DC: tl.constexpr, DR: tl.constexpr,
        FP8: tl.constexpr,
    ):
        pid = tl.program_id(0)
        pid = _xcd_remap(pid, BS * NS)
        b = pid // NS
        s = pid % NS

        hoffs = tl.arange(0, NH)
        coffs = tl.arange(0, DC)
        roffs = DC + tl.arange(0, DR)

        q_c = tl.load(Q + b * sq0 + hoffs[:, None] * sq1 + coffs[None, :])
        q_r = tl.load(Q + b * sq0 + hoffs[:, None] * sq1 + roffs[None, :])

        kv0 = tl.load(INDPTR + b)
        seq = tl.load(INDPTR + b + 1) - kv0
        tps = tl.cdiv(seq, NS)
        ss = s * tps
        se = tl.minimum(ss + tps, seq)

        kv_s = 1.0
        if FP8:
            kv_s = tl.load(KV_SCALE).to(tl.float32)

        emax = tl.zeros([NH], dtype=tl.float32) - float("inf")
        esum = tl.zeros([NH], dtype=tl.float32)
        acc = tl.zeros([NH, DC], dtype=tl.float32)

        if se > ss:
            for n0 in range(ss, se, BN):
                n0 = tl.multiple_of(n0, BN)
                noffs = n0 + tl.arange(0, BN)
                mask = noffs < se
                tidx = kv0 + noffs

                kv_c = tl.load(KV + tidx[:, None] * skv0 + coffs[None, :],
                               mask=mask[:, None], other=0.0)
                kv_r = tl.load(KV + tidx[:, None] * skv0 + roffs[None, :],
                               mask=mask[:, None], other=0.0)

                if FP8:
                    kv_c = (kv_c.to(tl.float32) * kv_s).to(tl.bfloat16)
                    kv_r = (kv_r.to(tl.float32) * kv_s).to(tl.bfloat16)

                qk = tl.dot(q_c, tl.trans(kv_c))
                qk += tl.dot(q_r, tl.trans(kv_r))
                qk *= scale_log2
                qk = tl.where(mask[None, :], qk, float("-inf"))

                new_max = tl.maximum(tl.max(qk, 1), emax)
                rescale = tl.exp2(emax - new_max)
                p = tl.exp2(qk - new_max[:, None])

                acc = acc * rescale[:, None] + tl.dot(p.to(tl.bfloat16), kv_c)
                esum = esum * rescale + tl.sum(p, 1)
                emax = new_max

            tl.store(
                MO + b * smo0 + hoffs[:, None] * smo1 + s * smo2 + coffs[None, :],
                (acc / esum[:, None]).to(tl.bfloat16),
            )
            tl.store(
                ML + b * sml0 + hoffs * sml1 + s * sml2,
                emax + tl.log2(esum),
            )

    @triton.jit
    def _stage2(
        MO, ML, OUT, INDPTR,
        smo0, smo1, smo2, sml0, sml1, sml2, so0, so1,
        BS: tl.constexpr, NS: tl.constexpr,
        NH: tl.constexpr, DC: tl.constexpr,
    ):
        pid = tl.program_id(0)
        pid = _xcd_remap(pid, BS * NH)
        b = pid // NH
        h = pid % NH

        coffs = tl.arange(0, DC)
        seq = tl.load(INDPTR + b + 1) - tl.load(INDPTR + b)
        tps = tl.cdiv(seq, NS)

        emax = -float("inf")
        esum = 0.0
        acc = tl.zeros([DC], dtype=tl.float32)

        mo_base = b * smo0 + h * smo1
        ml_base = b * sml0 + h * sml1

        for si in range(NS):
            ss = si * tps
            se = tl.minimum(ss + tps, seq)
            if se > ss:
                v = tl.load(MO + mo_base + si * smo2 + coffs).to(tl.float32)
                lse = tl.load(ML + ml_base + si * sml2)
                new_max = tl.maximum(lse, emax)
                old_s = tl.exp2(emax - new_max)
                new_s = tl.exp2(lse - new_max)
                acc = acc * old_s + new_s * v
                esum = esum * old_s + new_s
                emax = new_max

        tl.store(OUT + b * so0 + h * so1 + coffs, (acc / esum).to(tl.bfloat16))


def _launch_kw(stage):
    kw = {"num_warps": 4, "num_stages": 1 if stage == 1 else 2}
    if hasattr(torch.version, "hip") and torch.version.hip:
        kw["waves_per_eu"] = 1 if stage == 1 else 4
        kw["matrix_instr_nonkdim"] = 16
        kw["kpack"] = 2
    return kw


_unit_scale_cache: dict = {}

def _unit_scale(device):
    key = device.index if device.index is not None else -1
    if key not in _unit_scale_cache:
        _unit_scale_cache[key] = torch.ones(1, dtype=torch.float32, device=device)
    return _unit_scale_cache[key]


def _run_triton(q, kv_flat, kv_scale, kv_indptr, bs, kv_len, is_fp8):
    ns = _pick_triton_splits(bs, kv_len)
    mo, ml = _get_bufs(bs, ns, q.device)
    out = torch.empty((bs, NUM_HEADS, V_DIM), dtype=torch.bfloat16, device=q.device)

    _stage1[(bs * ns,)](
        q, kv_flat, kv_scale, kv_indptr, mo, ml,
        q.stride(0), q.stride(1), kv_flat.stride(0),
        mo.stride(0), mo.stride(1), mo.stride(2),
        ml.stride(0), ml.stride(1), ml.stride(2),
        SM_SCALE_LOG2,
        BS=bs, BN=BLOCK_N, NS=ns,
        NH=NUM_HEADS, DC=KV_LORA_RANK, DR=QK_ROPE_DIM,
        FP8=is_fp8,
        **_launch_kw(1),
    )

    _stage2[(bs * NUM_HEADS,)](
        mo, ml, out, kv_indptr,
        mo.stride(0), mo.stride(1), mo.stride(2),
        ml.stride(0), ml.stride(1), ml.stride(2),
        out.stride(0), out.stride(1),
        BS=bs, NS=ns, NH=NUM_HEADS, DC=V_DIM,
        **_launch_kw(2),
    )
    return out


# ── BMM path ──────────────────────────────────────────────────────────────
@torch.compile(mode="max-autotune-no-cudagraphs")
def _bmm_inner(q, kv, sm_scale):
    scores = torch.bmm(q, kv.transpose(1, 2))
    scores *= sm_scale
    attn = torch.softmax(scores, dim=-1)
    return torch.bmm(attn, kv[:, :, :V_DIM])


def _bmm_path(q, kv_data, bs, kv_len):
    kv = kv_data["bf16"].reshape(bs, kv_len, QK_DIM)
    return _bmm_inner(q, kv, SM_SCALE)


# ── AITER fp8 path ────────────────────────────────────────────────────────
def quantize_fp8(tensor):
    finfo = torch.finfo(FP8_DTYPE)
    amax = tensor.abs().amax().clamp(min=1e-12)
    scale = amax / finfo.max
    fp8 = (tensor / scale).clamp(min=finfo.min, max=finfo.max).to(FP8_DTYPE)
    return fp8, scale.to(torch.float32).reshape(1)


def _get_aiter_meta(bs, kv_len, num_splits, q_dtype, kv_dtype, qo_indptr, kv_indptr):
    key = (bs, kv_len, num_splits, q_dtype, kv_dtype)
    if key not in _meta_cache:
        kv_last = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
        _kv_last_page_cache[(bs, kv_len)] = kv_last
        info = get_mla_metadata_info_v1(
            bs, 1, NUM_HEADS, q_dtype, kv_dtype,
            is_sparse=False, fast_mode=False,
            num_kv_splits=num_splits, intra_batch_mode=True,
        )
        work = [torch.empty(s, dtype=d, device="cuda") for s, d in info]
        wm, wi, ws, ri, rf, rp = work
        get_mla_metadata_v1(
            qo_indptr, kv_indptr, kv_last,
            NUM_HEADS // NUM_KV_HEADS, NUM_KV_HEADS, True,
            wm, ws, wi, ri, rf, rp,
            page_size=PAGE_SIZE, kv_granularity=max(PAGE_SIZE, 16),
            max_seqlen_qo=1, uni_seqlen_qo=1, fast_mode=False,
            max_split_per_batch=num_splits, intra_batch_mode=True,
            dtype_q=q_dtype, dtype_kv=kv_dtype,
        )
        _meta_cache[key] = dict(
            work_meta_data=wm, work_indptr=wi, work_info_set=ws,
            reduce_indptr=ri, reduce_final_map=rf, reduce_partial_map=rp,
        )
        ik = (bs, kv_len, num_splits)
        if ik not in _kv_indices_cache:
            _kv_indices_cache[ik] = torch.arange(
                int(kv_indptr[-1].item()), dtype=torch.int32, device="cuda"
            )
    return _meta_cache[key], _kv_indices_cache[(bs, kv_len, num_splits)]


def _aiter_path(q, kv_data, qo_indptr, kv_indptr, bs, kv_len):
    if not _import_aiter():
        return _bmm_path(q, kv_data, bs, kv_len)

    ns = _pick_aiter_splits(bs, kv_len)
    q_fp8, q_scale = quantize_fp8(q)
    kv_fp8, kv_scale = kv_data["fp8"]
    kv_4d = kv_fp8.view(kv_fp8.shape[0], PAGE_SIZE, NUM_KV_HEADS, kv_fp8.shape[-1])

    meta, kv_idx = _get_aiter_meta(
        bs, kv_len, ns, q_fp8.dtype, kv_fp8.dtype, qo_indptr, kv_indptr,
    )
    if (bs, kv_len) not in _kv_last_page_cache:
        _kv_last_page_cache[(bs, kv_len)] = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)

    o = torch.empty((q.shape[0], NUM_HEADS, V_DIM), dtype=torch.bfloat16, device="cuda")
    mla_decode_fwd(
        q_fp8.view(-1, NUM_HEADS, QK_DIM), kv_4d, o,
        qo_indptr, kv_indptr, kv_idx,
        _kv_last_page_cache[(bs, kv_len)], 1,
        page_size=PAGE_SIZE, nhead_kv=NUM_KV_HEADS,
        sm_scale=SM_SCALE, logit_cap=0.0,
        num_kv_splits=ns, q_scale=q_scale, kv_scale=kv_scale,
        intra_batch_mode=True, **meta,
    )
    return o


# ── Triton fp8 path with fallback ─────────────────────────────────────────
def _triton_path(q, kv_data, kv_indptr, bs, kv_len):
    global _triton_fp8_ok, _triton_bf16_ok

    if _triton_fp8_ok is not False:
        kv_fp8, kv_scale = kv_data["fp8"]
        kv_flat = kv_fp8.reshape(-1, QK_DIM)
        if _triton_fp8_ok is None:
            try:
                out = _run_triton(q, kv_flat, kv_scale, kv_indptr, bs, kv_len, True)
                _triton_fp8_ok = True
                return out
            except Exception as e:
                print(f"[MLA] Triton fp8 failed: {e}", file=sys.stderr)
                _triton_fp8_ok = False
        else:
            return _run_triton(q, kv_flat, kv_scale, kv_indptr, bs, kv_len, True)

    if _triton_bf16_ok is not False:
        kv_flat = kv_data["bf16"].reshape(-1, QK_DIM)
        scale = _unit_scale(q.device)
        if _triton_bf16_ok is None:
            try:
                out = _run_triton(q, kv_flat, scale, kv_indptr, bs, kv_len, False)
                _triton_bf16_ok = True
                return out
            except Exception as e:
                print(f"[MLA] Triton bf16 failed: {e}", file=sys.stderr)
                _triton_bf16_ok = False
        else:
            return _run_triton(q, kv_flat, scale, kv_indptr, bs, kv_len, False)

    return None


# ── Entry point ────────────────────────────────────────────────────────────
# Routing based on benchmark data:
#   bs<=4 → bmm (16-18us, unbeatable for tiny batches)
#   bs>4, kv<=2048 → Triton fp8 (30-36us, beats AITER for short sequences)
#   bs>4, kv>2048 → AITER fp8 ASM (101-298us, bandwidth-optimal for long sequences)
KV_TRITON_THRESHOLD = 2048

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

    # Tiny batches: bmm wins
    if bs <= 4:
        return _bmm_path(q, kv_data, bs, kv_len)

    if q.device.type != "cuda":
        return _bmm_path(q, kv_data, bs, kv_len)

    # Short KV: Triton fp8 split-K (fused, lower overhead)
    if kv_len <= KV_TRITON_THRESHOLD and _HAS_TRITON:
        out = _triton_path(q, kv_data, kv_indptr, bs, kv_len)
        if out is not None:
            return out

    # Long KV: AITER fp8 ASM (bandwidth-optimized)
    return _aiter_path(q, kv_data, qo_indptr, kv_indptr, bs, kv_len)
scrolls · 463 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