Skip to content
KernelIndex
Search⌘K

submission 754770

Brian Sun · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

homestretch_splitk.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-754770?include=source"
interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, mxfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
AMD MXFP4 GEMMsuite of 6 cases
AMD Instinct MI355X
8.31µs
#53 of 1143
2026-04-07

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:df79d7d495574ad6dac76da0ee9bf8dbdeaad3f2f338454a6f1dd52469468fff
license declaredunknown
license concludedunknown
authorsBrian Sun
imported2026-08-15

Techniques

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

fp4FP4 quant + FP4 GEMM: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm -> bf16 C.
shared-memory__shared__ float lds_reduce[M_MAX][BLOCK_N][N_WF]; // [16][64][8] = 32KB
split-ktemplate <int K_TOTAL, int M_MAX, int N_TOTAL, int SPLIT_K>
tile-k = 128constexpr int TILE_K = 128;
tile-m = 16constexpr int TILE_M = 16;
tile-n = 16constexpr int TILE_N = 16;

Kernel source

homestretch_splitk.py1679 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

"""
FP4 quant + FP4 GEMM: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm -> bf16 C.
"""


import os
os.environ['PYTORCH_ROCM_ARCH'] = 'gfx950'
os.environ['CXX'] = 'clang++'

import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
from aiter import QuantType,dtypes
import aiter
from aiter.ops.shuffle import shuffle_weight
from aiter.ops.triton.quant import dynamic_mxfp4_quant  # #975-patched kernel
from aiter.utility.fp4_utils import e8m0_shuffle

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

// ============================================================
// Types
// ============================================================
typedef int v4i32 __attribute__((ext_vector_type(4)));
typedef int v8i32 __attribute__((ext_vector_type(8)));
typedef float v4f32 __attribute__((ext_vector_type(4)));
typedef float v16f32 __attribute__((ext_vector_type(16)));
typedef unsigned int ext_u32x2 __attribute__((ext_vector_type(2)));
typedef uint32_t __attribute__((address_space(3)))* as3_uint32_ptr;

extern "C" __device__ void llvm_amdgcn_raw_buffer_load_lds(
    v4i32 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__ v4i32 make_srsrc(const void* ptr, uint32_t range_bytes) {
    buffer_resource rsrc = {reinterpret_cast<uint64_t>(ptr), range_bytes, 0x110000};
    return *reinterpret_cast<const v4i32*>(&rsrc);
}

__device__ __forceinline__
v8i32 to_v8(v4i32 v) {
    // fp4 MFMA only reads lower 4 elements; upper 4 are ignored.
    // Union avoids 4 wasted v_mov_b32 zero-inits per call.
    union { v4i32 lo; v8i32 full; } u;
    u.lo = v;
    return u.full;
}

// ============================================================
// Constants
// ============================================================
constexpr int TILE_M  = 16;
constexpr int TILE_N  = 16;
constexpr int TILE_K  = 128;
constexpr int GROUP_SZ = 32;
constexpr int WF_SIZE  = 64;

// ============================================================
// Device helpers
// ============================================================

__device__ __forceinline__
uint8_t compute_e8m0_scale(const __hip_bfloat16* vals) {
    // Find max |val| using bf16 bit representation (no float conversion)
    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;

    // Promote to f32 bit space for rounding: bf16 is top 16 bits of f32
    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
}

// ============================================================
// All quant-done flags in a single cache line (128 bytes)
// ============================================================
struct alignas(128) FlagLine { unsigned char flags[128]; };  // 7 used, rest padding

// ============================================================
// Shape 2 split-K GEMM: 8 wf split K within CTA, each does N_SUBS=4 MFMAs.
// 16×64 output tile. A quantized once per wf, reused across 4 B tiles.
// LDS reduction across 8 K-split wf. Then workspace for cross-CTA split-K.
// Grid: (N_BLOCKS * SPLIT_K, 1, 1), Block: 8 * 64 = 512
// ============================================================
template <int K_TOTAL, int M_MAX, int N_TOTAL, int SPLIT_K>
__global__ __attribute__((amdgpu_flat_work_group_size(8 * WF_SIZE, 8 * WF_SIZE)))
void mxfp4_gemm_splitk_gemm(
    const __hip_bfloat16* __restrict__ A,
    const uint8_t*        __restrict__ B_shuf,
    const uint8_t*        __restrict__ B_scale,
    int B_scale_stride,
    float*                __restrict__ workspace  // [N_BLOCKS * SPLIT_K * M_MAX * BLOCK_N]
) {
    constexpr int K = K_TOTAL;
    constexpr int N = N_TOTAL;
    constexpr int N_WF = 8;
    constexpr int N_SUBS = 4;                      // B tiles per wf
    constexpr int BLOCK_N = TILE_N * N_SUBS;       // 64
    constexpr int N_BLOCKS = (N + BLOCK_N - 1) / BLOCK_N;  // 33
    constexpr int K_PER_SPLIT = K / SPLIT_K;       // 1024
    // 8 wf split K_PER_SPLIT: each wf handles 128 K = 1 MFMA per B tile

    const int tid       = threadIdx.x;
    const int wf_id     = tid / WF_SIZE;           // 0..7 — which K slice
    const int tid_in_wf = tid % WF_SIZE;
    const int lane      = tid_in_wf % 16;
    const int k_group   = tid_in_wf / 16;

    const int n_block   = blockIdx.x / SPLIT_K;
    const int split_idx = blockIdx.x % SPLIT_K;
    const int a_row     = lane;

    // A: each wf handles a different 128-element K slice
    const int k_start = split_idx * K_PER_SPLIT + wf_id * TILE_K;
    const __hip_bfloat16* a_ptr = A + a_row * K + k_start + k_group * GROUP_SZ;

    // B addressing base for this wf's K position
    const int c_cur = (k_start / 32) + k_group;
    const int bs_k_part = (c_cur >> 3) * 256 + (c_cur & 3) * 64 + ((c_cur >> 2) & 1) * 2;

    v4i32 b_buf[2]; uint32_t sb_buf[2];

    auto load_b_ns = [&](int buf, int ns) __attribute__((always_inline)) {
        int n_tile = n_block * N_SUBS + ns;
        if (n_tile >= N / TILE_N) return;
        int b_tile_n_base = n_tile * (K / 32);
        int b_n_sub = n_tile & 1;
        int bs_n_base = (n_tile >> 1) * (B_scale_stride * 32);
        int bs_lane_off = lane * 4 + b_n_sub;
        int b_off = (b_tile_n_base + c_cur) * 256 + lane * 16;
        b_buf[buf] = *reinterpret_cast<const v4i32*>(B_shuf + b_off);
        sb_buf[buf] = B_scale[bs_n_base + bs_k_part + bs_lane_off];
    };

    // Issue A load + first 2 B loads BEFORE quantize — B arrives during quantize
    v4i32 a_raw[4];
    const v4i32* sv = reinterpret_cast<const v4i32*>(a_ptr);
    #pragma unroll
    for (int v = 0; v < 4; v++) a_raw[v] = sv[v];
    load_b_ns(0, 0);
    if (N_SUBS > 1) load_b_ns(1, 1);

    // Quantize A (~103 cycles — B loads arrive during this window)
    v4i32 a_buf; uint32_t sa_buf;
    __hip_bfloat16* vals = reinterpret_cast<__hip_bfloat16*>(a_raw);
    quantize_group(vals, a_buf, sa_buf);

    v4f32 acc[N_SUBS];
    #pragma unroll
    for (int ns = 0; ns < N_SUBS; ns++)
        acc[ns] = {0.0f, 0.0f, 0.0f, 0.0f};

    int cur = 0;

    // Double-buffered N_SUBS loop: MFMA → prefetch next B → repeat
    #pragma unroll
    for (int ns = 0; ns < N_SUBS - 1; ns++) {
        int nxt = cur ^ 1;
        int n_tile = n_block * N_SUBS + ns;
        if (n_tile < N / TILE_N) {
            acc[ns] = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
                to_v8(a_buf), to_v8(b_buf[cur]), acc[ns],
                4, 4, 0, sa_buf, 0, sb_buf[cur]);
        }
        if (ns + 2 < N_SUBS) load_b_ns(cur, ns + 2);
        cur = nxt;
    }
    // Last N_SUB
    {
        int n_tile = n_block * N_SUBS + (N_SUBS - 1);
        if (n_tile < N / TILE_N) {
            acc[N_SUBS - 1] = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
                to_v8(a_buf), to_v8(b_buf[cur]), acc[N_SUBS - 1],
                4, 4, 0, sa_buf, 0, sb_buf[cur]);
        }
    }

    // ---- LDS reduction across 8 K-split wavefronts ----
    // Layout [row][col][wf]: 8 wf values contiguous for fast reduction reads
    __shared__ float lds_reduce[M_MAX][BLOCK_N][N_WF];  // [16][64][8] = 32KB

    {
        #pragma unroll
        for (int ns = 0; ns < N_SUBS; ns++) {
            int col = ns * TILE_N + lane;
            #pragma unroll
            for (int i = 0; i < 4; i++) {
                int row = k_group * 4 + i;
                lds_reduce[row][col][wf_id] = acc[ns][i];
            }
        }
    }
    __syncthreads();

    // All 512 threads reduce + store: 1024 elements / 512 threads = 2 per thread
    // Reduction reads 8 contiguous floats per element (2 × v4f32)
    {
        int ws_off = (n_block * SPLIT_K + split_idx) * M_MAX * BLOCK_N;
        #pragma unroll
        for (int e = tid; e < M_MAX * BLOCK_N; e += N_WF * WF_SIZE) {
            int row = e / BLOCK_N;
            int col = e % BLOCK_N;
            int abs_col = n_block * BLOCK_N + col;
            if (abs_col < N) {
                const float* src = &lds_reduce[row][col][0];
                v4f32 v0 = *reinterpret_cast<const v4f32*>(src);
                v4f32 v1 = *reinterpret_cast<const v4f32*>(src + 4);
                float sum = v0[0] + v0[1] + v0[2] + v0[3]
                          + v1[0] + v1[1] + v1[2] + v1[3];
                workspace[ws_off + e] = sum;
            }
        }
    }
}

