Skip to content
KernelIndex
Search⌘K

submission 596238

RyanWillie · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-596238?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
10.4µs
#261 of 1143
2026-03-20

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:60e41574fc9a0817ff8500e26e5fb9190a7ebfb3dc460ae9f38cabf6de3b025d
license declaredunknown
license concludedunknown
authorsRyanWillie
imported2026-08-26

Techniques

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

fp4"Fused MXFP4 quant+GEMM");
shared-memory__shared__ uint8_t lds_a[8192]; // 2 x (16 rows * 128 BF16 * 2 bytes)
split-kvoid mxfp4_fused_gemm_m16_splitk(

Kernel source

submission.py2034 lines
# Zeus Variant: v177_k512_2wave_only
# Experiment: #177
# Technique: Cherry-pick K=512 2-wave unrolled kernel for S3/S4 (M=32) from v174
# Hypothesis: v176 showed K=512 2-wave unrolled kernel improved S3 by -3.6% and S4 by -3.9%.
#   S1 (M=4) 1-wave regressed, S2 (M=16) 4-wave multiwave regressed badly.
#   This cherry-picks ONLY the winning change: K=512 2-wave for M>16 shapes.
# Base: submissions/v125_ds_bpermute_combined.py
"""
Hybrid HIP Quant + AITER ASM GEMM + Multi-Wave Fused M=16 Kernel + K=512 2-Wave:

Architecture:
  Shape 1 (M=4, K=512): Original crossbuf kernel (unchanged from v125)
  Shape 2 (M=16, N=2112, K=7168): Multi-wave fused 8-wave 16x16 MFMA kernel
    Uses ds_bpermute_b32 for A-data lane remapping instead of LDS.
  Shape 3/4 (M=32, K=512): NEW K=512 2-wave unrolled kernel (from v174)
    - Fully unrolled 4 K_STEP=128 iterations for K=512
    - Cross-buffer vmcnt(6) pipeline overlapping loads with compute
    - 2 waves (128 threads), TILE_M_GRID=32
  Shape 5/6 (K>=1024): Custom HIP quant kernel -> AITER ASM GEMM
"""
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

from task import input_t, output_t

import torch

# ═══════════════════════════════════════════════════════════════
# HIP C++ Source
# ═══════════════════════════════════════════════════════════════

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

// ═══════════════════════════════════════════════════════════════
// Type definitions for MFMA and CVT
// ═══════════════════════════════════════════════════════════════
typedef float  __attribute__((ext_vector_type(4)))  floatx4;
typedef int    __attribute__((ext_vector_type(4)))  intx4;
typedef __bf16 __attribute__((ext_vector_type(2)))  bf16x2_t;

// ═══════════════════════════════════════════════════════════════
// Constants
// ═══════════════════════════════════════════════════════════════
#define TILE_M_PER_WAVE 16
#define TILE_M_GRID 32       // 2 waves × 16 rows each
#define TILE_N 16
#define K_STEP 128
#define SCALE_GROUP 32
#define NUM_K_GROUPS 4       // K_STEP / SCALE_GROUP = 128 / 32
#define WAVESIZE 64
#define BLOCK_SIZE 128       // 2 waves

// ═══════════════════════════════════════════════════════════════
// Device helpers
// ═══════════════════════════════════════════════════════════════

__device__ __forceinline__ float bf16_to_f32(uint16_t x) {
    uint32_t bits = ((uint32_t)x) << 16;
    float f;
    __builtin_memcpy(&f, &bits, 4);
    return f;
}

__device__ __forceinline__ uint16_t f32_to_bf16(float f) {
    uint32_t bits;
    __builtin_memcpy(&bits, &f, 4);
    uint32_t lsb = (bits >> 16) & 1;
    uint32_t rounding_bias = 0x7FFF + lsb;
    bits += rounding_bias;
    return (uint16_t)(bits >> 16);
}

__device__ __forceinline__ uint32_t float_as_uint(float f) {
    uint32_t u;
    __builtin_memcpy(&u, &f, 4);
    return u;
}

__device__ __forceinline__ float uint_as_float(uint32_t u) {
    float f;
    __builtin_memcpy(&f, &u, 4);
    return f;
}

// ═══════════════════════════════════════════════════════════════
// Standalone MXFP4 Quantization Kernel (from v91)
// ═══════════════════════════════════════════════════════════════

__global__ __attribute__((amdgpu_flat_work_group_size(256, 256)))
void mxfp4_quant_kernel(
    const uint16_t* __restrict__ A,
    uint8_t*        __restrict__ A_q,
    uint8_t*        __restrict__ A_scale_sh,
    int M, int K,
    int scaleN,
    int scaleN_valid
) {
    int tid = blockIdx.x * 256 + threadIdx.x;
    int m = tid / scaleN_valid;
    int s = tid % scaleN_valid;

    if (m >= M) return;

    int k_start = s * SCALE_GROUP;

    int a_words[16];
    if (k_start + SCALE_GROUP <= K) {
        const intx4* a_vec = reinterpret_cast<const intx4*>(A + m * K + k_start);
        intx4 v0 = a_vec[0], v1 = a_vec[1], v2 = a_vec[2], v3 = a_vec[3];
        a_words[0]  = v0[0]; a_words[1]  = v0[1]; a_words[2]  = v0[2]; a_words[3]  = v0[3];
        a_words[4]  = v1[0]; a_words[5]  = v1[1]; a_words[6]  = v1[2]; a_words[7]  = v1[3];
        a_words[8]  = v2[0]; a_words[9]  = v2[1]; a_words[10] = v2[2]; a_words[11] = v2[3];
        a_words[12] = v3[0]; a_words[13] = v3[1]; a_words[14] = v3[2]; a_words[15] = v3[3];
    } else {
        for (int i = 0; i < 16; i++) a_words[i] = 0;
        for (int i = 0; i < SCALE_GROUP && k_start + i < K; i++) {
            uint16_t val = A[m * K + k_start + i];
            int w = i / 2;
            if (i % 2 == 0)
                a_words[w] = (int)((uint32_t)val);
            else
                a_words[w] |= (int)(((uint32_t)val) << 16);
        }
    }

    float av[32];
    #pragma unroll
    for (int w = 0; w < 16; w++) {
        uint32_t abs_word = ((uint32_t)a_words[w]) & 0x7FFF7FFFu;
        av[2*w]   = bf16_to_f32((uint16_t)(abs_word & 0xFFFF));
        av[2*w+1] = bf16_to_f32((uint16_t)(abs_word >> 16));
    }

    float L1_0, L1_1, L1_2, L1_3, L1_4, L1_5, L1_6, L1_7, L1_8, L1_9;
    asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_0) : "v"(av[0]),  "v"(av[1]),  "v"(av[2]));
    asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_1) : "v"(av[3]),  "v"(av[4]),  "v"(av[5]));
    asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_2) : "v"(av[6]),  "v"(av[7]),  "v"(av[8]));
    asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_3) : "v"(av[9]),  "v"(av[10]), "v"(av[11]));
    asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_4) : "v"(av[12]), "v"(av[13]), "v"(av[14]));
    asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_5) : "v"(av[15]), "v"(av[16]), "v"(av[17]));
    asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_6) : "v"(av[18]), "v"(av[19]), "v"(av[20]));
    asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_7) : "v"(av[21]), "v"(av[22]), "v"(av[23]));
    asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_8) : "v"(av[24]), "v"(av[25]), "v"(av[26]));
    asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_9) : "v"(av[27]), "v"(av[28]), "v"(av[29]));

    float L2_0, L2_1, L2_2, L2_3;
    asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L2_0) : "v"(L1_0), "v"(L1_1), "v"(L1_2));
    asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L2_1) : "v"(L1_3), "v"(L1_4), "v"(L1_5));
    asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L2_2) : "v"(L1_6), "v"(L1_7), "v"(L1_8));
    asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L2_3) : "v"(L1_9), "v"(av[30]), "v"(av[31]));

    float L3_0;
    asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L3_0) : "v"(L2_0), "v"(L2_1), "v"(L2_2));

    float amax_val;
    asm volatile("v_max_f32 %0, %1, %2" : "=v"(amax_val) : "v"(L3_0), "v"(L2_3));

    uint32_t amax_bits = float_as_uint(amax_val);
    amax_bits = ((amax_bits + 0x200000u) & 0xFF800000u);

    int scale_unbiased;
    if (amax_bits == 0) {
        scale_unbiased = -127;
    } else {
        int exp_biased = (int)((amax_bits >> 23) & 0xFF);
        scale_unbiased = exp_biased - 127 - 2;
        if (scale_unbiased < -127) scale_unbiased = -127;
        if (scale_unbiased > 127) scale_unbiased = 127;
    }
    uint8_t a_e8m0 = (uint8_t)(scale_unbiased + 127);
    float cvt_scale = uint_as_float(((uint32_t)a_e8m0) << 23);

    unsigned int dst0 = 0, dst1 = 0, dst2 = 0, dst3 = 0;
    dst0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst0, *reinterpret_cast<const bf16x2_t*>(&a_words[0]), cvt_scale, 0);
    dst0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst0, *reinterpret_cast<const bf16x2_t*>(&a_words[1]), cvt_scale, 1);
    dst0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst0, *reinterpret_cast<const bf16x2_t*>(&a_words[2]), cvt_scale, 2);
    dst0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst0, *reinterpret_cast<const bf16x2_t*>(&a_words[3]), cvt_scale, 3);
    dst1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst1, *reinterpret_cast<const bf16x2_t*>(&a_words[4]), cvt_scale, 0);
    dst1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst1, *reinterpret_cast<const bf16x2_t*>(&a_words[5]), cvt_scale, 1);
    dst1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst1, *reinterpret_cast<const bf16x2_t*>(&a_words[6]), cvt_scale, 2);
    dst1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst1, *reinterpret_cast<const bf16x2_t*>(&a_words[7]), cvt_scale, 3);
    dst2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst2, *reinterpret_cast<const bf16x2_t*>(&a_words[8]), cvt_scale, 0);
    dst2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst2, *reinterpret_cast<const bf16x2_t*>(&a_words[9]), cvt_scale, 1);
    dst2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst2, *reinterpret_cast<const bf16x2_t*>(&a_words[10]), cvt_scale, 2);
    dst2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst2, *reinterpret_cast<const bf16x2_t*>(&a_words[11]), cvt_scale, 3);
    dst3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst3, *reinterpret_cast<const bf16x2_t*>(&a_words[12]), cvt_scale, 0);
    dst3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst3, *reinterpret_cast<const bf16x2_t*>(&a_words[13]), cvt_scale, 1);
    dst3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst3, *reinterpret_cast<const bf16x2_t*>(&a_words[14]), cvt_scale, 2);
    dst3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst3, *reinterpret_cast<const bf16x2_t*>(&a_words[15]), cvt_scale, 3);

    uint8_t* out_ptr = A_q + m * (K / 2) + k_start / 2;
    reinterpret_cast<uint32_t*>(out_ptr)[0] = dst0;
    reinterpret_cast<uint32_t*>(out_ptr)[1] = dst1;
    reinterpret_cast<uint32_t*>(out_ptr)[2] = dst2;
    reinterpret_cast<uint32_t*>(out_ptr)[3] = dst3;

    int m_mod32 = m & 31;
    int m_div32 = m >> 5;
    int s_mod8 = s & 7;
    int s_div8 = s >> 3;

    int sh_idx = (m_mod32 >> 4)
              + ((s_mod8 >> 2) << 1)
              + ((m_mod32 & 15) << 2)
              + ((s_mod8 & 3) << 6)
              + (s_div8 << 8)
              + m_div32 * (32 * scaleN);

    A_scale_sh[sh_idx] = a_e8m0;
}


