Skip to content
KernelIndex
Search⌘K

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
AMD Instinct MI355X
342.4µs
#707 of 766
2026-03-29

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-softmaxconst float m_new = fmaxf(m_val, dot);
shared-memory__shared__ float s_m_global;
split-kScalar 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