Skip to content
KernelIndex
Search⌘K

submission 647149

npip99 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:2df11d0f8fc16cce4bdd05b3005dc42428e47c1a30ba5d65f98d787482fbecdc
license declaredunknown
license concludedunknown
authorsnpip99
imported2026-08-15

Techniques

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

num-warps = 4constexpr int NUM_WARPS = 4;
shared-memory__shared__ uint8_t s_a_scale[BLOCK_M][K / SCALE_GROUP_SIZE];
split-k- SPLIT_K: Same as general, but with SPLIT_K implemented.
tile-m = 1constexpr uint32_t WARP_TILE_M = 1;
tile-n = 1constexpr uint32_t WARP_TILE_N = 1;

Kernel source

submission.py3129 lines
"""
There are three main kernels implemented:

- General: Global->Register, 1 Warp : 1 Output tile
- SPLIT_K: Same as general, but with SPLIT_K implemented.
  - Used for Shape2 (16, 2112, 7168). Hyperparameters are selected to minimize L2 traffic.
- Global->LDS, LDS->Register, Warp tiling + Threadblock tiling
  - Used for the other 5 shapes

For the other 5 shapes, the C++ code is duplicated. The code is virtually identical, the only difference is the constexprs at the top (i.e. the tiling hyperparameters), and bounds checking.
TODO: Use -D and constexpr evaluations to prevent big copypaste. Constexpr could also evaluate the bounds checks at compile time.
"""

import os
from typing import Any
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950:xnack-"

import torch
from torch.utils.cpp_extension import load_inline

from task import input_t, output_t

CPP_WRAPPER = """
void entry(
    const uintptr_t a, const uintptr_t b,
    const uintptr_t b_q, const uintptr_t b_shuffle, const uintptr_t b_scale_sh,
    uintptr_t c,
    int M, int N, int K
);
"""

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

using f32x4_t = float __attribute__((ext_vector_type(4)));
using i32x8_t = int __attribute__((ext_vector_type(8)));
using i32x4_t = int __attribute__((ext_vector_type(4)));
using u32x8_t = uint32_t __attribute__((ext_vector_type(8)));
using u8x32_t = uint8_t __attribute__((ext_vector_type(32)));
using bf16x2_t = __bf16 __attribute__((ext_vector_type(2)));

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

// Quantize a group of 32 bf16 values -> E8M0 scale (biased uint8)
__device__ inline uint8_t bf16x32_to_scale_e8m0(const uint16_t* row) {
    uint16_t amax_u16 = 0;
    const uint32_t* row32 = (const uint32_t*)row;
#pragma unroll
    for (int i = 0; i < 16; i++) {
        uint32_t v = row32[i] & 0x7FFF7FFFu;   // abs both bf16 in one AND
        uint16_t lo = (uint16_t)(v);
        uint16_t hi = (uint16_t)(v >> 16);
        amax_u16 = max(amax_u16, max(lo, hi));
    }
    // Round exponent up when mantissa fraction >= 0.75
    // (bf16 has 7 mantissa bits; 0x20 = 2^5 matches the f32 0x200000 = 2^21 rounding)
    int biased_exp = ((amax_u16 + 0x20u) >> 7) & 0xFF;
    int e8m0 = biased_exp - 2;
    return (uint8_t)max(0, min(254, e8m0));
}
// Matches `fp4_utils.e8m0_to_f32`: scale_e8m0 << 23 as f32, with special cases
__device__ inline float e8m0_to_f32(uint8_t scale) {
    return __builtin_bit_cast(float, (uint32_t)scale << 23);
}

// Magic shuffle
__device__ inline int read_b_scale_base(int n_row, int k_group, int b_scale_stride) {
    int tr = n_row / 32, r = n_row % 32;
    int tc = k_group / 8, c = k_group % 8;
    int base  = tr * (32 * b_scale_stride) + tc * 256;
    int inner = (r / 16) + ((c / 4) % 2) * 2 + (r % 16) * 4 + (c % 4) * 64;
    return base + inner;
}

constexpr int MFMA_DIM_NM = 16;
constexpr int MFMA_DIM_K = 128;
constexpr int SCALE_GROUP_SIZE = 32;
constexpr int NUM_THREADS = 64;

template <int M, int N, int K>
__launch_bounds__(NUM_THREADS)
__global__ void kernel(
    const uint16_t* a, const uint16_t* b,
    const uint8_t* b_q, const uint8_t* b_shuffle, const uint8_t* b_scale_sh,
    uint16_t* c
) {
    constexpr bool BOUNDS_CHECK = (M % MFMA_DIM_NM != 0) || (N % MFMA_DIM_NM != 0);
    constexpr int b_scale_stride = CDIV(K / 32, 8) * 8;

    int tile_x = MFMA_DIM_NM * blockIdx.x; // N
    int tile_y = MFMA_DIM_NM * blockIdx.y; // M
    int lane_index = threadIdx.x;
    int lane_row_nm = lane_index % MFMA_DIM_NM;
    int lane_col_k = SCALE_GROUP_SIZE * (lane_index / MFMA_DIM_NM); // By chance, a lane sends exactly one scale group.

    // C[y][x] = Sum_{k=0}^{K-1} a[y][k] * b[x][k]
    f32x4_t result = {};
#pragma unroll
    for (int k_offset = 0; k_offset < K; k_offset += MFMA_DIM_K) {
        int k_idx_start = k_offset + lane_col_k;
        uint8_t a_scale = 0;
        u32x8_t a_reg = {};
        if (!BOUNDS_CHECK || tile_y + lane_row_nm < M) {
            // Quantize A
            const uint16_t* a_row = a + (tile_y + lane_row_nm) * K;
            a_scale = bf16x32_to_scale_e8m0(a_row + k_idx_start);
            float a_scale_f32 = e8m0_to_f32(a_scale);
#pragma unroll
            for (int reg_idx = 0; reg_idx < 4; reg_idx++) {
                const bf16x2_t* a_row_v2 = (const bf16x2_t*)(a_row + k_idx_start + reg_idx * 8);
                a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[0], a_scale_f32, 0);
                a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[1], a_scale_f32, 1);
                a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[2], a_scale_f32, 2);
                a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[3], a_scale_f32, 3);
            }
        }
        uint8_t b_scale = b_scale_sh[read_b_scale_base(tile_x + lane_row_nm, k_idx_start / SCALE_GROUP_SIZE, b_scale_stride)];
        i32x8_t b_reg;
        *(i32x4_t*)&b_reg = *(const i32x4_t*)&b_q[(tile_x + lane_row_nm) * (K / 2) + k_idx_start / 2];
        // This is not faster unless we use LDS
        // *(i32x4_t*)&b_reg = *(const i32x4_t*)&b_shuffle[
        //     tile_x * (K / 2)
        //     + (k_idx_start / 32) * 256
        //     + lane_row_nm * 16
        // ];
        result = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg, b_reg, result, 4, 4, 0, a_scale, 0, b_scale);
    }

    int result_row_offset = 4 * (lane_index / 16);
    int result_column_offset = lane_index % 16;
#pragma unroll
    for (int i = 0; i < 4; i++) {
        int row = tile_y + result_row_offset + i;
        if (!BOUNDS_CHECK || row < M) {
            c[row * N + (tile_x + result_column_offset)] = f32_to_bf16((float)result[i]);
        }
    }
}