// ============================================================
// Shape 2 split-K reduce: sum SPLIT_K partials for 16×64 tiles → bf16.
// Grid: (N_BLOCKS, 1, 1), Block: 256 (= 16 * 64 / 4 elements per thread)
// ============================================================
template <int M_MAX, int N_TOTAL, int SPLIT_K>
__global__ __attribute__((amdgpu_flat_work_group_size(256, 256)))
void mxfp4_splitk_reduce(
    const float*          __restrict__ workspace,
    __hip_bfloat16*       __restrict__ C_bf16
) {
    constexpr int N = N_TOTAL;
    constexpr int BLOCK_N = 64;
    constexpr int N_BLOCKS = (N + BLOCK_N - 1) / BLOCK_N;
    constexpr int TILE_FLOATS = M_MAX * BLOCK_N;  // 1024

    const int n_block = blockIdx.x;
    const int tid = threadIdx.x;  // 0..511
    const int base = tid * 4;     // each thread handles 4 consecutive floats

    if (base >= TILE_FLOATS) return;

    int row = base / BLOCK_N;
    int col = base % BLOCK_N;
    int abs_col = n_block * BLOCK_N + col;

    float s0 = 0, s1 = 0, s2 = 0, s3 = 0;
    #pragma unroll
    for (int s = 0; s < SPLIT_K; s++) {
        const float* src = &workspace[(n_block * SPLIT_K + s) * TILE_FLOATS + base];
        v4f32 v = *reinterpret_cast<const v4f32*>(src);
        s0 += v[0]; s1 += v[1]; s2 += v[2]; s3 += v[3];
    }

    // Bounds check for N not divisible by 128 (2112 % 128 = 64)
    if (abs_col + 3 < N) {
        int pk0, pk1;
        asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(pk0) : "v"(s0), "v"(s1));
        asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(pk1) : "v"(s2), "v"(s3));
        int gbl = row * N + abs_col;
        *reinterpret_cast<int*>(&C_bf16[gbl]) = pk0;
        *reinterpret_cast<int*>(&C_bf16[gbl + 2]) = pk1;
    } else {
        // Scalar fallback for edge
        if (abs_col < N) C_bf16[row * N + abs_col] = __float2bfloat16(s0);
        if (abs_col + 1 < N) C_bf16[row * N + abs_col + 1] = __float2bfloat16(s1);
        if (abs_col + 2 < N) C_bf16[row * N + abs_col + 2] = __float2bfloat16(s2);
        if (abs_col + 3 < N) C_bf16[row * N + abs_col + 3] = __float2bfloat16(s3);
    }
}

// ============================================================
// Shape 2 phased kernel: Phase 1 = all blocks quantize A,
// Phase 2 = blocks 0..131 run GEMM with prequantized A.
// Waterfall sync: block 0→1→2→...→6→master, all blocks wait on master.
// Grid: (256, 1, 1), Block: N_WF * 64
// ============================================================
template <int K_TOTAL, int M_MAX, int N_WF, int N_GEMM_BLOCKS, int N_GEMM_BLOCKS_EXTRA = 24, int FLAG_VAL = 1, uint64_t FLAG_PAT8 = 0x0101010101010101ULL>
__global__ __attribute__((amdgpu_flat_work_group_size(N_WF * WF_SIZE, N_WF * WF_SIZE)))
void mxfp4_gemm_shape2_phased(
    const __hip_bfloat16* __restrict__ A,
    const uint8_t*        __restrict__ B_shuf,
    const uint8_t*        __restrict__ B_scale,
    __hip_bfloat16*       __restrict__ C,
    int N, int B_scale_stride,
    v4i32*                __restrict__ A_q_out,
    uint32_t*             __restrict__ A_scale_out,
    FlagLine*             __restrict__ flag_line
) {
    constexpr int K = K_TOTAL;
    constexpr int GROUPS_PER_ROW = K / GROUP_SZ;
    constexpr int TOTAL_GROUPS = M_MAX * GROUPS_PER_ROW;
    constexpr int N_QUANT_BLOCKS = N_GEMM_BLOCKS_EXTRA;  // must be multiple of 8

    const int tid = threadIdx.x;
    const int threads_per_block = N_WF * WF_SIZE;

    // ======== Blocks 0..6: Quantize + signal + return ========
    if (blockIdx.x < N_QUANT_BLOCKS) {
        int idx = blockIdx.x * threads_per_block + tid;
        if (idx < TOTAL_GROUPS) {
            int row = idx / GROUPS_PER_ROW;
            int k_grp = idx % GROUPS_PER_ROW;
            const __hip_bfloat16* src = A + row * K + k_grp * GROUP_SZ;

            __hip_bfloat16 vals[GROUP_SZ];
            const v4i32* src_vec = reinterpret_cast<const v4i32*>(src);
            v4i32* dst_vec = reinterpret_cast<v4i32*>(vals);
            #pragma unroll
            for (int v = 0; v < 4; v++) dst_vec[v] = src_vec[v];

            v4i32 out; uint32_t scale;
            quantize_group(vals, out, scale);

            int out_idx = k_grp * M_MAX + row;
            *reinterpret_cast<volatile v4i32*>(&A_q_out[out_idx]) = out;
            *reinterpret_cast<volatile uint32_t*>(&A_scale_out[out_idx]) = scale;
        }

        __builtin_amdgcn_s_waitcnt(0);
        __syncthreads();

        if (tid == 0) {
            *reinterpret_cast<volatile unsigned char*>(&flag_line->flags[blockIdx.x]) = FLAG_VAL;
        }
        return;
    }

    // ======== Blocks 7..138: 4 bf16 iters, prefetch flag, check, branch ========
    const int n_tile = blockIdx.x - N_QUANT_BLOCKS;

    constexpr int K_ITERS_PER_WF = K / N_WF / TILE_K;  // 7
    constexpr int K_WF_TILES = N_WF * (TILE_K / 32);
    constexpr int WF_K_OFF   = TILE_K / 32;
    constexpr int A_ADVANCE = N_WF * TILE_K;
    constexpr int B_BYTE_INC = K_WF_TILES * 256;

    const int wf_id     = tid / WF_SIZE;
    const int tid_in_wf = tid % WF_SIZE;
    const int lane      = tid_in_wf % 16;
    const int k_group   = tid_in_wf / 16;

    const int m_tile  = 0;
    const int m_base  = 0;
    const int a_row   = lane % M_MAX;

    const int k_rotation = n_tile % K_ITERS_PER_WF;
    const int k_start = k_rotation * A_ADVANCE;

    const __hip_bfloat16* a_ptr = A + a_row * K + k_start + wf_id * TILE_K + k_group * GROUP_SZ;

    const int b_tile_n_base = n_tile * (K / 32);
    const int b_n_sub      = n_tile & 1;
    const int bs_n_base = (n_tile >> 1) * (B_scale_stride * 32);
    const int bs_lane_off = lane * 4 + b_n_sub;

    int c_cur = (k_start / 32) + wf_id * WF_K_OFF + k_group;
    int b_byte_off = (b_tile_n_base + c_cur) * 256 + lane * 16;

    v4f32 acc = {0.0f, 0.0f, 0.0f, 0.0f};
    v4i32    a_buf[2];
    uint32_t sa_buf[2];
    v4i32    b_buf[2];
    uint32_t sb_buf[2];
    v4i32    a_raw[2][4];

    auto issue_a_loads = [&](int buf) __attribute__((always_inline)) {
        const v4i32* src_vec = reinterpret_cast<const v4i32*>(a_ptr);
        #pragma unroll
        for (int v = 0; v < 4; v++) a_raw[buf][v] = src_vec[v];
    };
    auto quantize_a = [&](int buf) __attribute__((always_inline)) {
        __hip_bfloat16* vals = reinterpret_cast<__hip_bfloat16*>(a_raw[buf]);
        quantize_group(vals, a_buf[buf], sa_buf[buf]);
        if constexpr (M_MAX < TILE_M) { if (lane >= M_MAX) sa_buf[buf] = 0; }
    };
    auto load_a_pq = [&](int buf) __attribute__((always_inline)) {
        int pq_idx = c_cur * M_MAX + a_row;
        a_buf[buf] = A_q_out[pq_idx];
        sa_buf[buf] = A_scale_out[pq_idx];
        if constexpr (M_MAX < TILE_M) { if (lane >= M_MAX) sa_buf[buf] = 0; }
    };
    auto load_b = [&](int buf) __attribute__((always_inline)) {
        b_buf[buf] = *reinterpret_cast<const v4i32*>(B_shuf + b_byte_off);
        int k_inn = c_cur & 3;
        int k_sub_v = (c_cur >> 2) & 1;
        int k_blk = c_cur >> 3;
        sb_buf[buf] = B_scale[bs_n_base + k_blk * 256 + k_inn * 64 + bs_lane_off + k_sub_v * 2];
    };
    auto advance_k_bf16 = [&]() __attribute__((always_inline)) {
        a_ptr += A_ADVANCE;
        b_byte_off += B_BYTE_INC;
        c_cur += K_WF_TILES;
        if (a_ptr >= A + a_row * K + K) {
            a_ptr -= K;
            b_byte_off -= (K / 32) * 256;
            c_cur -= K / 32;
        }
    };
    auto advance_k_pq = [&]() __attribute__((always_inline)) {
        b_byte_off += B_BYTE_INC;
        c_cur += K_WF_TILES;
        if (c_cur >= K / 32) {
            b_byte_off -= (K / 32) * 256;
            c_cur -= K / 32;
        }
    };

    constexpr int K_THRESHOLD = 3;

    // ---- Iterations 0 to K_THRESHOLD-1: bf16 ----
    issue_a_loads(0); load_b(0);
    int cur = 0;

    #pragma unroll
    for (int ki = 0; ki < K_THRESHOLD; ki++) {
        int nxt = cur ^ 1;
        quantize_a(cur);
        advance_k_bf16();
        issue_a_loads(nxt); load_b(nxt);
        acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
            to_v8(a_buf[cur]), to_v8(b_buf[cur]), acc,
            4, 4, 0, sa_buf[cur], 0, sb_buf[cur]);
        cur = nxt;
    }

    // ---- Load flags as 3 × int64 (24 bytes), quantize (masks latency), compare ----
    constexpr int N_FLAG_WORDS = N_QUANT_BLOCKS / 8;  // must divide evenly
    const uint64_t* f64 = reinterpret_cast<const uint64_t*>(flag_line->flags);
    uint64_t fv[N_FLAG_WORDS];
    #pragma unroll
    for (int i = 0; i < N_FLAG_WORDS; i++) fv[i] = f64[i];

    quantize_a(cur);  // ~130 cycles, hides flag load latency

    bool use_pq = true;
    #pragma unroll
    for (int i = 0; i < N_FLAG_WORDS; i++) {
        if (fv[i] != FLAG_PAT8) { use_pq = false; break; }
    }

    if (use_pq) {
        // ---- Iteration K_THRESHOLD+1: cur already quantized, prefetch prequant ----
        {
            int nxt = cur ^ 1;
            advance_k_pq();
            load_a_pq(nxt); load_b(nxt);
            acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
                to_v8(a_buf[cur]), to_v8(b_buf[cur]), acc,
                4, 4, 0, sa_buf[cur], 0, sb_buf[cur]);
            cur = nxt;
        }
        // ---- Remaining prequant iterations ----
        #pragma unroll
        for (int ki = K_THRESHOLD + 1; ki < K_ITERS_PER_WF - 1; ki++) {
            int nxt = cur ^ 1;
            advance_k_pq();
            load_a_pq(nxt); load_b(nxt);
            acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
                to_v8(a_buf[cur]), to_v8(b_buf[cur]), acc,
                4, 4, 0, sa_buf[cur], 0, sb_buf[cur]);
            cur = nxt;
        }
        // Epilogue
        acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
            to_v8(a_buf[cur]), to_v8(b_buf[cur]), acc,
            4, 4, 0, sa_buf[cur], 0, sb_buf[cur]);
    } else {
        // ---- Remaining bf16 iterations (cur already quantized) ----
        // First: advance + load + MFMA (quantize already done above)
        {
            int nxt = cur ^ 1;
            advance_k_bf16();
            issue_a_loads(nxt); load_b(nxt);
            acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
                to_v8(a_buf[cur]), to_v8(b_buf[cur]), acc,
                4, 4, 0, sa_buf[cur], 0, sb_buf[cur]);
            cur = nxt;
        }
        // Remaining loop
        #pragma unroll
        for (int ki = K_THRESHOLD + 1; ki < K_ITERS_PER_WF - 1; ki++) {
            int nxt = cur ^ 1;
            quantize_a(cur);
            advance_k_bf16();
            issue_a_loads(nxt); load_b(nxt);
            acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
                to_v8(a_buf[cur]), to_v8(b_buf[cur]), acc,
                4, 4, 0, sa_buf[cur], 0, sb_buf[cur]);
            cur = nxt;
        }
        // Epilogue
        quantize_a(cur);
        acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
            to_v8(a_buf[cur]), to_v8(b_buf[cur]), acc,
            4, 4, 0, sa_buf[cur], 0, sb_buf[cur]);
    }

    // LDS reduction across N_WF wavefronts
    __shared__ float lds_reduce[N_WF][4][WF_SIZE];
    #pragma unroll
    for (int i = 0; i < 4; i++)
        lds_reduce[wf_id][i][tid_in_wf] = acc[i];
    __syncthreads();

    if (wf_id == 0) {
        float sum[4];
        #pragma unroll
        for (int i = 0; i < 4; i++) {
            sum[i] = 0.0f;
            #pragma unroll
            for (int w = 0; w < N_WF; w++)
                sum[i] += lds_reduce[w][i][tid_in_wf];
        }
        int out_col = n_tile * TILE_N + lane;
        #pragma unroll
        for (int i = 0; i < 4; i++) {
            int out_row = m_base + k_group * 4 + i;
            if ((M_MAX >= TILE_M) || (out_row < m_base + M_MAX))
                C[out_row * N + out_col] = __float2bfloat16(sum[i]);
        }
    }
}

