submission 634642
Akash Adsare · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 243 lines, June 9 Researcher Reciprocity License v1.0.
submissionv7.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-634642?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:a73f63b1c048ed1c7f3fc33ef4fbbff468373dcef76ce7c0e6fb26b0025b0485
license declaredunknown
license concludedunknown
authorsAkash Adsare
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_bf16, tl.trans(k_e1_bf16))num-warps = 4
num_warps = 4stages = 2
num_stages = 2tile-n = 64
BLOCK_N = 64Kernel source
submissionv7.py243 lines
import torch
import triton
import triton.language as tl
from task import input_t, output_t
@triton.jit
def fp4_to_float_bitwise(v_u8):
v = v_u8.to(tl.int32)
v_abs = v & 7
val_bits = 1056964608 + (v_abs << 22)
val_bits = tl.where(v_abs == 1, 1056964608, val_bits)
val_bits = tl.where(v_abs == 0, 0, val_bits)
val_bits |= (v & 8) << 28
return val_bits.to(tl.float32, bitcast=True)
@triton.jit
def mla_decode_mxfp4_kernel(
Q, KV_packed, KV_scales,
kv_indptr,
Partial_M, Partial_L, Partial_V,
sm_scale,
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)
split_idx = tl.program_id(1)
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 kv_len <= 0 or 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
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, :]
q_e1_bf16 = (tl.load(q_base + h_off * QK_HEAD_DIM + d_256 * 2).to(tl.float32) * sm_scale).to(tl.bfloat16)
q_o1_bf16 = (tl.load(q_base + h_off * QK_HEAD_DIM + d_256 * 2 + 1).to(tl.float32) * sm_scale).to(tl.bfloat16)
q_e2_bf16 = (tl.load(q_base + h_off * QK_HEAD_DIM + 512 + d_32 * 2).to(tl.float32) * sm_scale).to(tl.bfloat16)
q_o2_bf16 = (tl.load(q_base + h_off * QK_HEAD_DIM + 512 + d_32 * 2 + 1).to(tl.float32) * sm_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_v_even = tl.zeros([NUM_HEADS, 256], dtype=tl.float32)
acc_v_odd = tl.zeros([NUM_HEADS, 256], dtype=tl.float32)
PACKED_STRIDE = QK_HEAD_DIM // 2
for n_start in range(curr_start, curr_end, BLOCK_N):
offs = n_start + tl.arange(0, BLOCK_N)
n_mask = offs < curr_end
n_off = offs[:, None]
k_packed_1 = tl.load(
KV_packed + n_off * PACKED_STRIDE + d_256,
mask=n_mask[:, None], eviction_policy="evict_first",
)
scales_raw_1 = tl.load(
KV_scales + n_off * NUM_BLOCKS + tl.arange(0, 16)[None, :],
mask=n_mask[:, None], eviction_policy="evict_first",
)
scales_1 = ((scales_raw_1.to(tl.int32) & 0xFF) << 23).to(tl.float32, bitcast=True)
v_e1 = fp4_to_float_bitwise(k_packed_1 & 0xF)
k_e1_bf16 = tl.reshape(
tl.reshape(v_e1, [BLOCK_N, 16, 16]) * scales_1[:, :, None],
[BLOCK_N, 256]
).to(tl.bfloat16)
v_o1 = fp4_to_float_bitwise((k_packed_1 >> 4) & 0xF)
k_o1_bf16 = tl.reshape(
tl.reshape(v_o1, [BLOCK_N, 16, 16]) * scales_1[:, :, None],
[BLOCK_N, 256]
).to(tl.bfloat16)
k_packed_2 = tl.load(
KV_packed + n_off * PACKED_STRIDE + 256 + d_32,
mask=n_mask[:, None], eviction_policy="evict_first",
)
scales_raw_2 = tl.load(
KV_scales + n_off * NUM_BLOCKS + 16 + tl.arange(0, 2)[None, :],
mask=n_mask[:, None], eviction_policy="evict_first",
)
scales_2 = ((scales_raw_2.to(tl.int32) & 0xFF) << 23).to(tl.float32, bitcast=True)
v_e2 = fp4_to_float_bitwise(k_packed_2 & 0xF)
k_e2_bf16 = tl.reshape(
tl.reshape(v_e2, [BLOCK_N, 2, 16]) * scales_2[:, :, None],
[BLOCK_N, 32]
).to(tl.bfloat16)
v_o2 = fp4_to_float_bitwise((k_packed_2 >> 4) & 0xF)
k_o2_bf16 = tl.reshape(
tl.reshape(v_o2, [BLOCK_N, 2, 16]) * scales_2[:, :, None],
[BLOCK_N, 32]
).to(tl.bfloat16)
s = tl.dot(q_e1_bf16, tl.trans(k_e1_bf16))
s += tl.dot(q_o1_bf16, tl.trans(k_o1_bf16))
s += tl.dot(q_e2_bf16, tl.trans(k_e2_bf16))
s += tl.dot(q_o2_bf16, tl.trans(k_o2_bf16))
s = tl.where(n_mask[None, :], s, -float('inf'))
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)
p_bf16 = p.to(tl.bfloat16)
acc_v_even = acc_v_even * alpha[:, None] + tl.dot(p_bf16, k_e1_bf16)
acc_v_odd = acc_v_odd * alpha[:, None] + tl.dot(p_bf16, k_o1_bf16)
m_i = m_next
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)
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)
@triton.jit
def mla_reduce_kernel(
Partial_M, Partial_L, Partial_V,
Out,
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
acc_v_e = tl.zeros([256], dtype=tl.float32)
acc_v_o = tl.zeros([256], dtype=tl.float32)
d_256 = tl.arange(0, 256)
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
acc_v_e = acc_v_e * alpha_f + tl.load(Partial_V + off_p * V_HEAD_DIM + d_256) * alpha_s
acc_v_o = acc_v_o * alpha_f + tl.load(Partial_V + off_p * V_HEAD_DIM + 256 + d_256) * alpha_s
m_final = m_next
out_ptr = Out + (q_row_idx * NUM_HEADS + h_idx) * V_HEAD_DIM
inv_l = 1.0 / l_final
tl.store(out_ptr + d_256 * 2, (acc_v_e * inv_l).to(tl.bfloat16))
tl.store(out_ptr + d_256 * 2 + 1, (acc_v_o * inv_l).to(tl.bfloat16))
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
kv_packed, kv_scales = kv_data["mxfp4"]
kv_p_u8 = kv_packed.view(torch.uint8)
kv_s_u8 = kv_scales.view(torch.uint8)
total_q = q.shape[0]
NUM_HEADS = 16
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
NUM_BLOCKS = QK_HEAD_DIM // 32
batch_size = config["batch_size"]
kv_seqlen = config["kv_seq_len"]
BLOCK_N = 64
if kv_seqlen <= 1024:
if batch_size <= 4:
NUM_SPLITS = 8
elif batch_size <= 32:
NUM_SPLITS = 16
elif batch_size <= 64:
NUM_SPLITS = 4
else:
NUM_SPLITS = 2
else:
if batch_size <= 4:
NUM_SPLITS = 32
elif batch_size <= 32:
NUM_SPLITS = 8
elif batch_size <= 64:
NUM_SPLITS = 8
else:
NUM_SPLITS = 8
max_splits = max(1, kv_seqlen // BLOCK_N)
NUM_SPLITS = min(NUM_SPLITS, max_splits)
num_warps = 4
num_stages = 2
partial_m = torch.empty(total_q * NUM_HEADS * NUM_SPLITS,
dtype=torch.float32, device="cuda")
partial_l = torch.empty_like(partial_m)
partial_v = torch.empty(total_q * NUM_HEADS * NUM_SPLITS * V_HEAD_DIM,
dtype=torch.float32, device="cuda")
out = torch.empty((total_q, NUM_HEADS, V_HEAD_DIM),
dtype=torch.bfloat16, device="cuda")
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=num_warps, num_stages=num_stages,
)
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=2,
)
return outscrolls · 243 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