void entry(
    const uintptr_t a, const uintptr_t b,
    const uintptr_t b_q, const uintptr_t b_shuffle, const uintptr_t b_scale_sh,
    uintptr_t c,
    int M, int N, int K
) {
    assert(K % MFMA_DIM_K == 0);

#define LAUNCH(m,n,k) kernel<m,n,k><<<dim3(CDIV(N, 16), CDIV(M, 16)), NUM_THREADS>>>( \
    (const uint16_t*)a, (const uint16_t*)b, \
    (const uint8_t*)b_q, (const uint8_t*)b_shuffle, (const uint8_t*)b_scale_sh, \
    (uint16_t*)c)

    // Benchmark Shapes
    if (M == M_DIM && N == N_DIM && K == K_DIM) {
        LAUNCH(M_DIM, N_DIM, K_DIM);
    } else {
        assert(false && "Uncompiled (M, N, K) shape");
    }
}
"""

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

using u32x4_t = uint32_t __attribute__((ext_vector_type(4)));
using u16x32_t = uint16_t __attribute__((ext_vector_type(32)));
using bf16x2_t = __bf16 __attribute__((ext_vector_type(2)));
using f32x4_t = float __attribute__((ext_vector_type(4)));
using u32x8_t = uint32_t __attribute__((ext_vector_type(8)));
using i32x4_t = int __attribute__((ext_vector_type(4)));
using i32x8_t = int __attribute__((ext_vector_type(8)));

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

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

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

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

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

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

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

__device__ inline float fp4_e2m1_to_f32(uint8_t nibble) {
    uint8_t s = (nibble >> 3) & 1;
    uint8_t e = (nibble >> 1) & 3;
    uint8_t m = nibble & 1;
    float val = (e == 0) ? m * 0.5f : exp2f((float)e - 1.0f) * (1.0f + m * 0.5f);
    return s ? -val : val;
}

// STRIDE is the width of a row in bytes
// WIDTH is the number of bytes in a single indivisible "column group" (i.e. "greater column").
template <uint32_t WIDTH>
__device__ inline int swizzle(int r, int c) {
    static_assert(WIDTH && !(WIDTH & (WIDTH - 1)), "WIDTH must be power of 2");
    static_assert(WIDTH <= 256, "WIDTH must fit in one bank cycle");
    constexpr int POSITIONS = 256 / WIDTH;  // slots per bank cycle
    constexpr int MASK = POSITIONS - 1;
    int c_group = c / WIDTH;
    int c_offset = c % WIDTH;              // both compile to shift/mask
    int swizzled = (c_group & ~MASK) | ((c_group ^ r) & MASK);
    return swizzled * WIDTH + c_offset;
}

// Constants
constexpr uint32_t MFMA_DIM_NM = 16;
constexpr uint32_t MFMA_DIM_K = 128;
constexpr uint32_t SCALE_GROUP_SIZE = 32;
constexpr uint32_t THREADS_PER_WARP = 64;
constexpr uint32_t NUM_XCD = 8;
#ifdef __gfx950__
constexpr uint32_t NUM_CUS_PER_XCD = 32;

// Parameters
constexpr uint32_t WARP_TILE_M = 1;
constexpr uint32_t WARP_TILE_N = 1;
constexpr uint32_t NUM_WARPS_M = 1;
constexpr uint32_t NUM_WARPS_N = 4;
constexpr uint32_t NUM_CU_M = 1;
constexpr uint32_t NUM_CU_N = 9;
constexpr uint32_t NUM_XCD_M = 1;
constexpr uint32_t NUM_XCD_N = 5;

#else
constexpr uint32_t NUM_CUS_PER_XCD = 38;

// Parameters
constexpr uint32_t WARP_TILE_M = 2;
constexpr uint32_t WARP_TILE_N = 1;
constexpr uint32_t NUM_WARPS_M = 1;
constexpr uint32_t NUM_WARPS_N = 4;
constexpr uint32_t NUM_CU_M = 2;
constexpr uint32_t NUM_CU_N = 14;
constexpr uint32_t NUM_XCD_M = 1;
constexpr uint32_t NUM_XCD_N = 8;
#endif
constexpr uint32_t NUM_CUS = NUM_XCD_M * NUM_XCD_N * NUM_CUS_PER_XCD;

// Derived
constexpr uint32_t THREADS_PER_BLOCK = THREADS_PER_WARP * NUM_WARPS_M * NUM_WARPS_N;

// Magic shuffle
__device__ inline int read_b_scale_base(int n_row, int k_group, int b_scale_stride) {
    int tr = n_row / 32, r = n_row % 32;
    int tc = k_group / 8, c = k_group % 8;
    int base  = tr * (32 * b_scale_stride) + tc * 256;
    int inner = (r / 16) + (c / 4) * 2 + (r % 16) * 4 + (c % 4) * 64;
    return base + inner;
}

// Only used if timing is enabled.
__device__ uint32_t g_run_idx = 0;
__device__ inline uint32_t hash(uint32_t x) {
    x ^= x >> 16;
    x *= 0x45d9f3b;
    x ^= x >> 16;
    x *= 0x45d9f3b;
    x ^= x >> 16;
    return x;
}

__launch_bounds__(THREADS_PER_BLOCK, 1)
__global__ void kernel(
    const uint16_t* a,
    const uint8_t* b_q, const uint8_t* b_shuffle, const uint8_t* b_scale_sh,
    uint16_t* c
) {
    // Timing
    constexpr bool TIMEIT = false;
    uint64_t _ts[16];
    int _ti = 0;
    auto TIME = [&]() {
        if constexpr (TIMEIT) {
            __builtin_amdgcn_sched_barrier(0);
            _ts[_ti++] = __builtin_amdgcn_s_memrealtime();
            __builtin_amdgcn_sched_barrier(0);
        }
    };

    TIME();

    constexpr uint32_t M = 16;
    constexpr uint32_t M_ACTUAL = 4;
    constexpr uint32_t N = 2880;
    constexpr uint32_t K = 512;
    constexpr uint32_t SPLIT_K = 1;
    static_assert(MFMA_DIM_NM * WARP_TILE_M * NUM_WARPS_M * NUM_CU_M * NUM_XCD_M == M);
    static_assert(MFMA_DIM_NM * WARP_TILE_N * NUM_WARPS_N * NUM_CU_N * NUM_XCD_N == N);
    constexpr int b_scale_stride = CDIV(K / 32, 8) * 8;

    constexpr int BLOCK_M = MFMA_DIM_NM * WARP_TILE_M * NUM_WARPS_M;
    constexpr int BLOCK_N = MFMA_DIM_NM * WARP_TILE_N * NUM_WARPS_N;

    // Get XCD coordinate
    uint32_t xcd_id = blockIdx.x / NUM_CUS_PER_XCD;
    uint32_t xcd_worker_id = blockIdx.x % NUM_CUS_PER_XCD;
    uint32_t xcd_n = xcd_id % NUM_XCD_N;
    uint32_t xcd_m = xcd_id / NUM_XCD_N;

    // Get CU coordinate
    if (xcd_worker_id >= NUM_CU_M*NUM_CU_N) {
        return;
    }
    uint32_t cu_n = xcd_worker_id % NUM_CU_N;
    uint32_t cu_m = xcd_worker_id / NUM_CU_N;

    // Get warp id (Use SGPR for it)
    uint32_t warp_id = __builtin_amdgcn_readfirstlane(threadIdx.x / THREADS_PER_WARP);
    uint32_t warp_n = warp_id % NUM_WARPS_N;
    uint32_t warp_m = warp_id / NUM_WARPS_N;

    // Get lane assignment
    uint32_t lane_index = threadIdx.x % THREADS_PER_WARP;
    uint32_t lane_row_nm = lane_index % MFMA_DIM_NM;
    uint32_t lane_col_k = SCALE_GROUP_SIZE * (lane_index / MFMA_DIM_NM); // By chance, a lane sends exactly one scale group.

    // Get tile offset
    uint32_t tile_n_offset = xcd_n * (N / NUM_XCD_N) + cu_n * (N / NUM_XCD_N / NUM_CU_N);
    uint32_t tile_m_offset = xcd_m * (M / NUM_XCD_M) + cu_m * (M / NUM_XCD_M / NUM_CU_M);
    uint32_t warp_n_offset = warp_n * MFMA_DIM_NM * WARP_TILE_N;
    uint32_t warp_m_offset = warp_m * MFMA_DIM_NM * WARP_TILE_M;

    // Shared Memory
    __shared__ uint8_t s_a_scale[BLOCK_M][K / SCALE_GROUP_SIZE];
    __shared__ uint8_t s_a[BLOCK_M][K / 2];
    __shared__ uint8_t s_b_scale[BLOCK_N][K / SCALE_GROUP_SIZE];
    __shared__ uint8_t s_b[BLOCK_N][K / 2];
    auto swizzle_k = [](int r, int c) { return swizzle<SCALE_GROUP_SIZE / 2>(r, c); };

    // ======================
    // A: Global -> Register
    // ======================

    TIME();

    constexpr uint32_t A_CHUNKS_K = K / SCALE_GROUP_SIZE;
    constexpr uint32_t A_CHUNK_JOBS = BLOCK_M * A_CHUNKS_K;
    constexpr uint32_t A_CHUNK_JOBS_PER_THREAD = CDIV(A_CHUNK_JOBS, THREADS_PER_BLOCK);
    u16x32_t a_jobs[A_CHUNK_JOBS_PER_THREAD];
    uint8_t a_scale_jobs[A_CHUNK_JOBS_PER_THREAD];
#pragma unroll
    for (uint32_t a_chunk_job_idx = 0; a_chunk_job_idx < A_CHUNK_JOBS_PER_THREAD; a_chunk_job_idx++) {
        uint32_t a_chunk_job = a_chunk_job_idx * THREADS_PER_BLOCK + threadIdx.x;
        uint32_t a_transfer_k = (a_chunk_job % A_CHUNKS_K) * (K / A_CHUNKS_K);
        uint32_t a_transfer_m = a_chunk_job / A_CHUNKS_K;
        uint32_t a_row_idx = tile_m_offset + a_transfer_m;
        if (a_row_idx < M_ACTUAL) {
            *(u16x32_t*)(&a_jobs[a_chunk_job_idx]) = *(const u16x32_t*)(a + a_row_idx * K + a_transfer_k);
        } else {
            a_jobs[a_chunk_job_idx] = {};
        }
    }
    auto a_process_job_step1 = [&](int a_chunk_job_idx) {
        __builtin_amdgcn_sched_barrier(0);
        uint32_t a_chunk_job = a_chunk_job_idx * THREADS_PER_BLOCK + threadIdx.x;
        uint32_t a_transfer_k = (a_chunk_job % A_CHUNKS_K) * (K / A_CHUNKS_K);
        uint32_t a_transfer_m = a_chunk_job / A_CHUNKS_K;
        uint8_t a_scale = bf16x32_to_scale_e8m0(&a_jobs[a_chunk_job_idx]);
        s_a_scale[a_transfer_m][a_transfer_k / SCALE_GROUP_SIZE] = a_scale;
        a_scale_jobs[a_chunk_job_idx] = a_scale;
        __builtin_amdgcn_sched_barrier(0);
    };
    auto a_process_job_step2 = [&](int a_chunk_job_idx) {
        __builtin_amdgcn_sched_barrier(0);
        uint32_t a_chunk_job = a_chunk_job_idx * THREADS_PER_BLOCK + threadIdx.x;
        uint32_t a_transfer_k = (a_chunk_job % A_CHUNKS_K) * (K / A_CHUNKS_K);
        uint32_t a_transfer_m = a_chunk_job / A_CHUNKS_K;
        float a_scale_f32 = e8m0_to_f32(a_scale_jobs[a_chunk_job_idx]);
        u32x4_t a_reg;
#pragma unroll
        for (int reg_idx = 0; reg_idx < 4; reg_idx++) {
            const bf16x2_t* a_row_v2 = (const bf16x2_t*)(&a_jobs[a_chunk_job_idx]) + 4 * reg_idx;
#ifdef __gfx950__
            a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(0, a_row_v2[0], a_scale_f32, 0);
            a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[1], a_scale_f32, 1);
            a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[2], a_scale_f32, 2);
            a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[3], a_scale_f32, 3);
#else
            a_reg[reg_idx] = __builtin_bit_cast(uint32_t, a_row_v2[0]) ^ __builtin_bit_cast(uint32_t, a_row_v2[1])
                           ^ __builtin_bit_cast(uint32_t, a_row_v2[2]) ^ __builtin_bit_cast(uint32_t, a_row_v2[3])
                           ^ __builtin_bit_cast(uint32_t, a_scale_f32);
#endif
        }
        *(u32x4_t*)&s_a[a_transfer_m][swizzle_k(a_transfer_m, a_transfer_k / 2)] = a_reg;
        __builtin_amdgcn_sched_barrier(0);
    };
    uint32_t a_work_idx = 0;

    // ======================
    // B: Global -> Register
    // ======================

    TIME();

    constexpr uint32_t B_BYTES_PER_JOB = 16;
    constexpr uint32_t B_TOTAL_BYTES = BLOCK_N * (K / 2);
    constexpr uint32_t B_BYTES_PER_JOBGROUP = THREADS_PER_BLOCK * B_BYTES_PER_JOB;
    constexpr uint32_t B_NUM_JOBGROUPS = B_TOTAL_BYTES / B_BYTES_PER_JOBGROUP;
    static_assert(B_TOTAL_BYTES % B_BYTES_PER_JOBGROUP == 0);

    u32x4_t b_jobs[B_NUM_JOBGROUPS];
    uint8_t b_scale_jobs[B_NUM_JOBGROUPS];
    auto b_flat_byte = [&](uint32_t b_jobgroup_idx) -> uint32_t {
        return b_jobgroup_idx * B_BYTES_PER_JOBGROUP + threadIdx.x * B_BYTES_PER_JOB;
    };
    const uint8_t* b_base = b_shuffle + tile_n_offset * (K / 2);
    uint32_t b_thread_offset = threadIdx.x * B_BYTES_PER_JOB;

    auto b_process_job_step0 = [&](int b_jobgroup_idx) {
        __builtin_amdgcn_sched_barrier(0);
        uint32_t b_jobgroup_start_offset = b_jobgroup_idx * B_BYTES_PER_JOBGROUP;
        uint32_t b_job_offset = b_jobgroup_start_offset + b_thread_offset;
        uint32_t b_transfer_n = b_job_offset / (K / 2);
        uint32_t b_transfer_k_byte = b_job_offset % (K / 2);
        uint32_t b_row_idx = tile_n_offset + b_transfer_n;

        b_scale_jobs[b_jobgroup_idx] = b_scale_sh[read_b_scale_base(b_row_idx, b_transfer_k_byte / B_BYTES_PER_JOB, b_scale_stride)];

        uint32_t warp_base = b_jobgroup_start_offset + warp_id * (THREADS_PER_WARP * B_BYTES_PER_JOB);
#ifdef __gfx950__
        void* src = (void*)(b_base + warp_base + lane_index * 16);
        uint32_t lds_off = (uint32_t)(uintptr_t)((uint8_t*)s_b + warp_base);
        asm volatile(
            "s_mov_b32 m0, %0\n\t"
            "global_load_lds_dwordx4 %1, off\n\t"
            :: "s"(lds_off), "v"(src)
            : "memory", "m0"
        );
#else
#pragma unroll
        for (uint32_t sub = 0; sub < 4; sub++) {
            uint32_t chunk_base = warp_base + sub * 256;
            void* src = (void*)(b_base + chunk_base + lane_index * 4);
            uint32_t lds_off = (uint32_t)(uintptr_t)((uint8_t*)s_b + chunk_base);
            asm volatile(
                "s_mov_b32 m0, %0\n\t"
                "global_load_lds_dword %1, off\n\t"
                :: "s"(lds_off), "v"(src)
                : "memory", "m0"
            );
        }
#endif
        
        __builtin_amdgcn_sched_barrier(0);
    };

    auto b_process_job_step1 = [&](int b_jobgroup_idx) {
        __builtin_amdgcn_sched_barrier(0);
        uint32_t b_jobgroup_start_offset = b_jobgroup_idx * B_BYTES_PER_JOBGROUP;
        uint32_t b_job_offset = b_jobgroup_start_offset + b_thread_offset;
        uint32_t b_transfer_n = b_job_offset / (K / 2);
        uint32_t b_transfer_k_byte = b_job_offset % (K / 2);

        asm volatile("" ::: "memory");
        s_b_scale[b_transfer_n][b_transfer_k_byte / B_BYTES_PER_JOB] = b_scale_jobs[b_jobgroup_idx];
        asm volatile("" ::: "memory");
        __builtin_amdgcn_sched_barrier(0);
    };

#pragma unroll
    for (uint32_t b_jobgroup_idx = 0; b_jobgroup_idx < B_NUM_JOBGROUPS; b_jobgroup_idx++) {
        b_process_job_step0(b_jobgroup_idx);

        __builtin_amdgcn_sched_barrier(0);
        if (b_jobgroup_idx == 1) {
            a_process_job_step1(a_work_idx / 2);
        } else if (b_jobgroup_idx == 2) {
            a_process_job_step2(a_work_idx / 2);
        } else if (b_jobgroup_idx == 3) {
            b_process_job_step1(0);
        }
        __builtin_amdgcn_sched_barrier(0);
    }

    // ======================
    // B: Register -> LDS
    // ======================

    TIME();

#pragma unroll
    for (uint32_t b_jobgroup_idx = 1; b_jobgroup_idx < B_NUM_JOBGROUPS; b_jobgroup_idx++) {
        b_process_job_step1(b_jobgroup_idx);
    }

    // ======================
    // A: Register -> Quantize -> LDS
    // ======================

    TIME();

    // NOTE: This task was interwoven above. This while loop should not be hit.
// #pragma unroll
//     while (a_work_idx < A_CHUNK_JOBS_PER_THREAD * 2) {
//         if (a_work_idx % 2 == 0) {
//             a_process_job_step1(a_work_idx / 2);
//         } else {
//             a_process_job_step2(a_work_idx / 2);
//         }
//         a_work_idx += 1;
//     }

    // ======================
    // LDS->MFMA
    // ======================

    __syncthreads(); // Ensure LDS is populated

    TIME();

    f32x4_t result[WARP_TILE_M][WARP_TILE_N] = {};
    constexpr uint32_t K_ITERS = K / (SPLIT_K * MFMA_DIM_K);
    uint32_t k_start = (SPLIT_K > 1) ? blockIdx.z * (K / SPLIT_K) : 0;
    constexpr uint32_t PREFETCH = 2;

    uint8_t a_scale[PREFETCH][WARP_TILE_M];
    u32x8_t a_reg[PREFETCH][WARP_TILE_M];
    uint8_t b_scale[PREFETCH][WARP_TILE_N];
    u32x8_t b_reg[PREFETCH][WARP_TILE_N];
    auto load_lds = [&](uint32_t k_iter) {
        __builtin_amdgcn_sched_barrier(0);
        uint32_t k_offset = k_start + k_iter * MFMA_DIM_K;
        uint32_t k_idx_start = k_offset + lane_col_k;
        uint32_t buf_idx = k_iter % PREFETCH;

#pragma unroll
        for (uint32_t i = 0; i < WARP_TILE_M; i++) {
            uint32_t a_row = warp_m_offset + i * MFMA_DIM_NM + lane_row_nm;
            a_scale[buf_idx][i] = s_a_scale[a_row][k_idx_start / SCALE_GROUP_SIZE];
            *(u32x4_t*)&a_reg[buf_idx][i] = *(const u32x4_t*)&s_a[a_row][swizzle_k(a_row, k_idx_start / 2)];
        }
#pragma unroll
        for (uint32_t i = 0; i < WARP_TILE_N; i++) {
            uint32_t b_row = warp_n_offset + i * MFMA_DIM_NM + lane_row_nm;
            b_scale[buf_idx][i] = s_b_scale[b_row][k_idx_start / SCALE_GROUP_SIZE];
            uint32_t off = (warp_n_offset + i * MFMA_DIM_NM) * (K / 2)
                        + (k_idx_start / 32) * 256
                        + lane_row_nm * 16;
            *(u32x4_t*)&b_reg[buf_idx][i] = *(const u32x4_t*)((const uint8_t*)s_b + off);
        }
        __builtin_amdgcn_sched_barrier(0);
    };

#pragma unroll
    for (uint32_t i = 0; i < min(K_ITERS, PREFETCH); i++) {
        load_lds(i);
    }
#pragma unroll
    for (uint32_t k_iter = 0; k_iter < K_ITERS; k_iter++) {
        uint32_t k_prefetch_iter = k_iter + PREFETCH - 1;
        if (k_prefetch_iter < K_ITERS) {
            load_lds(k_prefetch_iter);
        }
        uint32_t k_offset = k_start + k_iter * MFMA_DIM_K;
        uint32_t k_idx_start = k_offset + lane_col_k;
        uint32_t buf_idx = k_iter % PREFETCH;

#pragma unroll
        for (uint32_t i = 0; i < WARP_TILE_M; i++) {
#pragma unroll
            for(uint32_t j = 0; j < WARP_TILE_N; j++) {
#ifdef __gfx950__
                result[i][j] = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg[buf_idx][i], b_reg[buf_idx][j], result[i][j], 4, 4, 0, a_scale[buf_idx][i], 0, b_scale[buf_idx][j]);
#else
                float scale = e8m0_to_f32(a_scale[buf_idx][i]) * e8m0_to_f32(b_scale[buf_idx][j]);
                f32x4_t mix;
                for (uint32_t r = 0; r < 4; r++)
                    mix[r] = scale * fp4_e2m1_to_f32(a_reg[buf_idx][i][r]) * fp4_e2m1_to_f32(b_reg[buf_idx][j][r]);
                result[i][j] += mix;
#endif
            }
        }
    }

    // ======================
    // Reg->Write to Global
    // ======================

    TIME();

#pragma unroll
    for (uint32_t i = 0; i < WARP_TILE_M; i++) {
#pragma unroll
        for (uint32_t j = 0; j < WARP_TILE_N; j++) {
            uint32_t result_row_offset = 4 * (lane_index / 16);
            uint32_t result_column_offset = lane_index % 16;
#pragma unroll
            for (uint32_t f = 0; f < 4; f++) {
                uint32_t row = tile_m_offset + warp_m_offset + i*MFMA_DIM_NM + result_row_offset + f;
                uint32_t col = tile_n_offset + warp_n_offset + j*MFMA_DIM_NM + result_column_offset;
                if (row < M_ACTUAL) {
                    c[row * N + col] = f32_to_bf16((float)result[i][j][f]);
                }
            }
        }
    }

    TIME();

    if constexpr (TIMEIT) {
        __shared__ uint32_t _run_idx;
        if (threadIdx.x == 0) {
            _run_idx = atomicAdd(&g_run_idx, 1) / NUM_CUS;
        }
        __syncthreads();
        uint32_t h = hash(_run_idx);
        if (threadIdx.x == (h % THREADS_PER_BLOCK) && blockIdx.x == (h % NUM_CUS)) {
            printf("%u:CU[%u,%u]:", g_run_idx, threadIdx.x, blockIdx.x);
            for (int i = 1; i < _ti; i++) {
                printf(" %llu", (unsigned long long)(_ts[i] - _ts[0]));
            }
            printf("\n");
        }
    }
}

void entry(
    const uintptr_t a, const uintptr_t b,
    const uintptr_t b_q, const uintptr_t b_shuffle, const uintptr_t b_scale_sh,
    uintptr_t c,
    int M, int N, int K
) {
    if (M == 4 && N == 2880 && K == 512) {
        kernel<<<dim3(NUM_CUS), dim3(THREADS_PER_BLOCK)>>>(
            (const uint16_t*)a,
            (const uint8_t*)b_q, (const uint8_t*)b_shuffle, (const uint8_t*)b_scale_sh,
            (uint16_t*)c
        );
    } else {
        // No impl
    }
}
"""

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

using f32x4_t = float __attribute__((ext_vector_type(4)));
using i32x4_t = int __attribute__((ext_vector_type(4)));
using u32x8_t = uint32_t __attribute__((ext_vector_type(8)));
using i32x8_t = int __attribute__((ext_vector_type(8)));
using bf16x2_t = __bf16 __attribute__((ext_vector_type(2)));

__device__ inline uint16_t f32_to_bf16(float v) {
    return (uint16_t)(__builtin_bit_cast(uint32_t, v) >> 16);
}

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

__device__ inline float e8m0_to_f32(uint8_t scale) {
    return __builtin_bit_cast(float, (uint32_t)scale << 23);
}

__device__ inline int read_b_scale_base(int n_row, int k_group, int b_scale_stride) {
    int tr = n_row / 32, r = n_row % 32;
    int tc = k_group / 8, c = k_group % 8;
    int base  = tr * (32 * b_scale_stride) + tc * 256;
    int inner = (r / 16) + ((c / 4) % 2) * 2 + (r % 16) * 4 + (c % 4) * 64;
    return base + inner;
}

// Problem
constexpr int M = 16;
constexpr int N_ACTUAL = 2112;
constexpr int N = 2176;
constexpr int K = 7168;
constexpr int SPLIT_K = 7;
constexpr int BLOCK_K = K / SPLIT_K;  // 1024

// Hardware
constexpr int MFMA_DIM_NM = 16;
constexpr int MFMA_DIM_K = 128;
constexpr int SCALE_GROUP_SIZE = 32;
constexpr int THREADS_PER_WARP = 64;
constexpr int NUM_WARPS = 4;
constexpr int NUM_THREADS = THREADS_PER_WARP * NUM_WARPS;  // 256
constexpr int NUM_XCD_N = 8;

#ifdef __gfx950__
constexpr int NUM_CUS_PER_XCD = 32;
#else
constexpr int NUM_CUS_PER_XCD = 38;
#endif

// Tiling
constexpr int NUM_CU_N = 17;  // N / (MFMA_DIM_NM * NUM_XCD_N) = 2176 / 128 = 17
constexpr int WORK_PER_XCD = NUM_CU_N * SPLIT_K;  // 17 * 7 = 119
constexpr int MAX_WORK_PER_CU = CDIV(WORK_PER_XCD, NUM_CUS_PER_XCD);  // 4
constexpr int K_ITERS = BLOCK_K / MFMA_DIM_K;  // 8
constexpr int b_scale_stride = CDIV(K / 32, 8) * 8;

static_assert(N == MFMA_DIM_NM * NUM_CU_N * NUM_XCD_N);
static_assert(K % SPLIT_K == 0);
static_assert(BLOCK_K % MFMA_DIM_K == 0);
static_assert(MAX_WORK_PER_CU == NUM_WARPS);

