Skip to content
KernelIndex
Search⌘K

submission 728245

divc13 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 767 lines, June 9 Researcher Reciprocity License v1.0.

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-728245?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
43.0µs
#127 of 766
2026-04-04

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:20476786658f918a6b6ba6f82604b33326d93924228e7e49c1151bd6dec8cbf0
license declaredunknown
license concludedunknown
authorsdivc13
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

num-warps = 4static constexpr int NUM_WARPS = 4;
shared-memory__shared__ __align__(16) unsigned char kv_lds[2][KV_TILE_BYTES];
split-kconst int split_kv_start, const int split_kv_end,

Kernel source

submission.py767 lines
import torch
import os
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t

# ---------------------------------------------------------------------------
# MLA decode v352: v_perm_b32 V byte extraction
# - Replace shift/mask/or V packing with v_perm_b32 chains (12 vs ~42 ALU)
# - Reduces exposed V extraction latency from ~18 to ~0 cycles per SV chunk
# - All other optimizations from v351 (adaptive grid, ds_swizzle, occ3)
# ---------------------------------------------------------------------------

os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"

HIP_SRC = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>

static constexpr int WARP_SIZE = 64;
static constexpr int NUM_WARPS = 4;
static constexpr int BLOCK_SIZE = WARP_SIZE * NUM_WARPS;

static constexpr int QK_DIM = 576;
static constexpr int V_DIM  = 512;
static constexpr int NUM_HEADS = 16;
static constexpr int NUM_K_CHUNKS = QK_DIM / 32;  // 18
static constexpr int NUM_K128_CHUNKS = (QK_DIM + 127) / 128;  // 5
static constexpr int SUPER_TILE = 32;
static constexpr int SV_CHUNKS = 8;
static constexpr int KV_TILE_BYTES = SUPER_TILE * QK_DIM;  // 18432

// 18432 bytes / 16 bytes per uint4 / 256 threads = 4.5 -> 5 rounds
static constexpr int PF_UINT4S = KV_TILE_BYTES / 16;  // 1152
static constexpr int PF_ROUNDS = (PF_UINT4S + BLOCK_SIZE - 1) / BLOCK_SIZE;  // 5

typedef float __attribute__((ext_vector_type(4))) v4f32;
typedef unsigned int __attribute__((ext_vector_type(4))) u32x4;
typedef int __attribute__((ext_vector_type(4))) i32x4;
typedef int __attribute__((ext_vector_type(8))) i32x8;
typedef unsigned int __attribute__((address_space(3)))* lds_ptr_t;

extern "C" __device__ void __llvm_amdgcn_raw_buffer_load_lds(
    i32x4 rsrc, lds_ptr_t lds_ptr, int size,
    int voffset, int soffset, int offset, int aux)
    __asm("llvm.amdgcn.raw.buffer.load.lds");

#define SWIZZLE_XOR(val, xor_val) \
    __int_as_float(__builtin_amdgcn_ds_swizzle(__float_as_int(val), 0xFC00 | (xor_val)))

struct buffer_resource { uint64_t ptr; uint32_t range; uint32_t config; };

__device__ __forceinline__ i32x4 make_buffer_rsrc(const void* p, uint32_t bytes) {
    buffer_resource r = {reinterpret_cast<uint64_t>(p), bytes, 0x110000};
    return *reinterpret_cast<i32x4*>(&r);
}

__device__ __forceinline__ float bf16_to_f32(unsigned short v) {
    return __uint_as_float(static_cast<unsigned int>(v) << 16);
}

__device__ __forceinline__ unsigned short f32_to_bf16(float v) {
    unsigned int bits = __float_as_uint(v);
    bits += 0x7FFF + ((bits >> 16) & 1);
    return static_cast<unsigned short>(bits >> 16);
}

__device__ __forceinline__ float fp8_to_f32(unsigned char b) {
    return __builtin_amdgcn_cvt_f32_fp8(static_cast<int>(b), 0);
}

__device__ __forceinline__ v4f32 mfma_f32_16x16x128_fp8(
    i32x8 A, i32x8 B, v4f32 C)
{
    v4f32 D;
    asm(
        "v_mfma_f32_16x16x128_f8f6f4 %0, %1, %2, %3 cbsz:0 blgp:0"
        : "=v"(D) : "v"(A), "v"(B), "v"(C));
    return D;
}

// =========================================================================
// Tile processing: template specialization for full vs partial tiles
// FULL_TILE=true: stcnt=32 (compile-time), all bounds checks eliminated
// FULL_TILE=false: runtime stcnt with full bounds checks
// =========================================================================