void launch_mxfp4_quant(
    torch::Tensor A,
    torch::Tensor A_q,
    torch::Tensor A_scale_sh,
    int M, int K,
    int scaleN,
    int scaleN_valid
) {
    TORCH_CHECK(A.is_cuda() && A_q.is_cuda() && A_scale_sh.is_cuda());

    int total_groups = M * scaleN_valid;
    int threads = 256;
    int blocks = (total_groups + threads - 1) / threads;

    mxfp4_quant_kernel<<<blocks, threads>>>(
        reinterpret_cast<const uint16_t*>(A.data_ptr<at::BFloat16>()),
        A_q.data_ptr<uint8_t>(),
        A_scale_sh.data_ptr<uint8_t>(),
        M, K, scaleN, scaleN_valid
    );
}


// ═══════════════════════════════════════════════════════════════
// LOAD macros and helpers for fused kernel (same as v91)
// ═══════════════════════════════════════════════════════════════

#define LOAD_A_DATA(a_tmp0, a_tmp1, a_tmp2, a_tmp3, a_path_var, \
                    A_ptr, a_row_global, a_k_start, stride_a, M, K) \
    do { \
        if ((a_row_global) < M && ((a_k_start) + SCALE_GROUP) <= K) { \
            const intx4* a_vec_ptr = reinterpret_cast<const intx4*>( \
                (A_ptr) + (a_row_global) * (stride_a) + (a_k_start)); \
            (a_tmp0) = a_vec_ptr[0]; \
            (a_tmp1) = a_vec_ptr[1]; \
            (a_tmp2) = a_vec_ptr[2]; \
            (a_tmp3) = a_vec_ptr[3]; \
            (a_path_var) = 0; \
        } else if ((a_row_global) < M) { \
            (a_tmp0) = intx4{0,0,0,0}; \
            (a_tmp1) = intx4{0,0,0,0}; \
            (a_tmp2) = intx4{0,0,0,0}; \
            (a_tmp3) = intx4{0,0,0,0}; \
            for (int i = 0; i < SCALE_GROUP; i++) { \
                int k_idx = (a_k_start) + i; \
                uint16_t bf16_val = 0; \
                if (k_idx < K) { \
                    bf16_val = (A_ptr)[(a_row_global) * (stride_a) + k_idx]; \
                } \
                int w = i / 2; \
                int slot = i % 2; \
                int* words; \
                if (w < 4) words = reinterpret_cast<int*>(&(a_tmp0)); \
                else if (w < 8) { words = reinterpret_cast<int*>(&(a_tmp1)); w -= 4; } \
                else if (w < 12) { words = reinterpret_cast<int*>(&(a_tmp2)); w -= 8; } \
                else { words = reinterpret_cast<int*>(&(a_tmp3)); w -= 12; } \
                if (slot == 0) { \
                    words[w] = (int)((uint32_t)bf16_val); \
                } else { \
                    words[w] |= (int)(((uint32_t)bf16_val) << 16); \
                } \
            } \
            (a_path_var) = 1; \
        } else { \
            (a_tmp0) = intx4{0,0,0,0}; \
            (a_tmp1) = intx4{0,0,0,0}; \
            (a_tmp2) = intx4{0,0,0,0}; \
            (a_tmp3) = intx4{0,0,0,0}; \
            (a_path_var) = 2; \
        } \
    } while(0)

#define LOAD_B_DATA(b_mfma_var, b_tile_row, k_base, lane, n_base, N, K, row_in_tile, k_group) \
    do { \
        const uint8_t* tile_base = (b_tile_row) + ((k_base) / 2) * 16; \
        if ((n_base) + TILE_N <= N && ((k_base) + K_STEP) <= K) { \
            const intx4* b_vec_ptr = reinterpret_cast<const intx4*>( \
                tile_base + (lane) * 16); \
            (b_mfma_var) = *b_vec_ptr; \
        } else { \
            int b_n = (n_base) + (row_in_tile); \
            if (b_n < N && ((k_base) + (k_group) * SCALE_GROUP + SCALE_GROUP) <= K) { \
                const intx4* b_vec_ptr = reinterpret_cast<const intx4*>( \
                    tile_base + (lane) * 16); \
                (b_mfma_var) = *b_vec_ptr; \
            } else { \
                (b_mfma_var)[0] = 0; (b_mfma_var)[1] = 0; \
                (b_mfma_var)[2] = 0; (b_mfma_var)[3] = 0; \
            } \
        } \
    } while(0)

#define LOAD_B_SCALE_SHUFFLED(b_e8m0_var, B_scale_sh_ptr, b_row_global, b_scale_k_idx, \
                              N, num_scale_k) \
    do { \
        if ((b_row_global) < N && (b_scale_k_idx) < (num_scale_k)) { \
            int _n = (b_row_global); \
            int _k = (b_scale_k_idx); \
            int _n_group  = _n >> 5; \
            int _n_half   = (_n >> 4) & 1; \
            int _n_row16  = _n & 15; \
            int _k_group8 = _k >> 3; \
            int _k_half   = (_k >> 2) & 1; \
            int _k_mod4   = _k & 3; \
            int _sh_idx   = _n_group * ((num_scale_k) * 32) \
                          + _k_group8 * 256 \
                          + _k_mod4 * 64 \
                          + _n_row16 * 4 \
                          + _k_half * 2 \
                          + _n_half; \
            (b_e8m0_var) = (B_scale_sh_ptr)[_sh_idx]; \
        } else { \
            (b_e8m0_var) = 0; \
        } \
    } while(0)

#define PROCESS_TILE(cur_a0, cur_a1, cur_a2, cur_a3, cur_b, cur_bs, \
                     cur_path, acc_var) \
    do { \
        int a_words[16]; \
        a_words[0]  = (cur_a0)[0]; a_words[1]  = (cur_a0)[1]; \
        a_words[2]  = (cur_a0)[2]; a_words[3]  = (cur_a0)[3]; \
        a_words[4]  = (cur_a1)[0]; a_words[5]  = (cur_a1)[1]; \
        a_words[6]  = (cur_a1)[2]; a_words[7]  = (cur_a1)[3]; \
        a_words[8]  = (cur_a2)[0]; a_words[9]  = (cur_a2)[1]; \
        a_words[10] = (cur_a2)[2]; a_words[11] = (cur_a2)[3]; \
        a_words[12] = (cur_a3)[0]; a_words[13] = (cur_a3)[1]; \
        a_words[14] = (cur_a3)[2]; a_words[15] = (cur_a3)[3]; \
        \
        float amax_val = 0.0f; \
        if ((cur_path) == 0 || (cur_path) == 1) { \
            float av[32]; \
            _Pragma("unroll") \
            for (int w = 0; w < 16; w++) { \
                uint32_t abs_word = ((uint32_t)a_words[w]) & 0x7FFF7FFFu; \
                av[2*w]   = bf16_to_f32((uint16_t)(abs_word & 0xFFFF)); \
                av[2*w+1] = bf16_to_f32((uint16_t)(abs_word >> 16)); \
            } \
            \
            float L1_0, L1_1, L1_2, L1_3, L1_4, L1_5, L1_6, L1_7, L1_8, L1_9; \
            asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_0) : "v"(av[0]),  "v"(av[1]),  "v"(av[2])); \
            asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_1) : "v"(av[3]),  "v"(av[4]),  "v"(av[5])); \
            asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_2) : "v"(av[6]),  "v"(av[7]),  "v"(av[8])); \
            asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_3) : "v"(av[9]),  "v"(av[10]), "v"(av[11])); \
            asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_4) : "v"(av[12]), "v"(av[13]), "v"(av[14])); \
            asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_5) : "v"(av[15]), "v"(av[16]), "v"(av[17])); \
            asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_6) : "v"(av[18]), "v"(av[19]), "v"(av[20])); \
            asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_7) : "v"(av[21]), "v"(av[22]), "v"(av[23])); \
            asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_8) : "v"(av[24]), "v"(av[25]), "v"(av[26])); \
            asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_9) : "v"(av[27]), "v"(av[28]), "v"(av[29])); \
            \
            float L2_0, L2_1, L2_2, L2_3; \
            asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L2_0) : "v"(L1_0), "v"(L1_1), "v"(L1_2)); \
            asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L2_1) : "v"(L1_3), "v"(L1_4), "v"(L1_5)); \
            asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L2_2) : "v"(L1_6), "v"(L1_7), "v"(L1_8)); \
            asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L2_3) : "v"(L1_9), "v"(av[30]), "v"(av[31])); \
            \
            float L3_0; \
            asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L3_0) : "v"(L2_0), "v"(L2_1), "v"(L2_2)); \
            \
            asm volatile("v_max_f32 %0, %1, %2" : "=v"(amax_val) : "v"(L3_0), "v"(L2_3)); \
        } \
        \
        uint32_t amax_bits = float_as_uint(amax_val); \
        amax_bits = ((amax_bits + 0x200000u) & 0xFF800000u); \
        int scale_unbiased; \
        if (amax_bits == 0) { \
            scale_unbiased = -127; \
        } else { \
            int exp_biased = (int)((amax_bits >> 23) & 0xFF); \
            scale_unbiased = exp_biased - 127 - 2; \
            if (scale_unbiased < -127) scale_unbiased = -127; \
            if (scale_unbiased > 127) scale_unbiased = 127; \
        } \
        uint8_t a_e8m0 = (uint8_t)(scale_unbiased + 127); \
        float cvt_scale = uint_as_float(((uint32_t)a_e8m0) << 23); \
        \
        unsigned int dst0 = 0, dst1 = 0, dst2 = 0, dst3 = 0; \
        dst0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
            dst0, *reinterpret_cast<const bf16x2_t*>(&a_words[0]), cvt_scale, 0); \
        dst0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
            dst0, *reinterpret_cast<const bf16x2_t*>(&a_words[1]), cvt_scale, 1); \
        dst0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
            dst0, *reinterpret_cast<const bf16x2_t*>(&a_words[2]), cvt_scale, 2); \
        dst0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
            dst0, *reinterpret_cast<const bf16x2_t*>(&a_words[3]), cvt_scale, 3); \
        dst1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
            dst1, *reinterpret_cast<const bf16x2_t*>(&a_words[4]), cvt_scale, 0); \
        dst1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
            dst1, *reinterpret_cast<const bf16x2_t*>(&a_words[5]), cvt_scale, 1); \
        dst1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
            dst1, *reinterpret_cast<const bf16x2_t*>(&a_words[6]), cvt_scale, 2); \
        dst1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
            dst1, *reinterpret_cast<const bf16x2_t*>(&a_words[7]), cvt_scale, 3); \
        dst2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
            dst2, *reinterpret_cast<const bf16x2_t*>(&a_words[8]), cvt_scale, 0); \
        dst2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
            dst2, *reinterpret_cast<const bf16x2_t*>(&a_words[9]), cvt_scale, 1); \
        dst2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
            dst2, *reinterpret_cast<const bf16x2_t*>(&a_words[10]), cvt_scale, 2); \
        dst2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
            dst2, *reinterpret_cast<const bf16x2_t*>(&a_words[11]), cvt_scale, 3); \
        dst3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
            dst3, *reinterpret_cast<const bf16x2_t*>(&a_words[12]), cvt_scale, 0); \
        dst3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
            dst3, *reinterpret_cast<const bf16x2_t*>(&a_words[13]), cvt_scale, 1); \
        dst3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
            dst3, *reinterpret_cast<const bf16x2_t*>(&a_words[14]), cvt_scale, 2); \
        dst3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
            dst3, *reinterpret_cast<const bf16x2_t*>(&a_words[15]), cvt_scale, 3); \
        \
        intx4 a_mfma; \
        a_mfma[0] = (int)dst0; \
        a_mfma[1] = (int)dst1; \
        a_mfma[2] = (int)dst2; \
        a_mfma[3] = (int)dst3; \
        \
        int scale_a_packed = (int)((uint32_t)a_e8m0); \
        int scale_b_packed = (int)((uint32_t)(cur_bs)); \
        \
        asm volatile("s_setprio 1" ::: "memory"); \
        asm volatile( \
            "v_mfma_scale_f32_16x16x128_f8f6f4 %0, %1, %2, %0, %3, %4 cbsz:4 blgp:4" \
            : "+v"(acc_var) \
            : "v"(a_mfma), "v"(cur_b), "v"(scale_a_packed), "v"(scale_b_packed) \
        ); \
        asm volatile("s_setprio 0" ::: "memory"); \
    } while(0)


