Skip to content
KernelIndex
Search⌘K

submission 607214

bigboiaadulla18 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_with_fusion_v3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-607214?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
9.35µs
#169 of 1143
2026-03-22

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:e054c8e406e16d8c24ead904066557b345c47fe63161f9400a546ac23c93c228
license declaredunknown
license concludedunknown
authorsbigboiaadulla18
imported2026-08-15

Techniques

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

fp4Fused BF16→MXFP4 quant + FP4 GEMM kernel.
shared-memory__shared__ uint8_t A_lds[4 * A_LDS_SLOT];
split-ktypename OutType, bool IS_SPLITK, bool kCheckOOB, bool B_IN_LDS = false>

Kernel source

submission_with_fusion_v3.py798 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

"""
Fused BF16→MXFP4 quant + FP4 GEMM kernel.
A (bf16) is quantized to MXFP4 on-the-fly inside the GEMM kernel.
B is pre-quantized and pre-shuffled. Scales in e8m0-shuffled layout.
Uses __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4 (gfx950).
"""

import os
import sys

import aiter
import torch
from aiter import dtypes, QuantType
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
from torch.utils.cpp_extension import load_inline


# K must be divisible by 64 (scale group 32 and fp4 pack 2)
SCALE_GROUP_SIZE = 32


os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"

