Skip to content
KernelIndex
Search⌘K

submission 597059

Ashwin Adulla · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:0340615111c479089c97614e37adeecc4b2b464bf26664214b1b4a57507d9fac
license declaredunknown
license concludedunknown
authorsAshwin Adulla
imported2026-08-26

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[3 * A_LDS_SLOT];
split-ktemplate <int MFMA_SIZE, int TILE_M_T, int TILE_N_T, typename OutType, bool IS_SPLITK,
tile-m = 16TILE_M = 16

Kernel source

submission_with_fusion.py750 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 ----
template <int N>
__device__ __forceinline__ void wait_vmcnt() {
    static_assert(N >= 0 && N <= 15, "vmcnt out of range");
    if constexpr (N == 0) asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
    else if constexpr (N == 1) asm volatile("s_waitcnt vmcnt(1)" ::: "memory");
    else if constexpr (N == 2) asm volatile("s_waitcnt vmcnt(2)" ::: "memory");
    else if constexpr (N == 3) asm volatile("s_waitcnt vmcnt(3)" ::: "memory");
    else if constexpr (N == 4) asm volatile("s_waitcnt vmcnt(4)" ::: "memory");
    else if constexpr (N == 5) asm volatile("s_waitcnt vmcnt(5)" ::: "memory");
    else if constexpr (N == 6) asm volatile("s_waitcnt vmcnt(6)" ::: "memory");
    else if constexpr (N == 7) asm volatile("s_waitcnt vmcnt(7)" ::: "memory");
    else if constexpr (N == 8) asm volatile("s_waitcnt vmcnt(8)" ::: "memory");
    else if constexpr (N == 9) asm volatile("s_waitcnt vmcnt(9)" ::: "memory");
    else if constexpr (N == 10) asm volatile("s_waitcnt vmcnt(10)" ::: "memory");
    else if constexpr (N == 11) asm volatile("s_waitcnt vmcnt(11)" ::: "memory");
    else if constexpr (N == 12) asm volatile("s_waitcnt vmcnt(12)" ::: "memory");
    else if constexpr (N == 13) asm volatile("s_waitcnt vmcnt(13)" ::: "memory");
    else if constexpr (N == 14) asm volatile("s_waitcnt vmcnt(14)" ::: "memory");
    else if constexpr (N == 15) asm volatile("s_waitcnt vmcnt(15)" ::: "memory");
}