// ============================================================
// Small-M kernel: N_WF wavefronts split K, LDS reduction.
// Each WG handles a single 16x16 output tile.
// Grid: (N/16, num_m_tiles, 1), Block: N_WF * 64
// ============================================================
template <int K_TOTAL, int M_MAX, int N_WF>
__global__ __attribute__((amdgpu_flat_work_group_size(N_WF * WF_SIZE, N_WF * WF_SIZE)))
void mxfp4_gemm_smallm(
    const __hip_bfloat16* __restrict__ A,
    const uint8_t*        __restrict__ B_shuf,
    const uint8_t*        __restrict__ B_scale,
    __hip_bfloat16*       __restrict__ C,
    int N, int B_scale_stride
) {
    constexpr int K = K_TOTAL;
    constexpr int K_ITERS_PER_WF = K / N_WF / TILE_K;

    const int tid       = threadIdx.x;
    const int wf_id     = tid / WF_SIZE;
    const int tid_in_wf = tid % WF_SIZE;
    const int lane      = tid_in_wf % 16;
    const int k_group   = tid_in_wf / 16;

    const int n_tile  = blockIdx.x;
    const int m_tile  = blockIdx.y;
    const int m_base  = m_tile * TILE_M;

    constexpr int WF_K_OFF   = TILE_K / 32;

    const int b_tile_n_base = n_tile * (K / 32);
    const int b_n_sub      = n_tile & 1;
    const int bs_n_base = (n_tile >> 1) * (B_scale_stride * 32);
    const int bs_lane_off = lane * 4 + b_n_sub;

    int c_cur = wf_id * WF_K_OFF + k_group;
    int b_byte_off = (b_tile_n_base + c_cur) * 256 + lane * 16;

    v4f32 acc = {0.0f, 0.0f, 0.0f, 0.0f};
    v4i32    a_buf;
    uint32_t sa_buf;
    v4i32    b_buf;
    uint32_t sb_buf;

    // ---- Load ALL of A into LDS via buffer_load_lds ----
    // A = 4 rows × 512 bf16 = 4KB. 256 threads × 16 bytes = 4KB. One load per thread.
    // Each wf loads one row (wf_id → row), perfectly coalesced.
    __shared__ __hip_bfloat16 lds_a[M_MAX][K];  // [4][512] = 4KB

    {
        v4i32 a_srsrc = make_srsrc(A, N * K * 2);
        int a_soff = (m_base + wf_id) * K * 2;
        int a_voff = tid_in_wf * 16;
        llvm_amdgcn_raw_buffer_load_lds(a_srsrc,
            reinterpret_cast<as3_uint32_ptr>(
                reinterpret_cast<uintptr_t>(&lds_a[wf_id][0])),
            16, a_voff, a_soff, 0, 0);
        // A load: 1 vmcnt in flight
    }

    // Issue B loads while A is in flight (2 more vmcnt)
    b_buf = *reinterpret_cast<const v4i32*>(B_shuf + b_byte_off);
    int k_inn = c_cur & 3;
    int k_sub_v = (c_cur >> 2) & 1;
    int k_blk = c_cur >> 3;
    sb_buf = B_scale[bs_n_base + k_blk * 256 + k_inn * 64 + bs_lane_off + k_sub_v * 2];

    // Wait for A buffer_load_lds (vmcnt(2) = let 2 B loads still be outstanding)
    asm volatile("s_waitcnt vmcnt(2)");
    __syncthreads();

    // Read A from LDS → registers, then quantize
    // Each thread reads its 32 bf16 from (row, K slice) in the full A tile
    {
        int my_row = lane % M_MAX;
        int my_k = wf_id * TILE_K + k_group * GROUP_SZ;
        const __hip_bfloat16* lds_src = &lds_a[my_row][my_k];
        v4i32 a_raw[4];
        const v4i32* src_vec = reinterpret_cast<const v4i32*>(lds_src);
        #pragma unroll
        for (int v = 0; v < 4; v++) a_raw[v] = src_vec[v];

        __hip_bfloat16* vals = reinterpret_cast<__hip_bfloat16*>(a_raw);
        quantize_group(vals, a_buf, sa_buf);
        sa_buf = (lane < M_MAX) ? sa_buf : 0u;
    }

    // Single MFMA
    acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
        to_v8(a_buf), to_v8(b_buf), acc,
        4, 4, 0, sa_buf, 0, sb_buf);

    // ---- LDS reduction across N_WF wavefronts ----
    if constexpr (M_MAX < TILE_M) {
        // Only k_group=0 holds valid rows. Compact LDS: [N_WF-1][M_MAX][16].
        // wf0 keeps its acc in registers — no LDS write or read for itself.
        __shared__ float lds_reduce[N_WF - 1][M_MAX][TILE_N];

        if (k_group == 0 && wf_id > 0) {
            #pragma unroll
            for (int i = 0; i < M_MAX; i++)
                lds_reduce[wf_id - 1][i][lane] = acc[i];
        }
        __syncthreads();

        if (wf_id == 0 && k_group == 0) {
            int out_col = n_tile * TILE_N + lane;
            #pragma unroll
            for (int i = 0; i < M_MAX; i++) {
                float sum = acc[i];  // start with own accumulator
                #pragma unroll
                for (int w = 0; w < N_WF - 1; w++)
                    sum += lds_reduce[w][i][lane];
                C[(m_base + i) * N + out_col] = __float2bfloat16(sum);
            }
        }
    } else {
        // wf0 skips its own LDS write/read — uses acc directly.
        __shared__ float lds_reduce[N_WF - 1][4][WF_SIZE];

        if (wf_id > 0) {
            #pragma unroll
            for (int i = 0; i < 4; i++)
                lds_reduce[wf_id - 1][i][tid_in_wf] = acc[i];
        }
        __syncthreads();

        if (wf_id == 0) {
            float sum[4];
            #pragma unroll
            for (int i = 0; i < 4; i++) {
                sum[i] = acc[i];  // start with own accumulator
                #pragma unroll
                for (int w = 0; w < N_WF - 1; w++)
                    sum[i] += lds_reduce[w][i][tid_in_wf];
            }
            int out_col = n_tile * TILE_N + lane;
            #pragma unroll
            for (int i = 0; i < 4; i++) {
                int out_row = m_base + k_group * 4 + i;
                C[out_row * N + out_col] = __float2bfloat16(sum[i]);
            }
        }
    }
}