constexpr int NUM_CUS = NUM_XCD_N * NUM_CUS_PER_XCD;

__launch_bounds__(NUM_THREADS, 1)
__global__ void kernel(
    const uint16_t* a, const uint16_t* b,
    const uint8_t* b_q, const uint8_t* b_shuffle, const uint8_t* b_scale_sh,
    uint16_t* c, float* c_f32, int* c_counter
) {
    // XCD coordinate
    int xcd_id = blockIdx.x / NUM_CUS_PER_XCD;
    int xcd_worker_id = blockIdx.x % NUM_CUS_PER_XCD;
    if (xcd_id >= NUM_XCD_N) return;

    int xcd_n = xcd_id;
    int warp_id = threadIdx.x / THREADS_PER_WARP;
    int lane_index = threadIdx.x % THREADS_PER_WARP;

    // Work assignment: k-major for L1 A reuse
    // Consecutive work_ids share k_slice -> same A rows -> L1 hits
    int work_id = xcd_worker_id * MAX_WORK_PER_CU + warp_id;
    if (work_id >= WORK_PER_XCD) return;

    int k_slice = work_id / NUM_CU_N;
    int n_tile  = work_id % NUM_CU_N;

    int tile_x = xcd_n * (N / NUM_XCD_N) + n_tile * MFMA_DIM_NM;
    int tile_y = 0;  // M=16 = one tile
    int k_base = k_slice * BLOCK_K;

    // Branchless OOB: clamp to last valid tile for B loads
    int tile_x_safe = min(tile_x, N_ACTUAL - MFMA_DIM_NM);

    int lane_row_nm = lane_index % MFMA_DIM_NM;
    int lane_col_k = SCALE_GROUP_SIZE * (lane_index / MFMA_DIM_NM);

    // Output info
    int result_row_offset = 4 * (lane_index / 16);
    int result_column_offset = lane_index % 16;
    int col = tile_x + result_column_offset;
    if (col >= N_ACTUAL) {
        return;
    }
    
    // Prefetch buffers: raw A bf16 (16 dwords = 32 bf16), B fp4 (4 dwords), B scale (1 byte)
    constexpr uint32_t PREFETCH = 3;
    uint32_t a_raw[PREFETCH][MFMA_DIM_NM];
    i32x4_t  b_raw[PREFETCH];
    uint8_t  b_sc[PREFETCH];

    int b_row = tile_x_safe + lane_row_nm;
    const uint16_t* a_base = a + (tile_y + lane_row_nm) * K;
    auto loadGlobal = [&](int k_iter) {
        __builtin_amdgcn_sched_barrier(0);
        int buf = k_iter % PREFETCH;
        int k_offset = k_base + k_iter * MFMA_DIM_K;
        int k_idx_start = k_offset + lane_col_k;

        const uint16_t* a_ptr = a_base + k_idx_start;
#pragma unroll
        for (int r = 0; r < 4; r++) {
            *(i32x4_t*)&a_raw[buf][r * 4] = *(const i32x4_t*)(a_ptr + r * 8);
        }

        b_raw[buf] = *(const i32x4_t*)&b_q[b_row * (K / 2) + k_idx_start / 2];
        b_sc[buf] = b_scale_sh[read_b_scale_base(b_row, k_idx_start / SCALE_GROUP_SIZE, b_scale_stride)];
        __builtin_amdgcn_sched_barrier(0);
    };

    // Fill pipeline
#pragma unroll
    for (int i = 0; i < PREFETCH - 1; i++) {
        if (i < K_ITERS) loadGlobal(i);
    }

    // Prefetch + Quant + MFMA
    f32x4_t result = {};
#pragma unroll
    for (int k_iter = 0; k_iter < K_ITERS; k_iter++) {
        // Prefetch next iteration
        int pf = k_iter + PREFETCH - 1;
        if (pf < K_ITERS) {
            asm volatile("" ::: "memory");
            loadGlobal(pf);
        }

        // Consume from prefetch buffer
        int buf = k_iter % PREFETCH;

        // Quantize A from buffered bf16
        uint8_t a_scale = bf16x32_to_scale_e8m0(a_raw[buf]);
        float a_scale_f32 = e8m0_to_f32(a_scale);
        u32x8_t a_reg = {};
#pragma unroll
        for (int reg_idx = 0; reg_idx < 4; reg_idx++) {
            const bf16x2_t* v2 = (const bf16x2_t*)&a_raw[buf][reg_idx * 4];
#ifdef __gfx950__
            a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], v2[0], a_scale_f32, 0);
            a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], v2[1], a_scale_f32, 1);
            a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], v2[2], a_scale_f32, 2);
            a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], v2[3], a_scale_f32, 3);
#else
            a_reg[reg_idx] = __builtin_bit_cast(uint32_t, v2[0]) ^ __builtin_bit_cast(uint32_t, v2[1])
                           ^ __builtin_bit_cast(uint32_t, v2[2]) ^ __builtin_bit_cast(uint32_t, v2[3])
                           ^ __builtin_bit_cast(uint32_t, a_scale_f32);
#endif
        }

        uint8_t b_scale = b_sc[buf];
        i32x8_t b_reg;
        *(i32x4_t*)&b_reg = b_raw[buf];

#ifdef __gfx950__
        result = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg, b_reg, result, 4, 4, 0, a_scale, 0, b_scale);
#else
        (void)a_reg; (void)b_reg; (void)a_scale; (void)b_scale;
#endif
    }

    // Write output
    int row_base = tile_y + result_row_offset;
#pragma unroll
    for (int i = 0; i < 4; i++) {
        int row = row_base + i;
        atomicAdd(&c_f32[row * N_ACTUAL + col], (float)result[i]);
    }
    // See who's job it is to write.
    asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
    int done = atomicAdd(&c_counter[row_base * N_ACTUAL + col], 1) + 1;
    if (done == SPLIT_K) {
#pragma unroll
        for (int i = 0; i < 4; i++) {
            int row = row_base + i;
            c[row * N_ACTUAL + col] = f32_to_bf16(c_f32[row * N_ACTUAL + col]);
            c_f32[row * N_ACTUAL + col] = 0.0f;
        }
        c_counter[row_base * N_ACTUAL + col] = 0;
    }
}

