submission 589625
Purple rain · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 486 lines, June 9 Researcher Reciprocity License v1.0.
submission_amd_mixed_mla_mxfp4_hip_splitk.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-589625?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:7d7a38996bfc85788aa29103e0515b40022d039877e56b5a312789fa114b0c3d
license declaredunknown
license concludedunknown
authorsPurple rain
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
kv_buffer_mxfp4, kv_scale_mxfp4 = kv_data["mxfp4"]shared-memory
__shared__ float smem_k[kQkHeadDim];split-k
name="mla_mxfp4_splitk_hip_v1",Kernel source
submission_amd_mixed_mla_mxfp4_hip_splitk.py486 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
import os
import threading
import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline
_HIP_LOCK = threading.Lock()
_HIP_MODULE = None
_HIP_BUILD_ERROR = None
_CHUNK_ENV = "MLA_MXFP4_CHUNK_SIZE"
_DEFAULT_CHUNK_SIZE = 128
CPP_WRAPPER = r"""
#include <torch/extension.h>
torch::Tensor mla_decode_forward(
torch::Tensor q,
torch::Tensor kv_buffer,
torch::Tensor kv_scale,
torch::Tensor qo_indptr,
torch::Tensor kv_indptr,
int64_t config_v_dim,
int64_t chunk_size,
double sm_scale);
"""
HIP_SRC = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <hip/hip_bfloat16.h>
#include <cstdint>
#include <cmath>
namespace {
constexpr int kNumHeads = 16;
constexpr int kQkHeadDim = 576;
constexpr int kVHeadDim = 512;
constexpr int kScaleBlock = 32;
__constant__ float kFp4Lut[16] = {
0.0f, 0.5f, 1.0f, 1.5f,
2.0f, 3.0f, 4.0f, 6.0f,
-0.0f, -0.5f, -1.0f, -1.5f,
-2.0f, -3.0f, -4.0f, -6.0f
};
__device__ __forceinline__ float e8m0_to_float(uint8_t e8m0_val) {
uint32_t bits = static_cast<uint32_t>(e8m0_val) << 23;
if (e8m0_val == 0) {
bits = 0x00400000u;
} else if (e8m0_val == 0xFF) {
bits = 0x7F800001u;
}
return __uint_as_float(bits);
}
__device__ __forceinline__ float unpack_mxfp4(uint8_t val, float scale, bool high_nibble) {
uint8_t nibble = high_nibble ? static_cast<uint8_t>(val >> 4) : static_cast<uint8_t>(val & 0x0F);
return kFp4Lut[nibble & 0x0F] * scale;
}
__device__ __forceinline__ float wavefront_sum_64(float v) {
for (int offset = 32; offset > 0; offset >>= 1) {
v += __shfl_down(v, offset);
}
return v;
}
__global__ void mla_decode_chunk_kernel(
const hip_bfloat16* __restrict__ q,
const uint8_t* __restrict__ kv_data,
const uint8_t* __restrict__ kv_scale,
const int32_t* __restrict__ qo_indptr,
const int32_t* __restrict__ kv_indptr,
float* __restrict__ out_partials,
float* __restrict__ lse_partials,
int kv_row_stride,
int kv_scale_row_stride,
int qk_head_dim,
int v_head_dim,
float sm_scale,
int chunk_size,
int num_chunks) {
const int chunk_idx = static_cast<int>(blockIdx.x);
const int batch_idx = static_cast<int>(blockIdx.y);
const int lane = static_cast<int>(threadIdx.x);
const int head = static_cast<int>(threadIdx.y);
if (head >= kNumHeads || chunk_idx >= num_chunks) {
return;
}
const int32_t q_start = qo_indptr[batch_idx];
const int32_t q_end = qo_indptr[batch_idx + 1];
if (q_end <= q_start) {
return;
}
const int32_t q_idx = q_start;
const int32_t kv_start = kv_indptr[batch_idx];
const int32_t kv_end = kv_indptr[batch_idx + 1];
const int32_t token_begin = kv_start + chunk_idx * chunk_size;
const int32_t token_end = token_begin + chunk_size < kv_end ? token_begin + chunk_size : kv_end;
const hip_bfloat16* q_row = q + (static_cast<int64_t>(q_idx) * kNumHeads + head) * qk_head_dim;
float q_frag[9];
#pragma unroll
for (int i = 0; i < 9; ++i) {
const int d = lane + i * 64;
q_frag[i] = (d < qk_head_dim) ? static_cast<float>(q_row[d]) : 0.0f;
}
float o_frag[8];
#pragma unroll
for (int i = 0; i < 8; ++i) {
o_frag[i] = 0.0f;
}
float m = -INFINITY;
float l = 0.0f;
__shared__ float smem_k[kQkHeadDim];
__shared__ float smem_v[kVHeadDim];
for (int32_t token = token_begin; token < token_end; ++token) {
if (head == 0) {
const uint8_t* kv_row = kv_data + static_cast<int64_t>(token) * kv_row_stride;
const uint8_t* scale_row = kv_scale + static_cast<int64_t>(token) * kv_scale_row_stride;
for (int d = lane; d < qk_head_dim; d += 64) {
const uint8_t packed = kv_row[d >> 1];
const float scale = e8m0_to_float(scale_row[d / kScaleBlock]);
const float val = unpack_mxfp4(packed, scale, (d & 1) != 0);
smem_k[d] = val;
if (d < v_head_dim) {
smem_v[d] = val;
}
}
}
__syncthreads();
float qk_local = 0.0f;
#pragma unroll
for (int i = 0; i < 9; ++i) {
const int d = lane + i * 64;
if (d < qk_head_dim) {
qk_local += q_frag[i] * smem_k[d];
}
}
float score = wavefront_sum_64(qk_local);
score = (lane == 0) ? score * sm_scale : 0.0f;
score = __shfl(score, 0);
const float new_m = fmaxf(m, score);
const float alpha = expf(m - new_m);
const float p = expf(score - new_m);
const float new_l = l * alpha + p;
#pragma unroll
for (int i = 0; i < 8; ++i) {
const int d = lane + i * 64;
o_frag[i] = o_frag[i] * alpha + p * smem_v[d];
}
m = new_m;
l = new_l;
__syncthreads();
}
const int64_t partial_base =
(((static_cast<int64_t>(batch_idx) * kNumHeads + head) * num_chunks + chunk_idx) * v_head_dim);
#pragma unroll
for (int i = 0; i < 8; ++i) {
const int d = lane + i * 64;
out_partials[partial_base + d] = o_frag[i];
}
if (lane == 0) {
const int64_t lse_idx =
(((static_cast<int64_t>(batch_idx) * kNumHeads + head) * num_chunks + chunk_idx) * 2);
lse_partials[lse_idx] = m;
lse_partials[lse_idx + 1] = l;
}
}
__global__ void mla_decode_reduce_kernel(
const float* __restrict__ out_partials,
const float* __restrict__ lse_partials,
hip_bfloat16* __restrict__ out,
const int32_t* __restrict__ qo_indptr,
int num_chunks,
int v_head_dim) {
const int batch_idx = static_cast<int>(blockIdx.x);
const int head = static_cast<int>(blockIdx.y);
const int tid = static_cast<int>(threadIdx.x);
__shared__ float s_m;
__shared__ float s_l;
if (tid == 0) {
float m = -INFINITY;
for (int c = 0; c < num_chunks; ++c) {
const int64_t lse_idx =
(((static_cast<int64_t>(batch_idx) * kNumHeads + head) * num_chunks + c) * 2);
const float mc = lse_partials[lse_idx];
const float lc = lse_partials[lse_idx + 1];
if (lc > 0.0f) {
m = fmaxf(m, mc);
}
}
if (!isfinite(m)) {
s_m = -INFINITY;
s_l = 0.0f;
} else {
float l = 0.0f;
for (int c = 0; c < num_chunks; ++c) {
const int64_t lse_idx =
(((static_cast<int64_t>(batch_idx) * kNumHeads + head) * num_chunks + c) * 2);
const float mc = lse_partials[lse_idx];
const float lc = lse_partials[lse_idx + 1];
if (lc > 0.0f) {
l += lc * expf(mc - m);
}
}
s_m = m;
s_l = l;
}
}
__syncthreads();
const int32_t q_idx = qo_indptr[batch_idx];
const int64_t out_base = (static_cast<int64_t>(q_idx) * kNumHeads + head) * v_head_dim;
for (int d = tid; d < v_head_dim; d += static_cast<int>(blockDim.x)) {
float acc = 0.0f;
if (s_l > 0.0f) {
for (int c = 0; c < num_chunks; ++c) {
const int64_t lse_idx =
(((static_cast<int64_t>(batch_idx) * kNumHeads + head) * num_chunks + c) * 2);
const float mc = lse_partials[lse_idx];
const float lc = lse_partials[lse_idx + 1];
if (lc > 0.0f) {
const float beta = expf(mc - s_m);
const int64_t p_idx =
(((static_cast<int64_t>(batch_idx) * kNumHeads + head) * num_chunks + c) * v_head_dim + d);
acc += out_partials[p_idx] * beta;
}
}
acc = acc / s_l;
}
out[out_base + d] = static_cast<hip_bfloat16>(acc);
}
}
inline void check_hip_error(const char* where) {
hipError_t err = hipGetLastError();
TORCH_CHECK(err == hipSuccess, where, " failed: ", hipGetErrorString(err));
}
} // namespace
torch::Tensor mla_decode_forward(
torch::Tensor q,
torch::Tensor kv_buffer,
torch::Tensor kv_scale,
torch::Tensor qo_indptr,
torch::Tensor kv_indptr,
int64_t config_v_dim,
int64_t chunk_size,
double sm_scale) {
TORCH_CHECK(q.is_cuda(), "q must be CUDA/HIP tensor");
TORCH_CHECK(kv_buffer.is_cuda(), "kv_buffer must be CUDA/HIP tensor");
TORCH_CHECK(kv_scale.is_cuda(), "kv_scale must be CUDA/HIP tensor");
TORCH_CHECK(qo_indptr.is_cuda(), "qo_indptr must be CUDA/HIP tensor");
TORCH_CHECK(kv_indptr.is_cuda(), "kv_indptr must be CUDA/HIP tensor");
TORCH_CHECK(q.scalar_type() == torch::kBFloat16, "q must be bfloat16");
TORCH_CHECK(q.dim() == 3, "q must be [total_q, 16, 576]");
TORCH_CHECK(q.size(1) == kNumHeads, "q num_heads must be 16");
TORCH_CHECK(q.size(2) == kQkHeadDim, "q head dim must be 576");
TORCH_CHECK(kv_buffer.element_size() == 1, "kv_buffer must be 1-byte packed fp4x2");
TORCH_CHECK(kv_buffer.dim() == 3, "kv_buffer must be [total_kv, 1, 288]");
TORCH_CHECK(kv_buffer.size(1) == 1, "kv_buffer num_kv_heads must be 1");
TORCH_CHECK(kv_buffer.size(2) == kQkHeadDim / 2, "kv_buffer packed width must be 288");
TORCH_CHECK(kv_scale.element_size() == 1, "kv_scale must be 1-byte e8m0");
TORCH_CHECK(kv_scale.dim() == 2, "kv_scale must be [total_kv, 18]");
TORCH_CHECK(kv_scale.size(1) == kQkHeadDim / kScaleBlock, "kv_scale width must be 18");
TORCH_CHECK(qo_indptr.scalar_type() == torch::kInt32, "qo_indptr must be int32");
TORCH_CHECK(kv_indptr.scalar_type() == torch::kInt32, "kv_indptr must be int32");
TORCH_CHECK(qo_indptr.dim() == 1, "qo_indptr must be 1D");
TORCH_CHECK(kv_indptr.dim() == 1, "kv_indptr must be 1D");
TORCH_CHECK(qo_indptr.size(0) == kv_indptr.size(0), "indptr size mismatch");
TORCH_CHECK(config_v_dim == kVHeadDim, "config_v_dim must be 512");
TORCH_CHECK(chunk_size > 0, "chunk_size must be > 0");
q = q.contiguous();
kv_buffer = kv_buffer.contiguous();
kv_scale = kv_scale.contiguous();
qo_indptr = qo_indptr.contiguous();
kv_indptr = kv_indptr.contiguous();
int current_device = 0;
hipDeviceProp_t prop;
TORCH_CHECK(hipGetDevice(¤t_device) == hipSuccess, "hipGetDevice failed");
TORCH_CHECK(hipGetDeviceProperties(&prop, current_device) == hipSuccess,
"hipGetDeviceProperties failed");
TORCH_CHECK(prop.warpSize == 64, "This kernel requires wavefront size 64, got ", prop.warpSize);
const int64_t total_q = q.size(0);
const int64_t batch_size = qo_indptr.size(0) - 1;
auto out = torch::empty({total_q, kNumHeads, config_v_dim}, q.options().dtype(torch::kBFloat16));
if (total_q == 0 || batch_size <= 0) {
return out.zero_();
}
auto q_lens = qo_indptr.slice(0, 1, batch_size + 1) - qo_indptr.slice(0, 0, batch_size);
TORCH_CHECK(q_lens.min().item<int64_t>() == 1 && q_lens.max().item<int64_t>() == 1,
"decode-only kernel requires q_seq_len == 1 per batch segment");
auto kv_lens = kv_indptr.slice(0, 1, batch_size + 1) - kv_indptr.slice(0, 0, batch_size);
const int64_t max_kv_len = kv_lens.max().item<int64_t>();
const int64_t num_chunks = (max_kv_len + chunk_size - 1) / chunk_size;
if (num_chunks <= 0) {
return out.zero_();
}
auto out_partials = torch::zeros(
{batch_size, kNumHeads, num_chunks, config_v_dim},
q.options().dtype(torch::kFloat32));
auto lse_partials = torch::zeros(
{batch_size, kNumHeads, num_chunks, 2},
q.options().dtype(torch::kFloat32));
const int kv_row_stride = static_cast<int>(kv_buffer.stride(0));
const int kv_scale_row_stride = static_cast<int>(kv_scale.stride(0));
hipLaunchKernelGGL(
mla_decode_chunk_kernel,
dim3(static_cast<unsigned int>(num_chunks), static_cast<unsigned int>(batch_size), 1),
dim3(64, 16, 1),
0,
0,
reinterpret_cast<const hip_bfloat16*>(q.data_ptr<at::BFloat16>()),
reinterpret_cast<const uint8_t*>(kv_buffer.data_ptr()),
reinterpret_cast<const uint8_t*>(kv_scale.data_ptr()),
qo_indptr.data_ptr<int32_t>(),
kv_indptr.data_ptr<int32_t>(),
out_partials.data_ptr<float>(),
lse_partials.data_ptr<float>(),
kv_row_stride,
kv_scale_row_stride,
kQkHeadDim,
static_cast<int>(config_v_dim),
static_cast<float>(sm_scale),
static_cast<int>(chunk_size),
static_cast<int>(num_chunks));
check_hip_error("mla_decode_chunk_kernel");
hipLaunchKernelGGL(
mla_decode_reduce_kernel,
dim3(static_cast<unsigned int>(batch_size), static_cast<unsigned int>(kNumHeads), 1),
dim3(256, 1, 1),
0,
0,
out_partials.data_ptr<float>(),
lse_partials.data_ptr<float>(),
reinterpret_cast<hip_bfloat16*>(out.data_ptr<at::BFloat16>()),
qo_indptr.data_ptr<int32_t>(),
static_cast<int>(num_chunks),
static_cast<int>(config_v_dim));
check_hip_error("mla_decode_reduce_kernel");
return out;
}
"""
def _get_chunk_size() -> int:
raw = os.getenv(_CHUNK_ENV)
if raw is None:
return _DEFAULT_CHUNK_SIZE
try:
return max(1, int(raw))
except ValueError:
return _DEFAULT_CHUNK_SIZE
def _normalize_mxfp4_scale(
kv_scale: torch.Tensor,
total_kv: int,
qk_head_dim: int,
) -> torch.Tensor:
num_blocks = qk_head_dim // 32
need = total_kv * num_blocks
flat = kv_scale.contiguous().reshape(-1)
if flat.numel() < need:
raise RuntimeError(
f"kv_scale has {flat.numel()} elements, need at least {need} for [{total_kv}, {num_blocks}]"
)
return flat[:need].view(total_kv, num_blocks).contiguous()
def _get_hip_module():
global _HIP_MODULE, _HIP_BUILD_ERROR
if _HIP_MODULE is not None:
return _HIP_MODULE
if _HIP_BUILD_ERROR is not None:
raise RuntimeError(f"HIP inline build failed previously: {_HIP_BUILD_ERROR}")
with _HIP_LOCK:
if _HIP_MODULE is not None:
return _HIP_MODULE
if _HIP_BUILD_ERROR is not None:
raise RuntimeError(f"HIP inline build failed previously: {_HIP_BUILD_ERROR}")
try:
os.environ.setdefault("CXX", "clang++")
_HIP_MODULE = load_inline(
name="mla_mxfp4_splitk_hip_v1",
cpp_sources=[CPP_WRAPPER],
cuda_sources=[HIP_SRC],
functions=["mla_decode_forward"],
verbose=False,
extra_cuda_cflags=["-O3", "-std=c++20", "-DUSE_ROCM"],
)
except Exception as e: # pragma: no cover
_HIP_BUILD_ERROR = e
raise RuntimeError(f"HIP inline build failed: {e}") from e
return _HIP_MODULE
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
kv_buffer_mxfp4, kv_scale_mxfp4 = kv_data["mxfp4"]
q = q.contiguous()
qo_indptr = qo_indptr.to(torch.int32).contiguous()
kv_indptr = kv_indptr.to(torch.int32).contiguous()
q_lens = qo_indptr[1:] - qo_indptr[:-1]
if int(q_lens.min().item()) != 1 or int(q_lens.max().item()) != 1:
raise RuntimeError("This decode kernel supports q_seq_len == 1 only")
total_kv = int(kv_buffer_mxfp4.size(0))
qk_head_dim = int(config["qk_head_dim"])
v_head_dim = int(config["v_head_dim"])
kv_buffer_mxfp4 = kv_buffer_mxfp4.contiguous()
kv_scale_mxfp4 = _normalize_mxfp4_scale(
kv_scale_mxfp4,
total_kv=total_kv,
qk_head_dim=qk_head_dim,
)
module = _get_hip_module()
return module.mla_decode_forward(
q,
kv_buffer_mxfp4,
kv_scale_mxfp4,
qo_indptr,
kv_indptr,
v_head_dim,
_get_chunk_size(),
float(config["sm_scale"]),
)
scrolls · 486 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