// ═══════════════════════════════════════════════════════════════
// NEW: Fused M=16 kernel with DOUBLE-BUFFERED K-loop
// For Shape 2 (M=16, N=2112, K=7168) only
//
// Architecture:
//   Grid: ceil(N/16) = 132 workgroups
//   Block: 64 threads (1 wavefront)
//   Tile: 16x16 output (one MFMA instruction per K=128 step)
//   K-loop: 56 iterations of K=128, DOUBLE-BUFFERED
//
// KEY CHANGE from v103 (serial):
//   v103: load_A → wait → LDS → read_LDS → quant → load_B → load_scale → MFMA → repeat
//   v108: Prefetch A[0] & B[0], then loop: {prefetch[n+1] → process[n] → wait → repeat}
//         This overlaps VMEM loads with MFMA compute.
//
// Pipeline per K-step:
//   1. Issue global loads for A[n+1] (non-blocking)
//   2. Issue global loads for B[n+1] and B_scale[n+1] (non-blocking)
//   3. Process current data: LDS write→wait→LDS read→quant→MFMA
//   4. s_waitcnt for next iteration's loads
//
// LDS: 2 x 4096 = 8192 bytes (double-buffered A data), trivially fits
// ═══════════════════════════════════════════════════════════════

// Helper: quantize 32 BF16 values (in 16 int words = 4 intx4) to FP4 (4 VGPRs)
// Returns a_mfma and a_e8m0 via out params
__device__ __forceinline__ void quantize_a_tile(
    const intx4& a_data0, const intx4& a_data1,
    const intx4& a_data2, const intx4& a_data3,
    intx4& a_mfma, uint8_t& a_e8m0
) {
    int a_words[16];
    a_words[0]  = a_data0[0]; a_words[1]  = a_data0[1];
    a_words[2]  = a_data0[2]; a_words[3]  = a_data0[3];
    a_words[4]  = a_data1[0]; a_words[5]  = a_data1[1];
    a_words[6]  = a_data1[2]; a_words[7]  = a_data1[3];
    a_words[8]  = a_data2[0]; a_words[9]  = a_data2[1];
    a_words[10] = a_data2[2]; a_words[11] = a_data2[3];
    a_words[12] = a_data3[0]; a_words[13] = a_data3[1];
    a_words[14] = a_data3[2]; a_words[15] = a_data3[3];

    // Compute amax via V_MAX3_F32 tree reduction
    float av[32];
    #pragma unroll
    for (int w = 0; w < 16; w++) {
        uint32_t abs_word = ((uint32_t)a_words[w]) & 0x7FFF7FFFu;
        av[2*w]   = bf16_to_f32((uint16_t)(abs_word & 0xFFFF));
        av[2*w+1] = bf16_to_f32((uint16_t)(abs_word >> 16));
    }

    float L1_0, L1_1, L1_2, L1_3, L1_4, L1_5, L1_6, L1_7, L1_8, L1_9;
    asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_0) : "v"(av[0]),  "v"(av[1]),  "v"(av[2]));
    asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_1) : "v"(av[3]),  "v"(av[4]),  "v"(av[5]));
    asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_2) : "v"(av[6]),  "v"(av[7]),  "v"(av[8]));
    asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_3) : "v"(av[9]),  "v"(av[10]), "v"(av[11]));
    asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_4) : "v"(av[12]), "v"(av[13]), "v"(av[14]));
    asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_5) : "v"(av[15]), "v"(av[16]), "v"(av[17]));
    asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_6) : "v"(av[18]), "v"(av[19]), "v"(av[20]));
    asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_7) : "v"(av[21]), "v"(av[22]), "v"(av[23]));
    asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_8) : "v"(av[24]), "v"(av[25]), "v"(av[26]));
    asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L1_9) : "v"(av[27]), "v"(av[28]), "v"(av[29]));

    float L2_0, L2_1, L2_2, L2_3;
    asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L2_0) : "v"(L1_0), "v"(L1_1), "v"(L1_2));
    asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L2_1) : "v"(L1_3), "v"(L1_4), "v"(L1_5));
    asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L2_2) : "v"(L1_6), "v"(L1_7), "v"(L1_8));
    asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L2_3) : "v"(L1_9), "v"(av[30]), "v"(av[31]));

    float L3_0;
    asm volatile("v_max3_f32 %0, %1, %2, %3" : "=v"(L3_0) : "v"(L2_0), "v"(L2_1), "v"(L2_2));

    float amax_val;
    asm volatile("v_max_f32 %0, %1, %2" : "=v"(amax_val) : "v"(L3_0), "v"(L2_3));

    uint32_t amax_bits = float_as_uint(amax_val);
    amax_bits = ((amax_bits + 0x200000u) & 0xFF800000u);

    int scale_unbiased;
    if (amax_bits == 0) {
        scale_unbiased = -127;
    } else {
        int exp_biased = (int)((amax_bits >> 23) & 0xFF);
        scale_unbiased = exp_biased - 127 - 2;
        if (scale_unbiased < -127) scale_unbiased = -127;
        if (scale_unbiased > 127) scale_unbiased = 127;
    }
    a_e8m0 = (uint8_t)(scale_unbiased + 127);
    float cvt_scale = uint_as_float(((uint32_t)a_e8m0) << 23);

    unsigned int dst0 = 0, dst1 = 0, dst2 = 0, dst3 = 0;
    dst0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst0, *reinterpret_cast<const bf16x2_t*>(&a_words[0]), cvt_scale, 0);
    dst0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst0, *reinterpret_cast<const bf16x2_t*>(&a_words[1]), cvt_scale, 1);
    dst0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst0, *reinterpret_cast<const bf16x2_t*>(&a_words[2]), cvt_scale, 2);
    dst0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst0, *reinterpret_cast<const bf16x2_t*>(&a_words[3]), cvt_scale, 3);
    dst1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst1, *reinterpret_cast<const bf16x2_t*>(&a_words[4]), cvt_scale, 0);
    dst1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst1, *reinterpret_cast<const bf16x2_t*>(&a_words[5]), cvt_scale, 1);
    dst1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst1, *reinterpret_cast<const bf16x2_t*>(&a_words[6]), cvt_scale, 2);
    dst1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst1, *reinterpret_cast<const bf16x2_t*>(&a_words[7]), cvt_scale, 3);
    dst2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst2, *reinterpret_cast<const bf16x2_t*>(&a_words[8]), cvt_scale, 0);
    dst2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst2, *reinterpret_cast<const bf16x2_t*>(&a_words[9]), cvt_scale, 1);
    dst2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst2, *reinterpret_cast<const bf16x2_t*>(&a_words[10]), cvt_scale, 2);
    dst2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst2, *reinterpret_cast<const bf16x2_t*>(&a_words[11]), cvt_scale, 3);
    dst3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst3, *reinterpret_cast<const bf16x2_t*>(&a_words[12]), cvt_scale, 0);
    dst3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst3, *reinterpret_cast<const bf16x2_t*>(&a_words[13]), cvt_scale, 1);
    dst3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst3, *reinterpret_cast<const bf16x2_t*>(&a_words[14]), cvt_scale, 2);
    dst3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(
        dst3, *reinterpret_cast<const bf16x2_t*>(&a_words[15]), cvt_scale, 3);

    a_mfma[0] = (int)dst0;
    a_mfma[1] = (int)dst1;
    a_mfma[2] = (int)dst2;
    a_mfma[3] = (int)dst3;
}


// ═══════════════════════════════════════════════════════════════
// Split-K Fused M=16 kernel: increases CU occupancy from 0.52 to 4.1 WGs/CU
//
// Grid: dim3(num_n_tiles, split_k) = (132, 8) = 1056 WGs for Shape 2
// Block: 64 threads (1 wavefront)
// Each WG processes K/split_k K-elements (e.g., 896 for K=7168, split_k=8)
// Output: F32 partial sums to workspace[split_idx * M * N + m * N + n]
// ═══════════════════════════════════════════════════════════════