constexpr uint32_t c_max_elems = M * N;
void entry(
    const uintptr_t a, const uintptr_t b,
    const uintptr_t b_q, const uintptr_t b_shuffle, const uintptr_t b_scale_sh,
    uintptr_t c,
    int M, int N, int K
) {
    static float* c_f32 = nullptr;
    if (!c_f32) {
        assert(hipMalloc(&c_f32, c_max_elems * sizeof(float)) == hipSuccess);
        assert(hipMemsetAsync(c_f32, 0, c_max_elems * sizeof(float)) == hipSuccess);
    }
    static int* c_counter = nullptr;
    if (!c_counter) {
        assert(hipMalloc(&c_counter, c_max_elems * sizeof(int)) == hipSuccess);
        assert(hipMemsetAsync((void*)c_counter, 0, c_max_elems * sizeof(int)) == hipSuccess);
    }

    if (M == 16 && N == 2112 && K == 7168) {
        kernel<<<dim3(NUM_CUS), dim3(NUM_THREADS)>>>(
            (const uint16_t*)a, (const uint16_t*)b,
            (const uint8_t*)b_q, (const uint8_t*)b_shuffle, (const uint8_t*)b_scale_sh,
            (uint16_t*)c, c_f32, c_counter
        );
    } else {
        // Do nothing
    }
}
"""

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

using u32x4_t = uint32_t __attribute__((ext_vector_type(4)));
using u16x32_t = uint16_t __attribute__((ext_vector_type(32)));
using bf16x2_t = __bf16 __attribute__((ext_vector_type(2)));
using f32x4_t = float __attribute__((ext_vector_type(4)));
using u32x8_t = uint32_t __attribute__((ext_vector_type(8)));
using i32x4_t = int __attribute__((ext_vector_type(4)));
using i32x8_t = int __attribute__((ext_vector_type(8)));

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

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

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

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

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

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

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

__device__ inline float fp4_e2m1_to_f32(uint8_t nibble) {
    uint8_t s = (nibble >> 3) & 1;
    uint8_t e = (nibble >> 1) & 3;
    uint8_t m = nibble & 1;
    float val = (e == 0) ? m * 0.5f : exp2f((float)e - 1.0f) * (1.0f + m * 0.5f);
    return s ? -val : val;
}

// STRIDE is the width of a row in bytes
// WIDTH is the number of bytes in a single indivisible "column group" (i.e. "greater column").
template <uint32_t WIDTH>
__device__ inline int swizzle(int r, int c) {
    static_assert(WIDTH && !(WIDTH & (WIDTH - 1)), "WIDTH must be power of 2");
    static_assert(WIDTH <= 256, "WIDTH must fit in one bank cycle");
    constexpr int POSITIONS = 256 / WIDTH;  // slots per bank cycle
    constexpr int MASK = POSITIONS - 1;
    int c_group = c / WIDTH;
    int c_offset = c % WIDTH;              // both compile to shift/mask
    int swizzled = (c_group & ~MASK) | ((c_group ^ r) & MASK);
    return swizzled * WIDTH + c_offset;
}

// Constants
constexpr uint32_t MFMA_DIM_NM = 16;
constexpr uint32_t MFMA_DIM_K = 128;
constexpr uint32_t SCALE_GROUP_SIZE = 32;
constexpr uint32_t THREADS_PER_WARP = 64;
constexpr uint32_t NUM_XCD = 8;
#ifdef __gfx950__
constexpr uint32_t NUM_CUS_PER_XCD = 32;

// Parameters
constexpr uint32_t WARP_TILE_M = 1;
constexpr uint32_t WARP_TILE_N = 1;
constexpr uint32_t NUM_WARPS_M = 1;
constexpr uint32_t NUM_WARPS_N = 4;
constexpr uint32_t NUM_CU_M = 2;
constexpr uint32_t NUM_CU_N = 8;
constexpr uint32_t NUM_XCD_M = 1;
constexpr uint32_t NUM_XCD_N = 8;

#else
constexpr uint32_t NUM_CUS_PER_XCD = 38;

// Parameters
constexpr uint32_t WARP_TILE_M = 2;
constexpr uint32_t WARP_TILE_N = 1;
constexpr uint32_t NUM_WARPS_M = 1;
constexpr uint32_t NUM_WARPS_N = 4;
constexpr uint32_t NUM_CU_M = 2;
constexpr uint32_t NUM_CU_N = 14;
constexpr uint32_t NUM_XCD_M = 1;
constexpr uint32_t NUM_XCD_N = 8;
#endif
constexpr uint32_t NUM_CUS = NUM_XCD * NUM_CUS_PER_XCD;

// Derived
constexpr uint32_t THREADS_PER_BLOCK = THREADS_PER_WARP * NUM_WARPS_M * NUM_WARPS_N;

// Magic shuffle
__device__ inline int read_b_scale_base(int n_row, int k_group, int b_scale_stride) {
    int tr = n_row / 32, r = n_row % 32;
    int tc = k_group / 8, c = k_group % 8;
    int base  = tr * (32 * b_scale_stride) + tc * 256;
    int inner = (r / 16) + (c / 4) * 2 + (r % 16) * 4 + (c % 4) * 64;
    return base + inner;
}

// Only used if timing is enabled.
__device__ uint32_t g_run_idx = 0;
__device__ inline uint32_t hash(uint32_t x) {
    x ^= x >> 16;
    x *= 0x45d9f3b;
    x ^= x >> 16;
    x *= 0x45d9f3b;
    x ^= x >> 16;
    return x;
}

__launch_bounds__(THREADS_PER_BLOCK, 1)
__global__ void kernel(
    const uint16_t* a,
    const uint8_t* b_q, const uint8_t* b_shuffle, const uint8_t* b_scale_sh,
    uint16_t* c
) {
    // Timing
    constexpr bool TIMEIT = false;
    uint64_t _ts[16];
    int _ti = 0;
    auto TIME = [&]() {
        if constexpr (TIMEIT) {
            __builtin_amdgcn_sched_barrier(0);
            _ts[_ti++] = __builtin_amdgcn_s_memrealtime();
            __builtin_amdgcn_sched_barrier(0);
        }
    };

    TIME();

    constexpr uint32_t M = 32;
    constexpr uint32_t N = 4096;
    constexpr uint32_t K = 512;
    constexpr uint32_t SPLIT_K = 1;
    static_assert(MFMA_DIM_NM * WARP_TILE_M * NUM_WARPS_M * NUM_CU_M * NUM_XCD_M == M);
    static_assert(MFMA_DIM_NM * WARP_TILE_N * NUM_WARPS_N * NUM_CU_N * NUM_XCD_N == N);
    constexpr int b_scale_stride = CDIV(K / 32, 8) * 8;

    constexpr int BLOCK_M = MFMA_DIM_NM * WARP_TILE_M * NUM_WARPS_M;
    constexpr int BLOCK_N = MFMA_DIM_NM * WARP_TILE_N * NUM_WARPS_N;

    // Get XCD coordinate
    uint32_t xcd_id = blockIdx.x / NUM_CUS_PER_XCD;
    uint32_t xcd_worker_id = blockIdx.x % NUM_CUS_PER_XCD;
    uint32_t xcd_n = xcd_id % NUM_XCD_N;
    uint32_t xcd_m = xcd_id / NUM_XCD_N;

    // Get CU coordinate
    if (xcd_worker_id >= NUM_CU_M*NUM_CU_N) {
        return;
    }
    uint32_t cu_n = xcd_worker_id % NUM_CU_N;
    uint32_t cu_m = xcd_worker_id / NUM_CU_N;

    // Get warp id (Use SGPR for it)
    uint32_t warp_id = __builtin_amdgcn_readfirstlane(threadIdx.x / THREADS_PER_WARP);
    uint32_t warp_n = warp_id % NUM_WARPS_N;
    uint32_t warp_m = warp_id / NUM_WARPS_N;

    // Get lane assignment
    uint32_t lane_index = threadIdx.x % THREADS_PER_WARP;
    uint32_t lane_row_nm = lane_index % MFMA_DIM_NM;
    uint32_t lane_col_k = SCALE_GROUP_SIZE * (lane_index / MFMA_DIM_NM); // By chance, a lane sends exactly one scale group.

    // Get tile offset
    uint32_t tile_n_offset = xcd_n * (N / NUM_XCD_N) + cu_n * (N / NUM_XCD_N / NUM_CU_N);
    uint32_t tile_m_offset = xcd_m * (M / NUM_XCD_M) + cu_m * (M / NUM_XCD_M / NUM_CU_M);
    uint32_t warp_n_offset = warp_n * MFMA_DIM_NM * WARP_TILE_N;
    uint32_t warp_m_offset = warp_m * MFMA_DIM_NM * WARP_TILE_M;

    // Shared Memory
    __shared__ uint8_t s_a_scale[BLOCK_M][K / SCALE_GROUP_SIZE];
    __shared__ uint8_t s_a[BLOCK_M][K / 2];
    __shared__ uint8_t s_b_scale[BLOCK_N][K / SCALE_GROUP_SIZE];
    __shared__ uint8_t s_b[BLOCK_N][K / 2];
    auto swizzle_k = [](int r, int c) { return swizzle<SCALE_GROUP_SIZE / 2>(r, c); };

    // ======================
    // A: Global -> Register
    // ======================

    TIME();

    constexpr uint32_t A_CHUNKS_K = K / SCALE_GROUP_SIZE;
    constexpr uint32_t A_CHUNK_JOBS = BLOCK_M * A_CHUNKS_K;
    constexpr uint32_t A_CHUNK_JOBS_PER_THREAD = CDIV(A_CHUNK_JOBS, THREADS_PER_BLOCK);
    u16x32_t a_jobs[A_CHUNK_JOBS_PER_THREAD];
    uint8_t a_scale_jobs[A_CHUNK_JOBS_PER_THREAD];
#pragma unroll
    for (uint32_t a_chunk_job_idx = 0; a_chunk_job_idx < A_CHUNK_JOBS_PER_THREAD; a_chunk_job_idx++) {
        uint32_t a_chunk_job = a_chunk_job_idx * THREADS_PER_BLOCK + threadIdx.x;
        uint32_t a_transfer_k = (a_chunk_job % A_CHUNKS_K) * (K / A_CHUNKS_K);
        uint32_t a_transfer_m = a_chunk_job / A_CHUNKS_K;
        uint32_t a_row_idx = tile_m_offset + a_transfer_m;
        *(u16x32_t*)(&a_jobs[a_chunk_job_idx]) = *(const u16x32_t*)(a + a_row_idx * K + a_transfer_k);
    }
    auto a_process_job_step1 = [&](int a_chunk_job_idx) {
        __builtin_amdgcn_sched_barrier(0);
        uint32_t a_chunk_job = a_chunk_job_idx * THREADS_PER_BLOCK + threadIdx.x;
        uint32_t a_transfer_k = (a_chunk_job % A_CHUNKS_K) * (K / A_CHUNKS_K);
        uint32_t a_transfer_m = a_chunk_job / A_CHUNKS_K;
        uint8_t a_scale = bf16x32_to_scale_e8m0(&a_jobs[a_chunk_job_idx]);
        s_a_scale[a_transfer_m][a_transfer_k / SCALE_GROUP_SIZE] = a_scale;
        a_scale_jobs[a_chunk_job_idx] = a_scale;
        __builtin_amdgcn_sched_barrier(0);
    };
    auto a_process_job_step2 = [&](int a_chunk_job_idx) {
        __builtin_amdgcn_sched_barrier(0);
        uint32_t a_chunk_job = a_chunk_job_idx * THREADS_PER_BLOCK + threadIdx.x;
        uint32_t a_transfer_k = (a_chunk_job % A_CHUNKS_K) * (K / A_CHUNKS_K);
        uint32_t a_transfer_m = a_chunk_job / A_CHUNKS_K;
        float a_scale_f32 = e8m0_to_f32(a_scale_jobs[a_chunk_job_idx]);
        u32x4_t a_reg;
#pragma unroll
        for (int reg_idx = 0; reg_idx < 4; reg_idx++) {
            const bf16x2_t* a_row_v2 = (const bf16x2_t*)(&a_jobs[a_chunk_job_idx]) + 4 * reg_idx;
#ifdef __gfx950__
            a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(0, a_row_v2[0], a_scale_f32, 0);
            a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[1], a_scale_f32, 1);
            a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[2], a_scale_f32, 2);
            a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[3], a_scale_f32, 3);
#else
            a_reg[reg_idx] = __builtin_bit_cast(uint32_t, a_row_v2[0]) ^ __builtin_bit_cast(uint32_t, a_row_v2[1])
                           ^ __builtin_bit_cast(uint32_t, a_row_v2[2]) ^ __builtin_bit_cast(uint32_t, a_row_v2[3])
                           ^ __builtin_bit_cast(uint32_t, a_scale_f32);
#endif
        }
        *(u32x4_t*)&s_a[a_transfer_m][swizzle_k(a_transfer_m, a_transfer_k / 2)] = a_reg;
        __builtin_amdgcn_sched_barrier(0);
    };
    uint32_t a_work_idx = 0;

    // ======================
    // B: Global -> Register
    // ======================

    TIME();

    constexpr uint32_t B_BYTES_PER_JOB = 16;
    constexpr uint32_t B_TOTAL_BYTES = BLOCK_N * (K / 2);
    constexpr uint32_t B_BYTES_PER_JOBGROUP = THREADS_PER_BLOCK * B_BYTES_PER_JOB;
    constexpr uint32_t B_NUM_JOBGROUPS = B_TOTAL_BYTES / B_BYTES_PER_JOBGROUP;
    static_assert(B_TOTAL_BYTES % B_BYTES_PER_JOBGROUP == 0);

    u32x4_t b_jobs[B_NUM_JOBGROUPS];
    uint8_t b_scale_jobs[B_NUM_JOBGROUPS];
    auto b_flat_byte = [&](uint32_t b_jobgroup_idx) -> uint32_t {
        return b_jobgroup_idx * B_BYTES_PER_JOBGROUP + threadIdx.x * B_BYTES_PER_JOB;
    };
    const uint8_t* b_base = b_shuffle + tile_n_offset * (K / 2);
    uint32_t b_thread_offset = threadIdx.x * B_BYTES_PER_JOB;

    auto b_process_job_step0 = [&](int b_jobgroup_idx) {
        __builtin_amdgcn_sched_barrier(0);
        uint32_t b_jobgroup_start_offset = b_jobgroup_idx * B_BYTES_PER_JOBGROUP;
        uint32_t b_job_offset = b_jobgroup_start_offset + b_thread_offset;
        uint32_t b_transfer_n = b_job_offset / (K / 2);
        uint32_t b_transfer_k_byte = b_job_offset % (K / 2);
        uint32_t b_row_idx = tile_n_offset + b_transfer_n;

        b_scale_jobs[b_jobgroup_idx] = b_scale_sh[read_b_scale_base(b_row_idx, b_transfer_k_byte / B_BYTES_PER_JOB, b_scale_stride)];

        uint32_t warp_base = b_jobgroup_start_offset + warp_id * (THREADS_PER_WARP * B_BYTES_PER_JOB);
#ifdef __gfx950__
        void* src = (void*)(b_base + warp_base + lane_index * 16);
        uint32_t lds_off = (uint32_t)(uintptr_t)((uint8_t*)s_b + warp_base);
        asm volatile(
            "s_mov_b32 m0, %0\n\t"
            "global_load_lds_dwordx4 %1, off\n\t"
            :: "s"(lds_off), "v"(src)
            : "memory", "m0"
        );
#else
#pragma unroll
        for (uint32_t sub = 0; sub < 4; sub++) {
            uint32_t chunk_base = warp_base + sub * 256;
            void* src = (void*)(b_base + chunk_base + lane_index * 4);
            uint32_t lds_off = (uint32_t)(uintptr_t)((uint8_t*)s_b + chunk_base);
            asm volatile(
                "s_mov_b32 m0, %0\n\t"
                "global_load_lds_dword %1, off\n\t"
                :: "s"(lds_off), "v"(src)
                : "memory", "m0"
            );
        }
#endif

        __builtin_amdgcn_sched_barrier(0);
    };

    auto b_process_job_step1 = [&](int b_jobgroup_idx) {
        __builtin_amdgcn_sched_barrier(0);
        uint32_t b_jobgroup_start_offset = b_jobgroup_idx * B_BYTES_PER_JOBGROUP;
        uint32_t b_job_offset = b_jobgroup_start_offset + b_thread_offset;
        uint32_t b_transfer_n = b_job_offset / (K / 2);
        uint32_t b_transfer_k_byte = b_job_offset % (K / 2);

        asm volatile("" ::: "memory");
        s_b_scale[b_transfer_n][b_transfer_k_byte / B_BYTES_PER_JOB] = b_scale_jobs[b_jobgroup_idx];
        asm volatile("" ::: "memory");
        __builtin_amdgcn_sched_barrier(0);
    };

#pragma unroll
    for (uint32_t b_jobgroup_idx = 0; b_jobgroup_idx < B_NUM_JOBGROUPS; b_jobgroup_idx++) {
        b_process_job_step0(b_jobgroup_idx);

        __builtin_amdgcn_sched_barrier(0);
        if (b_jobgroup_idx == 1) {
            a_process_job_step1(a_work_idx / 2);
        } else if (b_jobgroup_idx == 2) {
            a_process_job_step2(a_work_idx / 2);
        } else if (b_jobgroup_idx == 3) {
            b_process_job_step1(0);
        }
        __builtin_amdgcn_sched_barrier(0);
    }

    // ======================
    // B: Register -> LDS
    // ======================

    TIME();

#pragma unroll
    for (uint32_t b_jobgroup_idx = 1; b_jobgroup_idx < B_NUM_JOBGROUPS; b_jobgroup_idx++) {
        b_process_job_step1(b_jobgroup_idx);
    }

    // ======================
    // A: Register -> Quantize -> LDS
    // ======================

    TIME();

    // NOTE: This task was interwoven above. This while loop should not be hit.
// #pragma unroll
//     while (a_work_idx < A_CHUNK_JOBS_PER_THREAD * 2) {
//         if (a_work_idx % 2 == 0) {
//             a_process_job_step1(a_work_idx / 2);
//         } else {
//             a_process_job_step2(a_work_idx / 2);
//         }
//         a_work_idx += 1;
//     }

    // ======================
    // LDS->MFMA
    // ======================

    __syncthreads(); // Ensure LDS is populated

    TIME();

    f32x4_t result[WARP_TILE_M][WARP_TILE_N] = {};
    constexpr uint32_t K_ITERS = K / (SPLIT_K * MFMA_DIM_K);
    uint32_t k_start = (SPLIT_K > 1) ? blockIdx.z * (K / SPLIT_K) : 0;
    constexpr uint32_t PREFETCH = 2;

    uint8_t a_scale[PREFETCH][WARP_TILE_M];
    u32x8_t a_reg[PREFETCH][WARP_TILE_M];
    uint8_t b_scale[PREFETCH][WARP_TILE_N];
    u32x8_t b_reg[PREFETCH][WARP_TILE_N];
    auto load_lds = [&](uint32_t k_iter) {
        __builtin_amdgcn_sched_barrier(0);
        uint32_t k_offset = k_start + k_iter * MFMA_DIM_K;
        uint32_t k_idx_start = k_offset + lane_col_k;
        uint32_t buf_idx = k_iter % PREFETCH;

#pragma unroll
        for (uint32_t i = 0; i < WARP_TILE_M; i++) {
            uint32_t a_row = warp_m_offset + i * MFMA_DIM_NM + lane_row_nm;
            a_scale[buf_idx][i] = s_a_scale[a_row][k_idx_start / SCALE_GROUP_SIZE];
            *(u32x4_t*)&a_reg[buf_idx][i] = *(const u32x4_t*)&s_a[a_row][swizzle_k(a_row, k_idx_start / 2)];
        }
#pragma unroll
        for (uint32_t i = 0; i < WARP_TILE_N; i++) {
            uint32_t b_row = warp_n_offset + i * MFMA_DIM_NM + lane_row_nm;
            b_scale[buf_idx][i] = s_b_scale[b_row][k_idx_start / SCALE_GROUP_SIZE];
            uint32_t off = (warp_n_offset + i * MFMA_DIM_NM) * (K / 2)
                        + (k_idx_start / 32) * 256
                        + lane_row_nm * 16;
            *(u32x4_t*)&b_reg[buf_idx][i] = *(const u32x4_t*)((const uint8_t*)s_b + off);
        }
        __builtin_amdgcn_sched_barrier(0);
    };

#pragma unroll
    for (uint32_t i = 0; i < min(K_ITERS, PREFETCH); i++) {
        load_lds(i);
    }
#pragma unroll
    for (uint32_t k_iter = 0; k_iter < K_ITERS; k_iter++) {
        uint32_t k_prefetch_iter = k_iter + PREFETCH - 1;
        if (k_prefetch_iter < K_ITERS) {
            load_lds(k_prefetch_iter);
        }
        uint32_t k_offset = k_start + k_iter * MFMA_DIM_K;
        uint32_t k_idx_start = k_offset + lane_col_k;
        uint32_t buf_idx = k_iter % PREFETCH;

#pragma unroll
        for (uint32_t i = 0; i < WARP_TILE_M; i++) {
#pragma unroll
            for(uint32_t j = 0; j < WARP_TILE_N; j++) {
#ifdef __gfx950__
                result[i][j] = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg[buf_idx][i], b_reg[buf_idx][j], result[i][j], 4, 4, 0, a_scale[buf_idx][i], 0, b_scale[buf_idx][j]);
#else
                float scale = e8m0_to_f32(a_scale[buf_idx][i]) * e8m0_to_f32(b_scale[buf_idx][j]);
                f32x4_t mix;
                for (uint32_t r = 0; r < 4; r++)
                    mix[r] = scale * fp4_e2m1_to_f32(a_reg[buf_idx][i][r]) * fp4_e2m1_to_f32(b_reg[buf_idx][j][r]);
                result[i][j] += mix;
#endif
            }
        }
    }

    // ======================
    // Reg->Write to Global
    // ======================

    TIME();

#pragma unroll
    for (uint32_t i = 0; i < WARP_TILE_M; i++) {
#pragma unroll
        for (uint32_t j = 0; j < WARP_TILE_N; j++) {
            uint32_t result_row_offset = 4 * (lane_index / 16);
            uint32_t result_column_offset = lane_index % 16;
#pragma unroll
            for (uint32_t f = 0; f < 4; f++) {
                uint32_t row = tile_m_offset + warp_m_offset + i*MFMA_DIM_NM + result_row_offset + f;
                uint32_t col = tile_n_offset + warp_n_offset + j*MFMA_DIM_NM + result_column_offset;
                c[row * N + col] = f32_to_bf16((float)result[i][j][f]);
            }
        }
    }

    TIME();

    if constexpr (TIMEIT) {
        __shared__ uint32_t _run_idx;
        if (threadIdx.x == 0) {
            _run_idx = atomicAdd(&g_run_idx, 1) / NUM_CUS;
        }
        __syncthreads();
        uint32_t h = hash(_run_idx);
        if (threadIdx.x == (h % THREADS_PER_BLOCK) && blockIdx.x == (h % NUM_CUS)) {
            printf("%u:CU[%u,%u]:", g_run_idx, threadIdx.x, blockIdx.x);
            for (int i = 1; i < _ti; i++) {
                printf(" %llu", (unsigned long long)(_ts[i] - _ts[0]));
            }
            printf("\n");
        }
    }
}

void entry(
    const uintptr_t a, const uintptr_t b,
    const uintptr_t b_q, const uintptr_t b_shuffle, const uintptr_t b_scale_sh,
    uintptr_t c,
    int M, int N, int K
) {
    if (M == 32 && N == 4096 && K == 512) {
        kernel<<<dim3(NUM_CUS), dim3(THREADS_PER_BLOCK)>>>(
            (const uint16_t*)a,
            (const uint8_t*)b_q, (const uint8_t*)b_shuffle, (const uint8_t*)b_scale_sh,
            (uint16_t*)c
        );
    } else {
        // No impl
    }
}
"""