cuda_src = r"""
#include <hip/hip_runtime.h>
#include <hip/hip_bf16.h>
#include <cstdint>
#include <type_traits>

// ---- vector types ----
typedef int32_t  v8i32 __attribute__((ext_vector_type(8)));
typedef float    v4f32 __attribute__((ext_vector_type(4)));
typedef float    v16f32 __attribute__((ext_vector_type(16)));
typedef int32_t  i32x4 __attribute__((ext_vector_type(4)));
typedef uint32_t __attribute__((address_space(3)))* as3_uint32_ptr;

// ---- buffer_load_lds intrinsic (direct global → LDS) ----
extern "C" __device__ void llvm_amdgcn_raw_buffer_load_lds(
    i32x4 rsrc, as3_uint32_ptr lds_ptr, int size,
    int voffset, int soffset, int offset, int aux)
    __asm("llvm.amdgcn.raw.buffer.load.lds");

struct buffer_resource {
    uint64_t ptr;
    uint32_t range;
    uint32_t config;
};

__device__ inline i32x4 make_srsrc(const void* ptr, uint32_t range_bytes) {
    buffer_resource rsrc = {reinterpret_cast<uint64_t>(ptr), range_bytes, 0x110000};
    return *reinterpret_cast<const i32x4*>(&rsrc);
}

// ---- shared constants ----
static constexpr int WAVE_SIZE = 64;

// ---- accumulator type trait ----
template <int MFMA_SIZE> struct AccType;
template <> struct AccType<16> { using type = v4f32; };
template <> struct AccType<32> { using type = v16f32; };

// ---- e8m0-shuffled scale offset ----
__device__ __forceinline__
int shuffled_scale_offset(int row, int col, int sn_pad) {
    int m_block = row >> 5;
    int m_half  = (row >> 4) & 1;
    int m_in    = row & 15;
    int s_block = col >> 3;
    int s_half  = (col >> 2) & 1;
    int s_in    = col & 3;
    return m_block * (sn_pad << 5)
         + s_block * 256
         + s_in    * 64
         + m_in    * 4
         + s_half  * 2
         + m_half;
}

// ---- vmcnt helper ----
// GFX940+ s_waitcnt encoding: vmcnt is 6-bit (0-63), split across bits [3:0] and [15:14]
// expcnt is 3-bit [6:4], lgkmcnt is 6-bit [13:8] on GFX940+ (set to max 0x3F = don't wait)
template <int N>
__device__ __forceinline__ void wait_vmcnt() {
    static_assert(N >= 0 && N <= 63, "vmcnt out of range");
    constexpr int vmcnt_lo = N & 0xF;
    constexpr int vmcnt_hi = (N >> 4) & 0x3;
    constexpr int encoding = (vmcnt_hi << 14) | (0x3F << 8) | (0x7 << 4) | vmcnt_lo;
    __builtin_amdgcn_s_waitcnt(encoding);
}

// Per-sub-tile vmcnt dispatch: wait_vmcnt<BASE + (COUNT-1-sub) * STEP>
// BASE, STEP, COUNT are compile-time; sub is runtime but loop must be #pragma unroll
template <int BASE, int STEP, int COUNT>
__device__ __forceinline__ void wait_vmcnt_sub(int sub) {
    static_assert(COUNT >= 1 && COUNT <= 4, "K_MFMAS must be 1..4");
    if constexpr (COUNT == 1) {
        wait_vmcnt<BASE>();
    } else if constexpr (COUNT == 2) {
        if (sub == 0) wait_vmcnt<BASE + STEP>();
        else wait_vmcnt<BASE>();
    } else if constexpr (COUNT == 3) {
        if (sub == 0) wait_vmcnt<BASE + 2*STEP>();
        else if (sub == 1) wait_vmcnt<BASE + STEP>();
        else wait_vmcnt<BASE>();
    } else {
        if (sub == 0) wait_vmcnt<BASE + 3*STEP>();
        else if (sub == 1) wait_vmcnt<BASE + 2*STEP>();
        else if (sub == 2) wait_vmcnt<BASE + STEP>();
        else wait_vmcnt<BASE>();
    }
}

// ---- helpers ----

// Load B fragment (4×uint32) and B scale from global memory.
// For MFMA_SIZE=32, l can be 0..31 crossing two 16-row shuffle tiles.
template <int MFMA_SIZE>
__device__ __forceinline__
void load_b_frag(const uint8_t* __restrict__ B, const uint8_t* __restrict__ Bs,
                 int n_wave, int K_half, int k_byte, int k_group, int g, int l,
                 int sn_pad, v8i32& b_frag, int32_t& sb) {
    int l_in = l;
    int n_base = n_wave;
    if constexpr (MFMA_SIZE == 32) {
        l_in = l & 15;
        n_base = n_wave + (l >> 4) * 16;
    }
    const int64_t b_off = (int64_t)n_base * K_half
                        + ((k_byte >> 5) << 9)
                        + (g << 8) + (l_in << 4);
    // Single 16-byte vector load instead of 4 separate 4-byte loads
    typedef uint32_t u32x4 __attribute__((ext_vector_type(4)));
    const u32x4 b_vec = *reinterpret_cast<const u32x4*>(B + b_off);
    b_frag = {};
    b_frag[0] = b_vec[0]; b_frag[1] = b_vec[1];
    b_frag[2] = b_vec[2]; b_frag[3] = b_vec[3];
    sb = (int32_t)Bs[shuffled_scale_offset(n_wave + l, k_group + g, sn_pad)];
}

// Read A tile from LDS, quantize bf16 → MXFP4 (amax + scale + fp4 convert)
__device__ __forceinline__
void quantize_a_tile(const uint8_t* A_lds, uint32_t slot_off, uint32_t lds_read_base,
                     v8i32& a_frag, int32_t& sa) {
    const uint32_t* a_pairs = reinterpret_cast<const uint32_t*>(
        A_lds + slot_off + lds_read_base);
    uint32_t a_data[16];
    #pragma unroll
    for (int i = 0; i < 16; i++)
        a_data[i] = a_pairs[i];

    uint32_t max_packed = 0;
    #pragma unroll
    for (int i = 0; i < 16; i++) {
        uint32_t abs_pair = a_data[i] & 0x7FFF7FFFu;
        asm volatile("v_pk_max_u16 %0, %1, %2"
                     : "=v"(max_packed) : "v"(max_packed), "v"(abs_pair));
    }
    uint32_t max_abs = max(max_packed & 0xFFFFu, max_packed >> 16);
    float amax = __uint_as_float(max_abs << 16);

    uint32_t amax_u = __float_as_uint(amax);
    amax_u = (amax_u + 0x200000u) & 0xFF800000u;
    int exp_field = (int)((amax_u >> 23) & 0xFFu);
    int scale_unbiased = exp_field - 129;
    scale_unbiased = max(-127, min(127, scale_unbiased));
    sa = (int32_t)((uint8_t)(scale_unbiased + 127));

    float hw_scale = __uint_as_float((uint32_t)sa << 23);
    uint8_t fp4_bytes[16];
    #pragma unroll
    for (int i = 0; i < 16; i++) {
        uint32_t result;
        asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
                     : "=v"(result) : "v"(a_data[i]), "v"(hw_scale));
        fp4_bytes[i] = (uint8_t)result;
    }

    a_frag = {};
    const uint32_t* fp = reinterpret_cast<const uint32_t*>(fp4_bytes);
    a_frag[0] = fp[0]; a_frag[1] = fp[1];
    a_frag[2] = fp[2]; a_frag[3] = fp[3];
}

// Issue MFMA instruction (dispatches to correct intrinsic based on MFMA_SIZE)
template <int MFMA_SIZE>
__device__ __forceinline__
void do_mfma(v8i32 a_frag, v8i32 b_frag,
             typename AccType<MFMA_SIZE>::type& acc, int32_t sa, int32_t sb) {
    if constexpr (MFMA_SIZE == 16) {
        acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
            a_frag, b_frag, acc, 4, 4, 0, sa, 0, sb);
    } else {
        acc = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
            a_frag, b_frag, acc, 4, 4, 0, sa, 0, sb);
    }
}

// Store accumulator results. OutType = __hip_bfloat16 (normal) or float (SplitK partial).
// m_tile_base = block_m + wave_m * MFMA_SIZE (base M-row for this wave's tile)
template <int MFMA_SIZE, bool kCheckOOB, typename OutType>
__device__ __forceinline__
void store_acc(const typename AccType<MFMA_SIZE>::type& acc,
               OutType* __restrict__ C, int m_tile_base, int g, int l, int n_wave, int N, int M) {
    if constexpr (MFMA_SIZE == 16) {
        const int m_base = m_tile_base + 4 * g;
        #pragma unroll
        for (int i = 0; i < 4; i++) {
            if (!kCheckOOB || m_base + i < M) {
                if constexpr (std::is_same_v<OutType, float>)
                    C[(int64_t)(m_base + i) * N + n_wave + l] = acc[i];
                else
                    C[(int64_t)(m_base + i) * N + n_wave + l] = __float2bfloat16(acc[i]);
            }
        }
    } else {
        #pragma unroll
        for (int i = 0; i < 16; i++) {
            const int m_row = m_tile_base + g * 4 + (i / 4) * 8 + (i % 4);
            if (!kCheckOOB || m_row < M) {
                if constexpr (std::is_same_v<OutType, float>)
                    C[(int64_t)m_row * N + n_wave + l] = acc[i];
                else
                    C[(int64_t)m_row * N + n_wave + l] = __float2bfloat16(acc[i]);
            }
        }
    }
}

// ---- unified templated GEMM kernel ----
//
// Template params:
//   MFMA_SIZE: 16 or 32
//   TILE_M_T:  tile height (multiple of MFMA_SIZE)
//   TILE_N_T:  tile width  (multiple of MFMA_SIZE)
//   TILE_K_T:  K elements processed per iteration (multiple of MFMA_K; enables multiple MFMAs per iter)
//   OutType:   __hip_bfloat16 (normal) or float (SplitK partial sums)
//   IS_SPLITK: if true, each block processes a K-range subset; blockIdx.z = split index
//   B_IN_LDS:  if true, B data is prefetched through LDS (good for small M); if false, direct global load

template <int kM, int kN, int kK, int kNumSplits, int MFMA_SIZE, int TILE_M_T, int TILE_N_T, int TILE_K_T,
          typename OutType, bool IS_SPLITK, bool kCheckOOB, bool B_IN_LDS = false>
__global__ __launch_bounds__((TILE_M_T / MFMA_SIZE) * (TILE_N_T / MFMA_SIZE) * WAVE_SIZE)
void gemm_kernel(
    const __hip_bfloat16* __restrict__ A,
    const uint8_t*        __restrict__ B,
    const uint8_t*        __restrict__ Bs,
    OutType*              __restrict__ C,
    int M, int N, int K)
{
    // if kM, kN, and kK are specified, then overwrite M, N, and K. This will be known at compile time.
    // Having M, N, and K known at compile time allows for better optimizations.
    if constexpr (kM != 0) {
        M = kM;
    }
    if constexpr (kN != 0) {
        N = kN;
    }
    if constexpr (kK != 0) {
        K = kK;
    }

    // Compile-time constants
    constexpr int MFMA_K        = (MFMA_SIZE == 16) ? 128 : 64;
    constexpr int K_MFMAS       = TILE_K_T / MFMA_K;             // MFMAs per K-tile iteration
    constexpr int M_TILES       = TILE_M_T / MFMA_SIZE;
    constexpr int N_TILES       = TILE_N_T / MFMA_SIZE;
    constexpr int NUM_WAVES     = M_TILES * N_TILES;
    constexpr int A_LDS_ROW_T   = MFMA_K * 2;                    // 256 or 128 (per MFMA sub-tile)
    constexpr int G_SHIFT       = (MFMA_SIZE == 16) ? 4 : 5;
    constexpr int ROW_SHIFT     = (MFMA_SIZE == 16) ? 4 : 3;     // for a_voffset
    constexpr int COL_MASK      = (MFMA_SIZE == 16) ? 15 : 7;
    constexpr int ROWS_PER_WAVE = 1024 / A_LDS_ROW_T;            // rows each wave loads per chunk (4 or 8)
    constexpr int CHUNK_PER_WAVE = 1024;                          // bytes per wave per load
    constexpr int TOTAL_ROWS    = TILE_M_T;                       // rows to load
    constexpr int ROWS_PER_CHUNK = NUM_WAVES * ROWS_PER_WAVE;     // rows loaded per chunk by all waves
    constexpr int NUM_CHUNKS    = (TOTAL_ROWS + ROWS_PER_CHUNK - 1) / ROWS_PER_CHUNK; // loads per wave per sub-tile
    // A sub-slot: one MFMA_K-wide sub-tile of A
    constexpr int A_LDS_DATA_SUB = TILE_M_T * MFMA_K * 2;        // exact A data per sub-tile
    constexpr int A_LDS_LOAD_SUB = NUM_CHUNKS * NUM_WAVES * CHUNK_PER_WAVE; // load footprint per sub-tile
    constexpr int A_SUB_SLOT    = (A_LDS_DATA_SUB > A_LDS_LOAD_SUB) ? A_LDS_DATA_SUB : A_LDS_LOAD_SUB;
    // Full A slot: K_MFMAS sub-tiles
    constexpr int A_LDS_SLOT    = K_MFMAS * A_SUB_SLOT;

    static_assert(TILE_M_T % MFMA_SIZE == 0, "TILE_M_T must be multiple of MFMA_SIZE");
    static_assert(TILE_N_T % MFMA_SIZE == 0, "TILE_N_T must be multiple of MFMA_SIZE");
    static_assert(TILE_K_T % MFMA_K == 0, "TILE_K_T must be multiple of MFMA_K");
    // A_LDS_SLOT >= A_LDS_LOAD is guaranteed by max() above
    const int K_half = K >> 1;
    const int sn_pad = ((K / 32 + 7) >> 3) << 3;

    // SplitK: compute K range for this split from blockIdx.z
    int k_start = 0;
    int k_end = K;
    if constexpr (IS_SPLITK) {
        static_assert(kNumSplits != 0, "kNumSplits must be != 0 for SplitK");
        int num_splits = kNumSplits;
        int total_k_tiles = K / TILE_K_T;
        int tiles_per_split = (total_k_tiles + num_splits - 1) / num_splits;
        int my_tile_start = (int)blockIdx.z * tiles_per_split;
        int my_tile_end = min(my_tile_start + tiles_per_split, total_k_tiles);
        if (my_tile_start >= total_k_tiles) return;
        k_start = my_tile_start * TILE_K_T;
        k_end = my_tile_end * TILE_K_T;
        // Advance C to this split's slice of the workspace
        C = C + (int64_t)blockIdx.z * M * N;
    }
    const int num_k_tiles = (k_end - k_start) / TILE_K_T;

    // XCD-aware block mapping: gridDim.x is padded to multiple of 8
    if (blockIdx.x * TILE_N_T >= N) return;
    const int block_m = blockIdx.y * TILE_M_T;
    const int block_n = blockIdx.x * TILE_N_T;
    const int wave_id = threadIdx.x / WAVE_SIZE;
    const int lane    = threadIdx.x % WAVE_SIZE;

    // Wave-to-tile mapping
    const int wave_m     = wave_id / N_TILES;
    const int wave_n     = wave_id % N_TILES;
    // Lane-level mapping within MFMA tile
    const int g       = lane >> G_SHIFT;
    const int l       = lane & (MFMA_SIZE - 1);
    const int n_wave  = block_n + wave_n * MFMA_SIZE;

    typename AccType<MFMA_SIZE>::type acc = {};

    // B LDS constants (only used when B_IN_LDS)
    constexpr int B_SUB_SLOT    = B_IN_LDS ? NUM_WAVES * 1024 : 0;   // one MFMA_K sub-tile of B
    constexpr int B_LDS_SLOT    = K_MFMAS * B_SUB_SLOT;              // full TILE_K of B

    // B scale LDS constants (only used when B_IN_LDS)
    // Load a full contiguous 256-byte tile via buffer_load_lds (64 lanes × 4 bytes).
    // One tile covers 32 rows × 8 scale columns = 256 K-elements = one TILE_K iteration.
    // Reused across all K_MFMAS sub-tiles within the iteration (no per-sub load needed).
    constexpr int Bs_SCALE_SLOT = B_IN_LDS ? NUM_WAVES * 256 : 0;  // one 256B tile per wave

    // VMEM operation counts per TILE_K tile (per wave)
    constexpr int A_TILE_VMEM   = NUM_CHUNKS * K_MFMAS;              // A LDS loads per tile
    constexpr int B_TILE_VMEM   = K_MFMAS;                           // B LDS loads per tile (when B_IN_LDS)
    constexpr int BS_TILE_VMEM  = B_IN_LDS ? 1 : 0;                 // one 256B tile load per iteration
    constexpr int LDS_VMEM      = B_IN_LDS ? (A_TILE_VMEM + B_TILE_VMEM + BS_TILE_VMEM) : A_TILE_VMEM;
    // B_IN_LDS: everything through LDS (DIRECT_VMEM=0)
    // !B_IN_LDS: B data+scale loaded directly, retired per-sub in inner loop
    constexpr int DIRECT_VMEM   = B_IN_LDS ? 0 : 2 * K_MFMAS;
    constexpr int DIRECT_PER_SUB = B_IN_LDS ? 0 : 2;

    // Quad-buffered LDS for A (and B + B scales when B_IN_LDS)
    __shared__ uint8_t A_lds[4 * A_LDS_SLOT];
    __shared__ uint8_t B_lds[B_IN_LDS ? 4 * B_LDS_SLOT : 1];
    __shared__ uint8_t Bs_scale_lds[B_IN_LDS ? 4 * Bs_SCALE_SLOT : 1];
    static_assert(!B_IN_LDS || TILE_K_T == 256,
        "B scale tile load assumes TILE_K_T=256 (one 256B tile per iteration)");

    // Each wave reads its M-tile's rows from A LDS
    const uint32_t lds_read_base = (wave_m * MFMA_SIZE + l) * A_LDS_ROW_T + g * 64;

    // ---- A buffer resource ----
    // Clamp range to valid rows so OOB reads (when M % TILE_M != 0) return 0
    const int valid_m_rows = min(TILE_M_T, M - block_m);
    i32x4 a_srsrc = make_srsrc(
        reinterpret_cast<const void*>(A + (int64_t)block_m * K),
        (uint32_t)(valid_m_rows * K * 2));

    const int a_voffset_lane = (lane >> ROW_SHIFT) * (K * 2) + (lane & COL_MASK) * 16;

    // ---- B buffer resource (only for B_IN_LDS) ----
    [[maybe_unused]] i32x4 b_srsrc;
    [[maybe_unused]] int b_voffset = 0;
    if constexpr (B_IN_LDS) {
        b_srsrc = make_srsrc(
            reinterpret_cast<const void*>(B),
            (uint32_t)((uint64_t)N * K_half > 0xFFFFFFFFu ? 0xFFFFFFFFu : N * K_half));
        int l_in = l;
        int n_base = n_wave;
        if constexpr (MFMA_SIZE == 32) {
            l_in = l & 15;
            n_base = n_wave + (l >> 4) * 16;
        }
        b_voffset = n_base * K_half + (g << 8) + (l_in << 4);
    }

    // ---- B scale buffer resource (only for B_IN_LDS) ----
    // Loads a contiguous 256-byte tile via buffer_load_lds (64 lanes × 4 bytes).
    // One tile covers 32 rows × 8 scale columns = 256 K-elements.
    // bs_soffset_row = m_block*(sn_pad<<5), constant per wave (row-dependent base).
    // bs_lds_lane_base = g*64 + m_in*4 + m_half (per-lane constant for reading).
    [[maybe_unused]] i32x4 bs_srsrc;
    [[maybe_unused]] int bs_soffset_row = 0;
    [[maybe_unused]] int bs_lds_lane_base = 0;
    if constexpr (B_IN_LDS) {
        uint32_t bs_range = (uint32_t)(((N + 31u) >> 5) * ((uint32_t)sn_pad << 5));
        bs_srsrc = make_srsrc(reinterpret_cast<const void*>(Bs), bs_range);
        int row = n_wave + l;
        int m_block = row >> 5;
        int m_half  = (row >> 4) & 1;
        int m_in    = row & 15;
        bs_soffset_row = m_block * (sn_pad << 5);
        bs_lds_lane_base = g * 64 + m_in * 4 + m_half;
    }

    // Helper: load full A tile (K_MFMAS sub-tiles) into LDS
    // k_elem = starting K element offset for this tile
    auto load_a_tile = [&](uint32_t slot_off, int k_elem) __attribute__((always_inline)) {
        #pragma unroll
        for (int sub = 0; sub < K_MFMAS; sub++) {
            int soff = (k_elem + sub * MFMA_K) * 2;
            #pragma unroll
            for (int c = 0; c < NUM_CHUNKS; c++) {
                int base_row = c * ROWS_PER_CHUNK + wave_id * ROWS_PER_WAVE;
                int a_voffset = base_row * (K * 2) + a_voffset_lane;
                as3_uint32_ptr lds_dst = (as3_uint32_ptr)(
                    reinterpret_cast<uintptr_t>(A_lds) + slot_off
                    + sub * A_SUB_SLOT
                    + (uint32_t)(c * ROWS_PER_CHUNK * A_LDS_ROW_T + wave_id * CHUNK_PER_WAVE));
                llvm_amdgcn_raw_buffer_load_lds(a_srsrc, lds_dst, 16, a_voffset, soff, 0, 0);
            }
        }
    };

    // B-through-LDS helpers (only when B_IN_LDS)
    // k_elem = starting K element offset for this tile
    auto load_b_tile = [&](uint32_t slot_off, int k_elem) __attribute__((always_inline)) {
        if constexpr (B_IN_LDS) {
            #pragma unroll
            for (int sub = 0; sub < K_MFMAS; sub++) {
                int k_sub = k_elem + sub * MFMA_K;
                int b_soff = (k_sub >> 6) << 9;
                as3_uint32_ptr lds_dst = (as3_uint32_ptr)(
                    reinterpret_cast<uintptr_t>(B_lds) + slot_off
                    + sub * B_SUB_SLOT + (uint32_t)wave_id * 1024u);
                llvm_amdgcn_raw_buffer_load_lds(b_srsrc, lds_dst, 16, b_voffset, b_soff, 0, 0);
            }
        }
    };

    // Load B scales into LDS: single contiguous 256-byte tile per iteration.
    // 64 lanes × 4 bytes = 256 bytes. Covers 32 rows × 8 scale columns = 256 K-elements.
    // soffset = bs_soffset_row + s_block*256 where s_block = k_elem/256.
    auto load_bs_tile = [&](uint32_t slot_off, int k_elem) __attribute__((always_inline)) {
        if constexpr (B_IN_LDS) {
            int bs_soff = bs_soffset_row + (k_elem / 256) * 256;
            as3_uint32_ptr lds_dst = (as3_uint32_ptr)(
                reinterpret_cast<uintptr_t>(Bs_scale_lds) + slot_off
                + (uint32_t)wave_id * 256u);
            llvm_amdgcn_raw_buffer_load_lds(bs_srsrc, lds_dst, 4, lane * 4, bs_soff, 0, 0);
        }
    };

    // Read one MFMA_K sub-tile of B from LDS into registers
    auto read_b_frag_lds = [&](uint32_t sub_off, v8i32& b_frag) __attribute__((always_inline)) {
        if constexpr (B_IN_LDS) {
            const uint32_t* b_ptr = reinterpret_cast<const uint32_t*>(
                B_lds + sub_off + wave_id * 1024u + (uint32_t)lane * 16u);
            b_frag = {};
            b_frag[0] = b_ptr[0]; b_frag[1] = b_ptr[1];
            b_frag[2] = b_ptr[2]; b_frag[3] = b_ptr[3];
        }
    };

    // Read B scale byte from the contiguous 256-byte tile in LDS.
    // Layout within 256B tile: offset = g*64 + m_in*4 + s_half*2 + m_half
    // bs_lds_lane_base = g*64 + m_in*4 + m_half (constant per lane).
    // sub selects s_half: for MFMA_K=128, sub 0 → s_half=0, sub 1 → s_half=1.
    auto read_bs_scale_lds = [&](uint32_t slot_off, int sub) -> int32_t __attribute__((always_inline)) {
        if constexpr (B_IN_LDS) {
            return (int32_t)Bs_scale_lds[slot_off + wave_id * 256 + bs_lds_lane_base + sub * 2];
        }
        return 0;
    };

    // === Prologue: load first 2 TILE_K tiles (A + B data + B scales if B_IN_LDS) ===
    load_a_tile(0, k_start);
    if constexpr (B_IN_LDS) {
        load_b_tile(0, k_start);
        load_bs_tile(0, k_start);
    }
    if (num_k_tiles >= 2) {
        load_a_tile(A_LDS_SLOT, k_start + TILE_K_T);
        if constexpr (B_IN_LDS) {
            load_b_tile(B_LDS_SLOT, k_start + TILE_K_T);
            load_bs_tile(Bs_SCALE_SLOT, k_start + TILE_K_T);
        }
    }

    // =====================================================================
    // Standard K-loop with 4-stage quad-buffering, 2-ahead prefetch
    // =====================================================================
    //
    // Each iteration processes TILE_K_T elements = K_MFMAS sub-MFMAs.
    // Quad-buffered: slots 0,1,2,3. Prefetch is 2 tiles ahead.
    //
    // VMEM accounting (per iteration):
    //   LDS_VMEM = A_TILE_VMEM [+ B_TILE_VMEM + BS_TILE_VMEM] via buffer_load_lds
    //   B_IN_LDS: DIRECT_VMEM=0 (everything through LDS), no inner vmcnt waits
    //   !B_IN_LDS: DIRECT_VMEM=2*K_MFMAS (B data+scale), per-sub vmcnt waits

    // === Main loop: t=0..num_k_tiles-3 (prefetch t+2) ===
    for (int t = 0; t < num_k_tiles - 2; t++) {
        const int k = k_start + t * TILE_K_T;
        const uint32_t cur_a_slot = (t % 4) * A_LDS_SLOT;
        [[maybe_unused]] const uint32_t cur_b_slot = B_IN_LDS ? (t % 4) * B_LDS_SLOT : 0;
        [[maybe_unused]] const uint32_t cur_bs_slot = B_IN_LDS ? (t % 4) * Bs_SCALE_SLOT : 0;

        // !B_IN_LDS: load B data+scale from global into VGPRs
        [[maybe_unused]] v8i32 b_frags[K_MFMAS];
        [[maybe_unused]] int32_t sbs[K_MFMAS];
        if constexpr (!B_IN_LDS) {
            #pragma unroll
            for (int sub = 0; sub < K_MFMAS; sub++) {
                int k_sub = k + sub * MFMA_K;
                load_b_frag<MFMA_SIZE>(B, Bs, n_wave, K_half, k_sub / 2, k_sub / 32, g, l, sn_pad, b_frags[sub], sbs[sub]);
            }
        }

        // Prefetch tile t+2 (A + B data + B scales if B_IN_LDS)
        {
            const uint32_t pf_a_slot = ((t + 2) % 4) * A_LDS_SLOT;
            const int k_pf = k_start + (t + 2) * TILE_K_T;
            load_a_tile(pf_a_slot, k_pf);
            if constexpr (B_IN_LDS) {
                const uint32_t pf_b_slot = ((t + 2) % 4) * B_LDS_SLOT;
                const uint32_t pf_bs_slot = ((t + 2) % 4) * Bs_SCALE_SLOT;
                load_b_tile(pf_b_slot, k_pf);
                load_bs_tile(pf_bs_slot, k_pf);
            }
        }

        // Wait for current tile's LDS data (loaded 2 iters ago)
        wait_vmcnt<2 * LDS_VMEM + DIRECT_VMEM>();
        __syncthreads();

        // Inner loop: K_MFMAS quantize+MFMA operations
        #pragma unroll
        for (int sub = 0; sub < K_MFMAS; sub++) {
            v8i32 a_frag; int32_t sa;
            quantize_a_tile(A_lds, cur_a_slot + sub * A_SUB_SLOT, lds_read_base, a_frag, sa);

            if constexpr (B_IN_LDS) {
                v8i32 b_frag;
                read_b_frag_lds(cur_b_slot + sub * B_SUB_SLOT, b_frag);
                int32_t sb = read_bs_scale_lds(cur_bs_slot, sub);
                do_mfma<MFMA_SIZE>(a_frag, b_frag, acc, sa, sb);
            } else {
                wait_vmcnt_sub<LDS_VMEM, DIRECT_PER_SUB, K_MFMAS>(sub);
                do_mfma<MFMA_SIZE>(a_frag, b_frags[sub], acc, sa, sbs[sub]);
            }
        }
    }

    // === Second-to-last iteration (peeled): t=num_k_tiles-2, no prefetch ===
    if (num_k_tiles >= 2) {
        const int t = num_k_tiles - 2;
        const int k = k_start + t * TILE_K_T;
        const uint32_t cur_a_slot = (t % 4) * A_LDS_SLOT;
        [[maybe_unused]] const uint32_t cur_b_slot = B_IN_LDS ? (t % 4) * B_LDS_SLOT : 0;
        [[maybe_unused]] const uint32_t cur_bs_slot = B_IN_LDS ? (t % 4) * Bs_SCALE_SLOT : 0;

        [[maybe_unused]] v8i32 b_frags[K_MFMAS];
        [[maybe_unused]] int32_t sbs[K_MFMAS];
        if constexpr (!B_IN_LDS) {
            #pragma unroll
            for (int sub = 0; sub < K_MFMAS; sub++) {
                int k_sub = k + sub * MFMA_K;
                load_b_frag<MFMA_SIZE>(B, Bs, n_wave, K_half, k_sub / 2, k_sub / 32, g, l, sn_pad, b_frags[sub], sbs[sub]);
            }
        }

        wait_vmcnt<LDS_VMEM + DIRECT_VMEM>();
        __syncthreads();

        #pragma unroll
        for (int sub = 0; sub < K_MFMAS; sub++) {
            v8i32 a_frag; int32_t sa;
            quantize_a_tile(A_lds, cur_a_slot + sub * A_SUB_SLOT, lds_read_base, a_frag, sa);

            if constexpr (B_IN_LDS) {
                v8i32 b_frag;
                read_b_frag_lds(cur_b_slot + sub * B_SUB_SLOT, b_frag);
                int32_t sb = read_bs_scale_lds(cur_bs_slot, sub);
                do_mfma<MFMA_SIZE>(a_frag, b_frag, acc, sa, sb);
            } else {
                wait_vmcnt_sub<0, DIRECT_PER_SUB, K_MFMAS>(sub);
                do_mfma<MFMA_SIZE>(a_frag, b_frags[sub], acc, sa, sbs[sub]);
            }
        }
    }

    // === Last iteration (peeled): t=num_k_tiles-1, no outstanding prefetch ===
    {
        const int t = num_k_tiles - 1;
        const int k = k_start + t * TILE_K_T;
        const uint32_t cur_a_slot = (t % 4) * A_LDS_SLOT;
        [[maybe_unused]] const uint32_t cur_b_slot = B_IN_LDS ? (t % 4) * B_LDS_SLOT : 0;
        [[maybe_unused]] const uint32_t cur_bs_slot = B_IN_LDS ? (t % 4) * Bs_SCALE_SLOT : 0;

        [[maybe_unused]] v8i32 b_frags[K_MFMAS];
        [[maybe_unused]] int32_t sbs[K_MFMAS];
        if constexpr (!B_IN_LDS) {
            #pragma unroll
            for (int sub = 0; sub < K_MFMAS; sub++) {
                int k_sub = k + sub * MFMA_K;
                load_b_frag<MFMA_SIZE>(B, Bs, n_wave, K_half, k_sub / 2, k_sub / 32, g, l, sn_pad, b_frags[sub], sbs[sub]);
            }
        }

        wait_vmcnt<0>();
        __syncthreads();

        #pragma unroll
        for (int sub = 0; sub < K_MFMAS; sub++) {
            v8i32 a_frag; int32_t sa;
            quantize_a_tile(A_lds, cur_a_slot + sub * A_SUB_SLOT, lds_read_base, a_frag, sa);

            if constexpr (B_IN_LDS) {
                v8i32 b_frag;
                read_b_frag_lds(cur_b_slot + sub * B_SUB_SLOT, b_frag);
                int32_t sb = read_bs_scale_lds(cur_bs_slot, sub);
                do_mfma<MFMA_SIZE>(a_frag, b_frag, acc, sa, sb);
            } else {
                do_mfma<MFMA_SIZE>(a_frag, b_frags[sub], acc, sa, sbs[sub]);
            }
        }
    }

    store_acc<MFMA_SIZE, kCheckOOB, OutType>(acc, C, block_m + wave_m * MFMA_SIZE, g, l, n_wave, N, M);
}

// ---- SplitK reduction kernel ----
template <int kNumSplits>
__global__ void splitk_reduce(
    const float* __restrict__ workspace,  // [num_splits, M, N]
    __hip_bfloat16* __restrict__ C,       // [M, N]
    int MN)
{
    int idx = blockIdx.x * blockDim.x + threadIdx.x;
    if (idx >= MN) return;

    float sum = 0.0f;
    #pragma unroll
    for (int s = 0; s < kNumSplits; s++) {
        sum += workspace[(int64_t)s * MN + idx];
    }
    C[idx] = __float2bfloat16(sum);
}

// ---- torch wrapper ----
#include <torch/extension.h>

template <int kM, int kN, int kK, int kNumSplits, int MFMA_SIZE, int TM, int TN, int TK, bool kCheckOOB, bool B_LDS>
void launch_gemm(const __hip_bfloat16* A, const uint8_t* B, const uint8_t* Bs,
                 __hip_bfloat16* C, int m, int n, int k) {
    constexpr int NWAVES = (TM / MFMA_SIZE) * (TN / MFMA_SIZE);
    // constexpr int NUM_XCDS = 8;
    // int n_tiles = n / TN;
    // int n_tiles_padded = ((n_tiles + NUM_XCDS - 1) / NUM_XCDS) * NUM_XCDS;
    // dim3 grid(n_tiles_padded, (m + TM - 1) / TM, 1);
    dim3 grid((n + TN - 1) / TN, (m + TM - 1) / TM, 1);
    dim3 block(NWAVES * WAVE_SIZE);
    gemm_kernel<kM, kN, kK, kNumSplits, MFMA_SIZE, TM, TN, TK, __hip_bfloat16, false, kCheckOOB, B_LDS><<<grid, block>>>(A, B, Bs, C, m, n, k);
}

template <int kM, int kN, int kK, int kNumSplits, int MFMA_SIZE, int TM, int TN, int TK, bool kCheckOOB, bool B_LDS>
void launch_gemm_splitk(const __hip_bfloat16* A, const uint8_t* B, const uint8_t* Bs,
                        float* workspace, __hip_bfloat16* C, int m, int n, int k) {
    constexpr int NWAVES = (TM / MFMA_SIZE) * (TN / MFMA_SIZE);
    // constexpr int NUM_XCDS = 8;
    // int n_tiles = n / TN;
    // int n_tiles_padded = ((n_tiles + NUM_XCDS - 1) / NUM_XCDS) * NUM_XCDS;
    // dim3 grid(n_tiles_padded, (m + TM - 1) / TM, kNumSplits);
    dim3 grid((n + TN - 1) / TN, (m + TM - 1) / TM, kNumSplits);
    dim3 block(NWAVES * WAVE_SIZE);
    gemm_kernel<kM, kN, kK, kNumSplits, MFMA_SIZE, TM, TN, TK, float, true, kCheckOOB, B_LDS><<<grid, block>>>(A, B, Bs, workspace, m, n, k);

    // Reduce partial sums across splits
    constexpr int REDUCE_THREADS = 256;
    int mn = m * n;
    int reduce_blocks = (mn + REDUCE_THREADS - 1) / REDUCE_THREADS;
    splitk_reduce<kNumSplits><<<reduce_blocks, REDUCE_THREADS>>>(workspace, C, mn);
}

at::Tensor mxfp4_gemm(
    at::Tensor A, at::Tensor B_fp4, at::Tensor B_scale,
    at::Tensor C, int m, int n, int k)
{
    TORCH_CHECK(k % 128 == 0, "k must be divisible by 128");
    TORCH_CHECK(n % 16 == 0, "n must be divisible by 16");

    const auto* A_ptr  = reinterpret_cast<const __hip_bfloat16*>(A.data_ptr());
    const auto* B_ptr  = reinterpret_cast<const uint8_t*>(B_fp4.data_ptr());
    const auto* Bs_ptr = reinterpret_cast<const uint8_t*>(B_scale.data_ptr());
    auto* C_ptr        = reinterpret_cast<__hip_bfloat16*>(C.data_ptr());

    if (m == 4 && n == 2880 && k == 512) {
        launch_gemm<4, 2880, 512, 0, 16, 16, 32, 256, true, true>(A_ptr, B_ptr, Bs_ptr, C_ptr, m, n, k);
    }
    else if (m == 16 && n == 2112 && k == 7168) {
        static constexpr int num_splits = 7;
        auto workspace = at::empty({num_splits, m, n}, A.options().dtype(at::kFloat));
        auto* ws_ptr = reinterpret_cast<float*>(workspace.data_ptr());
        launch_gemm_splitk<16, 2112, 7168, num_splits, 16, 16, 64, 256, false, true>(A_ptr, B_ptr, Bs_ptr, ws_ptr, C_ptr, m, n, k);
    }
    else if (m == 32 && n == 4096 && k == 512) {
        launch_gemm<32, 4096, 512, 0, 16, 16, 32, 256, false, true>(A_ptr, B_ptr, Bs_ptr, C_ptr, m, n, k);
    }
    else if (m == 32 && n == 2880 && k == 512) {
        launch_gemm<32, 2880, 512, 0, 16, 16, 32, 256, false, true>(A_ptr, B_ptr, Bs_ptr, C_ptr, m, n, k);
    }
    else if (m == 64 && n == 7168 && k == 2048) {
        launch_gemm<64, 7168, 2048, 0, 16, 16, 64, 256, false, false>(A_ptr, B_ptr, Bs_ptr, C_ptr, m, n, k);
    }
    else if (m == 256 && n == 3072 && k == 1536) {
        launch_gemm<256, 3072, 1536, 0, 16, 16, 64, 256, false, false>(A_ptr, B_ptr, Bs_ptr, C_ptr, m, n, k);
    }
    else {
        launch_gemm<0, 0, 0, 0, 16, 16, 64, 256, true, false>(A_ptr, B_ptr, Bs_ptr, C_ptr, m, n, k);
    }

    return C;
}
"""

