Skip to content
KernelIndex
Search⌘K

submission 755166

Will Fisher · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

fettywap.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-755166?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.26µs
#51 of 1143
2026-04-07

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:0bb6b201e0e25612f9c5176bc9dc0c3eb6d6d0aadeba8c888f4504aa313d48c3
license declaredunknown
license concludedunknown
authorsWill Fisher
imported2026-08-15

Techniques

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

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

Kernel source

fettywap.py1038 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

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


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

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

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

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

extern "C" __device__ void llvm_amdgcn_raw_buffer_load_lds(
    v4i32 rsrc, as3_uint32_ptr lds_ptr, int size, int voffset, int soffset,
    int offset, int aux) __asm("llvm.amdgcn.raw.buffer.load.lds");

struct buffer_resource { uint64_t ptr; uint32_t range; uint32_t config; };
__device__ __forceinline__ v4i32 make_srsrc(const void* ptr, uint32_t range_bytes) {
    buffer_resource rsrc = {reinterpret_cast<uint64_t>(ptr), range_bytes, 0x110000};
    return *reinterpret_cast<const v4i32*>(&rsrc);
}

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

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

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

__device__ __forceinline__
uint8_t compute_e8m0_scale(const __hip_bfloat16* vals) {
    // Find max |val| using bf16 bit representation (no float conversion)
    uint16_t mx_bits = 0;
    #pragma unroll
    for (int i = 0; i < GROUP_SZ; i++) {
        uint16_t b = *reinterpret_cast<const uint16_t*>(&vals[i]) & 0x7FFF;
        mx_bits = max(mx_bits, b);
    }
    if (mx_bits == 0) return 0;

    // Promote to f32 bit space for rounding: bf16 is top 16 bits of f32
    uint32_t bits = (uint32_t)mx_bits << 16;
    bits = (bits + 0x200000u) & 0xFF800000u;
    int e8m0 = (int)(bits >> 23) - 2;
    return (uint8_t)max(0, min(254, e8m0));
}

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

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

    PACK_WORD(0);
    PACK_WORD(1);
    PACK_WORD(2);
    PACK_WORD(3);

    #undef PACK_WORD
}

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

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

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

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

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

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

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

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

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

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

    int cur = 0;

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

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

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

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

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

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

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

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

    // N=2112 = 33*64 exactly, base always 4-aligned → no edge case
    int pk0, pk1;
    asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(pk0) : "v"(s0), "v"(s1));
    asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(pk1) : "v"(s2), "v"(s3));
    int gbl = row * N + abs_col;
    *reinterpret_cast<int*>(&C_bf16[gbl]) = pk0;
    *reinterpret_cast<int*>(&C_bf16[gbl + 2]) = pk1;
}

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

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

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

    constexpr int WF_K_OFF   = TILE_K / 32;

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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


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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

        cur = nxt;
    }

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

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

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

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

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

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