CUDA_SRC_SHAPE4 = r"""

#define CDIV(a, b) (((a) + (b) - 1) / (b))

using u32x4_t = uint32_t __attribute__((ext_vector_type(4)));
using u16x32_t = uint16_t __attribute__((ext_vector_type(32)));
using bf16x2_t = __bf16 __attribute__((ext_vector_type(2)));
using f32x4_t = float __attribute__((ext_vector_type(4)));
using u32x8_t = uint32_t __attribute__((ext_vector_type(8)));
using i32x4_t = int __attribute__((ext_vector_type(4)));
using i32x8_t = int __attribute__((ext_vector_type(8)));

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

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

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

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

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

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

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

__device__ inline float fp4_e2m1_to_f32(uint8_t nibble) {
    uint8_t s = (nibble >> 3) & 1;
    uint8_t e = (nibble >> 1) & 3;
    uint8_t m = nibble & 1;
    float val = (e == 0) ? m * 0.5f : exp2f((float)e - 1.0f) * (1.0f + m * 0.5f);
    return s ? -val : val;
}

// STRIDE is the width of a row in bytes
// WIDTH is the number of bytes in a single indivisible "column group" (i.e. "greater column").
template <uint32_t WIDTH>
__device__ inline int swizzle(int r, int c) {
    static_assert(WIDTH && !(WIDTH & (WIDTH - 1)), "WIDTH must be power of 2");
    static_assert(WIDTH <= 256, "WIDTH must fit in one bank cycle");
    constexpr int POSITIONS = 256 / WIDTH;  // slots per bank cycle
    constexpr int MASK = POSITIONS - 1;
    int c_group = c / WIDTH;
    int c_offset = c % WIDTH;              // both compile to shift/mask
    int swizzled = (c_group & ~MASK) | ((c_group ^ r) & MASK);
    return swizzled * WIDTH + c_offset;
}

// Constants
constexpr uint32_t MFMA_DIM_NM = 16;
constexpr uint32_t MFMA_DIM_K = 128;
constexpr uint32_t SCALE_GROUP_SIZE = 32;
constexpr uint32_t THREADS_PER_WARP = 64;
constexpr uint32_t NUM_XCD = 8;
#ifdef __gfx950__
constexpr uint32_t NUM_CUS_PER_XCD = 32;

// Parameters
constexpr uint32_t WARP_TILE_M = 1;
constexpr uint32_t WARP_TILE_N = 1;
constexpr uint32_t NUM_WARPS_M = 1;
constexpr uint32_t NUM_WARPS_N = 4;
constexpr uint32_t NUM_CU_M = 2;
constexpr uint32_t NUM_CU_N = 9;
constexpr uint32_t NUM_XCD_M = 1;
constexpr uint32_t NUM_XCD_N = 5;

#else
constexpr uint32_t NUM_CUS_PER_XCD = 38;

// Parameters
constexpr uint32_t WARP_TILE_M = 2;
constexpr uint32_t WARP_TILE_N = 1;
constexpr uint32_t NUM_WARPS_M = 1;
constexpr uint32_t NUM_WARPS_N = 4;
constexpr uint32_t NUM_CU_M = 2;
constexpr uint32_t NUM_CU_N = 14;
constexpr uint32_t NUM_XCD_M = 1;
constexpr uint32_t NUM_XCD_N = 8;
#endif
constexpr uint32_t NUM_CUS = NUM_XCD_M * NUM_XCD_N * NUM_CUS_PER_XCD;

// Derived
constexpr uint32_t THREADS_PER_BLOCK = THREADS_PER_WARP * NUM_WARPS_M * NUM_WARPS_N;

// Magic shuffle
__device__ inline int read_b_scale_base(int n_row, int k_group, int b_scale_stride) {
    int tr = n_row / 32, r = n_row % 32;
    int tc = k_group / 8, c = k_group % 8;
    int base  = tr * (32 * b_scale_stride) + tc * 256;
    int inner = (r / 16) + (c / 4) * 2 + (r % 16) * 4 + (c % 4) * 64;
    return base + inner;
}

// Only used if timing is enabled.
__device__ uint32_t g_run_idx = 0;
__device__ inline uint32_t hash(uint32_t x) {
    x ^= x >> 16;
    x *= 0x45d9f3b;
    x ^= x >> 16;
    x *= 0x45d9f3b;
    x ^= x >> 16;
    return x;
}

__launch_bounds__(THREADS_PER_BLOCK, 1)
__global__ void kernel(
    const uint16_t* a,
    const uint8_t* b_q, const uint8_t* b_shuffle, const uint8_t* b_scale_sh,
    uint16_t* c
) {
    // Timing
    constexpr bool TIMEIT = false;
    uint64_t _ts[16];
    int _ti = 0;
    auto TIME = [&]() {
        if constexpr (TIMEIT) {
            __builtin_amdgcn_sched_barrier(0);
            _ts[_ti++] = __builtin_amdgcn_s_memrealtime();
            __builtin_amdgcn_sched_barrier(0);
        }
    };

    TIME();

    constexpr uint32_t M = 32;
    constexpr uint32_t N = 2880;
    constexpr uint32_t K = 512;
    constexpr uint32_t SPLIT_K = 1;
    static_assert(MFMA_DIM_NM * WARP_TILE_M * NUM_WARPS_M * NUM_CU_M * NUM_XCD_M == M);
    static_assert(MFMA_DIM_NM * WARP_TILE_N * NUM_WARPS_N * NUM_CU_N * NUM_XCD_N == N);
    constexpr int b_scale_stride = CDIV(K / 32, 8) * 8;

    constexpr int BLOCK_M = MFMA_DIM_NM * WARP_TILE_M * NUM_WARPS_M;
    constexpr int BLOCK_N = MFMA_DIM_NM * WARP_TILE_N * NUM_WARPS_N;

    // Get XCD coordinate
    uint32_t xcd_id = blockIdx.x / NUM_CUS_PER_XCD;
    uint32_t xcd_worker_id = blockIdx.x % NUM_CUS_PER_XCD;
    uint32_t xcd_n = xcd_id % NUM_XCD_N;
    uint32_t xcd_m = xcd_id / NUM_XCD_N;

    // Get CU coordinate
    if (xcd_worker_id >= NUM_CU_M*NUM_CU_N) {
        return;
    }
    uint32_t cu_n = xcd_worker_id % NUM_CU_N;
    uint32_t cu_m = xcd_worker_id / NUM_CU_N;

    // Get warp id (Use SGPR for it)
    uint32_t warp_id = __builtin_amdgcn_readfirstlane(threadIdx.x / THREADS_PER_WARP);
    uint32_t warp_n = warp_id % NUM_WARPS_N;
    uint32_t warp_m = warp_id / NUM_WARPS_N;

    // Get lane assignment
    uint32_t lane_index = threadIdx.x % THREADS_PER_WARP;
    uint32_t lane_row_nm = lane_index % MFMA_DIM_NM;
    uint32_t lane_col_k = SCALE_GROUP_SIZE * (lane_index / MFMA_DIM_NM); // By chance, a lane sends exactly one scale group.

    // Get tile offset
    uint32_t tile_n_offset = xcd_n * (N / NUM_XCD_N) + cu_n * (N / NUM_XCD_N / NUM_CU_N);
    uint32_t tile_m_offset = xcd_m * (M / NUM_XCD_M) + cu_m * (M / NUM_XCD_M / NUM_CU_M);
    uint32_t warp_n_offset = warp_n * MFMA_DIM_NM * WARP_TILE_N;
    uint32_t warp_m_offset = warp_m * MFMA_DIM_NM * WARP_TILE_M;

    // Shared Memory
    __shared__ uint8_t s_a_scale[BLOCK_M][K / SCALE_GROUP_SIZE];
    __shared__ uint8_t s_a[BLOCK_M][K / 2];
    __shared__ uint8_t s_b_scale[BLOCK_N][K / SCALE_GROUP_SIZE];
    __shared__ uint8_t s_b[BLOCK_N][K / 2];
    auto swizzle_k = [](int r, int c) { return swizzle<SCALE_GROUP_SIZE / 2>(r, c); };

    // ======================
    // A: Global -> Register
    // ======================

    TIME();

    constexpr uint32_t A_CHUNKS_K = K / SCALE_GROUP_SIZE;
    constexpr uint32_t A_CHUNK_JOBS = BLOCK_M * A_CHUNKS_K;
    constexpr uint32_t A_CHUNK_JOBS_PER_THREAD = CDIV(A_CHUNK_JOBS, THREADS_PER_BLOCK);
    u16x32_t a_jobs[A_CHUNK_JOBS_PER_THREAD];
    uint8_t a_scale_jobs[A_CHUNK_JOBS_PER_THREAD];
#pragma unroll
    for (uint32_t a_chunk_job_idx = 0; a_chunk_job_idx < A_CHUNK_JOBS_PER_THREAD; a_chunk_job_idx++) {
        uint32_t a_chunk_job = a_chunk_job_idx * THREADS_PER_BLOCK + threadIdx.x;
        uint32_t a_transfer_k = (a_chunk_job % A_CHUNKS_K) * (K / A_CHUNKS_K);
        uint32_t a_transfer_m = a_chunk_job / A_CHUNKS_K;
        uint32_t a_row_idx = tile_m_offset + a_transfer_m;
        *(u16x32_t*)(&a_jobs[a_chunk_job_idx]) = *(const u16x32_t*)(a + a_row_idx * K + a_transfer_k);
    }
    auto a_process_job_step1 = [&](int a_chunk_job_idx) {
        __builtin_amdgcn_sched_barrier(0);
        uint32_t a_chunk_job = a_chunk_job_idx * THREADS_PER_BLOCK + threadIdx.x;
        uint32_t a_transfer_k = (a_chunk_job % A_CHUNKS_K) * (K / A_CHUNKS_K);
        uint32_t a_transfer_m = a_chunk_job / A_CHUNKS_K;
        uint8_t a_scale = bf16x32_to_scale_e8m0(&a_jobs[a_chunk_job_idx]);
        s_a_scale[a_transfer_m][a_transfer_k / SCALE_GROUP_SIZE] = a_scale;
        a_scale_jobs[a_chunk_job_idx] = a_scale;
        __builtin_amdgcn_sched_barrier(0);
    };
    auto a_process_job_step2 = [&](int a_chunk_job_idx) {
        __builtin_amdgcn_sched_barrier(0);
        uint32_t a_chunk_job = a_chunk_job_idx * THREADS_PER_BLOCK + threadIdx.x;
        uint32_t a_transfer_k = (a_chunk_job % A_CHUNKS_K) * (K / A_CHUNKS_K);
        uint32_t a_transfer_m = a_chunk_job / A_CHUNKS_K;
        float a_scale_f32 = e8m0_to_f32(a_scale_jobs[a_chunk_job_idx]);
        u32x4_t a_reg;
#pragma unroll
        for (int reg_idx = 0; reg_idx < 4; reg_idx++) {
            const bf16x2_t* a_row_v2 = (const bf16x2_t*)(&a_jobs[a_chunk_job_idx]) + 4 * reg_idx;
#ifdef __gfx950__
            a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(0, a_row_v2[0], a_scale_f32, 0);
            a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[1], a_scale_f32, 1);
            a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[2], a_scale_f32, 2);
            a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[3], a_scale_f32, 3);
#else
            a_reg[reg_idx] = __builtin_bit_cast(uint32_t, a_row_v2[0]) ^ __builtin_bit_cast(uint32_t, a_row_v2[1])
                           ^ __builtin_bit_cast(uint32_t, a_row_v2[2]) ^ __builtin_bit_cast(uint32_t, a_row_v2[3])
                           ^ __builtin_bit_cast(uint32_t, a_scale_f32);
#endif
        }
        *(u32x4_t*)&s_a[a_transfer_m][swizzle_k(a_transfer_m, a_transfer_k / 2)] = a_reg;
        __builtin_amdgcn_sched_barrier(0);
    };
    uint32_t a_work_idx = 0;

    // ======================
    // B: Global -> Register
    // ======================

    TIME();

    constexpr uint32_t B_BYTES_PER_JOB = 16;
    constexpr uint32_t B_TOTAL_BYTES = BLOCK_N * (K / 2);
    constexpr uint32_t B_BYTES_PER_JOBGROUP = THREADS_PER_BLOCK * B_BYTES_PER_JOB;
    constexpr uint32_t B_NUM_JOBGROUPS = B_TOTAL_BYTES / B_BYTES_PER_JOBGROUP;
    static_assert(B_TOTAL_BYTES % B_BYTES_PER_JOBGROUP == 0);

    u32x4_t b_jobs[B_NUM_JOBGROUPS];
    uint8_t b_scale_jobs[B_NUM_JOBGROUPS];
    auto b_flat_byte = [&](uint32_t b_jobgroup_idx) -> uint32_t {
        return b_jobgroup_idx * B_BYTES_PER_JOBGROUP + threadIdx.x * B_BYTES_PER_JOB;
    };
    const uint8_t* b_base = b_shuffle + tile_n_offset * (K / 2);
    uint32_t b_thread_offset = threadIdx.x * B_BYTES_PER_JOB;

    auto b_process_job_step0 = [&](int b_jobgroup_idx) {
        __builtin_amdgcn_sched_barrier(0);
        uint32_t b_jobgroup_start_offset = b_jobgroup_idx * B_BYTES_PER_JOBGROUP;
        uint32_t b_job_offset = b_jobgroup_start_offset + b_thread_offset;
        uint32_t b_transfer_n = b_job_offset / (K / 2);
        uint32_t b_transfer_k_byte = b_job_offset % (K / 2);
        uint32_t b_row_idx = tile_n_offset + b_transfer_n;

        b_scale_jobs[b_jobgroup_idx] = b_scale_sh[read_b_scale_base(b_row_idx, b_transfer_k_byte / B_BYTES_PER_JOB, b_scale_stride)];

        uint32_t warp_base = b_jobgroup_start_offset + warp_id * (THREADS_PER_WARP * B_BYTES_PER_JOB);
#ifdef __gfx950__
        void* src = (void*)(b_base + warp_base + lane_index * 16);
        uint32_t lds_off = (uint32_t)(uintptr_t)((uint8_t*)s_b + warp_base);
        asm volatile(
            "s_mov_b32 m0, %0\n\t"
            "global_load_lds_dwordx4 %1, off\n\t"
            :: "s"(lds_off), "v"(src)
            : "memory", "m0"
        );
#else
#pragma unroll
        for (uint32_t sub = 0; sub < 4; sub++) {
            uint32_t chunk_base = warp_base + sub * 256;
            void* src = (void*)(b_base + chunk_base + lane_index * 4);
            uint32_t lds_off = (uint32_t)(uintptr_t)((uint8_t*)s_b + chunk_base);
            asm volatile(
                "s_mov_b32 m0, %0\n\t"
                "global_load_lds_dword %1, off\n\t"
                :: "s"(lds_off), "v"(src)
                : "memory", "m0"
            );
        }
#endif
        
        __builtin_amdgcn_sched_barrier(0);
    };

    auto b_process_job_step1 = [&](int b_jobgroup_idx) {
        __builtin_amdgcn_sched_barrier(0);
        uint32_t b_jobgroup_start_offset = b_jobgroup_idx * B_BYTES_PER_JOBGROUP;
        uint32_t b_job_offset = b_jobgroup_start_offset + b_thread_offset;
        uint32_t b_transfer_n = b_job_offset / (K / 2);
        uint32_t b_transfer_k_byte = b_job_offset % (K / 2);

        asm volatile("" ::: "memory");
        s_b_scale[b_transfer_n][b_transfer_k_byte / B_BYTES_PER_JOB] = b_scale_jobs[b_jobgroup_idx];
        asm volatile("" ::: "memory");
        __builtin_amdgcn_sched_barrier(0);
    };

#pragma unroll
    for (uint32_t b_jobgroup_idx = 0; b_jobgroup_idx < B_NUM_JOBGROUPS; b_jobgroup_idx++) {
        b_process_job_step0(b_jobgroup_idx);

        __builtin_amdgcn_sched_barrier(0);
        if (b_jobgroup_idx == 1) {
            a_process_job_step1(a_work_idx / 2);
        } else if (b_jobgroup_idx == 2) {
            a_process_job_step2(a_work_idx / 2);
        } else if (b_jobgroup_idx == 3) {
            b_process_job_step1(0);
        }
        __builtin_amdgcn_sched_barrier(0);
    }

    // ======================
    // B: Register -> LDS
    // ======================

    TIME();

#pragma unroll
    for (uint32_t b_jobgroup_idx = 1; b_jobgroup_idx < B_NUM_JOBGROUPS; b_jobgroup_idx++) {
        b_process_job_step1(b_jobgroup_idx);
    }

    // ======================
    // A: Register -> Quantize -> LDS
    // ======================

    TIME();

    // NOTE: This task was interwoven above. This while loop should not be hit.
// #pragma unroll
//     while (a_work_idx < A_CHUNK_JOBS_PER_THREAD * 2) {
//         if (a_work_idx % 2 == 0) {
//             a_process_job_step1(a_work_idx / 2);
//         } else {
//             a_process_job_step2(a_work_idx / 2);
//         }
//         a_work_idx += 1;
//     }

    // ======================
    // LDS->MFMA
    // ======================

    __syncthreads(); // Ensure LDS is populated

    TIME();

    f32x4_t result[WARP_TILE_M][WARP_TILE_N] = {};
    constexpr uint32_t K_ITERS = K / (SPLIT_K * MFMA_DIM_K);
    uint32_t k_start = (SPLIT_K > 1) ? blockIdx.z * (K / SPLIT_K) : 0;
    constexpr uint32_t PREFETCH = 2;

    uint8_t a_scale[PREFETCH][WARP_TILE_M];
    u32x8_t a_reg[PREFETCH][WARP_TILE_M];
    uint8_t b_scale[PREFETCH][WARP_TILE_N];
    u32x8_t b_reg[PREFETCH][WARP_TILE_N];
    auto load_lds = [&](uint32_t k_iter) {
        __builtin_amdgcn_sched_barrier(0);
        uint32_t k_offset = k_start + k_iter * MFMA_DIM_K;
        uint32_t k_idx_start = k_offset + lane_col_k;
        uint32_t buf_idx = k_iter % PREFETCH;

#pragma unroll
        for (uint32_t i = 0; i < WARP_TILE_M; i++) {
            uint32_t a_row = warp_m_offset + i * MFMA_DIM_NM + lane_row_nm;
            a_scale[buf_idx][i] = s_a_scale[a_row][k_idx_start / SCALE_GROUP_SIZE];
            *(u32x4_t*)&a_reg[buf_idx][i] = *(const u32x4_t*)&s_a[a_row][swizzle_k(a_row, k_idx_start / 2)];
        }
#pragma unroll
        for (uint32_t i = 0; i < WARP_TILE_N; i++) {
            uint32_t b_row = warp_n_offset + i * MFMA_DIM_NM + lane_row_nm;
            b_scale[buf_idx][i] = s_b_scale[b_row][k_idx_start / SCALE_GROUP_SIZE];
            uint32_t off = (warp_n_offset + i * MFMA_DIM_NM) * (K / 2)
                        + (k_idx_start / 32) * 256
                        + lane_row_nm * 16;
            *(u32x4_t*)&b_reg[buf_idx][i] = *(const u32x4_t*)((const uint8_t*)s_b + off);
        }
        __builtin_amdgcn_sched_barrier(0);
    };

#pragma unroll
    for (uint32_t i = 0; i < min(K_ITERS, PREFETCH); i++) {
        load_lds(i);
    }
#pragma unroll
    for (uint32_t k_iter = 0; k_iter < K_ITERS; k_iter++) {
        uint32_t k_prefetch_iter = k_iter + PREFETCH - 1;
        if (k_prefetch_iter < K_ITERS) {
            load_lds(k_prefetch_iter);
        }
        uint32_t k_offset = k_start + k_iter * MFMA_DIM_K;
        uint32_t k_idx_start = k_offset + lane_col_k;
        uint32_t buf_idx = k_iter % PREFETCH;

#pragma unroll
        for (uint32_t i = 0; i < WARP_TILE_M; i++) {
#pragma unroll
            for(uint32_t j = 0; j < WARP_TILE_N; j++) {
#ifdef __gfx950__
                result[i][j] = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg[buf_idx][i], b_reg[buf_idx][j], result[i][j], 4, 4, 0, a_scale[buf_idx][i], 0, b_scale[buf_idx][j]);
#else
                float scale = e8m0_to_f32(a_scale[buf_idx][i]) * e8m0_to_f32(b_scale[buf_idx][j]);
                f32x4_t mix;
                for (uint32_t r = 0; r < 4; r++)
                    mix[r] = scale * fp4_e2m1_to_f32(a_reg[buf_idx][i][r]) * fp4_e2m1_to_f32(b_reg[buf_idx][j][r]);
                result[i][j] += mix;
#endif
            }
        }
    }

    // ======================
    // Reg->Write to Global
    // ======================

    TIME();

#pragma unroll
    for (uint32_t i = 0; i < WARP_TILE_M; i++) {
#pragma unroll
        for (uint32_t j = 0; j < WARP_TILE_N; j++) {
            uint32_t result_row_offset = 4 * (lane_index / 16);
            uint32_t result_column_offset = lane_index % 16;
#pragma unroll
            for (uint32_t f = 0; f < 4; f++) {
                uint32_t row = tile_m_offset + warp_m_offset + i*MFMA_DIM_NM + result_row_offset + f;
                uint32_t col = tile_n_offset + warp_n_offset + j*MFMA_DIM_NM + result_column_offset;
                c[row * N + col] = f32_to_bf16((float)result[i][j][f]);
            }
        }
    }

    TIME();

    if constexpr (TIMEIT) {
        __shared__ uint32_t _run_idx;
        if (threadIdx.x == 0) {
            _run_idx = atomicAdd(&g_run_idx, 1) / NUM_CUS;
        }
        __syncthreads();
        uint32_t h = hash(_run_idx);
        if (threadIdx.x == (h % THREADS_PER_BLOCK) && blockIdx.x == (h % NUM_CUS)) {
            printf("%u:CU[%u,%u]:", g_run_idx, threadIdx.x, blockIdx.x);
            for (int i = 1; i < _ti; i++) {
                printf(" %llu", (unsigned long long)(_ts[i] - _ts[0]));
            }
            printf("\n");
        }
    }
}

void entry(
    const uintptr_t a, const uintptr_t b,
    const uintptr_t b_q, const uintptr_t b_shuffle, const uintptr_t b_scale_sh,
    uintptr_t c,
    int M, int N, int K
) {
    if (M == 32 && N == 2880 && K == 512) {
        kernel<<<dim3(NUM_CUS), dim3(THREADS_PER_BLOCK)>>>(
            (const uint16_t*)a,
            (const uint8_t*)b_q, (const uint8_t*)b_shuffle, (const uint8_t*)b_scale_sh,
            (uint16_t*)c
        );
    } else {
        // No impl
    }
}
"""