template<bool FULL_TILE>
__device__ __forceinline__ void process_tile(
    const unsigned char* __restrict__ kv_ptr,
    unsigned char kv_lds[][KV_TILE_BYTES],
    float s_W[][16][33],
    const i32x8* q_128,
    const float score_scale,
    float mv[4], float lv[4],
    float vacc[][4],
    int& cur_buf,
    const int mr, const int kg, const int tid, const int warp_id,
    const int split_kv_start, const int split_kv_end,
    const int total_tokens, const int num_st,
    const int st_idx, const int stcnt_arg)
{
    const int stcnt = FULL_TILE ? SUPER_TILE : stcnt_arg;
    const int ta = FULL_TILE ? 16 : min(16, stcnt);
    const int tb = FULL_TILE ? 16 : max(0, stcnt - 16);
    const unsigned char* kv_cur = kv_lds[cur_buf];

    // ---- INTERLEAVED DMA + QK: 1 DMA round per QK chunk ----
    // PF_ROUNDS == NUM_K128_CHUNKS == 5, so natural 1:1 interleaving.
    // DMA writes to kv_lds[nxt_buf] via VMEM, QK reads kv_lds[cur_buf] via LDS.
    __builtin_amdgcn_s_setprio(3);
    const int nxt_buf = cur_buf ^ 1;
    const bool has_next = (st_idx + 1 < num_st);

    i32x4 srsrc = {};
    if (has_next) {
        const int nxt_start = split_kv_start + (st_idx + 1) * SUPER_TILE;
        const int nxt_bytes = min(SUPER_TILE, split_kv_end - nxt_start) * QK_DIM;
        const unsigned char* __restrict__ nsrc = kv_ptr +
            static_cast<long long>(nxt_start) * QK_DIM;
        srsrc = make_buffer_rsrc(nsrc, nxt_bytes);
    }

    v4f32 ca = {0, 0, 0, 0};
    v4f32 cb = {0, 0, 0, 0};
    #pragma unroll
    for (int c = 0; c < NUM_K128_CHUNKS; c++) {
        if (has_next) {
            int dma_off = tid * 16 + c * BLOCK_SIZE * 16;
            lds_ptr_t ldp = (lds_ptr_t)(reinterpret_cast<uintptr_t>(kv_lds[nxt_buf]) + dma_off);
            __llvm_amdgcn_raw_buffer_load_lds(srsrc, ldp, 16, dma_off, 0, 0, 2);
        }
        i32x8 ba = {};
        if (FULL_TILE || mr < ta) {
            int base1 = mr * QK_DIM + c * 128 + 16 * kg;
            #pragma unroll
            for (int i = 0; i < 4; i++) {
                int off = base1 + i * 4;
                if (c * 128 + 16 * kg + i * 4 + 4 <= QK_DIM)
                    ba[i] = *reinterpret_cast<const int*>(&kv_cur[off]);
            }
            int base2 = mr * QK_DIM + c * 128 + 64 + 16 * kg;
            #pragma unroll
            for (int i = 0; i < 4; i++) {
                int off = base2 + i * 4;
                if (c * 128 + 64 + 16 * kg + i * 4 + 4 <= QK_DIM)
                    ba[4 + i] = *reinterpret_cast<const int*>(&kv_cur[off]);
            }
        }
        i32x8 bb = {};
        if (FULL_TILE || mr < tb) {
            int base1 = (16 + mr) * QK_DIM + c * 128 + 16 * kg;
            #pragma unroll
            for (int i = 0; i < 4; i++) {
                int off = base1 + i * 4;
                if (c * 128 + 16 * kg + i * 4 + 4 <= QK_DIM)
                    bb[i] = *reinterpret_cast<const int*>(&kv_cur[off]);
            }
            int base2 = (16 + mr) * QK_DIM + c * 128 + 64 + 16 * kg;
            #pragma unroll
            for (int i = 0; i < 4; i++) {
                int off = base2 + i * 4;
                if (c * 128 + 64 + 16 * kg + i * 4 + 4 <= QK_DIM)
                    bb[4 + i] = *reinterpret_cast<const int*>(&kv_cur[off]);
            }
        }
        ca = mfma_f32_16x16x128_fp8(q_128[c], ba, ca);
        cb = mfma_f32_16x16x128_fp8(q_128[c], bb, cb);
    }
    __builtin_amdgcn_s_setprio(0);

    // ---- Online softmax (16-lane reduce only: offsets 8,4,2,1) ----
    float sa0 = ca[0] * score_scale, sa1 = ca[1] * score_scale;
    float sa2 = ca[2] * score_scale, sa3 = ca[3] * score_scale;
    float sb0 = cb[0] * score_scale, sb1 = cb[1] * score_scale;
    float sb2 = cb[2] * score_scale, sb3 = cb[3] * score_scale;

    if (!FULL_TILE && mr >= ta) { sa0 = sa1 = sa2 = sa3 = -1e30f; }
    if (!FULL_TILE && mr >= tb) { sb0 = sb1 = sb2 = sb3 = -1e30f; }

    float tm0 = fmaxf(sa0, sb0), tm1 = fmaxf(sa1, sb1);
    float tm2 = fmaxf(sa2, sb2), tm3 = fmaxf(sa3, sb3);
    tm0 = fmaxf(tm0, SWIZZLE_XOR(tm0, 8));
    tm1 = fmaxf(tm1, SWIZZLE_XOR(tm1, 8));
    tm2 = fmaxf(tm2, SWIZZLE_XOR(tm2, 8));
    tm3 = fmaxf(tm3, SWIZZLE_XOR(tm3, 8));
    tm0 = fmaxf(tm0, SWIZZLE_XOR(tm0, 4));
    tm1 = fmaxf(tm1, SWIZZLE_XOR(tm1, 4));
    tm2 = fmaxf(tm2, SWIZZLE_XOR(tm2, 4));
    tm3 = fmaxf(tm3, SWIZZLE_XOR(tm3, 4));
    tm0 = fmaxf(tm0, SWIZZLE_XOR(tm0, 2));
    tm1 = fmaxf(tm1, SWIZZLE_XOR(tm1, 2));
    tm2 = fmaxf(tm2, SWIZZLE_XOR(tm2, 2));
    tm3 = fmaxf(tm3, SWIZZLE_XOR(tm3, 2));
    tm0 = fmaxf(tm0, SWIZZLE_XOR(tm0, 1));
    tm1 = fmaxf(tm1, SWIZZLE_XOR(tm1, 1));
    tm2 = fmaxf(tm2, SWIZZLE_XOR(tm2, 1));
    tm3 = fmaxf(tm3, SWIZZLE_XOR(tm3, 1));

    float nm0 = fmaxf(mv[0], tm0), nm1 = fmaxf(mv[1], tm1);
    float nm2 = fmaxf(mv[2], tm2), nm3 = fmaxf(mv[3], tm3);
    float rc0 = __expf(mv[0] - nm0), rc1 = __expf(mv[1] - nm1);
    float rc2 = __expf(mv[2] - nm2), rc3 = __expf(mv[3] - nm3);
    mv[0] = nm0; mv[1] = nm1; mv[2] = nm2; mv[3] = nm3;

    float wa0 = (FULL_TILE || mr < ta) ? __expf(sa0 - nm0) : 0.f;
    float wa1 = (FULL_TILE || mr < ta) ? __expf(sa1 - nm1) : 0.f;
    float wa2 = (FULL_TILE || mr < ta) ? __expf(sa2 - nm2) : 0.f;
    float wa3 = (FULL_TILE || mr < ta) ? __expf(sa3 - nm3) : 0.f;
    float wb0 = (FULL_TILE || mr < tb) ? __expf(sb0 - nm0) : 0.f;
    float wb1 = (FULL_TILE || mr < tb) ? __expf(sb1 - nm1) : 0.f;
    float wb2 = (FULL_TILE || mr < tb) ? __expf(sb2 - nm2) : 0.f;
    float wb3 = (FULL_TILE || mr < tb) ? __expf(sb3 - nm3) : 0.f;

    float dl0 = wa0 + wb0, dl1 = wa1 + wb1;
    float dl2 = wa2 + wb2, dl3 = wa3 + wb3;
    dl0 += SWIZZLE_XOR(dl0, 8); dl1 += SWIZZLE_XOR(dl1, 8);
    dl2 += SWIZZLE_XOR(dl2, 8); dl3 += SWIZZLE_XOR(dl3, 8);
    dl0 += SWIZZLE_XOR(dl0, 4); dl1 += SWIZZLE_XOR(dl1, 4);
    dl2 += SWIZZLE_XOR(dl2, 4); dl3 += SWIZZLE_XOR(dl3, 4);
    dl0 += SWIZZLE_XOR(dl0, 2); dl1 += SWIZZLE_XOR(dl1, 2);
    dl2 += SWIZZLE_XOR(dl2, 2); dl3 += SWIZZLE_XOR(dl3, 2);
    dl0 += SWIZZLE_XOR(dl0, 1); dl1 += SWIZZLE_XOR(dl1, 1);
    dl2 += SWIZZLE_XOR(dl2, 1); dl3 += SWIZZLE_XOR(dl3, 1);
    lv[0] = lv[0] * rc0 + dl0; lv[1] = lv[1] * rc1 + dl1;
    lv[2] = lv[2] * rc2 + dl2; lv[3] = lv[3] * rc3 + dl3;

    // ---- W to per-warp LDS ----
    s_W[warp_id][kg * 4    ][mr]      = wa0;
    s_W[warp_id][kg * 4 + 1][mr]      = wa1;
    s_W[warp_id][kg * 4 + 2][mr]      = wa2;
    s_W[warp_id][kg * 4 + 3][mr]      = wa3;
    s_W[warp_id][kg * 4    ][16 + mr] = wb0;
    s_W[warp_id][kg * 4 + 1][16 + mr] = wb1;
    s_W[warp_id][kg * 4 + 2][16 + mr] = wb2;
    s_W[warp_id][kg * 4 + 3][16 + mr] = wb3;

    asm volatile("s_waitcnt lgkmcnt(0)" ::: "memory");

    float wvals[8];
    #pragma unroll
    for (int i = 0; i < 8; i++)
        wvals[i] = s_W[warp_id][mr][i * 4 + kg];

    unsigned int wlo = __builtin_amdgcn_cvt_pk_fp8_f32(wvals[0], wvals[1], 0, false);
    wlo = __builtin_amdgcn_cvt_pk_fp8_f32(wvals[2], wvals[3], wlo, true);
    unsigned int whi = __builtin_amdgcn_cvt_pk_fp8_f32(wvals[4], wvals[5], 0, false);
    whi = __builtin_amdgcn_cvt_pk_fp8_f32(wvals[6], wvals[7], whi, true);
    long w_a = static_cast<long>(wlo) | (static_cast<long>(whi) << 32);

    // ---- SV MFMA from current LDS (each warp: 128 V dims) ----
    __builtin_amdgcn_s_setprio(1);
    const int v_warp_base = warp_id * SV_CHUNKS * 16;
    const unsigned int v_byte_idx = mr & 3;
    const unsigned int perm_extract = v_byte_idx | ((v_byte_idx + 4) << 8) | 0x80800000u;
    const unsigned int perm_merge = 0x05040100u;

    #pragma unroll
    for (int vc = 0; vc < SV_CHUNKS; vc += 2) {
        const int vd0 = v_warp_base + vc * 16 + mr;
        const int vd0_align = vd0 & ~3;

        unsigned int blo0 = 0, bhi0 = 0;
        unsigned int blo1 = 0, bhi1 = 0;
        if (vd0 < V_DIM) {
            unsigned int d0[8], d1[8];
            #pragma unroll
            for (int i = 0; i < 8; i++) {
                int tok = i * 4 + kg;
                if (FULL_TILE || tok < stcnt) {
                    const unsigned int* base = reinterpret_cast<const unsigned int*>(
                        &kv_cur[tok * QK_DIM + vd0_align]);
                    d0[i] = base[0];
                    d1[i] = base[4];
                } else {
                    d0[i] = 0u;
                    d1[i] = 0u;
                }
            }
            __builtin_amdgcn_sched_group_barrier(0x0040, 16, 0);
            __builtin_amdgcn_sched_group_barrier(0x0002, 16, 0);
            unsigned int p0, p1;
            p0 = __builtin_amdgcn_perm(d0[1], d0[0], perm_extract);
            p1 = __builtin_amdgcn_perm(d0[3], d0[2], perm_extract);
            blo0 = __builtin_amdgcn_perm(p1, p0, perm_merge);
            p0 = __builtin_amdgcn_perm(d0[5], d0[4], perm_extract);
            p1 = __builtin_amdgcn_perm(d0[7], d0[6], perm_extract);
            bhi0 = __builtin_amdgcn_perm(p1, p0, perm_merge);
            p0 = __builtin_amdgcn_perm(d1[1], d1[0], perm_extract);
            p1 = __builtin_amdgcn_perm(d1[3], d1[2], perm_extract);
            blo1 = __builtin_amdgcn_perm(p1, p0, perm_merge);
            p0 = __builtin_amdgcn_perm(d1[5], d1[4], perm_extract);
            p1 = __builtin_amdgcn_perm(d1[7], d1[6], perm_extract);
            bhi1 = __builtin_amdgcn_perm(p1, p0, perm_merge);
        }
        long v_b0 = static_cast<long>(blo0) | (static_cast<long>(bhi0) << 32);
        long v_b1 = static_cast<long>(blo1) | (static_cast<long>(bhi1) << 32);

        v4f32 sc0 = {vacc[vc][0]*rc0, vacc[vc][1]*rc1, vacc[vc][2]*rc2, vacc[vc][3]*rc3};
        v4f32 sc1 = {vacc[vc+1][0]*rc0, vacc[vc+1][1]*rc1, vacc[vc+1][2]*rc2, vacc[vc+1][3]*rc3};
        sc0 = __builtin_amdgcn_mfma_f32_16x16x32_fp8_fp8(w_a, v_b0, sc0, 0, 0, 0);
        sc1 = __builtin_amdgcn_mfma_f32_16x16x32_fp8_fp8(w_a, v_b1, sc1, 0, 0, 0);
        vacc[vc][0] = sc0[0]; vacc[vc][1] = sc0[1];
        vacc[vc][2] = sc0[2]; vacc[vc][3] = sc0[3];
        vacc[vc+1][0] = sc1[0]; vacc[vc+1][1] = sc1[1];
        vacc[vc+1][2] = sc1[2]; vacc[vc+1][3] = sc1[3];
    }

    // ---- Wait for GLOBAL_LOAD_LDS and flip ----
    __builtin_amdgcn_s_setprio(3);
    if (has_next) {
        asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
        __builtin_amdgcn_sched_barrier(0);
        __builtin_amdgcn_s_barrier();
        cur_buf = nxt_buf;
    }
}

