submission 601877
huismiling · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 383 lines, June 9 Researcher Reciprocity License v1.0.
submission_v0.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-601877?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:0f39169b3e84a244955d277cfe2919281f278b57152b8a6db2ee53553cde736c
license declaredunknown
license concludedunknown
authorshuismiling
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
extern __shared__ float q_smem[]; // qk_head_dim floatsKernel source
submission_v0.py383 lines
"""
MLA Decode Attention — HIP C++ kernel via torch.utils.cpp_extension.load_inline
Implements the same functionality as custom_kernel_bf16() from submission_0.py
using a fused HIP kernel with online softmax (Flash Attention style).
Kernel design:
- Grid: (num_heads, total_q) — one block per (head, query_position) pair
- Block: 256 threads
- Online softmax to avoid materializing the full attention score matrix
- Each thread handles 2 output dims (512 / 256 = 2)
- Threads cooperate on 576-dim dot product via warp + block reduction
- AMD wavefront size = 64 → 4 warps per block
"""
import os
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950;gfx942"
import torch
from torch.utils.cpp_extension import load_inline
# ---------------------------------------------------------------------------
# HIP kernel source (compiled with hipcc via PyTorch's ROCm backend)
# ---------------------------------------------------------------------------
hip_source = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <hip/hip_runtime.h>
#define BLOCK_SIZE 256
// AMD CDNA (MI300X) wavefront size
#define WARP_SZ 64
#define NUM_WARPS (BLOCK_SIZE / WARP_SZ) // 4
// ---- bf16 <-> fp32 helpers (bit-level, no header dependency) ----
__device__ __forceinline__ float bf16_to_fp32(const void* ptr, int idx) {
uint16_t raw = reinterpret_cast<const uint16_t*>(ptr)[idx];
union { uint32_t u; float f; } cvt;
cvt.u = (uint32_t)raw << 16;
return cvt.f;
}
__device__ __forceinline__ void fp32_to_bf16(void* ptr, int idx, float val) {
union { uint32_t u; float f; } cvt;
cvt.f = val;
// Round to nearest even
uint32_t rounding = ((cvt.u >> 16) & 1u) + 0x7FFFu;
cvt.u += rounding;
reinterpret_cast<uint16_t*>(ptr)[idx] = (uint16_t)(cvt.u >> 16);
}
// ---- warp-level sum reduction via shuffle ----
__device__ __forceinline__ float warp_reduce_sum(float val) {
#pragma unroll
for (int offset = WARP_SZ / 2; offset > 0; offset >>= 1) {
val += __shfl_xor(val, offset, WARP_SZ);
}
return val;
}
// =========================================================================
// MLA Decode Attention kernel (bf16 Q + bf16 KV, fp32 accumulation)
//
// Grid : (num_heads, total_q)
// Block: (BLOCK_SIZE)
//
// For each (head, query) pair the block iterates over all KV positions of
// the corresponding batch element, computing attention with online softmax
// so that the full score matrix is never materialised.
//
// Memory layout (row-major, contiguous):
// q : (total_q, num_heads, qk_head_dim) bf16
// kv : (total_kv, qk_head_dim) bf16 (kv_heads dim squeezed)
// out : (total_q, num_heads, v_head_dim) bf16
// =========================================================================
__global__ void mla_attention_bf16_kernel(
const void* __restrict__ q_ptr, // bf16
const void* __restrict__ kv_ptr, // bf16
void* __restrict__ out_ptr, // bf16
const int* __restrict__ qo_indptr, // (batch_size + 1)
const int* __restrict__ kv_indptr, // (batch_size + 1)
const int batch_size,
const int num_heads,
const int qk_head_dim, // 576
const int v_head_dim, // 512
const float sm_scale)
{
const int head_idx = blockIdx.x;
const int q_idx = blockIdx.y;
const int tid = threadIdx.x;
// ---- locate batch via binary search on qo_indptr ----
int lo = 0, hi = batch_size;
while (lo < hi) {
int mid = (lo + hi) >> 1;
if (qo_indptr[mid + 1] <= q_idx) lo = mid + 1;
else hi = mid;
}
const int batch_idx = lo;
if (batch_idx >= batch_size) return;
const int kv_start = kv_indptr[batch_idx];
const int kv_end = kv_indptr[batch_idx + 1];
// ---- load query vector into dynamic shared memory ----
extern __shared__ float q_smem[]; // qk_head_dim floats
__shared__ float warp_buf[NUM_WARPS]; // for block-level reduction
__shared__ float score_broadcast; // broadcast final score
const int q_base = q_idx * num_heads * qk_head_dim + head_idx * qk_head_dim;
for (int d = tid; d < qk_head_dim; d += BLOCK_SIZE) {
q_smem[d] = bf16_to_fp32(q_ptr, q_base + d);
}
__syncthreads();
// ---- online softmax state ----
float m_val = -1e38f; // running max
float d_sum = 0.0f; // running sum(exp)
// Output accumulators — each thread owns 2 output dims: tid and tid+BLOCK_SIZE
// v_head_dim=512, BLOCK_SIZE=256 → exactly 2 dims per thread
float o0 = 0.0f, o1 = 0.0f;
// ---- iterate over KV positions ----
for (int kv_pos = kv_start; kv_pos < kv_end; ++kv_pos) {
const int kv_base = kv_pos * qk_head_dim;
// Load KV values (reused for both dot-product and value accumulation)
float kv_val0 = bf16_to_fp32(kv_ptr, kv_base + tid); // dim tid (<256)
float kv_val1 = (tid + BLOCK_SIZE < qk_head_dim)
? bf16_to_fp32(kv_ptr, kv_base + tid + BLOCK_SIZE) // dim tid+256 (<512)
: 0.0f;
// ---- partial dot product q · k ----
float partial = q_smem[tid] * kv_val0;
if (tid + BLOCK_SIZE < qk_head_dim) {
partial += q_smem[tid + BLOCK_SIZE] * kv_val1;
}
// Handle remaining dims (512..575 → only threads 0..63)
for (int d = tid + 2 * BLOCK_SIZE; d < qk_head_dim; d += BLOCK_SIZE) {
partial += q_smem[d] * bf16_to_fp32(kv_ptr, kv_base + d);
}
// ---- warp-level reduction ----
partial = warp_reduce_sum(partial);
// ---- block-level reduction ----
const int warp_id = tid / WARP_SZ;
const int lane_id = tid % WARP_SZ;
if (lane_id == 0) {
warp_buf[warp_id] = partial;
}
__syncthreads();
if (tid == 0) {
float total = 0.0f;
#pragma unroll
for (int w = 0; w < NUM_WARPS; ++w) total += warp_buf[w];
score_broadcast = total * sm_scale;
}
__syncthreads();
const float score = score_broadcast;
// ---- online softmax update ----
float new_m = fmaxf(m_val, score);
float corr = expf(m_val - new_m); // correction factor (≤ 1)
float weight = expf(score - new_m); // new element weight (≤ 1)
d_sum = d_sum * corr + weight;
// ---- accumulate weighted values ----
if (tid < v_head_dim) {
o0 = o0 * corr + weight * kv_val0;
}
if (tid + BLOCK_SIZE < v_head_dim) {
o1 = o1 * corr + weight * kv_val1;
}
m_val = new_m;
}
// ---- finalize: divide by sum(exp) ----
float inv_d = (d_sum > 0.0f) ? (1.0f / d_sum) : 0.0f;
const int out_base = q_idx * num_heads * v_head_dim + head_idx * v_head_dim;
if (tid < v_head_dim) {
fp32_to_bf16(out_ptr, out_base + tid, o0 * inv_d);
}
if (tid + BLOCK_SIZE < v_head_dim) {
fp32_to_bf16(out_ptr, out_base + tid + BLOCK_SIZE, o1 * inv_d);
}
}
// =========================================================================
// C++ wrapper callable from Python (via pybind11)
// =========================================================================
torch::Tensor mla_attention_bf16(
torch::Tensor q, // (total_q, num_heads, qk_head_dim) bf16
torch::Tensor kv_buffer, // (total_kv, 1, qk_head_dim) bf16
torch::Tensor qo_indptr, // (batch_size + 1,) int32
torch::Tensor kv_indptr, // (batch_size + 1,) int32
int num_heads,
int kv_lora_rank, // v_head_dim = 512
float sm_scale)
{
TORCH_CHECK(q.is_cuda(), "q must be on GPU");
TORCH_CHECK(kv_buffer.is_cuda(), "kv_buffer must be on GPU");
TORCH_CHECK(q.dtype() == torch::kBFloat16, "q must be bf16");
TORCH_CHECK(kv_buffer.dtype() == torch::kBFloat16, "kv_buffer must be bf16");
q = q.contiguous();
kv_buffer = kv_buffer.contiguous();
qo_indptr = qo_indptr.contiguous();
kv_indptr = kv_indptr.contiguous();
const int total_q = q.size(0);
const int qk_head_dim = q.size(2); // 576
const int v_head_dim = kv_lora_rank; // 512
const int batch_size = qo_indptr.size(0) - 1;
// Squeeze the single kv-head dim: (total_kv, 1, D) → (total_kv, D)
auto kv = kv_buffer.squeeze(1);
auto output = torch::empty({total_q, num_heads, v_head_dim}, q.options());
if (total_q == 0) return output;
dim3 grid(num_heads, total_q);
dim3 block(BLOCK_SIZE);
// Dynamic shared memory for q_smem only (warp_buf + score_broadcast are static)
int shared_mem_bytes = qk_head_dim * (int)sizeof(float);
mla_attention_bf16_kernel<<<grid, block, shared_mem_bytes>>>(
q.data_ptr(),
kv.data_ptr(),
output.data_ptr(),
qo_indptr.data_ptr<int>(),
kv_indptr.data_ptr<int>(),
batch_size,
num_heads,
qk_head_dim,
v_head_dim,
sm_scale
);
return output;
}
"""
# ---------------------------------------------------------------------------
# C++ forward declaration (compiled with g++, used by pybind11 auto-binding)
# ---------------------------------------------------------------------------
cpp_source = """
#include <torch/extension.h>
torch::Tensor mla_attention_bf16(
torch::Tensor q,
torch::Tensor kv_buffer,
torch::Tensor qo_indptr,
torch::Tensor kv_indptr,
int num_heads,
int kv_lora_rank,
float sm_scale);
"""
# ---------------------------------------------------------------------------
# Compile & load
# ---------------------------------------------------------------------------
mla_module = load_inline(
name="mla_attention_hip",
cpp_sources=[cpp_source],
cuda_sources=[hip_source],
functions=["mla_attention_bf16"],
verbose=True,
extra_cuda_cflags=["-O3"],
)
# ---------------------------------------------------------------------------
# Drop-in replacement for custom_kernel_bf16 from submission_0.py
# ---------------------------------------------------------------------------
def custom_kernel(data):
q, kv_data, qo_indptr, kv_indptr, config = data
num_heads = config["num_heads"] # 16
kv_lora_rank = config["kv_lora_rank"] # 512
sm_scale = config["sm_scale"] # 1/sqrt(576)
kv_buffer_bf16 = kv_data["bf16"] # (total_kv, 1, 576) bf16
output = mla_module.mla_attention_bf16(
q, kv_buffer_bf16, qo_indptr, kv_indptr,
num_heads, kv_lora_rank, sm_scale,
)
return output
# ---------------------------------------------------------------------------
# Quick correctness test
# ---------------------------------------------------------------------------
if __name__ == "__main__":
import math
torch.manual_seed(42)
device = "cuda"
# MLA config
num_heads = 16
kv_lora_rank = 512
qk_head_dim = 576
sm_scale = 1.0 / math.sqrt(qk_head_dim)
# Small variable-length batch: 3 sequences
seq_q_lens = [1, 2, 1] # decode-like
seq_kv_lens = [128, 256, 64]
total_q = sum(seq_q_lens)
total_kv = sum(seq_kv_lens)
qo_indptr = torch.tensor(
[0] + list(torch.cumsum(torch.tensor(seq_q_lens), 0).numpy()),
dtype=torch.int32, device=device,
)
kv_indptr = torch.tensor(
[0] + list(torch.cumsum(torch.tensor(seq_kv_lens), 0).numpy()),
dtype=torch.int32, device=device,
)
q = torch.randn(total_q, num_heads, qk_head_dim, device=device, dtype=torch.bfloat16)
kv_buffer = torch.randn(total_kv, 1, qk_head_dim, device=device, dtype=torch.bfloat16)
config = {
"num_heads": num_heads,
"kv_lora_rank": kv_lora_rank,
"sm_scale": sm_scale,
}
kv_data = {"bf16": kv_buffer}
data = (q, kv_data, qo_indptr, kv_indptr, config)
# --- Reference (Python baseline from submission_0.py) ---
import torch.nn.functional as F
def ref_kernel(data):
q, kv_data, qo_indptr, kv_indptr, config = data
num_heads = config["num_heads"]
kv_lora_rank = config["kv_lora_rank"]
sm_scale_ = config["sm_scale"]
kv_buf = kv_data["bf16"]
batch_size = qo_indptr.shape[0] - 1
out_list = []
for i in range(batch_size):
q_s, q_e = int(qo_indptr[i].item()), int(qo_indptr[i + 1].item())
kv_s, kv_e = int(kv_indptr[i].item()), int(kv_indptr[i + 1].item())
qi = q[q_s:q_e]
kvc = kv_buf[kv_s:kv_e, 0]
ki = kvc
vi = kvc[:, :kv_lora_rank]
qi_t = qi.float().permute(1, 0, 2)
scores = torch.matmul(qi_t * sm_scale_, ki.float().T)
scores = F.softmax(scores, dim=-1)
oi = torch.matmul(scores, vi.float()).permute(1, 0, 2)
out_list.append(oi.to(torch.bfloat16))
return torch.cat(out_list, dim=0)
ref_out = ref_kernel(data)
hip_out = custom_kernel(data)
max_err = (ref_out.float() - hip_out.float()).abs().max().item()
mean_err = (ref_out.float() - hip_out.float()).abs().mean().item()
print(f"ref shape: {ref_out.shape} dtype: {ref_out.dtype}")
print(f"hip shape: {hip_out.shape} dtype: {hip_out.dtype}")
print(f"max abs error: {max_err:.6e}")
print(f"mean abs error: {mean_err:.6e}")
if max_err < 1e-2:
print("PASSED — HIP kernel matches reference")
else:
print("FAILED — error too large")
scrolls · 383 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