// 32×32×64 MFMA kernel: 32×(32*N_SUBS) output tile, N_WF wf split K.
// 2 K-chunks per step → 4 MFMAs per step (2 kc × N_SUBS).
// A quantized once per kc, reused across N_SUBS MFMAs.
// Pipelined: issue loads → MFMAs → quantize arrived A.
// LDS reduction across N_WF wf at end. Direct bf16 store.
// Grid: (N/(32*N_SUBS), M/32, 1), Block: N_WF * 64
// ============================================================
template <int K_TOTAL, int M_TOTAL, int N_TOTAL, int N_WF, int N_SUBS_T = 2>
__global__ __attribute__((amdgpu_flat_work_group_size(N_WF * WF_SIZE, N_WF * WF_SIZE)))
void mxfp4_gemm_32x32(
    const __hip_bfloat16* __restrict__ A,
    const uint8_t*        __restrict__ B_shuf,
    const uint8_t*        __restrict__ B_scale,
    __hip_bfloat16*       __restrict__ C,
    int B_scale_stride
) {
    constexpr int K = K_TOTAL;
    constexpr int M = M_TOTAL;
    constexpr int N = N_TOTAL;
    constexpr int TILE_32 = 32;
    constexpr int TILE_K_32 = 64;
    constexpr int GROUP_SZ_32 = 32;
    constexpr int N_SUBS = N_SUBS_T;
    constexpr int K_CHUNKS = 2;
    constexpr int K_STEP = K_CHUNKS * TILE_K_32;  // 128 fp4 per step
    constexpr int K_ITERS_PER_WF = K / N_WF / K_STEP;

    const int tid       = threadIdx.x;
    const int wf_id     = tid / WF_SIZE;
    const int tid_in_wf = tid % WF_SIZE;
    const int lane      = tid_in_wf % 32;
    const int k_chunk   = tid_in_wf / 32;

    const int n_block   = blockIdx.x;
    const int m_tile    = blockIdx.y;
    const int n_base    = n_block * (TILE_32 * N_SUBS);
    const int m_base    = m_tile * TILE_32;

    const int my_a_row  = m_base + lane;
    const __hip_bfloat16* A_row_ptr = A + my_a_row * K;

    const int b_lane_in_tile = lane % 16;
    const int b_n_sub_tile_off = lane / 16;
    int b_tile_n_idx[N_SUBS];
    int bs_n_blk_base[N_SUBS];
    int bs_n_sub_val[N_SUBS];
    #pragma unroll
    for (int ns = 0; ns < N_SUBS; ns++) {
        int n_off = n_base + ns * TILE_32 + lane;
        b_tile_n_idx[ns] = (n_base + ns * TILE_32) / 16 + b_n_sub_tile_off;
        bs_n_blk_base[ns] = (n_off / 32) * (B_scale_stride * 32);
        bs_n_sub_val[ns] = (n_off / 16) % 2;
    }
    const int bs_n_inn = (lane % 16) * 4;

    v16f32 acc[N_SUBS];
    #pragma unroll
    for (int ns = 0; ns < N_SUBS; ns++)
        acc[ns] = {0,0,0,0, 0,0,0,0, 0,0,0,0, 0,0,0,0};


    // A: base pointer for this thread's K range (wf spaced by K_STEP, not TILE_K_32)
    const __hip_bfloat16* a_load_ptr = A_row_ptr + wf_id * K_STEP + k_chunk * GROUP_SZ_32;

    // B: c_base uses K_STEP spacing
    const int c_base = (wf_id * K_STEP) / 32 + k_chunk;
    int b_byte_off[N_SUBS];
    #pragma unroll
    for (int ns = 0; ns < N_SUBS; ns++) {
        int tile_idx_0 = b_tile_n_idx[ns] * (K / 32) + c_base;
        b_byte_off[ns] = tile_idx_0 * 256 + b_lane_in_tile * 16;
    }

    // Double-buffered: quantized A + B data for 2 k_chunks
    v4i32    a_q[2][K_CHUNKS];           // [buf][kc] quantized fp4
    uint32_t a_sc_buf[2][K_CHUNKS];     // [buf][kc] scales
    v4i32    b_reg[2][K_CHUNKS][N_SUBS]; // [buf][kc][ns]
    uint32_t b_sc_buf[2][K_CHUNKS][N_SUBS];

    v4i32 a_raw[K_CHUNKS][4];  // raw bf16 before quantize

    // B scale: c%4, (c/4)%2 constant across steps. k_blk tracked incrementally.
    const int k_inn_kc0 = c_base % 4;
    const int k_inn_kc1 = (c_base + 2) % 4;
    const int k_sub_kc0 = (c_base / 4) % 2;
    const int k_sub_kc1 = ((c_base + 2) / 4) % 2;
    int bs_const[K_CHUNKS][N_SUBS];
    #pragma unroll
    for (int ns = 0; ns < N_SUBS; ns++) {
        bs_const[0][ns] = bs_n_blk_base[ns] + k_inn_kc0 * 64 + bs_n_inn
                        + k_sub_kc0 * 2 + bs_n_sub_val[ns];
        bs_const[1][ns] = bs_n_blk_base[ns] + k_inn_kc1 * 64 + bs_n_inn
                        + k_sub_kc1 * 2 + bs_n_sub_val[ns];
    }
    constexpr int K_ADVANCE = N_WF * K_STEP;
    constexpr int C_INC = K_ADVANCE / 32;
    constexpr int B_BYTE_INC = C_INC * 256;
    constexpr int K_BLK_INC = C_INC / 8;

    int k_blk_kc0 = c_base / 8;
    int k_blk_kc1 = (c_base + 2) / 8;

    auto load_a_raw = [&]() __attribute__((always_inline)) {
        const v4i32* src0 = reinterpret_cast<const v4i32*>(a_load_ptr);
        const v4i32* src1 = reinterpret_cast<const v4i32*>(a_load_ptr + TILE_K_32);
        #pragma unroll
        for (int v = 0; v < 4; v++) a_raw[0][v] = src0[v];
        #pragma unroll
        for (int v = 0; v < 4; v++) a_raw[1][v] = src1[v];
    };

    auto load_b = [&](int buf) __attribute__((always_inline)) {
        #pragma unroll
        for (int ns = 0; ns < N_SUBS; ns++) {
            b_reg[buf][0][ns] = *reinterpret_cast<const v4i32*>(B_shuf + b_byte_off[ns]);
            b_sc_buf[buf][0][ns] = B_scale[bs_const[0][ns] + k_blk_kc0 * 256];
            b_reg[buf][1][ns] = *reinterpret_cast<const v4i32*>(B_shuf + b_byte_off[ns] + 512);
            b_sc_buf[buf][1][ns] = B_scale[bs_const[1][ns] + k_blk_kc1 * 256];
        }
    };

    auto quantize_to = [&](int kc, int buf) __attribute__((always_inline)) {
        __hip_bfloat16* vals = reinterpret_cast<__hip_bfloat16*>(a_raw[kc]);
        quantize_group(vals, a_q[buf][kc], a_sc_buf[buf][kc]);
    };

    auto advance_ptrs = [&]() __attribute__((always_inline)) {
        a_load_ptr += K_ADVANCE;
        k_blk_kc0 += K_BLK_INC;
        k_blk_kc1 += K_BLK_INC;
        #pragma unroll
        for (int ns = 0; ns < N_SUBS; ns++)
            b_byte_off[ns] += B_BYTE_INC;
    };

    // Prologue: load + quantize step 0 into buf 0
    load_a_raw();
    load_b(0);
    quantize_to(0, 0);
    quantize_to(1, 0);
    int cur = 0;

    // Main K loop: MFMA first, then issue loads (overlap), then quantize
    #pragma unroll
    for (int ki = 0; ki < K_ITERS_PER_WF - 1; ki++) {
        int nxt = cur ^ 1;

        // MFMA[kc=0, ns=0] — start matrix core
        acc[0] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
            to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][0]), acc[0],
            4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][0]
        );

        // Issue loads for next step while MFMA[0] in flight
        advance_ptrs();
        load_a_raw();
        load_b(nxt);

        // Remaining kc=0 MFMAs (A kc=0 reused across n_subs)
        #pragma unroll
        for (int ns = 1; ns < N_SUBS; ns++)
            acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
                to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][ns]), acc[ns],
                4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][ns]
            );
        // All kc=1 MFMAs (A kc=1 reused across n_subs)
        #pragma unroll
        for (int ns = 0; ns < N_SUBS; ns++)
            acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
                to_v8(a_q[cur][1]), to_v8(b_reg[cur][1][ns]), acc[ns],
                4, 4, 0, a_sc_buf[cur][1], 0, b_sc_buf[cur][1][ns]
            );

        // Quantize arrived A data
        quantize_to(0, nxt);
        quantize_to(1, nxt);

        cur = nxt;
    }

    // Final step: just MFMAs
    #pragma unroll
    for (int ns = 0; ns < N_SUBS; ns++)
        acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
            to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][ns]), acc[ns],
            4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][ns]
        );
    #pragma unroll
    for (int ns = 0; ns < N_SUBS; ns++)
        acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
            to_v8(a_q[cur][1]), to_v8(b_reg[cur][1][ns]), acc[ns],
            4, 4, 0, a_sc_buf[cur][1], 0, b_sc_buf[cur][1][ns]
        );

    // Write acc directly to LDS in row-major layout [wf_id][row][col] as f32
    __shared__ float lds_f32[N_WF][TILE_32][TILE_32 * N_SUBS];

    {
        int k_group = tid_in_wf / 32;
        #pragma unroll
        for (int ns = 0; ns < N_SUBS; ns++) {
            int col_local = ns * TILE_32 + lane;
            #pragma unroll
            for (int i = 0; i < 16; i++) {
                int row_local = k_group * 4 + (i / 4) * 8 + (i % 4);
                lds_f32[wf_id][row_local][col_local] = acc[ns][i];
            }
        }
    }
    __syncthreads();

    // All threads: 128-bit LDS reads, reduce, convert, 128-bit global write
    {
        int local_idx = tid * 8;
        int row_local = local_idx / (TILE_32 * N_SUBS);
        int col_local = local_idx % (TILE_32 * N_SUBS);

        // Read 2 groups of 4 f32 from each wf slot, sum, convert to bf16
        __hip_bfloat16 out[8];
        #pragma unroll
        for (int g = 0; g < 2; g++) {
            float sum[4] = {0, 0, 0, 0};
            #pragma unroll
            for (int w = 0; w < N_WF; w++) {
                const float* src = &lds_f32[w][row_local][col_local + g * 4];
                #pragma unroll
                for (int j = 0; j < 4; j++)
                    sum[j] += src[j];
            }
            #pragma unroll
            for (int j = 0; j < 4; j++)
                out[g * 4 + j] = __float2bfloat16(sum[j]);
        }

        // 128-bit global write (8 bf16 = 16 bytes)
        *reinterpret_cast<v4i32*>(&C[(m_base + row_local) * N + n_base + col_local]) =
            *reinterpret_cast<v4i32*>(out);
    }
}