// =========================================================================
// Full MFMA pipeline kernel with double-buffered LDS + K=128 QK MFMA
// =========================================================================

__global__ __launch_bounds__(256, 3)
void mla_mfma_pipeline_kernel(
    const unsigned short* __restrict__ q_ptr,
    const unsigned char*  __restrict__ kv_ptr,
    float*                __restrict__ partial_m,
    float*                __restrict__ partial_l,
    float*                __restrict__ partial_acc,
    unsigned short*       __restrict__ out_ptr,
    const int*            __restrict__ qo_indptr,
    const int*            __restrict__ kv_indptr,
    const float*          __restrict__ kv_scale_ptr,
    const int  num_splits,
    const float sm_scale,
    const int grid_batch_first)
{
    const int batch_idx = grid_batch_first ? blockIdx.x : blockIdx.y;
    const int split_idx = grid_batch_first ? blockIdx.y : blockIdx.x;
    const int warp_id = threadIdx.x / WARP_SIZE;
    const int lane_id = threadIdx.x % WARP_SIZE;
    const int tid = threadIdx.x;

    const float score_scale = sm_scale * (*kv_scale_ptr);

    const int q_start = qo_indptr[batch_idx];
    const int kv_start = kv_indptr[batch_idx];
    const int kv_end = kv_indptr[batch_idx + 1];
    const int kv_len = kv_end - kv_start;

    const int tps = (kv_len + num_splits - 1) / num_splits;
    const int split_kv_start = kv_start + split_idx * tps;
    const int split_kv_end = min(split_kv_start + tps, kv_end);

    const int mr = lane_id & 0xF;
    const int kg = lane_id >> 4;

    if (split_kv_start >= kv_end) {
        if (lane_id < 16 && warp_id == 0) {
            int head = lane_id;
            int off = (batch_idx * NUM_HEADS + head) * num_splits + split_idx;
            partial_m[off] = -1e30f;
            partial_l[off] = 0.0f;
        }
        return;
    }

    // ===== LDS: double-buffered KV + per-warp W =====
    __shared__ __align__(16) unsigned char kv_lds[2][KV_TILE_BYTES];
    __shared__ float s_W[NUM_WARPS][16][33];

    // ===== Super-tile iteration setup =====
    const int total_tokens = split_kv_end - split_kv_start;
    const int num_st = (total_tokens + SUPER_TILE - 1) / SUPER_TILE;

    // ===== PROLOGUE: issue DMA FIRST, then Q prep overlaps with DMA =====
    {
        const int first_bytes = min(SUPER_TILE, total_tokens) * QK_DIM;
        const unsigned char* __restrict__ src0 = kv_ptr +
            static_cast<long long>(split_kv_start) * QK_DIM;
        i32x4 srsrc = make_buffer_rsrc(src0, first_bytes);
        #pragma unroll
        for (int r = 0; r < PF_ROUNDS; r++) {
            int off = tid * 16 + r * BLOCK_SIZE * 16;
            lds_ptr_t ldp = (lds_ptr_t)(reinterpret_cast<uintptr_t>(kv_lds[0]) + off);
            __llvm_amdgcn_raw_buffer_load_lds(srsrc, ldp, 16, off, 0, 0, 2);
        }
    }

    // ===== Q preload as FP8 (overlapped with DMA in flight) =====
    const unsigned short* qh = q_ptr +
        (static_cast<long long>(q_start) * NUM_HEADS + mr) * QK_DIM;

    i32x8 q_128[NUM_K128_CHUNKS];
    #pragma unroll
    for (int c = 0; c < NUM_K128_CHUNKS; c++) {
        unsigned int w[8];
        int base1 = c * 128 + 16 * kg;
        #pragma unroll
        for (int i = 0; i < 4; i++) {
            int d = base1 + i * 4;
            float f0 = (d     < QK_DIM) ? bf16_to_f32(qh[d])     : 0.f;
            float f1 = (d + 1 < QK_DIM) ? bf16_to_f32(qh[d + 1]) : 0.f;
            float f2 = (d + 2 < QK_DIM) ? bf16_to_f32(qh[d + 2]) : 0.f;
            float f3 = (d + 3 < QK_DIM) ? bf16_to_f32(qh[d + 3]) : 0.f;
            unsigned int pk = __builtin_amdgcn_cvt_pk_fp8_f32(f0, f1, 0, false);
            pk = __builtin_amdgcn_cvt_pk_fp8_f32(f2, f3, pk, true);
            w[i] = pk;
        }
        int base2 = c * 128 + 64 + 16 * kg;
        #pragma unroll
        for (int i = 0; i < 4; i++) {
            int d = base2 + i * 4;
            float f0 = (d     < QK_DIM) ? bf16_to_f32(qh[d])     : 0.f;
            float f1 = (d + 1 < QK_DIM) ? bf16_to_f32(qh[d + 1]) : 0.f;
            float f2 = (d + 2 < QK_DIM) ? bf16_to_f32(qh[d + 2]) : 0.f;
            float f3 = (d + 3 < QK_DIM) ? bf16_to_f32(qh[d + 3]) : 0.f;
            unsigned int pk = __builtin_amdgcn_cvt_pk_fp8_f32(f0, f1, 0, false);
            pk = __builtin_amdgcn_cvt_pk_fp8_f32(f2, f3, pk, true);
            w[4 + i] = pk;
        }
        q_128[c] = *reinterpret_cast<i32x8*>(w);
    }

    // ===== V accumulators + softmax state =====
    float vacc[SV_CHUNKS][4];
    #pragma unroll
    for (int i = 0; i < SV_CHUNKS; i++)
        vacc[i][0] = vacc[i][1] = vacc[i][2] = vacc[i][3] = 0.0f;

    float mv[4] = {-1e30f, -1e30f, -1e30f, -1e30f};
    float lv[4] = {0.0f, 0.0f, 0.0f, 0.0f};

    // ===== Wait for DMA (Q prep ran while DMA was in flight) =====
    asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
    __builtin_amdgcn_sched_barrier(0);
    __builtin_amdgcn_s_barrier();

    int cur_buf = 0;

    // ===== MAIN LOOP: Two-phase for compile-time optimization =====
    const int num_full = total_tokens / SUPER_TILE;
    const int has_partial = (total_tokens % SUPER_TILE) != 0;

    // Phase 1: Full tiles — stcnt=32 is compile-time constant
    for (int st_idx = 0; st_idx < num_full; st_idx++) {
        process_tile<true>(kv_ptr, kv_lds, s_W, q_128, score_scale,
            mv, lv, vacc, cur_buf, mr, kg, tid, warp_id,
            split_kv_start, split_kv_end, total_tokens, num_st,
            st_idx, SUPER_TILE);
    }

    // Phase 2: Last tile with runtime stcnt (if partial)
    if (has_partial) {
        int last_stcnt = total_tokens - num_full * SUPER_TILE;
        process_tile<false>(kv_ptr, kv_lds, s_W, q_128, score_scale,
            mv, lv, vacc, cur_buf, mr, kg, tid, warp_id,
            split_kv_start, split_kv_end, total_tokens, num_st,
            num_full, last_stcnt);
    }

    if (num_splits == 1) {
        const float kv_scale = *kv_scale_ptr;
        const int vwb = warp_id * SV_CHUNKS * 16;
        #pragma unroll
        for (int vc = 0; vc < SV_CHUNKS; vc++) {
            int vd = vwb + vc * 16 + mr;
            if (vd < V_DIM) {
                #pragma unroll
                for (int r = 0; r < 4; r++) {
                    int head = kg * 4 + r;
                    float inv_l = (lv[r] > 0.f) ? (kv_scale / lv[r]) : 0.f;
                    long long idx = (static_cast<long long>(q_start) * NUM_HEADS + head) * V_DIM + vd;
                    out_ptr[idx] = f32_to_bf16(vacc[vc][r] * inv_l);
                }
            }
        }
    } else {
        if (warp_id == 0 && mr == 0) {
            #pragma unroll
            for (int r = 0; r < 4; r++) {
                int head = kg * 4 + r;
                int off = (batch_idx * NUM_HEADS + head) * num_splits + split_idx;
                partial_m[off] = mv[r];
                partial_l[off] = lv[r];
            }
        }
        const int vwb = warp_id * SV_CHUNKS * 16;
        #pragma unroll
        for (int vc = 0; vc < SV_CHUNKS; vc++) {
            int vd = vwb + vc * 16 + mr;
            if (vd < V_DIM) {
                #pragma unroll
                for (int r = 0; r < 4; r++) {
                    int head = kg * 4 + r;
                    int off = (batch_idx * NUM_HEADS + head) * num_splits + split_idx;
                    partial_acc[static_cast<long long>(off) * V_DIM + vd] = vacc[vc][r];
                }
            }
        }
    }
}

