submission 755195
Apusx · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 443 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-755195?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:73cd89661e0eb0bac3eb8190c79d97f9ea9743b5fb77fb5d8b281f2d488e0848
license declaredunknown
license concludedunknown
authorsApusx
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
QKV_DTYPE = "mxfp4"mma
score = tl.dot(q1_even, tl.trans(low1_f16)) + tl.dot(q1_odd, tl.trans(high1_f16))num-warps = 8
num_warps=8,online-softmax
m_i_new = tl.maximum(m_i, tl.max(score, axis=1))persistent-kernel
num_head_groups = tl.num_programs(1)split-k
def mla_decode_mxfp4_kernel_splitk(stages = 3
num_stages=3tile-m = 4
BLOCK_M = 4Kernel source
submission.py443 lines
import torch
import torch.nn.functional as F
import triton
import triton.language as tl
from task import input_t, output_t
from aiter import dtypes as aiter_dtypes
# QKV dtype for custom_kernel dispatch: "bf16", "fp8", or "mxfp4"
QKV_DTYPE = "mxfp4"
@triton.jit
def fast_decompress(nibble):
x = (nibble & 7).to(tl.float32)
val = tl.where(x < 4, x * 0.5, tl.where(x == 7, 6.0, x - 2.0))
return tl.where((nibble & 8) != 0, -val, val)
@triton.jit
def mla_decode_mxfp4_kernel_splitk(
q_ptr, kv_ptr, kv_scale_ptr,
acc_even_ptr, acc_odd_ptr, m_ptr, l_ptr,
qo_indptr, kv_indptr,
stride_q_t, stride_q_h, stride_q_d,
stride_kv_t, stride_kv_h, stride_kv_d,
stride_kvs_t, stride_kvs_d,
stride_ws_b, stride_ws_hg, stride_ws_split, stride_ws_mh,
sm_scale, num_heads,
q_seq_len: tl.constexpr,
BLOCK_KV: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_H: tl.constexpr,
NUM_KV_SPLITS: tl.constexpr
):
batch_idx = tl.program_id(0)
head_group_idx = tl.program_id(1)
split_idx = tl.program_id(2)
q_start = tl.load(qo_indptr + batch_idx)
q_end = tl.load(qo_indptr + batch_idx + 1)
q_len = q_end - q_start
kv_start = tl.load(kv_indptr + batch_idx)
kv_end = tl.load(kv_indptr + batch_idx + 1)
kv_len = kv_end - kv_start
idx_mh = tl.arange(0, BLOCK_M * BLOCK_H)
idx_m = idx_mh // BLOCK_H
idx_h = head_group_idx * BLOCK_H + (idx_mh % BLOCK_H)
mask_mh = (idx_m < q_len) & (idx_h < num_heads)
offs_even_512 = tl.arange(0, 256) * 2
offs_odd_512 = tl.arange(0, 256) * 2 + 1
offs_even_64 = tl.arange(0, 32) * 2 + 512
offs_odd_64 = tl.arange(0, 32) * 2 + 513
q1_even = tl.load(q_ptr + (q_start + idx_m)[:, None] * stride_q_t + idx_h[:, None] * stride_q_h + offs_even_512[None, :] * stride_q_d, mask=mask_mh[:, None], other=0.0)
q1_odd = tl.load(q_ptr + (q_start + idx_m)[:, None] * stride_q_t + idx_h[:, None] * stride_q_h + offs_odd_512[None, :] * stride_q_d, mask=mask_mh[:, None], other=0.0)
q2_even = tl.load(q_ptr + (q_start + idx_m)[:, None] * stride_q_t + idx_h[:, None] * stride_q_h + offs_even_64[None, :] * stride_q_d, mask=mask_mh[:, None], other=0.0)
q2_odd = tl.load(q_ptr + (q_start + idx_m)[:, None] * stride_q_t + idx_h[:, None] * stride_q_h + offs_odd_64[None, :] * stride_q_d, mask=mask_mh[:, None], other=0.0)
q1_even = (q1_even * sm_scale).to(tl.float16)
q1_odd = (q1_odd * sm_scale).to(tl.float16)
q2_even = (q2_even * sm_scale).to(tl.float16)
q2_odd = (q2_odd * sm_scale).to(tl.float16)
m_i = tl.zeros([BLOCK_M * BLOCK_H], dtype=tl.float32) - float('inf')
l_i = tl.zeros([BLOCK_M * BLOCK_H], dtype=tl.float32)
acc_even = tl.zeros([BLOCK_M * BLOCK_H, 256], dtype=tl.float32)
acc_odd = tl.zeros([BLOCK_M * BLOCK_H, 256], dtype=tl.float32)
offs_kv = tl.arange(0, BLOCK_KV)
offs_v_uint8 = tl.arange(0, 256)
offs_r_uint8 = tl.arange(0, 32)
offs_scale1 = tl.arange(0, 16)
offs_scale2 = tl.arange(16, 18)
kv_buf_base = kv_ptr + offs_v_uint8[None, :] * stride_kv_d
kv_rot_base = kv_ptr + 256 * stride_kv_d + offs_r_uint8[None, :] * stride_kv_d
kvs_buf_base = kv_scale_ptr + offs_scale1[None, :] * stride_kvs_d
kvs_rot_base = kv_scale_ptr + offs_scale2[None, :] * stride_kvs_d
# SPLIT-K bounds
chunk_size = (kv_len + NUM_KV_SPLITS - 1) // NUM_KV_SPLITS
chunk_size = (chunk_size + BLOCK_KV - 1) // BLOCK_KV * BLOCK_KV
start_n = split_idx * chunk_size
end_n = start_n + chunk_size
if end_n > kv_len:
end_n = kv_len
if start_n > kv_len:
start_n = kv_len
# Pre-compute cursor pointers outside loop!
k1_ptrs = kv_buf_base + (kv_start + start_n + offs_kv)[:, None] * stride_kv_t
k2_ptrs = kv_rot_base + (kv_start + start_n + offs_kv)[:, None] * stride_kv_t
scale1_ptrs = kvs_buf_base + (kv_start + start_n + offs_kv)[:, None] * stride_kvs_t
scale2_ptrs = kvs_rot_base + (kv_start + start_n + offs_kv)[:, None] * stride_kvs_t
for k_step in range(start_n, end_n, BLOCK_KV):
mask_kv = (k_step + offs_kv) < kv_len
k1_uint8 = tl.load(k1_ptrs, mask=mask_kv[:, None], other=0)
k2_uint8 = tl.load(k2_ptrs, mask=mask_kv[:, None], other=0)
scale1_e8m0 = tl.load(scale1_ptrs, mask=mask_kv[:, None], other=127)
scale2_e8m0 = tl.load(scale2_ptrs, mask=mask_kv[:, None], other=127)
scale1_f32 = tl.exp2(scale1_e8m0.to(tl.float32) - 127.0)
scale2_f32 = tl.exp2(scale2_e8m0.to(tl.float32) - 127.0)
# Unpack via mathematical multi-polynomial to avoid array lookup
k1_low_f32 = fast_decompress(k1_uint8 & 0x0F)
k1_high_f32 = fast_decompress((k1_uint8 >> 4) & 0x0F)
k2_low_f32 = fast_decompress(k2_uint8 & 0x0F)
k2_high_f32 = fast_decompress((k2_uint8 >> 4) & 0x0F)
# Scaling using view broadcasts to avoid massive registry bloat
low1_scaled = tl.reshape(k1_low_f32, (BLOCK_KV, 16, 16)) * scale1_f32[:, :, None]
high1_scaled = tl.reshape(k1_high_f32, (BLOCK_KV, 16, 16)) * scale1_f32[:, :, None]
low1_f16 = tl.reshape(low1_scaled, (BLOCK_KV, 256)).to(tl.float16)
high1_f16 = tl.reshape(high1_scaled, (BLOCK_KV, 256)).to(tl.float16)
low2_scaled = tl.reshape(k2_low_f32, (BLOCK_KV, 2, 16)) * scale2_f32[:, :, None]
high2_scaled = tl.reshape(k2_high_f32, (BLOCK_KV, 2, 16)) * scale2_f32[:, :, None]
low2_f16 = tl.reshape(low2_scaled, (BLOCK_KV, 32)).to(tl.float16)
high2_f16 = tl.reshape(high2_scaled, (BLOCK_KV, 32)).to(tl.float16)
score = tl.dot(q1_even, tl.trans(low1_f16)) + tl.dot(q1_odd, tl.trans(high1_f16))
score += tl.dot(q2_even, tl.trans(low2_f16)) + tl.dot(q2_odd, tl.trans(high2_f16))
score = tl.where(mask_kv[None, :], score, float('-inf'))
m_i_new = tl.maximum(m_i, tl.max(score, axis=1))
alpha = tl.exp(m_i - m_i_new)
p = tl.exp(score - m_i_new[:, None])
l_i_new = alpha * l_i + tl.sum(p, axis=1)
p_f16 = p.to(tl.float16)
acc_even = acc_even * alpha[:, None] + tl.dot(p_f16, low1_f16)
acc_odd = acc_odd * alpha[:, None] + tl.dot(p_f16, high1_f16)
m_i = m_i_new
l_i = l_i_new
# Advance pointer cursors simply saving calculation overhead
k1_ptrs += BLOCK_KV * stride_kv_t
k2_ptrs += BLOCK_KV * stride_kv_t
scale1_ptrs += BLOCK_KV * stride_kvs_t
scale2_ptrs += BLOCK_KV * stride_kvs_t
num_head_groups = tl.num_programs(1)
offs_d = tl.arange(0, 256)
ws_base_acc = batch_idx * stride_ws_b + head_group_idx * stride_ws_hg + split_idx * stride_ws_split + idx_mh[:, None] * stride_ws_mh + offs_d[None, :]
ws_base_ml = batch_idx * (num_head_groups * NUM_KV_SPLITS * BLOCK_M * BLOCK_H) + head_group_idx * (NUM_KV_SPLITS * BLOCK_M * BLOCK_H) + split_idx * (BLOCK_M * BLOCK_H) + idx_mh
tl.store(acc_even_ptr + ws_base_acc, acc_even)
tl.store(acc_odd_ptr + ws_base_acc, acc_odd)
tl.store(m_ptr + ws_base_ml, m_i)
tl.store(l_ptr + ws_base_ml, l_i)
@triton.jit
def mla_decode_mxfp4_reduce(
acc_even_ptr, acc_odd_ptr, m_ptr, l_ptr, out_ptr,
qo_indptr,
stride_out_t, stride_out_h, stride_out_d,
stride_ws_b, stride_ws_hg, stride_ws_split, stride_ws_mh,
sm_scale, num_heads,
BLOCK_M: tl.constexpr,
BLOCK_H: tl.constexpr,
NUM_KV_SPLITS: tl.constexpr
):
batch_idx = tl.program_id(0)
head_group_idx = tl.program_id(1)
q_start = tl.load(qo_indptr + batch_idx)
q_len = tl.load(qo_indptr + batch_idx + 1) - q_start
idx_mh = tl.arange(0, BLOCK_M * BLOCK_H)
idx_m = idx_mh // BLOCK_H
idx_h_local = idx_mh % BLOCK_H
idx_h = head_group_idx * BLOCK_H + idx_h_local
mask_mh = (idx_m < q_len) & (idx_h < num_heads)
m_new = tl.zeros([BLOCK_M * BLOCK_H], dtype=tl.float32) - float('inf')
l_new = tl.zeros([BLOCK_M * BLOCK_H], dtype=tl.float32)
acc_even = tl.zeros([BLOCK_M * BLOCK_H, 256], dtype=tl.float32)
acc_odd = tl.zeros([BLOCK_M * BLOCK_H, 256], dtype=tl.float32)
num_head_groups = tl.num_programs(1)
offs_d = tl.arange(0, 256)
for split_idx in range(NUM_KV_SPLITS):
ws_base_acc = batch_idx * stride_ws_b + head_group_idx * stride_ws_hg + split_idx * stride_ws_split + idx_mh[:, None] * stride_ws_mh + offs_d[None, :]
ws_base_ml = batch_idx * (num_head_groups * NUM_KV_SPLITS * BLOCK_M * BLOCK_H) + head_group_idx * (NUM_KV_SPLITS * BLOCK_M * BLOCK_H) + split_idx * (BLOCK_M * BLOCK_H) + idx_mh
acc_even_k = tl.load(acc_even_ptr + ws_base_acc, mask=mask_mh[:, None], other=0.0)
acc_odd_k = tl.load(acc_odd_ptr + ws_base_acc, mask=mask_mh[:, None], other=0.0)
m_k = tl.load(m_ptr + ws_base_ml, mask=mask_mh, other=-float('inf'))
l_k = tl.load(l_ptr + ws_base_ml, mask=mask_mh, other=0.0)
m_next = tl.maximum(m_new, m_k)
alpha_new = tl.exp(m_new - m_next)
alpha_k = tl.exp(m_k - m_next)
alpha_new = tl.where(m_new == float('-inf'), 0.0, alpha_new)
alpha_k = tl.where(m_k == float('-inf'), 0.0, alpha_k)
l_new = l_new * alpha_new + l_k * alpha_k
acc_even = acc_even * alpha_new[:, None] + acc_even_k * alpha_k[:, None]
acc_odd = acc_odd * alpha_new[:, None] + acc_odd_k * alpha_k[:, None]
m_new = m_next
acc_even = acc_even / l_new[:, None]
acc_odd = acc_odd / l_new[:, None]
offs_even_512 = tl.arange(0, 256) * 2
offs_odd_512 = tl.arange(0, 256) * 2 + 1
out_ptrs_even = out_ptr + (q_start + idx_m)[:, None] * stride_out_t + idx_h[:, None] * stride_out_h + offs_even_512[None, :] * stride_out_d
out_ptrs_odd = out_ptr + (q_start + idx_m)[:, None] * stride_out_t + idx_h[:, None] * stride_out_h + offs_odd_512[None, :] * stride_out_d
tl.store(out_ptrs_even, acc_even.to(tl.bfloat16), mask=mask_mh[:, None])
tl.store(out_ptrs_odd, acc_odd.to(tl.bfloat16), mask=mask_mh[:, None])
@triton.jit
def mla_decode_mxfp4_kernel(
q_ptr, kv_ptr, kv_scale_ptr, out_ptr,
qo_indptr, kv_indptr,
stride_q_t, stride_q_h, stride_q_d,
stride_kv_t, stride_kv_h, stride_kv_d,
stride_kvs_t, stride_kvs_d,
stride_out_t, stride_out_h, stride_out_d,
sm_scale, num_heads,
q_seq_len: tl.constexpr,
BLOCK_KV: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_H: tl.constexpr
):
batch_idx = tl.program_id(0)
head_group_idx = tl.program_id(1)
q_start = tl.load(qo_indptr + batch_idx)
q_len = tl.load(qo_indptr + batch_idx + 1) - q_start
kv_start = tl.load(kv_indptr + batch_idx)
kv_len = tl.load(kv_indptr + batch_idx + 1) - kv_start
idx_mh = tl.arange(0, BLOCK_M * BLOCK_H)
idx_m = idx_mh // BLOCK_H
idx_h = head_group_idx * BLOCK_H + (idx_mh % BLOCK_H)
mask_mh = (idx_m < q_len) & (idx_h < num_heads)
offs_even_512 = tl.arange(0, 256) * 2
offs_odd_512 = tl.arange(0, 256) * 2 + 1
offs_even_64 = tl.arange(0, 32) * 2 + 512
offs_odd_64 = tl.arange(0, 32) * 2 + 513
q1_even = tl.load(q_ptr + (q_start + idx_m)[:, None] * stride_q_t + idx_h[:, None] * stride_q_h + offs_even_512[None, :] * stride_q_d, mask=mask_mh[:, None], other=0.0)
q1_odd = tl.load(q_ptr + (q_start + idx_m)[:, None] * stride_q_t + idx_h[:, None] * stride_q_h + offs_odd_512[None, :] * stride_q_d, mask=mask_mh[:, None], other=0.0)
q2_even = tl.load(q_ptr + (q_start + idx_m)[:, None] * stride_q_t + idx_h[:, None] * stride_q_h + offs_even_64[None, :] * stride_q_d, mask=mask_mh[:, None], other=0.0)
q2_odd = tl.load(q_ptr + (q_start + idx_m)[:, None] * stride_q_t + idx_h[:, None] * stride_q_h + offs_odd_64[None, :] * stride_q_d, mask=mask_mh[:, None], other=0.0)
q1_even = (q1_even * sm_scale).to(tl.float16)
q1_odd = (q1_odd * sm_scale).to(tl.float16)
q2_even = (q2_even * sm_scale).to(tl.float16)
q2_odd = (q2_odd * sm_scale).to(tl.float16)
m_i = tl.zeros([BLOCK_M * BLOCK_H], dtype=tl.float32) - float('inf')
l_i = tl.zeros([BLOCK_M * BLOCK_H], dtype=tl.float32)
acc_even = tl.zeros([BLOCK_M * BLOCK_H, 256], dtype=tl.float32)
acc_odd = tl.zeros([BLOCK_M * BLOCK_H, 256], dtype=tl.float32)
offs_kv = tl.arange(0, BLOCK_KV)
offs_v_uint8 = tl.arange(0, 256)
offs_r_uint8 = tl.arange(0, 32)
offs_scale1 = tl.arange(0, 16)
offs_scale2 = tl.arange(16, 18)
kv_buf_base = kv_ptr + offs_v_uint8[None, :] * stride_kv_d
kv_rot_base = kv_ptr + 256 * stride_kv_d + offs_r_uint8[None, :] * stride_kv_d
kvs_buf_base = kv_scale_ptr + offs_scale1[None, :] * stride_kvs_d
kvs_rot_base = kv_scale_ptr + offs_scale2[None, :] * stride_kvs_d
# Pre-compute cursor pointers
k1_ptrs = kv_buf_base + (kv_start + offs_kv)[:, None] * stride_kv_t
k2_ptrs = kv_rot_base + (kv_start + offs_kv)[:, None] * stride_kv_t
scale1_ptrs = kvs_buf_base + (kv_start + offs_kv)[:, None] * stride_kvs_t
scale2_ptrs = kvs_rot_base + (kv_start + offs_kv)[:, None] * stride_kvs_t
for k_step in range(0, kv_len, BLOCK_KV):
mask_kv = (k_step + offs_kv) < kv_len
k1_uint8 = tl.load(k1_ptrs, mask=mask_kv[:, None], other=0)
k2_uint8 = tl.load(k2_ptrs, mask=mask_kv[:, None], other=0)
scale1_e8m0 = tl.load(scale1_ptrs, mask=mask_kv[:, None], other=127)
scale2_e8m0 = tl.load(scale2_ptrs, mask=mask_kv[:, None], other=127)
scale1_f32 = tl.exp2(scale1_e8m0.to(tl.float32) - 127.0)
scale2_f32 = tl.exp2(scale2_e8m0.to(tl.float32) - 127.0)
k1_low_f32 = fast_decompress(k1_uint8 & 0x0F)
k1_high_f32 = fast_decompress((k1_uint8 >> 4) & 0x0F)
k2_low_f32 = fast_decompress(k2_uint8 & 0x0F)
k2_high_f32 = fast_decompress((k2_uint8 >> 4) & 0x0F)
low1_scaled = tl.reshape(k1_low_f32, (BLOCK_KV, 16, 16)) * scale1_f32[:, :, None]
high1_scaled = tl.reshape(k1_high_f32, (BLOCK_KV, 16, 16)) * scale1_f32[:, :, None]
low1_f16 = tl.reshape(low1_scaled, (BLOCK_KV, 256)).to(tl.float16)
high1_f16 = tl.reshape(high1_scaled, (BLOCK_KV, 256)).to(tl.float16)
low2_scaled = tl.reshape(k2_low_f32, (BLOCK_KV, 2, 16)) * scale2_f32[:, :, None]
high2_scaled = tl.reshape(k2_high_f32, (BLOCK_KV, 2, 16)) * scale2_f32[:, :, None]
low2_f16 = tl.reshape(low2_scaled, (BLOCK_KV, 32)).to(tl.float16)
high2_f16 = tl.reshape(high2_scaled, (BLOCK_KV, 32)).to(tl.float16)
score = tl.dot(q1_even, tl.trans(low1_f16)) + tl.dot(q1_odd, tl.trans(high1_f16))
score += tl.dot(q2_even, tl.trans(low2_f16)) + tl.dot(q2_odd, tl.trans(high2_f16))
score = tl.where(mask_kv[None, :], score, float('-inf'))
m_i_new = tl.maximum(m_i, tl.max(score, axis=1))
alpha = tl.exp(m_i - m_i_new)
p = tl.exp(score - m_i_new[:, None])
l_i_new = alpha * l_i + tl.sum(p, axis=1)
p_f16 = p.to(tl.float16)
acc_even = acc_even * alpha[:, None] + tl.dot(p_f16, low1_f16)
acc_odd = acc_odd * alpha[:, None] + tl.dot(p_f16, high1_f16)
m_i = m_i_new
l_i = l_i_new
k1_ptrs += BLOCK_KV * stride_kv_t
k2_ptrs += BLOCK_KV * stride_kv_t
scale1_ptrs += BLOCK_KV * stride_kvs_t
scale2_ptrs += BLOCK_KV * stride_kvs_t
acc_even = acc_even / l_i[:, None]
acc_odd = acc_odd / l_i[:, None]
out_ptrs_even = out_ptr + (q_start + idx_m)[:, None] * stride_out_t + idx_h[:, None] * stride_out_h + offs_even_512[None, :] * stride_out_d
out_ptrs_odd = out_ptr + (q_start + idx_m)[:, None] * stride_out_t + idx_h[:, None] * stride_out_h + offs_odd_512[None, :] * stride_out_d
tl.store(out_ptrs_even, acc_even.to(tl.bfloat16), mask=mask_mh[:, None])
tl.store(out_ptrs_odd, acc_odd.to(tl.bfloat16), mask=mask_mh[:, None])
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
num_heads = config["num_heads"]
sm_scale = config["sm_scale"]
kv_buffer_mxfp4, kv_scale_mxfp4 = kv_data["mxfp4"]
kv_buf = kv_buffer_mxfp4.view(torch.uint8)
kv_scale = kv_scale_mxfp4.view(torch.uint8)
batch_size = qo_indptr.shape[0] - 1
total_q = q.shape[0]
out = torch.empty((total_q, num_heads, 512), dtype=torch.bfloat16, device="cuda")
stride_q_t, stride_q_h, stride_q_d = q.stride()
stride_kv_t, stride_kv_h, stride_kv_d = kv_buf.stride()
stride_kvs_t, stride_kvs_d = kv_scale.stride()
stride_out_t, stride_out_h, stride_out_d = out.stride()
BLOCK_M = 4
BLOCK_H = 16
num_head_groups = (num_heads + BLOCK_H - 1) // BLOCK_H
num_blocks_base = batch_size * num_head_groups
kv_seq_len = config["kv_seq_len"]
target_blocks = 1024
NUM_KV_SPLITS = max(1, target_blocks // num_blocks_base)
max_splits_by_len = max(1, kv_seq_len // 128)
NUM_KV_SPLITS = min(NUM_KV_SPLITS, max_splits_by_len)
if kv_seq_len >= 8192:
NUM_KV_SPLITS = max(NUM_KV_SPLITS, 4)
if NUM_KV_SPLITS > 1:
# Split-K mapping
ws_shape_acc = (batch_size, num_head_groups, NUM_KV_SPLITS, BLOCK_M * BLOCK_H, 256)
ws_shape_ml = (batch_size, num_head_groups, NUM_KV_SPLITS, BLOCK_M * BLOCK_H)
acc_even_ws = torch.empty(ws_shape_acc, dtype=torch.float32, device="cuda")
acc_odd_ws = torch.empty(ws_shape_acc, dtype=torch.float32, device="cuda")
m_ws = torch.empty(ws_shape_ml, dtype=torch.float32, device="cuda")
l_ws = torch.empty(ws_shape_ml, dtype=torch.float32, device="cuda")
stride_ws_b, stride_ws_hg, stride_ws_split, stride_ws_mh, _ = acc_even_ws.stride()
grid_splitk = (batch_size, num_head_groups, NUM_KV_SPLITS)
mla_decode_mxfp4_kernel_splitk[grid_splitk](
q, kv_buf, kv_scale,
acc_even_ws, acc_odd_ws, m_ws, l_ws,
qo_indptr, kv_indptr,
stride_q_t, stride_q_h, stride_q_d,
stride_kv_t, stride_kv_h, stride_kv_d,
stride_kvs_t, stride_kvs_d,
stride_ws_b, stride_ws_hg, stride_ws_split, stride_ws_mh,
sm_scale, num_heads,
q_seq_len=BLOCK_M,
BLOCK_KV=128,
BLOCK_M=BLOCK_M,
BLOCK_H=BLOCK_H,
NUM_KV_SPLITS=NUM_KV_SPLITS,
num_warps=8,
num_stages=3
)
grid_reduce = (batch_size, num_head_groups)
mla_decode_mxfp4_reduce[grid_reduce](
acc_even_ws, acc_odd_ws, m_ws, l_ws, out,
qo_indptr,
stride_out_t, stride_out_h, stride_out_d,
stride_ws_b, stride_ws_hg, stride_ws_split, stride_ws_mh,
sm_scale, num_heads,
BLOCK_M=BLOCK_M,
BLOCK_H=BLOCK_H,
NUM_KV_SPLITS=NUM_KV_SPLITS,
num_warps=8
)
else:
grid = (batch_size, num_head_groups)
mla_decode_mxfp4_kernel[grid](
q, kv_buf, kv_scale, out,
qo_indptr, kv_indptr,
stride_q_t, stride_q_h, stride_q_d,
stride_kv_t, stride_kv_h, stride_kv_d,
stride_kvs_t, stride_kvs_d,
stride_out_t, stride_out_h, stride_out_d,
sm_scale, num_heads,
q_seq_len=BLOCK_M,
BLOCK_KV=128,
BLOCK_M=BLOCK_M,
BLOCK_H=BLOCK_H,
num_warps=8,
num_stages=3
)
return out
scrolls · 443 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