CUDA_SRC_SHAPE5 = r"""

#define CDIV(a, b) (((a) + (b) - 1) / (b))

using u32x4_t = uint32_t __attribute__((ext_vector_type(4)));
using u16x32_t = uint16_t __attribute__((ext_vector_type(32)));
using bf16x2_t = __bf16 __attribute__((ext_vector_type(2)));
using f32x4_t = float __attribute__((ext_vector_type(4)));
using u32x8_t = uint32_t __attribute__((ext_vector_type(8)));
using i32x4_t = int __attribute__((ext_vector_type(4)));
using i32x8_t = int __attribute__((ext_vector_type(8)));

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

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

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

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

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

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

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

__device__ inline float fp4_e2m1_to_f32(uint8_t nibble) {
    uint8_t s = (nibble >> 3) & 1;
    uint8_t e = (nibble >> 1) & 3;
    uint8_t m = nibble & 1;
    float val = (e == 0) ? m * 0.5f : exp2f((float)e - 1.0f) * (1.0f + m * 0.5f);
    return s ? -val : val;
}

// STRIDE is the width of a row in bytes
// WIDTH is the number of bytes in a single indivisible "column group" (i.e. "greater column").
template <uint32_t WIDTH>
__device__ inline int swizzle(int r, int c) {
    static_assert(WIDTH && !(WIDTH & (WIDTH - 1)), "WIDTH must be power of 2");
    static_assert(WIDTH <= 256, "WIDTH must fit in one bank cycle");
    constexpr int POSITIONS = 256 / WIDTH;  // slots per bank cycle
    constexpr int MASK = POSITIONS - 1;
    int c_group = c / WIDTH;
    int c_offset = c % WIDTH;              // both compile to shift/mask
    int swizzled = (c_group & ~MASK) | ((c_group ^ r) & MASK);
    return swizzled * WIDTH + c_offset;
}

// Constants
constexpr uint32_t MFMA_DIM_NM = 16;
constexpr uint32_t MFMA_DIM_K = 128;
constexpr uint32_t SCALE_GROUP_SIZE = 32;
constexpr uint32_t THREADS_PER_WARP = 64;
constexpr uint32_t NUM_XCD = 8;
#ifdef __gfx950__
constexpr uint32_t NUM_CUS_PER_XCD = 32;

// Parameters
constexpr uint32_t WARP_TILE_M = 1;
constexpr uint32_t WARP_TILE_N = 2;
constexpr uint32_t NUM_WARPS_M = 2;
constexpr uint32_t NUM_WARPS_N = 2;
constexpr uint32_t NUM_CU_M = 2;
constexpr uint32_t NUM_CU_N = 14;
constexpr uint32_t NUM_XCD_M = 1;
constexpr uint32_t NUM_XCD_N = 8;

#else
constexpr uint32_t NUM_CUS_PER_XCD = 38;

// Parameters
constexpr uint32_t WARP_TILE_M = 2;
constexpr uint32_t WARP_TILE_N = 1;
constexpr uint32_t NUM_WARPS_M = 1;
constexpr uint32_t NUM_WARPS_N = 4;
constexpr uint32_t NUM_CU_M = 2;
constexpr uint32_t NUM_CU_N = 14;
constexpr uint32_t NUM_XCD_M = 1;
constexpr uint32_t NUM_XCD_N = 8;
#endif
constexpr uint32_t NUM_CUS = NUM_XCD * NUM_CUS_PER_XCD;

// Derived
constexpr uint32_t THREADS_PER_BLOCK = THREADS_PER_WARP * NUM_WARPS_M * NUM_WARPS_N;

// Magic shuffle
__device__ inline int read_b_scale_base(int n_row, int k_group, int b_scale_stride) {
    int tr = n_row / 32, r = n_row % 32;
    int tc = k_group / 8, c = k_group % 8;
    int base  = tr * (32 * b_scale_stride) + tc * 256;
    int inner = (r / 16) + (c / 4) * 2 + (r % 16) * 4 + (c % 4) * 64;
    return base + inner;
}

// Only used if timing is enabled.
__device__ uint32_t g_run_idx = 0;
__device__ inline uint32_t hash(uint32_t x) {
    x ^= x >> 16;
    x *= 0x45d9f3b;
    x ^= x >> 16;
    x *= 0x45d9f3b;
    x ^= x >> 16;
    return x;
}

__launch_bounds__(THREADS_PER_BLOCK, 1)
__global__ void kernel(
    const uint16_t* a,
    const uint8_t* b_q, const uint8_t* b_shuffle, const uint8_t* b_scale_sh,
    uint16_t* c
) {
    // Timing
    constexpr bool TIMEIT = false;
    uint64_t _ts[16];
    int _ti = 0;
    auto TIME = [&]() {
        if constexpr (TIMEIT) {
            __builtin_amdgcn_sched_barrier(0);
            _ts[_ti++] = __builtin_amdgcn_s_memrealtime();
            __builtin_amdgcn_sched_barrier(0);
        }
    };

    TIME();

    constexpr uint32_t M = 64;
    constexpr uint32_t N = 7168;
    constexpr uint32_t K = 2048;
    constexpr uint32_t SPLIT_K = 1;
    static_assert(MFMA_DIM_NM * WARP_TILE_M * NUM_WARPS_M * NUM_CU_M * NUM_XCD_M == M);
    static_assert(MFMA_DIM_NM * WARP_TILE_N * NUM_WARPS_N * NUM_CU_N * NUM_XCD_N == N);
    constexpr int b_scale_stride = CDIV(K / 32, 8) * 8;

    constexpr int BLOCK_M = MFMA_DIM_NM * WARP_TILE_M * NUM_WARPS_M;
    constexpr int BLOCK_N = MFMA_DIM_NM * WARP_TILE_N * NUM_WARPS_N;

    // Get XCD coordinate
    uint32_t xcd_id = blockIdx.x / NUM_CUS_PER_XCD;
    uint32_t xcd_worker_id = blockIdx.x % NUM_CUS_PER_XCD;
    uint32_t xcd_n = xcd_id % NUM_XCD_N;
    uint32_t xcd_m = xcd_id / NUM_XCD_N;

    // Get CU coordinate
    if (xcd_worker_id >= NUM_CU_M*NUM_CU_N) {
        return;
    }
    uint32_t cu_n = xcd_worker_id % NUM_CU_N;
    uint32_t cu_m = xcd_worker_id / NUM_CU_N;

    // Get warp id (Use SGPR for it)
    uint32_t warp_id = __builtin_amdgcn_readfirstlane(threadIdx.x / THREADS_PER_WARP);
    uint32_t warp_n = warp_id % NUM_WARPS_N;
    uint32_t warp_m = warp_id / NUM_WARPS_N;

    // Get lane assignment
    uint32_t lane_index = threadIdx.x % THREADS_PER_WARP;
    uint32_t lane_row_nm = lane_index % MFMA_DIM_NM;
    uint32_t lane_col_k = SCALE_GROUP_SIZE * (lane_index / MFMA_DIM_NM); // By chance, a lane sends exactly one scale group.

    // Get tile offset
    uint32_t tile_n_offset = xcd_n * (N / NUM_XCD_N) + cu_n * (N / NUM_XCD_N / NUM_CU_N);
    uint32_t tile_m_offset = xcd_m * (M / NUM_XCD_M) + cu_m * (M / NUM_XCD_M / NUM_CU_M);
    uint32_t warp_n_offset = warp_n * MFMA_DIM_NM * WARP_TILE_N;
    uint32_t warp_m_offset = warp_m * MFMA_DIM_NM * WARP_TILE_M;

    // Shared Memory
    __shared__ uint8_t s_a_scale[BLOCK_M][K / SCALE_GROUP_SIZE];
    __shared__ uint8_t s_a[BLOCK_M][K / 2];
    __shared__ uint8_t s_b_scale[BLOCK_N][K / SCALE_GROUP_SIZE];
    __shared__ uint8_t s_b[BLOCK_N][K / 2];
    auto swizzle_k = [](int r, int c) { return swizzle<SCALE_GROUP_SIZE / 2>(r, c); };

    // ======================
    // A: Global -> Register
    // ======================

    TIME();

    constexpr uint32_t A_CHUNKS_K = K / SCALE_GROUP_SIZE;
    constexpr uint32_t A_CHUNK_JOBS = BLOCK_M * A_CHUNKS_K;
    constexpr uint32_t A_CHUNK_JOBS_PER_THREAD = CDIV(A_CHUNK_JOBS, THREADS_PER_BLOCK);
    u16x32_t a_jobs[A_CHUNK_JOBS_PER_THREAD];
    uint8_t a_scale_jobs[A_CHUNK_JOBS_PER_THREAD];
#pragma unroll
    for (uint32_t a_chunk_job_idx = 0; a_chunk_job_idx < A_CHUNK_JOBS_PER_THREAD; a_chunk_job_idx++) {
        uint32_t a_chunk_job = a_chunk_job_idx * THREADS_PER_BLOCK + threadIdx.x;
        uint32_t a_transfer_k = (a_chunk_job % A_CHUNKS_K) * (K / A_CHUNKS_K);
        uint32_t a_transfer_m = a_chunk_job / A_CHUNKS_K;
        uint32_t a_row_idx = tile_m_offset + a_transfer_m;
        *(u16x32_t*)(&a_jobs[a_chunk_job_idx]) = *(const u16x32_t*)(a + a_row_idx * K + a_transfer_k);
    }
    auto a_process_job_step1 = [&](int a_chunk_job_idx) {
        __builtin_amdgcn_sched_barrier(0);
        uint32_t a_chunk_job = a_chunk_job_idx * THREADS_PER_BLOCK + threadIdx.x;
        uint32_t a_transfer_k = (a_chunk_job % A_CHUNKS_K) * (K / A_CHUNKS_K);
        uint32_t a_transfer_m = a_chunk_job / A_CHUNKS_K;
        uint8_t a_scale = bf16x32_to_scale_e8m0(&a_jobs[a_chunk_job_idx]);
        s_a_scale[a_transfer_m][a_transfer_k / SCALE_GROUP_SIZE] = a_scale;
        a_scale_jobs[a_chunk_job_idx] = a_scale;
        __builtin_amdgcn_sched_barrier(0);
    };
    auto a_process_job_step2 = [&](int a_chunk_job_idx) {
        __builtin_amdgcn_sched_barrier(0);
        uint32_t a_chunk_job = a_chunk_job_idx * THREADS_PER_BLOCK + threadIdx.x;
        uint32_t a_transfer_k = (a_chunk_job % A_CHUNKS_K) * (K / A_CHUNKS_K);
        uint32_t a_transfer_m = a_chunk_job / A_CHUNKS_K;
        float a_scale_f32 = e8m0_to_f32(a_scale_jobs[a_chunk_job_idx]);
        u32x4_t a_reg;
#pragma unroll
        for (int reg_idx = 0; reg_idx < 4; reg_idx++) {
            const bf16x2_t* a_row_v2 = (const bf16x2_t*)(&a_jobs[a_chunk_job_idx]) + 4 * reg_idx;
#ifdef __gfx950__
            a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(0, a_row_v2[0], a_scale_f32, 0);
            a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[1], a_scale_f32, 1);
            a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[2], a_scale_f32, 2);
            a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[3], a_scale_f32, 3);
#else
            a_reg[reg_idx] = __builtin_bit_cast(uint32_t, a_row_v2[0]) ^ __builtin_bit_cast(uint32_t, a_row_v2[1])
                           ^ __builtin_bit_cast(uint32_t, a_row_v2[2]) ^ __builtin_bit_cast(uint32_t, a_row_v2[3])
                           ^ __builtin_bit_cast(uint32_t, a_scale_f32);
#endif
        }
        *(u32x4_t*)&s_a[a_transfer_m][swizzle_k(a_transfer_m, a_transfer_k / 2)] = a_reg;
        __builtin_amdgcn_sched_barrier(0);
    };
    uint32_t a_work_idx = 0;

    // ======================
    // B: Global -> Register
    // ======================

    TIME();

    constexpr uint32_t B_BYTES_PER_JOB = 16;
    constexpr uint32_t B_TOTAL_BYTES = BLOCK_N * (K / 2);
    constexpr uint32_t B_BYTES_PER_JOBGROUP = THREADS_PER_BLOCK * B_BYTES_PER_JOB;
    constexpr uint32_t B_NUM_JOBGROUPS = B_TOTAL_BYTES / B_BYTES_PER_JOBGROUP;
    static_assert(B_TOTAL_BYTES % B_BYTES_PER_JOBGROUP == 0);

    u32x4_t b_jobs[B_NUM_JOBGROUPS];
    uint8_t b_scale_jobs[B_NUM_JOBGROUPS];
    auto b_flat_byte = [&](uint32_t b_jobgroup_idx) -> uint32_t {
        return b_jobgroup_idx * B_BYTES_PER_JOBGROUP + threadIdx.x * B_BYTES_PER_JOB;
    };
    const uint8_t* b_base = b_shuffle + tile_n_offset * (K / 2);
    uint32_t b_thread_offset = threadIdx.x * B_BYTES_PER_JOB;

    auto b_process_job_step0 = [&](int b_jobgroup_idx) {
        __builtin_amdgcn_sched_barrier(0);
        uint32_t b_jobgroup_start_offset = b_jobgroup_idx * B_BYTES_PER_JOBGROUP;
        uint32_t b_job_offset = b_jobgroup_start_offset + b_thread_offset;
        uint32_t b_transfer_n = b_job_offset / (K / 2);
        uint32_t b_transfer_k_byte = b_job_offset % (K / 2);
        uint32_t b_row_idx = tile_n_offset + b_transfer_n;

        b_scale_jobs[b_jobgroup_idx] = b_scale_sh[read_b_scale_base(b_row_idx, b_transfer_k_byte / B_BYTES_PER_JOB, b_scale_stride)];

        uint32_t warp_base = b_jobgroup_start_offset + warp_id * (THREADS_PER_WARP * B_BYTES_PER_JOB);
#ifdef __gfx950__
        void* src = (void*)(b_base + warp_base + lane_index * 16);
        uint32_t lds_off = (uint32_t)(uintptr_t)((uint8_t*)s_b + warp_base);
        asm volatile(
            "s_mov_b32 m0, %0\n\t"
            "global_load_lds_dwordx4 %1, off\n\t"
            :: "s"(lds_off), "v"(src)
            : "memory", "m0"
        );
#else
#pragma unroll
        for (uint32_t sub = 0; sub < 4; sub++) {
            uint32_t chunk_base = warp_base + sub * 256;
            void* src = (void*)(b_base + chunk_base + lane_index * 4);
            uint32_t lds_off = (uint32_t)(uintptr_t)((uint8_t*)s_b + chunk_base);
            asm volatile(
                "s_mov_b32 m0, %0\n\t"
                "global_load_lds_dword %1, off\n\t"
                :: "s"(lds_off), "v"(src)
                : "memory", "m0"
            );
        }
#endif
        
        __builtin_amdgcn_sched_barrier(0);
    };

    auto b_process_job_step1 = [&](int b_jobgroup_idx) {
        __builtin_amdgcn_sched_barrier(0);
        uint32_t b_jobgroup_start_offset = b_jobgroup_idx * B_BYTES_PER_JOBGROUP;
        uint32_t b_job_offset = b_jobgroup_start_offset + b_thread_offset;
        uint32_t b_transfer_n = b_job_offset / (K / 2);
        uint32_t b_transfer_k_byte = b_job_offset % (K / 2);

        asm volatile("" ::: "memory");
        s_b_scale[b_transfer_n][b_transfer_k_byte / B_BYTES_PER_JOB] = b_scale_jobs[b_jobgroup_idx];
        asm volatile("" ::: "memory");
        __builtin_amdgcn_sched_barrier(0);
    };

#pragma unroll
    for (uint32_t b_jobgroup_idx = 0; b_jobgroup_idx < B_NUM_JOBGROUPS; b_jobgroup_idx++) {
        b_process_job_step0(b_jobgroup_idx);

        __builtin_amdgcn_sched_barrier(0);
        if (a_work_idx % 2 == 0) {
            a_process_job_step1(a_work_idx / 2);
        } else {
            a_process_job_step2(a_work_idx / 2);
        }
        a_work_idx += 1;
        __builtin_amdgcn_sched_barrier(0);
    }

    // ======================
    // B: Register -> LDS
    // ======================

    TIME();

#pragma unroll
    for (uint32_t b_jobgroup_idx = 0; b_jobgroup_idx < B_NUM_JOBGROUPS; b_jobgroup_idx++) {
        b_process_job_step1(b_jobgroup_idx);
    }

    // ======================
    // A: Register -> Quantize -> LDS
    // ======================

    TIME();

    // NOTE: This task was interwoven above. This while loop should not be hit.
// #pragma unroll
//     while (a_work_idx < A_CHUNK_JOBS_PER_THREAD * 2) {
//         if (a_work_idx % 2 == 0) {
//             a_process_job_step1(a_work_idx / 2);
//         } else {
//             a_process_job_step2(a_work_idx / 2);
//         }
//         a_work_idx += 1;
//     }

    // ======================
    // LDS->MFMA
    // ======================

    __syncthreads(); // Ensure LDS is populated

    TIME();

    f32x4_t result[WARP_TILE_M][WARP_TILE_N] = {};
    constexpr uint32_t K_ITERS = K / (SPLIT_K * MFMA_DIM_K);
    uint32_t k_start = (SPLIT_K > 1) ? blockIdx.z * (K / SPLIT_K) : 0;
    constexpr uint32_t PREFETCH = 2;

    uint8_t a_scale[PREFETCH][WARP_TILE_M];
    u32x8_t a_reg[PREFETCH][WARP_TILE_M];
    uint8_t b_scale[PREFETCH][WARP_TILE_N];
    u32x8_t b_reg[PREFETCH][WARP_TILE_N];
    auto load_lds = [&](uint32_t k_iter) {
        __builtin_amdgcn_sched_barrier(0);
        uint32_t k_offset = k_start + k_iter * MFMA_DIM_K;
        uint32_t k_idx_start = k_offset + lane_col_k;
        uint32_t buf_idx = k_iter % PREFETCH;

#pragma unroll
        for (uint32_t i = 0; i < WARP_TILE_M; i++) {
            uint32_t a_row = warp_m_offset + i * MFMA_DIM_NM + lane_row_nm;
            a_scale[buf_idx][i] = s_a_scale[a_row][k_idx_start / SCALE_GROUP_SIZE];
            *(u32x4_t*)&a_reg[buf_idx][i] = *(const u32x4_t*)&s_a[a_row][swizzle_k(a_row, k_idx_start / 2)];
        }
#pragma unroll
        for (uint32_t i = 0; i < WARP_TILE_N; i++) {
            uint32_t b_row = warp_n_offset + i * MFMA_DIM_NM + lane_row_nm;
            b_scale[buf_idx][i] = s_b_scale[b_row][k_idx_start / SCALE_GROUP_SIZE];
            uint32_t off = (warp_n_offset + i * MFMA_DIM_NM) * (K / 2)
                        + (k_idx_start / 32) * 256
                        + lane_row_nm * 16;
            *(u32x4_t*)&b_reg[buf_idx][i] = *(const u32x4_t*)((const uint8_t*)s_b + off);
        }
        __builtin_amdgcn_sched_barrier(0);
    };

#pragma unroll
    for (uint32_t i = 0; i < min(K_ITERS, PREFETCH); i++) {
        load_lds(i);
    }
#pragma unroll
    for (uint32_t k_iter = 0; k_iter < K_ITERS; k_iter++) {
        uint32_t k_prefetch_iter = k_iter + PREFETCH - 1;
        if (k_prefetch_iter < K_ITERS) {
            load_lds(k_prefetch_iter);
        }
        uint32_t k_offset = k_start + k_iter * MFMA_DIM_K;
        uint32_t k_idx_start = k_offset + lane_col_k;
        uint32_t buf_idx = k_iter % PREFETCH;

#pragma unroll
        for (uint32_t i = 0; i < WARP_TILE_M; i++) {
#pragma unroll
            for(uint32_t j = 0; j < WARP_TILE_N; j++) {
#ifdef __gfx950__
                result[i][j] = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg[buf_idx][i], b_reg[buf_idx][j], result[i][j], 4, 4, 0, a_scale[buf_idx][i], 0, b_scale[buf_idx][j]);
#else
                float scale = e8m0_to_f32(a_scale[buf_idx][i]) * e8m0_to_f32(b_scale[buf_idx][j]);
                f32x4_t mix;
                for (uint32_t r = 0; r < 4; r++)
                    mix[r] = scale * fp4_e2m1_to_f32(a_reg[buf_idx][i][r]) * fp4_e2m1_to_f32(b_reg[buf_idx][j][r]);
                result[i][j] += mix;
#endif
            }
        }
    }

    // ======================
    // Reg->Write to Global
    // ======================

    TIME();

#pragma unroll
    for (uint32_t i = 0; i < WARP_TILE_M; i++) {
#pragma unroll
        for (uint32_t j = 0; j < WARP_TILE_N; j++) {
            uint32_t result_row_offset = 4 * (lane_index / 16);
            uint32_t result_column_offset = lane_index % 16;
#pragma unroll
            for (uint32_t f = 0; f < 4; f++) {
                uint32_t row = tile_m_offset + warp_m_offset + i*MFMA_DIM_NM + result_row_offset + f;
                uint32_t col = tile_n_offset + warp_n_offset + j*MFMA_DIM_NM + result_column_offset;
                c[row * N + col] = f32_to_bf16((float)result[i][j][f]);
            }
        }
    }

    TIME();

    if constexpr (TIMEIT) {
        __shared__ uint32_t _run_idx;
        if (threadIdx.x == 0) {
            _run_idx = atomicAdd(&g_run_idx, 1) / NUM_CUS;
        }
        __syncthreads();
        uint32_t h = hash(_run_idx);
        if (threadIdx.x == (h % THREADS_PER_BLOCK) && blockIdx.x == (h % NUM_CUS)) {
            printf("%u:CU[%u,%u]:", g_run_idx, threadIdx.x, blockIdx.x);
            for (int i = 1; i < _ti; i++) {
                printf(" %llu", (unsigned long long)(_ts[i] - _ts[0]));
            }
            printf("\n");
        }
    }
}

void entry(
    const uintptr_t a, const uintptr_t b,
    const uintptr_t b_q, const uintptr_t b_shuffle, const uintptr_t b_scale_sh,
    uintptr_t c,
    int M, int N, int K
) {
    if (M == 64 && N == 7168 && K == 2048) {
        kernel<<<dim3(NUM_CUS), dim3(THREADS_PER_BLOCK)>>>(
            (const uint16_t*)a,
            (const uint8_t*)b_q, (const uint8_t*)b_shuffle, (const uint8_t*)b_scale_sh,
            (uint16_t*)c
        );
    } else {
        // No impl
    }
}
"""

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