// ============================================================
// Shape 5 phased kernel: 32×32 GEMM + prequant on idle CUs.
// Blocks 0..N_QUANT_BLOCKS-1: quantize A → global + volatile store + flag.
// Blocks N_QUANT_BLOCKS..N_QUANT_BLOCKS+N_GEMM_BLOCKS-1: GEMM with toggle flags.
// Grid: (N_QUANT_BLOCKS + N_GEMM_BLOCKS, 1, 1), Block: N_WF * 64
// ============================================================
template <int K_TOTAL, int M_TOTAL, int N_TOTAL, int N_WF, int N_SUBS_T,
          int N_GEMM_BLOCKS, int N_QUANT_BLOCKS = 16,
          int FLAG_VAL = 1, uint64_t FLAG_PAT8 = 0x0101010101010101ULL>
__global__ __attribute__((amdgpu_flat_work_group_size(N_WF * WF_SIZE, N_WF * WF_SIZE)))
void mxfp4_gemm_shape5_phased(
    const __hip_bfloat16* __restrict__ A,
    const uint8_t*        __restrict__ B_shuf,
    const uint8_t*        __restrict__ B_scale,
    __hip_bfloat16*       __restrict__ C,
    int B_scale_stride,
    v4i32*                __restrict__ A_q_out,
    uint32_t*             __restrict__ A_scale_out,
    FlagLine*             __restrict__ flag_line
) {
    constexpr int K = K_TOTAL;
    constexpr int M = M_TOTAL;
    constexpr int N = N_TOTAL;
    constexpr int GROUPS_PER_ROW = K / GROUP_SZ;
    constexpr int TOTAL_GROUPS = M * GROUPS_PER_ROW;

    const int tid = threadIdx.x;
    const int threads_per_block = N_WF * WF_SIZE;

    // ======== Quant blocks: quantize A, volatile store, set flag, return ========
    if (blockIdx.x < N_QUANT_BLOCKS) {
        int idx = blockIdx.x * threads_per_block + tid;
        if (idx < TOTAL_GROUPS) {
            int row = idx / GROUPS_PER_ROW;
            int k_grp = idx % GROUPS_PER_ROW;
            const __hip_bfloat16* src = A + row * K + k_grp * GROUP_SZ;

            __hip_bfloat16 vals[GROUP_SZ];
            const v4i32* src_vec = reinterpret_cast<const v4i32*>(src);
            v4i32* dst_vec = reinterpret_cast<v4i32*>(vals);
            #pragma unroll
            for (int v = 0; v < 4; v++) dst_vec[v] = src_vec[v];

            v4i32 out; uint32_t scale;
            quantize_group(vals, out, scale);

            int out_idx = k_grp * M + row;
            *reinterpret_cast<volatile v4i32*>(&A_q_out[out_idx]) = out;
            *reinterpret_cast<volatile uint32_t*>(&A_scale_out[out_idx]) = scale;
        }

        __builtin_amdgcn_s_waitcnt(0);
        __syncthreads();

        if (tid == 0) {
            *reinterpret_cast<volatile unsigned char*>(&flag_line->flags[blockIdx.x]) = FLAG_VAL;
        }
        return;
    }

    // ======== GEMM blocks with phased toggle ========
    constexpr int TILE_32 = 32;
    constexpr int TILE_K_32 = 64;
    constexpr int GROUP_SZ_32 = 32;
    constexpr int N_SUBS = N_SUBS_T;
    constexpr int K_CHUNKS = 2;
    constexpr int K_STEP = K_CHUNKS * TILE_K_32;  // 128 fp4 per step
    constexpr int K_ITERS_PER_WF = K / N_WF / K_STEP;

    const int wf_id     = tid / WF_SIZE;
    const int tid_in_wf = tid % WF_SIZE;
    const int lane      = tid_in_wf % 32;
    const int k_chunk   = tid_in_wf / 32;

    const int gemm_idx  = blockIdx.x - N_QUANT_BLOCKS;
    const int n_block   = gemm_idx % (N / (TILE_32 * N_SUBS));
    const int m_tile    = gemm_idx / (N / (TILE_32 * N_SUBS));
    const int n_base    = n_block * (TILE_32 * N_SUBS);
    const int m_base    = m_tile * TILE_32;

    const int my_a_row  = m_base + lane;
    const __hip_bfloat16* A_row_ptr = A + my_a_row * K;

    const int b_lane_in_tile = lane % 16;
    const int b_n_sub_tile_off = lane / 16;
    int b_tile_n_idx[N_SUBS];
    int bs_n_blk_base[N_SUBS];
    int bs_n_sub_val[N_SUBS];
    #pragma unroll
    for (int ns = 0; ns < N_SUBS; ns++) {
        int n_off = n_base + ns * TILE_32 + lane;
        b_tile_n_idx[ns] = (n_base + ns * TILE_32) / 16 + b_n_sub_tile_off;
        bs_n_blk_base[ns] = (n_off / 32) * (B_scale_stride * 32);
        bs_n_sub_val[ns] = (n_off / 16) % 2;
    }
    const int bs_n_inn = (lane % 16) * 4;

    v16f32 acc[N_SUBS];
    #pragma unroll
    for (int ns = 0; ns < N_SUBS; ns++)
        acc[ns] = {0,0,0,0, 0,0,0,0, 0,0,0,0, 0,0,0,0};

    const __hip_bfloat16* a_load_ptr = A_row_ptr + wf_id * K_STEP + k_chunk * GROUP_SZ_32;

    const int c_base = (wf_id * K_STEP) / 32 + k_chunk;
    int b_byte_off[N_SUBS];
    #pragma unroll
    for (int ns = 0; ns < N_SUBS; ns++) {
        int tile_idx_0 = b_tile_n_idx[ns] * (K / 32) + c_base;
        b_byte_off[ns] = tile_idx_0 * 256 + b_lane_in_tile * 16;
    }

    v4i32    a_q[2][K_CHUNKS];
    uint32_t a_sc_buf[2][K_CHUNKS];
    v4i32    b_reg[2][K_CHUNKS][N_SUBS];
    uint32_t b_sc_buf[2][K_CHUNKS][N_SUBS];
    v4i32 a_raw[K_CHUNKS][4];

    const int k_inn_kc0 = c_base % 4;
    const int k_inn_kc1 = (c_base + 2) % 4;
    const int k_sub_kc0 = (c_base / 4) % 2;
    const int k_sub_kc1 = ((c_base + 2) / 4) % 2;
    int bs_const[K_CHUNKS][N_SUBS];
    #pragma unroll
    for (int ns = 0; ns < N_SUBS; ns++) {
        bs_const[0][ns] = bs_n_blk_base[ns] + k_inn_kc0 * 64 + bs_n_inn
                        + k_sub_kc0 * 2 + bs_n_sub_val[ns];
        bs_const[1][ns] = bs_n_blk_base[ns] + k_inn_kc1 * 64 + bs_n_inn
                        + k_sub_kc1 * 2 + bs_n_sub_val[ns];
    }
    constexpr int K_ADVANCE = N_WF * K_STEP;
    constexpr int C_INC = K_ADVANCE / 32;
    constexpr int B_BYTE_INC = C_INC * 256;
    constexpr int K_BLK_INC = C_INC / 8;

    int k_blk_kc0 = c_base / 8;
    int k_blk_kc1 = (c_base + 2) / 8;

    // Track absolute K-group for prequant indexing
    int c_kc0 = c_base;
    int c_kc1 = c_base + 2;

    auto load_a_raw = [&]() __attribute__((always_inline)) {
        const v4i32* src0 = reinterpret_cast<const v4i32*>(a_load_ptr);
        const v4i32* src1 = reinterpret_cast<const v4i32*>(a_load_ptr + TILE_K_32);
        #pragma unroll
        for (int v = 0; v < 4; v++) a_raw[0][v] = src0[v];
        #pragma unroll
        for (int v = 0; v < 4; v++) a_raw[1][v] = src1[v];
    };

    auto load_b = [&](int buf) __attribute__((always_inline)) {
        #pragma unroll
        for (int ns = 0; ns < N_SUBS; ns++) {
            b_reg[buf][0][ns] = *reinterpret_cast<const v4i32*>(B_shuf + b_byte_off[ns]);
            b_sc_buf[buf][0][ns] = B_scale[bs_const[0][ns] + k_blk_kc0 * 256];
            b_reg[buf][1][ns] = *reinterpret_cast<const v4i32*>(B_shuf + b_byte_off[ns] + 512);
            b_sc_buf[buf][1][ns] = B_scale[bs_const[1][ns] + k_blk_kc1 * 256];
        }
    };

    auto quantize_to = [&](int kc, int buf) __attribute__((always_inline)) {
        __hip_bfloat16* vals = reinterpret_cast<__hip_bfloat16*>(a_raw[kc]);
        quantize_group(vals, a_q[buf][kc], a_sc_buf[buf][kc]);
    };

    auto read_prequant_a = [&](int buf) __attribute__((always_inline)) {
        int pq0 = c_kc0 * M + my_a_row;
        int pq1 = c_kc1 * M + my_a_row;
        a_q[buf][0] = A_q_out[pq0];
        a_sc_buf[buf][0] = A_scale_out[pq0];
        a_q[buf][1] = A_q_out[pq1];
        a_sc_buf[buf][1] = A_scale_out[pq1];
    };

    auto advance_ptrs_bf16 = [&]() __attribute__((always_inline)) {
        a_load_ptr += K_ADVANCE;
        k_blk_kc0 += K_BLK_INC;
        k_blk_kc1 += K_BLK_INC;
        c_kc0 += C_INC;
        c_kc1 += C_INC;
        #pragma unroll
        for (int ns = 0; ns < N_SUBS; ns++)
            b_byte_off[ns] += B_BYTE_INC;
    };

    auto advance_ptrs_pq = [&]() __attribute__((always_inline)) {
        k_blk_kc0 += K_BLK_INC;
        k_blk_kc1 += K_BLK_INC;
        c_kc0 += C_INC;
        c_kc1 += C_INC;
        #pragma unroll
        for (int ns = 0; ns < N_SUBS; ns++)
            b_byte_off[ns] += B_BYTE_INC;
    };

    constexpr int K_THRESHOLD = 1;

    // ---- Iterations 0 to K_THRESHOLD-1: bf16 ----
    load_a_raw();
    load_b(0);
    quantize_to(0, 0);
    quantize_to(1, 0);
    int cur = 0;

    #pragma unroll
    for (int ki = 0; ki < K_THRESHOLD; ki++) {
        int nxt = cur ^ 1;

        acc[0] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
            to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][0]), acc[0],
            4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][0]);

        advance_ptrs_bf16();
        load_a_raw();
        load_b(nxt);

        #pragma unroll
        for (int ns = 1; ns < N_SUBS; ns++)
            acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
                to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][ns]), acc[ns],
                4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][ns]);

        #pragma unroll
        for (int ns = 0; ns < N_SUBS; ns++)
            acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
                to_v8(a_q[cur][1]), to_v8(b_reg[cur][1][ns]), acc[ns],
                4, 4, 0, a_sc_buf[cur][1], 0, b_sc_buf[cur][1][ns]);

        quantize_to(0, nxt);
        quantize_to(1, nxt);

        cur = nxt;
    }

    // ---- Load flags, quantize current (masks flag latency), check ----
    constexpr int N_FLAG_WORDS = N_QUANT_BLOCKS / 8;
    const uint64_t* f64 = reinterpret_cast<const uint64_t*>(flag_line->flags);
    uint64_t fv[N_FLAG_WORDS];
    #pragma unroll
    for (int i = 0; i < N_FLAG_WORDS; i++) fv[i] = f64[i];

    // cur is already quantized from the loop above

    bool use_pq = true;
    #pragma unroll
    for (int i = 0; i < N_FLAG_WORDS; i++) {
        if (fv[i] != FLAG_PAT8) { use_pq = false; break; }
    }

    if constexpr (K_ITERS_PER_WF - K_THRESHOLD == 1) {
        // Only 1 iteration left after threshold — just epilogue, no transition needed
        // cur already has quantized data from the K_THRESHOLD loop
        #pragma unroll
        for (int ns = 0; ns < N_SUBS; ns++)
            acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
                to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][ns]), acc[ns],
                4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][ns]);
        #pragma unroll
        for (int ns = 0; ns < N_SUBS; ns++)
            acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
                to_v8(a_q[cur][1]), to_v8(b_reg[cur][1][ns]), acc[ns],
                4, 4, 0, a_sc_buf[cur][1], 0, b_sc_buf[cur][1][ns]);
    } else if (use_pq) {
        // ---- Transition iter: cur already quantized, prefetch prequant for next ----
        {
            int nxt = cur ^ 1;

            acc[0] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
                to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][0]), acc[0],
                4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][0]);

            advance_ptrs_pq();
            read_prequant_a(nxt);
            load_b(nxt);

            #pragma unroll
            for (int ns = 1; ns < N_SUBS; ns++)
                acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
                    to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][ns]), acc[ns],
                    4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][ns]);

            #pragma unroll
            for (int ns = 0; ns < N_SUBS; ns++)
                acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
                    to_v8(a_q[cur][1]), to_v8(b_reg[cur][1][ns]), acc[ns],
                    4, 4, 0, a_sc_buf[cur][1], 0, b_sc_buf[cur][1][ns]);

            cur = nxt;
        }

        // ---- Remaining prequant iterations ----
        #pragma unroll
        for (int ki = K_THRESHOLD + 1; ki < K_ITERS_PER_WF - 1; ki++) {
            int nxt = cur ^ 1;

            acc[0] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
                to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][0]), acc[0],
                4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][0]);

            advance_ptrs_pq();
            read_prequant_a(nxt);
            load_b(nxt);

            #pragma unroll
            for (int ns = 1; ns < N_SUBS; ns++)
                acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
                    to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][ns]), acc[ns],
                    4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][ns]);

            #pragma unroll
            for (int ns = 0; ns < N_SUBS; ns++)
                acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
                    to_v8(a_q[cur][1]), to_v8(b_reg[cur][1][ns]), acc[ns],
                    4, 4, 0, a_sc_buf[cur][1], 0, b_sc_buf[cur][1][ns]);

            cur = nxt;
        }

        // Epilogue
        #pragma unroll
        for (int ns = 0; ns < N_SUBS; ns++)
            acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
                to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][ns]), acc[ns],
                4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][ns]);
        #pragma unroll
        for (int ns = 0; ns < N_SUBS; ns++)
            acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
                to_v8(a_q[cur][1]), to_v8(b_reg[cur][1][ns]), acc[ns],
                4, 4, 0, a_sc_buf[cur][1], 0, b_sc_buf[cur][1][ns]);
    } else {
        // ---- Transition iter: cur already quantized, bf16 load for next ----
        {
            int nxt = cur ^ 1;

            acc[0] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
                to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][0]), acc[0],
                4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][0]);

            advance_ptrs_bf16();
            load_a_raw();
            load_b(nxt);

            #pragma unroll
            for (int ns = 1; ns < N_SUBS; ns++)
                acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
                    to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][ns]), acc[ns],
                    4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][ns]);

            #pragma unroll
            for (int ns = 0; ns < N_SUBS; ns++)
                acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
                    to_v8(a_q[cur][1]), to_v8(b_reg[cur][1][ns]), acc[ns],
                    4, 4, 0, a_sc_buf[cur][1], 0, b_sc_buf[cur][1][ns]);

            quantize_to(0, nxt);
            quantize_to(1, nxt);

            cur = nxt;
        }

        // ---- Remaining bf16 iterations ----
        #pragma unroll
        for (int ki = K_THRESHOLD + 1; ki < K_ITERS_PER_WF - 1; ki++) {
            int nxt = cur ^ 1;

            acc[0] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
                to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][0]), acc[0],
                4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][0]);

            advance_ptrs_bf16();
            load_a_raw();
            load_b(nxt);

            #pragma unroll
            for (int ns = 1; ns < N_SUBS; ns++)
                acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
                    to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][ns]), acc[ns],
                    4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][ns]);

            #pragma unroll
            for (int ns = 0; ns < N_SUBS; ns++)
                acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
                    to_v8(a_q[cur][1]), to_v8(b_reg[cur][1][ns]), acc[ns],
                    4, 4, 0, a_sc_buf[cur][1], 0, b_sc_buf[cur][1][ns]);

            quantize_to(0, nxt);
            quantize_to(1, nxt);

            cur = nxt;
        }

        // Epilogue
        quantize_to(0, cur);
        quantize_to(1, cur);
        #pragma unroll
        for (int ns = 0; ns < N_SUBS; ns++)
            acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
                to_v8(a_q[cur][0]), to_v8(b_reg[cur][0][ns]), acc[ns],
                4, 4, 0, a_sc_buf[cur][0], 0, b_sc_buf[cur][0][ns]);
        #pragma unroll
        for (int ns = 0; ns < N_SUBS; ns++)
            acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
                to_v8(a_q[cur][1]), to_v8(b_reg[cur][1][ns]), acc[ns],
                4, 4, 0, a_sc_buf[cur][1], 0, b_sc_buf[cur][1][ns]);
    }

    // LDS reduction (same as mxfp4_gemm_32x32)
    __shared__ float lds_f32[N_WF][TILE_32][TILE_32 * N_SUBS];

    {
        int k_group = tid_in_wf / 32;
        #pragma unroll
        for (int ns = 0; ns < N_SUBS; ns++) {
            int col_local = ns * TILE_32 + lane;
            #pragma unroll
            for (int i = 0; i < 16; i++) {
                int row_local = k_group * 4 + (i / 4) * 8 + (i % 4);
                lds_f32[wf_id][row_local][col_local] = acc[ns][i];
            }
        }
    }
    __syncthreads();

    {
        constexpr int TOTAL_OUT = TILE_32 * TILE_32 * N_SUBS;
        constexpr int THREADS = N_WF * WF_SIZE;
        constexpr int ELEMS = TOTAL_OUT / THREADS;

        if constexpr (ELEMS >= 8) {
            // 8 elements/thread, 128-bit write
            int local_idx = tid * 8;
            int row_local = local_idx / (TILE_32 * N_SUBS);
            int col_local = local_idx % (TILE_32 * N_SUBS);

            __hip_bfloat16 out[8];
            #pragma unroll
            for (int g = 0; g < 2; g++) {
                float sum[4] = {0, 0, 0, 0};
                #pragma unroll
                for (int w = 0; w < N_WF; w++) {
                    const float* src = &lds_f32[w][row_local][col_local + g * 4];
                    #pragma unroll
                    for (int j = 0; j < 4; j++)
                        sum[j] += src[j];
                }
                #pragma unroll
                for (int j = 0; j < 4; j++)
                    out[g * 4 + j] = __float2bfloat16(sum[j]);
            }
            *reinterpret_cast<v4i32*>(&C[(m_base + row_local) * N + n_base + col_local]) =
                *reinterpret_cast<v4i32*>(out);
        } else {
            // 4 elements/thread, 64-bit write
            int local_idx = tid * 4;
            if (local_idx < TOTAL_OUT) {
                int row_local = local_idx / (TILE_32 * N_SUBS);
                int col_local = local_idx % (TILE_32 * N_SUBS);

                float s0 = 0, s1 = 0, s2 = 0, s3 = 0;
                #pragma unroll
                for (int w = 0; w < N_WF; w++) {
                    v4f32 v = *reinterpret_cast<const v4f32*>(&lds_f32[w][row_local][col_local]);
                    s0 += v[0]; s1 += v[1]; s2 += v[2]; s3 += v[3];
                }
                int pk0, pk1;
                asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(pk0) : "v"(s0), "v"(s1));
                asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(pk1) : "v"(s2), "v"(s3));
                ext_u32x2 out_vec = {(uint32_t)pk0, (uint32_t)pk1};
                int gbl_off = (m_base + row_local) * N + n_base + col_local;
                *reinterpret_cast<ext_u32x2*>(&C[gbl_off]) = out_vec;
            }
        }
    }
}

