submission 596605
Jade · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 222 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-596605?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:459e1588e28abdaaba16b61bf837c376aae20bf8811097878f98a4d0df2b5131
license declaredunknown
license concludedunknown
authorsJade
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
qk = tl.dot(q1_bf16, tl.trans(k1_bf16)) + tl.dot(q2_bf16, tl.trans(k2_bf16))num-warps = 4
num_warps=4,online-softmax
m_new = tl.maximum(m_i, m_ij)persistent-kernel
num_splits = tl.num_programs(1)split-k
def mla_split_k_decode_kernel(stages = 2
num_stages=2,tile-k = 64
BLOCK_K=64,Kernel source
submission.py222 lines
"""
Optimized MLA Decode Kernel for MI355X.
V6 (The Apex Triton): Zero-Overhead Fast Path & Optimal CU Saturation.
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t
NUM_HEADS = 16
KV_LORA_RANK = 512
QK_ROPE_HEAD_DIM = 64
QK_HEAD_DIM = KV_LORA_RANK + QK_ROPE_HEAD_DIM
V_HEAD_DIM = KV_LORA_RANK
@triton.jit
def mla_split_k_decode_kernel(
Q, KV_FP8, KV_SCALE_VAL,
qo_indptr, kv_indptr,
Out_Partial, LSE_Partial,
stride_q_bs, stride_q_h, stride_q_d,
stride_kv_total, stride_kv_d,
stride_op_bs, stride_op_sp, stride_op_h, stride_op_d,
stride_lse_bs, stride_lse_sp, stride_lse_h,
sm_scale,
NUM_HEADS: tl.constexpr,
BLOCK_K: tl.constexpr,
):
batch_idx = tl.program_id(0)
split_idx = tl.program_id(1)
num_splits = tl.num_programs(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)
kv_len = kv_end - kv_start
if kv_len == 0: return
chunk_size = (kv_len + num_splits - 1) // num_splits
chunk_start = kv_start + split_idx * chunk_size
chunk_end = tl.minimum(kv_start + (split_idx + 1) * chunk_size, kv_end)
if chunk_start >= chunk_end:
lse_ptrs = LSE_Partial + batch_idx * stride_lse_bs + split_idx * stride_lse_sp + tl.arange(0, NUM_HEADS)
tl.store(lse_ptrs, tl.full([NUM_HEADS], -float('inf'), dtype=tl.float32))
return
offs_h = tl.arange(0, NUM_HEADS)
offs_d512 = tl.arange(0, 512)
offs_d64 = tl.arange(0, 64)
q1_ptrs = Q + q_start * stride_q_bs + offs_h[:, None] * stride_q_h + offs_d512[None, :] * stride_q_d
q2_ptrs = Q + q_start * stride_q_bs + offs_h[:, None] * stride_q_h + (512 + offs_d64)[None, :] * stride_q_d
q1 = tl.load(q1_ptrs)
q2 = tl.load(q2_ptrs)
combined_scale = sm_scale * KV_SCALE_VAL
q1_bf16 = (q1.to(tl.float32) * combined_scale).to(tl.bfloat16)
q2_bf16 = (q2.to(tl.float32) * combined_scale).to(tl.bfloat16)
m_i = tl.full([NUM_HEADS], -float('inf'), dtype=tl.float32)
l_i = tl.zeros([NUM_HEADS], dtype=tl.float32)
acc = tl.zeros([NUM_HEADS, 512], dtype=tl.float32)
for start_k in range(chunk_start, chunk_end, BLOCK_K):
offs_k = start_k + tl.arange(0, BLOCK_K)
valid_k = offs_k < chunk_end
safe_offs_k = tl.where(valid_k, offs_k, tl.zeros_like(offs_k))
k1_ptrs = KV_FP8 + safe_offs_k[:, None] * stride_kv_total + offs_d512[None, :] * stride_kv_d
k1_fp8 = tl.load(k1_ptrs, mask=valid_k[:, None], other=0.0)
k1_bf16 = k1_fp8.to(tl.float32).to(tl.bfloat16)
k2_ptrs = KV_FP8 + safe_offs_k[:, None] * stride_kv_total + (512 + offs_d64)[None, :] * stride_kv_d
k2_fp8 = tl.load(k2_ptrs, mask=valid_k[:, None], other=0.0)
k2_bf16 = k2_fp8.to(tl.float32).to(tl.bfloat16)
qk = tl.dot(q1_bf16, tl.trans(k1_bf16)) + tl.dot(q2_bf16, tl.trans(k2_bf16))
qk = tl.where(valid_k[None, :], qk, -float('inf'))
m_ij = tl.max(qk, axis=1)
m_new = tl.maximum(m_i, m_ij)
alpha = tl.exp(m_i - m_new)
p = tl.exp(qk - m_new[:, None])
l_i = l_i * alpha + tl.sum(p, axis=1)
v_ptrs = KV_FP8 + safe_offs_k[:, None] * stride_kv_total + offs_d512[None, :] * stride_kv_d
v_fp8 = tl.load(v_ptrs, mask=valid_k[:, None], other=0.0)
v_bf16 = v_fp8.to(tl.float32).to(tl.bfloat16)
acc = acc * alpha[:, None] + tl.dot(p.to(tl.bfloat16), v_bf16)
m_i = m_new
acc = (acc * KV_SCALE_VAL) / l_i[:, None]
lse = m_i + tl.math.log(l_i)
out_ptrs = Out_Partial + batch_idx * stride_op_bs + split_idx * stride_op_sp + offs_h[:, None] * stride_op_h + offs_d512[None, :] * stride_op_d
tl.store(out_ptrs, acc)
lse_ptrs = LSE_Partial + batch_idx * stride_lse_bs + split_idx * stride_lse_sp + offs_h
tl.store(lse_ptrs, lse)
@triton.jit
def reduce_kernel(
Out_Partial, LSE_Partial, Out,
stride_op_bs, stride_op_sp, stride_op_h, stride_op_d,
stride_lse_bs, stride_lse_sp, stride_lse_h,
stride_o_bs, stride_o_h, stride_o_d,
NUM_SPLITS: tl.constexpr,
NUM_HEADS: tl.constexpr,
):
batch_idx = tl.program_id(0)
head_idx = tl.program_id(1)
offs_d_v = tl.arange(0, 512)
offs_sp = tl.arange(0, NUM_SPLITS)
lse_ptrs = LSE_Partial + batch_idx * stride_lse_bs + offs_sp * stride_lse_sp + head_idx * stride_lse_h
lse = tl.load(lse_ptrs)
m_max = tl.max(lse, axis=0)
weights = tl.exp(lse - m_max)
sum_weights = tl.sum(weights, axis=0)
sum_weights = tl.where(sum_weights == 0.0, 1.0, sum_weights)
acc = tl.zeros([512], dtype=tl.float32)
for sp in range(NUM_SPLITS):
out_p_ptrs = Out_Partial + batch_idx * stride_op_bs + sp * stride_op_sp + head_idx * stride_op_h + offs_d_v * stride_op_d
val = tl.load(out_p_ptrs)
w = tl.load(LSE_Partial + batch_idx * stride_lse_bs + sp * stride_lse_sp + head_idx * stride_lse_h)
w_exp = tl.exp(w - m_max)
acc += val * w_exp
acc = acc / sum_weights
out_ptrs = Out + batch_idx * stride_o_bs + head_idx * stride_o_h + offs_d_v * stride_o_d
tl.store(out_ptrs, acc.to(tl.bfloat16))
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = qo_indptr.shape[0] - 1
total_q = q.shape[0]
kv_fp8_buffer, kv_scale = kv_data["fp8"]
kv_scale_val = float(kv_scale.item())
max_kv_len = int(torch.max(kv_indptr[1:] - kv_indptr[:-1]).item())
# 【极致物理调度】完美适配 MI355X 304个计算单元的切分策略
# 如果 Batch Size 够大,自带足够的并发度,绝对不切分!
if batch_size >= 128:
NUM_SPLITS = 1 # 产生 128 / 256 个 Block,占用率完美,无开销
elif batch_size >= 64:
NUM_SPLITS = 4 # 产生 256 个 Block
elif batch_size >= 32:
NUM_SPLITS = 8 # 产生 256 个 Block
else:
NUM_SPLITS = 32 # 小 Batch 强行切分 128 个 Block 防止 GPU 闲置
# 短序列无论如何都不切分
if max_kv_len <= 1024:
NUM_SPLITS = 1
out = torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), device=q.device, dtype=torch.bfloat16)
if NUM_SPLITS == 1:
# 【零开销快车道】彻底砍掉中间 33MB 内存的分配和第二道 Reduce 算子!
lse_dummy = torch.empty((batch_size, 1, NUM_HEADS), device=q.device, dtype=torch.float32)
grid_split = (batch_size, 1)
mla_split_k_decode_kernel[grid_split](
q, kv_fp8_buffer, kv_scale_val,
qo_indptr, kv_indptr,
out, lse_dummy, # 直接写入最终结果张量
q.stride(0), q.stride(1), q.stride(2),
kv_fp8_buffer.stride(0), kv_fp8_buffer.stride(2),
out.stride(0), 0, out.stride(1), out.stride(2), # 把 sp 维度的步长抹零
lse_dummy.stride(0), lse_dummy.stride(1), lse_dummy.stride(2),
config["sm_scale"],
NUM_HEADS=NUM_HEADS,
BLOCK_K=64,
num_warps=4,
num_stages=2,
)
return out
else:
# 【安全切分道】小 Batch 或极端长序列走经典切分归约
out_partial = torch.zeros((batch_size, NUM_SPLITS, NUM_HEADS, V_HEAD_DIM), device=q.device, dtype=torch.float32)
lse_partial = torch.full((batch_size, NUM_SPLITS, NUM_HEADS), -float('inf'), device=q.device, dtype=torch.float32)
grid_split = (batch_size, NUM_SPLITS)
mla_split_k_decode_kernel[grid_split](
q, kv_fp8_buffer, kv_scale_val,
qo_indptr, kv_indptr,
out_partial, lse_partial,
q.stride(0), q.stride(1), q.stride(2),
kv_fp8_buffer.stride(0), kv_fp8_buffer.stride(2),
out_partial.stride(0), out_partial.stride(1), out_partial.stride(2), out_partial.stride(3),
lse_partial.stride(0), lse_partial.stride(1), lse_partial.stride(2),
config["sm_scale"],
NUM_HEADS=NUM_HEADS,
BLOCK_K=64,
num_warps=4,
num_stages=2,
)
grid_reduce = (batch_size, NUM_HEADS)
reduce_kernel[grid_reduce](
out_partial, lse_partial, out,
out_partial.stride(0), out_partial.stride(1), out_partial.stride(2), out_partial.stride(3),
lse_partial.stride(0), lse_partial.stride(1), lse_partial.stride(2),
out.stride(0), out.stride(1), out.stride(2),
NUM_SPLITS=NUM_SPLITS,
NUM_HEADS=NUM_HEADS,
num_warps=4,
)
return outscrolls · 222 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