cpp_src = r"""
at::Tensor mxfp4_gemm(at::Tensor A, at::Tensor B_fp4, at::Tensor B_scale,
                      at::Tensor C, int m, int n, int k);
"""


def generate_input(m: int, n: int, k: int, seed: int):  # -> input_t:
    """
    Generate random bf16 inputs A [m, k], B [n, k] and quantized MXFP4 B, shuffled B and B_scale.

    Returns:
        Tuple of (A, B), both bf16 on cuda.
    """
    assert k % 64 == 0, "k must be divisible by 64 (scale group 32 and fp4 pack 2)"
    gen = torch.Generator(device="cuda")
    gen.manual_seed(seed)
    A = torch.randn((m, k), dtype=torch.bfloat16, device="cuda", generator=gen)
    B = torch.randn((n, k), dtype=torch.bfloat16, device="cuda", generator=gen)
    B_q, B_scale_sh = _quant_mxfp4(B, shuffle=True)
    # shuffle B(weight) to (16,16) tile coalesced
    B_shuffle = shuffle_weight(B_q, layout=(16, 16))
    return (A, B, B_q, B_shuffle, B_scale_sh)


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)


if sys.stdout is None:
    sys.stdout = open("/dev/stdout", "w")
if sys.stderr is None:
    sys.stderr = open("/dev/stderr", "w")

module = load_inline(
    name="A_bf16_B_mxfp4_C_bf16_gemm",
    cpp_sources=[cpp_src],
    cuda_sources=[cuda_src],
    functions=["mxfp4_gemm"],
    verbose=True,
    extra_cuda_cflags=[
        "-O3",
        "--offload-arch=gfx950",
        "-std=c++20",
        "-ffp-contract=fast",
        "-lhip_hcc",
    ],
)


def custom_kernel(data):
    """
    data is generated by generate_input()
    """
    A, B, B_q, B_shuffle, B_scale_sh = data
    A = A.contiguous()
    m, k = A.shape
    n, _ = B.shape

    C = torch.empty((m, n), dtype=torch.bfloat16, device="cuda")
    module.mxfp4_gemm(
        A,
        B_shuffle.view(torch.uint8),
        B_scale_sh.view(torch.uint8),
        C,
        m,
        n,
        k,
    )
    return C
scrolls · 798 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