__global__ __attribute__((amdgpu_flat_work_group_size(64, 64)))
void mxfp4_fused_gemm_m16_splitk(
    const uint16_t* __restrict__ A,           // [M, K] BF16 (M<=16)
    const uint8_t*  __restrict__ B_shuffle,   // [N//16, K/2*16] preshuffled
    const uint8_t*  __restrict__ B_scale_sh,  // shuffled e8m0 scales
    float*          __restrict__ workspace,   // [split_k, M, N] F32 partial sums
    int M, int N, int K,
    int k_per_split                           // K-elements per split (aligned to K_STEP)
) {
    // Double-buffered LDS for A data redistribution
    __shared__ uint8_t lds_a[8192];  // 2 x (16 rows * 128 BF16 * 2 bytes)

    const int n_tile = blockIdx.x;
    const int split_idx = blockIdx.y;
    const int n_base = n_tile * 16;
    if (n_base >= N) return;

    const int lane = threadIdx.x;  // 0..63
    const int row_in_tile = lane & 15;  // N-index within tile (0..15)
    const int k_group = lane >> 4;      // K-group index (0..3)

    const int K_half = K >> 1;
    const int num_scale_k = (K + 31) / 32;

    // B_shuffle tile base for this N-tile
    const uint8_t* b_tile_base = B_shuffle + (long long)n_tile * K_half * 16;

    // Accumulator
    floatx4 acc = {0.0f, 0.0f, 0.0f, 0.0f};

    const int stride_a = K;

    // ═══════════════════════════════════════════════════
    // Compute K-range for this split
    // ═══════════════════════════════════════════════════
    int k_start = split_idx * k_per_split;
    int k_end = k_start + k_per_split;
    if (k_end > K) k_end = K;
    if (k_start >= K) return;

    // Number of K_STEP=128 iterations for this split
    int first_k_step = k_start / K_STEP;
    int last_k_step = (k_end + K_STEP - 1) / K_STEP;
    int num_k_steps = last_k_step - first_k_step;
    if (num_k_steps <= 0) return;

    // ═══════════════════════════════════════════════════
    // Precompute lane's global A address components
    // Lane L loads A at: row = lane/4, col_base = (lane%4)*32
    // ═══════════════════════════════════════════════════
    const int a_row = lane >> 2;         // 0..15
    const int a_col_base = (lane & 3) * 32;  // 0, 32, 64, 96

    // ═══════════════════════════════════════════════════
    // Precompute lane's LDS read offset
    // ═══════════════════════════════════════════════════
    const int lds_read_base = row_in_tile * 256 + k_group * 64;

    // ═══════════════════════════════════════════════════
    // Precompute B_scale shuffle index components
    // ═══════════════════════════════════════════════════
    const int b_row_global = n_base + row_in_tile;
    const int _n_group  = b_row_global >> 5;
    const int _n_half   = (b_row_global >> 4) & 1;
    const int _n_row16  = b_row_global & 15;
    const int _scale_n_base = _n_group * (num_scale_k * 32) + _n_row16 * 4 + _n_half;

    // ═══════════════════════════════════════════════════
    // PROLOGUE: Issue loads for first K-step (non-blocking)
    // ═══════════════════════════════════════════════════
    int lds_buf = 0;
    int first_k_base = first_k_step * K_STEP;

    // Load A[first_k] into VGPRs
    intx4 a_load0, a_load1, a_load2, a_load3;
    {
        int a_k_global = first_k_base + a_col_base;
        if (a_row < M && a_k_global + 31 < K) {
            const intx4* a_vec = reinterpret_cast<const intx4*>(
                A + a_row * stride_a + a_k_global);
            a_load0 = a_vec[0];
            a_load1 = a_vec[1];
            a_load2 = a_vec[2];
            a_load3 = a_vec[3];
        } else {
            a_load0 = intx4{0,0,0,0};
            a_load1 = intx4{0,0,0,0};
            a_load2 = intx4{0,0,0,0};
            a_load3 = intx4{0,0,0,0};
        }
    }

    // Load B[first_k] directly to register
    intx4 b_mfma;
    {
        const uint8_t* b_k_ptr = b_tile_base + (long long)(first_k_base / 2) * 16;
        const intx4* b_vec = reinterpret_cast<const intx4*>(b_k_ptr + lane * 16);
        b_mfma = *b_vec;
    }

    // Load B_scale[first_k]
    uint8_t b_e8m0;
    {
        int b_scale_k_idx = first_k_base / SCALE_GROUP + k_group;
        int _k = b_scale_k_idx;
        int _k_group8 = _k >> 3;
        int _k_half   = (_k >> 2) & 1;
        int _k_mod4   = _k & 3;
        int _sh_idx   = _scale_n_base
                      + _k_group8 * 256
                      + _k_mod4 * 64
                      + _k_half * 2;
        b_e8m0 = B_scale_sh[_sh_idx];
    }

    // ═══════════════════════════════════════════════════
    // MAIN K-LOOP: Double-buffered over this split's K-range
    // ═══════════════════════════════════════════════════

    for (int step = 0; step < num_k_steps; step++) {
        int abs_k_step = first_k_step + step;
        int k_base = abs_k_step * K_STEP;
        int lds_offset = lds_buf * 4096;

        // ── Step 1: Write A data to LDS ──
        {
            intx4* lds_dst = reinterpret_cast<intx4*>(lds_a + lds_offset + lane * 64);
            lds_dst[0] = a_load0;
            lds_dst[1] = a_load1;
            lds_dst[2] = a_load2;
            lds_dst[3] = a_load3;
        }

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

        // ── Step 2: Read A from LDS with MFMA lane pattern ──
        intx4 a_data0, a_data1, a_data2, a_data3;
        {
            const intx4* lds_src = reinterpret_cast<const intx4*>(
                lds_a + lds_offset + lds_read_base);
            a_data0 = lds_src[0];
            a_data1 = lds_src[1];
            a_data2 = lds_src[2];
            a_data3 = lds_src[3];
        }

        // ── Step 3: Issue prefetch for NEXT K-step ──
        intx4 next_a_load0, next_a_load1, next_a_load2, next_a_load3;
        intx4 next_b_mfma;
        uint8_t next_b_e8m0;

        if (step + 1 < num_k_steps) {
            int next_abs_k_step = first_k_step + step + 1;
            int next_k_base = next_abs_k_step * K_STEP;

            // Prefetch A[k+1]
            {
                int a_k_global = next_k_base + a_col_base;
                if (a_row < M && a_k_global + 31 < K) {
                    const intx4* a_vec = reinterpret_cast<const intx4*>(
                        A + a_row * stride_a + a_k_global);
                    next_a_load0 = a_vec[0];
                    next_a_load1 = a_vec[1];
                    next_a_load2 = a_vec[2];
                    next_a_load3 = a_vec[3];
                } else {
                    next_a_load0 = intx4{0,0,0,0};
                    next_a_load1 = intx4{0,0,0,0};
                    next_a_load2 = intx4{0,0,0,0};
                    next_a_load3 = intx4{0,0,0,0};
                }
            }

            // Prefetch B[k+1]
            {
                const uint8_t* b_k_base_ptr = b_tile_base + (long long)(next_k_base / 2) * 16;
                const intx4* b_vec = reinterpret_cast<const intx4*>(b_k_base_ptr + lane * 16);
                next_b_mfma = *b_vec;
            }

            // Prefetch B_scale[k+1]
            {
                int b_scale_k_idx = next_k_base / SCALE_GROUP + k_group;
                int _k = b_scale_k_idx;
                int _k_group8 = _k >> 3;
                int _k_half   = (_k >> 2) & 1;
                int _k_mod4   = _k & 3;
                int _sh_idx   = _scale_n_base
                              + _k_group8 * 256
                              + _k_mod4 * 64
                              + _k_half * 2;
                next_b_e8m0 = B_scale_sh[_sh_idx];
            }
        }

        // ── Step 4: Quantize current A data (BF16 → FP4) ──
        intx4 a_mfma;
        uint8_t a_e8m0;
        quantize_a_tile(a_data0, a_data1, a_data2, a_data3, a_mfma, a_e8m0);

        // ── Step 5: MFMA execution ──
        int scale_a_packed = (int)((uint32_t)a_e8m0);
        int scale_b_packed = (int)((uint32_t)b_e8m0);

        asm volatile("s_setprio 1" ::: "memory");
        asm volatile(
            "v_mfma_scale_f32_16x16x128_f8f6f4 %0, %1, %2, %0, %3, %4 cbsz:4 blgp:4"
            : "+v"(acc)
            : "v"(a_mfma), "v"(b_mfma), "v"(scale_a_packed), "v"(scale_b_packed)
        );
        asm volatile("s_setprio 0" ::: "memory");

        // ── Step 6: Wait for prefetched data and swap buffers ──
        if (step + 1 < num_k_steps) {
            asm volatile("s_waitcnt vmcnt(0)" ::: "memory");

            a_load0 = next_a_load0;
            a_load1 = next_a_load1;
            a_load2 = next_a_load2;
            a_load3 = next_a_load3;
            b_mfma = next_b_mfma;
            b_e8m0 = next_b_e8m0;
        }

        lds_buf ^= 1;
    }

    // ════════════════════════════════════════════════
    // Epilogue: Write F32 partial sums to workspace
    // workspace layout: [split_k, M, N] in row-major
    // Lane L has acc[0..3] = C[row_base+0..3, col]
    // where col = L%16, row_base = (L/16)*4
    // ════════════════════════════════════════════════
    {
        int col = lane & 15;
        int row_base = (lane >> 4) * 4;
        int n_global = n_base + col;

        if (n_global < N) {
            float* ws_slice = workspace + (long long)split_idx * M * N;
            for (int v = 0; v < 4; v++) {
                int m_global = row_base + v;
                if (m_global < M) {
                    ws_slice[m_global * N + n_global] = acc[v];
                }
            }
        }
    }
}


// ═══════════════════════════════════════════════════════════════
// Reduction kernel: sum F32 partial sums across split_k, convert to BF16
// Grid: ceil(M * N / 256), Block: 256
// ═══════════════════════════════════════════════════════════════

__global__ __attribute__((amdgpu_flat_work_group_size(256, 256)))
void reduce_splitk_to_bf16(
    const float*    __restrict__ workspace,  // [split_k, M, N]
    uint16_t*       __restrict__ C,          // [M, N] BF16
    int M, int N, int split_k
) {
    int idx = blockIdx.x * 256 + threadIdx.x;
    int total = M * N;
    if (idx >= total) return;

    float sum = 0.0f;
    for (int s = 0; s < split_k; s++) {
        sum += workspace[(long long)s * M * N + idx];
    }

    C[idx] = f32_to_bf16(sum);
}


// ═══════════════════════════════════════════════════════════════
// NEW: Multi-wave fused M=16 kernel with ds_bpermute A-redistribution
//
// Grid: dim3(num_n_tiles) = 132 WGs for Shape 2
// Block: 512 threads (8 wavefronts per WG)
//
// Each wave independently processes K/8 K-range.
// KEY CHANGE (v124): Uses ds_bpermute_b32 for A-data lane remapping
//   instead of LDS write + barrier + LDS read.
//   - Eliminates per-wave LDS A double-buffers (saves 64KB LDS)
//   - No lgkmcnt barrier needed for A redistribution
//   - 16 ds_bpermute_b32 per K-step (one per int32 of 4 intx4)
//
// LDS layout (v124):
//   Reduction only: 8 × 4 × 64 × 4 = 8192 bytes (8KB)
//   Total: 8192 bytes — down from 72KB!
// ═══════════════════════════════════════════════════════════════

#define NUM_WAVES 8
#define MULTIWAVE_BLOCK_SIZE 512

