Skip to content
KernelIndex
Search⌘K

submission 692061

npip99 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:17d80bc7a2b400a1de62748b6c7255cf3353b7d288241c3e2851ea1ec78e4933
license declaredunknown
license concludedunknown
authorsnpip99
imported2026-08-15

Techniques

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

fp4kv_data["mxfp4"][0].data_ptr(),
num-warps = 4constexpr u32 NUM_WARPS = 4;
shared-memory__shared__ union {

Kernel source

submission.py1616 lines
# See {filename}.hip for details

import os
import torch
from torch.utils.cpp_extension import load_inline
from typing import Any

from task import input_t, output_t

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

CUDA_PRELUDE = """
#ifndef __gfx950__
#define __gfx950__
#endif
"""

CPP_WRAPPER = """
void entry(
    const uintptr_t q, // .............. (bs, 16, 576) bf16
    const uintptr_t kv_indptr, // ...... (bs+1,) int32
    const uintptr_t kv_bf16, // ........ (total_kv, 576) bf16
    const uintptr_t kv_fp8, // ......... (total_kv, 576) fp8
    const uintptr_t kv_fp8_scale, // ... (1,) f32
    const uintptr_t kv_mxfp4, // ....... (total_kv, 288) fp4x2
    const uintptr_t kv_mxfp4_scale, // . (total_kv, 18) uint8_t
    uintptr_t out, // .................. (bs, 16, 512) bf16
    uint32_t bs
);
"""

CUDA_SRC = r"""
#define CDIV(a, b) (((a) + (b) - 1) / (b))

using u8 = uint8_t;
using u16 = uint16_t;
using u16x32 = u16 __attribute__((ext_vector_type(32)));
using u32 = uint32_t;
using u32x4 = u32 __attribute__((ext_vector_type(4)));
using u32x8 = u32 __attribute__((ext_vector_type(8)));
using u64 = uint64_t;
using f32 = float;
using f32x4 = f32 __attribute__((ext_vector_type(4)));
using bf16x2 = __bf16 __attribute__((ext_vector_type(2)));
using bf16x4 = __bf16 __attribute__((ext_vector_type(4)));
using bf16x8 = __bf16 __attribute__((ext_vector_type(8)));

__device__ inline f32 fast_exp(f32 x) {
    constexpr f32 LOG2E = 1.4426950408889634f;
    return __builtin_amdgcn_exp2f(x * LOG2E);
}

__device__ inline u16 f32_to_bf16(f32 v) {
    return (u16)(__builtin_bit_cast(u32, v) >> 16);
}

__device__ inline float bf16_to_f32(uint16_t v) {
    return __builtin_bit_cast(float, (uint32_t)v << 16);
}

// Quantize a group of 32 bf16 values -> E8M0 scale (biased uint8)
__device__ inline uint8_t bf16x32_to_scale_e8m0(const u16x32* row) {
    uint32_t amax_u16x2_packed = 0;
    const uint32_t* row32 = (const uint32_t*)row;
#pragma unroll
    for (int i = 0; i < 16; i++) {
        // abs both bf16 in one instruction
        uint32_t v = row32[i] & 0x7FFF7FFFu;
        // max both bf16 in one instruction
        asm("v_pk_max_u16 %0, %0, %1" : "+v"(amax_u16x2_packed) : "v"(v));
    }
    uint16_t amax_u16 = max((uint16_t)amax_u16x2_packed, (uint16_t)(amax_u16x2_packed >> 16));
    // Round exponent up when mantissa fraction >= 0.75
    // (bf16 has 7 mantissa bits; 0x20 = 2^5 matches the f32 0x200000 = 2^21 rounding)
    int biased_exp = ((amax_u16 + 0x20u) >> 7) & 0xFF;
    int e8m0 = biased_exp - 2;
    return (uint8_t)max(0, min(254, e8m0));
}

// Matches `fp4_utils.e8m0_to_f32`: scale_e8m0 << 23 as f32, with special cases
__device__ inline float e8m0_to_f32(uint8_t scale) {
    return __builtin_bit_cast(float, (uint32_t)scale << 23);
}

__device__ inline __bf16 fp4_e2m1_to_bf16(u8 nibble, u8 scale) {
    u8 s = (nibble >> 3) & 1;
    u8 e = (nibble >> 1) & 3;
    u8 m = nibble & 1;
    f32 val = (e == 0) ? m * 0.5f : exp2f((f32)e - 1.0f) * (1.0f + m * 0.5f);
    val = s ? -val : val;
    u16 ret = f32_to_bf16(val * e8m0_to_f32(scale));
    return *(const __bf16*)&ret;
}

__device__ inline uint8_t f32_to_fp4_e2m1(float v) {
    // CAREFUL: We need to include the fix from https://github.com/ROCm/aiter/pull/975 / https://github.com/ROCm/aiter/pull/2249
    // FP4 e2m1: EBITS=2, MBITS=1, EXP_BIAS=1
    // denorm_exp = (127-1) + (23-1) + 1 = 149  =>  denorm_mask = 149 << 23 = 0x4A800000
    // val_to_add = ((1-127) << 23) + ((1<<21)-1) = 0xC11FFFFF  (int32 wrapping)
    constexpr float    FP4_MAX_NORMAL    = 6.0f;
    constexpr float    FP4_MIN_NORMAL    = 1.0f;
    constexpr uint32_t DENORM_MASK_INT   = 0x4A800000u;
    constexpr float    DENORM_MASK_FLOAT = __builtin_bit_cast(float, DENORM_MASK_INT);
    constexpr uint32_t VAL_TO_ADD        = 0xC11FFFFFu;

    uint8_t  s_bit    = (uint8_t)((__builtin_bit_cast(uint32_t, v) >> 28) & 0x8u);
    uint32_t abs_bits = __builtin_bit_cast(uint32_t, v) & 0x7FFFFFFFu;
    float    abs_v    = __builtin_bit_cast(float, abs_bits);

    uint8_t result;
    if (abs_v >= FP4_MAX_NORMAL) {
        // saturate branch
        result = 0x7u;
    } else if (abs_v < FP4_MIN_NORMAL) {
        // denormal branch: float trick to extract low bits
        uint32_t d = __builtin_bit_cast(uint32_t, abs_v + DENORM_MASK_FLOAT) - DENORM_MASK_INT;
        result = (uint8_t)d;
    } else {
        // normal branch: bias-adjust exponent + round-to-nearest
        uint32_t x     = abs_bits;
        uint32_t m_odd = (x >> 22) & 1u;
        x = x + VAL_TO_ADD + m_odd;
        result = (uint8_t)(x >> 22);
    }
    return s_bit | result;
}

// Take an f32 and its block's scaling term, and convert to fp4_e2m1
// The fp4_e2m1 will be in the low nibble of the return value
__device__ inline uint8_t f32_to_fp4_e2m1_scale(float v, uint8_t scale) {
    float f32_scale = e8m0_to_f32((uint8_t)(254u - scale)); // 2^(-(exp_unb)) = 2^(127 - scale)
    return f32_to_fp4_e2m1(v * f32_scale);
}

__device__ inline f32x4 mfma_bf16_16x16x32(bf16x8 a, bf16x8 b, f32x4 acc) {
#ifdef __gfx950__
    // return __builtin_amdgcn_mfma_f32_16x16x32_bf16(a, b, acc, 0, 0, 0);
    // AGPRs are bad, we can't multiply alpha into them. Pin VGPR.
    asm("v_mfma_f32_16x16x32_bf16 %0, %1, %2, %0"
        : "+v"(acc)
        : "v"(a), "v"(b));
    return acc;
#else
    acc = __builtin_amdgcn_mfma_f32_16x16x16bf16_1k(*(bf16x4*)&a, *(bf16x4*)&b, acc, 0, 0, 0);
    acc = __builtin_amdgcn_mfma_f32_16x16x16bf16_1k(*((bf16x4*)&a+1), *((bf16x4*)&b+1), acc, 0, 0, 0);
    return acc;
#endif
}

__device__ inline void global_to_lds_1024b(const u8* global_base, u32 lds_base, u32 lane_id) {
#ifdef __gfx950__
    asm volatile(
        "s_mov_b32 m0, %0\n\t"
        "global_load_lds_dwordx4 %1, off\n\t"
        :: "s"(lds_base), "v"((const void*)(global_base + lane_id * 16))
        : "memory", "m0"
    );
#else
    #pragma unroll
    for (u32 sub = 0; sub < 4; sub++) {
        asm volatile(
            "s_mov_b32 m0, %0\n\t"
            "global_load_lds_dword %1, off\n\t"
            :: "s"(lds_base + sub * 256), "v"((const void*)(global_base + sub * 256 + lane_id * 4))
            : "memory", "m0"
        );
    }
#endif
}

__device__ inline void global_to_lds_256b(const u8* global_base, u32 lds_base, u32 lane_id) {
    asm volatile(
        "s_mov_b32 m0, %0\n\t"
        "global_load_lds_dword %1, off\n\t"
        :: "s"(lds_base), "v"((const void*)(global_base + lane_id * 4))
        : "memory", "m0"
    );
}

__device__ inline f32 dpp_reduce_max_16(f32 v) {
    #define DPP_MOV(x, ctrl) __builtin_bit_cast(f32, \
        __builtin_amdgcn_mov_dpp(__builtin_bit_cast(u32, x), ctrl, 0xF, 0xF, false))
    v = fmaxf(v, DPP_MOV(v, 0xB1));   // quad XOR 1
    v = fmaxf(v, DPP_MOV(v, 0x4E));   // quad XOR 2
    v = fmaxf(v, DPP_MOV(v, 0x124));  // row_ror:4
    v = fmaxf(v, DPP_MOV(v, 0x128));  // row_ror:8
    return v;
    #undef DPP_MOV
}

__device__ inline f32 dpp_reduce_sum_16(f32 v) {
    #define DPP_MOV(x, ctrl) __builtin_bit_cast(f32, \
        __builtin_amdgcn_mov_dpp(__builtin_bit_cast(u32, x), ctrl, 0xF, 0xF, false))
    v += DPP_MOV(v, 0xB1);
    v += DPP_MOV(v, 0x4E);
    v += DPP_MOV(v, 0x124);
    v += DPP_MOV(v, 0x128);
    return v;
    #undef DPP_MOV
}

// Constants
constexpr f32 NEG_INF = -1e30f;
constexpr u32 THREADS_PER_WARP = 64;
#ifdef __gfx950__
constexpr u32 NUM_WARPS = 4;
#else
constexpr u32 NUM_WARPS = 2;
#endif
constexpr u32 THREADS_PER_BLOCK = THREADS_PER_WARP * NUM_WARPS;
constexpr u32 SCALE_GROUP_SIZE = 32;

// DeepSeek Parameters
constexpr u32 N_HEADS = 16;
constexpr u32 QK_HEAD_DIM = 576;
constexpr u32 V_HEAD_DIM = 512;

__global__ __launch_bounds__(THREADS_PER_BLOCK, 1)
void kernel(
    const u16* __restrict__ q,
    const u32* __restrict__ kv_indptr,
    const u16* __restrict__ kv_bf16,
    const u8* __restrict__ kv_fp8, const f32* __restrict__ kv_fp8_scale,
    const u8* __restrict__ kv_mxfp4, const u8* __restrict__ kv_mxfp4_scale,
    u16* __restrict__ out,
    int bs
) {
    u32 batch_idx = __builtin_amdgcn_readfirstlane(blockIdx.x);
    u32 warp_id = __builtin_amdgcn_readfirstlane(threadIdx.x / THREADS_PER_WARP);
    u32 lane_id = threadIdx.x % THREADS_PER_WARP;
    u32 lane_col = lane_id % 16;
    u32 lane_rowgroup = lane_id / 16; // rows \in 4*lane_rowgroup + {0,1,2,3}

    // Pre-load Q into registers (quantization deferred to after LDS prefetch)
    const u16* q_batch_item = q + batch_idx * N_HEADS * QK_HEAD_DIM;
    constexpr u32 SCORE_MFMA_DOT_DIM = 128;
    constexpr u32 SCORE_MFMA_ITERS = CDIV(QK_HEAD_DIM, SCORE_MFMA_DOT_DIM);

    u32x4 q_data[SCORE_MFMA_ITERS][4];
    #pragma unroll
    for (u32 k = 0; k < SCORE_MFMA_ITERS; k++) {
        u32 dim_base = k * SCORE_MFMA_DOT_DIM + lane_rowgroup * SCALE_GROUP_SIZE;
        const u16* q_chunk = q_batch_item + lane_col * QK_HEAD_DIM + dim_base;
        q_data[k][0] = *(const u32x4*)(q_chunk);
        q_data[k][1] = *(const u32x4*)(q_chunk + 8);
        q_data[k][2] = *(const u32x4*)(q_chunk + 16);
        q_data[k][3] = *(const u32x4*)(q_chunk + 24);
    }    

    u32 kv_start = kv_indptr[batch_idx];
    u32 kv_end = kv_indptr[batch_idx + 1];

    f32 softmax_max[4] = {NEG_INF, NEG_INF, NEG_INF, NEG_INF};
    f32 softmax_denom[4] = {};

    // KV Tiles
    constexpr u32 KV_TILE_DIM = 32;
    constexpr u32 PREFETCH = 2;
    constexpr u32 KV_SCALE_STRIDE = CDIV(QK_HEAD_DIM / SCALE_GROUP_SIZE, 8) * 8;

    // Shared Memory
    __shared__ union {
        // kv iterations
        struct {
            u8 kv_data[NUM_WARPS][PREFETCH][KV_TILE_DIM][QK_HEAD_DIM / 2];
            u8 kv_scale[NUM_WARPS][PREFETCH][KV_TILE_DIM][KV_SCALE_STRIDE];
            u16 weights[NUM_WARPS][N_HEADS][KV_TILE_DIM];
        } tile;
        // merge before write
        struct {
            f32 warp_max[NUM_WARPS][N_HEADS];
            f32 warp_denom[NUM_WARPS][N_HEADS];
            u16 values[NUM_WARPS][N_HEADS][V_HEAD_DIM];
        } merge;
    } lds;

    // (u32 kv_tile_start = kv_start; kv_tile_start < kv_end; kv_tile_start += KV_TILE_DIM) {
    constexpr u32 KV_DATA_BYTES = KV_TILE_DIM * (QK_HEAD_DIM / 2);
    constexpr u32 KV_SCALE_BYTES = KV_TILE_DIM * KV_SCALE_STRIDE;

    // These stay in SGPRs
    u32 lds_data_buf0 = (u32)(uintptr_t)&lds.tile.kv_data[warp_id][0];
    u32 lds_data_buf1 = (u32)(uintptr_t)&lds.tile.kv_data[warp_id][1];
    u32 lds_scale_buf0 = (u32)(uintptr_t)&lds.tile.kv_scale[warp_id][0];
    u32 lds_scale_buf1 = (u32)(uintptr_t)&lds.tile.kv_scale[warp_id][1];
    
    auto load_lds = [&](uint32_t kv_iter) -> void {
        u32 kv_tile_start = kv_start + kv_iter * KV_TILE_DIM;
        u32 kv_buf = kv_iter % 2;
        u32 lds_data_off = kv_buf ? lds_data_buf1 : lds_data_buf0;
        u32 lds_scale_off = kv_buf ? lds_scale_buf1 : lds_scale_buf0;

        constexpr u32 KV_DATA_U128 = KV_DATA_BYTES / 16;
        static_assert(KV_DATA_U128 % THREADS_PER_WARP == 0);
        constexpr u32 KV_DATA_ITERS = KV_DATA_U128 / THREADS_PER_WARP;
        const u8* kv_data_src = (const u8*)(kv_mxfp4 + (u64)kv_tile_start * (QK_HEAD_DIM / 2));
        #pragma unroll
        for (u32 i = 0; i < KV_DATA_ITERS; i++) {
            global_to_lds_1024b(kv_data_src + i * 1024, lds_data_off + i * 1024, lane_id);
        }

        constexpr u32 KV_SCALE_U32 = KV_SCALE_BYTES / 4;
        constexpr u32 KV_SCALE_ITERS = CDIV(KV_SCALE_U32, THREADS_PER_WARP);
        const u8* kv_scale_src = kv_mxfp4_scale + (u64)kv_tile_start * KV_SCALE_STRIDE;
        #pragma unroll
        for (u32 i = 0; i < KV_SCALE_ITERS; i++) {
            global_to_lds_256b(kv_scale_src + i * 256, lds_scale_off + i * 256, lane_id);
        }
    };

    u32 kv_range = kv_end - kv_start;
    u32 kv_num_iters = kv_range / KV_TILE_DIM / NUM_WARPS;
    if (kv_range % (NUM_WARPS * KV_TILE_DIM) != 0 || kv_num_iters < PREFETCH - 1) {
        __builtin_trap();
    }
    u32 warp_kv_offset = warp_id * kv_num_iters;

    #pragma unroll
    for (u32 i = 0; i < PREFETCH - 1; i++) {
        load_lds(warp_kv_offset + i);
    }

    constexpr u32 VALUE_MFMA_V_TILE_DIM = 16;
    constexpr u32 VALUE_MFMA_DOT_DIM = 32;
    constexpr u32 NUM_VALUE_TILES = V_HEAD_DIM / VALUE_MFMA_V_TILE_DIM;
    f32x4 value_out_lanes[NUM_VALUE_TILES] = {};

    __builtin_amdgcn_sched_barrier(0);
     
    // ========================================
    // Distributed Q quantization across warps
    // Each warp quantizes k iterations where k % NUM_WARPS == warp_id, shares via LDS
    // ========================================

    __shared__ u32 q_lds[SCORE_MFMA_ITERS][64][8];
    __shared__ u8 q_scale_lds[SCORE_MFMA_ITERS][64];

    u32x8 q_lanes[SCORE_MFMA_ITERS];
    u8 q_scale_lanes[SCORE_MFMA_ITERS];

    #pragma unroll
    for (u32 k = 0; k < SCORE_MFMA_ITERS; k++) {
        if (k % NUM_WARPS != warp_id) continue;
        u32 dim_base = k * SCORE_MFMA_DOT_DIM + lane_rowgroup * SCALE_GROUP_SIZE;
        if (k < SCORE_MFMA_ITERS - 1 || dim_base + SCALE_GROUP_SIZE <= QK_HEAD_DIM) {
            q_scale_lanes[k] = bf16x32_to_scale_e8m0((const u16x32*)&q_data[k]);
            f32 s = e8m0_to_f32(q_scale_lanes[k]);
            #pragma unroll
            for (u32 r = 0; r < 4; r++) {
#ifdef __gfx950__
                const bf16x2* v2 = (const bf16x2*)&q_data[k][r];
                q_lanes[k][r] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(0, v2[0], s, 0);
                q_lanes[k][r] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q_lanes[k][r], v2[1], s, 1);
                q_lanes[k][r] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q_lanes[k][r], v2[2], s, 2);
                q_lanes[k][r] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q_lanes[k][r], v2[3], s, 3);
#else
                const u16* q_r = (const u16*)&q_data[k][r];
                q_lanes[k][r] = 0;
                #pragma unroll
                for (u32 i = 0; i < 8; i++) {
                    u8 nibble = f32_to_fp4_e2m1_scale(bf16_to_f32(q_r[i]), q_scale_lanes[k]);
                    q_lanes[k][r] |= (u32)(nibble & 0xF) << (4 * i);
                }
#endif
            }
        } else {
            q_lanes[k] = {};
            q_scale_lanes[k] = 0;
        }
        // Write to LDS
        #pragma unroll
        for (u32 r = 0; r < 8; r++) {
            q_lds[k][lane_id][r] = q_lanes[k][r];
        }
        q_scale_lds[k][lane_id] = q_scale_lanes[k];
    }
    __syncthreads();

    // Read back iterations this warp didn't compute
    #pragma unroll
    for (u32 k = 0; k < SCORE_MFMA_ITERS; k++) {
        if (k % NUM_WARPS == warp_id) continue;
        #pragma unroll
        for (u32 r = 0; r < 8; r++) {
            q_lanes[k][r] = q_lds[k][lane_id][r];
        }
        q_scale_lanes[k] = q_scale_lds[k][lane_id];
    }

    __builtin_amdgcn_sched_barrier(0);
    
    for (u32 kv_iter = 0; kv_iter < kv_num_iters; kv_iter++) {
        // ==========
        // Cooperative load of KV Cache MXFP4 tile to LDS
        // ==========

        // Wait for the prefetched LDS to land
        // asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
        static_assert(PREFETCH <= 2); // Need custom vmcnt flag for PREFETCH > 2
        __builtin_amdgcn_s_waitcnt(0x3f70);

        // Prefetch the next LDS
        u32 kv_buf = kv_iter % 2;
        u32 kv_prefetch_iter = kv_iter + PREFETCH - 1;
        if (kv_prefetch_iter < kv_num_iters) {
            load_lds(warp_kv_offset + kv_prefetch_iter);
        }

        // ==========
        // scores = per-head query @ key * 1/sqrt(d)
        // scores = Q(N_HEADS,QK_HEAD_DIM) @ K(32,QK_HEAD_DIM)^T * 1/sqrt(QK_HEAD_DIM)
        // ==========

        constexpr u32 SCORE_MFMA_KV_TILE_DIM = 16;
        static_assert(KV_TILE_DIM % SCORE_MFMA_KV_TILE_DIM == 0);
        constexpr u32 MFMA_PER_KV_TILE_DIM = KV_TILE_DIM / SCORE_MFMA_KV_TILE_DIM;
        f32x4 scores[MFMA_PER_KV_TILE_DIM] = {};

        #pragma unroll
        for (u32 score_mfma_idx = 0; score_mfma_idx < SCORE_MFMA_ITERS; score_mfma_idx++) {
            #pragma unroll
            for (u32 kv_mfma_idx = 0; kv_mfma_idx < MFMA_PER_KV_TILE_DIM; kv_mfma_idx++) {
                u32 tok = kv_mfma_idx * 16 + lane_col;
                u32 kv_byte_offset = score_mfma_idx * (SCORE_MFMA_DOT_DIM / 2) + lane_rowgroup * 16;

                u32x8 kv_reg;
                u8 kv_scale = 0;
                if (score_mfma_idx < SCORE_MFMA_ITERS - 1 || kv_byte_offset + 16 <= QK_HEAD_DIM / 2) {
                    *(u32x4*)&kv_reg = *(const u32x4*)&lds.tile.kv_data[warp_id][kv_buf][tok][kv_byte_offset];
                    kv_scale = lds.tile.kv_scale[warp_id][kv_buf][tok][score_mfma_idx * (SCORE_MFMA_DOT_DIM / SCALE_GROUP_SIZE) + lane_rowgroup];
                } else {
                    *(u32x4*)&kv_reg = {};
                }

#ifdef __gfx950__
                scores[kv_mfma_idx] = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
                    q_lanes[score_mfma_idx],
                    kv_reg,
                    scores[kv_mfma_idx],
                    4, 4,
                    0, q_scale_lanes[score_mfma_idx],
                    0, kv_scale
                );
#else
                #pragma unroll
                for (u32 sub = 0; sub < 4; sub++) {
                    bf16x8 q_bf16, kv_bf16_lane;
                    #pragma unroll
                    for (u32 i = 0; i < 8; i++) {
                        u32 qi = sub * 8 + i;
                        u8 q_nibble = (q_lanes[score_mfma_idx][qi / 8] >> (4 * (qi % 8))) & 0xF;
                        q_bf16[i] = fp4_e2m1_to_bf16(q_nibble, q_scale_lanes[score_mfma_idx]);

                        u8 kv_byte = ((const u8*)&kv_reg)[sub * 4 + i / 2];
                        u8 kv_nibble = (i % 2 == 0) ? (kv_byte & 0xF) : (kv_byte >> 4);
                        kv_bf16_lane[i] = fp4_e2m1_to_bf16(kv_nibble, kv_scale);
                    }
                    scores[kv_mfma_idx] = mfma_bf16_16x16x32(q_bf16, kv_bf16_lane, scores[kv_mfma_idx]);
                }
#endif
            }
        }

        // Scale by 1 / sqrt(QK_HEAD_DIM)
        static_assert(QK_HEAD_DIM == 24 * 24); // Add constexpr sqrt to make this responsive w.r.t QK_HEAD_DIM
        constexpr f32 SM_SCALE = 1.0f / 24.0f;
        #pragma unroll
        for (u32 i = 0; i < MFMA_PER_KV_TILE_DIM; i++) {
            #pragma unroll
            for (u32 j = 0; j < 4; j++) {
                scores[i][j] *= SM_SCALE;
            }
        }

        // ==========
        // Online softmax
        // ==========

        // == MFMA thread-local KV-tile max reduction ==
        // From the MFMA, each lane's (lane_col, lane_rowgroup) holds values from 4 different heads.
        // We need to max-reduce over the MFMA KV Tiles, to get the per-lane max for each head
        f32 lane_max[4] = {scores[0][0], scores[0][1], scores[0][2], scores[0][3]};
        #pragma unroll
        for (u32 i = 1; i < MFMA_PER_KV_TILE_DIM; i++) {
            #pragma unroll
            for (u32 j = 0; j < 4; j++) {
                lane_max[j] = fmaxf(lane_max[j], scores[i][j]);
            }
        }

        // == MFMA intra-lanegroup max reduction ==
        // 16-consecutive lanes will share the same lane_rowgroup, so we reduce the max together.
        // Each lane will then get the same true max of the KV_TILE_DIM rows, for its 4 heads.
        #pragma unroll
        for (u32 i = 0; i < 4; i++) {
            lane_max[i] = dpp_reduce_max_16(lane_max[i]);
        }

        // == Update online state ==
        // new_head_max is the new global max (max of accumulated kv_tile softmax_max[i] and current kv_tile lane_max[i]).
        // alpha = exp(old_max - new_max) is the correction factor to rescale all previously accumulated values by
        // We update previously accumulated softmax_denom and softmax_max by this alpha right now.
        //   - Updating previously accumulated value_out_lanes will be done later.
        f32x4 alpha;
        #pragma unroll
        for (u32 i = 0; i < 4; i++) {
            f32 new_head_max = fmaxf(softmax_max[i], lane_max[i]);
            alpha[i] = fast_exp(softmax_max[i] - new_head_max);
            softmax_denom[i] *= alpha[i];
            softmax_max[i] = new_head_max;
        }

        // == Compute weights, thread-local KV-tile sum reduction ==
        // weight = exp(score-score_max)
        // For each score MFMA'ss output lane item, we calculate the weight for later value accumulation
        // Weights are stored in LDS (for intra-warp permutation, value lanes are transposed)
        // lane_sum is reduced across MFMA_PER_KV_TILE_DIM, for the 4 unique heads per lane.
        f32 weights[MFMA_PER_KV_TILE_DIM][4];
        f32 lane_sum[4] = {};
        #pragma unroll
        for (u32 i = 0; i < MFMA_PER_KV_TILE_DIM; i++) {
            #pragma unroll
            for (u32 j = 0; j < 4; j++) {
                weights[i][j] = fast_exp(scores[i][j] - softmax_max[j]);
                lds.tile.weights[warp_id][4 * lane_rowgroup + j][i * 16 + lane_col] = f32_to_bf16(weights[i][j]);
                lane_sum[j] += weights[i][j];
            }
        }

        // == accumulate per-lane softmax_denom ==
        // We defer cross-lane reduction to after loop
        #pragma unroll
        for (u32 i = 0; i < 4; i++) {
            softmax_denom[i] += lane_sum[i];
        }

        // ==========
        // value_out += lds_weights(N_HEADS,KV_TILE_DIM) @ kv(KV_TILE_DIM,V_HEAD_DIM)^T
        // ==========

        // Read weights for MFMA value lane: [head=lcol, tokens lgrp*8..+7]
        bf16x8 value_weight_lane = *(const bf16x8*)&lds.tile.weights[warp_id][lane_col][lane_rowgroup * 8];

        // MFMA for weights * values
        // We process in MFMA tiles of V_HEAD_DIM / VALUE_MFMA_V_TILE_DIM
        constexpr u32 FEATURES_PER_DWORD = 8;  // 4 bytes = 8 nibbles
        constexpr u32 VALUEGROUP_ITERS = V_HEAD_DIM / (VALUE_MFMA_V_TILE_DIM * FEATURES_PER_DWORD); // 512/(16*8) = 4
        static_assert(VALUEGROUP_ITERS == 4);

        // Scales: all 4 valuegroups fall in the same scale group (lane_col * 16 bytes = lane_col * 32 features → scale_idx = lane_col)
        u32 scale_base = (u32)(uintptr_t)&lds.tile.kv_scale[warp_id][kv_buf][lane_rowgroup * 8][lane_col];
        u32 scale_reg[8];
        __builtin_amdgcn_sched_barrier(0);
        #pragma unroll
        for (u32 row = 0; row < 8; row++) {
            asm volatile("ds_read_u8 %0, %1 offset:%c2" : "=v"(scale_reg[row]) : "v"(scale_base), "n"(row * KV_SCALE_STRIDE));
        }
        __builtin_amdgcn_sched_barrier(0);

        constexpr u32 VALUE_PREFETCH = 3;
        u32 data_reg[VALUE_PREFETCH][8];

        constexpr u32 VALUE_LOADS_PER_PRFETCH = 4;
        auto loadValueLDS = [&](u32 vg_idx) {
            u32 buf = vg_idx % VALUE_PREFETCH;
            u32 byte_col_base = lane_col * 16 + vg_idx * 4;
            u32 data_base = (u32)(uintptr_t)&lds.tile.kv_data[warp_id][kv_buf][lane_rowgroup * 8][byte_col_base];
            u32 data_base_hi = data_base + 4 * (QK_HEAD_DIM / 2);
            constexpr u32 QK_STRIDE_DWORD = (QK_HEAD_DIM / 2) / 4;
            __builtin_amdgcn_sched_barrier(0);
            asm volatile("ds_read2_b32 %0, %1 offset0:%c2 offset1:%c3"
                : "=v"(*(u64*)&data_reg[buf][0])
                : "v"(data_base), "n"((u32)0), "n"(QK_STRIDE_DWORD));
            asm volatile("ds_read2_b32 %0, %1 offset0:%c2 offset1:%c3"
                : "=v"(*(u64*)&data_reg[buf][2])
                : "v"(data_base), "n"(QK_STRIDE_DWORD * 2), "n"(QK_STRIDE_DWORD * 3));
            asm volatile("ds_read2_b32 %0, %1 offset0:%c2 offset1:%c3"
                : "=v"(*(u64*)&data_reg[buf][4])
                : "v"(data_base_hi), "n"((u32)0), "n"(QK_STRIDE_DWORD));
            asm volatile("ds_read2_b32 %0, %1 offset0:%c2 offset1:%c3"
                : "=v"(*(u64*)&data_reg[buf][6])
                : "v"(data_base_hi), "n"(QK_STRIDE_DWORD * 2), "n"(QK_STRIDE_DWORD * 3));
            __builtin_amdgcn_sched_barrier(0);
        };

        for (u32 valuegroup_idx = 0; valuegroup_idx < VALUE_PREFETCH - 1; valuegroup_idx++) {
            loadValueLDS(valuegroup_idx);
        }

        #pragma unroll
        for (u32 valuegroup_idx = 0; valuegroup_idx < VALUEGROUP_ITERS; valuegroup_idx++) {
            // 16 lanes x 4 bytes = 64 bytes = 128 features per value group
            // All 8 features in one dword share the same scale group
            // (8 features < SCALE_GROUP_SIZE=32)

            // Load from LDS
            // 4 bytes (data) + 1 byte (scale) from each of 8 rows
            __builtin_amdgcn_sched_barrier(0);
            if (valuegroup_idx + VALUE_PREFETCH - 1 < VALUEGROUP_ITERS) {
                loadValueLDS(valuegroup_idx + VALUE_PREFETCH - 1);
            }
            u32 inflight = (VALUEGROUP_ITERS - 1 - valuegroup_idx < VALUE_PREFETCH - 1)
                ? (VALUEGROUP_ITERS - 1 - valuegroup_idx)
                : (VALUE_PREFETCH - 1);
            static_assert(VALUE_PREFETCH <= 3); // Idk something weird happens here for >4. min(15, ..) doesn't fix it.
            asm volatile("s_waitcnt lgkmcnt(%c0)" :: "n"(inflight * VALUE_LOADS_PER_PRFETCH) : "memory");
            __builtin_amdgcn_sched_barrier(0);
            u32 buf = valuegroup_idx % VALUE_PREFETCH;

            // 4 byte columns, each producing 2 MFMAs (one for each nibble column)
            #pragma unroll
            for (u32 byte_col_offset = 0; byte_col_offset < 4; byte_col_offset++) {
                // the value tile index for lo nibble / hi nibble
                u32 lo_value_tile = valuegroup_idx * 8 + byte_col_offset * 2;
                u32 hi_value_tile = lo_value_tile + 1;
                value_out_lanes[lo_value_tile] *= alpha;
                value_out_lanes[hi_value_tile] *= alpha;

                bf16x8 value_lane_lo, value_lane_hi;
                #pragma unroll
                for (u32 row = 0; row < 8; row++) {
#ifdef __gfx950__
                    // f32 scale_f32 = e8m0_to_f32(scale);
                    // `e8m0_to_f32` requires an AND mask since it doesn't know ds_read_u8 will zero out the unused 3 bytes
                    f32 scale_f32 = __builtin_bit_cast(float, scale_reg[row] << 23);
                    u32 cvt;
                    switch (byte_col_offset) {
                        case 0: cvt = __builtin_bit_cast(u32, __builtin_amdgcn_cvt_scalef32_pk_bf16_fp4(data_reg[buf][row], scale_f32, 0)); break;
                        case 1: cvt = __builtin_bit_cast(u32, __builtin_amdgcn_cvt_scalef32_pk_bf16_fp4(data_reg[buf][row], scale_f32, 1)); break;
                        case 2: cvt = __builtin_bit_cast(u32, __builtin_amdgcn_cvt_scalef32_pk_bf16_fp4(data_reg[buf][row], scale_f32, 2)); break;
                        case 3: cvt = __builtin_bit_cast(u32, __builtin_amdgcn_cvt_scalef32_pk_bf16_fp4(data_reg[buf][row], scale_f32, 3)); break;
                    }
                    value_lane_lo[row] = __builtin_bit_cast(__bf16, (u16)cvt);
                    value_lane_hi[row] = __builtin_bit_cast(__bf16, (u16)(cvt >> 16));
#else
                    u8 scale = (u8)scale_reg[row];
                    u8 nibble_lo = (data_reg[buf][row] >> (byte_col_offset * 8)) & 0xF;
                    u8 nibble_hi = (data_reg[buf][row] >> (byte_col_offset * 8 + 4)) & 0xF;
                    value_lane_lo[row] = fp4_e2m1_to_bf16(nibble_lo, scale);
                    value_lane_hi[row] = fp4_e2m1_to_bf16(nibble_hi, scale);
#endif
                }

                value_out_lanes[lo_value_tile] = mfma_bf16_16x16x32(value_weight_lane, value_lane_lo, value_out_lanes[lo_value_tile]);
                value_out_lanes[hi_value_tile] = mfma_bf16_16x16x32(value_weight_lane, value_lane_hi, value_out_lanes[hi_value_tile]);
            }
        }
    }

    // ==========
    // Reduce online softmax parameters across all warps, and store results in LDS
    // lds value_out[NUM_WARPS] = value_out_lanes * interwarp_max_correction / interwarp_denom
    // ==========
    
    // Reduce softmax_denom across lanes in this rowgroup (same value for lane_id / 16)
    #pragma unroll
    for (u32 i = 0; i < 4; i++) {
        softmax_denom[i] = dpp_reduce_sum_16(softmax_denom[i]);
    }

    // == Store per-warp online softmax counters ==
    // We've already reduced within a warp, so only the per-head leader `lane_col == 0` needs to write.
    // There are 4 rowgroups, 4 heads each, 16 heads total.
    if (lane_col == 0) {
        #pragma unroll
        for (u32 i = 0; i < 4; i++) {
            lds.merge.warp_max[warp_id][4 * lane_rowgroup + i] = softmax_max[i];
            lds.merge.warp_denom[warp_id][4 * lane_rowgroup + i] = softmax_denom[i];
        }
    }
    __syncthreads();

    // == per-head online softmax correction ==
    // For the 4 heads, we get the max and denom by reducing the per-warp max/denom already written
    // exp(warp_max-interwarp_max) will be the max correction on this warp
    // We can also divide by the denominator, since inter-warp is the global data for this batch index.
    f32 global_correction[4];
    #pragma unroll
    for (u32 i = 0; i < 4; i++) {
        u32 head = 4 * lane_rowgroup + i;
        f32 interwarp_max = NEG_INF;
        #pragma unroll
        for (u32 w = 0; w < NUM_WARPS; w++) {
            interwarp_max = fmaxf(interwarp_max, lds.merge.warp_max[w][head]);
        }
        f32 interwarp_denom = 0;
        #pragma unroll
        for (u32 w = 0; w < NUM_WARPS; w++) {
            interwarp_denom += lds.merge.warp_denom[w][head] * fast_exp(lds.merge.warp_max[w][head] - interwarp_max);
        }
        // interwarp denom, is also the global denom
        global_correction[i] = fast_exp(softmax_max[i] - interwarp_max) / interwarp_denom;
    }

    // == per-head softmax'd values ==
    // By adjusting our value_out_lanes by the correction, we can write the corrected results to LDS
    // Interwarp reduction requires LDS communication.
    #pragma unroll
    for (u32 valuegroup_idx = 0; valuegroup_idx < 4; valuegroup_idx++) {
        #pragma unroll
        for (u32 value_tile_col_offset = 0; value_tile_col_offset < 8; value_tile_col_offset++) {
            u32 value_tile_idx = valuegroup_idx * 8 + value_tile_col_offset;
            #pragma unroll
            for (u32 i = 0; i < 4; i++) {
                u32 head = 4 * lane_rowgroup + i;
                u32 vdim = lane_col * 32 + valuegroup_idx * 8 + value_tile_col_offset;
                lds.merge.values[warp_id][head][vdim] = f32_to_bf16(value_out_lanes[value_tile_idx][i] * global_correction[i]);
            }
        }
    }
    __syncthreads();

    // ==========
    // Write to global
    // global out = \sum_warp value_out
    // ==========

    // == Use LDS to reduce sum over warps, and write to global ==
    // All threads can work together on reducing and writing to global
    // The original lane assignments are irrelevant now, we just distribute the work evenly and in order.
    u16* out_batch_item = out + batch_idx * N_HEADS * V_HEAD_DIM;
    constexpr u32 TOTAL_ELEMS = N_HEADS * V_HEAD_DIM;
    constexpr u32 CHUNK = 8;
    constexpr u32 ITERS = TOTAL_ELEMS / (THREADS_PER_BLOCK * CHUNK);
    static_assert(TOTAL_ELEMS % (THREADS_PER_BLOCK * CHUNK) == 0);

    // LDS -> reg
    u32x4 loaded[ITERS][NUM_WARPS];
    #pragma unroll
    for (u32 iter = 0; iter < ITERS; iter++) {
        u32 base = (iter * THREADS_PER_BLOCK + threadIdx.x) * CHUNK;
        #pragma unroll
        for (u32 w = 0; w < NUM_WARPS; w++) {
            loaded[iter][w] = *(const u32x4*)(&lds.merge.values[w][0][0] + base);
        }
    }

    // reg -> reduce -> store
    #pragma unroll
    for (u32 iter = 0; iter < ITERS; iter++) {
        f32 acc[CHUNK];
        #pragma unroll
        for (u32 i = 0; i < CHUNK; i++) {
            u32 word = loaded[iter][0][i / 2];
            acc[i] = bf16_to_f32((u16)(word >> (16 * (i & 1))));
        }
        #pragma unroll
        for (u32 w = 1; w < NUM_WARPS; w++) {
            #pragma unroll
            for (u32 i = 0; i < CHUNK; i++) {
                u32 word = loaded[iter][w][i / 2];
                acc[i] += bf16_to_f32((u16)(word >> (16 * (i & 1))));
            }
        }
        u32x4 out_packed;
        #pragma unroll
        for (u32 i = 0; i < 4; i++) {
            out_packed[i] = (u32)f32_to_bf16(acc[i * 2]) | ((u32)f32_to_bf16(acc[i * 2 + 1]) << 16);
        }
        u32 base = (iter * THREADS_PER_BLOCK + threadIdx.x) * CHUNK;
        *(u32x4*)(out_batch_item + base) = out_packed;
    }
}


void entry(
    const uintptr_t q, // .............. (bs, 16, 576) bf16
    const uintptr_t kv_indptr, // ...... (bs+1,) int32
    const uintptr_t kv_bf16, // ........ (total_kv, 576) bf16
    const uintptr_t kv_fp8, // ......... (total_kv, 576) fp8
    const uintptr_t kv_fp8_scale, // ... (1,) f32
    const uintptr_t kv_mxfp4, // ....... (total_kv, 288) fp4x2
    const uintptr_t kv_mxfp4_scale, // . (total_kv, 18) uint8_t
    uintptr_t out, // .................. (bs, 16, 512) bf16
    u32 bs
) {
    kernel<<<bs, THREADS_PER_BLOCK>>>(
        (const u16*)q,
        (const u32*)kv_indptr,
        (const u16*)kv_bf16,
        (const u8*)kv_fp8, (const f32*)kv_fp8_scale,
        (const u8*)kv_mxfp4, (const u8*)kv_mxfp4_scale,
        (u16*)out,
        bs
    );
}
"""

CUDA_SRC_2 = r"""
#define CDIV(a, b) (((a) + (b) - 1) / (b))

using u8 = uint8_t;
using u16 = uint16_t;
using u16x32 = u16 __attribute__((ext_vector_type(32)));
using u32 = uint32_t;
using u32x4 = u32 __attribute__((ext_vector_type(4)));
using u32x8 = u32 __attribute__((ext_vector_type(8)));
using u64 = uint64_t;
using f32 = float;
using f32x4 = f32 __attribute__((ext_vector_type(4)));
using bf16x2 = __bf16 __attribute__((ext_vector_type(2)));
using bf16x4 = __bf16 __attribute__((ext_vector_type(4)));
using bf16x8 = __bf16 __attribute__((ext_vector_type(8)));

__device__ inline f32 fast_exp(f32 x) {
    constexpr f32 LOG2E = 1.4426950408889634f;
    return __builtin_amdgcn_exp2f(x * LOG2E);
}

__device__ inline u16 f32_to_bf16(f32 v) {
    return (u16)(__builtin_bit_cast(u32, v) >> 16);
}

__device__ inline float bf16_to_f32(uint16_t v) {
    return __builtin_bit_cast(float, (uint32_t)v << 16);
}

// Quantize a group of 32 bf16 values -> E8M0 scale (biased uint8)
__device__ inline uint8_t bf16x32_to_scale_e8m0(const u16x32* row) {
    uint32_t amax_u16x2_packed = 0;
    const uint32_t* row32 = (const uint32_t*)row;
#pragma unroll
    for (int i = 0; i < 16; i++) {
        // abs both bf16 in one instruction
        uint32_t v = row32[i] & 0x7FFF7FFFu;
        // max both bf16 in one instruction
        asm("v_pk_max_u16 %0, %0, %1" : "+v"(amax_u16x2_packed) : "v"(v));
    }
    uint16_t amax_u16 = max((uint16_t)amax_u16x2_packed, (uint16_t)(amax_u16x2_packed >> 16));
    // Round exponent up when mantissa fraction >= 0.75
    // (bf16 has 7 mantissa bits; 0x20 = 2^5 matches the f32 0x200000 = 2^21 rounding)
    int biased_exp = ((amax_u16 + 0x20u) >> 7) & 0xFF;
    int e8m0 = biased_exp - 2;
    return (uint8_t)max(0, min(254, e8m0));
}

// Matches `fp4_utils.e8m0_to_f32`: scale_e8m0 << 23 as f32, with special cases
__device__ inline float e8m0_to_f32(uint8_t scale) {
    return __builtin_bit_cast(float, (uint32_t)scale << 23);
}

__device__ inline __bf16 fp4_e2m1_to_bf16(u8 nibble, u8 scale) {
    u8 s = (nibble >> 3) & 1;
    u8 e = (nibble >> 1) & 3;
    u8 m = nibble & 1;
    f32 val = (e == 0) ? m * 0.5f : exp2f((f32)e - 1.0f) * (1.0f + m * 0.5f);
    val = s ? -val : val;
    u16 ret = f32_to_bf16(val * e8m0_to_f32(scale));
    return *(const __bf16*)&ret;
}

__device__ inline uint8_t f32_to_fp4_e2m1(float v) {
    // CAREFUL: We need to include the fix from https://github.com/ROCm/aiter/pull/975 / https://github.com/ROCm/aiter/pull/2249
    // FP4 e2m1: EBITS=2, MBITS=1, EXP_BIAS=1
    // denorm_exp = (127-1) + (23-1) + 1 = 149  =>  denorm_mask = 149 << 23 = 0x4A800000
    // val_to_add = ((1-127) << 23) + ((1<<21)-1) = 0xC11FFFFF  (int32 wrapping)
    constexpr float    FP4_MAX_NORMAL    = 6.0f;
    constexpr float    FP4_MIN_NORMAL    = 1.0f;
    constexpr uint32_t DENORM_MASK_INT   = 0x4A800000u;
    constexpr float    DENORM_MASK_FLOAT = __builtin_bit_cast(float, DENORM_MASK_INT);
    constexpr uint32_t VAL_TO_ADD        = 0xC11FFFFFu;

    uint8_t  s_bit    = (uint8_t)((__builtin_bit_cast(uint32_t, v) >> 28) & 0x8u);
    uint32_t abs_bits = __builtin_bit_cast(uint32_t, v) & 0x7FFFFFFFu;
    float    abs_v    = __builtin_bit_cast(float, abs_bits);

    uint8_t result;
    if (abs_v >= FP4_MAX_NORMAL) {
        // saturate branch
        result = 0x7u;
    } else if (abs_v < FP4_MIN_NORMAL) {
        // denormal branch: float trick to extract low bits
        uint32_t d = __builtin_bit_cast(uint32_t, abs_v + DENORM_MASK_FLOAT) - DENORM_MASK_INT;
        result = (uint8_t)d;
    } else {
        // normal branch: bias-adjust exponent + round-to-nearest
        uint32_t x     = abs_bits;
        uint32_t m_odd = (x >> 22) & 1u;
        x = x + VAL_TO_ADD + m_odd;
        result = (uint8_t)(x >> 22);
    }
    return s_bit | result;
}

// Take an f32 and its block's scaling term, and convert to fp4_e2m1
// The fp4_e2m1 will be in the low nibble of the return value
__device__ inline uint8_t f32_to_fp4_e2m1_scale(float v, uint8_t scale) {
    float f32_scale = e8m0_to_f32((uint8_t)(254u - scale)); // 2^(-(exp_unb)) = 2^(127 - scale)
    return f32_to_fp4_e2m1(v * f32_scale);
}

__device__ inline f32x4 mfma_bf16_16x16x32(bf16x8 a, bf16x8 b, f32x4 acc) {
#ifdef __gfx950__
    // return __builtin_amdgcn_mfma_f32_16x16x32_bf16(a, b, acc, 0, 0, 0);
    // AGPRs are bad, we can't multiply alpha into them. Pin VGPR.
    asm("v_mfma_f32_16x16x32_bf16 %0, %1, %2, %0"
        : "+v"(acc)
        : "v"(a), "v"(b));
    return acc;
#else
    acc = __builtin_amdgcn_mfma_f32_16x16x16bf16_1k(*(bf16x4*)&a, *(bf16x4*)&b, acc, 0, 0, 0);
    acc = __builtin_amdgcn_mfma_f32_16x16x16bf16_1k(*((bf16x4*)&a+1), *((bf16x4*)&b+1), acc, 0, 0, 0);
    return acc;
#endif
}

__device__ inline void global_to_lds_1024b(const u8* global_base, u8* lds_base, u32 lane_id) {
    u32 lds_off = __builtin_amdgcn_readfirstlane((u32)(uintptr_t)lds_base);
#ifdef __gfx950__
    asm volatile(
        "s_mov_b32 m0, %0\n\t"
        "global_load_lds_dwordx4 %1, off\n\t"
        :: "s"(lds_off), "v"((const void*)(global_base + lane_id * 16))
        : "memory", "m0"
    );
#else
    #pragma unroll
    for (u32 sub = 0; sub < 4; sub++) {
        asm volatile(
            "s_mov_b32 m0, %0\n\t"
            "global_load_lds_dword %1, off\n\t"
            :: "s"(lds_off + sub * 256), "v"((const void*)(global_base + sub * 256 + lane_id * 4))
            : "memory", "m0"
        );
    }
#endif
}

__device__ inline void global_to_lds_256b(const u8* global_base, u8* lds_base, u32 lane_id) {
    u32 lds_off = __builtin_amdgcn_readfirstlane((u32)(uintptr_t)lds_base);
    asm volatile(
        "s_mov_b32 m0, %0\n\t"
        "global_load_lds_dword %1, off\n\t"
        :: "s"(lds_off), "v"((const void*)(global_base + lane_id * 4))
        : "memory", "m0"
    );
}

__device__ inline f32 dpp_reduce_max_16(f32 v) {
    #define DPP_MOV(x, ctrl) __builtin_bit_cast(f32, \
        __builtin_amdgcn_mov_dpp(__builtin_bit_cast(u32, x), ctrl, 0xF, 0xF, false))
    v = fmaxf(v, DPP_MOV(v, 0xB1));   // quad XOR 1
    v = fmaxf(v, DPP_MOV(v, 0x4E));   // quad XOR 2
    v = fmaxf(v, DPP_MOV(v, 0x124));  // row_ror:4
    v = fmaxf(v, DPP_MOV(v, 0x128));  // row_ror:8
    return v;
    #undef DPP_MOV
}

__device__ inline f32 dpp_reduce_sum_16(f32 v) {
    #define DPP_MOV(x, ctrl) __builtin_bit_cast(f32, \
        __builtin_amdgcn_mov_dpp(__builtin_bit_cast(u32, x), ctrl, 0xF, 0xF, false))
    v += DPP_MOV(v, 0xB1);
    v += DPP_MOV(v, 0x4E);
    v += DPP_MOV(v, 0x124);
    v += DPP_MOV(v, 0x128);
    return v;
    #undef DPP_MOV
}

// Constants
constexpr f32 NEG_INF = -1e30f;
constexpr u32 THREADS_PER_WARP = 64;
#ifdef __gfx950__
constexpr u32 NUM_WARPS = 4;
#else
constexpr u32 NUM_WARPS = 2;
#endif
constexpr u32 THREADS_PER_BLOCK = THREADS_PER_WARP * NUM_WARPS;
constexpr u32 SCALE_GROUP_SIZE = 32;
constexpr u32 KV_TILE_DIM = 32;

// DeepSeek Parameters
constexpr u32 N_HEADS = 16;
constexpr u32 QK_HEAD_DIM = 576;
constexpr u32 V_HEAD_DIM = 512;

__global__ __launch_bounds__(THREADS_PER_BLOCK, 1)
void kernel(
    const u16* __restrict__ q,
    const u32* __restrict__ kv_indptr,
    const u16* __restrict__ kv_bf16,
    const u8* __restrict__ kv_fp8, const f32* __restrict__ kv_fp8_scale,
    const u8* __restrict__ kv_mxfp4, const u8* __restrict__ kv_mxfp4_scale,
    u16* __restrict__ out,
    int bs,
    u32 num_splits,
    u16* __restrict__ partial_values,
    f32* __restrict__ partial_max,
    f32* __restrict__ partial_denom,
    u32* __restrict__ counter
) {
    u32 batch_idx = blockIdx.x / num_splits;
    u32 split_idx = blockIdx.x % num_splits;
    u32 warp_id = threadIdx.x / THREADS_PER_WARP;
    u32 lane_id = threadIdx.x % THREADS_PER_WARP;
    u32 lane_col = lane_id % 16;
    u32 lane_rowgroup = lane_id / 16; // rows \in 4*lane_rowgroup + {0,1,2,3}

    // Pre-load and quantize Q into FP4 registers (constant across all KV tiles)
    const u16* q_batch_item = q + batch_idx * N_HEADS * QK_HEAD_DIM;
    constexpr u32 SCORE_MFMA_DOT_DIM = 128;
    constexpr u32 SCORE_MFMA_ITERS = CDIV(QK_HEAD_DIM, SCORE_MFMA_DOT_DIM);
    u32x8 q_lanes[SCORE_MFMA_ITERS];
    u8 q_scale_lanes[SCORE_MFMA_ITERS];
    #pragma unroll
    for (u32 k = 0; k < SCORE_MFMA_ITERS; k++) {
        u32 dim_base = k * SCORE_MFMA_DOT_DIM + lane_rowgroup * SCALE_GROUP_SIZE;
        if (k < SCORE_MFMA_ITERS - 1 || dim_base + SCALE_GROUP_SIZE <= QK_HEAD_DIM) {
            const u16* q_chunk = q_batch_item + lane_col * QK_HEAD_DIM + dim_base;
            q_scale_lanes[k] = bf16x32_to_scale_e8m0((const u16x32*)q_chunk);
            f32 s = e8m0_to_f32(q_scale_lanes[k]);
            #pragma unroll
            for (u32 r = 0; r < 4; r++) {
#ifdef __gfx950__
                const bf16x2* v2 = (const bf16x2*)(q_chunk + r * 8);
                q_lanes[k][r] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(0, v2[0], s, 0);
                q_lanes[k][r] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q_lanes[k][r], v2[1], s, 1);
                q_lanes[k][r] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q_lanes[k][r], v2[2], s, 2);
                q_lanes[k][r] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q_lanes[k][r], v2[3], s, 3);
#else
                q_lanes[k][r] = 0;
                #pragma unroll
                for (u32 i = 0; i < 8; i++) {
                    u8 nibble = f32_to_fp4_e2m1_scale(bf16_to_f32(q_chunk[r * 8 + i]), q_scale_lanes[k]);
                    q_lanes[k][r] |= (u32)(nibble & 0xF) << (4 * i);
                }
#endif
            }
        } else {
            q_lanes[k] = {};
            q_scale_lanes[k] = 0;
        }
    }

    u32 full_kv_start = kv_indptr[batch_idx];
    u32 full_kv_end = kv_indptr[batch_idx + 1];
    u32 full_kv_range = full_kv_end - full_kv_start;
    u32 split_len = full_kv_range / num_splits;
    u32 kv_start = full_kv_start + split_idx * split_len;
    u32 kv_end = kv_start + split_len;

    f32 softmax_max[4] = {NEG_INF, NEG_INF, NEG_INF, NEG_INF};
    f32 softmax_denom[4] = {};

    constexpr u32 VALUE_MFMA_V_TILE_DIM = 16;
    constexpr u32 VALUE_MFMA_DOT_DIM = 32;
    constexpr u32 NUM_VALUE_TILES = V_HEAD_DIM / VALUE_MFMA_V_TILE_DIM;
    f32x4 value_out_lanes[NUM_VALUE_TILES] = {};

    // KV Tiles
    constexpr u32 PREFETCH = 2;
    constexpr u32 KV_SCALE_STRIDE = CDIV(QK_HEAD_DIM / SCALE_GROUP_SIZE, 8) * 8;

    // Shared Memory
    __shared__ union {
        // kv iterations
        struct {
            u8 kv_data[NUM_WARPS][PREFETCH][KV_TILE_DIM][QK_HEAD_DIM / 2];
            u8 kv_scale[NUM_WARPS][PREFETCH][KV_TILE_DIM][KV_SCALE_STRIDE];
            u16 weights[NUM_WARPS][N_HEADS][KV_TILE_DIM];
        } tile;
        // merge before write
        struct {
            f32 warp_max[NUM_WARPS][N_HEADS];
            f32 warp_denom[NUM_WARPS][N_HEADS];
            u16 values[NUM_WARPS][N_HEADS][V_HEAD_DIM];
        } merge;
    } lds;

    u32 kv_range = kv_end - kv_start;
    u32 kv_num_iters = kv_range / KV_TILE_DIM / NUM_WARPS;
    if (kv_range % (NUM_WARPS * KV_TILE_DIM) != 0 || kv_num_iters < PREFETCH - 1) {
        __builtin_trap();
    }
    u32 warp_kv_offset = warp_id * kv_num_iters;

    // (u32 kv_tile_start = kv_start; kv_tile_start < kv_end; kv_tile_start += KV_TILE_DIM) {
    auto load_lds = [&](uint32_t kv_iter) -> void {
        u32 kv_tile_start = kv_start + (warp_kv_offset + kv_iter) * KV_TILE_DIM;
        u32 kv_buf = kv_iter % 2;
        constexpr u32 KV_DATA_BYTES = KV_TILE_DIM * (QK_HEAD_DIM / 2);
        constexpr u32 KV_DATA_U128 = KV_DATA_BYTES / 16;
        static_assert(KV_DATA_U128 % THREADS_PER_WARP == 0);
        constexpr u32 KV_DATA_ITERS = KV_DATA_U128 / THREADS_PER_WARP;
        const u8* kv_data_src = (const u8*)(kv_mxfp4 + (u64)kv_tile_start * (QK_HEAD_DIM / 2));
        u8* lds_data_dst = (u8*)&lds.tile.kv_data[warp_id][kv_buf];
        #pragma unroll
        for (u32 i = 0; i < KV_DATA_ITERS; i++) {
            global_to_lds_1024b(kv_data_src + i * 1024, lds_data_dst + i * 1024, lane_id);
        }

        constexpr u32 KV_SCALE_BYTES = KV_TILE_DIM * KV_SCALE_STRIDE;
        constexpr u32 KV_SCALE_U32 = KV_SCALE_BYTES / 4;
        constexpr u32 KV_SCALE_ITERS = CDIV(KV_SCALE_U32, THREADS_PER_WARP);
        const u8* kv_scale_src = (const u8*)(kv_mxfp4_scale + (u64)kv_tile_start * KV_SCALE_STRIDE);
        u8* lds_scale_data_dst = (u8*)&lds.tile.kv_scale[warp_id][kv_buf];
        #pragma unroll
        for (u32 i = 0; i < KV_SCALE_ITERS; i++) {
            global_to_lds_256b(kv_scale_src + i * 256, lds_scale_data_dst + i * 256, lane_id);
        }
    };

    #pragma unroll
    for (u32 i = 0; i < PREFETCH - 1; i++) {
        load_lds(i);
    }
    for (u32 kv_iter = 0; kv_iter < kv_num_iters; kv_iter++) {
        // ==========
        // Cooperative load of KV Cache MXFP4 tile to LDS
        // ==========

        // Wait for the prefetched LDS to land
        // asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
        static_assert(PREFETCH <= 2); // Need custom vmcnt flag for PREFETCH > 2
        __builtin_amdgcn_s_waitcnt(0x3f70);

        // Prefetch the next LDS
        u32 kv_buf = kv_iter % 2;
        u32 kv_prefetch_iter = kv_iter + PREFETCH - 1;
        if (kv_prefetch_iter < kv_num_iters) {
            load_lds(kv_prefetch_iter);
        }

        // ==========
        // scores = per-head query @ key * 1/sqrt(d)
        // scores = Q(N_HEADS,QK_HEAD_DIM) @ K(32,QK_HEAD_DIM)^T * 1/sqrt(QK_HEAD_DIM)
        // ==========

        constexpr u32 SCORE_MFMA_KV_TILE_DIM = 16;
        static_assert(KV_TILE_DIM % SCORE_MFMA_KV_TILE_DIM == 0);
        constexpr u32 MFMA_PER_KV_TILE_DIM = KV_TILE_DIM / SCORE_MFMA_KV_TILE_DIM;
        f32x4 scores[MFMA_PER_KV_TILE_DIM] = {};

        #pragma unroll
        for (u32 score_mfma_idx = 0; score_mfma_idx < SCORE_MFMA_ITERS; score_mfma_idx++) {
            #pragma unroll
            for (u32 kv_mfma_idx = 0; kv_mfma_idx < MFMA_PER_KV_TILE_DIM; kv_mfma_idx++) {
                u32 tok = kv_mfma_idx * 16 + lane_col;
                u32 kv_byte_offset = score_mfma_idx * (SCORE_MFMA_DOT_DIM / 2) + lane_rowgroup * 16;

                u32x8 kv_reg;
                u8 kv_scale = 0;
                if (score_mfma_idx < SCORE_MFMA_ITERS - 1 || kv_byte_offset + 16 <= QK_HEAD_DIM / 2) {
                    *(u32x4*)&kv_reg = *(const u32x4*)&lds.tile.kv_data[warp_id][kv_buf][tok][kv_byte_offset];
                    kv_scale = lds.tile.kv_scale[warp_id][kv_buf][tok][score_mfma_idx * (SCORE_MFMA_DOT_DIM / SCALE_GROUP_SIZE) + lane_rowgroup];
                } else {
                    *(u32x4*)&kv_reg = {};
                }

#ifdef __gfx950__
                scores[kv_mfma_idx] = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
                    q_lanes[score_mfma_idx],
                    kv_reg,
                    scores[kv_mfma_idx],
                    4, 4,
                    0, q_scale_lanes[score_mfma_idx],
                    0, kv_scale
                );
#else
                #pragma unroll
                for (u32 sub = 0; sub < 4; sub++) {
                    bf16x8 q_bf16, kv_bf16_lane;
                    #pragma unroll
                    for (u32 i = 0; i < 8; i++) {
                        u32 qi = sub * 8 + i;
                        u8 q_nibble = (q_lanes[score_mfma_idx][qi / 8] >> (4 * (qi % 8))) & 0xF;
                        q_bf16[i] = fp4_e2m1_to_bf16(q_nibble, q_scale_lanes[score_mfma_idx]);

                        u8 kv_byte = ((const u8*)&kv_reg)[sub * 4 + i / 2];
                        u8 kv_nibble = (i % 2 == 0) ? (kv_byte & 0xF) : (kv_byte >> 4);
                        kv_bf16_lane[i] = fp4_e2m1_to_bf16(kv_nibble, kv_scale);
                    }
                    scores[kv_mfma_idx] = mfma_bf16_16x16x32(q_bf16, kv_bf16_lane, scores[kv_mfma_idx]);
                }
#endif
            }
        }

        // Scale by 1 / sqrt(QK_HEAD_DIM)
        static_assert(QK_HEAD_DIM == 24 * 24); // Add constexpr sqrt to make this responsive w.r.t QK_HEAD_DIM
        constexpr f32 SM_SCALE = 1.0f / 24.0f;
        #pragma unroll
        for (u32 i = 0; i < MFMA_PER_KV_TILE_DIM; i++) {
            #pragma unroll
            for (u32 j = 0; j < 4; j++) {
                scores[i][j] *= SM_SCALE;
            }
        }

        // ==========
        // Online softmax
        // ==========

        // == MFMA thread-local KV-tile max reduction ==
        // From the MFMA, each lane's (lane_col, lane_rowgroup) holds values from 4 different heads.
        // We need to max-reduce over the MFMA KV Tiles, to get the per-lane max for each head
        f32 lane_max[4] = {scores[0][0], scores[0][1], scores[0][2], scores[0][3]};
        #pragma unroll
        for (u32 i = 1; i < MFMA_PER_KV_TILE_DIM; i++) {
            #pragma unroll
            for (u32 j = 0; j < 4; j++) {
                lane_max[j] = fmaxf(lane_max[j], scores[i][j]);
            }
        }

        // == MFMA intra-lanegroup max reduction ==
        // 16-consecutive lanes will share the same lane_rowgroup, so we reduce the max together.
        // Each lane will then get the same true max of the KV_TILE_DIM rows, for its 4 heads.
        #pragma unroll
        for (u32 i = 0; i < 4; i++) {
            lane_max[i] = dpp_reduce_max_16(lane_max[i]);
        }

        // == Update online state ==
        // new_head_max is the new global max (max of accumulated kv_tile softmax_max[i] and current kv_tile lane_max[i]).
        // alpha = exp(old_max - new_max) is the correction factor to rescale all previously accumulated values by
        // We update previously accumulated softmax_denom and softmax_max by this alpha right now.
        //   - Updating previously accumulated value_out_lanes will be done later.
        f32x4 alpha;
        #pragma unroll
        for (u32 i = 0; i < 4; i++) {
            f32 new_head_max = fmaxf(softmax_max[i], lane_max[i]);
            alpha[i] = fast_exp(softmax_max[i] - new_head_max);
            softmax_denom[i] *= alpha[i];
            softmax_max[i] = new_head_max;
        }

        // == Compute weights, thread-local KV-tile sum reduction ==
        // weight = exp(score-score_max)
        // For each score MFMA'ss output lane item, we calculate the weight for later value accumulation
        // Weights are stored in LDS (for intra-warp permutation, value lanes are transposed)
        // lane_sum is reduced across MFMA_PER_KV_TILE_DIM, for the 4 unique heads per lane.
        f32 weights[MFMA_PER_KV_TILE_DIM][4];
        f32 lane_sum[4] = {};
        #pragma unroll
        for (u32 i = 0; i < MFMA_PER_KV_TILE_DIM; i++) {
            #pragma unroll
            for (u32 j = 0; j < 4; j++) {
                weights[i][j] = fast_exp(scores[i][j] - softmax_max[j]);
                lds.tile.weights[warp_id][4 * lane_rowgroup + j][i * 16 + lane_col] = f32_to_bf16(weights[i][j]);
                lane_sum[j] += weights[i][j];
            }
        }

        // == accumulate per-lane softmax_denom ==
        // We defer cross-lane reduction to after loop
        #pragma unroll
        for (u32 i = 0; i < 4; i++) {
            softmax_denom[i] += lane_sum[i];
        }

        // ==========
        // value_out += lds_weights(N_HEADS,KV_TILE_DIM) @ kv(KV_TILE_DIM,V_HEAD_DIM)^T
        // ==========

        // Read weights for MFMA value lane: [head=lcol, tokens lgrp*8..+7]
        bf16x8 value_weight_lane = *(const bf16x8*)&lds.tile.weights[warp_id][lane_col][lane_rowgroup * 8];

        // MFMA for weights * values
        // We process in MFMA tiles of V_HEAD_DIM / VALUE_MFMA_V_TILE_DIM
        constexpr u32 FEATURES_PER_DWORD = 8;  // 4 bytes = 8 nibbles
        constexpr u32 VALUEGROUP_ITERS = V_HEAD_DIM / (VALUE_MFMA_V_TILE_DIM * FEATURES_PER_DWORD); // 512/(16*8) = 4
        static_assert(VALUEGROUP_ITERS == 4);

        #pragma unroll
        for (u32 valuegroup_idx = 0; valuegroup_idx < VALUEGROUP_ITERS; valuegroup_idx++) {
            // 16 lanes x 4 bytes = 64 bytes = 128 features per value group
            u32 byte_col_base = valuegroup_idx * 64 + lane_col * 4;
            // All 8 features in one dword share the same scale group
            // (8 features < SCALE_GROUP_SIZE=32)
            u32 scale_idx = valuegroup_idx * 4 + lane_col / 4;

            u32 data_base = (u32)(uintptr_t)&lds.tile.kv_data[warp_id][kv_buf][lane_rowgroup * 8][byte_col_base];
            u32 scale_base = (u32)(uintptr_t)&lds.tile.kv_scale[warp_id][kv_buf][lane_rowgroup * 8][scale_idx];

            // Load from LDS
            // 4 bytes (data) + 1 byte (scale) from each of 8 rows
            u32 data_reg[8];
            u32 scale_reg[8];
            __builtin_amdgcn_sched_barrier(0);
            #pragma unroll
            for (u32 row = 0; row < 8; row++) {
                asm volatile(
                    "ds_read_b32 %0, %1 offset:%c2"
                    : "=v"(data_reg[row])
                    : "v"(data_base),
                    "n"(row * (QK_HEAD_DIM / 2))
                );
                asm volatile(
                    "ds_read_u8 %0, %1 offset:%c2"
                    : "=v"(scale_reg[row])
                    : "v"(scale_base),
                    "n"(row * KV_SCALE_STRIDE)
                );
            }
            __builtin_amdgcn_sched_barrier(0);
            asm volatile("s_waitcnt lgkmcnt(0)" ::: "memory");
            __builtin_amdgcn_sched_barrier(0);

            // 4 byte columns, each producing 2 MFMAs (one for each nibble column)
            #pragma unroll
            for (u32 byte_col_offset = 0; byte_col_offset < 4; byte_col_offset++) {
                // the value tile index for lo nibble / hi nibble
                u32 lo_value_tile = valuegroup_idx * 8 + byte_col_offset * 2;
                u32 hi_value_tile = lo_value_tile + 1;
                value_out_lanes[lo_value_tile] *= alpha;
                value_out_lanes[hi_value_tile] *= alpha;

                bf16x8 value_lane_lo, value_lane_hi;
                #pragma unroll
                for (u32 row = 0; row < 8; row++) {
                    u8 scale = (u8)scale_reg[row];
#ifdef __gfx950__
                    f32 scale_f32 = e8m0_to_f32(scale);
                    u32 cvt;
                    switch (byte_col_offset) {
                        case 0: cvt = __builtin_bit_cast(u32, __builtin_amdgcn_cvt_scalef32_pk_bf16_fp4(data_reg[row], scale_f32, 0)); break;
                        case 1: cvt = __builtin_bit_cast(u32, __builtin_amdgcn_cvt_scalef32_pk_bf16_fp4(data_reg[row], scale_f32, 1)); break;
                        case 2: cvt = __builtin_bit_cast(u32, __builtin_amdgcn_cvt_scalef32_pk_bf16_fp4(data_reg[row], scale_f32, 2)); break;
                        case 3: cvt = __builtin_bit_cast(u32, __builtin_amdgcn_cvt_scalef32_pk_bf16_fp4(data_reg[row], scale_f32, 3)); break;
                    }
                    value_lane_lo[row] = __builtin_bit_cast(__bf16, (u16)cvt);
                    value_lane_hi[row] = __builtin_bit_cast(__bf16, (u16)(cvt >> 16));
#else
                    u8 nibble_lo = (data_reg[row] >> (byte_col_offset * 8)) & 0xF;
                    u8 nibble_hi = (data_reg[row] >> (byte_col_offset * 8 + 4)) & 0xF;
                    value_lane_lo[row] = fp4_e2m1_to_bf16(nibble_lo, scale);
                    value_lane_hi[row] = fp4_e2m1_to_bf16(nibble_hi, scale);
#endif
                }

                value_out_lanes[lo_value_tile] = mfma_bf16_16x16x32(value_weight_lane, value_lane_lo, value_out_lanes[lo_value_tile]);
                value_out_lanes[hi_value_tile] = mfma_bf16_16x16x32(value_weight_lane, value_lane_hi, value_out_lanes[hi_value_tile]);
            }
        }
    }

    // ==========
    // Reduce online softmax parameters across all warps, and store results in LDS
    // ==========

    // Reduce softmax_denom across lanes in this rowgroup (same value for lane_id / 16)
    #pragma unroll
    for (u32 i = 0; i < 4; i++)
        softmax_denom[i] = dpp_reduce_sum_16(softmax_denom[i]);

    // == Store per-warp online softmax counters ==
    // We've already reduced within a warp, so only the per-head leader `lane_col == 0` needs to write.
    // There are 4 rowgroups, 4 heads each, 16 heads total.
    if (lane_col == 0) {
        #pragma unroll
        for (u32 i = 0; i < 4; i++) {
            lds.merge.warp_max[warp_id][4 * lane_rowgroup + i] = softmax_max[i];
            lds.merge.warp_denom[warp_id][4 * lane_rowgroup + i] = softmax_denom[i];
        }
    }
    __syncthreads();

    f32 split_max[4];
    f32 split_denom[4];
    f32 warp_correction[4];
    #pragma unroll
    for (u32 i = 0; i < 4; i++) {
        u32 head = 4 * lane_rowgroup + i;
        split_max[i] = NEG_INF;
        #pragma unroll
        for (u32 w = 0; w < NUM_WARPS; w++) {
            split_max[i] = fmaxf(split_max[i], lds.merge.warp_max[w][head]);
        }
        split_denom[i] = 0;
        #pragma unroll
        for (u32 w = 0; w < NUM_WARPS; w++) {
            split_denom[i] += lds.merge.warp_denom[w][head] * fast_exp(lds.merge.warp_max[w][head] - split_max[i]);
        }
        warp_correction[i] = fast_exp(softmax_max[i] - split_max[i]);
    }

    #pragma unroll
    for (u32 valuegroup_idx = 0; valuegroup_idx < 4; valuegroup_idx++) {
        #pragma unroll
        for (u32 value_tile_col_offset = 0; value_tile_col_offset < 8; value_tile_col_offset++) {
            u32 value_tile_idx = valuegroup_idx * 8 + value_tile_col_offset;
            #pragma unroll
            for (u32 i = 0; i < 4; i++) {
                u32 head = 4 * lane_rowgroup + i;
                u32 vdim = valuegroup_idx * 128 + lane_col * 8 + value_tile_col_offset;
                lds.merge.values[warp_id][head][vdim] = f32_to_bf16(value_out_lanes[value_tile_idx][i] * warp_correction[i]);
            }
        }
    }
    __syncthreads();

    // ==========
    // Write partials to global
    // ==========

    u32 split_base = batch_idx * num_splits + split_idx;
    for (u32 idx = threadIdx.x; idx < N_HEADS * V_HEAD_DIM; idx += THREADS_PER_BLOCK) {
        u32 head = idx / V_HEAD_DIM;
        u32 vdim = idx % V_HEAD_DIM;
        f32 sum = 0;
        #pragma unroll
        for (u32 w = 0; w < NUM_WARPS; w++) {
            sum += bf16_to_f32(lds.merge.values[w][head][vdim]);
        }
        partial_values[split_base * N_HEADS * V_HEAD_DIM + idx] = f32_to_bf16(sum);
    }
    if (lane_col == 0 && warp_id == 0) {
        #pragma unroll
        for (u32 i = 0; i < 4; i++) {
            u32 head = 4 * lane_rowgroup + i;
            partial_max[split_base * N_HEADS + head] = split_max[i];
            partial_denom[split_base * N_HEADS + head] = split_denom[i];
        }
    }
}
template<u32 NUM_SPLITS>
__global__ __launch_bounds__(256)
void reduce_splits(
    const u16* __restrict__ partial_values,
    const f32* __restrict__ partial_max,
    const f32* __restrict__ partial_denom,
    u16* __restrict__ out,
    u32 bs
) {
    constexpr u32 ELEMENTS = N_HEADS * V_HEAD_DIM;
    u32 global_idx = blockIdx.x * 256 + threadIdx.x;
    u32 total = bs * ELEMENTS;
    if (global_idx >= total) return;

    u32 batch_idx = global_idx / ELEMENTS;
    u32 idx = global_idx % ELEMENTS;
    u32 head = idx / V_HEAD_DIM;

    // Load all data into registers first
    f32 split_max[NUM_SPLITS];
    f32 split_denom[NUM_SPLITS];
    u16 split_val[NUM_SPLITS];

    #pragma unroll
    for (u32 s = 0; s < NUM_SPLITS; s++) {
        u32 sb = batch_idx * NUM_SPLITS + s;
        split_max[s] = partial_max[sb * N_HEADS + head];
    }

    #pragma unroll
    for (u32 s = 0; s < NUM_SPLITS; s++) {
        u32 sb = batch_idx * NUM_SPLITS + s;
        split_denom[s] = partial_denom[sb * N_HEADS + head];
    }

    #pragma unroll
    for (u32 s = 0; s < NUM_SPLITS; s++) {
        u32 sb = batch_idx * NUM_SPLITS + s;
        split_val[s] = partial_values[sb * ELEMENTS + idx];
    }

    // Compute
    f32 global_max = NEG_INF;
    #pragma unroll
    for (u32 s = 0; s < NUM_SPLITS; s++)
        global_max = fmaxf(global_max, split_max[s]);

    f32 global_denom = 0;
    #pragma unroll
    for (u32 s = 0; s < NUM_SPLITS; s++)
        global_denom += split_denom[s] * fast_exp(split_max[s] - global_max);

    f32 val_sum = 0;
    #pragma unroll
    for (u32 s = 0; s < NUM_SPLITS; s++) {
        f32 correction = fast_exp(split_max[s] - global_max) / global_denom;
        val_sum += bf16_to_f32(split_val[s]) * correction;
    }

    out[global_idx] = f32_to_bf16(val_sum);
}

#define HIP_CALL(val) check((val), #val, __FILE__, __LINE__)
template <typename T> void check(T err, const char *const func, const char *const file, const int line) {
    if (err != hipSuccess) {
        fprintf(stderr, "HIP Runtime Error at: %s:%d\n", file, line);
        fprintf(stderr, "%s %s\n", hipGetErrorString(err), func);
        exit(1);
    }
}

void entry(
    const uintptr_t q, // .............. (bs, 16, 576) bf16
    const uintptr_t kv_indptr, // ...... (bs+1,) int32
    const uintptr_t kv_bf16, // ........ (total_kv, 576) bf16
    const uintptr_t kv_fp8, // ......... (total_kv, 576) fp8
    const uintptr_t kv_fp8_scale, // ... (1,) f32
    const uintptr_t kv_mxfp4, // ....... (total_kv, 288) fp4x2
    const uintptr_t kv_mxfp4_scale, // . (total_kv, 18) uint8_t
    uintptr_t out, // .................. (bs, 16, 512) bf16
    u32 bs
) {
    // Reduction Buffers
    constexpr u32 MAX_BS = 256;
    constexpr u32 MAX_SPLITS = 256;
    static u16* partial_values = nullptr;
    static f32* partial_max = nullptr;
    static f32* partial_denom = nullptr;
    static u32* counter = nullptr;
    if (!partial_values) {
        HIP_CALL(hipMalloc(&partial_values, MAX_BS * MAX_SPLITS * N_HEADS * V_HEAD_DIM * sizeof(u16)));
        HIP_CALL(hipMalloc(&partial_max, MAX_BS * MAX_SPLITS * N_HEADS * sizeof(f32)));
        HIP_CALL(hipMalloc(&partial_denom, MAX_BS * MAX_SPLITS * N_HEADS * sizeof(f32)));
        HIP_CALL(hipMalloc(&counter, MAX_BS * sizeof(u32)));
        HIP_CALL(hipMemset(counter, 0, MAX_BS * sizeof(u32)));
    }

    // Get power of 2
    u32 min_kvlen = ((const u32*)kv_indptr)[1] - ((const u32*)kv_indptr)[0];
    u32 max_splits_kv = min_kvlen / (KV_TILE_DIM * NUM_WARPS);
    u32 raw_num_splits = min(max_splits_kv, max(1u, 256u / bs));
    u32 num_splits = 1u;
    while (num_splits * 2 <= raw_num_splits) num_splits *= 2;

    kernel<<<bs * num_splits, THREADS_PER_BLOCK>>>(
        (const u16*)q,
        (const u32*)kv_indptr,
        (const u16*)kv_bf16,
        (const u8*)kv_fp8, (const f32*)kv_fp8_scale,
        (const u8*)kv_mxfp4, (const u8*)kv_mxfp4_scale,
        (u16*)out,
        bs,
        num_splits,
        partial_values,
        partial_max,
        partial_denom,
        counter
    );

    u32 total = bs * N_HEADS * V_HEAD_DIM;
    u32 reduce_blocks = CDIV(total, 256);
    if (num_splits == 1) {
        reduce_splits<1><<<reduce_blocks, 256>>>(partial_values, partial_max, partial_denom, (u16*)out, bs);
    } else if (num_splits == 4) {
        reduce_splits<4><<<reduce_blocks, 256>>>(partial_values, partial_max, partial_denom, (u16*)out, bs);
    } else if (num_splits == 8) {
        reduce_splits<8><<<reduce_blocks, 256>>>(partial_values, partial_max, partial_denom, (u16*)out, bs);
    } else if (num_splits == 16) {
        reduce_splits<16><<<reduce_blocks, 256>>>(partial_values, partial_max, partial_denom, (u16*)out, bs);
    } else if (num_splits == 64) {
        reduce_splits<64><<<reduce_blocks, 256>>>(partial_values, partial_max, partial_denom, (u16*)out, bs);
    } else {
        assert(false);
    }
}
"""


class CompiledModule:
    M: int
    module: Any
    out: torch.Tensor

    def __init__(self, M: int) -> None:
        CUDA_PRELUDE = """
#ifndef __gfx950__
#define __gfx950__
#endif
"""
        cflags = ["--offload-arch=gfx950", "-std=c++20", "-O3", "-ffast-math", "-march=native", "-funroll-loops", "-fomit-frame-pointer"]
        cflags.extend([f"-DM_DIM={M}"])
        self.M = M
        self.out = torch.empty((M, 16, 512), dtype=torch.bfloat16, device='cuda')
        cuda_src = None
        if M == 256:
            cuda_src = CUDA_SRC
        else:
            cuda_src = CUDA_SRC_2
        self.module = load_inline(
            name=f"solution_{M}",
            cpp_sources=[CPP_WRAPPER],
            cuda_sources=[CUDA_PRELUDE + cuda_src],
            functions=['entry'],
            with_cuda=True,
            verbose=False,
            extra_cuda_cflags=cflags,
            extra_cflags=cflags,
        )

    def inference(self, data: input_t) -> output_t:
        q, kv_data, qo_indptr, kv_indptr, config = data
        bs = qo_indptr.numel() - 1
        self.module.entry(
            q.data_ptr(),
            kv_indptr.data_ptr(),
            kv_data["bf16"].data_ptr(),
            kv_data["fp8"][0].data_ptr(),
            kv_data["fp8"][1].data_ptr(),
            kv_data["mxfp4"][0].data_ptr(),
            kv_data["mxfp4"][1].data_ptr(),
            self.out.data_ptr(),
            bs,
        )
        return self.out

_compiled_modules: dict[int, CompiledModule] = {}

def custom_kernel(data: input_t) -> output_t:
    global _compiled_modules
    q, kv_data, qo_indptr, kv_indptr, config = data
    bs = qo_indptr.numel() - 1

    if bs not in _compiled_modules:
        _compiled_modules[bs] = CompiledModule(bs)

    return _compiled_modules[bs].inference(data)

scrolls · 1616 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