// ============================================================
// Shape 6: 12 WFs, 1 K_step each, no loop.
// Each WF: load A+B → quantize → 6 MFMAs → LDS reduce.
// Grid: (N/(32*N_SUBS), M/32, 1), Block: 12 * 64 = 768
// ============================================================
template <int K_TOTAL, int M_TOTAL, int N_TOTAL, int N_WF, int N_SUBS_T>
__global__ __attribute__((amdgpu_flat_work_group_size(N_WF * WF_SIZE, N_WF * WF_SIZE)))
void mxfp4_gemm_shape6(
    const __hip_bfloat16* __restrict__ A,
    const uint8_t*        __restrict__ B_shuf,
    const uint8_t*        __restrict__ B_scale,
    __hip_bfloat16*       __restrict__ C,
    int B_scale_stride
) {
    constexpr int K = K_TOTAL;
    constexpr int M = M_TOTAL;
    constexpr int N = N_TOTAL;
    constexpr int TILE_32 = 32;
    constexpr int TILE_K_32 = 64;
    constexpr int GROUP_SZ_32 = 32;
    constexpr int N_SUBS = N_SUBS_T;
    constexpr int K_CHUNKS = 2;
    constexpr int K_STEP = K_CHUNKS * TILE_K_32;  // 128

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

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

    // This WF's K position (1 K_step per WF, 12 WFs cover all K)
    const int k_off = wf_id * K_STEP + k_chunk * GROUP_SZ_32;
    const int c_base = (wf_id * K_STEP) / 32 + k_chunk;

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

    // ======== Load A + B ========
    v4i32 a_raw[K_CHUNKS][4];
    const __hip_bfloat16* A_row = A + my_a_row * K;
    {
        const v4i32* src0 = reinterpret_cast<const v4i32*>(A_row + k_off);
        const v4i32* src1 = reinterpret_cast<const v4i32*>(A_row + k_off + TILE_K_32);
        #pragma unroll
        for (int v = 0; v < 4; v++) a_raw[0][v] = src0[v];
        #pragma unroll
        for (int v = 0; v < 4; v++) a_raw[1][v] = src1[v];
    }

    v4i32    b_data[K_CHUNKS][N_SUBS];
    uint32_t b_sc[K_CHUNKS][N_SUBS];
    #pragma unroll
    for (int ns = 0; ns < N_SUBS; ns++) {
        int tile_idx = b_tile_n_idx[ns] * (K / 32) + c_base;
        b_data[0][ns] = *reinterpret_cast<const v4i32*>(B_shuf + tile_idx * 256 + b_lane_in_tile * 16);
        int c0 = c_base;
        b_sc[0][ns] = B_scale[bs_n_blk_base[ns] + (c0/8)*256 + (c0%4)*64 + bs_n_inn + ((c0/4)%2)*2 + bs_n_sub_val[ns]];
        b_data[1][ns] = *reinterpret_cast<const v4i32*>(B_shuf + (tile_idx + 2) * 256 + b_lane_in_tile * 16);
        int c1 = c_base + 2;
        b_sc[1][ns] = B_scale[bs_n_blk_base[ns] + (c1/8)*256 + (c1%4)*64 + bs_n_inn + ((c1/4)%2)*2 + bs_n_sub_val[ns]];
    }

    // ======== Quantize A ========
    v4i32 a_q[K_CHUNKS];
    uint32_t a_s[K_CHUNKS];
    #pragma unroll
    for (int kc = 0; kc < K_CHUNKS; kc++) {
        __hip_bfloat16* vals = reinterpret_cast<__hip_bfloat16*>(a_raw[kc]);
        quantize_group(vals, a_q[kc], a_s[kc]);
    }

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

    #pragma unroll
    for (int ns = 0; ns < N_SUBS; ns++)
        acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
            to_v8(a_q[0]), to_v8(b_data[0][ns]), acc[ns],
            4, 4, 0, a_s[0], 0, b_sc[0][ns]);
    #pragma unroll
    for (int ns = 0; ns < N_SUBS; ns++)
        acc[ns] = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
            to_v8(a_q[1]), to_v8(b_data[1][ns]), acc[ns],
            4, 4, 0, a_s[1], 0, b_sc[1][ns]);

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

    // 768 threads, 3072 output elements → 4 elements per thread
    {
        constexpr int TILE_COLS = TILE_32 * N_SUBS;
        constexpr int TOTAL_ELEMS = TILE_32 * TILE_COLS;
        constexpr int ELEMS_PER_THREAD = TOTAL_ELEMS / (N_WF * WF_SIZE);

        int local_idx = tid * ELEMS_PER_THREAD;
        int row_local = local_idx / TILE_COLS;
        int col_local = local_idx % TILE_COLS;

        __hip_bfloat16 out[ELEMS_PER_THREAD];
        #pragma unroll
        for (int e = 0; e < ELEMS_PER_THREAD; e++) {
            float s = 0.0f;
            #pragma unroll
            for (int w = 0; w < N_WF; w++)
                s += lds_f32[w][row_local][col_local + e];
            out[e] = __float2bfloat16(s);
        }

        if constexpr (ELEMS_PER_THREAD == 4) {
            *reinterpret_cast<uint64_t*>(&C[(m_base + row_local) * N + n_base + col_local]) =
                *reinterpret_cast<uint64_t*>(out);
        } else if constexpr (ELEMS_PER_THREAD == 8) {
            *reinterpret_cast<v4i32*>(&C[(m_base + row_local) * N + n_base + col_local]) =
                *reinterpret_cast<v4i32*>(out);
        }
    }
}

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

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

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

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

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

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

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

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

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

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

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

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

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

        int out_col = n_tile * TILE_N + lane;
        #pragma unroll
        for (int i = 0; i < 4; i++) {
            int out_row = m_base + k_group * 4 + i;
            C[out_row * N + out_col] = __float2bfloat16(sum[i]);
        }
    }
}
// Host entry point — raw pointers, C++ dispatch
// ============================================================
void run(
    uintptr_t a_ptr,
    uintptr_t b_shuf_ptr,
    uintptr_t b_scale_ptr,
    uintptr_t c_ptr,
    int M, int N, int K,
    int B_scale_stride,
    uintptr_t ws_ptr
) {
    const auto* A      = reinterpret_cast<const __hip_bfloat16*>(a_ptr);
    const auto* B_shuf = reinterpret_cast<const uint8_t*>(b_shuf_ptr);
    const auto* B_scale = reinterpret_cast<const uint8_t*>(b_scale_ptr);
    auto* C            = reinterpret_cast<__hip_bfloat16*>(c_ptr);

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

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

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

    // Shape 5: M=64, N=7168, K=2048 — 32×64 tile, 32x32x64 MFMA, 4 wf, 2 n_subs
    if (M == 64 && K == 2048) {
        mxfp4_gemm_32x32<2048, 64, 7168, 4><<<dim3(7168 / 64, 64 / 32, 1), 4 * WF_SIZE>>>(
            A, B_shuf, B_scale, C, B_scale_stride
        );
        return;
    }

    // Shape 6: M=256, N=3072, K=1536 — 12 WFs, 1 K_step each, no loop
    if (M == 256 && K == 1536) {
        mxfp4_gemm_shape6<1536, 256, 3072, 12, 3><<<dim3(3072 / 96, 256 / 32, 1), 12 * WF_SIZE>>>(
            A, B_shuf, B_scale, C, B_scale_stride
        );
        return;
    }

}

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

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

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

# Split-K workspace (allocated once)
_ws_buf = None    # Shape 2 split-K: workspace [33 * 7 * 16 * 64] floats

def _ensure_ws_buf(device):
    global _ws_buf
    if _ws_buf is None:
        _ws_buf = torch.empty(33 * 7 * 16 * 64, dtype=torch.float32, device=device)  # N_BLOCKS * SPLIT_K * M * BLOCK_N

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

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

    _ensure_ws_buf(A.device)

    module.run(
        A.data_ptr(),
        B_shuffle.data_ptr(),
        B_scale_sh.data_ptr(),
        C.data_ptr(),
        M, N, K,
        B_scale_sh.size(1),
        _ws_buf.data_ptr(),
    )
    return C
scrolls · 1038 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