// =========================================================================
// Reduce kernel (template-specialized for compile-time loop unrolling)
// =========================================================================

__global__ __launch_bounds__(512)
void mla_reduce_kernel_generic(
    const float*          __restrict__ partial_m,
    const float*          __restrict__ partial_l,
    const float*          __restrict__ partial_acc,
    unsigned short*       __restrict__ out_ptr,
    const float*          __restrict__ kv_scale_ptr,
    const int  num_splits)
{
    const int item_idx = blockIdx.x;
    const int tid = threadIdx.x;
    const float kv_scale = *kv_scale_ptr;

    __shared__ float s_corr[128];
    __shared__ float s_inv_l;

    const int base = item_idx * num_splits;

    float my_m = -1e30f;
    float my_l = 0.0f;
    if (tid < num_splits) {
        my_m = partial_m[base + tid];
        my_l = partial_l[base + tid];
    }

    float merged_m = my_m;
    #pragma unroll
    for (int off = 32; off >= 1; off >>= 1)
        merged_m = fmaxf(merged_m, __shfl_xor(merged_m, off));

    float my_c = 0.0f;
    if (tid < num_splits && my_l > 0.f)
        my_c = __expf(my_m - merged_m);
    float weighted_l = my_l * my_c;

    if (tid < num_splits)
        s_corr[tid] = my_c;

    float merged_l = weighted_l;
    #pragma unroll
    for (int off = 32; off >= 1; off >>= 1)
        merged_l += __shfl_xor(merged_l, off);

    if (tid == 0)
        s_inv_l = (merged_l > 0.f) ? (kv_scale / merged_l) : 0.f;
    __syncthreads();

    if (tid < V_DIM) {
        float val = 0.0f;
        for (int s = 0; s < num_splits; ++s) {
            val += partial_acc[(static_cast<long long>(base + s)) * V_DIM + tid]
                   * s_corr[s];
        }
        out_ptr[static_cast<long long>(item_idx) * V_DIM + tid] = f32_to_bf16(val * s_inv_l);
    }
}

