submission 661945
mega-dmitriy · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 298 lines, June 9 Researcher Reciprocity License v1.0.
submission_v104_pro1.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-661945?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:8c2b8a48e00bd1f5b84b52e3959990b36e724956dc81f04211f8d183d0132d61
license declaredunknown
license concludedunknown
authorsmega-dmitriy
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp8
q_lat = tl.load(Q + q_idx * stride_q_tok + offs_h[:, None] * stride_q_head + offs_dk_lat[None, :]).to(tl.float8e4nv)mma
scores = (tl.dot(q_lat, tl.trans(k_lat_fp8)) + tl.dot(q_rope, tl.trans(k_rope_fp8))).to(tl.float32) * score_scaleonline-softmax
m_new = tl.maximum(m_i, m_ij)persistent-kernel
out_base = (split_id * tl.num_programs(1) + batch_id) * NUM_HEADS_CONSTsplit-k
split_kv_start = kv_start + split_id * TILES_PER_SPLIT * BLOCK_KVstages = 2
TILES_TOTAL=kv_seq_len // BKV, num_stages=2, allow_flush_denorm=True)Kernel source
submission_v104_pro1.py298 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
V104_pro1: Rec 5 + Rec 4 — AITER num_kv_splits sweep (24 instead of 32) + exp2.
- Try num_kv_splits=24 for bs=256/kv=8k (researcher suggests 24 as first sweep point)
- exp2 softmax + allow_flush_denorm on all Triton kernels
- Based on v69_c (current best 46.69us)
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t
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
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
LOG2E = 1.4426950408889634
SM_SCALE_LOG2E = SM_SCALE * LOG2E
PAGE_SIZE = 1
FP8_DTYPE = aiter_dtypes.fp8
@triton.jit
def _fused_fp8qk_singlepass(
Q, KV_FP8, KV_SCALE, O,
kv_indptr, qo_indptr, sm_scale_log2e,
stride_q_tok, stride_q_head, stride_kv_tok,
BLOCK_KV: tl.constexpr, BLOCK_DV: tl.constexpr,
NUM_HEADS_CONST: tl.constexpr, TILES_TOTAL: tl.constexpr,
):
batch_id = tl.program_id(0)
kv_start = tl.load(kv_indptr + batch_id)
q_idx = tl.load(qo_indptr + batch_id)
offs_h = tl.arange(0, NUM_HEADS_CONST)
offs_dv = tl.arange(0, BLOCK_DV)
offs_dk_lat = tl.arange(0, 512)
offs_dk_rope = tl.arange(0, 64)
offs_kv = tl.arange(0, BLOCK_KV)
q_lat = tl.load(Q + q_idx * stride_q_tok + offs_h[:, None] * stride_q_head + offs_dk_lat[None, :]).to(tl.float8e4nv)
q_rope = tl.load(Q + q_idx * stride_q_tok + offs_h[:, None] * stride_q_head + 512 + offs_dk_rope[None, :]).to(tl.float8e4nv)
kv_s = tl.load(KV_SCALE)
score_scale = sm_scale_log2e * kv_s
m_i = tl.full([NUM_HEADS_CONST], value=float("-inf"), dtype=tl.float32)
l_i = tl.zeros([NUM_HEADS_CONST], dtype=tl.float32)
acc = tl.zeros([NUM_HEADS_CONST, BLOCK_DV], dtype=tl.float32)
for tile_idx in range(TILES_TOTAL):
kv_offset = kv_start + tile_idx * BLOCK_KV
k_lat_fp8 = tl.load(KV_FP8 + (kv_offset + offs_kv[:, None]) * stride_kv_tok + offs_dk_lat[None, :])
k_rope_fp8 = tl.load(KV_FP8 + (kv_offset + offs_kv[:, None]) * stride_kv_tok + 512 + offs_dk_rope[None, :])
scores = (tl.dot(q_lat, tl.trans(k_lat_fp8)) + tl.dot(q_rope, tl.trans(k_rope_fp8))).to(tl.float32) * score_scale
m_ij = tl.max(scores, axis=1)
m_new = tl.maximum(m_i, m_ij)
alpha = tl.math.exp2(m_i - m_new)
p = tl.math.exp2(scores - m_new[:, None])
l_i = alpha * l_i + tl.sum(p, axis=1)
k_lat_bf16 = k_lat_fp8.to(tl.bfloat16)
acc = acc * alpha[:, None] + tl.dot(p.to(tl.bfloat16), k_lat_bf16)
m_i = m_new
acc = (acc * kv_s) / tl.maximum(l_i[:, None], 1e-12)
out_ptr = O + q_idx * NUM_HEADS_CONST * BLOCK_DV
tl.store(out_ptr + offs_h[:, None] * BLOCK_DV + offs_dv[None, :], acc.to(tl.bfloat16))
@triton.jit
def _flash_decode_fp8qk_stage1(
Q, KV_FP8, KV_SCALE, partial_O, partial_m, partial_l,
kv_indptr, qo_indptr, sm_scale_log2e,
stride_q_tok, stride_q_head, stride_kv_tok,
NUM_SPLITS: tl.constexpr, BLOCK_KV: tl.constexpr,
BLOCK_DV: tl.constexpr, NUM_HEADS_CONST: tl.constexpr,
TILES_PER_SPLIT: tl.constexpr,
):
split_id = tl.program_id(0)
batch_id = tl.program_id(1)
kv_start = tl.load(kv_indptr + batch_id)
q_idx = tl.load(qo_indptr + batch_id)
split_kv_start = kv_start + split_id * TILES_PER_SPLIT * BLOCK_KV
out_base = (split_id * tl.num_programs(1) + batch_id) * NUM_HEADS_CONST
offs_h = tl.arange(0, NUM_HEADS_CONST)
offs_dv = tl.arange(0, BLOCK_DV)
offs_dk_lat = tl.arange(0, 512)
offs_dk_rope = tl.arange(0, 64)
offs_kv = tl.arange(0, BLOCK_KV)
q_lat = tl.load(Q + q_idx * stride_q_tok + offs_h[:, None] * stride_q_head + offs_dk_lat[None, :]).to(tl.float8e4nv)
q_rope = tl.load(Q + q_idx * stride_q_tok + offs_h[:, None] * stride_q_head + 512 + offs_dk_rope[None, :]).to(tl.float8e4nv)
kv_s = tl.load(KV_SCALE)
score_scale = sm_scale_log2e * kv_s
m_i = tl.full([NUM_HEADS_CONST], value=float("-inf"), dtype=tl.float32)
l_i = tl.zeros([NUM_HEADS_CONST], dtype=tl.float32)
acc = tl.zeros([NUM_HEADS_CONST, BLOCK_DV], dtype=tl.float32)
for tile_idx in range(TILES_PER_SPLIT):
kv_offset = split_kv_start + tile_idx * BLOCK_KV
k_lat_fp8 = tl.load(KV_FP8 + (kv_offset + offs_kv[:, None]) * stride_kv_tok + offs_dk_lat[None, :])
k_rope_fp8 = tl.load(KV_FP8 + (kv_offset + offs_kv[:, None]) * stride_kv_tok + 512 + offs_dk_rope[None, :])
scores = (tl.dot(q_lat, tl.trans(k_lat_fp8)) + tl.dot(q_rope, tl.trans(k_rope_fp8))).to(tl.float32) * score_scale
m_ij = tl.max(scores, axis=1)
m_new = tl.maximum(m_i, m_ij)
alpha = tl.math.exp2(m_i - m_new)
p = tl.math.exp2(scores - m_new[:, None])
l_i = alpha * l_i + tl.sum(p, axis=1)
k_lat_bf16 = k_lat_fp8.to(tl.bfloat16)
acc = acc * alpha[:, None] + tl.dot(p.to(tl.bfloat16), k_lat_bf16)
m_i = m_new
acc = acc * kv_s
tl.store(partial_m + out_base + offs_h, m_i)
tl.store(partial_l + out_base + offs_h, l_i)
tl.store(partial_O + (out_base + offs_h[:, None]) * BLOCK_DV + offs_dv[None, :], acc)
@triton.jit
def _flash_decode_stage2(
partial_O, partial_m, partial_l, O, qo_indptr,
NUM_SPLITS: tl.constexpr, BLOCK_DV: tl.constexpr,
NUM_HEADS_CONST: tl.constexpr, BATCH_SIZE: tl.constexpr,
):
pid = tl.program_id(0)
batch_id = pid // NUM_HEADS_CONST
head_id = pid % NUM_HEADS_CONST
q_idx = tl.load(qo_indptr + batch_id)
offs_dv = tl.arange(0, BLOCK_DV)
m_global = tl.full([], value=float("-inf"), dtype=tl.float32)
for s in range(NUM_SPLITS):
base = (s * BATCH_SIZE + batch_id) * NUM_HEADS_CONST + head_id
m_global = tl.maximum(m_global, tl.load(partial_m + base))
acc = tl.zeros([BLOCK_DV], dtype=tl.float32)
l_total = tl.full([], value=0.0, dtype=tl.float32)
for s in range(NUM_SPLITS):
base = (s * BATCH_SIZE + batch_id) * NUM_HEADS_CONST + head_id
alpha = tl.math.exp2(tl.load(partial_m + base) - m_global)
acc += alpha * tl.load(partial_O + base * BLOCK_DV + offs_dv)
l_total += alpha * tl.load(partial_l + base)
acc = acc / tl.maximum(l_total, 1e-12)
tl.store(O + q_idx * NUM_HEADS_CONST * BLOCK_DV + head_id * BLOCK_DV + offs_dv, acc.to(tl.bfloat16))
@triton.jit
def _q_scale_cast_fp8(Q_IN, Q_OUT, inv_scale, N_ELEM: tl.constexpr, BLOCK: tl.constexpr):
pid = tl.program_id(0)
offs = pid * BLOCK + tl.arange(0, BLOCK)
mask = offs < N_ELEM
q = tl.load(Q_IN + offs, mask=mask).to(tl.float32) * inv_scale
tl.store(Q_OUT + offs, q.to(tl.float8e4nv), mask=mask)
_triton_buf = {}
_aiter_meta = {}
_aiter_idx = {}
_aiter_kv_last = {}
_aiter_meta_ready = {}
_out_buf = {}
_q_fp8_buf = {}
_finfo = torch.finfo(FP8_DTYPE)
_FIXED_INV_SCALE = _finfo.max / 16.0
_FIXED_Q_SCALE = torch.tensor([16.0 / _finfo.max], dtype=torch.float32, device="cuda")
def _run_fused_fp8qk(q, kv_data, qo_indptr, kv_indptr, batch_size, nq, BKV, kv_seq_len):
kv_fp8, kv_scale = kv_data["fp8"]
kv_flat = kv_fp8.view(-1, QK_HEAD_DIM)
total_q = q.shape[0]
ok = (total_q, nq)
if ok not in _out_buf:
_out_buf[ok] = torch.empty((total_q, nq, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
o = _out_buf[ok]
_fused_fp8qk_singlepass[(batch_size,)](
q, kv_flat, kv_scale, o, kv_indptr, qo_indptr, SM_SCALE_LOG2E,
q.stride(0), q.stride(1), kv_flat.stride(0),
BLOCK_KV=BKV, BLOCK_DV=V_HEAD_DIM, NUM_HEADS_CONST=nq,
TILES_TOTAL=kv_seq_len // BKV, num_stages=2, allow_flush_denorm=True)
return o
def _run_fp8qk(q, kv_data, qo_indptr, kv_indptr, batch_size, nq, NS, BKV, kv_seq_len, stages=2):
kv_fp8, kv_scale = kv_data["fp8"]
kv_flat = kv_fp8.view(-1, QK_HEAD_DIM)
total_q = q.shape[0]
total_e = NS * batch_size * nq
ck = ("fp8qk", NS, batch_size, BKV)
if ck not in _triton_buf:
_triton_buf[ck] = (
torch.empty((total_e, V_HEAD_DIM), dtype=torch.float32, device="cuda"),
torch.empty((total_e,), dtype=torch.float32, device="cuda"),
torch.empty((total_e,), dtype=torch.float32, device="cuda"),
)
pO, pm, pl = _triton_buf[ck]
tiles_per_split = kv_seq_len // (NS * BKV)
_flash_decode_fp8qk_stage1[(NS, batch_size)](
q, kv_flat, kv_scale, pO, pm, pl, kv_indptr, qo_indptr, SM_SCALE_LOG2E,
q.stride(0), q.stride(1), kv_flat.stride(0),
NUM_SPLITS=NS, BLOCK_KV=BKV, BLOCK_DV=V_HEAD_DIM, NUM_HEADS_CONST=nq,
TILES_PER_SPLIT=tiles_per_split, num_stages=stages, allow_flush_denorm=True)
ok = (total_q, nq)
if ok not in _out_buf:
_out_buf[ok] = torch.empty((total_q, nq, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
o = _out_buf[ok]
_flash_decode_stage2[(batch_size * nq,)](
pO, pm, pl, o, qo_indptr,
NUM_SPLITS=NS, BLOCK_DV=V_HEAD_DIM, NUM_HEADS_CONST=nq, BATCH_SIZE=batch_size,
allow_flush_denorm=True)
return o
def _run_aiter(q, kv_data, qo_indptr, kv_indptr, config, num_splits=32):
batch_size = config["batch_size"]
nq = config["num_heads"]
nkv = config["num_kv_heads"]
dq = config["qk_head_dim"]
dv = config["v_head_dim"]
q_seq_len = config["q_seq_len"]
kv_seq_len = config["kv_seq_len"]
NUM_SPLITS = num_splits
n_elem = q.numel()
qk = ("q_fp8", n_elem)
if qk not in _q_fp8_buf:
_q_fp8_buf[qk] = torch.empty(q.shape, dtype=FP8_DTYPE, device="cuda")
q_fp8 = _q_fp8_buf[qk]
BLOCK = 1024
grid = ((n_elem + BLOCK - 1) // BLOCK,)
_q_scale_cast_fp8[grid](q.view(-1), q_fp8.view(-1), _FIXED_INV_SCALE, N_ELEM=n_elem, BLOCK=BLOCK)
q_scale = _FIXED_Q_SCALE
kv_fp8, kv_s = kv_data["fp8"]
total_kv_len = batch_size * kv_seq_len
kv_4d = kv_fp8.view(kv_fp8.shape[0], PAGE_SIZE, nkv, kv_fp8.shape[-1])
kl_key = (batch_size, kv_seq_len)
if kl_key not in _aiter_kv_last:
_aiter_kv_last[kl_key] = torch.full(
(batch_size,), kv_seq_len, dtype=torch.int32, device="cuda")
kv_last = _aiter_kv_last[kl_key]
ck = (batch_size, total_kv_len, nq, q_fp8.dtype, kv_fp8.dtype, NUM_SPLITS)
if ck not in _aiter_meta:
info = get_mla_metadata_info_v1(
batch_size, q_seq_len, nq, q_fp8.dtype, kv_fp8.dtype,
is_sparse=False, fast_mode=False,
num_kv_splits=NUM_SPLITS, intra_batch_mode=True)
_aiter_meta[ck] = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
work = _aiter_meta[ck]
(wm, wi, wis, ri, rfm, rpm) = work
if ck not in _aiter_meta_ready:
get_mla_metadata_v1(
qo_indptr, kv_indptr, kv_last,
nq // nkv, nkv, True,
wm, wis, wi, ri, rfm, rpm,
page_size=PAGE_SIZE, kv_granularity=max(PAGE_SIZE, 16),
max_seqlen_qo=q_seq_len, uni_seqlen_qo=q_seq_len,
fast_mode=False, max_split_per_batch=NUM_SPLITS,
intra_batch_mode=True, dtype_q=q_fp8.dtype, dtype_kv=kv_fp8.dtype)
_aiter_meta_ready[ck] = True
if total_kv_len not in _aiter_idx:
_aiter_idx[total_kv_len] = torch.arange(total_kv_len, dtype=torch.int32, device="cuda")
ok = (q.shape[0], nq, dv)
if ok not in _out_buf:
_out_buf[ok] = torch.empty(ok, dtype=torch.bfloat16, device="cuda")
o = _out_buf[ok]
mla_decode_fwd(
q_fp8.view(-1, nq, dq), kv_4d, o,
qo_indptr, kv_indptr, _aiter_idx[total_kv_len],
kv_last, q_seq_len,
page_size=PAGE_SIZE, nhead_kv=nkv,
sm_scale=SM_SCALE, logit_cap=0.0,
num_kv_splits=NUM_SPLITS,
q_scale=q_scale, kv_scale=kv_s,
intra_batch_mode=True,
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 custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = config["batch_size"]
kv_seq_len = config["kv_seq_len"]
nq = config["num_heads"]
if kv_seq_len <= 1024:
if batch_size <= 4:
return _run_fp8qk(q, kv_data, qo_indptr, kv_indptr, batch_size, nq, NS=16, BKV=64, kv_seq_len=kv_seq_len, stages=1)
elif batch_size <= 32:
return _run_fp8qk(q, kv_data, qo_indptr, kv_indptr, batch_size, nq, NS=8, BKV=64, kv_seq_len=kv_seq_len)
elif batch_size <= 64:
return _run_fp8qk(q, kv_data, qo_indptr, kv_indptr, batch_size, nq, NS=4, BKV=64, kv_seq_len=kv_seq_len)
else:
return _run_fused_fp8qk(q, kv_data, qo_indptr, kv_indptr, batch_size, nq, BKV=64, kv_seq_len=kv_seq_len)
else:
if batch_size <= 4:
return _run_fp8qk(q, kv_data, qo_indptr, kv_indptr, batch_size, nq, NS=32, BKV=64, kv_seq_len=kv_seq_len)
elif batch_size <= 32:
return _run_fp8qk(q, kv_data, qo_indptr, kv_indptr, batch_size, nq, NS=8, BKV=64, kv_seq_len=kv_seq_len)
elif batch_size <= 64:
return _run_fp8qk(q, kv_data, qo_indptr, kv_indptr, batch_size, nq, NS=4, BKV=64, kv_seq_len=kv_seq_len)
else:
# AITER with num_kv_splits=24 (sweeping down from 32)
return _run_aiter(q, kv_data, qo_indptr, kv_indptr, config, num_splits=24)
scrolls · 298 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