Skip to content
KernelIndex
Search⌘K

submission 755183

Will Fisher · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

attempt_v103.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-755183?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
29.3µs
#14 of 766
2026-04-07

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:e5f28330db81c8f0105f24a06d64e48f87d98dc42039091c576297f1c76524ba
license declaredunknown
license concludedunknown
authorsWill Fisher
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 f = __half22float2(v_acc[i]);

Kernel source

attempt_v103.py906 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 <hip/amd_detail/amd_hip_fp16.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 float SM_SCALE     = 1.0f / 24.0f;         // 1/sqrt(576), hardcoded for DeepSeek R1
constexpr float SM_LOG2E     = SM_SCALE * LOG2E;      // pre-multiplied, compile-time constant
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;
}

constexpr int GROUP_SZ = 32;

__device__ __forceinline__
uint8_t compute_e8m0_scale(const __hip_bfloat16* vals) {
    uint16_t mx_bits = 0;
    #pragma unroll
    for (int i = 0; i < GROUP_SZ; i++) {
        uint16_t b = *reinterpret_cast<const uint16_t*>(&vals[i]) & 0x7FFF;
        mx_bits = max(mx_bits, b);
    }
    if (mx_bits == 0) return 0;
    uint32_t bits = (uint32_t)mx_bits << 16;
    bits = (bits + 0x200000u) & 0xFF800000u;
    int e8m0 = (int)(bits >> 23) - 2;
    return (uint8_t)max(0, min(254, e8m0));
}

__device__ __forceinline__
void quantize_group(const __hip_bfloat16* vals, v4i32& out, uint32_t& scale_out) {
    uint8_t scale = compute_e8m0_scale(vals);
    scale_out = scale;
    float scale_f = __uint_as_float((uint32_t)scale << 23);

    #define PACK_WORD(w) do { \
        unsigned int packed = 0; \
        packed = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
            packed, __hip_bfloat162{vals[(w)*8+0], vals[(w)*8+1]}, scale_f, 0); \
        packed = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
            packed, __hip_bfloat162{vals[(w)*8+2], vals[(w)*8+3]}, scale_f, 1); \
        packed = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
            packed, __hip_bfloat162{vals[(w)*8+4], vals[(w)*8+5]}, scale_f, 2); \
        packed = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
            packed, __hip_bfloat162{vals[(w)*8+6], vals[(w)*8+7]}, scale_f, 3); \
        out[w] = (int)packed; \
    } while(0)

    PACK_WORD(0);
    PACK_WORD(1);
    PACK_WORD(2);
    PACK_WORD(3);
    #undef PACK_WORD
}

