submission 742585
lgc0338 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 260 lines, June 9 Researcher Reciprocity License v1.0.
submission_monkey.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-742585?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:d4c2f622c027ea2f6ad837007ea1565b35d7e8f9c9f83b0cc6fbb728c06c54cf
license declaredunknown
license concludedunknown
authorslgc0338
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
qk += tl.dot(qq, tl.trans(kk))num-warps = 4
q.stride(0), q.stride(1), kv.stride(0), BN=64, NS=NS, num_warps=4)online-softmax
m_ij = tl.max(qk, axis=1); m_new = tl.maximum(m_i, m_ij)persistent-kernel
None, # non-persistent indptr (None = persistent mode)tile-n = 64
q.stride(0), q.stride(1), kv.stride(0), BN=64, NS=NS, num_warps=4)Kernel source
submission_monkey.py260 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""Monkey-patch: bypass mla_decode_fwd wrapper overhead.
Key findings from aiter source analysis:
1. Wrapper allocates splitData/splitLse EVERY call (~10μs)
2. Wrapper does module lookup every call
3. Wrapper has Python if/elif dispatch logic
Our patch: call stage1_asm + reduce directly with ALL buffers pre-cached.
Previous attempts said "3-6% slower" — but those may not have cached everything properly.
Combined with: Triton for small cases, patched ASM for large cases.
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t
from aiter import mla_decode_stage1_asm_fwd, mla_reduce_v1
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
FP8 = aiter_dtypes.fp8
FI = torch.finfo(FP8)
QS_V = 5.0 / FI.max
IV = float(1.0 / QS_V)
SM_SCALE = tl.constexpr(1.0 / (576 ** 0.5))
LOG2E = tl.constexpr(1.44269504)
SM = 1.0 / (576 ** 0.5)
BF16 = torch.bfloat16
_c = {}
_asm = {}
_qs_ready = False
def _g(k, s, d):
if k not in _c or _c[k].shape != s:
_c[k] = torch.empty(s, dtype=d, device="cuda")
return _c[k]
# ============================================================
# Triton kernels (same as final, proven fast for small cases)
# ============================================================
@triton.jit
def _fused_quant(O, I, iv, N: tl.constexpr, B: tl.constexpr):
pid = tl.program_id(0)
offs = pid * B + tl.arange(0, B)
mask = offs < N
x = tl.load(I + offs, mask=mask).to(tl.float32) * iv
tl.store(O + offs, tl.clamp(x, -240.0, 240.0).to(O.dtype.element_ty), mask=mask)
@triton.jit
def _attn_fp8(
Q, KV, PO, PLSE, kv_indptr, kv_scale_ptr,
stride_qb, stride_qh, stride_kvt,
BN: tl.constexpr, NS: tl.constexpr,
):
bid = tl.program_id(0); sid = tl.program_id(1)
ks = tl.load(kv_indptr + bid); ke = tl.load(kv_indptr + bid + 1)
split_size = tl.cdiv(ke - ks, NS)
my_start = ks + sid * split_size
my_end = tl.minimum(my_start + split_size, ke)
h = tl.arange(0, 16); v = tl.arange(0, 512)
m_i = tl.full((16,), float("-inf"), dtype=tl.float32)
l_i = tl.zeros((16,), dtype=tl.float32)
acc = tl.zeros((16, 512), dtype=tl.float32)
kvs = tl.load(kv_scale_ptr); eff_scale = SM_SCALE * kvs
if my_start >= my_end:
po = PO + (bid * NS + sid) * 16 * 512
tl.store(po + h[:, None] * 512 + v[None, :], acc)
tl.store(PLSE + (bid * NS + sid) * 16 + h, m_i); return
for tile_start in range(my_start, my_end, BN):
ti = tile_start + tl.arange(0, BN); tm = ti < my_end
qk = tl.zeros((16, BN), dtype=tl.float32)
for dk in range(0, 576, 64):
do = dk + tl.arange(0, 64)
qq = tl.load(Q + bid * stride_qb + h[:, None] * stride_qh + do[None, :])
kk = tl.load(KV + ti[:, None] * stride_kvt + do[None, :],
mask=tm[:, None], other=0.0).to(tl.bfloat16)
qk += tl.dot(qq, tl.trans(kk))
qk = qk * eff_scale; qk = tl.where(tm[None, :], qk, float("-inf"))
m_ij = tl.max(qk, axis=1); m_new = tl.maximum(m_i, m_ij)
alpha = tl.math.exp2((m_i - m_new) * LOG2E)
p = tl.math.exp2((qk - m_new[:, None]) * LOG2E)
l_i = alpha * l_i + tl.sum(p, axis=1); acc = acc * alpha[:, None]
vv = tl.load(KV + ti[:, None] * stride_kvt + v[None, :],
mask=tm[:, None], other=0.0).to(tl.bfloat16)
acc += tl.dot(p.to(tl.bfloat16), vv); m_i = m_new
acc = (acc * kvs) / l_i[:, None]; lse = m_i + tl.log(l_i)
po = PO + (bid * NS + sid) * 16 * 512
tl.store(po + h[:, None] * 512 + v[None, :], acc)
tl.store(PLSE + (bid * NS + sid) * 16 + h, lse)
@triton.jit
def _attn_bf16(
Q, KV, PO, PLSE, kv_indptr,
stride_qb, stride_qh, stride_kvt,
BN: tl.constexpr, NS: tl.constexpr,
):
bid = tl.program_id(0); sid = tl.program_id(1)
h = tl.arange(0, 16); v = tl.arange(0, 512)
ks = tl.load(kv_indptr + bid); ke = tl.load(kv_indptr + bid + 1)
split_size = tl.cdiv(ke - ks, NS)
my_start = ks + sid * split_size
my_end = tl.minimum(my_start + split_size, ke)
m_i = tl.full((16,), float("-inf"), dtype=tl.float32)
l_i = tl.zeros((16,), dtype=tl.float32)
acc = tl.zeros((16, 512), dtype=tl.float32)
if my_start >= my_end:
po = PO + (bid * NS + sid) * 16 * 512
tl.store(po + h[:, None] * 512 + v[None, :], acc)
tl.store(PLSE + (bid * NS + sid) * 16 + h, m_i); return
for tile_start in range(my_start, my_end, BN):
ti = tile_start + tl.arange(0, BN); tm = ti < my_end
qk = tl.zeros((16, BN), dtype=tl.float32)
for dk in range(0, 576, 64):
do = dk + tl.arange(0, 64)
qq = tl.load(Q + bid * stride_qb + h[:, None] * stride_qh + do[None, :])
kk = tl.load(KV + ti[:, None] * stride_kvt + do[None, :],
mask=tm[:, None], other=0.0)
qk += tl.dot(qq, tl.trans(kk.to(qq.dtype)))
qk = qk * SM_SCALE; qk = tl.where(tm[None, :], qk, float("-inf"))
m_ij = tl.max(qk, axis=1); m_new = tl.maximum(m_i, m_ij)
alpha = tl.math.exp2((m_i - m_new) * LOG2E)
p = tl.math.exp2((qk - m_new[:, None]) * LOG2E)
l_i = alpha * l_i + tl.sum(p, axis=1); acc = acc * alpha[:, None]
vv = tl.load(KV + ti[:, None] * stride_kvt + v[None, :],
mask=tm[:, None], other=0.0)
acc += tl.dot(p.to(vv.dtype), vv); m_i = m_new
acc = acc / l_i[:, None]; lse = m_i + tl.log(l_i)
po = PO + (bid * NS + sid) * 16 * 512
tl.store(po + h[:, None] * 512 + v[None, :], acc)
tl.store(PLSE + (bid * NS + sid) * 16 + h, lse)
@triton.jit
def _reduce(PO, PLSE, O, stride_ob, stride_oh, NS: tl.constexpr):
b = tl.program_id(0); h = tl.program_id(1); v = tl.arange(0, 512)
gm = tl.full((1,), float("-inf"), dtype=tl.float32)
for s in tl.static_range(NS):
gm = tl.maximum(gm, tl.load(PLSE + (b * NS + s) * 16 + h))
acc = tl.zeros((512,), dtype=tl.float32); tw = tl.zeros((1,), dtype=tl.float32)
for s in tl.static_range(NS):
lse = tl.load(PLSE + (b * NS + s) * 16 + h)
w = tl.math.exp2((lse - gm) * LOG2E)
acc += w * tl.load(PO + (b * NS + s) * 16 * 512 + h * 512 + v); tw += w
tl.store(O + b * stride_ob + h * stride_oh + v, (acc / tw).to(tl.bfloat16))
# ============================================================
# Direct ASM call with ALL buffers pre-cached (bypass wrapper)
# ============================================================
def _setup_direct_asm(bs, kv_len, total_kv, ns=32):
key = (bs, kv_len)
if key in _asm:
return _asm[key]
ki = torch.arange(total_kv, dtype=torch.int32, device="cuda")
lpl = torch.full((bs,), kv_len, dtype=torch.int32, device="cuda")
# Metadata
info = get_mla_metadata_info_v1(bs, 1, 16, FP8, FP8,
is_sparse=False, fast_mode=False, num_kv_splits=ns, intra_batch_mode=True)
bufs = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
wm, wi, wis, ri, rfm, rpm = bufs
dqo = torch.arange(0, bs + 1, dtype=torch.int32, device="cuda")
dkv = dqo * kv_len
get_mla_metadata_v1(dqo, dkv, lpl, 16, 1, True, wm, wis, wi, ri, rfm, rpm,
page_size=1, kv_granularity=16, max_seqlen_qo=1, uni_seqlen_qo=1,
fast_mode=False, max_split_per_batch=ns, intra_batch_mode=True,
dtype_q=FP8, dtype_kv=FP8)
# Pre-allocate split buffers (wrapper allocates these every call!)
sd = torch.empty((ns * bs, 16, 512), dtype=torch.float32, device="cuda")
sl = torch.empty((ns * bs, 16, 1), dtype=torch.float32, device="cuda")
_asm[key] = {
'ki': ki, 'lpl': lpl,
'wm': wm, 'wi': wi, 'wis': wis,
'ri': ri, 'rfm': rfm, 'rpm': rpm,
'sd': sd, 'sl': sl,
}
return _asm[key]
# ============================================================
# Dispatch
# ============================================================
@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
global _qs_ready
q, kv_data, qo_indptr, kv_indptr, config = data
bs = config["batch_size"]
kv_len = config["kv_seq_len"]
total_kv = bs * kv_len
o = _g("o", (bs, 16, 512), BF16)
# === Triton bf16: bs<=4/kv<=1024 ===
if bs <= 4 and kv_len <= 1024:
kv = kv_data["bf16"].view(-1, 576); NS = 16
po = _g("tpo", (bs, NS, 16, 512), torch.float32)
pl = _g("tpl", (bs, NS, 16), torch.float32)
_attn_bf16[(bs, NS)](q, kv, po, pl, kv_indptr,
q.stride(0), q.stride(1), kv.stride(0), BN=64, NS=NS, num_warps=4)
_reduce[(bs, 16)](po, pl, o, o.stride(0), o.stride(1), NS=NS, num_warps=4)
return o
kv8, kvs = kv_data["fp8"]
kv_flat = kv8.view(-1, 576)
# === Triton FP8 for small/medium ===
if bs <= 4:
NS = 64; nw = 4
elif bs <= 32 and kv_len <= 1024:
NS = 8; nw = 4
elif bs <= 64 and kv_len <= 1024:
NS = 4; nw = 8
else:
# === Direct ASM (bypass wrapper) for large cases ===
q8 = _g("q8", q.shape, FP8)
N = q.numel()
_fused_quant[(N + 1023) // 1024,](q8, q, IV, N=N, B=1024, num_warps=4)
qs = _g("qs", (1,), torch.float32)
if not _qs_ready:
qs.fill_(QS_V)
_qs_ready = True
kv4d = kv8.view(total_kv, 1, 1, -1)
sc = _setup_direct_asm(bs, kv_len, total_kv)
# Direct stage1 + reduce (bypass mla_decode_fwd wrapper)
mla_decode_stage1_asm_fwd(
q8.view(-1, 16, 576), kv4d,
qo_indptr, kv_indptr, sc['ki'], sc['lpl'],
None, # non-persistent indptr (None = persistent mode)
sc['wm'], sc['wi'], sc['wis'],
1, 1, 1, SM,
sc['sd'], sc['sl'], o,
q_scale=qs, kv_scale=kvs,
)
mla_reduce_v1(sc['sd'], sc['sl'], sc['ri'], sc['rfm'], sc['rpm'], 1, o, None)
return o
# Triton FP8 path
po = _g(f"po_{bs}_{NS}", (bs, NS, 16, 512), torch.float32)
pl = _g(f"pl_{bs}_{NS}", (bs, NS, 16), torch.float32)
_attn_fp8[(bs, NS)](q, kv_flat, po, pl, kv_indptr, kvs,
q.stride(0), q.stride(1), kv_flat.stride(0),
BN=128, NS=NS, num_warps=nw)
_reduce[(bs, 16)](po, pl, o, o.stride(0), o.stride(1), NS=NS, num_warps=4)
return o
scrolls · 260 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