using u32x4_t = uint32_t __attribute__((ext_vector_type(4)));
using u16x32_t = uint16_t __attribute__((ext_vector_type(32)));
using bf16x2_t = __bf16 __attribute__((ext_vector_type(2)));
using f32x4_t = float __attribute__((ext_vector_type(4)));
using u32x8_t = uint32_t __attribute__((ext_vector_type(8)));
using i32x4_t = int __attribute__((ext_vector_type(4)));
using i32x8_t = int __attribute__((ext_vector_type(8)));

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

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

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

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

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

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

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

__device__ inline float fp4_e2m1_to_f32(uint8_t nibble) {
    uint8_t s = (nibble >> 3) & 1;
    uint8_t e = (nibble >> 1) & 3;
    uint8_t m = nibble & 1;
    float val = (e == 0) ? m * 0.5f : exp2f((float)e - 1.0f) * (1.0f + m * 0.5f);
    return s ? -val : val;
}

// STRIDE is the width of a row in bytes
// WIDTH is the number of bytes in a single indivisible "column group" (i.e. "greater column").
template <uint32_t WIDTH>
__device__ inline int swizzle(int r, int c) {
    static_assert(WIDTH && !(WIDTH & (WIDTH - 1)), "WIDTH must be power of 2");
    static_assert(WIDTH <= 256, "WIDTH must fit in one bank cycle");
    constexpr int POSITIONS = 256 / WIDTH;  // slots per bank cycle
    constexpr int MASK = POSITIONS - 1;
    int c_group = c / WIDTH;
    int c_offset = c % WIDTH;              // both compile to shift/mask
    int swizzled = (c_group & ~MASK) | ((c_group ^ r) & MASK);
    return swizzled * WIDTH + c_offset;
}

// Constants
constexpr uint32_t MFMA_DIM_NM = 16;
constexpr uint32_t MFMA_DIM_K = 128;
constexpr uint32_t SCALE_GROUP_SIZE = 32;
constexpr uint32_t THREADS_PER_WARP = 64;
constexpr uint32_t NUM_XCD = 8;
#ifdef __gfx950__
constexpr uint32_t NUM_CUS_PER_XCD = 32;

// Parameters
constexpr uint32_t WARP_TILE_M = 1;
constexpr uint32_t WARP_TILE_N = 3;
constexpr uint32_t NUM_WARPS_M = 2;
constexpr uint32_t NUM_WARPS_N = 2;
constexpr uint32_t NUM_CU_M = 4;
constexpr uint32_t NUM_CU_N = 8;
constexpr uint32_t NUM_XCD_M = 2;
constexpr uint32_t NUM_XCD_N = 4;

#else
constexpr uint32_t NUM_CUS_PER_XCD = 38 * 2;

// Parameters
constexpr uint32_t WARP_TILE_M = 1;
constexpr uint32_t WARP_TILE_N = 3;
constexpr uint32_t NUM_WARPS_M = 2;
constexpr uint32_t NUM_WARPS_N = 1;
constexpr uint32_t NUM_CU_M = 4;
constexpr uint32_t NUM_CU_N = 16;
constexpr uint32_t NUM_XCD_M = 2;
constexpr uint32_t NUM_XCD_N = 4;
#endif
constexpr uint32_t NUM_CUS = NUM_XCD * NUM_CUS_PER_XCD;

// Derived
constexpr uint32_t THREADS_PER_BLOCK = THREADS_PER_WARP * NUM_WARPS_M * NUM_WARPS_N;

// Magic shuffle
__device__ inline int read_b_scale_base(int n_row, int k_group, int b_scale_stride) {
    int tr = n_row / 32, r = n_row % 32;
    int tc = k_group / 8, c = k_group % 8;
    int base  = tr * (32 * b_scale_stride) + tc * 256;
    int inner = (r / 16) + (c / 4) * 2 + (r % 16) * 4 + (c % 4) * 64;
    return base + inner;
}

// Only used if timing is enabled.
__device__ uint32_t g_run_idx = 0;
__device__ inline uint32_t hash(uint32_t x) {
    x ^= x >> 16;
    x *= 0x45d9f3b;
    x ^= x >> 16;
    x *= 0x45d9f3b;
    x ^= x >> 16;
    return x;
}