// 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,
    __half* __restrict__ partial_out,
    float* __restrict__ partial_max,
    float* __restrict__ partial_sum,
    __hip_bfloat16* __restrict__ final_output,
    const float* __restrict__ kv_scale_ptr
) {
    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;

    // Read fp8 kv_scale once, round UP to next power of 2 to avoid NaN
    // Hardware cvt instructions drop mantissa (E8M0), rounding scale DOWN.
    // Smaller divisor → larger quotient → overflow. Fix: round UP.
    float kv_scale_raw = *kv_scale_ptr;
    uint32_t ks_bits = __float_as_uint(kv_scale_raw);
    uint32_t ks_exp = (ks_bits >> 23) & 0xFF;
    uint32_t ks_mant = ks_bits & 0x7FFFFF;
    int safe_exp = (int)ks_exp + (ks_mant != 0 ? 1 : 0);
    float safe_scale = __uint_as_float((uint32_t)safe_exp << 23);
    int mfma_v_scale = safe_exp;
    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);

    // Buffer resource descriptors hoisted — use soffset for per-tile offset
    const int kv_len = kv_end - kv_start;
    i32x4 kv_srsrc = make_srsrc(kv_data + kv_start * KV_PACKED, kv_len * KV_PACKED);
    i32x4 sc_srsrc = make_srsrc(kv_scales + kv_start * SCALE_COLS, kv_len * SCALE_COLS);

    auto load_tile = [&](int tile_tok_offset, 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;
        const int kv_soff = tile_tok_offset * KV_PACKED;
        const int sc_soff = tile_tok_offset * SCALE_COLS;
        constexpr int KV_LOADS = (BUF_KV_BYTES + 15) / 16;
        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, kv_soff, 0, 0);
        }
        constexpr int SC_LOADS = (BUF_SC_BYTES + 15) / 16;
        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, sc_soff, 0, 0);
        }
    };

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

    // ---- Q quantize ----
    {
        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) {
                // Vectorized 128-bit loads (4 × 16 bytes = 32 bf16)
                __hip_bfloat16 qv[GROUP_SZ];
                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);
                v4i32 ar;
                uint32_t e8;
                quantize_group(qv, ar, e8);
                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] = (uint8_t)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; }

    v4i32 q_v4[FP4_K_CHUNKS];
    int q_scale_cached[FP4_K_CHUNKS];
    // SM_LOG2E = SM_LOG2E is constexpr (compile-time constant)
    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_v4[d] = qr;
            }
            if (k_group >= 2) { q_v4[4] = {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((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(
                to_v8(q_v4[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_fp8_f16(ztmp, h0, safe_scale, false);
                    ev_short2 b0123 = __builtin_amdgcn_cvt_scalef32_pk_fp8_f16(b01, h1, safe_scale, true);
                    ev_short2 b45 = __builtin_amdgcn_cvt_scalef32_pk_fp8_f16(ztmp, h2, safe_scale, false);
                    ev_short2 b4567 = __builtin_amdgcn_cvt_scalef32_pk_fp8_f16(b45, h3, safe_scale, 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]);
        }

        // Cross-warp softmax exchange
        float* scratch = reinterpret_cast<float*>(smem + OFF_SCORE);
        const int sg = k_group << 2;
        if (lane_mod16 == 0) {
            *reinterpret_cast<f32x4*>(&scratch[warp_id * NUM_HEADS + sg]) =
                (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 + sg]) =
                (f32x4){warp_sum[0], warp_sum[1], warp_sum[2], warp_sum[3]};
        }
        __syncthreads();

        // Global max: init from own warp (already in regs, skip -INFINITY)
        float my_tm[HEADS_PER_LANE] = {warp_max[0], warp_max[1], warp_max[2], warp_max[3]};
        #pragma unroll
        for (int w = 0; w < NUM_WARPS; w++) {
            f32x4 wm = *reinterpret_cast<const f32x4*>(&scratch[w * NUM_HEADS + sg]);
            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]);
        }
        // Corrected sum: peel warp 0
        f32x4 wm0 = *reinterpret_cast<const f32x4*>(&scratch[sg]);
        f32x4 ws0 = *reinterpret_cast<const f32x4*>(&scratch[NUM_WARPS * NUM_HEADS + sg]);
        float my_ts[HEADS_PER_LANE] = {
            ws0[0] * __builtin_amdgcn_exp2f(wm0[0] - my_tm[0]),
            ws0[1] * __builtin_amdgcn_exp2f(wm0[1] - my_tm[1]),
            ws0[2] * __builtin_amdgcn_exp2f(wm0[2] - my_tm[2]),
            ws0[3] * __builtin_amdgcn_exp2f(wm0[3] - my_tm[3])
        };
        #pragma unroll
        for (int w = 1; w < NUM_WARPS; w++) {
            f32x4 wm = *reinterpret_cast<const f32x4*>(&scratch[w * NUM_HEADS + sg]);
            f32x4 ws = *reinterpret_cast<const f32x4*>(&scratch[NUM_WARPS * NUM_HEADS + w * NUM_HEADS + sg]);
            #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();

        // ---- Pre-compute merge params (VALU, doesn't need v_out) ----
        float co_arr[HEADS_PER_LANE], cn_arr[HEADS_PER_LANE];
        if (ti > 0) {
            #pragma unroll
            for (int h = 0; h < HEADS_PER_LANE; h++) {
                float nm = fmaxf(running_max[h], my_tm[h]);
                co_arr[h] = __builtin_amdgcn_exp2f(running_max[h] - nm);
                cn_arr[h] = __builtin_amdgcn_exp2f(my_tm[h] - nm);
                running_sum[h] = fmaf(running_sum[h], co_arr[h], my_ts[h] * cn_arr[h]);
                running_max[h] = nm;
            }
        }

        // ---- V multiply + interleaved merge pre-scale ----
        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
            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
                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));
                }

                // Issue MFMA (matrix core)
                v_out[r] = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
                    a_reg, b_reg[cur], v_out[r], 0, 0, 0, 127, 0, mfma_v_scale);

                // Interleave merge pre-scale on VALU during MFMA execution
                if (ti > 0) {
                    #pragma unroll
                    for (int h = 0; h < HEADS_PER_LANE; h++)
                        running_v[r][h] *= co_arr[h];
                }
            }
        }

        // ---- Merge: post-MFMA (just add v_out * cn, or copy for ti==0) ----
        if (ti == 0) {
            #pragma unroll
            for (int h = 0; h < HEADS_PER_LANE; h++) {
                running_max[h] = my_tm[h];
                running_sum[h] = my_ts[h];
                #pragma unroll
                for (int i = 0; i < 4; i++) running_v[i][h] = v_out[i][h];
            }
        } else {
            #pragma unroll
            for (int i = 0; i < 4; i++)
                #pragma unroll
                for (int h = 0; h < HEADS_PER_LANE; h++)
                    running_v[i][h] += v_out[i][h] * cn_arr[h];
        }

    }

    // ---- 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;
        __half* 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] = __float2half(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 __half* __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;
    constexpr int PV_STRIDE = NUM_HEADS * V_HEAD_DIM;

    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 __half* pv = partial_out + bo * NUM_HEADS * V_HEAD_DIM
        + head_idx * V_HEAD_DIM + lane_id * ELEMS;

    // Peel first iteration
    float local_max = *pm; pm += NUM_HEADS;
    float local_sum = *ps; ps += NUM_HEADS;
    union { u32x4_vec raw; __half2 h2[4]; } cur;
    cur.raw = *reinterpret_cast<const u32x4_vec*>(pv); pv += PV_STRIDE;
    __half2 v_acc[4];
    #pragma unroll
    for (int i = 0; i < 4; i++) v_acc[i] = cur.h2[i];

    #pragma unroll
    for (int p = 1; p < NUM_PARTIALS; p++) {
        float cm = *pm; pm += NUM_HEADS;
        float cs = *ps; ps += NUM_HEADS;
        cur.raw = *reinterpret_cast<const u32x4_vec*>(pv); pv += PV_STRIDE;

        float nm = fmaxf(local_max, cm);
        float co = __builtin_amdgcn_exp2f(local_max - nm);
        float cn = __builtin_amdgcn_exp2f(cm - nm);
        local_sum = fmaf(local_sum, co, cs * cn);
        local_max = nm;
        __half2 co2 = __float2half2_rn(co);
        __half2 cn2 = __float2half2_rn(cn);
        #pragma unroll
        for (int i = 0; i < 4; i++)
            v_acc[i] = __hfma2(cur.h2[i], cn2, __hmul2(v_acc[i], co2));
    }

    // Normalize in f32, convert to bf16
    float inv = (local_sum > 0.0f) ? __builtin_amdgcn_rcpf(local_sum) : 0.0f;
    u32x4_vec out_pack;
    #pragma unroll
    for (int i = 0; i < 4; i++) {
        float2 f = __half22float2(v_acc[i]);
        asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(out_pack[i]) : "v"(f.x * inv), "v"(f.y * inv));
    }

    *reinterpret_cast<u32x4_vec*>(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 __half* __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;

    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 __half* pv_base = partial_out + bo * NUM_HEADS * V_HEAD_DIM
        + head_idx * V_HEAD_DIM + lane_id * ELEMS;

    // Peel first iteration per warp
    const int p_start = warp_id * PARTIALS_PER_WARP;
    const float* wpm = pm + p_start * NUM_HEADS;
    const float* wps = ps + p_start * NUM_HEADS;
    const __half* pv = pv_base + p_start * PV_STRIDE;
    float local_max = *wpm; wpm += NUM_HEADS;
    float local_sum = *wps; wps += NUM_HEADS;
    union { u32x4_vec raw; __half2 h2[4]; } cur;
    cur.raw = *reinterpret_cast<const u32x4_vec*>(pv); pv += PV_STRIDE;
    __half2 v_acc[4];
    #pragma unroll
    for (int i = 0; i < 4; i++) v_acc[i] = cur.h2[i];

    #pragma unroll
    for (int p = 1; p < PARTIALS_PER_WARP; p++) {
        float cm = *wpm; wpm += NUM_HEADS;
        float cs = *wps; wps += NUM_HEADS;
        cur.raw = *reinterpret_cast<const u32x4_vec*>(pv); pv += PV_STRIDE;

        float nm = fmaxf(local_max, cm);
        float co = __builtin_amdgcn_exp2f(local_max - nm);
        float cn = __builtin_amdgcn_exp2f(cm - nm);
        local_sum = fmaf(local_sum, co, cs * cn);
        local_max = nm;
        __half2 co2 = __float2half2_rn(co);
        __half2 cn2 = __float2half2_rn(cn);
        #pragma unroll
        for (int i = 0; i < 4; i++)
            v_acc[i] = __hfma2(cur.h2[i], cn2, __hmul2(v_acc[i], co2));
    }

    // F16 LDS merge: pre-scale in f16, write/read __half2, halved LDS bandwidth
    extern __shared__ char reduce_smem[];
    float* smem_max = reinterpret_cast<float*>(reduce_smem);
    float* smem_sum = smem_max + NWARPS;
    __half2* smem_vh = reinterpret_cast<__half2*>(smem_sum + NWARPS);

    if (lane_id == 0) {
        smem_max[warp_id] = local_max;
        smem_sum[warp_id] = local_sum;
    }
    __syncthreads();

    if constexpr (NWARPS <= 4) {
        // === Flat merge: pre-scale in f16, warp 0 sums in f16 ===
        float global_max = -INFINITY;
        #pragma unroll
        for (int w = 0; w < NWARPS; w++)
            global_max = fmaxf(global_max, smem_max[w]);

        __half2 scale2 = __float2half2_rn(__builtin_amdgcn_exp2f(local_max - global_max));
        local_sum *= __builtin_amdgcn_exp2f(local_max - global_max);
        #pragma unroll
        for (int i = 0; i < 4; i++) v_acc[i] = __hmul2(v_acc[i], scale2);

        __half2* wv_dst = &smem_vh[(warp_id * WAVEFRONT_SIZE + lane_id) * (ELEMS/2)];
        #pragma unroll
        for (int i = 0; i < 4; i++) wv_dst[i] = v_acc[i];
        if (lane_id == 0) smem_sum[warp_id] = local_sum;
        __syncthreads();

        if (warp_id == 0) {
            __half2 gv[4];
            #pragma unroll
            for (int i = 0; i < 4; i++) gv[i] = __float2half2_rn(0.0f);
            float gs = 0;
            #pragma unroll
            for (int w = 0; w < NWARPS; w++) {
                const __half2* wv = &smem_vh[(w * WAVEFRONT_SIZE + lane_id) * (ELEMS/2)];
                #pragma unroll
                for (int i = 0; i < 4; i++) gv[i] = __hadd2(gv[i], wv[i]);
                gs += smem_sum[w];
            }
            float inv = (gs > 0.0f) ? __builtin_amdgcn_rcpf(gs) : 0.0f;
            u32x4_vec out_pack;
            #pragma unroll
            for (int i = 0; i < 4; i++) {
                float2 f = __half22float2(gv[i]);
                asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(out_pack[i]) : "v"(f.x * inv), "v"(f.y * inv));
            }
            *reinterpret_cast<u32x4_vec*>(output + (q_start * NUM_HEADS + head_idx) * V_HEAD_DIM + lane_id * ELEMS) = out_pack;
        }
    } else {
        // === Two-tier merge in f16 ===
        constexpr int GROUP_SIZE = 4;
        constexpr int NUM_GROUPS = NWARPS / GROUP_SIZE;
        const int group_id = warp_id / GROUP_SIZE;
        const int local_warp = warp_id & (GROUP_SIZE - 1);

        float group_max = -INFINITY;
        #pragma unroll
        for (int w = group_id * GROUP_SIZE; w < (group_id + 1) * GROUP_SIZE; w++)
            group_max = fmaxf(group_max, smem_max[w]);

        __half2 scale2 = __float2half2_rn(__builtin_amdgcn_exp2f(local_max - group_max));
        local_sum *= __builtin_amdgcn_exp2f(local_max - group_max);
        #pragma unroll
        for (int i = 0; i < 4; i++) v_acc[i] = __hmul2(v_acc[i], scale2);

        __half2* wv_dst = &smem_vh[(warp_id * WAVEFRONT_SIZE + lane_id) * (ELEMS/2)];
        #pragma unroll
        for (int i = 0; i < 4; i++) wv_dst[i] = v_acc[i];
        if (lane_id == 0) smem_sum[warp_id] = local_sum;
        __syncthreads();

        // Group leader sums in f16
        if (local_warp == 0) {
            __half2 gv[4];
            #pragma unroll
            for (int i = 0; i < 4; i++) gv[i] = __float2half2_rn(0.0f);
            float gsum = 0;
            #pragma unroll
            for (int w = group_id * GROUP_SIZE; w < (group_id + 1) * GROUP_SIZE; w++) {
                const __half2* wv = &smem_vh[(w * WAVEFRONT_SIZE + lane_id) * (ELEMS/2)];
                #pragma unroll
                for (int i = 0; i < 4; i++) gv[i] = __hadd2(gv[i], wv[i]);
                gsum += smem_sum[w];
            }
            if (lane_id == 0) { smem_max[warp_id] = group_max; smem_sum[warp_id] = gsum; }
            __half2* gv_dst = &smem_vh[(warp_id * WAVEFRONT_SIZE + lane_id) * (ELEMS/2)];
            #pragma unroll
            for (int i = 0; i < 4; i++) gv_dst[i] = gv[i];
        }
        __syncthreads();

        // Warp 0 merges groups with rescaling in f16
        if (warp_id == 0) {
            float fmax = -INFINITY;
            #pragma unroll
            for (int g = 0; g < NUM_GROUPS; g++)
                fmax = fmaxf(fmax, smem_max[g * GROUP_SIZE]);

            __half2 fv[4];
            #pragma unroll
            for (int i = 0; i < 4; i++) fv[i] = __float2half2_rn(0.0f);
            float fsum = 0;
            #pragma unroll
            for (int g = 0; g < NUM_GROUPS; g++) {
                __half2 s2 = __float2half2_rn(__builtin_amdgcn_exp2f(smem_max[g * GROUP_SIZE] - fmax));
                const __half2* gv = &smem_vh[(g * GROUP_SIZE * WAVEFRONT_SIZE + lane_id) * (ELEMS/2)];
                #pragma unroll
                for (int i = 0; i < 4; i++) fv[i] = __hfma2(gv[i], s2, fv[i]);
                fsum += smem_sum[g * GROUP_SIZE] * __builtin_amdgcn_exp2f(smem_max[g * GROUP_SIZE] - fmax);
            }

            float inv = (fsum > 0.0f) ? __builtin_amdgcn_rcpf(fsum) : 0.0f;
            u32x4_vec out_pack;
            #pragma unroll
            for (int i = 0; i < 4; i++) {
                float2 f = __half22float2(fv[i]);
                asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(out_pack[i]) : "v"(f.x * inv), "v"(f.y * inv));
            }
            *reinterpret_cast<u32x4_vec*>(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,
    int batch_size, uintptr_t kv_scale_ptr) {
    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<__half*>(pout_ptr);
    auto* pmax = reinterpret_cast<float*>(pmax_ptr);
    auto* psum = reinterpret_cast<float*>(psum_ptr);
    const auto* kvsc = reinterpret_cast<const float*>(kv_scale_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, kvsc);
    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 * 2 + 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,
    int i, int max_kvseqlen, uintptr_t kv_sc) {
    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, i, kv_sc);
    else if (np == 1 && tpb == 64) run_impl<1, 32>(a, b, c, d, e, f, po, pm, ps, i, kv_sc);
    else if (np == 4 && tpb == 2) run_impl<4, 2>(a, b, c, d, e, f, po, pm, ps, i, kv_sc);
    else if (np == 4 && tpb == 16) run_impl<4, 8>(a, b, c, d, e, f, po, pm, ps, i, kv_sc);
    else if (np == 8 && tpb == 1) run_impl<8, 1>(a, b, c, d, e, f, po, pm, ps, i, kv_sc);
    else if (np == 8 && tpb == 8) run_impl<8, 4>(a, b, c, d, e, f, po, pm, ps, i, kv_sc);
    else if (np == 64 && tpb == 1) run_impl<64, 1>(a, b, c, d, e, f, po, pm, ps, i, kv_sc);
    else run_impl<1, 8>(a, b, c, d, e, f, po, pm, ps, i, kv_sc);
}

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"]
    _, fp8_scale = kv_data["fp8"]
    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.float16, 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(),
        batch_size, max_kvseqlen, fp8_scale.data_ptr())
    return output
scrolls · 906 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