// ============================================================
// M=32 kernel: 8 wf = 4 K-splits × 2 N-subs. 16×32 output tile.
// A reused across n_sub via L1 cache. Each wf loads its own B.
// Grid: (N/32, 2, 1), Block: 512
// ============================================================
template <int K_TOTAL, int N_TOTAL>
__global__ __attribute__((amdgpu_flat_work_group_size(512, 512)))
void mxfp4_gemm_m32(
    const __hip_bfloat16* __restrict__ A,
    const uint8_t*        __restrict__ B_shuf,
    const uint8_t*        __restrict__ B_scale,
    __hip_bfloat16*       __restrict__ C,
    int B_scale_stride
) {
    constexpr int K = K_TOTAL;
    constexpr int N = N_TOTAL;
    constexpr int K_SPLITS = K / TILE_K;   // 4
    constexpr int N_SUBS = 2;

    const int tid       = threadIdx.x;
    const int wf_id     = tid / WF_SIZE;           // 0..7
    const int tid_in_wf = tid % WF_SIZE;
    const int lane      = tid_in_wf % TILE_N;      // 0..15
    const int k_group   = tid_in_wf / TILE_N;      // 0..3

    const int k_split   = wf_id & 3;               // 0..3
    const int n_sub     = wf_id >> 2;               // 0 or 1

    const int m_base    = blockIdx.y * TILE_M;      // 0 or 16
    const int n_tile    = blockIdx.x * N_SUBS + n_sub;

    const int a_row = m_base + lane;
    const int k_off = k_split * TILE_K + k_group * GROUP_SZ;
    const int scale_thread_off = k_group * 64 + lane * 4;

    // Load A (both n_sub wf load same A — n_sub=1 hits L1 cache)
    __hip_bfloat16 vals[GROUP_SZ];
    const v4i32* src_vec = reinterpret_cast<const v4i32*>(A + a_row * K + k_off);
    v4i32* dst_vec = reinterpret_cast<v4i32*>(vals);
    #pragma unroll
    for (int v = 0; v < 4; v++) dst_vec[v] = src_vec[v];

    // Load B (each wf loads its own n_tile's B)
    int tile_idx = n_tile * (K / 32) + k_split * 4 + k_group;
    int byte_offset = tile_idx * 256 + lane * 16;
    v4i32 b_reg = *reinterpret_cast<const v4i32*>(B_shuf + byte_offset);

    int scale_base = (n_tile >> 1) * (B_scale_stride * 32) + ((k_split >> 1) * 256);
    uint32_t scale_b = B_scale[scale_base + scale_thread_off + ((k_split & 1) << 1) + (n_tile & 1)];

    // Quantize A
    v4i32 a_reg;
    uint32_t scale_a;
    quantize_group(vals, a_reg, scale_a);

    // Single MFMA
    v4f32 acc = {0.0f, 0.0f, 0.0f, 0.0f};
    acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
        to_v8(a_reg), to_v8(b_reg), acc,
        4, 4, 0, scale_a, 0, scale_b
    );

    // K-split reduction via LDS
    __shared__ float reduce_lds[N_SUBS][K_SPLITS][4][WF_SIZE];

    #pragma unroll
    for (int i = 0; i < 4; i++)
        reduce_lds[n_sub][k_split][i][tid_in_wf] = acc[i];
    __syncthreads();

    if (k_split == 0) {
        float sum[4];
        #pragma unroll
        for (int i = 0; i < 4; i++) {
            sum[i] = reduce_lds[n_sub][0][i][tid_in_wf]
                   + reduce_lds[n_sub][1][i][tid_in_wf]
                   + reduce_lds[n_sub][2][i][tid_in_wf]
                   + reduce_lds[n_sub][3][i][tid_in_wf];
        }

        int out_col = n_tile * TILE_N + lane;
        #pragma unroll
        for (int i = 0; i < 4; i++) {
            int out_row = m_base + k_group * 4 + i;
            C[out_row * N + out_col] = __float2bfloat16(sum[i]);
        }
    }
}
// Host entry point — raw pointers, C++ dispatch
// ============================================================
static int d_call_count = 0;

