Skip to content
KernelIndex
Search⌘K

submission 734135

willfisher · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

harry.jpeg.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-734135?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
34.3µs
#55 of 766
2026-04-05

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:129b13052addc70370781c18c39a00acb5ee4bb38dc64fd6eca044d68cbd9c0d
license declaredunknown
license concludedunknown
authorswillfisher
imported2026-08-15

Techniques

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

fp4kv_buffer, kv_scales = kv_data["mxfp4"]
num-warps = 8constexpr int NUM_WARPS = 8;
online-softmaxfloat running_max[HEADS_PER_LANE], running_sum[HEADS_PER_LANE];
shared-memoryextern __shared__ char smem[];
vector-width = float2float2 f01 = __bfloat1622float2(bp[0]);

Kernel source

harry.jpeg.py770 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X

import os
if "PYTORCH_ROCM_ARCH" not in os.environ:
    os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"

import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t

NUM_HEADS = 16
KV_LORA_RANK = 512
QK_ROPE_HEAD_DIM = 64
QK_HEAD_DIM = KV_LORA_RANK + QK_ROPE_HEAD_DIM
V_HEAD_DIM = KV_LORA_RANK
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)

CUDA_SRC = r"""
#include <hip/hip_runtime.h>
#include <hip/amd_detail/amd_hip_bf16.h>
#include <pybind11/pybind11.h>

constexpr int NUM_HEADS      = 16;
constexpr int QK_HEAD_DIM    = 576;
constexpr int V_HEAD_DIM     = 512;
constexpr int KV_PACKED      = 288;
constexpr int SCALE_COLS     = 18;
constexpr int WAVEFRONT_SIZE = 64;
constexpr int NUM_WARPS      = 8;
constexpr int BLOCK_THREADS  = NUM_WARPS * WAVEFRONT_SIZE;
constexpr int HEADS_PER_LANE = 4;
constexpr float LOG2E        = 1.4426950408889634f;
constexpr int FP4_K          = 128;                           // dims per K chunk
constexpr int FP4_K_CHUNKS   = 5;                             // ceil(576/128)
constexpr int V_K_CHUNKS     = 4;                             // 512/128 — chunks that contain V data
constexpr int TILE_TOKENS    = 128;                           // tokens per tile

// LDS layout: KV double-buf | V_fp8/Q reuse | score | attn_fp8 | stats
// KV loaded via coalesced DMA → LDS, then each thread reads to registers
constexpr int BUF_KV_BYTES    = TILE_TOKENS * KV_PACKED;         // 36864
constexpr int BUF_SC_BYTES    = TILE_TOKENS * SCALE_COLS;        // 2304
constexpr int BUF_TOTAL       = BUF_KV_BYTES + BUF_SC_BYTES;    // 39168
constexpr int V_FP8_STRIDE    = V_HEAD_DIM + 8;                  // 520: 8-byte aligned, breaks bank conflicts (520/4%64=2)
constexpr int V_FP8_BYTES     = TILE_TOKENS * V_FP8_STRIDE;      // 66560
constexpr int SCORE_STRIDE    = NUM_HEADS + 2;                   // 18 bf16/row: pad for bank conflicts (gcd(9,64)=1)
constexpr int SCORE_BYTES     = TILE_TOKENS * SCORE_STRIDE * 2;  // 4608
constexpr int ATTN_FP8_BYTES  = TILE_TOKENS * NUM_HEADS + 136;   // 2184 (fp8 attn + non-linear bank shift)
constexpr int STATS_BYTES     = NUM_HEADS * 2 * 4;               // 128

// Q reuses the V_fp8 area (Q written in prologue, cached to regs, then V_fp8 overwrites)
constexpr int OFF_KV0         = 0;
constexpr int OFF_KV1         = BUF_TOTAL;                       // 39168
constexpr int OFF_VFP8        = BUF_TOTAL * 2;                   // 78336 (also OFF_Q during prologue)
constexpr int OFF_Q           = OFF_VFP8;                        // reuses V_fp8 area
constexpr int OFF_QSC         = OFF_Q + NUM_HEADS * KV_PACKED;   // 82944
constexpr int OFF_SCORE       = OFF_VFP8 + V_FP8_BYTES;          // 143872
constexpr int OFF_ATTN        = OFF_SCORE + SCORE_BYTES;          // 147968
constexpr int OFF_STATS       = OFF_ATTN + ATTN_FP8_BYTES;       // 150016
constexpr int TOTAL_LDS       = OFF_STATS + STATS_BYTES;          // 150144 (~147KB)

typedef __bf16 bf16x2 __attribute__((__vector_size__(2 * sizeof(__bf16))));
typedef float f32x4 __attribute__((__vector_size__(4 * sizeof(float))));
typedef int v4i32 __attribute__((__vector_size__(4 * sizeof(int))));
typedef int v2i32 __attribute__((__vector_size__(2 * sizeof(int))));
typedef int v8i32 __attribute__((__vector_size__(8 * sizeof(int))));
typedef uint32_t u32x4_vec __attribute__((__vector_size__(4 * sizeof(uint32_t))));
typedef __bf16 bf16x4 __attribute__((__vector_size__(4 * sizeof(__bf16))));

__device__ __forceinline__ v8i32 to_v8(v4i32 x) {
    v8i32 r = {x[0], x[1], x[2], x[3], 0, 0, 0, 0};
    return r;
}

// Fused DPP mirror butterfly: single-instruction v_max/v_add with DPP modifier
#define DPP_MAX(m, ctrl) asm volatile( \
    "v_max_f32_dpp %0, %0, %0 " ctrl " row_mask:0xf bank_mask:0xf" : "+v"(m))
#define DPP_ADD(s, ctrl) asm volatile( \
    "v_add_f32_dpp %0, %0, %0 " ctrl " row_mask:0xf bank_mask:0xf" : "+v"(s))

__device__ __forceinline__ float dpp_row_max(float m) {
    DPP_MAX(m, "quad_perm:[1,0,3,2]");
    DPP_MAX(m, "quad_perm:[2,3,0,1]");
    DPP_MAX(m, "row_half_mirror");
    DPP_MAX(m, "row_mirror");
    return m;
}

__device__ __forceinline__ float dpp_row_sum(float s) {
    DPP_ADD(s, "quad_perm:[1,0,3,2]");
    DPP_ADD(s, "quad_perm:[2,3,0,1]");
    DPP_ADD(s, "row_half_mirror");
    DPP_ADD(s, "row_mirror");
    return s;
}

using i32x4 = int32_t __attribute__((ext_vector_type(4)));
using i32x2 = int32_t __attribute__((ext_vector_type(2)));
using as3_uint32_ptr = uint32_t __attribute__((address_space(3)))*;
using as3_bf16x4_ptr = bf16x4 __attribute__((address_space(3)))*;
using as3_i32x2_ptr = i32x2 __attribute__((address_space(3)))*;
using as3_v2i32_ptr = v2i32 __attribute__((address_space(3)))*;
typedef __bf16 bf16x8 __attribute__((__vector_size__(8 * sizeof(__bf16))));
typedef short ev_short2 __attribute__((ext_vector_type(2)));

extern "C" __device__ __attribute__((const)) __bf16 llvm_exp2_bf16(__bf16) __asm("llvm.exp2.bf16");
extern "C" __device__ void llvm_amdgcn_raw_buffer_load_lds(
    i32x4 rsrc, as3_uint32_ptr lds_ptr, int size, int voffset, int soffset, int offset, int aux)
    __asm("llvm.amdgcn.raw.buffer.load.lds");

struct buffer_resource { uint64_t ptr; uint32_t range; uint32_t config; };
__device__ __forceinline__ i32x4 make_srsrc(const void* ptr, uint32_t range_bytes) {
    buffer_resource rsrc = {reinterpret_cast<uint64_t>(ptr), range_bytes, 0x110000};
    return *reinterpret_cast<const i32x4*>(&rsrc);
}

// ============================================================
// Grouped kernel: double-buffered KV LDS, per-warp V dequant
// ============================================================
template<int NUM_PARTIALS, int TILES_PER_BLOCK>
__global__ void __launch_bounds__(BLOCK_THREADS)
mla_grouped(
    const __hip_bfloat16* __restrict__ q,
    const uint8_t* __restrict__ kv_data,
    const uint8_t* __restrict__ kv_scales,
    const int32_t* __restrict__ qo_indptr,
    const int32_t* __restrict__ kv_indptr,
    __hip_bfloat16* __restrict__ partial_out,
    float* __restrict__ partial_max,
    float* __restrict__ partial_sum,
    __hip_bfloat16* __restrict__ final_output,
    float sm_scale
) {
    const int batch_idx   = blockIdx.x / NUM_PARTIALS;
    const int partial_idx = blockIdx.x - batch_idx * NUM_PARTIALS;
    const int tile_base   = partial_idx * TILES_PER_BLOCK;
    const int warp_id     = threadIdx.x >> 6;
    const int lane_id     = threadIdx.x & 63;
    const int lane_mod16  = lane_id & 15;
    const int k_group     = lane_id >> 4;
    const int tok         = warp_id * 16 + lane_mod16;

    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];

    extern __shared__ char smem[];
    uint8_t* q_lds       = reinterpret_cast<uint8_t*>(smem + OFF_Q);
    uint8_t* qsc_lds     = reinterpret_cast<uint8_t*>(smem + OFF_QSC);
    uint8_t* v_fp8_lds   = reinterpret_cast<uint8_t*>(smem + OFF_VFP8);
    uint8_t* attn_fp8_lds = reinterpret_cast<uint8_t*>(smem + OFF_ATTN);
    float*   stats_lds   = reinterpret_cast<float*>(smem + OFF_STATS);

    // Helper: coalesced DMA load of full KV tile into LDS buffer
    auto load_tile = [&](int hbm_start, int buf) {
        uint8_t* dst_kv = reinterpret_cast<uint8_t*>(smem + (buf == 0 ? OFF_KV0 : OFF_KV1));
        uint8_t* dst_sc = dst_kv + BUF_KV_BYTES;
        constexpr int KV_LOADS = (BUF_KV_BYTES + 15) / 16;
        i32x4 kv_srsrc = make_srsrc(kv_data + hbm_start * KV_PACKED, BUF_KV_BYTES);
        for (int off = threadIdx.x; off < KV_LOADS; off += BLOCK_THREADS) {
            int voff = off * 16;
            llvm_amdgcn_raw_buffer_load_lds(kv_srsrc,
                reinterpret_cast<as3_uint32_ptr>(reinterpret_cast<uintptr_t>(dst_kv) + voff),
                16, voff, 0, 0, 0);
        }
        constexpr int SC_LOADS = (BUF_SC_BYTES + 15) / 16;
        i32x4 sc_srsrc = make_srsrc(kv_scales + hbm_start * SCALE_COLS, BUF_SC_BYTES);
        for (int off = threadIdx.x; off < SC_LOADS; off += BLOCK_THREADS) {
            int voff = off * 16;
            llvm_amdgcn_raw_buffer_load_lds(sc_srsrc,
                reinterpret_cast<as3_uint32_ptr>(reinterpret_cast<uintptr_t>(dst_sc) + voff),
                16, voff, 0, 0, 0);
        }
    };

    // Issue tile 0 DMA (non-blocking, in flight during Q quantize)
    load_tile(kv_start + tile_base * TILE_TOKENS, 0);

    // ---- Q quantize (same as before) ----
    {
        const __hip_bfloat16* q_base = q + q_start * NUM_HEADS * QK_HEAD_DIM;
        const int dc = warp_id;
        if (dc < FP4_K_CHUNKS) {
            const int qds = dc * FP4_K + (k_group << 5);
            if (qds < QK_HEAD_DIM) {
                __hip_bfloat16 qv[32];
                const __hip_bfloat16* qp = q_base + lane_mod16 * QK_HEAD_DIM + qds;
                #pragma unroll
                for (int i = 0; i < 4; i++)
                    *reinterpret_cast<u32x4_vec*>(&qv[i*8]) = *reinterpret_cast<const u32x4_vec*>(qp + i*8);
                uint16_t mx = 0;
                #pragma unroll
                for (int i = 0; i < 32; i++) { uint16_t b = *reinterpret_cast<const uint16_t*>(&qv[i]) & 0x7FFF; mx = max(mx, b); }
                uint8_t e8 = 0;
                if (mx != 0) { uint32_t bits = (uint32_t)mx << 16; bits = (bits + 0x200000u) & 0xFF800000u; e8 = (uint8_t)max(0, min(254, (int)(bits >> 23) - 2)); }
                float sf = __uint_as_float((uint32_t)e8 << 23);
                v4i32 ar;
                #pragma unroll
                for (int w = 0; w < 4; w++) {
                    unsigned int pk = 0;
                    pk = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk, __hip_bfloat162{qv[w*8+0],qv[w*8+1]}, sf, 0);
                    pk = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk, __hip_bfloat162{qv[w*8+2],qv[w*8+3]}, sf, 1);
                    pk = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk, __hip_bfloat162{qv[w*8+4],qv[w*8+5]}, sf, 2);
                    pk = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk, __hip_bfloat162{qv[w*8+6],qv[w*8+7]}, sf, 3);
                    ar[w] = (int)pk;
                }
                const int qb = qds >> 5;
                *reinterpret_cast<u32x4_vec*>(q_lds + lane_mod16 * KV_PACKED + (qb << 4)) = *reinterpret_cast<const u32x4_vec*>(&ar);
                qsc_lds[lane_mod16 * SCALE_COLS + qb] = e8;
            }
        }
    }

    // Running accumulators
    f32x4 running_v[4];
    float running_max[HEADS_PER_LANE], running_sum[HEADS_PER_LANE];
    #pragma unroll
    for (int i = 0; i < 4; i++) running_v[i] = {0,0,0,0};
    #pragma unroll
    for (int h = 0; h < HEADS_PER_LANE; h++) { running_max[h] = -INFINITY; running_sum[h] = 0; }

    v8i32 q_v8[FP4_K_CHUNKS];
    int q_scale_cached[FP4_K_CHUNKS];
    const float sm_log2e = sm_scale * LOG2E;
    using f16x2 = _Float16 __attribute__((ext_vector_type(2)));

    // ============ TILE LOOP ============
    for (int ti = 0; ti < TILES_PER_BLOCK; ti++) {
        const int cur = ti & 1;
        const int tile_start = kv_start + (tile_base + ti) * TILE_TOKENS;
        const int tile_len = max(0, min(TILE_TOKENS, kv_end - tile_start));

        // Wait for KV DMA + ensure all threads' LDS writes visible
        asm volatile("s_waitcnt vmcnt(0)");
        __syncthreads();

        // Cache Q on first iteration (after sync ensures q_lds visible)
        if (ti == 0) {
            #pragma unroll
            for (int d = 0; d < FP4_K_CHUNKS; d++) {
                const int qb = d * (FP4_K / 32) + k_group;
                v4i32 qr = *reinterpret_cast<const v4i32*>(q_lds + lane_mod16 * KV_PACKED + (qb << 4));
                q_scale_cached[d] = (int)qsc_lds[lane_mod16 * SCALE_COLS + qb];
                q_v8[d] = to_v8(qr);
            }
            if (k_group >= 2) { q_v8[4] = {0,0,0,0,0,0,0,0}; q_scale_cached[4] = 0; }
        }

        const uint8_t* kv_lds = reinterpret_cast<const uint8_t*>(smem + (cur == 0 ? OFF_KV0 : OFF_KV1));
        const uint8_t* sc_lds = kv_lds + BUF_KV_BYTES;
        const int tok_kv_base = tok * KV_PACKED + (k_group << 4);
        const int tok_sc_base = tok * SCALE_COLS + k_group;

        // Prefetch next tile (overlaps with entire tile compute)
        if (ti + 1 < TILES_PER_BLOCK)
            load_tile(kv_start + (tile_base + ti + 1) * TILE_TOKENS, 1 - cur);

        // ---- Fused Score + V dequant, double-buffered LDS reads ----
        f32x4 scores = {0,0,0,0};
        const bool tok_valid = tok < tile_len;

        // Prefetch chunk 0
        v4i32 kv_next = {0,0,0,0};
        int sc_next = 0;
        if (tok_valid) {
            kv_next = *reinterpret_cast<const v4i32*>(kv_lds + tok_kv_base);
            sc_next = (int)sc_lds[tok_sc_base];
        }

        #pragma unroll
        for (int d = 0; d < FP4_K_CHUNKS; d++) {
            v4i32 kv = kv_next;
            int sc = sc_next;

            // Score MFMA first (fires on matrix core immediately)
            scores = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
                q_v8[d], to_v8(kv),
                scores, 4, 4, 0, q_scale_cached[d], 0, sc);

            // Prefetch chunk d+1 (LDS read during MFMA + dequant)
            if (d + 1 < FP4_K_CHUNKS) {
                const int next_dim = (d + 1) * (FP4_K / 2);
                if (tok_valid && next_dim + (k_group << 4) + 16 <= KV_PACKED) {
                    kv_next = *reinterpret_cast<const v4i32*>(kv_lds + tok_kv_base + next_dim);
                    sc_next = (int)sc_lds[tok_sc_base + (d + 1) * (FP4_K / 32)];
                } else {
                    kv_next = {0,0,0,0}; sc_next = 0;
                }
            }

            // V dequant: fp4 → f16 → bf8 (VALU, overlaps with MFMA on matrix core)
            if (d < V_K_CHUNKS) {
                uint8_t* bf8_dst = v_fp8_lds + tok * V_FP8_STRIDE + d * FP4_K + (k_group << 5);
                float block_sf = tok_valid ? __uint_as_float((uint32_t)sc << 23) : 0.0f;
                ev_short2 ztmp = {0, 0};

                #pragma unroll
                for (int w = 0; w < 4; w++) {
                    uint32_t pk = (uint32_t)kv[w];
                    f16x2 h0 = __builtin_amdgcn_cvt_scalef32_pk_f16_fp4(pk, block_sf, 0);
                    f16x2 h1 = __builtin_amdgcn_cvt_scalef32_pk_f16_fp4(pk, block_sf, 1);
                    f16x2 h2 = __builtin_amdgcn_cvt_scalef32_pk_f16_fp4(pk, block_sf, 2);
                    f16x2 h3 = __builtin_amdgcn_cvt_scalef32_pk_f16_fp4(pk, block_sf, 3);
                    ev_short2 b01 = __builtin_amdgcn_cvt_scalef32_pk_bf8_f16(ztmp, h0, 1.0f, false);
                    ev_short2 b0123 = __builtin_amdgcn_cvt_scalef32_pk_bf8_f16(b01, h1, 1.0f, true);
                    ev_short2 b45 = __builtin_amdgcn_cvt_scalef32_pk_bf8_f16(ztmp, h2, 1.0f, false);
                    ev_short2 b4567 = __builtin_amdgcn_cvt_scalef32_pk_bf8_f16(b45, h3, 1.0f, true);
                    *reinterpret_cast<ev_short2*>(bf8_dst + w * 8) = b0123;
                    *reinterpret_cast<ev_short2*>(bf8_dst + w * 8 + 4) = b4567;
                }
            }
        }

        // ---- Register-based softmax: no score_lds ----
        if (!tok_valid) { scores[0] = scores[1] = scores[2] = scores[3] = -INFINITY; }
        else { scores[0] *= sm_log2e; scores[1] *= sm_log2e; scores[2] *= sm_log2e; scores[3] *= sm_log2e; }

        // Warp-local max over 16 tokens (DPP mirror butterfly within k_group row)
        float warp_max[HEADS_PER_LANE];
        #pragma unroll
        for (int h = 0; h < HEADS_PER_LANE; h++)
            warp_max[h] = dpp_row_max(scores[h]);

        // Warp-local exp + sum (DPP mirror butterfly, no * LOG2E — already in log2 space)
        float attn_w[HEADS_PER_LANE], warp_sum[HEADS_PER_LANE];
        #pragma unroll
        for (int h = 0; h < HEADS_PER_LANE; h++) {
            attn_w[h] = __builtin_amdgcn_exp2f(scores[h] - warp_max[h]);
            warp_sum[h] = dpp_row_sum(attn_w[h]);
        }

        // Vectorized LDS write: 4 heads at once (float4)
        float* scratch = reinterpret_cast<float*>(smem + OFF_SCORE);
        if (lane_mod16 == 0) {
            *reinterpret_cast<f32x4*>(&scratch[warp_id * NUM_HEADS + (k_group << 2)]) =
                (f32x4){warp_max[0], warp_max[1], warp_max[2], warp_max[3]};
            *reinterpret_cast<f32x4*>(&scratch[NUM_WARPS * NUM_HEADS + warp_id * NUM_HEADS + (k_group << 2)]) =
                (f32x4){warp_sum[0], warp_sum[1], warp_sum[2], warp_sum[3]};
        }
        __syncthreads();

        // Vectorized LDS reads: global max + corrected sum (float4 per warp)
        float my_tm[HEADS_PER_LANE] = {-INFINITY, -INFINITY, -INFINITY, -INFINITY};
        float my_ts[HEADS_PER_LANE] = {0, 0, 0, 0};
        #pragma unroll
        for (int w = 0; w < NUM_WARPS; w++) {
            f32x4 wm = *reinterpret_cast<const f32x4*>(&scratch[w * NUM_HEADS + (k_group << 2)]);
            my_tm[0] = fmaxf(my_tm[0], wm[0]); my_tm[1] = fmaxf(my_tm[1], wm[1]);
            my_tm[2] = fmaxf(my_tm[2], wm[2]); my_tm[3] = fmaxf(my_tm[3], wm[3]);
        }
        #pragma unroll
        for (int w = 0; w < NUM_WARPS; w++) {
            f32x4 wm = *reinterpret_cast<const f32x4*>(&scratch[w * NUM_HEADS + (k_group << 2)]);
            f32x4 ws = *reinterpret_cast<const f32x4*>(&scratch[NUM_WARPS * NUM_HEADS + w * NUM_HEADS + (k_group << 2)]);
            #pragma unroll
            for (int h = 0; h < HEADS_PER_LANE; h++)
                my_ts[h] += ws[h] * __builtin_amdgcn_exp2f(wm[h] - my_tm[h]);
        }

        // Correct attn weights with global max + write fp8 to LDS
        {
            #pragma unroll
            for (int h = 0; h < HEADS_PER_LANE; h++)
                attn_w[h] *= __builtin_amdgcn_exp2f(warp_max[h] - my_tm[h]);

            ev_short2 tmp = {0, 0};
            ev_short2 fp8_lo = __builtin_amdgcn_cvt_scalef32_pk_fp8_f32(tmp, attn_w[0], attn_w[1], 1.0f, false);
            ev_short2 fp8_all = __builtin_amdgcn_cvt_scalef32_pk_fp8_f32(fp8_lo, attn_w[2], attn_w[3], 1.0f, true);
            int pad = ((tok & 16) >> 1) + ((tok & 32) << 2);
            *reinterpret_cast<ev_short2*>(&attn_fp8_lds[tok * NUM_HEADS + pad + (k_group << 2)]) = fp8_all;
        }
        __syncthreads();

        // ---- V multiply: fp8 attn × fp8 V, using transpose reads ----
        // A = attn fp8 [16 heads × 128 tokens], B = V fp8 [16 dims × 128 tokens]
        // 4 MFMAs per warp (512 dims / 16 per MFMA / 8 warps = 4 rounds)
        f32x4 v_out[4];
        #pragma unroll
        for (int i = 0; i < 4; i++) v_out[i] = {0,0,0,0};
        {
            const int half = lane_mod16 & 1;
            const int tok_in_grp = lane_mod16 >> 1;
            const int half8 = half << 3;
            const int kg16 = k_group << 4;
            int tok_arr[4];
            tok_arr[0] = kg16 + tok_in_grp;
            tok_arr[1] = kg16 + 8 + tok_in_grp;
            tok_arr[2] = 64 + kg16 + tok_in_grp;
            tok_arr[3] = 64 + kg16 + 8 + tok_in_grp;

            // Load A (attn fp8): 4 transpose reads, reused across all rounds
            const uintptr_t attn_base = reinterpret_cast<uintptr_t>(attn_fp8_lds) + half8;
            v8i32 a_reg;
            {
                v2i32* a_parts = reinterpret_cast<v2i32*>(&a_reg);
                #pragma unroll
                for (int c = 0; c < 4; c++) {
                    int t_c = tok_arr[c];
                    int pad_c = ((t_c & 16) >> 1) + ((t_c & 32) << 2);
                    a_parts[c] = __builtin_amdgcn_ds_read_tr8_b64_v2i32(
                        reinterpret_cast<as3_v2i32_ptr>(attn_base + t_c * NUM_HEADS + pad_c));
                }
            }

            // V fp8: double-buffered B reads to hide LDS latency behind MFMA
            const uintptr_t vfp8_base = reinterpret_cast<uintptr_t>(v_fp8_lds) + half8;
            v8i32 b_reg[2];

            // Pre-load round 0
            {
                const uintptr_t vb0 = vfp8_base + ((0 * 8 + warp_id) << 4);
                v2i32* bp = reinterpret_cast<v2i32*>(&b_reg[0]);
                #pragma unroll
                for (int c = 0; c < 4; c++)
                    bp[c] = __builtin_amdgcn_ds_read_tr8_b64_v2i32(
                        reinterpret_cast<as3_v2i32_ptr>(vb0 + tok_arr[c] * V_FP8_STRIDE));
            }

            #pragma unroll
            for (int r = 0; r < 4; r++) {
                int cur = r & 1, nxt = (r + 1) & 1;

                // Prefetch next round while MFMA executes current
                if (r + 1 < 4) {
                    const uintptr_t vb_next = vfp8_base + (((r + 1) * 8 + warp_id) << 4);
                    v2i32* bp = reinterpret_cast<v2i32*>(&b_reg[nxt]);
                    #pragma unroll
                    for (int c = 0; c < 4; c++)
                        bp[c] = __builtin_amdgcn_ds_read_tr8_b64_v2i32(
                            reinterpret_cast<as3_v2i32_ptr>(vb_next + tok_arr[c] * V_FP8_STRIDE));
                }

                v_out[r] = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
                    a_reg, b_reg[cur], v_out[r], 0, 1, 0, 127, 0, 127);
            }
        }

        // ---- Merge ----
        #pragma unroll
        for (int h = 0; h < HEADS_PER_LANE; h++) {
            float nm = fmaxf(running_max[h], my_tm[h]);
            float co = __builtin_amdgcn_exp2f(running_max[h] - nm);
            float cn = __builtin_amdgcn_exp2f(my_tm[h] - nm);
            running_sum[h] = running_sum[h] * co + my_ts[h] * cn;
            running_max[h] = nm;
            #pragma unroll
            for (int i = 0; i < 4; i++)
                running_v[i][h] = running_v[i][h] * co + v_out[i][h] * cn;
        }

    }

    // ---- Write output ----
    if (NUM_PARTIALS == 1) {
        // Direct output: divide by sum, write final result (skip reduce kernel)
        __hip_bfloat16* ob = final_output + (q_start * NUM_HEADS) * V_HEAD_DIM;
        #pragma unroll
        for (int h = 0; h < HEADS_PER_LANE; h++) {
            float inv = (running_sum[h] > 0.0f) ? (1.0f / running_sum[h]) : 0.0f;
            #pragma unroll
            for (int i = 0; i < 4; i++) {
                const int vd = (i*8 + warp_id)*16 + lane_mod16;
                ob[((k_group<<2)+h)*V_HEAD_DIM + vd] = __float2bfloat16(running_v[i][h] * inv);
            }
        }
    } else {
        // Write partial for reduce kernel
        const int pidx = batch_idx * NUM_PARTIALS + partial_idx;
        __hip_bfloat16* ob = partial_out + pidx * (NUM_HEADS * V_HEAD_DIM);
        #pragma unroll
        for (int i = 0; i < 4; i++) {
            const int vd = (i*8 + warp_id)*16 + lane_mod16;
            #pragma unroll
            for (int h = 0; h < HEADS_PER_LANE; h++)
                ob[((k_group<<2)+h)*V_HEAD_DIM + vd] = __float2bfloat16(running_v[i][h]);
        }
        if (lane_mod16 == 0 && warp_id == 0) {
            float* pm = partial_max + pidx * NUM_HEADS;
            float* ps = partial_sum + pidx * NUM_HEADS;
            #pragma unroll
            for (int h = 0; h < HEADS_PER_LANE; h++) {
                pm[(k_group<<2)+h] = running_max[h];
                ps[(k_group<<2)+h] = running_sum[h];
            }
        }
    }
}

// ============================================================
// Reduce kernel
// ============================================================
template<int NUM_PARTIALS>
__global__ void __launch_bounds__(WAVEFRONT_SIZE)
mla_reduce(
    const __hip_bfloat16* __restrict__ partial_out,
    const float* __restrict__ partial_max,
    const float* __restrict__ partial_sum,
    const int32_t* __restrict__ qo_indptr,
    __hip_bfloat16* __restrict__ output
) {
    const int bid = blockIdx.x, batch_idx = bid / NUM_HEADS, head_idx = bid - batch_idx * NUM_HEADS;
    const int lane_id = threadIdx.x & 63, q_start = qo_indptr[batch_idx];
    constexpr int ELEMS = V_HEAD_DIM / WAVEFRONT_SIZE;  // 8
    constexpr int PV_STRIDE = NUM_HEADS * V_HEAD_DIM;

    typedef float ext_f32x4 __attribute__((ext_vector_type(4)));
    typedef uint32_t ext_u32x4 __attribute__((ext_vector_type(4)));

    float local_max = -INFINITY, local_sum = 0;
    ext_f32x4 v_acc0 = {0,0,0,0}, v_acc1 = {0,0,0,0};

    const int bo = batch_idx * NUM_PARTIALS;
    const float* pm = partial_max + bo * NUM_HEADS + head_idx;
    const float* ps = partial_sum + bo * NUM_HEADS + head_idx;
    const __hip_bfloat16* pv_base = partial_out + bo * NUM_HEADS * V_HEAD_DIM
        + head_idx * V_HEAD_DIM + lane_id * ELEMS;

    #pragma unroll
    for (int p = 0; p < NUM_PARTIALS; p++) {
        // Scalar loads for wave-uniform values (SALU path)
        float cm = __builtin_amdgcn_readfirstlane(pm[p * NUM_HEADS]);
        float cs = __builtin_amdgcn_readfirstlane(ps[p * NUM_HEADS]);

        const __hip_bfloat16* pv = pv_base + p * PV_STRIDE;

        // Vectorized bf16 load → native bf16x2→float2 conversion
        const __hip_bfloat162* bp = reinterpret_cast<const __hip_bfloat162*>(pv);
        float2 f01 = __bfloat1622float2(bp[0]);
        float2 f23 = __bfloat1622float2(bp[1]);
        float2 f45 = __bfloat1622float2(bp[2]);
        float2 f67 = __bfloat1622float2(bp[3]);
        ext_f32x4 v0 = {f01.x, f01.y, f23.x, f23.y};
        ext_f32x4 v1 = {f45.x, f45.y, f67.x, f67.y};

        // Wave-uniform branch (SALU)
        if (cm >= local_max) {
            float co = __builtin_amdgcn_exp2f(local_max - cm);
            local_sum = local_sum * co + cs;
            v_acc0 = v_acc0 * co + v0;
            v_acc1 = v_acc1 * co + v1;
            local_max = cm;
        } else {
            float cn = __builtin_amdgcn_exp2f(cm - local_max);
            local_sum += cs * cn;
            v_acc0 += v0 * cn;
            v_acc1 += v1 * cn;
        }
    }

    float inv = (local_sum > 0.0f) ? __builtin_amdgcn_rcpf(local_sum) : 0.0f;
    v_acc0 *= inv;
    v_acc1 *= inv;

    // Pack f32 → bf16 via v_cvt_pk_bf16_f32 (2 f32 → 1 packed bf16x2 per instruction)
    ext_u32x4 out_pack;
    asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(out_pack[0]) : "v"(v_acc0[0]), "v"(v_acc0[1]));
    asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(out_pack[1]) : "v"(v_acc0[2]), "v"(v_acc0[3]));
    asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(out_pack[2]) : "v"(v_acc1[0]), "v"(v_acc1[1]));
    asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(out_pack[3]) : "v"(v_acc1[2]), "v"(v_acc1[3]));

    *reinterpret_cast<ext_u32x4*>(output + (q_start * NUM_HEADS + head_idx) * V_HEAD_DIM + lane_id * ELEMS) = out_pack;
}

// ============================================================
// Multi-wavefront reduce for high-partial shapes (e.g. 64 partials)
// Each wavefront handles NUM_PARTIALS/NWARPS partials, then LDS merge
// ============================================================
template<int NUM_PARTIALS, int NWARPS>
__global__ void __launch_bounds__(NWARPS * WAVEFRONT_SIZE)
mla_reduce_multi(
    const __hip_bfloat16* __restrict__ partial_out,
    const float* __restrict__ partial_max,
    const float* __restrict__ partial_sum,
    const int32_t* __restrict__ qo_indptr,
    __hip_bfloat16* __restrict__ output
) {
    const int bid = blockIdx.x, batch_idx = bid / NUM_HEADS, head_idx = bid - batch_idx * NUM_HEADS;
    const int warp_id = threadIdx.x >> 6;
    const int lane_id = threadIdx.x & 63, q_start = qo_indptr[batch_idx];
    constexpr int ELEMS = V_HEAD_DIM / WAVEFRONT_SIZE;
    constexpr int PV_STRIDE = NUM_HEADS * V_HEAD_DIM;
    constexpr int PARTIALS_PER_WARP = NUM_PARTIALS / NWARPS;

    typedef float ext_f32x4 __attribute__((ext_vector_type(4)));
    typedef uint32_t ext_u32x4 __attribute__((ext_vector_type(4)));

    float local_max = -INFINITY, local_sum = 0;
    ext_f32x4 v_acc0 = {0,0,0,0}, v_acc1 = {0,0,0,0};

    const int bo = batch_idx * NUM_PARTIALS;
    const float* pm = partial_max + bo * NUM_HEADS + head_idx;
    const float* ps = partial_sum + bo * NUM_HEADS + head_idx;
    const __hip_bfloat16* pv_base = partial_out + bo * NUM_HEADS * V_HEAD_DIM
        + head_idx * V_HEAD_DIM + lane_id * ELEMS;

    // Each warp processes its slice of partials
    const int p_start = warp_id * PARTIALS_PER_WARP;
    #pragma unroll
    for (int p = p_start; p < p_start + PARTIALS_PER_WARP; p++) {
        float cm = __builtin_amdgcn_readfirstlane(pm[p * NUM_HEADS]);
        float cs = __builtin_amdgcn_readfirstlane(ps[p * NUM_HEADS]);
        const __hip_bfloat162* bp = reinterpret_cast<const __hip_bfloat162*>(pv_base + p * PV_STRIDE);
        float2 f01 = __bfloat1622float2(bp[0]);
        float2 f23 = __bfloat1622float2(bp[1]);
        float2 f45 = __bfloat1622float2(bp[2]);
        float2 f67 = __bfloat1622float2(bp[3]);
        ext_f32x4 v0 = {f01.x, f01.y, f23.x, f23.y};
        ext_f32x4 v1 = {f45.x, f45.y, f67.x, f67.y};

        if (cm >= local_max) {
            float co = __builtin_amdgcn_exp2f(local_max - cm);
            local_sum = local_sum * co + cs;
            v_acc0 = v_acc0 * co + v0;
            v_acc1 = v_acc1 * co + v1;
            local_max = cm;
        } else {
            float cn = __builtin_amdgcn_exp2f(cm - local_max);
            local_sum += cs * cn;
            v_acc0 += v0 * cn;
            v_acc1 += v1 * cn;
        }
    }

    // === Cooperative pre-scaled LDS reduction ===
    extern __shared__ char reduce_smem[];
    float* smem_max = reinterpret_cast<float*>(reduce_smem);
    float* smem_sum = smem_max + NWARPS;
    float* smem_v   = smem_sum + NWARPS;

    // Step 1: Write max/sum (predicated — only lane 0, avoids 64-way bank conflict)
    if (lane_id == 0) {
        smem_max[warp_id] = local_max;
        smem_sum[warp_id] = local_sum;
    }
    __syncthreads();

    // Step 2: ALL warps compute global max (SALU via readfirstlane)
    float global_max = -INFINITY;
    #pragma unroll
    for (int w = 0; w < NWARPS; w++)
        global_max = fmaxf(global_max, __builtin_amdgcn_readfirstlane(smem_max[w]));

    // Step 3: ALL warps pre-scale their own vectors + sum
    float scale = __builtin_amdgcn_exp2f(local_max - global_max);
    v_acc0 *= scale;
    v_acc1 *= scale;
    local_sum *= scale;

    // Step 4: Write pre-scaled vectors + sum to LDS
    float* wv_dst = &smem_v[warp_id * V_HEAD_DIM + lane_id * ELEMS];
    *reinterpret_cast<ext_f32x4*>(wv_dst) = v_acc0;
    *reinterpret_cast<ext_f32x4*>(wv_dst + 4) = v_acc1;
    if (lane_id == 0) smem_sum[warp_id] = local_sum;
    __syncthreads();

    // Step 5: Warp 0 does pure branchless vector addition (no exp2f)
    if (warp_id == 0) {
        ext_f32x4 gv0 = {0,0,0,0}, gv1 = {0,0,0,0};
        float gs = 0;
        #pragma unroll
        for (int w = 0; w < NWARPS; w++) {
            const float* wv_base = &smem_v[w * V_HEAD_DIM + lane_id * ELEMS];
            gv0 += *reinterpret_cast<const ext_f32x4*>(wv_base);
            gv1 += *reinterpret_cast<const ext_f32x4*>(wv_base + 4);
            gs += __builtin_amdgcn_readfirstlane(smem_sum[w]);
        }

        float inv = (gs > 0.0f) ? __builtin_amdgcn_rcpf(gs) : 0.0f;
        gv0 *= inv;
        gv1 *= inv;

        ext_u32x4 out_pack;
        asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(out_pack[0]) : "v"(gv0[0]), "v"(gv0[1]));
        asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(out_pack[1]) : "v"(gv0[2]), "v"(gv0[3]));
        asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(out_pack[2]) : "v"(gv1[0]), "v"(gv1[1]));
        asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(out_pack[3]) : "v"(gv1[2]), "v"(gv1[3]));
        *reinterpret_cast<ext_u32x4*>(output + (q_start * NUM_HEADS + head_idx) * V_HEAD_DIM + lane_id * ELEMS) = out_pack;
    }
}

// ============================================================
// Dispatch
// ============================================================
template<int NUM_PARTIALS, int TILES_PER_BLOCK>
void run_impl(uintptr_t q_ptr, uintptr_t kv_data_ptr, uintptr_t kv_scales_ptr,
    uintptr_t qo_indptr_ptr, uintptr_t kv_indptr_ptr, uintptr_t output_ptr,
    uintptr_t pout_ptr, uintptr_t pmax_ptr, uintptr_t psum_ptr,
    float sm_scale, int batch_size) {
    const auto* q_d = reinterpret_cast<const __hip_bfloat16*>(q_ptr);
    const auto* kvd = reinterpret_cast<const uint8_t*>(kv_data_ptr);
    const auto* kvs = reinterpret_cast<const uint8_t*>(kv_scales_ptr);
    const auto* qoi = reinterpret_cast<const int32_t*>(qo_indptr_ptr);
    const auto* kvi = reinterpret_cast<const int32_t*>(kv_indptr_ptr);
    auto* out = reinterpret_cast<__hip_bfloat16*>(output_ptr);
    auto* pout = reinterpret_cast<__hip_bfloat16*>(pout_ptr);
    auto* pmax = reinterpret_cast<float*>(pmax_ptr);
    auto* psum = reinterpret_cast<float*>(psum_ptr);
    const int tb = batch_size * NUM_PARTIALS;
    hipLaunchKernelGGL((mla_grouped<NUM_PARTIALS, TILES_PER_BLOCK>),
        dim3(tb), dim3(BLOCK_THREADS), TOTAL_LDS, 0,
        q_d, kvd, kvs, qoi, kvi, pout, pmax, psum, out, sm_scale);
    if (NUM_PARTIALS > 1) {
        if (NUM_PARTIALS >= 8) {
            // Multi-wavefront reduce: 4 warps for 64 partials, 2 warps for 8
            constexpr int REDUCE_WARPS = (NUM_PARTIALS >= 64) ? 16 : 2;
            constexpr int REDUCE_THREADS = REDUCE_WARPS * WAVEFRONT_SIZE;
            constexpr int REDUCE_LDS = REDUCE_WARPS * (V_HEAD_DIM * 4 + 2 * 4);
            hipLaunchKernelGGL((mla_reduce_multi<NUM_PARTIALS, REDUCE_WARPS>),
                dim3(batch_size * NUM_HEADS), dim3(REDUCE_THREADS), REDUCE_LDS, 0,
                pout, pmax, psum, qoi, out);
        } else {
            hipLaunchKernelGGL((mla_reduce<NUM_PARTIALS>),
                dim3(batch_size * NUM_HEADS), dim3(WAVEFRONT_SIZE), 0, 0,
                pout, pmax, psum, qoi, out);
        }
    }
}

void run_s1(uintptr_t a, uintptr_t b, uintptr_t c, uintptr_t d, uintptr_t e, uintptr_t f,
    uintptr_t po, uintptr_t pm, uintptr_t ps,
    float g, int i, int max_kvseqlen) {
    const int total_tiles = max_kvseqlen / 128;
    const int tpb = max(1, i * total_tiles / 256);
    const int np = total_tiles / tpb;

    if (np == 1 && tpb == 8) run_impl<1, 8>(a, b, c, d, e, f, po, pm, ps, g, i);
    else if (np == 1 && tpb == 64) run_impl<1, 64>(a, b, c, d, e, f, po, pm, ps, g, i);
    else if (np == 4 && tpb == 2) run_impl<4, 2>(a, b, c, d, e, f, po, pm, ps, g, i);
    else if (np == 4 && tpb == 16) run_impl<4, 16>(a, b, c, d, e, f, po, pm, ps, g, i);
    else if (np == 8 && tpb == 1) run_impl<8, 1>(a, b, c, d, e, f, po, pm, ps, g, i);
    else if (np == 8 && tpb == 8) run_impl<8, 8>(a, b, c, d, e, f, po, pm, ps, g, i);
    else if (np == 64 && tpb == 1) run_impl<64, 1>(a, b, c, d, e, f, po, pm, ps, g, i);
    else run_impl<1, 8>(a, b, c, d, e, f, po, pm, ps, g, i);
}

PYBIND11_MODULE(mla_new_kernel_module, m) { m.def("run_s1", &run_s1); }
"""

