submission 607604
Haolin · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 265 lines, June 9 Researcher Reciprocity License v1.0.
test_0321_1.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-607604?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:b6ec742b5cb0a9cdb7675d784cf61b8ab0f5638da9e7ee383ffed36ea3ba09f1
license declaredunknown
license concludedunknown
authorsHaolin
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
kv_buffer, kv_scale = kv_data["mxfp4"]mma
S += tl.dot(q_even_256, tl.trans(K_even_256))online-softmax
m_new = tl.maximum(m_global, m_i)split-k
MAX_SPLIT_K,Kernel source
test_0321_1.py265 lines
import torch
import triton
import triton.language as tl
# ---------------------------------------------------------------------------
# MXFP4 Dequantization Kernel Utils (E2M1)
# ---------------------------------------------------------------------------
@triton.jit
def dequant_nibble(n):
sign = (n >> 3) & 1
abs_val = n & 0x7
# FIX: Remove dtype= kwarg and use .to() explicit casting
f = tl.zeros_like(abs_val).to(tl.float32)
f = tl.where(abs_val == 1, 0.5, f)
f = tl.where(abs_val == 2, 1.0, f)
f = tl.where(abs_val == 3, 1.5, f)
f = tl.where(abs_val == 4, 2.0, f)
f = tl.where(abs_val == 5, 3.0, f)
f = tl.where(abs_val == 6, 4.0, f)
f = tl.where(abs_val == 7, 6.0, f)
return tl.where(sign, -f, f)
# ---------------------------------------------------------------------------
# Phase 1: Chunked Attention computation (Split-K)
# ---------------------------------------------------------------------------
@triton.jit
def mla_decode_mxfp4_kernel(
q_even_ptr, q_odd_ptr,
kv_buffer_ptr, kv_scale_ptr,
kv_indptr,
workspace_O_even, workspace_O_odd,
workspace_m, workspace_l,
stride_q_even_b, stride_q_even_h, stride_q_even_d,
stride_q_odd_b, stride_q_odd_h, stride_q_odd_d,
stride_kv_row, stride_scale_row,
sm_scale,
MAX_SPLIT_K,
BLOCK_KV: tl.constexpr
):
seq_idx = tl.program_id(0)
split_idx = tl.program_id(1)
kv_start = tl.load(kv_indptr + seq_idx)
kv_end = tl.load(kv_indptr + seq_idx + 1)
kv_len = kv_end - kv_start
chunk_start = split_idx * BLOCK_KV
if chunk_start >= kv_len:
return
# MQA sharing - load 16 heads. Split 288-dim loads into 256 + 32 to satisfy Triton power-of-2 rules
heads = tl.arange(0, 16)
cols_256 = tl.arange(0, 256)
cols_32 = tl.arange(0, 32)
# --- Load Q (256 chunk) ---
q_even_ptrs_256 = q_even_ptr + seq_idx * stride_q_even_b + heads[:, None] * stride_q_even_h + cols_256[None, :] * stride_q_even_d
q_odd_ptrs_256 = q_odd_ptr + seq_idx * stride_q_odd_b + heads[:, None] * stride_q_odd_h + cols_256[None, :] * stride_q_odd_d
q_even_256 = tl.load(q_even_ptrs_256)
q_odd_256 = tl.load(q_odd_ptrs_256)
# --- Load Q (32 chunk) ---
q_even_ptrs_32 = q_even_ptr + seq_idx * stride_q_even_b + heads[:, None] * stride_q_even_h + (256 + cols_32)[None, :] * stride_q_even_d
q_odd_ptrs_32 = q_odd_ptr + seq_idx * stride_q_odd_b + heads[:, None] * stride_q_odd_h + (256 + cols_32)[None, :] * stride_q_odd_d
q_even_32 = tl.load(q_even_ptrs_32)
q_odd_32 = tl.load(q_odd_ptrs_32)
# --- Setup KV Offsets ---
offs_kv = tl.arange(0, BLOCK_KV)
mask_kv = (chunk_start + offs_kv) < kv_len
row_offsets = kv_start + chunk_start + offs_kv
# --- Load KV and Scales (256 chunk) ---
kv_ptrs_256 = kv_buffer_ptr + row_offsets[:, None] * stride_kv_row + cols_256[None, :]
K_byte_256 = tl.load(kv_ptrs_256, mask=mask_kv[:, None])
scale_cols_256 = cols_256 // 16
scale_ptrs_256 = kv_scale_ptr + row_offsets[:, None] * stride_scale_row + scale_cols_256[None, :]
scale_uint8_256 = tl.load(scale_ptrs_256, mask=mask_kv[:, None])
scale_f32_256 = (scale_uint8_256.to(tl.uint32) << 23).to(tl.float32, bitcast=True)
low_256 = K_byte_256 & 0x0F
high_256 = (K_byte_256 >> 4) & 0x0F
K_even_256 = (dequant_nibble(low_256) * scale_f32_256).to(tl.bfloat16)
K_odd_256 = (dequant_nibble(high_256) * scale_f32_256).to(tl.bfloat16)
# --- Load KV and Scales (32 chunk) ---
kv_ptrs_32 = kv_buffer_ptr + row_offsets[:, None] * stride_kv_row + (256 + cols_32)[None, :]
K_byte_32 = tl.load(kv_ptrs_32, mask=mask_kv[:, None])
scale_cols_32 = (256 + cols_32) // 16
scale_ptrs_32 = kv_scale_ptr + row_offsets[:, None] * stride_scale_row + scale_cols_32[None, :]
scale_uint8_32 = tl.load(scale_ptrs_32, mask=mask_kv[:, None])
scale_f32_32 = (scale_uint8_32.to(tl.uint32) << 23).to(tl.float32, bitcast=True)
low_32 = K_byte_32 & 0x0F
high_32 = (K_byte_32 >> 4) & 0x0F
K_even_32 = (dequant_nibble(low_32) * scale_f32_32).to(tl.bfloat16)
K_odd_32 = (dequant_nibble(high_32) * scale_f32_32).to(tl.bfloat16)
# --- Compute attention scores (GEMM) ---
S = tl.zeros((16, BLOCK_KV), dtype=tl.float32)
S += tl.dot(q_even_256, tl.trans(K_even_256))
S += tl.dot(q_odd_256, tl.trans(K_odd_256))
S += tl.dot(q_even_32, tl.trans(K_even_32))
S += tl.dot(q_odd_32, tl.trans(K_odd_32))
S = S * sm_scale
# --- Masking and Softmax ---
S = tl.where(mask_kv[None, :], S, float('-inf'))
m_i = tl.max(S, axis=1)
p = tl.exp(S - m_i[:, None])
l_i = tl.sum(p, axis=1)
p_bf16 = p.to(tl.bfloat16)
# --- Compute Output ---
O_even = tl.dot(p_bf16, K_even_256)
O_odd = tl.dot(p_bf16, K_odd_256)
# Store intermediate results
cols_out = tl.arange(0, 256)
out_offs = seq_idx * (MAX_SPLIT_K * 16 * 256) + split_idx * (16 * 256) + heads[:, None] * 256 + cols_out[None, :]
tl.store(workspace_O_even + out_offs, O_even)
tl.store(workspace_O_odd + out_offs, O_odd)
ml_offs = seq_idx * (MAX_SPLIT_K * 16) + split_idx * 16 + heads
tl.store(workspace_m + ml_offs, m_i)
tl.store(workspace_l + ml_offs, l_i)
# ---------------------------------------------------------------------------
# Phase 2: Reduction Kernel
# ---------------------------------------------------------------------------
@triton.jit
def mla_reduce_kernel(
workspace_O_even, workspace_O_odd, workspace_m, workspace_l,
out_even, out_odd,
kv_indptr,
MAX_SPLIT_K, BLOCK_KV: tl.constexpr
):
seq_idx = tl.program_id(0)
head_idx = tl.program_id(1)
kv_start = tl.load(kv_indptr + seq_idx)
kv_end = tl.load(kv_indptr + seq_idx + 1)
kv_len = kv_end - kv_start
if kv_len == 0:
return
num_splits = (kv_len + BLOCK_KV - 1) // BLOCK_KV
m_global = float('-inf')
l_global = 0.0
ml_base = workspace_m + seq_idx * (MAX_SPLIT_K * 16) + head_idx
# Pass 1: Max m and sum l over all sequence splits
for i in range(num_splits):
m_i = tl.load(ml_base + i * 16)
l_i = tl.load(workspace_l + seq_idx * (MAX_SPLIT_K * 16) + i * 16 + head_idx)
m_new = tl.maximum(m_global, m_i)
l_new = l_global * tl.exp(m_global - m_new) + l_i * tl.exp(m_i - m_new)
m_global = m_new
l_global = l_new
# Pass 2: Combine partial Attention Outputs
acc_even = tl.zeros((256,), dtype=tl.float32)
acc_odd = tl.zeros((256,), dtype=tl.float32)
cols = tl.arange(0, 256)
O_base_even = workspace_O_even + seq_idx * (MAX_SPLIT_K * 16 * 256) + head_idx * 256 + cols
O_base_odd = workspace_O_odd + seq_idx * (MAX_SPLIT_K * 16 * 256) + head_idx * 256 + cols
for i in range(num_splits):
m_i = tl.load(ml_base + i * 16)
weight = tl.exp(m_i - m_global)
O_even_i = tl.load(O_base_even + i * (16 * 256))
O_odd_i = tl.load(O_base_odd + i * (16 * 256))
acc_even += O_even_i * weight
acc_odd += O_odd_i * weight
acc_even = acc_even / l_global
acc_odd = acc_odd / l_global
out_even_ptr = out_even + seq_idx * (16 * 256) + head_idx * 256 + cols
out_odd_ptr = out_odd + seq_idx * (16 * 256) + head_idx * 256 + cols
tl.store(out_even_ptr, acc_even.to(tl.bfloat16))
tl.store(out_odd_ptr, acc_odd.to(tl.bfloat16))
# ---------------------------------------------------------------------------
# Host Function Wrapper
# ---------------------------------------------------------------------------
def custom_kernel(data) -> torch.Tensor:
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = config["batch_size"]
# 1. Extract MXFP4 cache
kv_buffer, kv_scale = kv_data["mxfp4"]
# View uint8 cache flat
kv_buffer = kv_buffer.view(torch.uint8).view(-1, 288)
# E8M0 Scale normalization
kv_scale = kv_scale.view(torch.uint8)
if kv_scale.dim() == 3:
kv_scale = kv_scale.view(kv_scale.shape[0], kv_scale.shape[-1])
# 2. Re-stride Query
q_even = q[:, :, 0::2].contiguous()
q_odd = q[:, :, 1::2].contiguous()
# 3. Dynamic Flash-Decoding Sizing
max_kv_len = 0
if batch_size > 0:
max_kv_len = int((kv_indptr[1:] - kv_indptr[:-1]).max().item())
BLOCK_KV = 64
MAX_SPLIT_K = max(1, (max_kv_len + BLOCK_KV - 1) // BLOCK_KV)
# 4. Global Workspace Allocation
workspace_O_even = torch.empty((batch_size, MAX_SPLIT_K, 16, 256), dtype=torch.float32, device='cuda')
workspace_O_odd = torch.empty((batch_size, MAX_SPLIT_K, 16, 256), dtype=torch.float32, device='cuda')
workspace_m = torch.empty((batch_size, MAX_SPLIT_K, 16), dtype=torch.float32, device='cuda')
workspace_l = torch.empty((batch_size, MAX_SPLIT_K, 16), dtype=torch.float32, device='cuda')
out_even = torch.empty((batch_size, 16, 256), dtype=torch.bfloat16, device='cuda')
out_odd = torch.empty((batch_size, 16, 256), dtype=torch.bfloat16, device='cuda')
# 5. Dispatch
if batch_size > 0 and max_kv_len > 0:
grid_1 = (batch_size, MAX_SPLIT_K)
mla_decode_mxfp4_kernel[grid_1](
q_even, q_odd,
kv_buffer, kv_scale,
kv_indptr,
workspace_O_even, workspace_O_odd,
workspace_m, workspace_l,
q_even.stride(0), q_even.stride(1), q_even.stride(2),
q_odd.stride(0), q_odd.stride(1), q_odd.stride(2),
kv_buffer.stride(0), kv_scale.stride(0),
config["sm_scale"],
MAX_SPLIT_K,
BLOCK_KV=BLOCK_KV
)
grid_2 = (batch_size, 16)
mla_reduce_kernel[grid_2](
workspace_O_even, workspace_O_odd, workspace_m, workspace_l,
out_even, out_odd,
kv_indptr,
MAX_SPLIT_K, BLOCK_KV
)
# 6. Reconstruct sequence dimensions seamlessly
out = torch.zeros((batch_size, 16, 512), dtype=torch.bfloat16, device='cuda')
out[:, :, 0::2] = out_even
out[:, :, 1::2] = out_odd
return outscrolls · 265 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