submission 754086
PromptForcePrime · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 220 lines, June 9 Researcher Reciprocity License v1.0.
solution.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-754086?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:9461a2065034781f89c887ebeb7c6c13e1f1d2551088e8137a226af0db91ef45
license declaredunknown
license concludedunknown
authorsPromptForcePrime
imported2026-08-15
Kernel source
solution.py220 lines
"""
solution.py — v94: right-sized workspace (ws_ms=ns).
v93 had ws_ms=32/64 but actual splits (ns) often just 1.
This caused 32x oversized logits_buf/lse_buf (e.g. 256MB vs 8MB).
v94 sets ws_ms=ns so _meta_info sizes workspace exactly for the
actual split count, reducing allocation overhead and cache waste.
"""
import sys as _sys
import torch
_NUM_CUS = 304
_MAX_TOKS_PER_SPLIT = 2048
_MAX_KV_TOKENS = 256 * 8192
_mla_fwd = None
_meta_info = None
_meta_fn = None
_stage1 = None
_reduce = None
def _ensure_aiter():
global _mla_fwd, _meta_info, _meta_fn, _stage1, _reduce
if _mla_fwd is not None:
return
from aiter.mla import mla_decode_fwd
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
_mla_fwd = mla_decode_fwd
_meta_info = get_mla_metadata_info_v1
_meta_fn = get_mla_metadata_v1
try:
import aiter as _a
s1 = getattr(_a, 'mla_decode_stage1_asm_fwd', None)
rd = getattr(_a, 'mla_reduce_v1', None)
if callable(s1) and callable(rd):
_stage1 = s1
_reduce = rd
except Exception:
pass
_cache = {}
_kvi = {}
_qs = None
_direct_ok = True
_last_q_ptr = {}
def _get_config(bs, kvl):
if bs <= 8:
ps, ibm = 1, True
elif kvl >= 4096:
ps, ibm = 8, False
else:
ps, ibm = 2, False
if 32 <= bs < 256:
ns = 1
elif bs >= 256:
if kvl >= 4096:
ns = max(1, (kvl + _MAX_TOKS_PER_SPLIT - 1) // _MAX_TOKS_PER_SPLIT)
else:
ns = 1
else:
ms = 64 if (bs <= 8 and kvl >= 4096) else 32
tok_ceil = kvl // 64
if kvl >= 4096:
ns = min(tok_ceil, ms)
else:
one_wave = (_NUM_CUS + bs - 1) // bs
half_wave = max(1, one_wave // 2)
ns = min(tok_ceil, ms)
for s in (1, 2, 4, 8, 16):
if s >= half_wave and s <= tok_ceil and s <= ms:
ns = s
break
gran = max(1, (64 if kvl >= 4096 else 16) // ps)
return ps, ibm, ns, gran
def _aiter_mla_decode(q, kv_fp8, kv_scale, qo_indptr, kv_indptr,
bs, nh, nkv, dq, dv, sm, qsl):
global _qs, _direct_ok
_ensure_aiter()
fp8d = kv_fp8.dtype
tkv = kv_fp8.shape[0]
kvl = tkv // bs
ck = (bs, tkv)
if ck not in _cache:
ps, ibm, ns, gran = _get_config(bs, kvl)
if ps not in _kvi:
_kvi[ps] = torch.arange(_MAX_KV_TOKENS // ps,
dtype=torch.int32, device=q.device)
if _qs is None:
_qs = torch.ones(1, dtype=torch.float32, device="cuda")
info = _meta_info(
bs, qsl, nh, fp8d, fp8d,
is_sparse=False, fast_mode=False,
num_kv_splits=ns, intra_batch_mode=ibm,
)
wt = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
wm, wi, wis, ri, rfm, rpm = wt
tq = q.shape[0]
ot = torch.empty((tq, nh, dv), dtype=torch.bfloat16, device="cuda")
sl = kv_indptr[1:] - kv_indptr[:-1]
lpl = ((sl - 1) % ps + 1).to(torch.int32)
ip = (kv_indptr // ps).to(torch.int32)
_meta_fn(
qo_indptr, ip, lpl,
nh // nkv, nkv, False,
wm, wis, wi, ri, rfm, rpm,
page_size=ps, kv_granularity=gran,
max_seqlen_qo=qsl, uni_seqlen_qo=qsl,
fast_mode=False, max_split_per_batch=ns,
intra_batch_mode=ibm,
dtype_q=fp8d, dtype_kv=fp8d,
)
qv_buf = torch.empty((tq, nh, dq), dtype=fp8d, device="cuda")
rpm_rows = rpm.size(0)
logits_buf = torch.empty((rpm_rows, 1, nh, dv),
dtype=torch.float32, device="cuda")
lse_buf = torch.empty((rpm_rows, 1, nh, 1),
dtype=torch.float32, device="cuda")
_cache[ck] = (ps, ns, _kvi[ps], wm, wi, wis, ri, rfm, rpm,
ot, lpl, ip, ibm, fp8d, qv_buf, logits_buf, lse_buf)
(ps, ns, kvi, wm, wi, wis, ri, rfm, rpm,
o, lpl, ip, ibm, fp8d, qv_buf, logits_buf, lse_buf) = _cache[ck]
kv4d = kv_fp8.view(tkv // ps, ps, nkv, dq)
q_ptr = q.data_ptr()
if _last_q_ptr.get(ck) != q_ptr:
qv_buf.copy_(q.view(-1, nh, dq))
_last_q_ptr[ck] = q_ptr
# --- Direct ops (primary path) ---
if _stage1 is not None and _direct_ok:
try:
_stage1(
qv_buf, kv4d,
qo_indptr, ip, kvi, lpl,
None, wm, wi, wis,
qsl, ps, nkv, sm,
logits_buf, lse_buf, o,
_qs, kv_scale,
)
_reduce(
logits_buf, lse_buf,
ri, rfm, rpm,
qsl, o, None,
)
return o
except Exception as e:
_direct_ok = False
print(f"[v92] direct ops failed: {type(e).__name__}: {e}",
file=_sys.stderr)
# --- Wrapper fallback ---
_mla_fwd(
qv_buf, kv4d, o,
qo_indptr, ip, kvi, lpl,
qsl,
page_size=ps, nhead_kv=nkv,
sm_scale=sm, logit_cap=0.0,
num_kv_splits=ns,
q_scale=_qs, kv_scale=kv_scale,
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 o
def _naive_fallback(q, kv_data, qo_indptr, kv_indptr, config):
kv = kv_data["bf16"]
bs = config["batch_size"]
sm = config["sm_scale"]
dv = config["v_head_dim"]
out = []
for b in range(bs):
qs, qe = int(qo_indptr[b].item()), int(qo_indptr[b + 1].item())
ks, ke = int(kv_indptr[b].item()), int(kv_indptr[b + 1].item())
qb = q[qs:qe].float()
kb = kv[ks:ke, 0, :].float()
sc = torch.einsum("qhd,kd->qhk", qb, kb) * sm
at = torch.softmax(sc, dim=-1)
out.append(torch.einsum("qhk,kd->qhd", at, kb[:, :dv]))
return torch.cat(out, dim=0).to(torch.bfloat16)
def custom_kernel(data):
q, kv_data, qo_indptr, kv_indptr, config = data
if q.is_cuda and "fp8" in kv_data:
fp8_pair = kv_data["fp8"]
if isinstance(fp8_pair, (tuple, list)) and len(fp8_pair) == 2:
try:
return _aiter_mla_decode(
q, fp8_pair[0], fp8_pair[1],
qo_indptr, kv_indptr,
config["batch_size"], config["num_heads"],
config["num_kv_heads"], config["qk_head_dim"],
config["v_head_dim"], config["sm_scale"],
config["q_seq_len"],
)
except Exception as e:
print(f"[aiter fail] {type(e).__name__}: {e}", file=_sys.stderr)
return _naive_fallback(q, kv_data, qo_indptr, kv_indptr, config)
scrolls · 220 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