void run(
    uintptr_t a_ptr,
    uintptr_t b_shuf_ptr,
    uintptr_t b_scale_ptr,
    uintptr_t c_ptr,
    int M, int N, int K,
    int B_scale_stride,
    uintptr_t aq_ptr,
    uintptr_t as_ptr,
    uintptr_t flag_ptr,
    uintptr_t aq5_ptr,
    uintptr_t as5_ptr,
    uintptr_t flag5_ptr,
    uintptr_t ws_ptr
) {
    d_call_count++;
    const auto* A      = reinterpret_cast<const __hip_bfloat16*>(a_ptr);
    const auto* B_shuf = reinterpret_cast<const uint8_t*>(b_shuf_ptr);
    const auto* B_scale = reinterpret_cast<const uint8_t*>(b_scale_ptr);
    auto* C            = reinterpret_cast<__hip_bfloat16*>(c_ptr);

    // Shape 1: M<=4, K=512 — 4 wavefronts split K
    if (M <= 4 && K == 512) {
        int num_n_tiles = (N + TILE_N - 1) / TILE_N;
        mxfp4_gemm_smallm<512, 4, 4><<<dim3(num_n_tiles, 1, 1), 4 * WF_SIZE>>>(
            A, B_shuf, B_scale, C,
            N, B_scale_stride
        );
        return;
    }

    // Shape 2: M<=16, K=7168 — split-K=7, 8wf K-split, N_SUBS=4 (16×64 tile)
    if (M <= 16 && K == 7168) {
        auto* ws = reinterpret_cast<float*>(ws_ptr);
        constexpr int SPLIT_K = 7;
        constexpr int BLOCK_N = 64;
        constexpr int N_BLOCKS = (2112 + BLOCK_N - 1) / BLOCK_N;  // 33
        // Kernel 1: 8 wf split K, each does 4 MFMAs (N_SUBS=4), LDS reduce
        mxfp4_gemm_splitk_gemm<7168, 16, 2112, SPLIT_K>
            <<<dim3(N_BLOCKS * SPLIT_K, 1, 1), 8 * WF_SIZE>>>(
            A, B_shuf, B_scale, B_scale_stride, ws);
        // Kernel 2: reduce 7 partials per 16×64 block
        mxfp4_splitk_reduce<16, 2112, SPLIT_K>
            <<<dim3(N_BLOCKS, 1, 1), 256>>>(ws, C);
        return;
    }

    // Shapes 3 & 4: M=32, K=512 — 8 wf (2 m_tiles × 4 K-splits), B in LDS
    if (M == 32 && K == 512 && N == 4096) {
        mxfp4_gemm_m32<512, 4096><<<dim3(4096 / (TILE_N * 2), 2, 1), 512>>>(
            A, B_shuf, B_scale, C, B_scale_stride
        );
        return;
    }
    if (M == 32 && K == 512) {
        mxfp4_gemm_m32<512, 2880><<<dim3(2880 / (TILE_N * 2), 2, 1), 512>>>(
            A, B_shuf, B_scale, C, B_scale_stride
        );
        return;
    }

    // Shape 5: M=64, N=7168, K=2048 — phased: 16 quant blocks + 224 GEMM blocks
    if (M == 64 && K == 2048) {
        auto* A_q_out5     = reinterpret_cast<v4i32*>(aq5_ptr);
        auto* A_scale_out5 = reinterpret_cast<uint32_t*>(as5_ptr);
        auto* flag_line5   = reinterpret_cast<FlagLine*>(flag5_ptr);
        constexpr int N_GEMM_BLOCKS5 = 224;  // (7168/64) * (64/32) = 112 * 2
        constexpr int N_QUANT5 = 16;         // multiple of 8
        constexpr int TOTAL5 = N_QUANT5 + N_GEMM_BLOCKS5;
        if (d_call_count & 1)
            mxfp4_gemm_shape5_phased<2048, 64, 7168, 8, 2, N_GEMM_BLOCKS5, N_QUANT5, 1, 0x0101010101010101ULL>
                <<<dim3(TOTAL5, 1, 1), 8 * WF_SIZE>>>(
                A, B_shuf, B_scale, C, B_scale_stride, A_q_out5, A_scale_out5, flag_line5);
        else
            mxfp4_gemm_shape5_phased<2048, 64, 7168, 8, 2, N_GEMM_BLOCKS5, N_QUANT5, 2, 0x0202020202020202ULL>
                <<<dim3(TOTAL5, 1, 1), 8 * WF_SIZE>>>(
                A, B_shuf, B_scale, C, B_scale_stride, A_q_out5, A_scale_out5, flag_line5);
        return;
    }

    // Shape 6: M=256, N=3072, K=1536 — 32×96 tile, 6 wf K-split, 3 n_subs
    if (M == 256 && K == 1536) {
        mxfp4_gemm_32x32<1536, 256, 3072, 6, 3><<<dim3(3072 / 96, 256 / 32, 1), 6 * WF_SIZE>>>(
            A, B_shuf, B_scale, C, B_scale_stride
        );
        return;
    }

}

