submission 702905
Bortlesboat · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 96 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-702905?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:3c4a1917d07681f7c1bbcf75d108cf3cffee4d256e110fc7e8c5f4663156aec3
license declaredunknown
license concludedunknown
authorsBortlesboat
imported2026-08-15
Kernel source
submission.py96 lines
"""V606: V590 + JIT prewarm + a8w8 pg1 for (64,1024) only."""
import os, sys
os.environ.setdefault("HIP_FORCE_DEV_KERNARG", "1")
import torch
from task import input_t, output_t
import aiter
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
NH, NKV, QD, VD = 16, 1, 576, 512
SM = 1.0 / (QD ** 0.5)
FP8 = aiter_dtypes.fp8; _fi = torch.finfo(FP8)
FIXED_Q_SCALE = torch.tensor([4.0 / _fi.max], dtype=torch.float32, device="cuda")
_sfq = getattr(aiter, 'static_per_tensor_quant', None)
def qfp8_fast(t):
if _sfq:
try:
r = _sfq(t, FIXED_Q_SCALE)
return (r, FIXED_Q_SCALE) if not isinstance(r, tuple) else r
except: pass
return (t / FIXED_Q_SCALE).clamp(min=_fi.min, max=_fi.max).to(FP8), FIXED_Q_SCALE
_qc_id = None; _qc_numel = 0; _qc_result = None; _qc_ref = None
def _get_cached_qfp8(q):
global _qc_id, _qc_numel, _qc_result, _qc_ref
qid = id(q)
if qid == _qc_id and q.numel() == _qc_numel and _qc_ref is q: return _qc_result
r = qfp8_fast(q)
_qc_id = qid; _qc_numel = q.numel(); _qc_result = r; _qc_ref = q; return r
_DISPATCH = {
(4, 1024): (0, 0, 2), (4, 8192): (3, 16, 8),
(32, 1024): (1, 8, 2), (64, 1024): (3, 1, 2),
(32, 8192): (3, 16, 8), (64, 8192): (3, 32, 8),
(256, 1024): (2, 1, 2), (256, 8192): (3, 16, 8),
}
BF16_AITER = {(32, 1024)}
def _bmm(q, kv_data, cfg):
bs, kl = cfg["batch_size"], cfg["kv_seq_len"]
kv = kv_data["bf16"].view(bs, kl, QD); q3 = q.view(bs, NH, QD)
s = torch.bmm(q3, kv.transpose(1, 2)) * SM; a = torch.softmax(s, dim=-1)
return torch.bmm(a, kv[:, :, :VD]).bfloat16().contiguous()
def _prewarm(ps):
try:
bs, kl = 1, ps * 4
q_bf16 = torch.randn(bs, NH, QD, dtype=torch.bfloat16, device="cuda")
qi, qs = qfp8_fast(q_bf16)
kf = torch.randn(bs * kl, 1, QD, dtype=torch.bfloat16, device="cuda").to(FP8)
npp = kl // ps
qoi = torch.tensor([0, 1], dtype=torch.int32, device="cuda")
kpi = torch.tensor([0, npp], dtype=torch.int32, device="cuda")
ki = torch.arange(npp, dtype=torch.int32, device="cuda")
klp = torch.tensor([ps], dtype=torch.int32, device="cuda")
info = get_mla_metadata_info_v1(bs, 1, NH, qi.dtype, kf.dtype, is_sparse=False, fast_mode=True, num_kv_splits=1, intra_batch_mode=True)
w = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
wm, wi, wis, ri, rfm, rpm = w
get_mla_metadata_v1(qoi, kpi, klp, NH, NKV, True, wm, wis, wi, ri, rfm, rpm, page_size=ps, kv_granularity=max(ps, 16), max_seqlen_qo=1, uni_seqlen_qo=1, fast_mode=True, max_split_per_batch=1, intra_batch_mode=True, dtype_q=qi.dtype, dtype_kv=kf.dtype)
o = torch.empty(bs, NH, VD, dtype=torch.bfloat16, device="cuda")
lg = torch.empty(1, 1, NH, VD, dtype=torch.float32, device="cuda")
lse = torch.empty(1, 1, NH, 1, dtype=torch.float32, device="cuda")
aiter.mla_decode_stage1_asm_fwd(qi.view(-1, NH, QD), kf.view(-1, ps, NKV, QD), qoi, kpi, ki, klp, None, wm, wi, wis, 1, ps, NKV, SM, lg, lse, o, qs, torch.ones(1, dtype=torch.float32, device="cuda"))
torch.cuda.synchronize()
except Exception as e:
print(f"[V606] prewarm pg{ps}: {e}", file=sys.stderr)
for _ps in [1, 2, 4, 8]: _prewarm(_ps)
_c = {}
def _build(bs, ql, kl, ns, qd, kd, qoi, kvi, ps):
k = (bs, ql, kl, ns, qd, kd, ps)
if k in _c: return _c[k]
tq = bs * ql; npp = (kl + ps - 1) // ps; tp = bs * npp
ki = torch.arange(tp, dtype=torch.int32, device="cuda")
lpl = kl % ps; lpl = ps if lpl == 0 else lpl
klp = torch.full((bs,), lpl, dtype=torch.int32, device="cuda")
kpi = torch.arange(0, (bs + 1) * npp, npp, dtype=torch.int32, device="cuda")
info = get_mla_metadata_info_v1(bs, ql, NH, qd, kd, is_sparse=False, fast_mode=True, num_kv_splits=ns, intra_batch_mode=True)
w = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
wm, wi, wis, ri, rfm, rpm = w
get_mla_metadata_v1(qoi, kpi, klp, NH // NKV, NKV, True, wm, wis, wi, ri, rfm, rpm, page_size=ps, kv_granularity=max(ps, 16), max_seqlen_qo=ql, uni_seqlen_qo=ql, fast_mode=True, max_split_per_batch=ns, intra_batch_mode=True, dtype_q=qd, dtype_kv=kd)
np_ = rpm.size(0) * ql
_c[k] = dict(ki=ki, klp=klp, kpi=kpi, wm=wm, wi=wi, wis=wis, ri=ri, rfm=rfm, rpm=rpm, lg=torch.empty((np_, 1, NH, VD), dtype=torch.float32, device="cuda"), lse=torch.empty((np_, 1, NH, 1), dtype=torch.float32, device="cuda"), o=torch.empty((tq, NH, VD), dtype=torch.bfloat16, device="cuda"), mq=ql, ps=ps)
return _c[k]
def custom_kernel(data):
q, kv_data, qoi, kvi, cfg = data
bs = cfg["batch_size"]; kl = cfg["kv_seq_len"]
mode, ns, ps = _DISPATCH[(bs, kl)]
if mode == 0: return _bmm(q, kv_data, cfg)
if mode == 1:
kf = kv_data["bf16"]; ks = None; qi = q; qs = None
elif mode == 2:
kf, ks = kv_data["fp8"]; qi = q; qs = None
else:
kf, ks = kv_data["fp8"]; qi, qs = _get_cached_qfp8(q)
c = _build(bs, cfg["q_seq_len"], kl, ns, qi.dtype, kf.dtype, qoi, kvi, ps)
k4 = kf.view(-1, ps, NKV, kf.shape[-1]); o = c["o"]
aiter.mla_decode_stage1_asm_fwd(qi.view(-1, NH, QD), k4, qoi, c["kpi"], c["ki"], c["klp"], None, c["wm"], c["wi"], c["wis"], c["mq"], ps, NKV, SM, c["lg"], c["lse"], o, qs, ks)
if ns > 1:
aiter.mla_reduce_v1(c["lg"], c["lse"], c["ri"], c["rfm"], c["rpm"], c["mq"], o, None)
return o
scrolls · 96 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