template<int NUM_SPLITS>
__global__ __launch_bounds__(512)
void mla_reduce_kernel(
    const float*          __restrict__ partial_m,
    const float*          __restrict__ partial_l,
    const float*          __restrict__ partial_acc,
    unsigned short*       __restrict__ out_ptr,
    const float*          __restrict__ kv_scale_ptr)
{
    const int item_idx = blockIdx.x;
    const int tid = threadIdx.x;
    const float kv_scale = *kv_scale_ptr;

    __shared__ float s_corr[NUM_SPLITS < 128 ? 128 : NUM_SPLITS];
    __shared__ float s_inv_l;

    const int base = item_idx * NUM_SPLITS;

    float my_m = -1e30f;
    float my_l = 0.0f;
    if (tid < NUM_SPLITS) {
        my_m = partial_m[base + tid];
        my_l = partial_l[base + tid];
    }

    float merged_m = my_m;
    #pragma unroll
    for (int off = 32; off >= 1; off >>= 1)
        merged_m = fmaxf(merged_m, __shfl_xor(merged_m, off));

    float my_c = 0.0f;
    if (tid < NUM_SPLITS && my_l > 0.f)
        my_c = __expf(my_m - merged_m);
    float weighted_l = my_l * my_c;

    if (tid < NUM_SPLITS)
        s_corr[tid] = my_c;

    float merged_l = weighted_l;
    #pragma unroll
    for (int off = 32; off >= 1; off >>= 1)
        merged_l += __shfl_xor(merged_l, off);

    if (tid == 0)
        s_inv_l = (merged_l > 0.f) ? (kv_scale / merged_l) : 0.f;
    __syncthreads();

    if (tid < V_DIM) {
        float val = 0.0f;
        #pragma unroll
        for (int s = 0; s < NUM_SPLITS; ++s) {
            val += partial_acc[(static_cast<long long>(base + s)) * V_DIM + tid]
                   * s_corr[s];
        }
        out_ptr[static_cast<long long>(item_idx) * V_DIM + tid] = f32_to_bf16(val * s_inv_l);
    }
}

