submission 695464
sunfj · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 289 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-695464?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:ede5329a22a86d033bf0c88d6835cc2629556ce1008751b024d663278b0e168a
license declaredunknown
license concludedunknown
authorssunfj
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
autotune
@triton.autotune(mma
qk = tl.dot(q_512, tl.trans(k_512))num-warps = 8
triton.Config({'BLOCK_H': 64, 'BLOCK_N': 64}, num_stages=3, num_warps=8),split-k
key=['nq', 'num_batches', 'SPLIT_K']stages = 3
triton.Config({'BLOCK_H': 64, 'BLOCK_N': 64}, num_stages=3, num_warps=8),Kernel source
submission.py289 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
import torch
import triton
import triton.language as tl
class FlashDecodeCache:
mid_o = None
mid_m = None
mid_l = None
fp32_one = None
# =====================================================================
# Stage 1: 纯血 MLA 极速内核 (消除 V 矩阵冗余加载)
# =====================================================================
@triton.autotune(
configs=[
triton.Config({'BLOCK_H': 64, 'BLOCK_N': 64}, num_stages=3, num_warps=8),
triton.Config({'BLOCK_H': 32, 'BLOCK_N': 128}, num_stages=3, num_warps=4),
triton.Config({'BLOCK_H': 32, 'BLOCK_N': 64}, num_stages=4, num_warps=4),
triton.Config({'BLOCK_H': 16, 'BLOCK_N': 128}, num_stages=4, num_warps=4),
],
key=['nq', 'num_batches', 'SPLIT_K']
)
@triton.jit
def mla_decode_stage1(
Q, KV, Out, mid_O, mid_M, mid_L,
qo_indptr, kv_indptr,
kv_scale_ptr, sm_scale,
stride_qz, stride_qh, stride_qd,
stride_kz, stride_kh, stride_kd,
# 【核心优化 1】: 彻底删除了 stride_vz, vh, vd,因为不需要了!
stride_oz, stride_oh, stride_od,
nq, num_batches,
SPLIT_K: tl.constexpr,
BLOCK_H: tl.constexpr,
BLOCK_N: tl.constexpr,
):
pid_h = tl.program_id(0)
batch_idx = tl.program_id(1)
pid_sk = tl.program_id(2)
if batch_idx >= num_batches:
return
q_idx = tl.load(qo_indptr + batch_idx)
kv_start_base = tl.load(kv_indptr + batch_idx)
kv_end_base = tl.load(kv_indptr + batch_idx + 1)
total_kv = kv_end_base - kv_start_base
if total_kv <= 0:
return
chunk_size = tl.cdiv(total_kv, SPLIT_K)
kv_start = kv_start_base + pid_sk * chunk_size
kv_end = tl.minimum(kv_end_base, kv_start + chunk_size)
offs_h = pid_h * BLOCK_H + tl.arange(0, BLOCK_H)
h_mask = offs_h < nq
offs_dv = tl.arange(0, 512)
if kv_start >= kv_end:
if SPLIT_K > 1:
m_ptrs = mid_M + batch_idx * (nq * SPLIT_K) + offs_h * SPLIT_K + pid_sk
tl.store(m_ptrs, -float("inf"), mask=h_mask)
l_ptrs = mid_L + batch_idx * (nq * SPLIT_K) + offs_h * SPLIT_K + pid_sk
tl.store(l_ptrs, 0.0, mask=h_mask)
o_ptrs = mid_O + batch_idx * (nq * SPLIT_K * 512) + offs_h[:, None] * (SPLIT_K * 512) + pid_sk * 512 + offs_dv[None, :]
tl.store(tl.multiple_of(o_ptrs, [1, 16]), 0.0, mask=h_mask[:, None])
return
kv_scale = tl.load(kv_scale_ptr)
combined_scale = kv_scale * sm_scale
offs_d_512 = tl.arange(0, 512)
offs_d_64 = tl.arange(512, 576)
q_ptrs_base_512 = Q + q_idx * stride_qz + offs_h[:, None] * stride_qh + offs_d_512[None, :] * stride_qd
q_512 = tl.load(tl.multiple_of(q_ptrs_base_512, [1,16]), mask=h_mask[:, None], other=0.0).to(tl.float32)
q_512 = (q_512 * combined_scale).to(tl.bfloat16)
q_ptrs_base_64 = Q + q_idx * stride_qz + offs_h[:, None] * stride_qh + offs_d_64[None, :] * stride_qd
q_64 = tl.load(tl.multiple_of(q_ptrs_base_64, [1,16]), mask=h_mask[:, None], other=0.0).to(tl.float32)
q_64 = (q_64 * combined_scale).to(tl.bfloat16)
m_i = tl.zeros([BLOCK_H], dtype=tl.float32) - float("inf")
l_i = tl.zeros([BLOCK_H], dtype=tl.float32)
acc = tl.zeros([BLOCK_H, 512], dtype=tl.float32)
offs_n = tl.arange(0, BLOCK_N)
for current_n in range(kv_start, kv_end, BLOCK_N):
n_idx = current_n + offs_n
kv_mask = n_idx < kv_end
k_ptrs_512 = KV + n_idx[:, None] * stride_kz + 0 * stride_kh + offs_d_512[None, :] * stride_kd
k_512 = tl.load(tl.multiple_of(k_ptrs_512, [1,16]), mask=kv_mask[:, None], other=0.0).to(tl.bfloat16)
qk = tl.dot(q_512, tl.trans(k_512))
k_ptrs_64 = KV + n_idx[:, None] * stride_kz + 0 * stride_kh + offs_d_64[None, :] * stride_kd
k_64 = tl.load(tl.multiple_of(k_ptrs_64, [1,16]), mask=kv_mask[:, None], other=0.0).to(tl.bfloat16)
qk += tl.dot(q_64, tl.trans(k_64))
qk = tl.where(kv_mask[None, :], qk, float("-inf"))
qk = tl.where(h_mask[:, None], qk, float("-inf"))
m_ij = tl.maximum(m_i, tl.max(qk, axis=1))
m_ij = tl.where(m_ij == float("-inf"), 0.0, m_ij)
p = tl.exp(qk - m_ij[:, None])
l_ij = tl.sum(p, axis=1)
alpha = tl.exp(m_i - m_ij)
l_i = l_i * alpha + l_ij
p_bf16 = p.to(tl.bfloat16)
acc = acc * alpha[:, None]
# =====================================================================
# 【核心优化 2】: 绝境逢生!不再去显存读取 v_chunk!
# 在 MLA 中,V 就是 K 的前 512 维!直接复用 SRAM 里的 k_512 进行矩阵乘加!
# 直接省去 50% 全局显存带宽,耗时瞬间暴跌!
# =====================================================================
acc += tl.dot(p_bf16, k_512)
m_i = m_ij
if SPLIT_K == 1:
acc = (acc / l_i[:, None]) * kv_scale
out_ptrs = Out + q_idx * stride_oz + offs_h[:, None] * stride_oh + offs_dv[None, :] * stride_od
tl.store(tl.multiple_of(out_ptrs, [1, 16]), acc.to(Out.dtype.element_ty), mask=h_mask[:, None])
else:
m_ptrs = mid_M + batch_idx * (nq * SPLIT_K) + offs_h * SPLIT_K + pid_sk
tl.store(m_ptrs, m_i, mask=h_mask)
l_ptrs = mid_L + batch_idx * (nq * SPLIT_K) + offs_h * SPLIT_K + pid_sk
tl.store(l_ptrs, l_i, mask=h_mask)
o_ptrs = mid_O + batch_idx * (nq * SPLIT_K * 512) + offs_h[:, None] * (SPLIT_K * 512) + pid_sk * 512 + offs_dv[None, :]
tl.store(tl.multiple_of(o_ptrs, [1, 16]), acc, mask=h_mask[:, None])
# =====================================================================
# Stage 2: 全局归约 (Reduction)
# =====================================================================
@triton.autotune(
configs=[
triton.Config({'BLOCK_H': 64}, num_stages=2, num_warps=4),
triton.Config({'BLOCK_H': 32}, num_stages=2, num_warps=4),
],
key=['nq', 'num_batches']
)
@triton.jit
def mla_decode_stage2(
mid_O, mid_M, mid_L, Out,
qo_indptr, kv_indptr, kv_scale_ptr,
stride_oz, stride_oh, stride_od,
nq, num_batches,
SPLIT_K: tl.constexpr,
BLOCK_H: tl.constexpr
):
pid_h = tl.program_id(0)
batch_idx = tl.program_id(1)
if batch_idx >= num_batches: return
offs_h = pid_h * BLOCK_H + tl.arange(0, BLOCK_H)
h_mask = offs_h < nq
offs_dv = tl.arange(0, 512)
kv_start = tl.load(kv_indptr + batch_idx)
kv_end = tl.load(kv_indptr + batch_idx + 1)
if kv_start >= kv_end:
q_idx = tl.load(qo_indptr + batch_idx)
out_ptrs = Out + q_idx * stride_oz + offs_h[:, None] * stride_oh + offs_dv[None, :] * stride_od
tl.store(tl.multiple_of(out_ptrs, [1,16]), 0.0, mask=h_mask[:, None])
return
kv_scale = tl.load(kv_scale_ptr)
global_m = tl.zeros([BLOCK_H], dtype=tl.float32) - float("inf")
global_l = tl.zeros([BLOCK_H], dtype=tl.float32)
acc = tl.zeros([BLOCK_H, 512], dtype=tl.float32)
for sk in range(SPLIT_K):
m_ptrs = mid_M + batch_idx * (nq * SPLIT_K) + offs_h * SPLIT_K + sk
m_sk = tl.load(m_ptrs, mask=h_mask, other=-float("inf"))
global_m = tl.maximum(global_m, m_sk)
global_m = tl.where(global_m == -float("inf"), 0.0, global_m)
for sk in range(SPLIT_K):
m_ptrs = mid_M + batch_idx * (nq * SPLIT_K) + offs_h * SPLIT_K + sk
m_sk = tl.load(m_ptrs, mask=h_mask, other=-float("inf"))
l_ptrs = mid_L + batch_idx * (nq * SPLIT_K) + offs_h * SPLIT_K + sk
l_sk = tl.load(l_ptrs, mask=h_mask, other=0.0)
alpha = tl.exp(m_sk - global_m)
global_l += l_sk * alpha
o_ptrs = mid_O + batch_idx * (nq * SPLIT_K * 512) + offs_h[:, None] * (SPLIT_K * 512) + sk * 512 + offs_dv[None, :]
o_vals = tl.load(tl.multiple_of(o_ptrs, [1, 16]), mask=h_mask[:, None], other=0.0)
acc += o_vals * alpha[:, None]
acc = (acc / global_l[:, None]) * kv_scale
q_idx = tl.load(qo_indptr + batch_idx)
out_ptrs = Out + q_idx * stride_oz + offs_h[:, None] * stride_oh + offs_dv[None, :] * stride_od
tl.store(tl.multiple_of(out_ptrs, [1,16]), acc.to(Out.dtype.element_ty), mask=h_mask[:, None])
def mla_decode_triton(q: torch.Tensor, kv_input: torch.Tensor, qo_indptr: torch.Tensor, kv_indptr: torch.Tensor, config: dict, kv_scale_tensor: torch.Tensor):
batch_size = config["batch_size"]
nq = config["num_heads"]
dq = config["qk_head_dim"]
dv = config["v_head_dim"]
sm_scale = 1.0 / (dq ** 0.5)
total_q = q.size(0)
outputs = torch.empty((total_q, nq, dv), device=q.device, dtype=torch.bfloat16)
# =====================================================================
# 并发激进嗅探器:哪怕只有 1024 的序列,也强行切割保证占满 GPU!
# =====================================================================
total_kv = kv_input.shape[0]
avg_seq_len = total_kv // max(1, batch_size)
# 调大目标块数,强行触发 MI355X 并发
target_blocks = 512
base_blocks = max(1, batch_size * max(1, nq // 64))
# 只要序列长于 128,就允许切割!
max_split = max(1, avg_seq_len // 128)
desired_split = triton.cdiv(target_blocks, base_blocks)
# 动态确定最佳分割度
SPLIT_K = min(16, min(max_split, desired_split))
MAX_SPLIT_K = 16
if SPLIT_K > 1:
if FlashDecodeCache.mid_o is None or FlashDecodeCache.mid_o.shape[0] < batch_size or FlashDecodeCache.mid_o.shape[2] < MAX_SPLIT_K:
FlashDecodeCache.mid_o = torch.zeros((batch_size, nq, MAX_SPLIT_K, dv), dtype=torch.float32, device=q.device)
FlashDecodeCache.mid_m = torch.full((batch_size, nq, MAX_SPLIT_K), float("-inf"), dtype=torch.float32, device=q.device)
FlashDecodeCache.mid_l = torch.zeros((batch_size, nq, MAX_SPLIT_K), dtype=torch.float32, device=q.device)
grid1 = lambda META: (triton.cdiv(nq, META['BLOCK_H']), batch_size, SPLIT_K)
mla_decode_stage1[grid1](
q, kv_input, outputs, FlashDecodeCache.mid_o, FlashDecodeCache.mid_m, FlashDecodeCache.mid_l,
qo_indptr, kv_indptr, kv_scale_tensor, sm_scale,
q.stride(0), q.stride(1), q.stride(2),
kv_input.stride(0), kv_input.stride(1), kv_input.stride(2),
# 移除了 kv_input 第三组 stride,内核签名已同步更新
outputs.stride(0), outputs.stride(1), outputs.stride(2),
nq, batch_size,
SPLIT_K=SPLIT_K,
)
if SPLIT_K > 1:
grid2 = lambda META: (triton.cdiv(nq, META['BLOCK_H']), batch_size)
mla_decode_stage2[grid2](
FlashDecodeCache.mid_o, FlashDecodeCache.mid_m, FlashDecodeCache.mid_l, outputs,
qo_indptr, kv_indptr, kv_scale_tensor,
outputs.stride(0), outputs.stride(1), outputs.stride(2),
nq, batch_size,
SPLIT_K=SPLIT_K,
)
return outputs
def custom_kernel(data):
if FlashDecodeCache.fp32_one is None or FlashDecodeCache.fp32_one.device != data[0].device:
FlashDecodeCache.fp32_one = torch.tensor([1.0], dtype=torch.float32, device=data[0].device)
q, kv_data, qo_indptr, kv_indptr, config = data
if "fp8" in kv_data:
kv_input, kv_scale_tensor = kv_data["fp8"]
if kv_scale_tensor is None:
kv_scale_tensor = FlashDecodeCache.fp32_one
else:
kv_input = kv_data.get("bf16", kv_data[list(kv_data.keys())[0]])
if isinstance(kv_input, tuple): kv_input = kv_input[0]
kv_scale_tensor = FlashDecodeCache.fp32_one
return mla_decode_triton(q, kv_input, qo_indptr, kv_indptr, config, kv_scale_tensor)scrolls · 289 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