submission 733177
Shellmia0 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 189 lines, June 9 Researcher Reciprocity License v1.0.
submission_20260406_v249a_kv8k.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-733177?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:65b7e55c8fb5bdec9d0e4f08976aae0f5213898cc81b5d0e326c655de1f36b1a
license declaredunknown
license concludedunknown
authorsShellmia0
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
s = tl.dot(q0, tl.trans(kv0)) + tl.dot(q1, tl.trans(kv1))num-warps = 4
KVLEN=kvlen, BKV=bkv, sm_scale=SM_SCALE, NS=ns, num_warps=4,Kernel source
submission_20260406_v249a_kv8k.py189 lines
"""
v143: Optimal hybrid routing — Triton fp8 where it wins, aiter elsewhere.
Routing table (benchmark-validated):
- bs<=4, kv=1024: Triton NS=16 BKV=32 → 16μs
- bs<=4, kv=8192: Triton NS=16 BKV=64 → 34μs
- bs=32, kv=1024: Triton NS=16 BKV=32 → 27μs
- bs=32, kv=8192: aiter page8 gran=32 → 32μs
- bs>=64: aiter page1/page8 → 38-89μs
"""
import os as _os
import sys as _sys
_devnull_fd = _os.open(_os.devnull, _os.O_WRONLY)
_os.dup2(_devnull_fd, 2)
_sys.stderr = open(_os.devnull, 'w')
import torch
import triton
import triton.language as tl
from task import input_t, output_t
NUM_HEADS = 16
V_HEAD_DIM = 512
QK_HEAD_DIM = 576
SM_SCALE = QK_HEAD_DIM ** -0.5
from aiter.mla import mla_decode_fwd
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
FP8_DTYPE = aiter_dtypes.fp8
_cache_triton = {}
_cache_aiter = {}
@triton.jit
def _fd_s1(
Q, KV, KVsc, Part, Lse,
stride_qb, stride_qh, stride_kvt,
KVLEN: tl.constexpr, BKV: tl.constexpr, sm_scale,
NS: tl.constexpr,
):
pid = tl.program_id(0)
sid = pid % NS
bid = pid // NS
kps = KVLEN // NS
ks = sid * kps
h = tl.arange(0, 16)
d0 = tl.arange(0, 512)
d1 = tl.arange(0, 64)
q0 = tl.load(Q + bid * stride_qb + h[:, None] * stride_qh + d0[None, :]).to(tl.bfloat16)
q1 = tl.load(Q + bid * stride_qb + h[:, None] * stride_qh + (512 + d1[None, :])).to(tl.bfloat16)
sc = tl.load(KVsc).to(tl.float32)
kvbase = bid * KVLEN * stride_kvt
m = tl.full([16], value=-float('inf'), dtype=tl.float32)
l = tl.zeros([16], dtype=tl.float32)
a = tl.zeros([16, 512], dtype=tl.float32)
rows = tl.arange(0, BKV)
n_iters = kps // BKV
for i in range(n_iters):
o = ks + i * BKV
kv0 = tl.load(KV + kvbase + (o + rows[:, None]) * stride_kvt + d0[None, :]).to(tl.bfloat16)
kv1 = tl.load(KV + kvbase + (o + rows[:, None]) * stride_kvt + (512 + d1[None, :])).to(tl.bfloat16)
s = tl.dot(q0, tl.trans(kv0)) + tl.dot(q1, tl.trans(kv1))
s = s.to(tl.float32) * sc * sm_scale
bm = tl.max(s, axis=1)
mn = tl.maximum(m, bm)
al = tl.exp(m - mn)
p = tl.exp(s - mn[:, None])
l = l * al + tl.sum(p, axis=1)
m = mn
a = a * al[:, None] + tl.dot(p.to(tl.bfloat16), kv0).to(tl.float32) * sc
a = a / l[:, None]
lse = m + tl.log(l)
base = (bid * NS * 16 + sid * 16) * 512
tl.store(Part + base + h[:, None] * 512 + d0[None, :], a.to(tl.bfloat16))
tl.store(Lse + bid * NS * 16 + sid * 16 + h, lse)
@triton.jit
def _fd_red(Part, Lse, Out, NS: tl.constexpr):
pid = tl.program_id(0)
hid = pid % 16
bid = pid // 16
acc = tl.zeros([512], dtype=tl.float32)
mg = -float('inf')
lg = 0.0
for s in range(NS):
lse = tl.load(Lse + bid * NS * 16 + s * 16 + hid)
part = tl.load(Part + (bid * NS * 16 + s * 16 + hid) * 512 + tl.arange(0, 512)).to(tl.float32)
mn = tl.maximum(mg, lse)
a = tl.exp(mg - mn)
b = tl.exp(lse - mn)
lg = lg * a + b
mg = mn
acc = acc * a + b * part
acc = acc / lg
tl.store(Out + bid * 16 * 512 + hid * 512 + tl.arange(0, 512), acc.to(tl.bfloat16))
def _run_triton(q, kv_fp8, kv_scale, bs, kvlen, dev, ns, bkv):
key = ('tri', bs, kvlen, ns, bkv)
if key not in _cache_triton:
part = torch.empty(bs * ns * NUM_HEADS * V_HEAD_DIM, dtype=torch.bfloat16, device=dev)
lse = torch.empty(bs * ns * NUM_HEADS, dtype=torch.float32, device=dev)
out = torch.empty(bs, NUM_HEADS, V_HEAD_DIM, dtype=torch.bfloat16, device=dev)
_cache_triton[key] = (part, lse, out)
part, lse, out = _cache_triton[key]
kv_flat = kv_fp8.view(bs * kvlen, QK_HEAD_DIM)
_fd_s1[(bs * ns,)](
q, kv_flat, kv_scale, part, lse,
q.stride(0), q.stride(1), kv_flat.stride(0),
KVLEN=kvlen, BKV=bkv, sm_scale=SM_SCALE, NS=ns, num_warps=4,
)
_fd_red[(bs * NUM_HEADS,)](part, lse, out, NS=ns)
return out
def _build_aiter(bs, kvlen, dev, ps, gran, fm, nsplits):
pc = kvlen // ps if ps > 1 else kvlen
np_ = bs * pc if ps > 1 else bs * kvlen
lpl = ps if ps > 1 else kvlen
qoi = torch.arange(bs + 1, dtype=torch.int32, device=dev)
kvi = torch.arange(bs + 1, dtype=torch.int32, device=dev) * pc
klp = torch.full((bs,), lpl, dtype=torch.int32, device=dev)
ki = torch.arange(np_, dtype=torch.int32, device=dev)
out = torch.empty((bs, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=dev)
info = get_mla_metadata_info_v1(bs, 1, NUM_HEADS, torch.bfloat16, FP8_DTYPE,
is_sparse=False, fast_mode=fm, num_kv_splits=nsplits, intra_batch_mode=False)
w = [torch.empty(s, dtype=d, device=dev) for s, d in info]
get_mla_metadata_v1(qoi, kvi, klp, 16, 1, True, w[0], w[2], w[1], w[3], w[4], w[5],
page_size=ps, kv_granularity=gran, max_seqlen_qo=1, uni_seqlen_qo=1,
fast_mode=fm, max_split_per_batch=nsplits, intra_batch_mode=False,
dtype_q=torch.bfloat16, dtype_kv=FP8_DTYPE)
return qoi, kvi, klp, ki, out, w, np_, ps
def _run_aiter(q, kv_fp8, kv_scale, bs, kvlen, dev, ps, gran, fm, nsplits):
key = ('ait', bs, kvlen, ps, gran, fm, nsplits)
if key not in _cache_aiter:
_cache_aiter[key] = _build_aiter(bs, kvlen, dev, ps, gran, fm, nsplits)
qoi, kvi, klp, ki, out, w, np_, ps = _cache_aiter[key]
kv4d = kv_fp8.view(np_, ps, 1, QK_HEAD_DIM)
mla_decode_fwd(q, kv4d, out, qoi, kvi, ki, klp, 1,
page_size=ps, nhead_kv=1, sm_scale=SM_SCALE, logit_cap=0.0,
num_kv_splits=nsplits, q_scale=None, kv_scale=kv_scale,
intra_batch_mode=False, work_meta_data=w[0], work_indptr=w[1],
work_info_set=w[2], reduce_indptr=w[3], reduce_final_map=w[4],
reduce_partial_map=w[5])
return out
def custom_kernel(data):
q, kv_data, qo_indptr, _, _ = data
kv_fp8, kv_scale = kv_data["fp8"]
bs = qo_indptr.numel() - 1
kvlen = kv_fp8.shape[0] // bs
dev = q.device
# Triton fp8: wins for small batch + short kv
if bs <= 4:
if kvlen <= 1024:
return _run_triton(q, kv_fp8, kv_scale, bs, kvlen, dev, ns=16, bkv=32)
else:
return _run_triton(q, kv_fp8, kv_scale, bs, kvlen, dev, ns=16, bkv=64)
if bs <= 4 and kvlen <= 1024:
return _run_triton(q, kv_fp8, kv_scale, bs, kvlen, dev, ns=16, bkv=32)
if bs <= 32 and kvlen <= 1024:
return _run_triton(q, kv_fp8, kv_scale, bs, kvlen, dev, ns=4, bkv=32)
# v247: Triton for bs=64, pg2 for bs>=128
if bs >= 128 and kvlen <= 1024 and kvlen % 2 == 0:
ps, gran = 2, 16
return _run_aiter(q, kv_fp8, kv_scale, bs, kvlen, dev, ps, gran, fm=False, nsplits=32)
if bs >= 64 and kvlen <= 1024:
return _run_triton(q, kv_fp8, kv_scale, bs, kvlen, dev, ns=4, bkv=32)
if kvlen >= 8192:
return _run_aiter(q, kv_fp8, kv_scale, bs, kvlen, dev, ps=8, gran=16, fm=False, nsplits=64)
ps, gran = 1, 16
fm = bs <= 32
return _run_aiter(q, kv_fp8, kv_scale, bs, kvlen, dev, ps, gran, fm, 32)
scrolls · 189 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