Skip to content
KernelIndex
Search⌘K

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
AMD Instinct MI355X
4.23ms
#744 of 766
2026-03-21

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-memoryextern __shared__ float q_smem[]; // qk_head_dim floats

Kernel 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