__global__ __attribute__((amdgpu_flat_work_group_size(512, 512)))
void mxfp4_fused_gemm_m16_multiwave(
    const uint16_t* __restrict__ A,           // [M, K] BF16 (M<=16)
    const uint8_t*  __restrict__ B_shuffle,   // [N//16, K/2*16] preshuffled
    const uint8_t*  __restrict__ B_scale_sh,  // shuffled e8m0 scales
    uint16_t*       __restrict__ C,           // [M, N] BF16 output (direct!)
    int M, int N, int K,
    int k_per_wave                            // K-elements per wave (aligned to K_STEP)
) {
    // LDS: reduction buffer only (no A double-buffers needed with ds_bpermute!)
    __shared__ uint8_t lds_raw[8192];

    // Reduction buffer at offset 0
    float* lds_reduce = reinterpret_cast<float*>(lds_raw);
    // Layout: lds_reduce[wave_id * 256 + v * 64 + lane]

    const int n_tile = blockIdx.x;
    const int n_base = n_tile * 16;
    if (n_base >= N) return;

    const int wave_id = threadIdx.x / 64;  // 0..7
    const int lane = threadIdx.x % 64;     // 0..63
    const int row_in_tile = lane & 15;     // N-index within tile (0..15)
    const int k_group = lane >> 4;         // K-group index (0..3)

    const int K_half = K >> 1;
    const int num_scale_k = (K + 31) / 32;
    const int stride_a = K;

    // B_shuffle tile base for this N-tile
    const uint8_t* b_tile_base = B_shuffle + (long long)n_tile * K_half * 16;

    // Accumulator (per-wave, independent)
    floatx4 acc = {0.0f, 0.0f, 0.0f, 0.0f};

    // ═══════════════════════════════════════════════════
    // Precompute ds_bpermute source lane offset
    //
    // A load pattern: lane L has row=L/4, kgrp=L%4
    // MFMA needs:     lane L wants row=L%16, kgrp=L/16
    // Source lane:    src = (L%16)*4 + (L/16)
    // ds_bpermute takes byte offset = src_lane * 4
    // ═══════════════════════════════════════════════════
    const int bpermute_src_lane = (row_in_tile << 2) | k_group;  // (L%16)*4 + (L/16)
    const int bpermute_byte_offset = bpermute_src_lane << 2;      // * 4 for byte offset

    // ═══════════════════════════════════════════════════
    // Compute K-range for this wave
    // ═══════════════════════════════════════════════════
    int k_start = wave_id * k_per_wave;
    int k_end = k_start + k_per_wave;
    if (k_end > K) k_end = K;
    if (k_start >= K) goto do_reduction;  // This wave has no work

    {
        int first_k_step = k_start / K_STEP;
        int last_k_step = (k_end + K_STEP - 1) / K_STEP;
        int num_k_steps = last_k_step - first_k_step;
        if (num_k_steps <= 0) goto do_reduction;

        // ═══════════════════════════════════════════════════
        // Precompute lane's global A address components
        // Lane L loads A at: row = lane/4, col_base = (lane%4)*32
        // ═══════════════════════════════════════════════════
        const int a_row = lane >> 2;              // 0..15
        const int a_col_base = (lane & 3) * 32;   // 0, 32, 64, 96

        // Precompute B_scale shuffle index components
        const int b_row_global = n_base + row_in_tile;
        const int _n_group  = b_row_global >> 5;
        const int _n_half   = (b_row_global >> 4) & 1;
        const int _n_row16  = b_row_global & 15;
        const int _scale_n_base = _n_group * (num_scale_k * 32) + _n_row16 * 4 + _n_half;

        // ═══════════════════════════════════════════════════
        // PROLOGUE: Load first K-step data
        // ═══════════════════════════════════════════════════
        int first_k_base = first_k_step * K_STEP;

        // Load A[first_k]
        intx4 a_load0, a_load1, a_load2, a_load3;
        {
            int a_k_global = first_k_base + a_col_base;
            if (a_row < M && a_k_global + 31 < K) {
                const intx4* a_vec = reinterpret_cast<const intx4*>(
                    A + a_row * stride_a + a_k_global);
                a_load0 = a_vec[0];
                a_load1 = a_vec[1];
                a_load2 = a_vec[2];
                a_load3 = a_vec[3];
            } else {
                a_load0 = intx4{0,0,0,0};
                a_load1 = intx4{0,0,0,0};
                a_load2 = intx4{0,0,0,0};
                a_load3 = intx4{0,0,0,0};
            }
        }

        // Load B[first_k]
        intx4 b_mfma;
        {
            const uint8_t* b_k_ptr = b_tile_base + (long long)(first_k_base / 2) * 16;
            const intx4* b_vec = reinterpret_cast<const intx4*>(b_k_ptr + lane * 16);
            b_mfma = *b_vec;
        }

        // Load B_scale[first_k]
        uint8_t b_e8m0;
        {
            int b_scale_k_idx = first_k_base / SCALE_GROUP + k_group;
            int _k = b_scale_k_idx;
            int _k_group8 = _k >> 3;
            int _k_half   = (_k >> 2) & 1;
            int _k_mod4   = _k & 3;
            int _sh_idx   = _scale_n_base
                          + _k_group8 * 256
                          + _k_mod4 * 64
                          + _k_half * 2;
            b_e8m0 = B_scale_sh[_sh_idx];
        }

        // ═══════════════════════════════════════════════════
        // MAIN K-LOOP: Each wave processes its own K-range
        // Uses ds_bpermute for A-data redistribution (no LDS needed!)
        // ═══════════════════════════════════════════════════
        for (int step = 0; step < num_k_steps; step++) {

            // ── Redistribute A data via ds_bpermute ──
            // Each lane has a_load[0..3] in coalesced pattern (row=L/4, kgrp=L%4)
            // MFMA needs (row=L%16, kgrp=L/16)
            // ds_bpermute_b32: lane L reads from source lane = (L%16)*4 + (L/16)
            intx4 a_data0, a_data1, a_data2, a_data3;
            {
                int d0, d1, d2, d3;
                // Permute a_load0 (4 int32s)
                d0 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load0[0]);
                d1 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load0[1]);
                d2 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load0[2]);
                d3 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load0[3]);
                a_data0 = intx4{d0, d1, d2, d3};

                // Permute a_load1
                d0 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load1[0]);
                d1 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load1[1]);
                d2 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load1[2]);
                d3 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load1[3]);
                a_data1 = intx4{d0, d1, d2, d3};

                // Permute a_load2
                d0 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load2[0]);
                d1 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load2[1]);
                d2 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load2[2]);
                d3 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load2[3]);
                a_data2 = intx4{d0, d1, d2, d3};

                // Permute a_load3
                d0 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load3[0]);
                d1 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load3[1]);
                d2 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load3[2]);
                d3 = __builtin_amdgcn_ds_bpermute(bpermute_byte_offset, a_load3[3]);
                a_data3 = intx4{d0, d1, d2, d3};
            }

            // ── Issue prefetch for NEXT K-step ──
            intx4 next_a_load0, next_a_load1, next_a_load2, next_a_load3;
            intx4 next_b_mfma;
            uint8_t next_b_e8m0;

            if (step + 1 < num_k_steps) {
                int next_abs_k_step = first_k_step + step + 1;
                int next_k_base = next_abs_k_step * K_STEP;

                // Prefetch A[k+1]
                {
                    int a_k_global = next_k_base + a_col_base;
                    if (a_row < M && a_k_global + 31 < K) {
                        const intx4* a_vec = reinterpret_cast<const intx4*>(
                            A + a_row * stride_a + a_k_global);
                        next_a_load0 = a_vec[0];
                        next_a_load1 = a_vec[1];
                        next_a_load2 = a_vec[2];
                        next_a_load3 = a_vec[3];
                    } else {
                        next_a_load0 = intx4{0,0,0,0};
                        next_a_load1 = intx4{0,0,0,0};
                        next_a_load2 = intx4{0,0,0,0};
                        next_a_load3 = intx4{0,0,0,0};
                    }
                }

                // Prefetch B[k+1]
                {
                    const uint8_t* b_k_base_ptr = b_tile_base + (long long)(next_k_base / 2) * 16;
                    const intx4* b_vec = reinterpret_cast<const intx4*>(b_k_base_ptr + lane * 16);
                    next_b_mfma = *b_vec;
                }

                // Prefetch B_scale[k+1]
                {
                    int b_scale_k_idx = next_k_base / SCALE_GROUP + k_group;
                    int _k = b_scale_k_idx;
                    int _k_group8 = _k >> 3;
                    int _k_half   = (_k >> 2) & 1;
                    int _k_mod4   = _k & 3;
                    int _sh_idx   = _scale_n_base
                                  + _k_group8 * 256
                                  + _k_mod4 * 64
                                  + _k_half * 2;
                    next_b_e8m0 = B_scale_sh[_sh_idx];
                }
            }

            // ── Quantize current A data (BF16 → FP4) ──
            intx4 a_mfma;
            uint8_t a_e8m0;
            quantize_a_tile(a_data0, a_data1, a_data2, a_data3, a_mfma, a_e8m0);

            // ── MFMA execution ──
            int scale_a_packed = (int)((uint32_t)a_e8m0);
            int scale_b_packed = (int)((uint32_t)b_e8m0);

            asm volatile("s_setprio 1" ::: "memory");
            asm volatile(
                "v_mfma_scale_f32_16x16x128_f8f6f4 %0, %1, %2, %0, %3, %4 cbsz:4 blgp:4"
                : "+v"(acc)
                : "v"(a_mfma), "v"(b_mfma), "v"(scale_a_packed), "v"(scale_b_packed)
            );
            asm volatile("s_setprio 0" ::: "memory");

            // ── Swap to next prefetched data ──
            if (step + 1 < num_k_steps) {
                asm volatile("s_waitcnt vmcnt(0)" ::: "memory");

                a_load0 = next_a_load0;
                a_load1 = next_a_load1;
                a_load2 = next_a_load2;
                a_load3 = next_a_load3;
                b_mfma = next_b_mfma;
                b_e8m0 = next_b_e8m0;
            }
        }
    }  // end K-loop scope

do_reduction:
    // ════════════════════════════════════════════════
    // LDS Reduction: Sum accumulators across 8 waves
    //
    // Each wave's acc[0..3] holds partial sums for its K-range.
    // We need to sum across all 8 waves for the same output position.
    //
    // Step 1: Each wave writes its acc[4] to LDS reduction buffer
    // Step 2: __syncthreads() to ensure all waves have written
    // Step 3: Wave 0 reads and sums all 8 waves' values
    // Step 4: Wave 0 writes BF16 to output
    // ════════════════════════════════════════════════
    {
        // Write acc to LDS reduction buffer
        // Layout: lds_reduce[wave_id * 256 + v * 64 + lane]
        int reduce_base = wave_id * 256 + lane;
        lds_reduce[reduce_base + 0 * 64] = acc[0];
        lds_reduce[reduce_base + 1 * 64] = acc[1];
        lds_reduce[reduce_base + 2 * 64] = acc[2];
        lds_reduce[reduce_base + 3 * 64] = acc[3];

        __syncthreads();

        // Only wave 0 does the final reduction and writes output
        if (wave_id == 0) {
            float sum0 = 0.0f, sum1 = 0.0f, sum2 = 0.0f, sum3 = 0.0f;

            #pragma unroll
            for (int w = 0; w < NUM_WAVES; w++) {
                int w_base = w * 256 + lane;
                sum0 += lds_reduce[w_base + 0 * 64];
                sum1 += lds_reduce[w_base + 1 * 64];
                sum2 += lds_reduce[w_base + 2 * 64];
                sum3 += lds_reduce[w_base + 3 * 64];
            }

            // Write BF16 directly to output
            int col = lane & 15;
            int row_base = (lane >> 4) * 4;
            int n_global = n_base + col;

            if (n_global < N) {
                for (int v = 0; v < 4; v++) {
                    int m_global = row_base + v;
                    if (m_global < M) {
                        float val = (v == 0) ? sum0 : (v == 1) ? sum1 : (v == 2) ? sum2 : sum3;
                        C[m_global * N + n_global] = f32_to_bf16(val);
                    }
                }
            }
        }
    }
}


// ═══════════════════════════════════════════════════════════════
// PyTorch wrapper for multi-wave fused M=16 kernel (single kernel!)
// ═══════════════════════════════════════════════════════════════
void launch_mxfp4_fused_gemm_m16_multiwave(
    torch::Tensor A,
    torch::Tensor B_shuffle,
    torch::Tensor B_scale_sh,
    torch::Tensor C,           // [M, N] bfloat16 (direct output!)
    int M, int N, int K
) {
    TORCH_CHECK(A.is_cuda() && B_shuffle.is_cuda() && B_scale_sh.is_cuda() && C.is_cuda());

    int num_n_tiles = (N + 15) / 16;

    // Compute k_per_wave aligned to K_STEP=128
    int k_steps_total = (K + K_STEP - 1) / K_STEP;
    int k_steps_per_wave = (k_steps_total + NUM_WAVES - 1) / NUM_WAVES;
    int k_per_wave = k_steps_per_wave * K_STEP;

    dim3 grid(num_n_tiles);
    dim3 block(MULTIWAVE_BLOCK_SIZE);  // 512 threads = 8 wavefronts

    mxfp4_fused_gemm_m16_multiwave<<<grid, block>>>(
        reinterpret_cast<const uint16_t*>(A.data_ptr<at::BFloat16>()),
        B_shuffle.data_ptr<uint8_t>(),
        B_scale_sh.data_ptr<uint8_t>(),
        reinterpret_cast<uint16_t*>(C.data_ptr<at::BFloat16>()),
        M, N, K,
        k_per_wave
    );
}