torch::Tensor mla_decode(
    torch::Tensor q, torch::Tensor kv_buffer,
    torch::Tensor qo_indptr, torch::Tensor kv_indptr,
    torch::Tensor kv_scale_tensor,
    int64_t num_heads, int64_t num_splits,
    float sm_scale,
    torch::Tensor partial_m, torch::Tensor partial_l,
    torch::Tensor partial_acc,
    torch::Tensor output)
{
    const int batch_size = qo_indptr.size(0) - 1;
    const int num_items = batch_size * static_cast<int>(num_heads);

    const int kv_len_est = (kv_indptr.data_ptr<int>()[1] - kv_indptr.data_ptr<int>()[0]);
    const int gbf = (kv_len_est > 2048) ? 1 : 0;
    dim3 grid1 = gbf ? dim3(batch_size, static_cast<int>(num_splits))
                      : dim3(static_cast<int>(num_splits), batch_size);
    dim3 block1(BLOCK_SIZE);
    mla_mfma_pipeline_kernel<<<grid1, block1>>>(
        reinterpret_cast<const unsigned short*>(q.data_ptr()),
        reinterpret_cast<const unsigned char*>(kv_buffer.data_ptr()),
        partial_m.data_ptr<float>(), partial_l.data_ptr<float>(),
        partial_acc.data_ptr<float>(),
        reinterpret_cast<unsigned short*>(output.data_ptr()),
        qo_indptr.data_ptr<int>(), kv_indptr.data_ptr<int>(),
        kv_scale_tensor.data_ptr<float>(),
        static_cast<int>(num_splits), sm_scale, gbf);
    if (num_splits == 1) return output;

    dim3 grid2(num_items);
    dim3 block2(512);

    #define REDUCE_DISPATCH(N) \
        mla_reduce_kernel<N><<<grid2, block2>>>( \
            partial_m.data_ptr<float>(), partial_l.data_ptr<float>(), \
            partial_acc.data_ptr<float>(), \
            reinterpret_cast<unsigned short*>(output.data_ptr()), \
            kv_scale_tensor.data_ptr<float>())

    switch (static_cast<int>(num_splits)) {
        case 3:  REDUCE_DISPATCH(3);  break;
        case 8:  REDUCE_DISPATCH(8);  break;
        case 12: REDUCE_DISPATCH(12); break;
        case 16: REDUCE_DISPATCH(16); break;
        case 24: REDUCE_DISPATCH(24); break;
        case 64: REDUCE_DISPATCH(64); break;
        default:
            mla_reduce_kernel_generic<<<grid2, block2>>>(
                partial_m.data_ptr<float>(), partial_l.data_ptr<float>(),
                partial_acc.data_ptr<float>(),
                reinterpret_cast<unsigned short*>(output.data_ptr()),
                kv_scale_tensor.data_ptr<float>(),
                static_cast<int>(num_splits));
            break;
    }
    #undef REDUCE_DISPATCH

    return output;
}
"""

CPP_DECL = """
torch::Tensor mla_decode(
    torch::Tensor q, torch::Tensor kv_buffer,
    torch::Tensor qo_indptr, torch::Tensor kv_indptr,
    torch::Tensor kv_scale_tensor,
    int64_t num_heads, int64_t num_splits,
    float sm_scale,
    torch::Tensor partial_m, torch::Tensor partial_l,
    torch::Tensor partial_acc,
    torch::Tensor output);
