submission 585650
Arseni Ivanov · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 333 lines, June 9 Researcher Reciprocity License v1.0.
submission_3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-585650?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:28dd9afcb5f3ed1e8c7c16fb71adbd51568554aaaf7619d3edce889b41ed95c8
license declaredunknown
license concludedunknown
authorsArseni Ivanov
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
qk1 = tl.dot(q1_bf16, tl.trans(v_bf16))num-warps = 4
num_warps=4,online-softmax
m_new = tl.maximum(m_i, m_ij)split-k
return custom_kernel_fp8_splitk(data)stages = 3
num_stages=3Kernel source
submission_3.py333 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
import torch
import torch.nn.functional as F
import triton
import triton.language as tl
from task import input_t, output_t
from aiter import dtypes as aiter_dtypes
FP8_DTYPE = aiter_dtypes.fp8
QKV_DTYPE = "fp8"
def custom_kernel(data: input_t) -> output_t:
"""Dispatch to the appropriate kernel based on QKV_DTYPE."""
if QKV_DTYPE == "fp8":
qo_indptr = data[2]
batch_size = qo_indptr.shape[0] - 1
if batch_size > 64:
return custom_kernel_fp8_nosplit(data)
else:
return custom_kernel_fp8_splitk(data)
elif QKV_DTYPE == "bf16":
return custom_kernel_bf16(data)
else:
raise ValueError(f"Invalid QKV_DTYPE: {QKV_DTYPE}")
@triton.jit
def mla_decode_fp8_nosplit_kernel(
q_ptr, kv_ptr, out_ptr,
qo_indptr, kv_indptr,
stride_q_tok, stride_q_h, stride_q_d,
stride_kv_tok, stride_kv_h, stride_kv_d,
stride_out_tok, stride_out_h, stride_out_d,
sm_scale, kv_scale,
BLOCK_KV: tl.constexpr,
):
batch_idx = tl.program_id(0)
q_start = tl.load(qo_indptr + batch_idx)
kv_start = tl.load(kv_indptr + batch_idx)
kv_end = tl.load(kv_indptr + batch_idx + 1)
seq_kv = kv_end - kv_start
offs_h = tl.arange(0, 16)
offs_d1 = tl.arange(0, 512)
offs_d2 = tl.arange(0, 64)
q_base = q_ptr + q_start * stride_q_tok + offs_h[:, None] * stride_q_h
q1_ptrs = q_base + offs_d1[None, :] * stride_q_d
q2_ptrs = q_base + (512 + offs_d2[None, :]) * stride_q_d
q1_bf16 = (tl.load(q1_ptrs).to(tl.float32) * sm_scale).to(tl.bfloat16)
q2_bf16 = (tl.load(q2_ptrs).to(tl.float32) * sm_scale).to(tl.bfloat16)
m_i = tl.full([16], float("-inf"), dtype=tl.float32)
l_i = tl.full([16], 1.0, dtype=tl.float32)
acc = tl.zeros([16, 512], dtype=tl.float32)
offs_kv = tl.arange(0, BLOCK_KV)
kv_base = kv_ptr + kv_start * stride_kv_tok
for start_n in range(0, seq_kv, BLOCK_KV):
start_n = tl.multiple_of(start_n, BLOCK_KV)
mask_kv = (start_n + offs_kv) < seq_kv
curr_kv_ptrs = kv_base + (start_n + offs_kv)[:, None] * stride_kv_tok
v_ptrs = curr_kv_ptrs + offs_d1[None, :] * stride_kv_d
k2_ptrs = curr_kv_ptrs + (512 + offs_d2[None, :]) * stride_kv_d
v_fp8 = tl.load(v_ptrs, mask=mask_kv[:, None], other=0.0)
k2_fp8 = tl.load(k2_ptrs, mask=mask_kv[:, None], other=0.0)
v_bf16 = (v_fp8.to(tl.float32) * kv_scale).to(tl.bfloat16)
k2_bf16 = (k2_fp8.to(tl.float32) * kv_scale).to(tl.bfloat16)
qk1 = tl.dot(q1_bf16, tl.trans(v_bf16))
qk2 = tl.dot(q2_bf16, tl.trans(k2_bf16))
qk = qk1 + qk2
qk = tl.where(mask_kv[None, :], qk, float("-inf"))
m_ij = tl.max(qk, 1)
m_new = tl.maximum(m_i, m_ij)
alpha = tl.exp(m_i - m_new)
p = tl.exp(qk - m_new[:, None])
l_ij = tl.sum(p, 1)
l_i = l_i * alpha + l_ij
acc = acc * alpha[:, None]
acc += tl.dot(p.to(tl.bfloat16), v_bf16)
m_i = m_new
acc = acc / l_i[:, None]
out_base = out_ptr + q_start * stride_out_tok + offs_h[:, None] * stride_out_h
out_ptrs = out_base + offs_d1[None, :] * stride_out_d
tl.store(out_ptrs, acc.to(tl.bfloat16))
def custom_kernel_fp8_nosplit(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
sm_scale = config["sm_scale"]
kv_buffer_fp8, kv_scale_fp8 = kv_data["fp8"]
kv_scale_val = kv_scale_fp8.item()
batch_size = qo_indptr.shape[0] - 1
total_q = q.shape[0]
out = torch.empty((total_q, 16, 512), dtype=torch.bfloat16, device=q.device)
BLOCK_KV = 128
mla_decode_fp8_nosplit_kernel[(batch_size,)](
q, kv_buffer_fp8, out,
qo_indptr, kv_indptr,
q.stride(0), q.stride(1), q.stride(2),
kv_buffer_fp8.stride(0), kv_buffer_fp8.stride(1), kv_buffer_fp8.stride(2),
out.stride(0), out.stride(1), out.stride(2),
sm_scale, kv_scale_val,
BLOCK_KV=BLOCK_KV,
)
return out
@triton.jit
def mla_decode_fp8_splitk_kernel(
q_ptr, kv_ptr,
workspace_acc, workspace_m, workspace_l,
out_ptr,
qo_indptr, kv_indptr,
stride_q_tok, stride_q_h, stride_q_d,
stride_kv_tok, stride_kv_h, stride_kv_d,
stride_out_tok, stride_out_h, stride_out_d,
combined_scale, kv_scale,
SPLIT_K: tl.constexpr,
BLOCK_KV: tl.constexpr,
):
batch_idx = tl.program_id(0)
split_idx = tl.program_id(1)
q_start = tl.load(qo_indptr + batch_idx)
kv_start = tl.load(kv_indptr + batch_idx)
kv_end = tl.load(kv_indptr + batch_idx + 1)
seq_kv = kv_end - kv_start
chunk_size = tl.cdiv(seq_kv, SPLIT_K)
chunk_start_idx = split_idx * chunk_size
chunk_end_idx = tl.minimum((split_idx + 1) * chunk_size, seq_kv)
offs_h = tl.arange(0, 16)
offs_d1 = tl.arange(0, 512)
offs_d2 = tl.arange(0, 64)
q_base = q_ptr + q_start * stride_q_tok + offs_h[:, None] * stride_q_h
q1_ptrs = q_base + offs_d1[None, :] * stride_q_d
q2_ptrs = q_base + (512 + offs_d2[None, :]) * stride_q_d
q1_bf16 = (tl.load(q1_ptrs).to(tl.float32) * combined_scale).to(tl.bfloat16)
q2_bf16 = (tl.load(q2_ptrs).to(tl.float32) * combined_scale).to(tl.bfloat16)
m_i = tl.full([16], float("-inf"), dtype=tl.float32)
l_i = tl.full([16], 0.0, dtype=tl.float32)
acc = tl.zeros([16, 512], dtype=tl.float32)
offs_kv = tl.arange(0, BLOCK_KV)
kv_base = kv_ptr + kv_start * stride_kv_tok
for start_n in range(chunk_start_idx, chunk_end_idx, BLOCK_KV):
start_n = tl.multiple_of(start_n, BLOCK_KV)
mask_kv = (start_n + offs_kv) < chunk_end_idx
curr_kv_ptrs = kv_base + (start_n + offs_kv)[:, None] * stride_kv_tok
v_ptrs = curr_kv_ptrs + offs_d1[None, :] * stride_kv_d
k2_ptrs = curr_kv_ptrs + (512 + offs_d2[None, :]) * stride_kv_d
v_fp8 = tl.load(v_ptrs, mask=mask_kv[:, None], other=0.0)
k2_fp8 = tl.load(k2_ptrs, mask=mask_kv[:, None], other=0.0)
v_bf16 = v_fp8.to(tl.bfloat16)
k2_bf16 = k2_fp8.to(tl.bfloat16)
qk1 = tl.dot(q1_bf16, tl.trans(v_bf16))
qk2 = tl.dot(q2_bf16, tl.trans(k2_bf16))
qk = qk1 + qk2
qk = tl.where(mask_kv[None, :], qk, float("-inf"))
m_ij = tl.max(qk, 1)
m_new = tl.maximum(m_i, m_ij)
alpha = tl.exp(m_i - m_new)
p = tl.exp(qk - m_new[:, None])
l_ij = tl.sum(p, 1)
l_i = l_i * alpha + l_ij
acc = acc * alpha[:, None]
acc += tl.dot(p.to(tl.bfloat16), v_bf16)
m_i = m_new
if SPLIT_K == 1:
final_out = (acc / l_i[:, None]) * kv_scale
out_base = out_ptr + q_start * stride_out_tok + offs_h[:, None] * stride_out_h
out_ptrs = out_base + offs_d1[None, :] * stride_out_d
tl.store(out_ptrs, final_out.to(tl.bfloat16))
else:
ws_idx = batch_idx * SPLIT_K + split_idx
tl.store(workspace_m + ws_idx * 16 + offs_h, m_i)
tl.store(workspace_l + ws_idx * 16 + offs_h, l_i)
ws_acc_base = workspace_acc + ws_idx * 16 * 512 + offs_h[:, None] * 512
tl.store(ws_acc_base + offs_d1[None, :], acc)
@triton.jit
def mla_decode_reduce_kernel(
workspace_acc, workspace_m, workspace_l, out_ptr,
qo_indptr, stride_out_tok, stride_out_h, stride_out_d,
kv_scale,
SPLIT_K: tl.constexpr
):
batch_idx = tl.program_id(0)
q_start = tl.load(qo_indptr + batch_idx)
offs_h = tl.arange(0, 16)
offs_d1 = tl.arange(0, 512)
m_global = tl.full([16], float("-inf"), dtype=tl.float32)
l_global = tl.full([16], 0.0, dtype=tl.float32)
acc_global = tl.zeros([16, 512], dtype=tl.float32)
for split_idx in range(SPLIT_K):
ws_idx = batch_idx * SPLIT_K + split_idx
m_j = tl.load(workspace_m + ws_idx * 16 + offs_h)
l_j = tl.load(workspace_l + ws_idx * 16 + offs_h)
ws_acc_base = workspace_acc + ws_idx * 16 * 512 + offs_h[:, None] * 512
acc_j = tl.load(ws_acc_base + offs_d1[None, :])
m_new = tl.maximum(m_global, m_j)
alpha_global = tl.exp(m_global - m_new)
alpha_j = tl.exp(m_j - m_new)
l_global = l_global * alpha_global + l_j * alpha_j
acc_global = acc_global * alpha_global[:, None] + acc_j * alpha_j[:, None]
m_global = m_new
out = (acc_global / l_global[:, None]) * kv_scale
out_base = out_ptr + q_start * stride_out_tok + offs_h[:, None] * stride_out_h
tl.store(out_base + offs_d1[None, :], out.to(tl.bfloat16))
def custom_kernel_fp8_splitk(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
sm_scale = config["sm_scale"]
kv_buffer_fp8, kv_scale_fp8 = kv_data["fp8"]
kv_scale_val = kv_scale_fp8.item()
batch_size = qo_indptr.shape[0] - 1
total_q = q.shape[0]
combined_scale = sm_scale * kv_scale_val
target_blocks = 256
split_k = max(1, min(64, target_blocks // batch_size))
out = torch.empty((total_q, 16, 512), dtype=torch.bfloat16, device=q.device)
if split_k > 1:
workspace_acc = torch.empty((batch_size, split_k, 16, 512), dtype=torch.float32, device=q.device)
workspace_m = torch.empty((batch_size, split_k, 16), dtype=torch.float32, device=q.device)
workspace_l = torch.empty((batch_size, split_k, 16), dtype=torch.float32, device=q.device)
else:
workspace_acc = q
workspace_m = q
workspace_l = q
BLOCK_KV = 128
grid_compute = (batch_size, split_k)
mla_decode_fp8_splitk_kernel[grid_compute](
q, kv_buffer_fp8,
workspace_acc, workspace_m, workspace_l,
out,
qo_indptr, kv_indptr,
q.stride(0), q.stride(1), q.stride(2),
kv_buffer_fp8.stride(0), kv_buffer_fp8.stride(1), kv_buffer_fp8.stride(2),
out.stride(0), out.stride(1), out.stride(2),
combined_scale, kv_scale_val,
SPLIT_K=split_k,
BLOCK_KV=BLOCK_KV,
num_warps=4,
num_stages=3
)
if split_k > 1:
grid_reduce = (batch_size,)
mla_decode_reduce_kernel[grid_reduce](
workspace_acc, workspace_m, workspace_l, out,
qo_indptr, out.stride(0), out.stride(1), out.stride(2),
kv_scale_val, SPLIT_K=split_k, num_warps=4
)
return out
def custom_kernel_bf16(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
kv_lora_rank = config["kv_lora_rank"]
sm_scale = config["sm_scale"]
kv_buffer_bf16 = kv_data["bf16"]
batch_size = qo_indptr.shape[0] - 1
out_list = []
for i in range(batch_size):
q_s, q_e = int(qo_indptr[i].item()), int(qo_indptr[i + 1].item())
kv_s, kv_e = int(kv_indptr[i].item()), int(kv_indptr[i + 1].item())
qi = q[q_s:q_e]
kvc = kv_buffer_bf16[kv_s:kv_e, 0]
ki = kvc
vi = kvc[:, :kv_lora_rank]
qi_t = qi.float().permute(1, 0, 2)
scores = torch.matmul(qi_t * sm_scale, ki.float().T)
scores = F.softmax(scores, dim=-1)
oi = torch.matmul(scores, vi.float())
oi = oi.permute(1, 0, 2)
out_list.append(oi.to(torch.bfloat16))
return torch.cat(out_list, dim=0)
scrolls · 333 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