__launch_bounds__(THREADS_PER_BLOCK, 1)
__global__ void kernel(
    const uint16_t* a,
    const uint8_t* b_q, const uint8_t* b_shuffle, const uint8_t* b_scale_sh,
    uint16_t* c
) {
    // Timing
    constexpr bool TIMEIT = false;
    uint64_t _ts[16];
    int _ti = 0;
    auto TIME = [&]() {
        if constexpr (TIMEIT) {
            __builtin_amdgcn_sched_barrier(0);
            _ts[_ti++] = __builtin_amdgcn_s_memrealtime();
            __builtin_amdgcn_sched_barrier(0);
        }
    };

    TIME();

    constexpr uint32_t M = 256;
    constexpr uint32_t N = 3072;
    constexpr uint32_t K = 1536;
    constexpr uint32_t SPLIT_K = 1;
    static_assert(MFMA_DIM_NM * WARP_TILE_M * NUM_WARPS_M * NUM_CU_M * NUM_XCD_M == M);
    static_assert(MFMA_DIM_NM * WARP_TILE_N * NUM_WARPS_N * NUM_CU_N * NUM_XCD_N == N);
    constexpr int b_scale_stride = CDIV(K / 32, 8) * 8;

    constexpr int BLOCK_M = MFMA_DIM_NM * WARP_TILE_M * NUM_WARPS_M;
    constexpr int BLOCK_N = MFMA_DIM_NM * WARP_TILE_N * NUM_WARPS_N;

    // Get XCD coordinate
    uint32_t xcd_id = blockIdx.x / NUM_CUS_PER_XCD;
    uint32_t xcd_worker_id = blockIdx.x % NUM_CUS_PER_XCD;
    uint32_t xcd_n = xcd_id % NUM_XCD_N;
    uint32_t xcd_m = xcd_id / NUM_XCD_N;

    // Get CU coordinate
    if (xcd_worker_id >= NUM_CU_M*NUM_CU_N) {
        return;
    }
    uint32_t cu_n = xcd_worker_id % NUM_CU_N;
    uint32_t cu_m = xcd_worker_id / NUM_CU_N;

    // Get warp id (Use SGPR for it)
    uint32_t warp_id = __builtin_amdgcn_readfirstlane(threadIdx.x / THREADS_PER_WARP);
    uint32_t warp_n = warp_id % NUM_WARPS_N;
    uint32_t warp_m = warp_id / NUM_WARPS_N;

    // Get lane assignment
    uint32_t lane_index = threadIdx.x % THREADS_PER_WARP;
    uint32_t lane_row_nm = lane_index % MFMA_DIM_NM;
    uint32_t lane_col_k = SCALE_GROUP_SIZE * (lane_index / MFMA_DIM_NM); // By chance, a lane sends exactly one scale group.

    // Get tile offset
    uint32_t tile_n_offset = xcd_n * (N / NUM_XCD_N) + cu_n * (N / NUM_XCD_N / NUM_CU_N);
    uint32_t tile_m_offset = xcd_m * (M / NUM_XCD_M) + cu_m * (M / NUM_XCD_M / NUM_CU_M);
    uint32_t warp_n_offset = warp_n * MFMA_DIM_NM * WARP_TILE_N;
    uint32_t warp_m_offset = warp_m * MFMA_DIM_NM * WARP_TILE_M;

    // Shared Memory
    __shared__ uint8_t s_a_scale[BLOCK_M][K / SCALE_GROUP_SIZE];
    __shared__ uint8_t s_a[BLOCK_M][K / 2];
    __shared__ uint8_t s_b_scale[BLOCK_N][K / SCALE_GROUP_SIZE];
    __shared__ uint8_t s_b[BLOCK_N][K / 2];
    auto swizzle_k = [](int r, int c) { return swizzle<SCALE_GROUP_SIZE / 2>(r, c); };

    // ======================
    // A: Global -> Register
    // ======================

    TIME();

    constexpr uint32_t A_CHUNKS_K = K / SCALE_GROUP_SIZE;
    constexpr uint32_t A_CHUNK_JOBS = BLOCK_M * A_CHUNKS_K;
    constexpr uint32_t A_CHUNK_JOBS_PER_THREAD = CDIV(A_CHUNK_JOBS, THREADS_PER_BLOCK);
    u16x32_t a_jobs[A_CHUNK_JOBS_PER_THREAD];
    uint8_t a_scale_jobs[A_CHUNK_JOBS_PER_THREAD];
#pragma unroll
    for (uint32_t a_chunk_job_idx = 0; a_chunk_job_idx < A_CHUNK_JOBS_PER_THREAD; a_chunk_job_idx++) {
        uint32_t a_chunk_job = a_chunk_job_idx * THREADS_PER_BLOCK + threadIdx.x;
        uint32_t a_transfer_k = (a_chunk_job % A_CHUNKS_K) * (K / A_CHUNKS_K);
        uint32_t a_transfer_m = a_chunk_job / A_CHUNKS_K;
        uint32_t a_row_idx = tile_m_offset + a_transfer_m;
        *(u16x32_t*)(&a_jobs[a_chunk_job_idx]) = *(const u16x32_t*)(a + a_row_idx * K + a_transfer_k);
    }
    auto a_process_job_step1 = [&](int a_chunk_job_idx) {
        __builtin_amdgcn_sched_barrier(0);
        uint32_t a_chunk_job = a_chunk_job_idx * THREADS_PER_BLOCK + threadIdx.x;
        uint32_t a_transfer_k = (a_chunk_job % A_CHUNKS_K) * (K / A_CHUNKS_K);
        uint32_t a_transfer_m = a_chunk_job / A_CHUNKS_K;
        uint8_t a_scale = bf16x32_to_scale_e8m0(&a_jobs[a_chunk_job_idx]);
        s_a_scale[a_transfer_m][a_transfer_k / SCALE_GROUP_SIZE] = a_scale;
        a_scale_jobs[a_chunk_job_idx] = a_scale;
        __builtin_amdgcn_sched_barrier(0);
    };
    auto a_process_job_step2 = [&](int a_chunk_job_idx) {
        __builtin_amdgcn_sched_barrier(0);
        uint32_t a_chunk_job = a_chunk_job_idx * THREADS_PER_BLOCK + threadIdx.x;
        uint32_t a_transfer_k = (a_chunk_job % A_CHUNKS_K) * (K / A_CHUNKS_K);
        uint32_t a_transfer_m = a_chunk_job / A_CHUNKS_K;
        float a_scale_f32 = e8m0_to_f32(a_scale_jobs[a_chunk_job_idx]);
        u32x4_t a_reg;
#pragma unroll
        for (int reg_idx = 0; reg_idx < 4; reg_idx++) {
            const bf16x2_t* a_row_v2 = (const bf16x2_t*)(&a_jobs[a_chunk_job_idx]) + 4 * reg_idx;
#ifdef __gfx950__
            a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(0, a_row_v2[0], a_scale_f32, 0);
            a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[1], a_scale_f32, 1);
            a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[2], a_scale_f32, 2);
            a_reg[reg_idx] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(a_reg[reg_idx], a_row_v2[3], a_scale_f32, 3);
#else
            a_reg[reg_idx] = __builtin_bit_cast(uint32_t, a_row_v2[0]) ^ __builtin_bit_cast(uint32_t, a_row_v2[1])
                           ^ __builtin_bit_cast(uint32_t, a_row_v2[2]) ^ __builtin_bit_cast(uint32_t, a_row_v2[3])
                           ^ __builtin_bit_cast(uint32_t, a_scale_f32);
#endif
        }
        *(u32x4_t*)&s_a[a_transfer_m][swizzle_k(a_transfer_m, a_transfer_k / 2)] = a_reg;
        __builtin_amdgcn_sched_barrier(0);
    };
    uint32_t a_work_idx = 0;

    // ======================
    // B: Global -> Register
    // ======================

    TIME();

    constexpr uint32_t B_BYTES_PER_JOB = 16;
    constexpr uint32_t B_TOTAL_BYTES = BLOCK_N * (K / 2);
    constexpr uint32_t B_BYTES_PER_JOBGROUP = THREADS_PER_BLOCK * B_BYTES_PER_JOB;
    constexpr uint32_t B_NUM_JOBGROUPS = B_TOTAL_BYTES / B_BYTES_PER_JOBGROUP;
    static_assert(B_TOTAL_BYTES % B_BYTES_PER_JOBGROUP == 0);

    u32x4_t b_jobs[B_NUM_JOBGROUPS];
    uint8_t b_scale_jobs[B_NUM_JOBGROUPS];
    auto b_flat_byte = [&](uint32_t b_jobgroup_idx) -> uint32_t {
        return b_jobgroup_idx * B_BYTES_PER_JOBGROUP + threadIdx.x * B_BYTES_PER_JOB;
    };
    const uint8_t* b_base = b_shuffle + tile_n_offset * (K / 2);
    uint32_t b_thread_offset = threadIdx.x * B_BYTES_PER_JOB;

    auto b_process_job_step0 = [&](int b_jobgroup_idx) {
        __builtin_amdgcn_sched_barrier(0);
        uint32_t b_jobgroup_start_offset = b_jobgroup_idx * B_BYTES_PER_JOBGROUP;
        uint32_t b_job_offset = b_jobgroup_start_offset + b_thread_offset;
        uint32_t b_transfer_n = b_job_offset / (K / 2);
        uint32_t b_transfer_k_byte = b_job_offset % (K / 2);
        uint32_t b_row_idx = tile_n_offset + b_transfer_n;

        b_scale_jobs[b_jobgroup_idx] = b_scale_sh[read_b_scale_base(b_row_idx, b_transfer_k_byte / B_BYTES_PER_JOB, b_scale_stride)];

        uint32_t warp_base = b_jobgroup_start_offset + warp_id * (THREADS_PER_WARP * B_BYTES_PER_JOB);
#ifdef __gfx950__
        void* src = (void*)(b_base + warp_base + lane_index * 16);
        uint32_t lds_off = (uint32_t)(uintptr_t)((uint8_t*)s_b + warp_base);
        asm volatile(
            "s_mov_b32 m0, %0\n\t"
            "global_load_lds_dwordx4 %1, off\n\t"
            :: "s"(lds_off), "v"(src)
            : "memory", "m0"
        );
#else
#pragma unroll
        for (uint32_t sub = 0; sub < 4; sub++) {
            uint32_t chunk_base = warp_base + sub * 256;
            void* src = (void*)(b_base + chunk_base + lane_index * 4);
            uint32_t lds_off = (uint32_t)(uintptr_t)((uint8_t*)s_b + chunk_base);
            asm volatile(
                "s_mov_b32 m0, %0\n\t"
                "global_load_lds_dword %1, off\n\t"
                :: "s"(lds_off), "v"(src)
                : "memory", "m0"
            );
        }
#endif

        __builtin_amdgcn_sched_barrier(0);
    };

    auto b_process_job_step1 = [&](int b_jobgroup_idx) {
        __builtin_amdgcn_sched_barrier(0);
        uint32_t b_jobgroup_start_offset = b_jobgroup_idx * B_BYTES_PER_JOBGROUP;
        uint32_t b_job_offset = b_jobgroup_start_offset + b_thread_offset;
        uint32_t b_transfer_n = b_job_offset / (K / 2);
        uint32_t b_transfer_k_byte = b_job_offset % (K / 2);

        s_b_scale[b_transfer_n][b_transfer_k_byte / B_BYTES_PER_JOB] = b_scale_jobs[b_jobgroup_idx];
        __builtin_amdgcn_sched_barrier(0);
    };

#pragma unroll
    for (uint32_t b_jobgroup_idx = 0; b_jobgroup_idx < B_NUM_JOBGROUPS; b_jobgroup_idx++) {
        b_process_job_step0(b_jobgroup_idx);

        if (b_jobgroup_idx % 3 == 2) {
            if (a_work_idx % 2 == 0) {
                a_process_job_step1(a_work_idx / 2);
            } else {
                a_process_job_step2(a_work_idx / 2);
            }
            a_work_idx += 1;
        }
    }

    // ======================
    // B: Register -> LDS
    // ======================

    TIME();

#pragma unroll
    for (uint32_t b_jobgroup_idx = 0; b_jobgroup_idx < B_NUM_JOBGROUPS; b_jobgroup_idx++) {
        if (b_jobgroup_idx % 3 == 2) {
            if (a_work_idx % 2 == 0) {
                a_process_job_step1(a_work_idx / 2);
            } else {
                a_process_job_step2(a_work_idx / 2);
            }
            a_work_idx += 1;
        }

        b_process_job_step1(b_jobgroup_idx);
    }

    // ======================
    // A: Register -> Quantize -> LDS
    // ======================

    TIME();

    // NOTE: This task was interwoven above. This while loop should not be hit.
#pragma unroll
    while (a_work_idx < A_CHUNK_JOBS_PER_THREAD * 2) {
        if (a_work_idx % 2 == 0) {
            a_process_job_step1(a_work_idx / 2);
        } else {
            a_process_job_step2(a_work_idx / 2);
        }
        a_work_idx += 1;
    }

    // ======================
    // LDS->MFMA
    // ======================

    __syncthreads(); // Ensure LDS is populated

    TIME();

    f32x4_t result[WARP_TILE_M][WARP_TILE_N] = {};
    constexpr uint32_t K_ITERS = K / (SPLIT_K * MFMA_DIM_K);
    uint32_t k_start = (SPLIT_K > 1) ? blockIdx.z * (K / SPLIT_K) : 0;
    constexpr uint32_t PREFETCH = 2;

    uint8_t a_scale[PREFETCH][WARP_TILE_M];
    u32x8_t a_reg[PREFETCH][WARP_TILE_M];
    uint8_t b_scale[PREFETCH][WARP_TILE_N];
    u32x8_t b_reg[PREFETCH][WARP_TILE_N];
    auto load_lds = [&](uint32_t k_iter) {
        __builtin_amdgcn_sched_barrier(0);
        uint32_t k_offset = k_start + k_iter * MFMA_DIM_K;
        uint32_t k_idx_start = k_offset + lane_col_k;
        uint32_t buf_idx = k_iter % PREFETCH;

#pragma unroll
        for (uint32_t i = 0; i < WARP_TILE_M; i++) {
            uint32_t a_row = warp_m_offset + i * MFMA_DIM_NM + lane_row_nm;
            a_scale[buf_idx][i] = s_a_scale[a_row][k_idx_start / SCALE_GROUP_SIZE];
            *(u32x4_t*)&a_reg[buf_idx][i] = *(const u32x4_t*)&s_a[a_row][swizzle_k(a_row, k_idx_start / 2)];
        }
#pragma unroll
        for (uint32_t i = 0; i < WARP_TILE_N; i++) {
            uint32_t b_row = warp_n_offset + i * MFMA_DIM_NM + lane_row_nm;
            b_scale[buf_idx][i] = s_b_scale[b_row][k_idx_start / SCALE_GROUP_SIZE];
            uint32_t off = (warp_n_offset + i * MFMA_DIM_NM) * (K / 2)
                        + (k_idx_start / 32) * 256
                        + lane_row_nm * 16;
            *(u32x4_t*)&b_reg[buf_idx][i] = *(const u32x4_t*)((const uint8_t*)s_b + off);
        }
        __builtin_amdgcn_sched_barrier(0);
    };

#pragma unroll
    for (uint32_t i = 0; i < min(K_ITERS, PREFETCH); i++) {
        load_lds(i);
    }
#pragma unroll
    for (uint32_t k_iter = 0; k_iter < K_ITERS; k_iter++) {
        uint32_t k_prefetch_iter = k_iter + PREFETCH - 1;
        if (k_prefetch_iter < K_ITERS) {
            load_lds(k_prefetch_iter);
        }
        uint32_t k_offset = k_start + k_iter * MFMA_DIM_K;
        uint32_t k_idx_start = k_offset + lane_col_k;
        uint32_t buf_idx = k_iter % PREFETCH;

#pragma unroll
        for (uint32_t i = 0; i < WARP_TILE_M; i++) {
#pragma unroll
            for(uint32_t j = 0; j < WARP_TILE_N; j++) {
#ifdef __gfx950__
                result[i][j] = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg[buf_idx][i], b_reg[buf_idx][j], result[i][j], 4, 4, 0, a_scale[buf_idx][i], 0, b_scale[buf_idx][j]);
#else
                float scale = e8m0_to_f32(a_scale[buf_idx][i]) * e8m0_to_f32(b_scale[buf_idx][j]);
                f32x4_t mix;
                for (uint32_t r = 0; r < 4; r++)
                    mix[r] = scale * fp4_e2m1_to_f32(a_reg[buf_idx][i][r]) * fp4_e2m1_to_f32(b_reg[buf_idx][j][r]);
                result[i][j] += mix;
#endif
            }
        }
    }

    // ======================
    // Reg->Write to Global
    // ======================

    TIME();

#pragma unroll
    for (uint32_t i = 0; i < WARP_TILE_M; i++) {
#pragma unroll
        for (uint32_t j = 0; j < WARP_TILE_N; j++) {
            uint32_t result_row_offset = 4 * (lane_index / 16);
            uint32_t result_column_offset = lane_index % 16;
#pragma unroll
            for (uint32_t f = 0; f < 4; f++) {
                uint32_t row = tile_m_offset + warp_m_offset + i*MFMA_DIM_NM + result_row_offset + f;
                uint32_t col = tile_n_offset + warp_n_offset + j*MFMA_DIM_NM + result_column_offset;
                c[row * N + col] = f32_to_bf16((float)result[i][j][f]);
            }
        }
    }

    TIME();

    if constexpr (TIMEIT) {
        __shared__ uint32_t _run_idx;
        if (threadIdx.x == 0) {
            _run_idx = atomicAdd(&g_run_idx, 1) / NUM_CUS;
        }
        __syncthreads();
        uint32_t h = hash(_run_idx);
        if (threadIdx.x == (h % THREADS_PER_BLOCK) && blockIdx.x == (h % NUM_CUS)) {
            printf("%u:CU[%u,%u]:", g_run_idx, threadIdx.x, blockIdx.x);
            for (int i = 1; i < _ti; i++) {
                printf(" %llu", (unsigned long long)(_ts[i] - _ts[0]));
            }
            printf("\n");
        }
    }
}

void entry(
    const uintptr_t a, const uintptr_t b,
    const uintptr_t b_q, const uintptr_t b_shuffle, const uintptr_t b_scale_sh,
    uintptr_t c,
    int M, int N, int K
) {
    if (M == 256 && N == 3072 && K == 1536) {
        kernel<<<dim3(NUM_CUS), dim3(THREADS_PER_BLOCK)>>>(
            (const uint16_t*)a,
            (const uint8_t*)b_q, (const uint8_t*)b_shuffle, (const uint8_t*)b_scale_sh,
            (uint16_t*)c
        );
    } else {
        // No impl
    }
}
"""

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

    def __init__(self, M, N, K):
        CUDA_PRELUDE = """
#ifndef __gfx950__
#define __gfx950__
#endif
"""
        cflags = ["--offload-arch=gfx950", "-std=c++20", "-O3", "-ffast-math", "-march=native", "-funroll-loops", "-fomit-frame-pointer"]
        cflags.extend([f"-DM_DIM={M}", f"-DN_DIM={N}", f"-DK_DIM={K}"])
        self.M = M
        self.N = N
        self.K = K
        self.out = torch.empty((M, N), dtype=torch.bfloat16, device='cuda')
        cuda_src = None
        if M == 4 and N == 2880 and K == 512:
            cuda_src = CUDA_SRC_SHAPE1
        elif M == 16 and N == 2112 and K == 7168:
            cuda_src = CUDA_SRC_SHAPE2
        elif M == 32 and N == 4096 and K == 512:
            cuda_src = CUDA_SRC_SHAPE3
        elif M == 32 and N == 2880 and K == 512:
            cuda_src = CUDA_SRC_SHAPE4
        elif M == 64 and N == 7168 and K == 2048:
            cuda_src = CUDA_SRC_SHAPE5
        elif M == 256 and N == 3072 and K == 1536:
            cuda_src = CUDA_SRC_SHAPE6
        else:
            cuda_src = CUDA_SRC_SHAPE_GENERAL
        self.module = load_inline(
            name=f"solution_{M}_{N}_{K}",
            cpp_sources=[CPP_WRAPPER],
            cuda_sources=[CUDA_PRELUDE + cuda_src],
            functions=['entry'],
            with_cuda=True,
            verbose=False,
            extra_cuda_cflags=cflags,
            extra_cflags=cflags,
        )
    
    def inference(self, A, B, B_q, B_shuffle, B_scale_sh):
        self.module.entry(
            A.data_ptr(), B.data_ptr(),
            B_q.data_ptr(), B_shuffle.data_ptr(), B_scale_sh.data_ptr(),
            self.out.data_ptr(),
            self.M, self.N, self.K,
        )
        return self.out

_compiled_modules: dict[tuple[int, int, int], CompiledModule] = {}

def custom_kernel(data: input_t) -> output_t:
    global _compiled_modules
    A, B, B_q, B_shuffle, B_scale_sh = data
    M, K = A.shape
    N, _  = B.shape

    key = (M, N, K)
    if key not in _compiled_modules:
        _compiled_modules[key] = CompiledModule(M, N, K)
    
    return  _compiled_modules[key].inference(A, B, B_q, B_shuffle, B_scale_sh)
scrolls · 3129 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