submission 737458
Jayluci4 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 191 lines, June 9 Researcher Reciprocity License v1.0.
submission_mla_hybrid.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-737458?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:264b18534353576efa1abe856907da461cb1d8bd252a1e69aee72e702895d79a
license declaredunknown
license concludedunknown
authorsJayluci4
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
scores = tl.dot(q_nope, tl.trans(kv_nope_bf))num-warps = 4
NUM_SPLITS=splits, BLOCK_KV=64, num_warps=4, num_stages=2)online-softmax
m_new = tl.maximum(m_i, m_block)stages = 2
NUM_SPLITS=splits, BLOCK_KV=64, num_warps=4, num_stages=2)Kernel source
submission_mla_hybrid.py191 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
Hybrid MLA decode: Triton fused for bs<=32, ASM for bs>=64.
Triton: single launch, bf16 Q (no quant), MQA, fp8 KV.
ASM: hand-tuned BW, splits=1 for bs=256 (no reduce).
"""
import gc
import sys
import os
gc.disable()
sys.setswitchinterval(1000.0)
os.environ["HIP_FORCE_DEV_KERNARG"] = "1"
import torch
import triton
import triton.language as tl
from task import input_t, output_t
import aiter
from aiter import dtypes as aiter_dtypes
NUM_HEADS = 16
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
SM_SCALE = 1.0 / (576 ** 0.5)
PAGE_SIZE = 1
FP8 = aiter_dtypes.fp8
# ASM only for shapes where it beats Triton (large bs × large kv)
_ASM_SPLITS = {(64, 8192): 4, (256, 8192): 1}
# Triton for everything else (kv_scale absorbed, single-launch wins)
_TRI_SPLITS = {
(4, 1024): 64, (4, 8192): 64,
(32, 1024): 8, (32, 8192): 8,
(64, 1024): 4, (256, 1024): 1,
}
# Route: use ASM when shape is in _ASM_SPLITS
_USE_ASM = {(64, 8192), (256, 8192)}
_cache = {}
@triton.jit
def _mla_fused_kernel(
Q_ptr, KV_fp8_ptr, O_ptr, Mid_O_ptr, Mid_LSE_ptr,
kv_indptr_ptr, kv_scale_ptr, sm_scale,
NUM_SPLITS: tl.constexpr, BLOCK_KV: tl.constexpr,
):
pid = tl.program_id(0)
batch_idx = pid // NUM_SPLITS
split_idx = pid % NUM_SPLITS
kv_start = tl.load(kv_indptr_ptr + batch_idx)
kv_end = tl.load(kv_indptr_ptr + batch_idx + 1)
kv_len = kv_end - kv_start
split_size = (kv_len + NUM_SPLITS - 1) // NUM_SPLITS
s_start = kv_start + split_idx * split_size
s_end = tl.minimum(s_start + split_size, kv_end)
kv_scale = tl.load(kv_scale_ptr)
sm_scale_adj = sm_scale * kv_scale
q_base = batch_idx * 16 * 576
offs_h = tl.arange(0, 16)
offs_512 = tl.arange(0, 512)
offs_64 = tl.arange(0, 64)
q_nope = tl.load(Q_ptr + q_base + offs_h[:, None] * 576 + offs_512[None, :]).to(tl.bfloat16)
q_rope = tl.load(Q_ptr + q_base + offs_h[:, None] * 576 + 512 + offs_64[None, :]).to(tl.bfloat16)
m_i = tl.full([16], float("-inf"), dtype=tl.float32)
l_i = tl.zeros([16], dtype=tl.float32)
o_acc = tl.zeros([16, 512], dtype=tl.float32)
offs_kv = tl.arange(0, BLOCK_KV)
for kv_off in range(0, split_size, BLOCK_KV):
kv_pos = s_start + kv_off
valid = (kv_pos + offs_kv) < s_end
kv_rows = tl.minimum(kv_pos + offs_kv, s_end - 1)
kv_nope_bf = tl.load(KV_fp8_ptr + kv_rows[:, None] * 576 + offs_512[None, :]).to(tl.bfloat16)
kv_rope_bf = tl.load(KV_fp8_ptr + kv_rows[:, None] * 576 + 512 + offs_64[None, :]).to(tl.bfloat16)
scores = tl.dot(q_nope, tl.trans(kv_nope_bf))
scores += tl.dot(q_rope, tl.trans(kv_rope_bf))
scores = scores.to(tl.float32) * sm_scale_adj
scores = tl.where(valid[None, :], scores, float("-inf"))
m_block = tl.max(scores, axis=1)
m_new = tl.maximum(m_i, m_block)
alpha = tl.exp(m_i - m_new)
p = tl.exp(scores - m_new[:, None])
l_new = l_i * alpha + tl.sum(p, axis=1)
p_bf = p.to(tl.bfloat16)
o_acc = o_acc * alpha[:, None] + tl.dot(p_bf, kv_nope_bf).to(tl.float32)
m_i = m_new
l_i = l_new
safe_l = tl.where(l_i > 0, l_i, 1.0)
o_acc = (o_acc / safe_l[:, None]) * kv_scale
if NUM_SPLITS == 1:
o_base = batch_idx * 16 * 512
tl.store(O_ptr + o_base + offs_h[:, None] * 512 + offs_512[None, :], o_acc.to(tl.bfloat16))
else:
mid_idx = batch_idx * NUM_SPLITS + split_idx
mid_base = mid_idx * 16 * 512
tl.store(Mid_O_ptr + mid_base + offs_h[:, None] * 512 + offs_512[None, :], o_acc)
lse = tl.where(l_i > 0, m_i + tl.log(l_i), float("-inf"))
tl.store(Mid_LSE_ptr + mid_idx * 16 + offs_h, lse)
@triton.jit
def _reduce_kernel(Mid_O_ptr, Mid_LSE_ptr, O_ptr, NUM_SPLITS: tl.constexpr):
pid_b = tl.program_id(0)
pid_h = tl.program_id(1)
offs_v = tl.arange(0, 512)
m_max = tl.full([1], float("-inf"), dtype=tl.float32)
for s in range(NUM_SPLITS):
lse = tl.load(Mid_LSE_ptr + (pid_b * NUM_SPLITS + s) * 16 + pid_h)
m_max = tl.maximum(m_max, lse)
acc = tl.zeros([512], dtype=tl.float32)
l_sum = tl.zeros([1], dtype=tl.float32)
for s in range(NUM_SPLITS):
idx = pid_b * NUM_SPLITS + s
lse = tl.load(Mid_LSE_ptr + idx * 16 + pid_h)
w = tl.exp(lse - m_max)
l_sum += w
partial = tl.load(Mid_O_ptr + idx * 16 * 512 + pid_h * 512 + offs_v)
acc += w * partial
safe_l = tl.where(l_sum > 0.0, l_sum, 1.0)
tl.store(O_ptr + pid_b * 16 * 512 + pid_h * 512 + offs_v, (acc / safe_l).to(tl.bfloat16))
def _build_triton(bs, kv_len):
splits = _TRI_SPLITS.get((bs, kv_len), 8)
output = torch.empty((bs, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
mid_o = torch.empty((bs * splits, NUM_HEADS, V_HEAD_DIM), dtype=torch.float32, device="cuda") if splits > 1 else None
mid_lse = torch.empty((bs * splits, NUM_HEADS), dtype=torch.float32, device="cuda") if splits > 1 else None
return ("t", splits, output, mid_o, mid_lse)
def _build_asm(bs, kv_len):
splits = _ASM_SPLITS.get((bs, kv_len), 1)
total_kv = bs * kv_len
kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")
kv_last_page_len = torch.ones(bs, dtype=torch.int32, device="cuda")
q_fp8 = torch.empty((bs, NUM_HEADS, QK_HEAD_DIM), dtype=FP8, device="cuda")
q_scale = torch.ones(1, dtype=torch.float32, device="cuda")
output = torch.empty((bs, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
splits_indptr = torch.arange(0, (bs + 1) * splits, splits, dtype=torch.int32, device="cuda")
if splits == 1:
logits = output.view(bs, 1, NUM_HEADS, V_HEAD_DIM)
attn_lse = torch.empty((bs, 1, NUM_HEADS, 1), dtype=torch.float32, device="cuda")
else:
logits = torch.empty((bs, splits, NUM_HEADS, V_HEAD_DIM), dtype=torch.float32, device="cuda")
attn_lse = torch.empty((bs, splits, NUM_HEADS, 1), dtype=torch.float32, device="cuda")
return ("a", splits, kv_indices, kv_last_page_len, q_fp8, q_scale,
output, splits_indptr, logits, attn_lse, total_kv)
def custom_kernel(data: input_t) -> output_t:
cfg = data[4]
bs = cfg["batch_size"]
kv_len = cfg["kv_seq_len"]
key = (bs, kv_len)
if key not in _cache:
_cache[key] = _build_asm(bs, kv_len) if key in _USE_ASM else _build_triton(bs, kv_len)
c = _cache[key]
if c[0] == "t":
_, splits, output, mid_o, mid_lse = c
kv_fp8, kv_scale = data[1]["fp8"]
mo = mid_o if mid_o is not None else output
ml = mid_lse if mid_lse is not None else output
_mla_fused_kernel[(bs * splits,)](
data[0], kv_fp8, output, mo, ml, data[3], kv_scale, SM_SCALE,
NUM_SPLITS=splits, BLOCK_KV=64, num_warps=4, num_stages=2)
if splits > 1:
_reduce_kernel[(bs, NUM_HEADS)](mid_o, mid_lse, output,
NUM_SPLITS=splits, num_warps=4, num_stages=1)
return output
else:
_, splits, kv_indices, kv_last_page_len, q_fp8, q_scale, \
output, splits_indptr, logits, attn_lse, total_kv = c
q_fp8.copy_(data[0])
kv_fp8, kv_scale = data[1]["fp8"]
aiter.mla_decode_stage1_asm_fwd(
q_fp8, kv_fp8.view(total_kv, PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM),
data[2], data[3], kv_indices, kv_last_page_len, splits_indptr,
None, None, None, 1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
logits, attn_lse, output, q_scale, kv_scale)
if splits == 1:
return output
reduce_o = logits.view(bs * splits, NUM_HEADS, V_HEAD_DIM)
reduce_lse = attn_lse.view(bs * splits, NUM_HEADS)
_reduce_kernel[(bs, NUM_HEADS)](reduce_o, reduce_lse, output,
NUM_SPLITS=splits, num_warps=4, num_stages=1)
return output
scrolls · 191 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