submission 608791
XiaomingFun233 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 736 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-608791?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:c57282b48ff0b6f9e8cd55e75bfcd9e1567b2e1a8881898b1f61b290806b38c6
license declaredunknown
license concludedunknown
authorsXiaomingFun233
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__ __align__(16) float smem_kv[kSplitTileKv][kQkHeadDim];split-k
name="mla_mxfp4_splitk_hip_v5_metadata_graph",vector-width = float4
float4* smem_ptr = reinterpret_cast<float4*>(&smem_kv[token_local][dim_base]);Kernel source
submission_amd_mixed_mla_mxfp4_hip_splitk.py736 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
import math
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
_METADATA_CACHE: dict[tuple[object, ...], tuple[torch.Tensor, torch.Tensor, torch.Tensor]] = {}
_BUFFER_CACHE: dict[tuple[int, int, int], tuple[torch.Tensor, torch.Tensor, torch.Tensor]] = {}
_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 work_info,
torch::Tensor reduce_indptr,
torch::Tensor reduce_partial_map,
torch::Tensor partial_output,
torch::Tensor partial_lse,
int64_t config_v_dim,
double sm_scale);
"""
# The inline implementation below is reorganized to mirror the original MLA decode
# graph we documented in `e2e/cmd/mla_decode_fwd_csrc_graph.md`:
# 1. metadata/work-item construction on the Python side
# 2. asm_mla.cu-like stage1 partial accumulation
# 3. reduce.cu-like partial reduction into the final output
# The HK path from hk_decode_fwd.cu is intentionally not recreated here because
# this benchmark always uses 16 query heads with MXFP4 KV, while the HK kernel is
# an experimental fp8 h128 path.
HIP_SRC_COMMON = r"""
#include <torch/extension.h>
#include <hip/hip_bfloat16.h>
#include <hip/hip_runtime.h>
#include <algorithm>
#include <cmath>
#include <cstdint>
namespace {
constexpr int kNumHeads = 16;
constexpr int kQkHeadDim = 576;
constexpr int kVHeadDim = 512;
constexpr int kScaleBlock = 32;
constexpr int kSplitTileKv = 16;
constexpr int kPackedBytesPerToken = kQkHeadDim / 2;
constexpr int kVecBytes = 4;
constexpr int kVecPerToken = kPackedBytesPerToken / kVecBytes;
constexpr int kWavefrontSize = 64;
constexpr int kSplitWarpsPerBlock = kNumHeads;
constexpr int kSplitThreadsPerBlock = kWavefrontSize * kSplitWarpsPerBlock;
struct MlaShape {
int64_t total_q;
int64_t batch_size;
int64_t num_work;
};
__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 packed, float scale, bool high_nibble) {
const uint8_t nibble =
high_nibble ? static_cast<uint8_t>(packed >> 4) : static_cast<uint8_t>(packed & 0x0F);
return kFp4Lut[nibble & 0x0F] * scale;
}
__device__ __forceinline__ float wavefront_sum_64(float value) {
for (int offset = 32; offset > 0; offset >>= 1) {
value += __shfl_down(value, offset);
}
return value;
}
inline void check_hip_error(const char* where) {
const hipError_t err = hipGetLastError();
TORCH_CHECK(err == hipSuccess, where, " failed: ", hipGetErrorString(err));
}
inline MlaShape check_mla_inputs(
torch::Tensor q,
torch::Tensor kv_buffer,
torch::Tensor kv_scale,
torch::Tensor qo_indptr,
torch::Tensor work_info,
torch::Tensor reduce_indptr,
torch::Tensor reduce_partial_map,
int64_t config_v_dim) {
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(work_info.is_cuda(), "work_info must be CUDA/HIP tensor");
TORCH_CHECK(reduce_indptr.is_cuda(), "reduce_indptr must be CUDA/HIP tensor");
TORCH_CHECK(reduce_partial_map.is_cuda(), "reduce_partial_map 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 packed fp4x2 bytes");
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 e8m0 bytes");
TORCH_CHECK(kv_scale.dim() == 2, "kv_scale must be [total_kv, 18]");
TORCH_CHECK(kv_scale.size(0) == kv_buffer.size(0), "kv_scale row count mismatch");
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(work_info.scalar_type() == torch::kInt32, "work_info must be int32");
TORCH_CHECK(reduce_indptr.scalar_type() == torch::kInt32, "reduce_indptr must be int32");
TORCH_CHECK(reduce_partial_map.scalar_type() == torch::kInt32, "reduce_partial_map must be int32");
TORCH_CHECK(config_v_dim == kVHeadDim, "config_v_dim must be 512");
TORCH_CHECK(q.is_contiguous(), "q must be contiguous");
TORCH_CHECK(kv_buffer.is_contiguous(), "kv_buffer must be contiguous");
TORCH_CHECK(kv_scale.is_contiguous(), "kv_scale must be contiguous");
TORCH_CHECK(qo_indptr.is_contiguous(), "qo_indptr must be contiguous");
TORCH_CHECK(work_info.is_contiguous(), "work_info must be contiguous");
TORCH_CHECK(reduce_indptr.is_contiguous(), "reduce_indptr must be contiguous");
TORCH_CHECK(reduce_partial_map.is_contiguous(), "reduce_partial_map must be contiguous");
MlaShape shape{};
shape.total_q = q.size(0);
shape.batch_size = qo_indptr.size(0) - 1;
TORCH_CHECK(shape.batch_size >= 0, "batch_size must be non-negative");
shape.num_work = work_info.size(0);
TORCH_CHECK(
reduce_indptr.size(0) == shape.batch_size + 1,
"reduce_indptr size mismatch");
TORCH_CHECK(
reduce_partial_map.size(0) == shape.num_work,
"reduce_partial_map size mismatch");
return shape;
}
} // namespace
"""
HIP_SRC_STAGE1 = r"""
namespace {
__global__ void mla_stage1_split_mxfp4_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,
const int32_t* __restrict__ work_info,
float* __restrict__ partial_output,
float* __restrict__ partial_lse,
int kv_row_stride,
int kv_scale_row_stride,
int v_head_dim,
float sm_scale,
int split_size,
int uniform_num_splits,
int work_info_stride) {
const int work_idx = static_cast<int>(blockIdx.x);
const int lane = static_cast<int>(threadIdx.x);
const int head = static_cast<int>(threadIdx.y);
const int linear_tid = head * kWavefrontSize + lane;
if (head >= kNumHeads) {
return;
}
int32_t partial_idx = work_idx;
int32_t q_idx = 0;
int32_t token_begin = 0;
int32_t token_end = 0;
if (uniform_num_splits > 0) {
const int32_t batch_idx = work_idx / uniform_num_splits;
const int32_t split_idx = work_idx - batch_idx * uniform_num_splits;
q_idx = qo_indptr[batch_idx];
const int32_t kv_start = kv_indptr[batch_idx];
const int32_t kv_end = kv_indptr[batch_idx + 1];
token_begin = kv_start + split_idx * split_size;
token_end = token_begin + split_size < kv_end ? token_begin + split_size : kv_end;
} else {
const int32_t* work_row = work_info + static_cast<int64_t>(work_idx) * work_info_stride;
partial_idx = work_row[1];
q_idx = work_row[2];
token_begin = work_row[4];
token_end = work_row[5];
}
const hip_bfloat16* q_row = q + (static_cast<int64_t>(q_idx) * kNumHeads + head) * kQkHeadDim;
float q_frag[(kQkHeadDim + kWavefrontSize - 1) / kWavefrontSize];
#pragma unroll
for (int i = 0; i < (kQkHeadDim + kWavefrontSize - 1) / kWavefrontSize; ++i) {
const int d = lane + i * kWavefrontSize;
q_frag[i] = (d < kQkHeadDim) ? static_cast<float>(q_row[d]) : 0.0f;
}
float out_frag[kVHeadDim / kWavefrontSize];
#pragma unroll
for (int i = 0; i < kVHeadDim / kWavefrontSize; ++i) {
out_frag[i] = 0.0f;
}
float max_score = -INFINITY;
float score_sum = 0.0f;
__shared__ __align__(16) float smem_kv[kSplitTileKv][kQkHeadDim];
for (int32_t tile_begin = token_begin; tile_begin < token_end; tile_begin += kSplitTileKv) {
const int32_t tile_tokens =
std::min(kSplitTileKv, static_cast<int>(token_end - tile_begin));
const int total_vec_tasks = tile_tokens * kVecPerToken;
for (int task = linear_tid; task < total_vec_tasks; task += kSplitThreadsPerBlock) {
const int token_local = task / kVecPerToken;
const int vec_idx = task % kVecPerToken;
const int32_t token = tile_begin + token_local;
const int scale_idx = vec_idx / 4;
const uint8_t* scale_row = kv_scale + static_cast<int64_t>(token) * kv_scale_row_stride;
const float scale = e8m0_to_float(scale_row[scale_idx]);
const uint8_t* kv_row = kv_data + static_cast<int64_t>(token) * kv_row_stride;
const uint32_t packed4 = reinterpret_cast<const uint32_t*>(kv_row)[vec_idx];
const uint8_t* packed_bytes = reinterpret_cast<const uint8_t*>(&packed4);
const int dim_base = vec_idx * 8;
float4* smem_ptr = reinterpret_cast<float4*>(&smem_kv[token_local][dim_base]);
float4 low4;
float4 high4;
low4.x = unpack_mxfp4(packed_bytes[0], scale, false);
low4.y = unpack_mxfp4(packed_bytes[0], scale, true);
low4.z = unpack_mxfp4(packed_bytes[1], scale, false);
low4.w = unpack_mxfp4(packed_bytes[1], scale, true);
high4.x = unpack_mxfp4(packed_bytes[2], scale, false);
high4.y = unpack_mxfp4(packed_bytes[2], scale, true);
high4.z = unpack_mxfp4(packed_bytes[3], scale, false);
high4.w = unpack_mxfp4(packed_bytes[3], scale, true);
smem_ptr[0] = low4;
smem_ptr[1] = high4;
}
__syncthreads();
#pragma unroll 1
for (int token_local = 0; token_local < tile_tokens; ++token_local) {
float qk_local = 0.0f;
#pragma unroll
for (int i = 0; i < (kQkHeadDim + kWavefrontSize - 1) / kWavefrontSize; ++i) {
const int d = lane + i * kWavefrontSize;
if (d < kQkHeadDim) {
qk_local += q_frag[i] * smem_kv[token_local][d];
}
}
float score = wavefront_sum_64(qk_local);
score = (lane == 0) ? score * sm_scale : 0.0f;
score = __shfl(score, 0);
const float new_max = fmaxf(max_score, score);
const float alpha = __expf(max_score - new_max);
const float prob = __expf(score - new_max);
const float new_sum = score_sum * alpha + prob;
#pragma unroll
for (int i = 0; i < kVHeadDim / kWavefrontSize; ++i) {
const int d = lane + i * kWavefrontSize;
out_frag[i] = out_frag[i] * alpha + prob * smem_kv[token_local][d];
}
max_score = new_max;
score_sum = new_sum;
}
__syncthreads();
}
const int64_t split_base =
((static_cast<int64_t>(partial_idx) * kNumHeads + head) * v_head_dim);
#pragma unroll
for (int i = 0; i < kVHeadDim / kWavefrontSize; ++i) {
const int d = lane + i * kWavefrontSize;
partial_output[split_base + d] = out_frag[i];
}
if (lane == 0) {
const int64_t lse_base =
((static_cast<int64_t>(partial_idx) * kNumHeads + head) * 2);
partial_lse[lse_base] = max_score;
partial_lse[lse_base + 1] = score_sum;
}
}
void mla_decode_stage1_fp32_split(
torch::Tensor q,
torch::Tensor kv_buffer,
torch::Tensor kv_scale,
torch::Tensor qo_indptr,
torch::Tensor work_info,
torch::Tensor partial_output,
torch::Tensor partial_lse,
double sm_scale,
const MlaShape& shape) {
if (shape.num_work <= 0) {
return;
}
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));
const int work_info_stride = static_cast<int>(work_info.stride(0));
hipLaunchKernelGGL(
mla_stage1_split_mxfp4_kernel,
dim3(static_cast<unsigned int>(shape.num_work), 1, 1),
dim3(kWavefrontSize, kNumHeads, 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>(),
qo_indptr.data_ptr<int32_t>(),
work_info.data_ptr<int32_t>(),
partial_output.data_ptr<float>(),
partial_lse.data_ptr<float>(),
kv_row_stride,
kv_scale_row_stride,
kVHeadDim,
static_cast<float>(sm_scale),
0,
0,
work_info_stride);
check_hip_error("mla_stage1_split_mxfp4_kernel");
}
} // namespace
"""
HIP_SRC_REDUCE = r"""
namespace {
__global__ void mla_reduce_v1_like_kernel(
const float* __restrict__ partial_output,
const float* __restrict__ partial_lse,
hip_bfloat16* __restrict__ out,
const int32_t* __restrict__ qo_indptr,
const int32_t* __restrict__ reduce_indptr,
const int32_t* __restrict__ reduce_partial_map,
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);
const int num_threads = static_cast<int>(blockDim.x);
const int32_t reduce_begin = reduce_indptr[batch_idx];
const int32_t reduce_end = reduce_indptr[batch_idx + 1];
const int num_splits = reduce_end - reduce_begin;
// Shared memory: [0..num_threads-1] for warp-reduce temporaries
extern __shared__ float smem_reduce[];
float* smem_max = smem_reduce; // [num_threads] for max reduction
float* smem_sum = smem_reduce + num_threads; // [num_threads] for sum reduction
// --- Phase 1: parallel max over splits ---
float local_max = -INFINITY;
for (int i = tid; i < num_splits; i += num_threads) {
const int32_t partial_idx = reduce_partial_map[reduce_begin + i];
const int64_t lse_base = (static_cast<int64_t>(partial_idx) * kNumHeads + head) * 2;
const float split_sum = partial_lse[lse_base + 1];
if (split_sum > 0.0f) {
local_max = fmaxf(local_max, partial_lse[lse_base]);
}
}
smem_max[tid] = local_max;
__syncthreads();
for (int stride = num_threads >> 1; stride > 0; stride >>= 1) {
if (tid < stride) {
smem_max[tid] = fmaxf(smem_max[tid], smem_max[tid + stride]);
}
__syncthreads();
}
const float global_max = smem_max[0];
if (!isfinite(global_max)) {
// All splits empty — write zeros
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 += num_threads) {
out[out_base + d] = static_cast<hip_bfloat16>(0.0f);
}
return;
}
// --- Phase 2: parallel sum of rescaled weights ---
float local_sum = 0.0f;
for (int i = tid; i < num_splits; i += num_threads) {
const int32_t partial_idx = reduce_partial_map[reduce_begin + i];
const int64_t lse_base = (static_cast<int64_t>(partial_idx) * kNumHeads + head) * 2;
const float split_sum = partial_lse[lse_base + 1];
if (split_sum > 0.0f) {
local_sum += __expf(partial_lse[lse_base] - global_max) * split_sum;
}
}
smem_sum[tid] = local_sum;
__syncthreads();
for (int stride = num_threads >> 1; stride > 0; stride >>= 1) {
if (tid < stride) {
smem_sum[tid] += smem_sum[tid + stride];
}
__syncthreads();
}
const float global_sum = smem_sum[0];
// --- Phase 3: parallel V accumulation ---
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 += num_threads) {
float acc = 0.0f;
for (int i = 0; i < num_splits; i++) {
const int32_t partial_idx = reduce_partial_map[reduce_begin + i];
const int64_t lse_base = (static_cast<int64_t>(partial_idx) * kNumHeads + head) * 2;
const float split_sum_w = partial_lse[lse_base + 1];
if (split_sum_w > 0.0f) {
const float rescale = __expf(partial_lse[lse_base] - global_max);
const int64_t split_base =
(static_cast<int64_t>(partial_idx) * kNumHeads + head) * v_head_dim + d;
acc += partial_output[split_base] * rescale;
}
}
out[out_base + d] = static_cast<hip_bfloat16>(acc / global_sum);
}
}
void mla_reduce_v1_like(
torch::Tensor partial_output,
torch::Tensor partial_lse,
torch::Tensor qo_indptr,
torch::Tensor reduce_indptr,
torch::Tensor reduce_partial_map,
torch::Tensor out,
const MlaShape& shape) {
if (shape.batch_size <= 0) {
return;
}
constexpr int kReduceThreads = 256;
const int smem_bytes = 2 * kReduceThreads * sizeof(float);
hipLaunchKernelGGL(
mla_reduce_v1_like_kernel,
dim3(static_cast<unsigned int>(shape.batch_size), static_cast<unsigned int>(kNumHeads), 1),
dim3(kReduceThreads, 1, 1),
smem_bytes,
0,
partial_output.data_ptr<float>(),
partial_lse.data_ptr<float>(),
reinterpret_cast<hip_bfloat16*>(out.data_ptr<at::BFloat16>()),
qo_indptr.data_ptr<int32_t>(),
reduce_indptr.data_ptr<int32_t>(),
reduce_partial_map.data_ptr<int32_t>(),
kVHeadDim);
check_hip_error("mla_reduce_v1_like_kernel");
}
} // namespace
"""
HIP_SRC_ENTRY = r"""
torch::Tensor mla_decode_forward(
torch::Tensor q,
torch::Tensor kv_buffer,
torch::Tensor kv_scale,
torch::Tensor qo_indptr,
torch::Tensor work_info,
torch::Tensor reduce_indptr,
torch::Tensor reduce_partial_map,
torch::Tensor partial_output,
torch::Tensor partial_lse,
int64_t config_v_dim,
double sm_scale) {
const MlaShape shape =
check_mla_inputs(
q,
kv_buffer,
kv_scale,
qo_indptr,
work_info,
reduce_indptr,
reduce_partial_map,
config_v_dim);
auto out = torch::empty({shape.total_q, kNumHeads, config_v_dim}, q.options().dtype(torch::kBFloat16));
if (shape.total_q == 0 || shape.batch_size <= 0) {
return out.zero_();
}
mla_decode_stage1_fp32_split(
q,
kv_buffer,
kv_scale,
qo_indptr,
work_info,
partial_output,
partial_lse,
sm_scale,
shape);
mla_reduce_v1_like(partial_output, partial_lse, qo_indptr, reduce_indptr, reduce_partial_map, out, shape);
return out;
}
"""
HIP_SRC = "\n".join(
[
HIP_SRC_COMMON,
HIP_SRC_STAGE1,
HIP_SRC_REDUCE,
HIP_SRC_ENTRY,
]
)
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 _build_decode_metadata(
qo_indptr: torch.Tensor,
kv_indptr: torch.Tensor,
split_size: int,
device: torch.device,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
qo_host = qo_indptr.detach().cpu().tolist()
kv_host = kv_indptr.detach().cpu().tolist()
cache_key = (
str(device),
split_size,
tuple(qo_host),
tuple(kv_host),
)
cached = _METADATA_CACHE.get(cache_key)
if cached is not None:
return cached
batch_size = len(qo_host) - 1
work_rows: list[list[int]] = []
reduce_partial_map: list[int] = []
reduce_indptr = [0]
for batch_idx in range(batch_size):
q_start = int(qo_host[batch_idx])
q_end = int(qo_host[batch_idx + 1])
kv_start = int(kv_host[batch_idx])
kv_end = int(kv_host[batch_idx + 1])
kv_len = kv_end - kv_start
num_splits = math.ceil(kv_len / split_size) if kv_len > 0 else 0
for split_idx in range(num_splits):
split_begin = kv_start + split_idx * split_size
split_end = min(split_begin + split_size, kv_end)
partial_idx = len(work_rows)
work_rows.append(
[batch_idx, partial_idx, q_start, q_end, split_begin, split_end, 0, 0]
)
reduce_partial_map.append(partial_idx)
reduce_indptr.append(len(reduce_partial_map))
if work_rows:
work_info = torch.tensor(work_rows, dtype=torch.int32, device=device)
else:
work_info = torch.empty((0, 8), dtype=torch.int32, device=device)
metadata = (
work_info.contiguous(),
torch.tensor(reduce_indptr, dtype=torch.int32, device=device).contiguous(),
torch.tensor(reduce_partial_map, dtype=torch.int32, device=device).contiguous(),
)
_METADATA_CACHE[cache_key] = metadata
return metadata
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_v5_metadata_graph",
cpp_sources=[CPP_WRAPPER],
cuda_sources=[HIP_SRC],
functions=["mla_decode_forward"],
verbose=False,
extra_cuda_cflags=["-O3", "-ffast-math", "-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"])
batch_size = int(qo_indptr.size(0)) - 1
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,
)
split_size = _get_chunk_size()
work_info, reduce_indptr, reduce_partial_map = _build_decode_metadata(
qo_indptr,
kv_indptr,
split_size=split_size,
device=q.device,
)
num_work = int(work_info.size(0))
num_heads = 16
buf_key = (num_work, num_heads, v_head_dim)
if buf_key not in _BUFFER_CACHE:
_BUFFER_CACHE[buf_key] = (
torch.empty((num_work, num_heads, v_head_dim), dtype=torch.float32, device=q.device),
torch.empty((num_work, num_heads, 2), dtype=torch.float32, device=q.device),
)
partial_output, partial_lse = _BUFFER_CACHE[buf_key]
module = _get_hip_module()
return module.mla_decode_forward(
q,
kv_buffer_mxfp4,
kv_scale_mxfp4,
qo_indptr,
work_info,
reduce_indptr,
reduce_partial_map,
partial_output,
partial_lse,
v_head_dim,
float(config["sm_scale"]),
)
scrolls · 736 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