submission 608589
akasha_08267 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 316 lines, June 9 Researcher Reciprocity License v1.0.
submissionv2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-608589?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:31d456cdc6596a584eb08021b1855404b1b509a0b1456f6bd4d1ed9a787899b7
license declaredunknown
license concludedunknown
authorsakasha_08267
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
kv_packed, kv_scales = kv_data["mxfp4"]mma
s = tl.dot(q_e1, tl.trans(k_e1))num-warps = 4
num_warps=4,stages = 2
num_stages=2,tile-n = 64
BLOCK_N = 64Kernel source
submissionv2.py316 lines
import math
import torch
import triton
import triton.language as tl
from task import input_t, output_t
# ---------------------------------------------------------------------------
# DeepSeek R1 MLA constants
# ---------------------------------------------------------------------------
NUM_HEADS = 16
NUM_KV_HEADS = 1
KV_LORA_RANK = 512
QK_ROPE_HEAD_DIM = 64
QK_HEAD_DIM = KV_LORA_RANK + QK_ROPE_HEAD_DIM # 576
V_HEAD_DIM = KV_LORA_RANK # 512
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
# MXFP4 specifics
BLOCK_D = 32
NUM_BLOCKS = QK_HEAD_DIM // BLOCK_D # 18
NUM_SPLITS = 32
BLOCK_N = 64
# ---------------------------------------------------------------------------
# 4-bit OCP-MX to float32 conversion
# ---------------------------------------------------------------------------
@triton.jit
def fp4_to_float(v):
"""Convert 4-bit E2M1 (OCP MX) values to float32."""
sign = (v >> 3) & 1
exp = (v >> 1) & 3
mant = v & 1
# Denormal: exponent == 0
denorm = mant.to(tl.float32) * 0.5
# Normal: (1 + mant*0.5) * 2^(exp-1)
norm = (1.0 + mant.to(tl.float32) * 0.5) * tl.exp2((exp - 1).to(tl.float32))
val = tl.where(exp == 0, denorm, norm)
return tl.where(sign > 0, -val, val)
# ---------------------------------------------------------------------------
# Main MXFP4 decode kernel (persistent KV splits)
# ---------------------------------------------------------------------------
@triton.jit
def mla_decode_mxfp4_kernel(
Q, # (total_q, NUM_HEADS, QK_HEAD_DIM) bf16
KV_packed, # (total_kv, 1, QK_HEAD_DIM//2) uint8 / fp4x2
KV_scales, # (total_kv, 1, NUM_BLOCKS) uint8 (E8M0)
kv_indptr, # (batch_size+1,) int32
Partial_M, # (total_q * NUM_HEADS * NUM_SPLITS) float32
Partial_L, # (total_q * NUM_HEADS * NUM_SPLITS) float32
Partial_V, # (total_q * NUM_HEADS * NUM_SPLITS * V_HEAD_DIM) float32
sm_scale, # float32
NUM_HEADS: tl.constexpr,
QK_HEAD_DIM: tl.constexpr,
V_HEAD_DIM: tl.constexpr,
NUM_SPLITS: tl.constexpr,
BLOCK_N: tl.constexpr,
NUM_BLOCKS: tl.constexpr,
):
q_row_idx = tl.program_id(0) # which query row
split_idx = tl.program_id(1) # which split within that row
# -----------------------------
# 1. Determine KV window
# -----------------------------
kv_start = tl.load(kv_indptr + q_row_idx)
kv_end = tl.load(kv_indptr + q_row_idx + 1)
kv_len = kv_end - kv_start
kv_per_split = (kv_len + NUM_SPLITS - 1) // NUM_SPLITS
curr_start = kv_start + split_idx * kv_per_split
curr_end = tl.minimum(curr_start + kv_per_split, kv_end)
# If no KV tokens for this row or this split, write neutral partials and exit.
if (kv_len <= 0) | (curr_start >= curr_end):
h_range = tl.arange(0, NUM_HEADS)
off_p = (q_row_idx * NUM_HEADS + h_range) * NUM_SPLITS + split_idx
tl.store(Partial_M + off_p, tl.full([NUM_HEADS], -float('inf'), dtype=tl.float32))
tl.store(Partial_L + off_p, tl.zeros([NUM_HEADS], dtype=tl.float32))
return
# -----------------------------
# 2. Load Q for all heads into registers
# -----------------------------
q_base = Q + q_row_idx * NUM_HEADS * QK_HEAD_DIM
h_off = tl.arange(0, NUM_HEADS)[:, None]
d_256 = tl.arange(0, 256)[None, :]
d_32 = tl.arange(0, 32)[None, :]
# 576 dims = 2*256 + 2*32 (even/odd)
q_e1 = tl.load(q_base + h_off * QK_HEAD_DIM + d_256 * 2).to(tl.float32)
q_o1 = tl.load(q_base + h_off * QK_HEAD_DIM + d_256 * 2 + 1).to(tl.float32)
q_e2 = tl.load(q_base + h_off * QK_HEAD_DIM + (256 + d_32) * 2).to(tl.float32)
q_o2 = tl.load(q_base + h_off * QK_HEAD_DIM + (256 + d_32) * 2 + 1).to(tl.float32)
# -----------------------------
# 3. Online-softmax accumulators
# -----------------------------
m_i = tl.full([NUM_HEADS], -float('inf'), dtype=tl.float32)
l_i = tl.zeros([NUM_HEADS], dtype=tl.float32)
acc_v_even = tl.zeros([NUM_HEADS, 256], dtype=tl.float32) # first 256 of 512
acc_v_odd = tl.zeros([NUM_HEADS, 256], dtype=tl.float32) # last 256 of 512
# -----------------------------
# 4. Iterate over KV tokens in BLOCK_N tiles
# -----------------------------
for n_start in range(curr_start, curr_end, BLOCK_N):
n_offsets = n_start + tl.arange(0, BLOCK_N)
n_mask = n_offsets < curr_end
# 4.1 Load packed MXFP4 KV
# Each row has QK_HEAD_DIM//2 bytes (two 4-bit values per byte).
kv_row_stride_bytes = QK_HEAD_DIM // 2
# First 256 dims use the first 256 bytes (512 values)
k_packed_1 = tl.load(
KV_packed + n_offsets[:, None] * kv_row_stride_bytes + d_256,
mask=n_mask[:, None],
other=0,
eviction_policy="evict_first"
)
# Last 64 dims use the next 32 bytes
k_packed_2 = tl.load(
KV_packed + n_offsets[:, None] * kv_row_stride_bytes + 256 + d_32,
mask=n_mask[:, None],
other=0,
eviction_policy="evict_first"
)
# 4.2 Load block scales (E8M0, 18 blocks total)
# First 8 blocks cover the first 256 dims, the next 2 for the 64 dims.
scales_raw_1 = tl.load(
KV_scales + n_offsets[:, None] * NUM_BLOCKS + tl.arange(0, 8)[None, :],
mask=n_mask[:, None],
other=0,
eviction_policy="evict_first"
)
scales_raw_2 = tl.load(
KV_scales + n_offsets[:, None] * NUM_BLOCKS + 8 + tl.arange(0, 2)[None, :],
mask=n_mask[:, None],
other=0,
eviction_policy="evict_first"
)
# Convert E8M0 to float scalars
scales_1 = tl.exp2(scales_raw_1.to(tl.float32) - 127.0) # (BLOCK_N, 8)
scales_2 = tl.exp2(scales_raw_2.to(tl.float32) - 127.0) # (BLOCK_N, 2)
# 4.3 Unpack MXFP4 → float32 and apply scales
# nibble low/high
v_e1 = fp4_to_float((k_packed_1 & 0xF).to(tl.int32))
v_o1 = fp4_to_float(((k_packed_1 >> 4) & 0xF).to(tl.int32))
v_e2 = fp4_to_float((k_packed_2 & 0xF).to(tl.int32))
v_o2 = fp4_to_float(((k_packed_2 >> 4) & 0xF).to(tl.int32))
# First 256 dims: 8 blocks x 32 features = 256.
k_e1 = tl.reshape(
tl.reshape(v_e1, [BLOCK_N, 8, 32]) * scales_1[:, :, None],
[BLOCK_N, 256]
)
k_o1 = tl.reshape(
tl.reshape(v_o1, [BLOCK_N, 8, 32]) * scales_1[:, :, None],
[BLOCK_N, 256]
)
# Last 64 dims: 2 blocks x 16 features = 32 per half.
k_e2 = tl.reshape(
tl.reshape(v_e2, [BLOCK_N, 2, 16]) * scales_2[:, :, None],
[BLOCK_N, 32]
)
k_o2 = tl.reshape(
tl.reshape(v_o2, [BLOCK_N, 2, 16]) * scales_2[:, :, None],
[BLOCK_N, 32]
)
# 4.4 Compute scores for each head: s = q·k^T
# q_e1: (16, 256), k_e1: (BLOCK_N, 256) → s1: (16, BLOCK_N)
s = tl.dot(q_e1, tl.trans(k_e1))
s += tl.dot(q_e2, tl.trans(k_e2))
s += tl.dot(q_o1, tl.trans(k_o1))
s += tl.dot(q_o2, tl.trans(k_o2))
s = s * sm_scale
s = tl.where(n_mask[None, :], s, -float('inf'))
# 4.5 Online softmax update
m_ij = tl.max(s, axis=1)
m_next = tl.maximum(m_i, m_ij)
alpha = tl.exp(m_i - m_next)
p = tl.exp(s - m_next[:, None])
l_i = l_i * alpha + tl.sum(p, axis=1)
# Value accumulation uses only first 512 dims → first 256 even + 256 odd.
acc_v_even = acc_v_even * alpha[:, None] + tl.dot(p, k_e1)
acc_v_odd = acc_v_odd * alpha[:, None] + tl.dot(p, k_o1)
m_i = m_next
# -----------------------------
# 5. Write partial results
# -----------------------------
h_range = tl.arange(0, NUM_HEADS)
off_p = (q_row_idx * NUM_HEADS + h_range) * NUM_SPLITS + split_idx
tl.store(Partial_M + off_p, m_i)
tl.store(Partial_L + off_p, l_i)
# Flatten V-head dimension: each head/split has V_HEAD_DIM = 512 values.
tl.store(Partial_V + off_p[:, None] * V_HEAD_DIM + d_256, acc_v_even)
tl.store(Partial_V + off_p[:, None] * V_HEAD_DIM + 256 + d_256, acc_v_odd)
# ---------------------------------------------------------------------------
# Reduction kernel over splits
# ---------------------------------------------------------------------------
@triton.jit
def mla_reduce_kernel(
Partial_M, Partial_L, Partial_V,
Out, # (total_q, NUM_HEADS, V_HEAD_DIM) bf16
NUM_HEADS: tl.constexpr,
V_HEAD_DIM: tl.constexpr,
NUM_SPLITS: tl.constexpr,
):
q_row_idx = tl.program_id(0)
h_idx = tl.program_id(1)
m_final = -float('inf')
l_final = 0.0
d_256 = tl.arange(0, 256)
acc_v_e = tl.zeros([256], dtype=tl.float32)
acc_v_o = tl.zeros([256], dtype=tl.float32)
# Reduce across splits
for s in range(NUM_SPLITS):
off_p = (q_row_idx * NUM_HEADS + h_idx) * NUM_SPLITS + s
m_s = tl.load(Partial_M + off_p)
l_s = tl.load(Partial_L + off_p)
m_next = tl.maximum(m_final, m_s)
alpha_f = tl.exp(m_final - m_next)
alpha_s = tl.exp(m_s - m_next)
l_final = l_final * alpha_f + l_s * alpha_s
v_e_s = tl.load(Partial_V + off_p * V_HEAD_DIM + d_256)
v_o_s = tl.load(Partial_V + off_p * V_HEAD_DIM + 256 + d_256)
acc_v_e = acc_v_e * alpha_f + v_e_s * alpha_s
acc_v_o = acc_v_o * alpha_f + v_o_s * alpha_s
m_final = m_next
# Normalize
acc_v_e = acc_v_e / l_final
acc_v_o = acc_v_o / l_final
# Write out
out_ptr = Out + (q_row_idx * NUM_HEADS + h_idx) * V_HEAD_DIM
tl.store(out_ptr + d_256 * 2, acc_v_e.to(tl.bfloat16))
tl.store(out_ptr + d_256 * 2 + 1, acc_v_o.to(tl.bfloat16))
# ---------------------------------------------------------------------------
# Python wrapper matching (q, kv_data, qo_indptr, kv_indptr, config) -> out
# ---------------------------------------------------------------------------
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
kv_packed, kv_scales = kv_data["mxfp4"]
# Ensure uint8 for Triton loads
kv_p_u8 = kv_packed.view(torch.uint8)
kv_s_u8 = kv_scales.view(torch.uint8)
total_q = q.shape[0]
device = q.device
# Allocate partial buffers
partial_m = torch.empty(total_q * NUM_HEADS * NUM_SPLITS, dtype=torch.float32, device=device)
partial_l = torch.empty_like(partial_m)
partial_v = torch.empty(total_q * NUM_HEADS * NUM_SPLITS * V_HEAD_DIM, dtype=torch.float32, device=device)
out = torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=device)
# Launch decode kernel: grid = (total_q, NUM_SPLITS)
mla_decode_mxfp4_kernel[(total_q, NUM_SPLITS)](
q, kv_p_u8, kv_s_u8, kv_indptr,
partial_m, partial_l, partial_v,
float(config["sm_scale"]),
NUM_HEADS=NUM_HEADS,
QK_HEAD_DIM=QK_HEAD_DIM,
V_HEAD_DIM=V_HEAD_DIM,
NUM_SPLITS=NUM_SPLITS,
BLOCK_N=BLOCK_N,
NUM_BLOCKS=NUM_BLOCKS,
num_warps=4,
num_stages=2,
)
# Launch reduction kernel: grid = (total_q, NUM_HEADS)
mla_reduce_kernel[(total_q, NUM_HEADS)](
partial_m, partial_l, partial_v, out,
NUM_HEADS=NUM_HEADS,
V_HEAD_DIM=V_HEAD_DIM,
NUM_SPLITS=NUM_SPLITS,
num_warps=4,
num_stages=2,
)
return outscrolls · 316 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