"""

_module = load_inline(
    name="mla_hip_v352_perm_extract",
    cpp_sources=CPP_DECL,
    cuda_sources=HIP_SRC,
    functions=["mla_decode"],
    extra_cuda_cflags=[
        "-O3", "-std=c++17",
        "-ffast-math", "-funsafe-math-optimizations", "-ffp-contract=fast",
        "-fno-gpu-rdc",
        "-mllvm", "-amdgpu-early-inline-all=true",
        "-mllvm", "-amdgpu-function-calls=false",
        "-mllvm", "-amdgpu-max-memory-clause=64",
        "-mllvm", "-amdgpu-load-store-vectorizer",
        "-mllvm", "-amdgpu-early-ifcvt",
        "-mllvm", "-amdgpu-internalize-symbols",
        "-mllvm", "-amdgpu-scalarize-global-loads",
        "-mllvm", "-amdgpu-dpp-combine",
        "-mllvm", "-amdgpu-enable-pre-ra-optimizations",
        "-mllvm", "-amdgpu-promote-alloca-to-vector-limit=256",
    ],
    verbose=False,
)



def _choose_splits(batch_size, kv_len):
    tiles = max(1, kv_len // 32)
    max_useful = max(1, tiles // 2)

    if tiles > 64:
        ideal = max(1, -(-768 // batch_size))
    else:
        target_wgs = max(512, batch_size * 8)
        ideal = max(1, target_wgs // batch_size)

    splits = max(1, min(ideal, max_useful, 64))

    while splits > 1 and batch_size * splits > 912:
        splits -= 1

    return splits


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data

    kv_buffer_fp8, kv_scale = kv_data["fp8"]
    kv_buffer = kv_buffer_fp8.view(-1, 576)

    batch_size = config["batch_size"]
    num_heads = config["num_heads"]
    sm_scale = config["sm_scale"]
    total_q = q.size(0)
    num_items = batch_size * num_heads

    total_kv = kv_buffer.shape[0]
    kv_len = total_kv // batch_size

    num_splits = _choose_splits(batch_size, kv_len)

    dev = q.device
    pm = torch.empty((num_items * num_splits,), dtype=torch.float32, device=dev)
    pl = torch.empty((num_items * num_splits,), dtype=torch.float32, device=dev)
    pa = torch.empty((num_items * num_splits, 512), dtype=torch.float32, device=dev)
    out = torch.empty((total_q, num_heads, 512), dtype=torch.bfloat16, device=dev)

    return _module.mla_decode(
        q, kv_buffer, qo_indptr, kv_indptr,
        kv_scale,
        num_heads, num_splits,
        sm_scale,
        pm, pl, pa, out)
scrolls · 767 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