// ═══════════════════════════════════════════════════════════════
// PyTorch wrapper for split-K fused M=16 kernel (legacy, kept for fallback)
// Launches main kernel then reduction kernel
// ═══════════════════════════════════════════════════════════════
void launch_mxfp4_fused_gemm_m16_splitk(
    torch::Tensor A,
    torch::Tensor B_shuffle,
    torch::Tensor B_scale_sh,
    torch::Tensor workspace,   // [split_k, M, N] float32
    torch::Tensor C,           // [M, N] bfloat16
    int M, int N, int K,
    int split_k
) {
    TORCH_CHECK(A.is_cuda() && B_shuffle.is_cuda() && B_scale_sh.is_cuda());
    TORCH_CHECK(workspace.is_cuda() && C.is_cuda());
    TORCH_CHECK(workspace.dtype() == torch::kFloat32, "workspace must be float32");

    int num_n_tiles = (N + 15) / 16;

    // Compute k_per_split aligned to K_STEP=128
    int k_steps_total = (K + K_STEP - 1) / K_STEP;
    int k_steps_per_split = (k_steps_total + split_k - 1) / split_k;
    int k_per_split = k_steps_per_split * K_STEP;

    // Main kernel: grid(num_n_tiles, split_k)
    dim3 grid_main(num_n_tiles, split_k);
    dim3 block_main(64);  // 1 wavefront

    mxfp4_fused_gemm_m16_splitk<<<grid_main, block_main>>>(
        reinterpret_cast<const uint16_t*>(A.data_ptr<at::BFloat16>()),
        B_shuffle.data_ptr<uint8_t>(),
        B_scale_sh.data_ptr<uint8_t>(),
        workspace.data_ptr<float>(),
        M, N, K,
        k_per_split
    );

    // Reduction kernel: sum across split_k, convert to BF16
    int total_elements = M * N;
    int red_blocks = (total_elements + 255) / 256;

    reduce_splitk_to_bf16<<<red_blocks, 256>>>(
        workspace.data_ptr<float>(),
        reinterpret_cast<uint16_t*>(C.data_ptr<at::BFloat16>()),
        M, N, split_k
    );
}


// ═══════════════════════════════════════════════════════════════
// K=512 specialized 2-wave kernel with fully unrolled K-loop
// For S3/S4 (M=32, K=512): 4 K_STEP=128 iterations, cross-buffer pipeline
// ═══════════════════════════════════════════════════════════════

__global__ __attribute__((amdgpu_flat_work_group_size(128, 128)))
void mxfp4_fused_gemm_k512_2wave(
    const uint16_t* __restrict__ A,
    const uint8_t*  __restrict__ B_preshuffle,
    uint16_t*       __restrict__ C,
    const uint8_t*  __restrict__ B_scale_sh,
    int M, int N, int K,
    int stride_bps
) {
    const int wave_id = threadIdx.x / 64;
    const int lane = threadIdx.x % 64;

    int num_tile_n = (N + TILE_N - 1) / TILE_N;
    int tile_m = blockIdx.x / num_tile_n;
    int tile_n = blockIdx.x % num_tile_n;

    int m_base = tile_m * TILE_M_GRID;
    int m_wave = m_base + wave_id * TILE_M_PER_WAVE;
    int n_base = tile_n * TILE_N;

    int stride_a = K;
    int num_scale_k = (K + 31) / 32;

    floatx4 acc = {0.0f, 0.0f, 0.0f, 0.0f};

    int row_in_tile = lane & 15;
    int k_group     = lane >> 4;

    const uint8_t* b_tile_row = B_preshuffle + (n_base / 16) * stride_bps;

    // Precompute invariant A row
    int a_row_global = m_wave + row_in_tile;

    // ═══════════════════════════════════════════════════
    // K=512: exactly 4 K_STEP=128 iterations, fully unrolled
    // ═══════════════════════════════════════════════════

    // ── Load step 0 ──
    intx4 a_tmp0_0, a_tmp1_0, a_tmp2_0, a_tmp3_0;
    intx4 b_mfma_0;
    uint8_t b_e8m0_0;
    int a_path_0;
    {
        int a_k_start = k_group * SCALE_GROUP;
        LOAD_A_DATA(a_tmp0_0, a_tmp1_0, a_tmp2_0, a_tmp3_0,
                    a_path_0, A, a_row_global, a_k_start, stride_a, M, K);
        LOAD_B_DATA(b_mfma_0, b_tile_row, 0, lane, n_base, N, K,
                    row_in_tile, k_group);
        int b_row_global = n_base + row_in_tile;
        int b_scale_k_idx = k_group;
        LOAD_B_SCALE_SHUFFLED(b_e8m0_0, B_scale_sh, b_row_global, b_scale_k_idx,
                     N, num_scale_k);
    }

    // ── Load step 1, process step 0 ──
    intx4 a_tmp0_1, a_tmp1_1, a_tmp2_1, a_tmp3_1;
    intx4 b_mfma_1;
    uint8_t b_e8m0_1;
    int a_path_1;
    {
        int k_base_1 = K_STEP;
        int a_k_start = k_base_1 + k_group * SCALE_GROUP;
        LOAD_A_DATA(a_tmp0_1, a_tmp1_1, a_tmp2_1, a_tmp3_1,
                    a_path_1, A, a_row_global, a_k_start, stride_a, M, K);
        LOAD_B_DATA(b_mfma_1, b_tile_row, k_base_1, lane, n_base, N, K,
                    row_in_tile, k_group);
        int b_row_global = n_base + row_in_tile;
        int b_scale_k_idx = k_base_1 / SCALE_GROUP + k_group;
        LOAD_B_SCALE_SHUFFLED(b_e8m0_1, B_scale_sh, b_row_global, b_scale_k_idx,
                     N, num_scale_k);

        asm volatile("s_waitcnt vmcnt(6)" ::: "memory");
        PROCESS_TILE(a_tmp0_0, a_tmp1_0, a_tmp2_0, a_tmp3_0,
                     b_mfma_0, b_e8m0_0, a_path_0, acc);
    }

    // ── Load step 2, process step 1 ──
    intx4 a_tmp0_2, a_tmp1_2, a_tmp2_2, a_tmp3_2;
    intx4 b_mfma_2;
    uint8_t b_e8m0_2;
    int a_path_2;
    {
        int k_base_2 = 2 * K_STEP;
        int a_k_start = k_base_2 + k_group * SCALE_GROUP;
        LOAD_A_DATA(a_tmp0_2, a_tmp1_2, a_tmp2_2, a_tmp3_2,
                    a_path_2, A, a_row_global, a_k_start, stride_a, M, K);
        LOAD_B_DATA(b_mfma_2, b_tile_row, k_base_2, lane, n_base, N, K,
                    row_in_tile, k_group);
        int b_row_global = n_base + row_in_tile;
        int b_scale_k_idx = k_base_2 / SCALE_GROUP + k_group;
        LOAD_B_SCALE_SHUFFLED(b_e8m0_2, B_scale_sh, b_row_global, b_scale_k_idx,
                     N, num_scale_k);

        asm volatile("s_waitcnt vmcnt(6)" ::: "memory");
        PROCESS_TILE(a_tmp0_1, a_tmp1_1, a_tmp2_1, a_tmp3_1,
                     b_mfma_1, b_e8m0_1, a_path_1, acc);
    }

    // ── Load step 3, process step 2 ──
    intx4 a_tmp0_3, a_tmp1_3, a_tmp2_3, a_tmp3_3;
    intx4 b_mfma_3;
    uint8_t b_e8m0_3;
    int a_path_3;
    {
        int k_base_3 = 3 * K_STEP;
        int a_k_start = k_base_3 + k_group * SCALE_GROUP;
        LOAD_A_DATA(a_tmp0_3, a_tmp1_3, a_tmp2_3, a_tmp3_3,
                    a_path_3, A, a_row_global, a_k_start, stride_a, M, K);
        LOAD_B_DATA(b_mfma_3, b_tile_row, k_base_3, lane, n_base, N, K,
                    row_in_tile, k_group);
        int b_row_global = n_base + row_in_tile;
        int b_scale_k_idx = k_base_3 / SCALE_GROUP + k_group;
        LOAD_B_SCALE_SHUFFLED(b_e8m0_3, B_scale_sh, b_row_global, b_scale_k_idx,
                     N, num_scale_k);

        asm volatile("s_waitcnt vmcnt(6)" ::: "memory");
        PROCESS_TILE(a_tmp0_2, a_tmp1_2, a_tmp2_2, a_tmp3_2,
                     b_mfma_2, b_e8m0_2, a_path_2, acc);
    }

    // ── Process step 3 (final) ──
    {
        asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
        PROCESS_TILE(a_tmp0_3, a_tmp1_3, a_tmp2_3, a_tmp3_3,
                     b_mfma_3, b_e8m0_3, a_path_3, acc);
    }

    // ── Store output ──
    {
        int col = lane & 15;
        int row_base = (lane >> 4) * 4;

        int n_global = n_base + col;
        if (n_global < N) {
            for (int v = 0; v < 4; v++) {
                int m_global = m_wave + row_base + v;
                if (m_global < M) {
                    C[m_global * N + n_global] = f32_to_bf16(acc[v]);
                }
            }
        }
    }
}


// ═══════════════════════════════════════════════════════════════
// FUSED kernel: bf16 A quantize + FP4×FP4 MFMA GEMM (from v47/v91)
// Cross-buffer vmcnt(6) pipeline + 2-wave + s_setprio
// ═══════════════════════════════════════════════════════════════

