submission 661291
Zephyr Zhao · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 277 lines, June 9 Researcher Reciprocity License v1.0.
exp_v2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-661291?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:2f8a3b66954764f67d5c5cc8b927185e217dc1a218bf6d7b2b2226f0006a2b76
license declaredunknown
license concludedunknown
authorsZephyr Zhao
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
online-softmax
const float m_new = fmaxf(m_val, dot);shared-memory
__shared__ float s_m_global;split-k
Scalar dot products with flash-decoding split-K. Pre-allocated buffers.Kernel source
exp_v2.py277 lines
"""
HIP C++ MLA decode v2: BF16 KV, stride-64 coalesced loads, shared K/V read.
Scalar dot products with flash-decoding split-K. Pre-allocated buffers.
"""
import os
os.environ['PYTORCH_ROCM_ARCH'] = 'gfx950'
os.environ['CXX'] = 'clang++'
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
NUM_SPLITS = 32
CUDA_SRC = r"""
#include <hip/hip_runtime.h>
#include <hip/amd_detail/amd_hip_bf16.h>
#include <cmath>
#define WARP_SIZE 64
#define NH 16
#define DK 576
#define DV 512
#define ELEMS_K 9 // 576/64
#define ELEMS_V 8 // 512/64
// Phase 1: Partial attention — BF16 KV, stride-64 coalesced, shared K/V load
// Grid: (B * NS, 1, 1)
// Block: (NH * WARP_SIZE = 1024, 1, 1)
__global__ void mla_partial_attn(
const __hip_bfloat16* __restrict__ Q, // [B, NH, DK]
const __hip_bfloat16* __restrict__ KV, // [total_kv, DK]
const int* __restrict__ kv_indptr, // [B+1]
float* __restrict__ partial_O, // [B * NH * NS, DV]
float* __restrict__ partial_m, // [B * NH * NS]
float* __restrict__ partial_l, // [B * NH * NS]
const int B, const int NS, const float sm_scale
) {
const int block_id = blockIdx.x;
const int split_id = block_id % NS;
const int batch_id = block_id / NS;
if (batch_id >= B) return;
const int warp_id = threadIdx.x / WARP_SIZE;
const int lane_id = threadIdx.x % WARP_SIZE;
const int head_id = warp_id;
// KV range for this batch
const int kv_start = kv_indptr[batch_id];
const int kv_end = kv_indptr[batch_id + 1];
const int kv_len = kv_end - kv_start;
const int chunk = (kv_len + NS - 1) / NS;
const int my_start = kv_start + split_id * chunk;
const int my_end = min(my_start + chunk, kv_end);
// Load Q[batch, head, :] with stride-64 coalesced layout, pre-scaled
// Lane handles indices: lane, lane+64, lane+128, ..., lane+512
float qr[ELEMS_K];
#pragma unroll
for (int i = 0; i < ELEMS_K; i++) {
const int idx = i * WARP_SIZE + lane_id;
qr[i] = (idx < DK)
? __bfloat162float(Q[batch_id * NH * DK + head_id * DK + idx]) * sm_scale
: 0.0f;
}
// Softmax state and V accumulator
float m_val = -1e30f;
float l_val = 0.0f;
float v_acc[ELEMS_V];
#pragma unroll
for (int i = 0; i < ELEMS_V; i++) v_acc[i] = 0.0f;
// Main loop: iterate over KV tokens
for (int t = my_start; t < my_end; t++) {
const __hip_bfloat16* kv_ptr = KV + (long long)t * DK;
// Load KV once (stride-64 coalesced), compute dot product, save for V
float dot = 0.0f;
float kv_vals[ELEMS_K];
#pragma unroll
for (int i = 0; i < ELEMS_K; i++) {
const int idx = i * WARP_SIZE + lane_id;
// Coalesced: consecutive lanes read consecutive addresses
float k_val = __bfloat162float(kv_ptr[idx]); // DK=576=9*64, always in bounds
kv_vals[i] = k_val;
dot += qr[i] * k_val;
}
// Warp reduction (64 lanes)
#pragma unroll
for (int offset = 32; offset > 0; offset >>= 1) {
dot += __shfl_xor(dot, offset, WARP_SIZE);
}
// Online softmax
const float m_new = fmaxf(m_val, dot);
const float exp_old = __expf(m_val - m_new);
const float exp_cur = __expf(dot - m_new);
// V accumulation: reuse first 8 of 9 KV values (first 512 dims)
#pragma unroll
for (int i = 0; i < ELEMS_V; i++) {
v_acc[i] = v_acc[i] * exp_old + exp_cur * kv_vals[i];
}
m_val = m_new;
l_val = exp_old * l_val + exp_cur;
}
// Store partial results
const int out_idx = (batch_id * NH + head_id) * NS + split_id;
#pragma unroll
for (int i = 0; i < ELEMS_V; i++) {
const int idx = i * WARP_SIZE + lane_id;
partial_O[(long long)out_idx * DV + idx] = v_acc[i];
}
if (lane_id == 0) {
partial_m[out_idx] = m_val;
partial_l[out_idx] = l_val;
}
}
// Phase 2: Reduce
__global__ void mla_reduce(
const float* __restrict__ partial_O,
const float* __restrict__ partial_m,
const float* __restrict__ partial_l,
__hip_bfloat16* __restrict__ O,
const int NS
) {
const int bh = blockIdx.x;
const int tid = threadIdx.x;
__shared__ float s_m_global;
__shared__ float s_inv_l;
if (tid == 0) {
float mg = -1e30f;
for (int s = 0; s < NS; s++)
mg = fmaxf(mg, partial_m[bh * NS + s]);
s_m_global = mg;
float lt = 0.0f;
for (int s = 0; s < NS; s++)
lt += __expf(partial_m[bh * NS + s] - mg) * partial_l[bh * NS + s];
s_inv_l = (lt > 0.0f) ? (1.0f / lt) : 0.0f;
}
__syncthreads();
const float mg = s_m_global;
const float inv_l = s_inv_l;
for (int d = tid; d < DV; d += blockDim.x) {
float o_sum = 0.0f;
for (int s = 0; s < NS; s++) {
const int pidx = bh * NS + s;
o_sum += __expf(partial_m[pidx] - mg) * partial_O[(long long)pidx * DV + d];
}
O[(long long)bh * DV + d] = __float2bfloat16(o_sum * inv_l);
}
}
// Static buffer cache (avoids per-call allocation)
static float* g_partial_O = nullptr;
static float* g_partial_m = nullptr;
static float* g_partial_l = nullptr;
static int64_t g_buf_size_o = 0;
static int64_t g_buf_size_ml = 0;
static void ensure_buffers(int64_t size_o, int64_t size_ml) {
if (size_o > g_buf_size_o) {
if (g_partial_O) hipFree(g_partial_O);
hipMalloc(&g_partial_O, size_o * sizeof(float));
g_buf_size_o = size_o;
}
if (size_ml > g_buf_size_ml) {
if (g_partial_m) hipFree(g_partial_m);
if (g_partial_l) hipFree(g_partial_l);
hipMalloc(&g_partial_m, size_ml * sizeof(float));
hipMalloc(&g_partial_l, size_ml * sizeof(float));
g_buf_size_ml = size_ml;
}
}
void mla_decode(
torch::Tensor Q, // [B, NH, DK] bf16
torch::Tensor KV, // [total_kv, DK] bf16
torch::Tensor kv_indptr, // [B+1] int32
torch::Tensor O, // [B, NH, DV] bf16 (pre-allocated)
int NS,
float sm_scale
) {
const int B = Q.size(0);
const int nh = Q.size(1);
const int dv = DV;
int64_t n = (int64_t)B * nh * NS;
ensure_buffers(n * dv, n);
// Phase 1
dim3 grid1(B * NS);
dim3 block1(nh * 64);
mla_partial_attn<<<grid1, block1>>>(
reinterpret_cast<const __hip_bfloat16*>(Q.data_ptr()),
reinterpret_cast<const __hip_bfloat16*>(KV.data_ptr()),
kv_indptr.data_ptr<int>(),
g_partial_O,
g_partial_m,
g_partial_l,
B, NS, sm_scale
);
// Phase 2: write into pre-allocated O
dim3 grid2(B * nh);
dim3 block2(256);
mla_reduce<<<grid2, block2>>>(
g_partial_O,
g_partial_m,
g_partial_l,
reinterpret_cast<__hip_bfloat16*>(O.data_ptr()),
NS
);
}
"""
CPP_SRC = """
void mla_decode(
torch::Tensor Q,
torch::Tensor KV,
torch::Tensor kv_indptr,
torch::Tensor O,
int NS,
float sm_scale
);
"""
module = load_inline(
name='mla_decode_hip_v3c',
cpp_sources=[CPP_SRC],
cuda_sources=[CUDA_SRC],
functions=['mla_decode'],
verbose=True,
extra_cuda_cflags=["--offload-arch=gfx950", "-std=c++20"],
)
# Pre-allocated output tensor cache
_output_cache = {}
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
B = config["batch_size"]
NH = config["num_heads"]
DK = config["qk_head_dim"]
DV = config["v_head_dim"]
sm = config["sm_scale"]
# BF16 KV — squeeze the middle dim=1
kv_bf16 = kv_data["bf16"]
kv_flat = kv_bf16.reshape(-1, DK)
q_3d = q.reshape(B, NH, DK)
# Pre-cached output tensor (avoids torch::empty per call)
key = (B, NH, DV)
if key not in _output_cache:
_output_cache[key] = torch.empty(B, NH, DV, dtype=torch.bfloat16, device=q.device)
O = _output_cache[key]
module.mla_decode(q_3d, kv_flat, kv_indptr, O, NUM_SPLITS, sm)
return O
scrolls · 277 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