hip_module = load_inline(
    name="mla_new_kernel_module",
    cpp_sources="",
    cuda_sources=CUDA_SRC,
    with_cuda=True,
    verbose=True,
    extra_cuda_cflags=["-std=c++20", "-O3"],
    no_implicit_headers=True,
)


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    batch_size = config["batch_size"]
    kv_buffer, kv_scales = kv_data["mxfp4"]
    total_q_len = q.size(0)
    max_kvseqlen = kv_buffer.size(0) // batch_size
    total_tiles = max_kvseqlen // 128
    tpb = max(1, batch_size * total_tiles // 256)
    np_ = total_tiles // tpb
    tb = batch_size * np_
    pout = torch.empty(tb, NUM_HEADS, V_HEAD_DIM, dtype=torch.bfloat16, device=q.device)
    pmax = torch.empty(tb, NUM_HEADS, dtype=torch.float32, device=q.device)
    psum = torch.empty(tb, NUM_HEADS, dtype=torch.float32, device=q.device)
    output = torch.empty(total_q_len, NUM_HEADS, V_HEAD_DIM, dtype=torch.bfloat16, device=q.device)
    hip_module.run_s1(
        q.data_ptr(), kv_buffer.data_ptr(), kv_scales.data_ptr(),
        qo_indptr.data_ptr(), kv_indptr.data_ptr(), output.data_ptr(),
        pout.data_ptr(), pmax.data_ptr(), psum.data_ptr(),
        SM_SCALE, batch_size, max_kvseqlen)
    return output
scrolls · 770 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