// ---- 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, 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) {
    if constexpr (MFMA_SIZE == 16) {
        const int m_base = m_tile_base + 4 * g;
        #pragma unroll
        for (int i = 0; i < 4; i++) {
            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 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)
//   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
//   PINGPONG:  if true, use 8-wave ping-pong scheduling (requires WAVES_PER_WG == 2, B_IN_LDS == true)

template <int MFMA_SIZE, int TILE_M_T, int TILE_N_T, typename OutType, bool IS_SPLITK,
          bool B_IN_LDS = false, bool PINGPONG = 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)
{
    // Compile-time constants
    constexpr int MFMA_K        = (MFMA_SIZE == 16) ? 128 : 64;
    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 WAVES_PER_WG  = (NUM_WAVES <= 4) ? 1 : 2;
    constexpr int A_LDS_ROW_T   = MFMA_K * 2;                    // 256 or 128
    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
    constexpr int A_LDS_DATA    = TILE_M_T * MFMA_K * 2;         // exact A data per slot
    constexpr int A_LDS_LOAD    = NUM_CHUNKS * NUM_WAVES * CHUNK_PER_WAVE; // load footprint
    constexpr int A_LDS_SLOT    = (A_LDS_DATA > A_LDS_LOAD) ? A_LDS_DATA : A_LDS_LOAD;

    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");
    // A_LDS_SLOT >= A_LDS_LOAD is guaranteed by max() above
    static_assert(!PINGPONG || WAVES_PER_WG == 2,
                  "PINGPONG requires 8 waves (WAVES_PER_WG == 2)");
    static_assert(!PINGPONG || B_IN_LDS,
                  "PINGPONG requires B_IN_LDS");

    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) {
        int total_k_tiles = K / MFMA_K;
        int tiles_per_split = (total_k_tiles + (int)gridDim.z - 1) / (int)gridDim.z;
        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 * MFMA_K;
        k_end = my_tile_end * MFMA_K;
        // 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) / MFMA_K;

    // 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;
    [[maybe_unused]] const int wavegroup  = wave_id / WAVES_PER_WG;
    [[maybe_unused]] const int wave_in_wg = wave_id % WAVES_PER_WG;

    // 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_LDS_SLOT    = B_IN_LDS ? NUM_WAVES * 1024 : 0;

    // VMEM operation counts per K tile (per wave)
    constexpr int A_VMEM_OPS  = NUM_CHUNKS;   // buffer_load_lds calls for A

    // Triple-buffered LDS for A (and B when B_IN_LDS)
    __shared__ uint8_t A_lds[3 * A_LDS_SLOT];
    __shared__ uint8_t B_lds[B_IN_LDS ? 3 * B_LDS_SLOT : 1];

    // 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 ----
    i32x4 a_srsrc = make_srsrc(
        reinterpret_cast<const void*>(A + (int64_t)block_m * K),
        (uint32_t)(TILE_M_T * 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);
    }

    // Helper: load A tile into LDS (issues NUM_CHUNKS buffer_load_lds)
    auto load_a_tile = [&](uint32_t slot_off, int soff) __attribute__((always_inline)) {
        #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
                + (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)
    auto load_b_tile = [&](uint32_t slot_off, int b_soff) __attribute__((always_inline)) {
        if constexpr (B_IN_LDS) {
            as3_uint32_ptr lds_dst = (as3_uint32_ptr)(
                reinterpret_cast<uintptr_t>(B_lds) + slot_off + (uint32_t)wave_id * 1024u);
            llvm_amdgcn_raw_buffer_load_lds(b_srsrc, lds_dst, 16, b_voffset, b_soff, 0, 0);
        }
    };

    auto read_b_frag_lds = [&](uint32_t slot_off, v8i32& b_frag) __attribute__((always_inline)) {
        if constexpr (B_IN_LDS) {
            const uint32_t* b_ptr = reinterpret_cast<const uint32_t*>(
                B_lds + slot_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];
        }
    };

    // === Prologue: load first A tile (and B if B_IN_LDS) into LDS slot 0 ===
    load_a_tile(0, k_start * 2);
    if constexpr (B_IN_LDS) {
        load_b_tile(0, (k_start >> 6) << 9);
    }

  if constexpr (!PINGPONG) {
    // =====================================================================
    // Standard K-loop (no ping-pong)
    // =====================================================================
    //
    // B_IN_LDS=true:  B prefetched to LDS one K-tile ahead, read from LDS
    //   Issue: B_scale(1), A_next(A_VMEM_OPS), B_next(1)
    //   Outstanding after prev: prev_A + prev_B + new = 2*A_VMEM_OPS + 3
    //   Keep = A_VMEM_OPS + 2, second wait = A_VMEM_OPS + 1
    //
    // B_IN_LDS=false: B loaded directly from global in current iteration
    //   Issue: B_data(1) + B_scale(1), A_next(A_VMEM_OPS)
    //   Outstanding after prev: prev_A + new = 2*A_VMEM_OPS + 2
    //   Keep = A_VMEM_OPS + 2, second wait = A_VMEM_OPS

    // === Main loop: all iterations except the last ===
    for (int t = 0; t < num_k_tiles - 1; t++) {
        const int k = k_start + t * MFMA_K;
        const uint32_t cur_a_slot = (t % 3) * A_LDS_SLOT;
        [[maybe_unused]] const uint32_t cur_b_slot = B_IN_LDS ? (t % 3) * B_LDS_SLOT : 0;

        v8i32 b_frag; int32_t sb;

        if constexpr (B_IN_LDS) {
            // B scale from global (1 VMEM)
            sb = (int32_t)Bs[shuffled_scale_offset(n_wave + l, (k >> 5) + g, sn_pad)];
        } else {
            // B data + B scale from global (2 VMEM)
            load_b_frag<MFMA_SIZE>(B, Bs, n_wave, K_half, k / 2, k / 32, g, l, sn_pad, b_frag, sb);
        }

        // Load A_next (and B_next if B_IN_LDS) to LDS
        {
            const uint32_t next_a_slot = ((t + 1) % 3) * A_LDS_SLOT;
            const int k_next = k + MFMA_K;
            load_a_tile(next_a_slot, k_next * 2);
            if constexpr (B_IN_LDS) {
                const uint32_t next_b_slot = ((t + 1) % 3) * B_LDS_SLOT;
                load_b_tile(next_b_slot, (k_next >> 6) << 9);
            }
        }

        // Wait for prev A (and prev B if B_IN_LDS) LDS loads to complete.
        // Both paths: keep A_VMEM_OPS + 2 newest ops in flight
        wait_vmcnt<A_VMEM_OPS + 2>();
        __syncthreads();

        // Read A from LDS → quantize
        v8i32 a_frag; int32_t sa;
        quantize_a_tile(A_lds, cur_a_slot, lds_read_base, a_frag, sa);

        if constexpr (B_IN_LDS) {
            // Read B from LDS
            read_b_frag_lds(cur_b_slot, b_frag);
            // Wait for B scale. Keep A_next + B_next = A_VMEM_OPS + 1
            wait_vmcnt<A_VMEM_OPS + 1>();
        } else {
            // Wait for B data + B scale. Keep A_next = A_VMEM_OPS
            wait_vmcnt<A_VMEM_OPS>();
        }

        do_mfma<MFMA_SIZE>(a_frag, b_frag, acc, sa, sb);
    }

    // === Last iteration (peeled): no next A/B load ===
    {
        const int t = num_k_tiles - 1;
        const int k = k_start + t * MFMA_K;
        const uint32_t cur_a_slot = (t % 3) * A_LDS_SLOT;
        [[maybe_unused]] const uint32_t cur_b_slot = B_IN_LDS ? (t % 3) * B_LDS_SLOT : 0;

        v8i32 b_frag; int32_t sb;

        if constexpr (B_IN_LDS) {
            sb = (int32_t)Bs[shuffled_scale_offset(n_wave + l, (k >> 5) + g, sn_pad)];
        } else {
            load_b_frag<MFMA_SIZE>(B, Bs, n_wave, K_half, k / 2, k / 32, g, l, sn_pad, b_frag, sb);
        }

        wait_vmcnt<0>();
        __syncthreads();

        v8i32 a_frag; int32_t sa;
        quantize_a_tile(A_lds, cur_a_slot, lds_read_base, a_frag, sa);

        if constexpr (B_IN_LDS) {
            read_b_frag_lds(cur_b_slot, b_frag);
        }

        do_mfma<MFMA_SIZE>(a_frag, b_frag, acc, sa, sb);
    }

  } else {
    // =====================================================================
    // Ping-Pong K-loop: 8 waves, 2 per SIMD, staggered compute/memory
    // =====================================================================
    // wave_in_wg==0: compute-first  (quantize+MFMA, then load A_next+B_next)
    // wave_in_wg==1: memory-first   (load B_scale+A_next+B_next, then quantize+MFMA)
    // s_setprio(1) boosts compute wave, s_setprio(0) yields to partner
    // sched_barrier prevents compiler from reordering across phase boundaries

    for (int t = 0; t < num_k_tiles - 1; t++) {
        const int k = k_start + t * MFMA_K;
        const uint32_t cur_a_slot = (t % 3) * A_LDS_SLOT;
        const uint32_t cur_b_slot = (t % 3) * B_LDS_SLOT;
        const uint32_t next_a_slot = ((t + 1) % 3) * A_LDS_SLOT;
        const uint32_t next_b_slot = ((t + 1) % 3) * B_LDS_SLOT;
        const int k_next = k + MFMA_K;

        // Wait for all outstanding VMEM from previous iteration
        wait_vmcnt<0>();
        __syncthreads();

        if (wave_in_wg == 0) {
            // --- Compute-first path ---
            asm volatile("s_setprio 1" ::: "memory");

            // Quantize A from LDS (no VMEM)
            v8i32 a_frag; int32_t sa;
            quantize_a_tile(A_lds, cur_a_slot, lds_read_base, a_frag, sa);

            // Read B from LDS (no VMEM)
            v8i32 b_frag;
            read_b_frag_lds(cur_b_slot, b_frag);

            // Load B scale (1 VMEM, only outstanding op)
            int32_t sb = (int32_t)Bs[shuffled_scale_offset(n_wave + l, (k >> 5) + g, sn_pad)];
            wait_vmcnt<0>();

            do_mfma<MFMA_SIZE>(a_frag, b_frag, acc, sa, sb);

            asm volatile("s_setprio 0" ::: "memory");
            asm volatile("s_nop 0" ::: "memory");

            // Memory phase: load next A and B into LDS
            load_a_tile(next_a_slot, k_next * 2);
            load_b_tile(next_b_slot, (k_next >> 6) << 9);
        } else {
            // --- Memory-first path ---
            asm volatile("s_setprio 0" ::: "memory");

            // Load B scale first (becomes oldest VMEM op)
            int32_t sb = (int32_t)Bs[shuffled_scale_offset(n_wave + l, (k >> 5) + g, sn_pad)];

            // Load next A and B into LDS
            load_a_tile(next_a_slot, k_next * 2);
            load_b_tile(next_b_slot, (k_next >> 6) << 9);

            asm volatile("s_nop 0" ::: "memory");
            asm volatile("s_setprio 1" ::: "memory");

            // Quantize A from LDS (no VMEM)
            v8i32 a_frag; int32_t sa;
            quantize_a_tile(A_lds, cur_a_slot, lds_read_base, a_frag, sa);

            // Read B from LDS (no VMEM)
            v8i32 b_frag;
            read_b_frag_lds(cur_b_slot, b_frag);

            // Wait for B scale (oldest). Keep A_next + B_next in flight.
            // Outstanding: B_scale(1) + A_next(A_VMEM_OPS) + B_next(1) = A_VMEM_OPS + 2
            // Retire 1 oldest (B_scale), keep A_VMEM_OPS + 1
            wait_vmcnt<A_VMEM_OPS + 1>();

            do_mfma<MFMA_SIZE>(a_frag, b_frag, acc, sa, sb);

            asm volatile("s_setprio 0" ::: "memory");
            asm volatile("s_nop 0" ::: "memory");
        }
    }

    // === Last iteration (peeled): both waves do compute, no next loads ===
    {
        const int t = num_k_tiles - 1;
        const int k = k_start + t * MFMA_K;
        const uint32_t cur_a_slot = (t % 3) * A_LDS_SLOT;
        const uint32_t cur_b_slot = (t % 3) * B_LDS_SLOT;

        wait_vmcnt<0>();
        __syncthreads();

        asm volatile("s_setprio 1" ::: "memory");

        v8i32 a_frag; int32_t sa;
        quantize_a_tile(A_lds, cur_a_slot, lds_read_base, a_frag, sa);

        v8i32 b_frag;
        read_b_frag_lds(cur_b_slot, b_frag);

        int32_t sb = (int32_t)Bs[shuffled_scale_offset(n_wave + l, (k >> 5) + g, sn_pad)];
        wait_vmcnt<0>();

        do_mfma<MFMA_SIZE>(a_frag, b_frag, acc, sa, sb);

        asm volatile("s_setprio 0" ::: "memory");
    }

  } // end PINGPONG

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

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

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

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

template <int MFMA_SIZE, int TM, int TN, bool B_LDS = false, bool PP = false>
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);
    dim3 block(NWAVES * WAVE_SIZE);
    gemm_kernel<MFMA_SIZE, TM, TN, __hip_bfloat16, false, B_LDS, PP><<<grid, block>>>(A, B, Bs, C, m, n, k);
}

