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
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")mma
qk = tl.dot(q_c, tl.trans(kv_c))split-k
2. Triton fp8 split-K: bs>4, kv<=2048 — fused attention beats AITER for short kvtile-n = 32
BLOCK_N = 32Kernel 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