PYBIND11_MODULE(mxfp4_gemm_pt, m) {
    m.def("run", &run, "MXFP4 GEMM kernel with phased quantization");
}
"""

module = load_inline(
    name='mxfp4_gemm_pt',
    cpp_sources='',
    cuda_sources=HIP_SRC,
    with_cuda=True,
    verbose=True,
    extra_cuda_cflags=["-std=c++20", "-O3", "--offload-arch=gfx950"],
    no_implicit_headers=True,
)

def _quant_mxfp4(x, shuffle=True):
    x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
    if shuffle:
        bs_e8m0 = e8m0_shuffle(bs_e8m0)
    return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)

# Phased buffers (allocated once)
_aq_buf = None    # Shape 2: v4i32 per group (3584 × 16 bytes)
_as_buf = None    # Shape 2: uint32 per group (3584 × 4 bytes)
_flag_buf = None  # Shape 2: FlagLine: 128 bytes
_aq5_buf = None   # Shape 5: v4i32 per group (4096 × 16 bytes)
_as5_buf = None   # Shape 5: uint32 per group (4096 × 4 bytes)
_flag5_buf = None # Shape 5: FlagLine: 128 bytes
_ws_buf = None    # Shape 2 split-K: workspace [14 * 132 * 16 * 16] floats

def _ensure_phased_bufs(device):
    global _aq_buf, _as_buf, _flag_buf, _aq5_buf, _as5_buf, _flag5_buf, _ws_buf
    if _aq_buf is None:
        n = 16 * (7168 // 32)  # 3584
        _aq_buf = torch.empty(n * 4, dtype=torch.int32, device=device)
        _as_buf = torch.empty(n, dtype=torch.int32, device=device)
        _flag_buf = torch.empty(128, dtype=torch.uint8, device=device)
    if _aq5_buf is None:
        n5 = 64 * (2048 // 32)  # 4096
        _aq5_buf = torch.empty(n5 * 4, dtype=torch.int32, device=device)
        _as5_buf = torch.empty(n5, dtype=torch.int32, device=device)
        _flag5_buf = torch.empty(128, dtype=torch.uint8, device=device)
    if _ws_buf is None:
        _ws_buf = torch.empty(33 * 7 * 16 * 64, dtype=torch.float32, device=device)  # N_BLOCKS * SPLIT_K * M * BLOCK_N

def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data
    M, K = A.shape
    N = B_shuffle.size(0)
    C = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)

    if (M == 8) or (M == 16 and N == 3072) or \
       (M == 64 and N == 3072) or (M == 256 and N == 2880):
        A_q, A_scale_sh = _quant_mxfp4(A, shuffle=True)
        out_gemm = aiter.gemm_a4w4(
            A_q, B_shuffle, A_scale_sh, B_scale_sh,
            dtype=dtypes.bf16, bpreshuffle=True)
        return out_gemm

    _ensure_phased_bufs(A.device)

    module.run(
        A.data_ptr(),
        B_shuffle.data_ptr(),
        B_scale_sh.data_ptr(),
        C.data_ptr(),
        M, N, K,
        B_scale_sh.size(1),
        _aq_buf.data_ptr(),
        _as_buf.data_ptr(),
        _flag_buf.data_ptr(),
        _aq5_buf.data_ptr(),
        _as5_buf.data_ptr(),
        _flag5_buf.data_ptr(),
        _ws_buf.data_ptr(),
    )
    return C
scrolls · 1679 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