__global__ __attribute__((amdgpu_flat_work_group_size(128, 128)))
void mxfp4_fused_gemm_crossbuf(
    const uint16_t* __restrict__ A,
    const uint8_t*  __restrict__ B_preshuffle,
    uint16_t*       __restrict__ C,
    const uint8_t*  __restrict__ B_scale_sh,
    int M, int N, int K,
    int stride_bps
) {
    const int wave_id = threadIdx.x / 64;
    const int lane = threadIdx.x % 64;

    int num_tile_n = (N + TILE_N - 1) / TILE_N;
    int tile_m = blockIdx.x / num_tile_n;
    int tile_n = blockIdx.x % num_tile_n;

    int m_base = tile_m * TILE_M_GRID;
    int m_wave = m_base + wave_id * TILE_M_PER_WAVE;
    int n_base = tile_n * TILE_N;

    int stride_a = K;
    int num_scale_k = (K + 31) / 32;

    floatx4 acc = {0.0f, 0.0f, 0.0f, 0.0f};

    int row_in_tile = lane & 15;
    int k_group     = lane >> 4;

    const uint8_t* b_tile_row = B_preshuffle + (n_base / 16) * stride_bps;

    intx4 buf0_a_tmp0, buf0_a_tmp1, buf0_a_tmp2, buf0_a_tmp3;
    intx4 buf0_b_mfma;
    uint8_t buf0_b_e8m0;
    int buf0_a_path;

    intx4 buf1_a_tmp0, buf1_a_tmp1, buf1_a_tmp2, buf1_a_tmp3;
    intx4 buf1_b_mfma;
    uint8_t buf1_b_e8m0;
    int buf1_a_path;

    int num_k_tiles = (K + K_STEP - 1) / K_STEP;

    if (num_k_tiles == 0) goto store_output;

    {
        int k_base = 0;
        int a_row_global = m_wave + row_in_tile;
        int a_k_start = k_base + k_group * SCALE_GROUP;

        LOAD_A_DATA(buf0_a_tmp0, buf0_a_tmp1, buf0_a_tmp2, buf0_a_tmp3,
                    buf0_a_path, A, a_row_global, a_k_start, stride_a, M, K);

        LOAD_B_DATA(buf0_b_mfma, b_tile_row, k_base, lane, n_base, N, K,
                    row_in_tile, k_group);

        int b_row_global = n_base + row_in_tile;
        int b_scale_k_idx = k_base / SCALE_GROUP + k_group;
        LOAD_B_SCALE_SHUFFLED(buf0_b_e8m0, B_scale_sh, b_row_global, b_scale_k_idx,
                     N, num_scale_k);
    }

    if (num_k_tiles == 1) {
        asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
        PROCESS_TILE(buf0_a_tmp0, buf0_a_tmp1, buf0_a_tmp2, buf0_a_tmp3,
                     buf0_b_mfma, buf0_b_e8m0, buf0_a_path, acc);
        goto store_output;
    }

    {
        int k_tile_idx = 0;

        for (; k_tile_idx < num_k_tiles - 1; k_tile_idx += 2) {
            {
                int next_k_base = (k_tile_idx + 1) * K_STEP;
                int a_row_global = m_wave + row_in_tile;
                int a_k_start = next_k_base + k_group * SCALE_GROUP;

                LOAD_A_DATA(buf1_a_tmp0, buf1_a_tmp1, buf1_a_tmp2, buf1_a_tmp3,
                            buf1_a_path, A, a_row_global, a_k_start, stride_a, M, K);

                LOAD_B_DATA(buf1_b_mfma, b_tile_row, next_k_base, lane, n_base, N, K,
                            row_in_tile, k_group);

                int b_row_global = n_base + row_in_tile;
                int b_scale_k_idx = next_k_base / SCALE_GROUP + k_group;
                LOAD_B_SCALE_SHUFFLED(buf1_b_e8m0, B_scale_sh, b_row_global, b_scale_k_idx,
                             N, num_scale_k);

                asm volatile("s_waitcnt vmcnt(6)" ::: "memory");

                PROCESS_TILE(buf0_a_tmp0, buf0_a_tmp1, buf0_a_tmp2, buf0_a_tmp3,
                             buf0_b_mfma, buf0_b_e8m0, buf0_a_path, acc);
            }

            {
                int next_k_tile = k_tile_idx + 2;
                if (next_k_tile < num_k_tiles) {
                    int next_k_base = next_k_tile * K_STEP;
                    int a_row_global = m_wave + row_in_tile;
                    int a_k_start = next_k_base + k_group * SCALE_GROUP;

                    LOAD_A_DATA(buf0_a_tmp0, buf0_a_tmp1, buf0_a_tmp2, buf0_a_tmp3,
                                buf0_a_path, A, a_row_global, a_k_start, stride_a, M, K);

                    LOAD_B_DATA(buf0_b_mfma, b_tile_row, next_k_base, lane, n_base, N, K,
                                row_in_tile, k_group);

                    int b_row_global = n_base + row_in_tile;
                    int b_scale_k_idx = next_k_base / SCALE_GROUP + k_group;
                    LOAD_B_SCALE_SHUFFLED(buf0_b_e8m0, B_scale_sh, b_row_global, b_scale_k_idx,
                                 N, num_scale_k);

                    asm volatile("s_waitcnt vmcnt(6)" ::: "memory");
                } else {
                    asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
                }

                PROCESS_TILE(buf1_a_tmp0, buf1_a_tmp1, buf1_a_tmp2, buf1_a_tmp3,
                             buf1_b_mfma, buf1_b_e8m0, buf1_a_path, acc);
            }
        }

        if (k_tile_idx == num_k_tiles - 1) {
            asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
            PROCESS_TILE(buf0_a_tmp0, buf0_a_tmp1, buf0_a_tmp2, buf0_a_tmp3,
                         buf0_b_mfma, buf0_b_e8m0, buf0_a_path, acc);
        }
    }

store_output:
    {
        int col = lane & 15;
        int row_base = (lane >> 4) * 4;

        int n_global = n_base + col;
        if (n_global < N) {
            for (int v = 0; v < 4; v++) {
                int m_global = m_wave + row_base + v;
                if (m_global < M) {
                    C[m_global * N + n_global] = f32_to_bf16(acc[v]);
                }
            }
        }
    }
}


// ═══════════════════════════════════════════════════════════════
// SPLIT-K kernel (from v91, for splitk shapes)
// ═══════════════════════════════════════════════════════════════

__global__ __attribute__((amdgpu_flat_work_group_size(128, 128)))
void mxfp4_fused_gemm_splitk(
    const uint16_t* __restrict__ A,
    const uint8_t*  __restrict__ B_preshuffle,
    float*          __restrict__ C_f32,
    const uint8_t*  __restrict__ B_scale_sh,
    int M, int N, int K,
    int stride_bps,
    int num_spatial_blocks,
    int k_per_split
) {
    const int wave_id = threadIdx.x / 64;
    const int lane = threadIdx.x % 64;

    int split_id = blockIdx.x / num_spatial_blocks;
    int spatial_idx = blockIdx.x % num_spatial_blocks;

    int num_tile_n = (N + TILE_N - 1) / TILE_N;
    int tile_m = spatial_idx / num_tile_n;
    int tile_n = spatial_idx % num_tile_n;

    int m_base = tile_m * TILE_M_GRID;
    int m_wave = m_base + wave_id * TILE_M_PER_WAVE;
    int n_base = tile_n * TILE_N;

    int k_start = split_id * k_per_split;
    int k_end = k_start + k_per_split;
    if (k_end > K) k_end = K;
    if (k_start >= K) return;

    int stride_a = K;
    int num_scale_k = (K + 31) / 32;

    floatx4 acc = {0.0f, 0.0f, 0.0f, 0.0f};

    int row_in_tile = lane & 15;
    int k_group     = lane >> 4;

    const uint8_t* b_tile_row = B_preshuffle + (n_base / 16) * stride_bps;

    intx4 buf0_a_tmp0, buf0_a_tmp1, buf0_a_tmp2, buf0_a_tmp3;
    intx4 buf0_b_mfma;
    uint8_t buf0_b_e8m0;
    int buf0_a_path;

    intx4 buf1_a_tmp0, buf1_a_tmp1, buf1_a_tmp2, buf1_a_tmp3;
    intx4 buf1_b_mfma;
    uint8_t buf1_b_e8m0;
    int buf1_a_path;

    int first_k_tile = k_start / K_STEP;
    int last_k_tile = (k_end + K_STEP - 1) / K_STEP;
    int num_k_tiles = last_k_tile - first_k_tile;

    if (num_k_tiles == 0) return;

    {
        int k_base = first_k_tile * K_STEP;
        int a_row_global = m_wave + row_in_tile;
        int a_k_start = k_base + k_group * SCALE_GROUP;

        LOAD_A_DATA(buf0_a_tmp0, buf0_a_tmp1, buf0_a_tmp2, buf0_a_tmp3,
                    buf0_a_path, A, a_row_global, a_k_start, stride_a, M, K);

        LOAD_B_DATA(buf0_b_mfma, b_tile_row, k_base, lane, n_base, N, K,
                    row_in_tile, k_group);

        int b_row_global = n_base + row_in_tile;
        int b_scale_k_idx = k_base / SCALE_GROUP + k_group;
        LOAD_B_SCALE_SHUFFLED(buf0_b_e8m0, B_scale_sh, b_row_global, b_scale_k_idx,
                     N, num_scale_k);
    }

    if (num_k_tiles == 1) {
        asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
        PROCESS_TILE(buf0_a_tmp0, buf0_a_tmp1, buf0_a_tmp2, buf0_a_tmp3,
                     buf0_b_mfma, buf0_b_e8m0, buf0_a_path, acc);
        goto store_splitk_output;
    }

    {
        int k_tile_offset = 0;

        for (; k_tile_offset < num_k_tiles - 1; k_tile_offset += 2) {
            {
                int next_abs_k_tile = first_k_tile + k_tile_offset + 1;
                int next_k_base = next_abs_k_tile * K_STEP;
                int a_row_global = m_wave + row_in_tile;
                int a_k_start = next_k_base + k_group * SCALE_GROUP;

                LOAD_A_DATA(buf1_a_tmp0, buf1_a_tmp1, buf1_a_tmp2, buf1_a_tmp3,
                            buf1_a_path, A, a_row_global, a_k_start, stride_a, M, K);

                LOAD_B_DATA(buf1_b_mfma, b_tile_row, next_k_base, lane, n_base, N, K,
                            row_in_tile, k_group);

                int b_row_global = n_base + row_in_tile;
                int b_scale_k_idx = next_k_base / SCALE_GROUP + k_group;
                LOAD_B_SCALE_SHUFFLED(buf1_b_e8m0, B_scale_sh, b_row_global, b_scale_k_idx,
                             N, num_scale_k);

                asm volatile("s_waitcnt vmcnt(6)" ::: "memory");

                PROCESS_TILE(buf0_a_tmp0, buf0_a_tmp1, buf0_a_tmp2, buf0_a_tmp3,
                             buf0_b_mfma, buf0_b_e8m0, buf0_a_path, acc);
            }

            {
                int next_k_tile_offset = k_tile_offset + 2;
                if (next_k_tile_offset < num_k_tiles) {
                    int next_abs_k_tile = first_k_tile + next_k_tile_offset;
                    int next_k_base = next_abs_k_tile * K_STEP;
                    int a_row_global = m_wave + row_in_tile;
                    int a_k_start = next_k_base + k_group * SCALE_GROUP;

                    LOAD_A_DATA(buf0_a_tmp0, buf0_a_tmp1, buf0_a_tmp2, buf0_a_tmp3,
                                buf0_a_path, A, a_row_global, a_k_start, stride_a, M, K);

                    LOAD_B_DATA(buf0_b_mfma, b_tile_row, next_k_base, lane, n_base, N, K,
                                row_in_tile, k_group);

                    int b_row_global = n_base + row_in_tile;
                    int b_scale_k_idx = next_k_base / SCALE_GROUP + k_group;
                    LOAD_B_SCALE_SHUFFLED(buf0_b_e8m0, B_scale_sh, b_row_global, b_scale_k_idx,
                                 N, num_scale_k);

                    asm volatile("s_waitcnt vmcnt(6)" ::: "memory");
                } else {
                    asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
                }

                PROCESS_TILE(buf1_a_tmp0, buf1_a_tmp1, buf1_a_tmp2, buf1_a_tmp3,
                             buf1_b_mfma, buf1_b_e8m0, buf1_a_path, acc);
            }
        }

        if (k_tile_offset == num_k_tiles - 1) {
            asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
            PROCESS_TILE(buf0_a_tmp0, buf0_a_tmp1, buf0_a_tmp2, buf0_a_tmp3,
                         buf0_b_mfma, buf0_b_e8m0, buf0_a_path, acc);
        }
    }

store_splitk_output:
    {
        int col = lane & 15;
        int row_base = (lane >> 4) * 4;

        int n_global = n_base + col;
        if (n_global < N) {
            for (int v = 0; v < 4; v++) {
                int m_global = m_wave + row_base + v;
                if (m_global < M) {
                    atomicAdd(&C_f32[m_global * N + n_global], acc[v]);
                }
            }
        }
    }
}


// ═══════════════════════════════════════════════════════════════
// PyTorch wrappers
// ═══════════════════════════════════════════════════════════════

