submission 629960
ưhat's up · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 634 lines, June 9 Researcher Reciprocity License v1.0.
v0.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-629960?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:dd9a3db6a2b6f47fc7e665498f5ab627f2250fa150d02e5b48e398e0f567f8e7
license declaredunknown
license concludedunknown
authorsưhat's up
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
kv_data["mxfp4"] — (Tensor, Tensor): kv_buffer fp4x2 (total_kv,1,288) + fp8_e8m0 scalepersistent-kernel
The reference uses aiter's a8w8 persistent MLA kernel (fp8 Q + fp8 KV),shared-memory
__shared__ float shared_qk_max[N_WARPS][16];Kernel source
v0.py634 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
MLA (Multi-head Latent Attention) decode kernel — submission template.
Implement custom_kernel() to beat the aiter a8w8 reference (fp8 Q + fp8 KV).
DeepSeek R1 forward_absorb MLA config:
num_heads = 16 (query heads, after TP split)
num_kv_heads = 1 (shared latent KV head)
kv_lora_rank = 512 (latent dim)
qk_rope_head_dim = 64 (RoPE dim)
qk_head_dim = 576 (kv_lora_rank + qk_rope_head_dim, absorbed q/k dim)
v_head_dim = 512 (= kv_lora_rank, output dim)
sm_scale = 1/sqrt(576)
KV buffer format (forward_absorb):
- Full 576 dims used as keys (for Q@K^T score computation)
- First 512 dims (kv_lora_rank) used as values (for output computation)
Input tuple:
q: (total_q, 16, 576) bfloat16 — absorbed query
kv_data: dict with three KV cache formats:
kv_data["bf16"] — Tensor (total_kv, 1, 576) bfloat16
kv_data["fp8"] — (Tensor, Tensor): kv_buffer fp8 (total_kv,1,576) + scalar scale
kv_data["mxfp4"] — (Tensor, Tensor): kv_buffer fp4x2 (total_kv,1,288) + fp8_e8m0 scale
qo_indptr: (batch_size + 1,) int32 — query segment pointers
kv_indptr: (batch_size + 1,) int32 — KV segment pointers
config: dict with MLA parameters
Output:
attention output: (total_q, 16, 512) bfloat16
The reference uses aiter's a8w8 persistent MLA kernel (fp8 Q + fp8 KV),
which is ~2-3x faster than bf16. To beat it, consider:
1. Use mxfp4 KV cache for even lower memory bandwidth
- Fuse dequantization with attention to avoid bf16 materialization
2. Custom kernel with tighter memory access patterns
3. MQA: 1 KV head shared across 16 query heads — minimize redundant memory loads
4. Variable-length batching: indptr-based segmented attention
5. Split K/V from buffer: full 576 dims for keys, first 512 dims for values
"""
from task import input_t, output_t
import torch
from torch.utils.cpp_extension import load_inline
import os
if "PYTORCH_ROCM_ARCH" not in os.environ:
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950:xnack-"
# ---------------------------------------------------------------------------
# MXFP4 quantization (aiter native: block-32, fp4x2 + fp8_e8m0 dtypes)
# Uses aiter.utility.fp4_utils.dynamic_mxfp4_quant
# ---------------------------------------------------------------------------
def quantize_mxfp4(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
"""
MXFP4 block-wise quantization using aiter's dynamic_mxfp4_quant.
Block size = 32. Each block gets an E8M0 scale factor.
Two FP4 E2M1 values are packed per byte.
Args:
tensor: bf16 tensor of shape [B, M, N] (N must be divisible by 32)
Returns:
(fp4_data, scale_e8m0)
- fp4_data: shape [B, M, N//2] in aiter_dtypes.fp4x2
- scale_e8m0: shape [B*M, ceil(N/32)] padded, in aiter_dtypes.fp8_e8m0
"""
from aiter.ops.triton.quant import dynamic_mxfp4_quant
orig_shape = tensor.shape # (B, M, N)
B, M, N = orig_shape
# dynamic_mxfp4_quant expects 2D: (B*M, N)
tensor_2d = tensor.reshape(B * M, N)
fp4_data_2d, scale_e8m0 = dynamic_mxfp4_quant(tensor_2d)
# Reshape fp4_data back to 3D: (B, M, N//2)
fp4_data = fp4_data_2d.view(B, M, N // 2)
return fp4_data, scale_e8m0
kernel_cpp = r"""
#undef __HIP_NO_HALF_OPERATORS__
#undef __HIP_NO_HALF_CONVERSIONS__
#include <hip/hip_bf16.h>
#include <hip/hip_ext_ocp.h>
#include <hip/hip_fp16.h>
#include <hip/hip_fp4.h>
#include <hip/hip_fp8.h>
#include <hip/hip_runtime.h>
#include <pybind11/pybind11.h>
#include <math.h>
using fp4x2_t = __amd_fp4x2_storage_t;
using fp4x16_t = fp4x2_t __attribute__((ext_vector_type(8)));
using fp4x32_t = fp4x2_t __attribute__((ext_vector_type(16)));
using fp4x64_t = fp4x2_t __attribute__((ext_vector_type(32)));
using uint8x2_t = uint8_t __attribute__((ext_vector_type(2)));
using uint8x16_t = uint8_t __attribute__((ext_vector_type(16)));
using uint8x32_t = uint8_t __attribute__((ext_vector_type(32)));
using bfp16 = __hip_bfloat16;
using bfp16x2_t = ushort __attribute__((ext_vector_type(2)));
using bfp16x4_t = ushort __attribute__((ext_vector_type(4)));
using bfp16x8_t = ushort __attribute__((ext_vector_type(8)));
using floatx4_t = float __attribute__((ext_vector_type(4)));
using floatx8_t = float __attribute__((ext_vector_type(8)));
using floatx16_t = float __attribute__((ext_vector_type(16)));
using floatx32_t = float __attribute__((ext_vector_type(32)));
#define WARP_SIZE 64
#define CEIL_DIV(x, y) (((x) + ((y) - 1)) / (y))
__device__ inline ushort float_2_bfloatraw(float f) {
#if HIP_BF16_AVX512_OP
union {
__bf16 bf16;
unsigned short us;
} u = {_mm_cvtness_sbh(f)};
return u.us;
#else
union {
float fp32;
unsigned int u32;
} u = {f};
if (~u.u32 & 0x7f800000) {
u.u32 += 0x7fff + ((u.u32 >> 16) & 1); // Round to nearest, round to even
} else if (u.u32 & 0xffff) {
u.u32 |= 0x10000; // Preserve signaling NaN
}
return static_cast<unsigned short>(u.u32 >> 16);
#endif
}
__device__ inline fp4x64_t fp4x32_to_fp4x64(const fp4x32_t &p) {
fp4x64_t v;
reinterpret_cast<fp4x32_t *>(&v)[0] = p;
reinterpret_cast<fp4x32_t *>(&v)[1] = 0;
return v;
}
__device__ inline uint8x32_t fp4x32_to_uint8x32(const fp4x32_t &p) {
uint8x32_t v;
#pragma unroll
for (int i = 0; i < 16; ++i) {
uint8x2_t tmp;
tmp[0] = p[i] & 0x0f;
tmp[1] = (p[i] & 0xf0) >> 4;
reinterpret_cast<uint8x2_t *>(&v)[i] = tmp;
}
return v;
}
__device__ inline fp4x32_t uint8x32_to_fp4x32(const uint8x32_t &p) {
fp4x32_t v;
#pragma unroll
for (int i = 0; i < 16; ++i) {
v[i] = (p[i << 1 | 1] << 4) | (p[i << 1]);
}
return v;
}
__device__ __forceinline__ bfp16x2_t fp4x2_to_bfp16x2(uint8_t packed) {
ushort em0 = packed & 0x07; // Low nibble E+M
ushort em1 = packed & 0x70; // High nibble E+M
// BFloat16 Constants
// Bias difference (127 - 1) = 126. Shifted to BF16 exp position: 126 << 7
constexpr ushort bfp16_bias = 0x3F00;
constexpr ushort bfp16_0p5 = 0x3F00; // 0.5 in Float32 hex
// 2. Direct alignment to BFloat16 positions
// Shift E+M left by 6. This perfectly drops the Mantissa into bit 6,
// and the Exponent into bits 7 and 8.
// Sign bits (Bit 3 and Bit 7) are shifted to the 15st bit.
ushort x0 = (em0 << 6) | ((packed & 0x08) << 12);
ushort x1 = (em1 << 2) | ((packed & 0x80) << 8);
// 3. Apply Bias via Integer Addition (Triton's trick)
// CMOV (Conditional Move): Only add bias if exponent is not zero
x0 = ((em0 & 0x06) != 0) ? (x0 + bfp16_bias) : x0;
x1 = ((em1 & 0x60) != 0) ? (x1 + bfp16_bias) : x1;
// 4. Handle Subnormals
// CMOV: Force to +/- 0.5 if it's a subnormal (em == 0x01)
x0 = (em0 == 0x01) ? (bfp16_0p5 | (x0 & 0x8000)) : x0;
x1 = (em1 == 0x10) ? (bfp16_0p5 | (x1 & 0x8000)) : x1;
return {x0, x1};
}
__device__ inline floatx4_t mfma_16x16x32(bfp16x8_t a, bfp16x8_t b,
floatx4_t c) {
return __builtin_amdgcn_mfma_f32_16x16x32_bf16(a, b, c, 0, 0, 0);
}
__device__ inline floatx4_t mfma_16x16x128(fp4x64_t a, fp4x64_t b, floatx4_t c,
uint8_t scale_a, uint8_t scale_b) {
return __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a, b, c, 4, 4, 0,
scale_a, 0, scale_b);
}
__device__ inline floatx16_t mfma_32x32x64(fp4x64_t a, fp4x64_t b, floatx16_t c,
uint8_t scale_a, uint8_t scale_b) {
return __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(a, b, c, 4, 4, 0,
scale_a, 0, scale_b);
}
__global__ void mla_decode(
const fp4x2_t *q, // [total_q, num_heads, (kv_lora + rope_dim) / 2]
const uint8_t *q_scale, // [total_q, num_heads, (kv_lora + rope_dim) / 32]
const fp4x2_t *kv_cache, // [total_kv, 1, (kv_lora + rope_dim) / 2]
const uint8_t *kv_scale, // [total_kv, 1, (kv_lora + rope_dim) / 32]
const int *qo_indptr, // [num_seqs + 1]
const int *kv_indptr, // [num_seqs + 1]
bfp16 *output, // [total_q, num_heads, num_partition, kv_lora]
float *lse, // [total_q, num_heads, num_partition]
const int num_seqs,
const int max_q_seq_len, // = 1
const int max_kv_seq_len, const float scale, const int max_num_partitions) {
// template arguments
constexpr int N_WARPS = 4;
constexpr int PARTITION_SIZE = 256;
constexpr int QK_HEAD_DIM = 576;
constexpr int V_HEAD_DIM = 512;
constexpr int N_HEADS = 16;
//
constexpr int QKHE_PER_LOOP = 128;
constexpr int QKHE_LOOP = CEIL_DIV(QK_HEAD_DIM, QKHE_PER_LOOP);
constexpr int VHE_PER_LOOP = 16;
constexpr int VHE_LOOP = CEIL_DIV(V_HEAD_DIM, 16);
constexpr int TOKENS_PER_WARP = PARTITION_SIZE / N_WARPS;
constexpr int TOKENS_PER_LOOP = 16;
constexpr int TOKEN_LOOP = CEIL_DIV(TOKENS_PER_WARP, 16);
constexpr int VT_PER_LOOP = 32;
constexpr int VT_LOOP = TOKENS_PER_WARP / 32;
const int seq_id = blockIdx.y;
const int query_idx = qo_indptr[seq_id];
const int part_id = blockIdx.x;
const int kv_idx = kv_indptr[seq_id];
const int ctx_len = kv_indptr[seq_id + 1] - kv_idx;
const int start_token_idx = part_id * PARTITION_SIZE;
const int num_partitions = CEIL_DIV(ctx_len, PARTITION_SIZE);
if (part_id >= num_partitions)
return;
const int end_token_idx = min(start_token_idx + PARTITION_SIZE, ctx_len);
const int tid = threadIdx.x;
const int wid = tid / WARP_SIZE;
const int lane_id = tid % WARP_SIZE;
const int lane_id16 = lane_id % 16;
const int lane4_id = lane_id / 16;
// load q: vram -> reg
const fp4x2_t *q_ptr =
q + (query_idx * N_HEADS + lane_id16) * (QK_HEAD_DIM / 2);
fp4x64_t q_local[QKHE_LOOP];
for (int i = 0; i < QKHE_LOOP; ++i) {
const int head_elem = i * QKHE_PER_LOOP + lane4_id * 32;
fp4x32_t tmp;
if (head_elem < QK_HEAD_DIM) {
tmp = *reinterpret_cast<const fp4x32_t *>(q_ptr + head_elem / 2);
} else {
tmp = 0;
}
q_local[i] = fp4x32_to_fp4x64(tmp);
}
// load q_scale : vram->reg
uint8_t q_scale_local[QKHE_LOOP];
const uint8_t *q_scale_ptr =
q_scale + (query_idx * N_HEADS + lane_id16) * (QK_HEAD_DIM / 32);
for (int i = 0; i < QKHE_LOOP; ++i) {
const int head_elem = i * QKHE_PER_LOOP + lane4_id * 32;
if (head_elem < QK_HEAD_DIM) {
q_scale_local[i] = q_scale_ptr[head_elem / 32];
} else {
q_scale_local[i] = 0;
}
}
// load k: vram -> reg
fp4x64_t k_local[TOKEN_LOOP][QKHE_LOOP];
for (int i = 0; i < TOKEN_LOOP; ++i) {
const int token_idx = start_token_idx + wid * TOKENS_PER_WARP +
i * TOKENS_PER_LOOP + lane_id16;
const fp4x2_t *k_ptr = kv_cache + (kv_idx + token_idx) * (QK_HEAD_DIM / 2);
for (int j = 0; j < QKHE_LOOP; ++j) {
const int head_elem = j * QKHE_PER_LOOP + lane4_id * 32;
fp4x32_t tmp;
if (token_idx < end_token_idx && head_elem < QK_HEAD_DIM) {
tmp = *reinterpret_cast<const fp4x32_t *>(k_ptr + head_elem / 2);
} else {
tmp = 0;
}
k_local[i][j] = fp4x32_to_fp4x64(tmp);
}
}
// load k_scale: vram -> reg
uint8_t k_scale_local[TOKEN_LOOP][QKHE_LOOP];
for (int i = 0; i < TOKEN_LOOP; ++i) {
const int token_idx = start_token_idx + wid * TOKENS_PER_WARP +
i * TOKENS_PER_LOOP + lane_id16;
const uint8_t *k_scale_ptr =
kv_scale + (kv_idx + token_idx) * (QK_HEAD_DIM / 32);
for (int j = 0; j < QKHE_LOOP; ++j) {
const int head_elem = j * QKHE_PER_LOOP + lane4_id * 32;
if (token_idx < end_token_idx && head_elem < QK_HEAD_DIM) {
k_scale_local[i][j] = k_scale_ptr[head_elem / 32];
} else {
k_scale_local[i][j] = 0;
}
}
}
// q @ k
floatx4_t acc[TOKEN_LOOP];
for (int i = 0; i < TOKEN_LOOP; ++i) {
acc[i] = 0;
for (int j = 0; j < QKHE_LOOP; ++j) {
acc[i] = mfma_16x16x128(k_local[i][j], q_local[j], acc[i],
k_scale_local[i][j], q_scale_local[j]);
}
acc[i] *= scale;
}
float qk_max = -INFINITY;
for (int i = 0; i < TOKEN_LOOP; ++i) {
for (int j = 0; j < 4; ++j) {
qk_max = fmaxf(qk_max, acc[i][j]);
}
}
qk_max = fmaxf(__shfl_xor(qk_max, 16), qk_max);
qk_max = fmaxf(__shfl_xor(qk_max, 32), qk_max);
__shared__ float shared_qk_max[N_WARPS][16];
if (lane_id < 16) {
shared_qk_max[wid][lane_id16] = qk_max;
}
__syncthreads();
for (int w = 0; w < N_WARPS; ++w) {
qk_max = fmaxf(qk_max, shared_qk_max[w][lane_id16]);
}
float exp_sum = 0.f;
for (int i = 0; i < TOKEN_LOOP; ++i) {
const int token_idx = start_token_idx + wid * TOKENS_PER_WARP +
i * TOKENS_PER_LOOP + lane4_id * 4;
for (int j = 0; j < 4; ++j) {
float tmp =
(token_idx + j < end_token_idx) ? __expf(acc[i][j] - qk_max) : 0.f;
acc[i][j] = tmp;
exp_sum += tmp;
}
}
exp_sum += __shfl_xor(exp_sum, 16);
exp_sum += __shfl_xor(exp_sum, 32);
__shared__ float shared_exp_sum[N_WARPS][16];
if (lane_id < 16) {
shared_exp_sum[wid][lane_id16] = exp_sum;
}
__syncthreads();
exp_sum = 0.f;
for (int w = 0; w < N_WARPS; ++w) {
exp_sum += shared_exp_sum[w][lane_id16];
}
if (lse != nullptr && wid == 0 && lane_id < 16) {
lse[(query_idx * N_HEADS + lane_id) * max_num_partitions + part_id] =
logf(exp_sum) + qk_max;
}
bfp16x8_t p[VT_LOOP];
for (int i = 0; i < VT_LOOP; ++i) {
for (int j = 0; j < 4; ++j) {
p[i][j] = float_2_bfloatraw(acc[i * 2][j]);
p[i][j + 4] = float_2_bfloatraw(acc[i * 2 + 1][j]);
}
}
bfp16x8_t v_local[VHE_LOOP][VT_LOOP];
__shared__ bfp16 transposed_v[N_WARPS][QKHE_PER_LOOP][VT_PER_LOOP];
for (int qkhe_d = 0; qkhe_d < V_HEAD_DIM / QKHE_PER_LOOP; ++qkhe_d) {
const int qkh_elem = lane4_id * 32;
for (int vt_d = 0; vt_d < VT_LOOP; ++vt_d) {
for (int i = 0; i < (VT_PER_LOOP / TOKENS_PER_LOOP); ++i) {
int t_d = vt_d * (VT_PER_LOOP / TOKENS_PER_LOOP) + i;
const int token_id = i * 4 + lane_id16 / 4 * 8 + lane_id16 % 4;
fp4x32_t tmp = *reinterpret_cast<fp4x32_t *>(&k_local[t_d][qkhe_d]);
ushort raw_v_scale = (ushort)k_scale_local[t_d][qkhe_d] << 7;
bfp16 v_scale = *reinterpret_cast<bfp16 *>(&raw_v_scale);
for (int j = 0; j < 16; ++j) {
const int head_elem = lane4_id * 32 + j * 2;
bfp16x2_t quan_v = fp4x2_to_bfp16x2(tmp[j]);
bfp16 v0 = reinterpret_cast<bfp16 *>(&quan_v)[0] * v_scale;
bfp16 v1 = reinterpret_cast<bfp16 *>(&quan_v)[1] * v_scale;
transposed_v[wid][head_elem][token_id] = v0;
transposed_v[wid][head_elem + 1][token_id] = v1;
}
}
__syncthreads();
for (int i = 0; i < (QKHE_PER_LOOP / VHE_PER_LOOP); ++i) {
int vhe_d = qkhe_d * (QKHE_PER_LOOP / VHE_PER_LOOP) + i;
const int head_elem = i * VHE_PER_LOOP + lane_id16;
const int token_id = lane4_id * 8;
v_local[vhe_d][vt_d] = *reinterpret_cast<bfp16x8_t *>(
&transposed_v[wid][head_elem][token_id]);
}
__syncthreads();
}
}
floatx4_t o_acc[VHE_LOOP];
for (int i = 0; i < VHE_LOOP; ++i) {
o_acc[i] = 0;
for (int j = 0; j < VT_LOOP; ++j) {
o_acc[i] = mfma_16x16x32(v_local[i][j], p[j], o_acc[i]);
}
}
floatx4_t *shared_out = reinterpret_cast<floatx4_t *>(transposed_v);
constexpr int chunk_vhe = (sizeof(bfp16) * QKHE_PER_LOOP * VT_PER_LOOP) /
(sizeof(floatx4_t) * WARP_SIZE);
for (int i = 0; i < VHE_LOOP; i += chunk_vhe) {
for (int j = 0; j < chunk_vhe; ++j) {
shared_out[(j * N_WARPS + wid) * WARP_SIZE + lane_id] = o_acc[i + j];
}
__syncthreads();
if (wid == 0) {
for (int j = 0; j < chunk_vhe; ++j) {
for (int w = 1; w < N_WARPS; ++w) {
o_acc[i + j] += shared_out[(j * N_WARPS + w) * WARP_SIZE + lane_id];
}
}
}
__syncthreads();
}
if (wid == 0) {
float inv_exp_sum = exp_sum == 0 ? 1.f : 1.f / exp_sum;
for (int i = 0; i < VHE_LOOP; ++i) {
o_acc[i] *= inv_exp_sum;
}
bfp16 *out_ptr =
output +
(query_idx * N_HEADS + lane_id16) * max_num_partitions * V_HEAD_DIM +
part_id * V_HEAD_DIM;
for (int i = 0; i < VHE_LOOP; ++i) {
const int head_elem = i * VHE_PER_LOOP + lane4_id * 4;
bfp16x4_t tmp = {
float_2_bfloatraw(o_acc[i][0]), float_2_bfloatraw(o_acc[i][1]),
float_2_bfloatraw(o_acc[i][2]), float_2_bfloatraw(o_acc[i][3])};
*reinterpret_cast<bfp16x4_t *>(out_ptr + head_elem) = tmp;
}
}
}
__global__ void mla_decode_reduce(
const bfp16 *input, // [total_q, n_heads, max_num_partitions, kv_lora]
const float *lse, // [num_seqs, n_heads, max_num_partitions]
bfp16 *output, // [total_q, num_heads, kv_lora]
const int *qo_indptr, // [num_seqs + 1]
const int *kv_indptr, // [num_seqs + 1]
const int num_seqs, const int max_num_partitions) {
// HEAD_DIM = kv_lora = 512
constexpr int HEAD_DIM = 512;
constexpr int PACKED_SIZE = 8;
constexpr int NUM_HEADS = 16;
constexpr int PARTITION_SIZE = 256;
using packed_io_t = ushort __attribute__((ext_vector_type(PACKED_SIZE)));
using packed_process_t = float __attribute__((ext_vector_type(PACKED_SIZE)));
const int tid = blockIdx.x * blockDim.x + threadIdx.x;
const int global_idx = tid * PACKED_SIZE;
const int seq_id = global_idx / (NUM_HEADS * HEAD_DIM);
const int head_id = (global_idx / HEAD_DIM) % NUM_HEADS;
const int head_element = global_idx % HEAD_DIM;
const int query_id = qo_indptr[seq_id];
const int context_len = kv_indptr[seq_id + 1] - kv_indptr[seq_id];
const int num_partition = CEIL_DIV(context_len, PARTITION_SIZE);
float accum_lse = -INFINITY;
packed_process_t accum_out = 0.0f;
for (int part_id = 0; part_id < num_partition; ++part_id) {
const int input_idx =
((query_id * NUM_HEADS + head_id) * max_num_partitions + part_id) *
HEAD_DIM +
head_element;
packed_io_t local_out_bf16 =
reinterpret_cast<const packed_io_t *>(&input[input_idx])[0];
packed_process_t local_out;
#pragma unroll
for (int i = 0; i < PACKED_SIZE; ++i) {
local_out[i] = __bfloat162float(__ushort_as_bfloat16(local_out_bf16[i]));
}
float local_lse =
lse[(query_id * NUM_HEADS + head_id) * max_num_partitions + part_id];
const float new_lse = logf(expf(accum_lse) + expf(local_lse));
accum_out = accum_out * expf(accum_lse - new_lse) +
local_out * expf(local_lse - new_lse);
accum_lse = new_lse;
}
packed_io_t output_bf16;
#pragma unroll
for (int i = 0; i < PACKED_SIZE; ++i) {
output_bf16[i] = float_2_bfloatraw(accum_out[i]);
}
const int output_idx =
(query_id * NUM_HEADS + head_id) * HEAD_DIM + head_element;
reinterpret_cast<packed_io_t *>(&output[output_idx])[0] = output_bf16;
}
void launch_mla_decode(uintptr_t q, //
uintptr_t q_scale, //
uintptr_t kv, //
uintptr_t kv_scale, //
uintptr_t qo_indptr, //
uintptr_t kv_indptr, //
uintptr_t out, //
uintptr_t lse, //
uintptr_t tmp_out, //
int num_seqs, int max_q_seq_len, int max_kv_seq_len,
float scale, int max_num_partitions) {
const auto *d_q = reinterpret_cast<fp4x2_t *>(q);
const auto *d_q_scale = reinterpret_cast<uint8_t *>(q_scale);
const auto *d_kv = reinterpret_cast<fp4x2_t *>(kv);
const auto *d_kv_scale = reinterpret_cast<uint8_t *>(kv_scale);
const auto *d_qo_indptr = reinterpret_cast<int *>(qo_indptr);
const auto *d_kv_indptr = reinterpret_cast<int *>(kv_indptr);
auto *d_out = reinterpret_cast<bfp16 *>(out);
auto *d_lse = lse == 0 ? nullptr : reinterpret_cast<float *>(lse);
auto *d_tmp_out = tmp_out == 0 ? d_out : reinterpret_cast<bfp16 *>(tmp_out);
dim3 block_size(256);
dim3 grid_size(max_num_partitions, num_seqs);
hipLaunchKernelGGL(mla_decode, grid_size, block_size, 0, 0, d_q, d_q_scale,
d_kv, d_kv_scale, d_qo_indptr, d_kv_indptr, d_tmp_out,
d_lse, num_seqs, max_q_seq_len, max_kv_seq_len, scale,
max_num_partitions);
if (max_num_partitions == 1)
return;
block_size = dim3(256);
grid_size = dim3(num_seqs * 4);
hipLaunchKernelGGL(mla_decode_reduce, grid_size, block_size, 0, 0, d_tmp_out,
d_lse, d_out, d_qo_indptr, d_kv_indptr, num_seqs,
max_num_partitions);
}
PYBIND11_MODULE(mxfp4, m) {
m.def("launch_mla_decode", &launch_mla_decode, "HIP MLA Decode kernel");
}
"""
hip_module = load_inline(
name="mxfp4",
cpp_sources="",
cuda_sources=kernel_cpp,
with_cuda=True,
verbose=False,
extra_cuda_cflags=["-std=c++20", "-O3"],
no_implicit_headers=True,
)
def custom_kernel(data: input_t) -> output_t:
PARITTION_SIZE = 256
q_bf16, kv_data, qo_indptr, kv_indptr, config = data
q, q_scale = quantize_mxfp4(q_bf16)
kv_cache, kv_scale = kv_data["mxfp4"]
num_seqs = config["batch_size"]
max_q_seq_len = config["q_seq_len"]
max_kv_seq_len = config["kv_seq_len"]
scale = config["sm_scale"]
out = torch.empty((num_seqs, 16, 512), dtype=torch.bfloat16, device=q.device)
max_num_partitions = (max_kv_seq_len + PARITTION_SIZE - 1) // PARITTION_SIZE
lse = None
tmp_out = None
if max_num_partitions > 1:
lse = torch.empty(
(num_seqs, 16, max_num_partitions), dtype=torch.float32, device=q.device
)
tmp_out = torch.empty(
(num_seqs, 16, max_num_partitions, 512),
dtype=torch.bfloat16,
device=q.device,
)
hip_module.launch_mla_decode(
q.data_ptr(),
q_scale.data_ptr(),
kv_cache.data_ptr(),
kv_scale.data_ptr(),
qo_indptr.data_ptr(),
kv_indptr.data_ptr(),
out.data_ptr(),
lse.data_ptr() if lse is not None else 0,
tmp_out.data_ptr() if lse is not None else 0,
num_seqs,
max_q_seq_len,
max_kv_seq_len,
scale,
max_num_partitions,
)
return out
scrolls · 634 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