submission 608517
Amanpreet Singh · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 171 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-608517?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:77930926974fc42e6322126278d340c30fc026e87319c227077e60cdfcbfa5de
license declaredunknown
license concludedunknown
authorsAmanpreet Singh
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
qk = tl.dot(q_v_bf16, tl.trans(v_bf16), out_dtype=tl.float32) + tl.dot(q_r_bf16, tl.trans(r_bf16), out_dtype=tl.float32)num-warps = 4
NUM_SPLITS=NS, V_DIM=v_dim, num_warps=4, num_stages=1persistent-kernel
n_splits = tl.num_programs(1)split-k
def _mla_fp8_mqa_splitk(stages = 3
V_DIM=v_dim, ROPE_DIM=r_dim, BLOCK_KV=B_KV, num_warps=nw, num_stages=3Kernel source
submission.py171 lines
import torch
import triton
import triton.language as tl
from task import input_t, output_t
@triton.jit
def _mla_fp8_mqa_splitk(
Q, KV_Data, P_Acc, P_Lse,
qo_indptr, kv_indptr,
stride_q_b, stride_q_h, stride_q_d,
stride_kv_b, stride_kv_d,
stride_pa_b, stride_pa_h, stride_pa_s, stride_pa_d,
stride_pl_b, stride_pl_h, stride_pl_s,
KV_Scale, sm_scale,
V_DIM: tl.constexpr, ROPE_DIM: tl.constexpr, BLOCK_KV: tl.constexpr
):
b_idx = tl.program_id(0)
s_idx = tl.program_id(1)
n_splits = tl.num_programs(1)
q_st = tl.load(qo_indptr + b_idx)
kv_st = tl.load(kv_indptr + b_idx)
kv_end = tl.load(kv_indptr + b_idx + 1)
seq_len = kv_end - kv_st
chunk = (seq_len + n_splits - 1) // n_splits
start_n = s_idx * chunk
end_n = tl.minimum(start_n + chunk, seq_len)
offs_h = tl.arange(0, 16)
offs_v = tl.arange(0, V_DIM)
offs_r = tl.arange(0, ROPE_DIM) + V_DIM
offs_kv = tl.arange(0, BLOCK_KV)
# Natively load scalars straight from VRAM, bypassing Python JIT stalls
kv_scale_val = tl.load(KV_Scale)
combined_qk_scale = sm_scale * kv_scale_val
q_v_ptr = Q + (q_st * stride_q_b) + offs_h[:, None] * stride_q_h + offs_v[None, :] * stride_q_d
q_r_ptr = Q + (q_st * stride_q_b) + offs_h[:, None] * stride_q_h + offs_r[None, :] * stride_q_d
# No PyTorch conversion overhead; we load Q natively as BFLOAT16 in the registers
q_v_bf16 = tl.load(q_v_ptr)
q_r_bf16 = tl.load(q_r_ptr)
kv_v_ptr = KV_Data + ((kv_st + start_n) * stride_kv_b) + offs_kv[:, None] * stride_kv_b + offs_v[None, :] * stride_kv_d
kv_r_ptr = KV_Data + ((kv_st + start_n) * stride_kv_b) + offs_kv[:, None] * stride_kv_b + offs_r[None, :] * stride_kv_d
m_i = tl.zeros([16], dtype=tl.float32) - float('inf')
l_i = tl.zeros([16], dtype=tl.float32)
acc = tl.zeros([16, V_DIM], dtype=tl.float32)
for n in range(start_n, end_n, BLOCK_KV):
n_align = tl.multiple_of(n, BLOCK_KV)
mask = (n_align + offs_kv) < end_n
v_fp8 = tl.load(kv_v_ptr, mask=mask[:, None], other=0.0)
r_fp8 = tl.load(kv_r_ptr, mask=mask[:, None], other=0.0)
v_bf16 = v_fp8.to(tl.bfloat16)
r_bf16 = r_fp8.to(tl.bfloat16)
qk = tl.dot(q_v_bf16, tl.trans(v_bf16), out_dtype=tl.float32) + tl.dot(q_r_bf16, tl.trans(r_bf16), out_dtype=tl.float32)
qk = qk * combined_qk_scale
qk = tl.where(mask[None, :], qk, -float('inf'))
m_ij = tl.maximum(m_i, tl.max(qk, axis=1))
alpha = tl.exp(m_i - m_ij)
p = tl.exp(qk - m_ij[:, None])
acc = acc * alpha[:, None] + tl.dot(p.to(tl.bfloat16), v_bf16, out_dtype=tl.float32)
l_i = l_i * alpha + tl.sum(p, axis=1)
m_i = m_ij
kv_v_ptr += BLOCK_KV * stride_kv_b
kv_r_ptr += BLOCK_KV * stride_kv_b
acc = tl.where((l_i > 0.0)[:, None], (acc / l_i[:, None]) * kv_scale_val, 0.0)
lse = m_i + tl.math.log(l_i)
lse = tl.where(l_i > 0.0, lse, -float('inf'))
pa_ptr = P_Acc + (q_st * stride_pa_b) + offs_h[:, None] * stride_pa_h + (s_idx * stride_pa_s) + offs_v[None, :] * stride_pa_d
pl_ptr = P_Lse + (q_st * stride_pl_b) + offs_h * stride_pl_h + (s_idx * stride_pl_s)
tl.store(pa_ptr, acc)
tl.store(pl_ptr, lse)
@triton.jit
def _mla_mqa_reduce(
P_Acc, P_Lse, Out, qo_indptr,
s_pa_b, s_pa_h, s_pa_s, s_pa_d,
s_pl_b, s_pl_h, s_pl_s,
s_o_b, s_o_h, s_o_d,
NUM_SPLITS: tl.constexpr, V_DIM: tl.constexpr
):
b_idx = tl.program_id(0)
h_idx = tl.program_id(1)
q_st = tl.load(qo_indptr + b_idx)
offs_s = tl.arange(0, NUM_SPLITS)
offs_v = tl.arange(0, V_DIM)
pl_ptr = P_Lse + (q_st * s_pl_b) + (h_idx * s_pl_h) + offs_s * s_pl_s
lse = tl.load(pl_ptr)
max_lse = tl.max(lse, axis=0)
w = tl.exp(lse - max_lse)
sum_w = tl.sum(w, axis=0)
pa_ptr = P_Acc + (q_st * s_pa_b) + (h_idx * s_pa_h) + offs_s[:, None] * s_pa_s + offs_v[None, :] * s_pa_d
acc = tl.load(pa_ptr)
out = tl.sum(acc * w[:, None], axis=0) / sum_w
o_ptr = Out + (q_st * s_o_b) + (h_idx * s_o_h) + offs_v * s_o_d
tl.store(o_ptr, out.to(Out.dtype.element_ty))
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
kv_fp8, scale_fp8 = kv_data["fp8"]
bs = config["batch_size"]
nh = config["num_heads"]
v_dim = config["v_head_dim"]
r_dim = config["qk_rope_head_dim"]
seq_len = config["kv_seq_len"]
tq = q.shape[0]
out = torch.empty((tq, nh, v_dim), dtype=torch.bfloat16, device=q.device)
# ------------------------------------------------------------------
# THE OCCUPANCY ENGINE
# ------------------------------------------------------------------
if seq_len <= 1024:
# Dynamically scale splits to wake up all 304 Compute Units
target_ns = max(1, 256 // bs)
NS = min(target_ns, 16) # Cap at 16 to keep loop sizes healthy
B_KV = 64
nw = 4
else:
# UNLEASHED SPLIT-K: Small inner loops = massive speed on big sequences
NS = 32
B_KV = 128
nw = 8
p_acc = torch.empty((tq, nh, NS, v_dim), dtype=torch.float32, device=q.device)
p_lse = torch.empty((tq, nh, NS), dtype=torch.float32, device=q.device)
_mla_fp8_mqa_splitk[(bs, NS)](
q, kv_fp8, p_acc, p_lse,
qo_indptr, kv_indptr,
q.stride(0), q.stride(1), q.stride(2),
kv_fp8.stride(0), kv_fp8.stride(2),
p_acc.stride(0), p_acc.stride(1), p_acc.stride(2), p_acc.stride(3),
p_lse.stride(0), p_lse.stride(1), p_lse.stride(2),
scale_fp8, config["sm_scale"],
V_DIM=v_dim, ROPE_DIM=r_dim, BLOCK_KV=B_KV, num_warps=nw, num_stages=3
)
if NS == 1:
out.copy_(p_acc.squeeze(2).to(torch.bfloat16))
else:
_mla_mqa_reduce[(bs, nh)](
p_acc, p_lse, out, qo_indptr,
p_acc.stride(0), p_acc.stride(1), p_acc.stride(2), p_acc.stride(3),
p_lse.stride(0), p_lse.stride(1), p_lse.stride(2),
out.stride(0), out.stride(1), out.stride(2),
NUM_SPLITS=NS, V_DIM=v_dim, num_warps=4, num_stages=1
)
return outscrolls · 171 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