template <int MFMA_SIZE, int TM, int TN, bool B_LDS = false, bool PP = false>
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, int num_splits) {
    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, num_splits);
    dim3 block(NWAVES * WAVE_SIZE);
    gemm_kernel<MFMA_SIZE, TM, TN, float, true, B_LDS, PP><<<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<<<reduce_blocks, REDUCE_THREADS>>>(workspace, C, mn, num_splits);
}

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(m % 16 == 0, "m must be divisible by 16");
    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<16, 16, 16>(A_ptr, B_ptr, Bs_ptr, C_ptr, m, n, k);
    }
    else if (m == 16 && n == 2112 && k == 7168) {
        int num_splits = 8;
        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, 16, 64>(A_ptr, B_ptr, Bs_ptr, ws_ptr, C_ptr, m, n, k, num_splits);
    }
    else if (m == 32 && n == 4096 && k == 512) {
        launch_gemm<16, 16, 32, true>(A_ptr, B_ptr, Bs_ptr, C_ptr, m, n, k);
    }
    else if (m == 32 && n == 2880 && k == 512) {
        launch_gemm<16, 16, 32, true>(A_ptr, B_ptr, Bs_ptr, C_ptr, m, n, k);
    }
    else if (m == 64 && n == 7168 && k == 2048) {
        launch_gemm<16, 16, 64>(A_ptr, B_ptr, Bs_ptr, C_ptr, m, n, k);
    }
    else if (m == 256 && n == 3072 && k == 1536) {
        // Best so far: <16, 16, 64>
        launch_gemm<16, 16, 64>(A_ptr, B_ptr, Bs_ptr, C_ptr, m, n, k);
    }
    else {
        launch_gemm<16, 16, 64>(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

    TILE_M = 16
    m_padded = ((m + TILE_M - 1) // TILE_M) * TILE_M
    if m_padded != m:
        A = torch.nn.functional.pad(A, (0, 0, 0, m_padded - m))

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