void launch_mxfp4_fused_gemm(
    torch::Tensor A,
    torch::Tensor B_preshuffle,
    torch::Tensor C,
    torch::Tensor B_scale_sh,
    int M, int N, int K,
    int stride_bps
) {
    TORCH_CHECK(A.is_cuda() && B_preshuffle.is_cuda() && C.is_cuda() && B_scale_sh.is_cuda());

    int num_tile_m = (M + TILE_M_GRID - 1) / TILE_M_GRID;
    int num_tile_n = (N + TILE_N - 1) / TILE_N;
    int num_blocks = num_tile_m * num_tile_n;

    dim3 grid(num_blocks);
    dim3 block(BLOCK_SIZE);

    mxfp4_fused_gemm_crossbuf<<<grid, block>>>(
        reinterpret_cast<const uint16_t*>(A.data_ptr<at::BFloat16>()),
        B_preshuffle.data_ptr<uint8_t>(),
        reinterpret_cast<uint16_t*>(C.data_ptr<at::BFloat16>()),
        B_scale_sh.data_ptr<uint8_t>(),
        M, N, K,
        stride_bps
    );
}

void launch_mxfp4_fused_gemm_splitk(
    torch::Tensor A,
    torch::Tensor B_preshuffle,
    torch::Tensor C_f32,
    torch::Tensor B_scale_sh,
    int M, int N, int K,
    int stride_bps,
    int split_k
) {
    TORCH_CHECK(A.is_cuda() && B_preshuffle.is_cuda() && C_f32.is_cuda() && B_scale_sh.is_cuda());
    TORCH_CHECK(C_f32.dtype() == torch::kFloat32, "C_f32 must be float32 for atomicAdd");

    int num_tile_m = (M + TILE_M_GRID - 1) / TILE_M_GRID;
    int num_tile_n = (N + TILE_N - 1) / TILE_N;
    int num_spatial_blocks = num_tile_m * num_tile_n;

    int k_tiles_total = (K + K_STEP - 1) / K_STEP;
    int k_tiles_per_split = (k_tiles_total + split_k - 1) / split_k;
    int k_per_split = k_tiles_per_split * K_STEP;

    int total_blocks = num_spatial_blocks * split_k;

    dim3 grid(total_blocks);
    dim3 block(BLOCK_SIZE);

    mxfp4_fused_gemm_splitk<<<grid, block>>>(
        reinterpret_cast<const uint16_t*>(A.data_ptr<at::BFloat16>()),
        B_preshuffle.data_ptr<uint8_t>(),
        C_f32.data_ptr<float>(),
        B_scale_sh.data_ptr<uint8_t>(),
        M, N, K,
        stride_bps,
        num_spatial_blocks,
        k_per_split
    );
}


// ═══════════════════════════════════════════════════════════════
// PyTorch wrapper for K=512 2-wave unrolled kernel
// ═══════════════════════════════════════════════════════════════
void launch_mxfp4_fused_gemm_k512_2wave(
    torch::Tensor A,
    torch::Tensor B_preshuffle,
    torch::Tensor C,
    torch::Tensor B_scale_sh,
    int M, int N, int K,
    int stride_bps
) {
    TORCH_CHECK(A.is_cuda() && B_preshuffle.is_cuda() && C.is_cuda() && B_scale_sh.is_cuda());

    // 2-wave kernel: TILE_M_GRID=32 per WG (same as crossbuf)
    int num_tile_m = (M + TILE_M_GRID - 1) / TILE_M_GRID;
    int num_tile_n = (N + TILE_N - 1) / TILE_N;
    int num_blocks = num_tile_m * num_tile_n;

    dim3 grid(num_blocks);
    dim3 block(BLOCK_SIZE);  // 128 threads = 2 wavefronts

    mxfp4_fused_gemm_k512_2wave<<<grid, block>>>(
        reinterpret_cast<const uint16_t*>(A.data_ptr<at::BFloat16>()),
        B_preshuffle.data_ptr<uint8_t>(),
        reinterpret_cast<uint16_t*>(C.data_ptr<at::BFloat16>()),
        B_scale_sh.data_ptr<uint8_t>(),
        M, N, K,
        stride_bps
    );
}


PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("launch_mxfp4_fused_gemm", &launch_mxfp4_fused_gemm,
          "Fused MXFP4 quant+GEMM");
    m.def("launch_mxfp4_fused_gemm_splitk", &launch_mxfp4_fused_gemm_splitk,
          "Fused MXFP4 quant+GEMM Split-K");
    m.def("launch_mxfp4_quant", &launch_mxfp4_quant,
          "Standalone MXFP4 quantization with shuffled scale output");
    m.def("launch_mxfp4_fused_gemm_m16_splitk", &launch_mxfp4_fused_gemm_m16_splitk,
          "Split-K fused MXFP4 quant+GEMM for M=16 with 1-wave 16x16 tile");
    m.def("launch_mxfp4_fused_gemm_m16_multiwave", &launch_mxfp4_fused_gemm_m16_multiwave,
          "Multi-wave fused MXFP4 quant+GEMM for M=16 — single kernel, no workspace");
    m.def("launch_mxfp4_fused_gemm_k512_2wave", &launch_mxfp4_fused_gemm_k512_2wave,
          "K=512 specialized 2-wave unrolled MXFP4 quant+GEMM");
}
"""

# ═══════════════════════════════════════════════════════════════
# Module-level HIP compilation (outside timed region)
# ═══════════════════════════════════════════════════════════════

import os
os.environ['PYTORCH_ROCM_ARCH'] = 'gfx950'
import torch.utils.cpp_extension as cpp_ext
_hip_module = cpp_ext.load_inline(
    name="v177_k512_2wave_only_gfx950",
    cpp_sources="",
    cuda_sources=_HIP_SOURCE,
    extra_cuda_cflags=[
        "--offload-arch=gfx950",
        "-O3",
        "-std=c++17",
    ],
    with_cuda=True,
    verbose=True,
)

# ═══════════════════════════════════════════════════════════════
# AITER ASM GEMM import (outside timed region)
# ═══════════════════════════════════════════════════════════════

from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
from aiter import dtypes as aiter_dtypes

# ═══════════════════════════════════════════════════════════════
# Shape-specific AITER ASM tile/splitK configuration (from v91)
# ═══════════════════════════════════════════════════════════════

SHAPE_CONFIG = {
    # Shape 5: M=64, N=7168, K=2048
    (64, 7168, 2048): (
        "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
        1
    ),
    # Shape 6: M=256, N=3072, K=1536 — log2_k_split=1 from v119 (12.7→12.4µs)
    (256, 3072, 1536): (
        "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
        1
    ),
}

DEFAULT_CONFIG = (
    "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
    3
)


# ═══════════════════════════════════════════════════════════════
# NEW: Fused M=16 path for Shape 2
# ═══════════════════════════════════════════════════════════════

def _fused_m16_path(A, B_shuffle, B_scale_sh, m, k, n):
    """Multi-wave fused quant+GEMM for M=16 with 8-wave 16x16 MFMA tile.
    
    8 wavefronts per WG (512 threads). Each wave processes K/8 K-elements.
    LDS-based reduction across waves. Direct BF16 output.
    NO workspace, NO reduction kernel — single kernel launch!
    """
    B_sh_u8 = B_shuffle if B_shuffle.dtype == torch.uint8 else B_shuffle.view(torch.uint8)

    # B_shuffle is [N, K/2] but preshuffled. We need [N//16, K/2*16] view
    n_tile_groups = n >> 4
    k_half_times_16 = (k >> 1) << 4
    B_preshuffle = B_sh_u8.reshape(n_tile_groups, k_half_times_16)

    B_scale_flat = B_scale_sh if B_scale_sh.dtype == torch.uint8 else B_scale_sh.view(torch.uint8)

    # Output BF16 tensor (direct — no workspace needed!)
    C = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)

    _hip_module.launch_mxfp4_fused_gemm_m16_multiwave(
        A, B_preshuffle, B_scale_flat, C, m, n, k
    )

    return C


# ═══════════════════════════════════════════════════════════════
# HIP kernel dispatch — K=512 path
# v177: K=512 with M>16 uses specialized 2-wave unrolled kernel
# ═══════════════════════════════════════════════════════════════

def _hip_path(A, B_shuffle, B_scale_sh, m, k, n):
    """HIP fused quant+GEMM path for K=512 shapes."""
    B_sh_u8 = B_shuffle if B_shuffle.dtype == torch.uint8 else B_shuffle.view(torch.uint8)

    n_tile_groups = n >> 4
    k_half_times_16 = (k >> 1) << 4
    B_preshuffle = B_sh_u8.reshape(n_tile_groups, k_half_times_16)
    stride_bps = k_half_times_16

    B_scale_flat = B_scale_sh if B_scale_sh.dtype == torch.uint8 else B_scale_sh.view(torch.uint8)

    # K=512 with M>16: use specialized unrolled 2-wave kernel (S3/S4)
    if k == 512 and m > 16:
        C = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
        _hip_module.launch_mxfp4_fused_gemm_k512_2wave(
            A, B_preshuffle, C, B_scale_flat, m, n, k, stride_bps
        )
        return C

    # Default: original crossbuf (handles S1 M=4, and any other K=512 with small M)
    split_k = 4 if (m <= 16 and k >= 4096) else 1

    if split_k > 1:
        C_f32 = torch.zeros((m, n), dtype=torch.float32, device=A.device)
        _hip_module.launch_mxfp4_fused_gemm_splitk(
            A, B_preshuffle, C_f32, B_scale_flat, m, n, k, stride_bps, split_k
        )
        C = C_f32.to(torch.bfloat16)
        return C
    else:
        C = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
        _hip_module.launch_mxfp4_fused_gemm(
            A, B_preshuffle, C, B_scale_flat, m, n, k, stride_bps
        )
        return C


# ═══════════════════════════════════════════════════════════════
# AITER ASM path — HIP quant + ASM GEMM with tuned tile/splitK
# ═══════════════════════════════════════════════════════════════

def _aiter_asm_path(A, B_shuffle, B_scale_sh, m, k, n):
    """HIP quant kernel + AITER ASM GEMM with per-shape tuning."""
    A_q = torch.empty((m, k // 2), dtype=torch.uint8, device=A.device)

    scaleN_valid = (k + 31) // 32
    scaleN = ((scaleN_valid + 7) // 8) * 8
    M_pad256 = ((m + 255) // 256) * 256

    A_scale_sh = torch.empty((M_pad256, scaleN), dtype=torch.uint8, device=A.device)

    _hip_module.launch_mxfp4_quant(A, A_q, A_scale_sh, m, k, scaleN, scaleN_valid)

    config = SHAPE_CONFIG.get((m, n, k), DEFAULT_CONFIG)
    kernelName, log2_k_split = config

    M_pad32 = ((m + 31) // 32) * 32
    out = torch.empty((M_pad32, n), dtype=torch.bfloat16, device=A.device)

    A_q_fp4 = A_q.view(aiter_dtypes.fp4x2)
    B_shuffle_fp4 = B_shuffle.view(aiter_dtypes.fp4x2)

    gemm_a4w4_asm(
        A_q_fp4, B_shuffle_fp4, A_scale_sh, B_scale_sh,
        out, kernelName, None, 1.0, 0.0, log2_k_split
    )

    return out[:m]


# ═══════════════════════════════════════════════════════════════
# Main dispatch
# ═══════════════════════════════════════════════════════════════

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.shape[0]

    if m == 16 and k >= 1024:
        # NEW: Fused M=16 path for Shape 2 (M=16, N=2112, K=7168)
        return _fused_m16_path(A, B_shuffle, B_scale_sh, m, k, n)
    elif k >= 1024:
        # AITER ASM path with shape-specific tile/splitK
        return _aiter_asm_path(A, B_shuffle, B_scale_sh, m, k, n)
    else:
        # Custom HIP fused quant+GEMM for K=512 shapes
        return _hip_path(A, B_shuffle, B_scale_sh, m, k, n)
scrolls · 2034 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