Skip to content
KernelIndex
Search⌘K

submission 638406

Divyanshsingh1910 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:6014338e610bf736704fe30a6638f80a8b572141bab449e60f1ad0001b2718c4
license declaredunknown
license concludedunknown
authorsDivyanshsingh1910
imported2026-08-26

Techniques

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

fp4- A: load bf16 from global, quant to FP4 in registers, write to LDS
shared-memory__shared__ char A_lds[2][BM * A_LDS_ROW];
tile-k = 256- BM=32, BN=128, BK=256, 256 threads (4 waves)
tile-m = 32- BM=32, BN=128, BK=256, 256 threads (4 waves)
tile-n = 128- BM=32, BN=128, BK=256, 256 threads (4 waves)
vector-width = uint4_tusing uint4_t = uint4;

Kernel source

submission.py439 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

"""
=== v21_fused: Single fused kernel ===
bf16 A -> quantize in registers -> MFMA 16x16x128 with B_shuffle -> bf16 C
Eliminates 3 kernel launches (quant + shuffle + gemm) into 1.

Architecture:
- BM=32, BN=128, BK=256, 256 threads (4 waves)
- A: load bf16 from global, quant to FP4 in registers, write to LDS
- B: load directly from B_shuffle to VGPRs (no LDS)
- B_scale: cooperative load to LDS with shuffled index decode
- A_scale: computed during quant, stays in registers
- Double-buffered A LDS + B_scale LDS
"""

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

HIP_SRC = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <hip/hip_bf16.h>

using fp4x2_t = unsigned char;
using fp4x64_t = fp4x2_t __attribute__((ext_vector_type(32)));
using f32x4_t = float __attribute__((ext_vector_type(4)));
using i32x4_t = int __attribute__((ext_vector_type(4)));
using i32x8_t = int __attribute__((ext_vector_type(8)));
using uint4_t = uint4;

constexpr int BM = 32;
constexpr int BN = 128;
constexpr int BK = 256;
constexpr int A_LDS_ROW = BK / 2 + 4;  // 132 bytes padded

// ============================================================
// Branchless FP4 E2M1 quantization (RNE, matches aiter exactly)
// ============================================================

// Quantize 32 bf16 values to FP4, return packed 16 bytes + E8M0 scale byte
// Input: 32 bf16 values in vals[0..31]
// Output: packed[0..15] (fp4x2 bytes), scale_byte
__device__ __forceinline__ void quant_group_32(
    const __hip_bfloat16* vals, unsigned char* packed, unsigned char& scale_byte
) {
    // Step 1: Find amax
    float amax = 0.0f;
    float fvals[32];
    #pragma unroll
    for (int i = 0; i < 32; i++) {
        fvals[i] = __bfloat162float(vals[i]);
        float av = fvals[i] < 0.0f ? -fvals[i] : fvals[i];
        if (av > amax) amax = av;
    }

    // Step 2: E8M0 scale = round amax to nearest power of 2, subtract 2 from exponent
    // aiter: scale_e8m0_unbiased = floor(log2(amax_rounded)) - 2
    // The -2 maps FP4 E2M1 range [0,6] to normalized [0, ~4-6] for full utilization
    unsigned int amax_u32 = __float_as_uint(amax);
    if (amax == 0.0f) {
        scale_byte = 0;
    } else {
        unsigned int rounded_bits = (amax_u32 + 0x200000u) & 0xFF800000u;
        unsigned int exp_raw = (rounded_bits >> 23) & 0xFFu;
        scale_byte = (exp_raw >= 2u) ? (unsigned char)(exp_raw - 2u) : (unsigned char)0;
    }

    // Reconstruct scale from adjusted E8M0 for quantization: scale = 2^(scale_byte - 127)
    float scale_float = (scale_byte > 0) ? __uint_as_float(((unsigned int)scale_byte) << 23) : 0.0f;
    float inv_scale = (scale_float > 0.0f) ? (1.0f / scale_float) : 0.0f;

    // Step 3: Quantize each value to FP4 and pack pairs
    #pragma unroll
    for (int i = 0; i < 16; i++) {
        unsigned char lo, hi;

        // Even element (low nibble)
        {
            float x = fvals[2 * i];
            float ax = x < 0.0f ? -x : x;
            float normed = ax * inv_scale;

            // Branchless 7-comparison RNE
            int fp4_val = (normed > 0.25f) + (normed >= 0.75f) + (normed > 1.25f)
                        + (normed >= 1.75f) + (normed > 2.5f) + (normed >= 3.5f)
                        + (normed > 5.0f);
            // Sign bit
            int sign = (x < 0.0f) ? 8 : 0;
            lo = (unsigned char)(fp4_val | sign);
        }

        // Odd element (high nibble)
        {
            float x = fvals[2 * i + 1];
            float ax = x < 0.0f ? -x : x;
            float normed = ax * inv_scale;

            int fp4_val = (normed > 0.25f) + (normed >= 0.75f) + (normed > 1.25f)
                        + (normed >= 1.75f) + (normed > 2.5f) + (normed >= 3.5f)
                        + (normed > 5.0f);
            int sign = (x < 0.0f) ? 8 : 0;
            hi = (unsigned char)(fp4_val | sign);
        }

        packed[i] = (hi << 4) | (lo & 0x0F);
    }
}

// ============================================================
// B_shuffle direct VGPR load (same as v20)
// ============================================================
__device__ __forceinline__ fp4x64_t load_b_direct(
    const char* __restrict__ B_sh_ptr,
    int global_n_tile, int k_tile_base, int total_k_tiles,
    int group4, int lane16
) {
    int k_tile_off = group4 >> 1;
    int sub_group = group4 & 1;
    const char* tile_ptr = B_sh_ptr +
        (int64_t)(global_n_tile * total_k_tiles + k_tile_base + k_tile_off) * 512;
    const char* src = tile_ptr + sub_group * 256 + lane16 * 16;
    i32x4_t raw = *reinterpret_cast<const i32x4_t*>(src);
    i32x8_t full = {raw[0], raw[1], raw[2], raw[3], 0, 0, 0, 0};
    return __builtin_bit_cast(fp4x64_t, full);
}

// ============================================================
// Read A from LDS for 16x16x128 MFMA
// ============================================================
__device__ __forceinline__ fp4x64_t load_a_from_lds(
    const char* __restrict__ lds,
    int mt, int ki, int lane16, int group4
) {
    int row = mt * 16 + lane16;
    int col = ki * 64 + group4 * 16;
    const char* ptr = lds + row * A_LDS_ROW + col;
    i32x4_t raw = *reinterpret_cast<const i32x4_t*>(ptr);
    i32x8_t full = {raw[0], raw[1], raw[2], raw[3], 0, 0, 0, 0};
    return __builtin_bit_cast(fp4x64_t, full);
}

// ============================================================
// B_scale unshuffle index (inline, from shuffled layout)
// ============================================================
__device__ __forceinline__ unsigned char load_b_scale_unshuffled(
    const unsigned char* __restrict__ B_scale_sh,
    int n_row, int k_scale_idx, int N, int K_div_32
) {
    int sm = ((N + 255) / 256) * 256;
    int sn = ((K_div_32 + 7) / 8) * 8;
    (void)sm;

    int d0 = n_row / 32;
    int d1 = (n_row & 31) / 16;
    int d2 = n_row & 15;
    int d3 = k_scale_idx / 8;
    int d4 = (k_scale_idx & 7) / 4;
    int d5 = k_scale_idx & 3;

    int idx = d0 * (sn / 8 * 4 * 16 * 2 * 2)
            + d3 * (4 * 16 * 2 * 2)
            + d5 * (16 * 2 * 2)
            + d2 * (2 * 2)
            + d4 * 2
            + d1;

    return B_scale_sh[idx];
}

// ============================================================
// Cooperative A quant: load bf16 from global, quant to FP4, write to LDS
// Also writes A_scale bytes to a_scale_buf for later use
//
// BM=32 rows, BK=256 cols of bf16 A
// Each row: 256 bf16 = 8 groups of 32 = 8 scale bytes, 128 packed FP4 bytes
// 256 threads total, BM*BK = 32*256 = 8192 bf16 values
// Each thread: 8192/256 = 32 bf16 values = 1 group of 32
// So each thread quants exactly 1 group and produces 16 packed bytes + 1 scale byte
// ============================================================
__device__ __forceinline__ void coop_quant_A(
    const __hip_bfloat16* __restrict__ A_bf16,
    char* __restrict__ a_lds,
    unsigned char* __restrict__ a_scale_lds,  // BM * (BK/32) = 32 * 8 = 256 bytes
    int tile_m, int k, int M, int K, int tid
) {
    // tid in [0,256): each handles 1 group of 32 bf16 values
    // 256 groups = 32 rows * 8 groups_per_row
    int row = tid / 8;          // 0..31
    int grp = tid % 8;          // 0..7 (which group of 32 within the row)
    int g_row = tile_m + row;

    __hip_bfloat16 vals[32];
    if (g_row < M) {
        int col_start = k + grp * 32;
        const __hip_bfloat16* src = A_bf16 + (int64_t)g_row * K + col_start;
        // Load 32 bf16 values (64 bytes = 4 x uint4)
        if (col_start + 31 < K) {
            *reinterpret_cast<uint4_t*>(&vals[0]) = *reinterpret_cast<const uint4_t*>(&src[0]);
            *reinterpret_cast<uint4_t*>(&vals[8]) = *reinterpret_cast<const uint4_t*>(&src[8]);
            *reinterpret_cast<uint4_t*>(&vals[16]) = *reinterpret_cast<const uint4_t*>(&src[16]);
            *reinterpret_cast<uint4_t*>(&vals[24]) = *reinterpret_cast<const uint4_t*>(&src[24]);
        } else {
            #pragma unroll
            for (int i = 0; i < 32; i++)
                vals[i] = (col_start + i < K) ? src[i] : __float2bfloat16(0.0f);
        }
    } else {
        #pragma unroll
        for (int i = 0; i < 32; i++)
            vals[i] = __float2bfloat16(0.0f);
    }

    // Quantize
    unsigned char packed[16];
    unsigned char scale_byte;
    quant_group_32(vals, packed, scale_byte);

    // Write packed FP4 to A LDS
    // Row layout in LDS: A_LDS_ROW bytes per row (132 bytes, 128 data + 4 pad)
    // Group grp occupies bytes [grp*16 .. grp*16+15] within the row
    char* dst = a_lds + row * A_LDS_ROW + grp * 16;
    *reinterpret_cast<i32x4_t*>(dst) = *reinterpret_cast<i32x4_t*>(packed);

    // Write scale byte
    // Layout: a_scale_lds[row * 8 + grp]
    a_scale_lds[row * 8 + grp] = scale_byte;
}

// ============================================================
// Main fused kernel
// ============================================================
__global__
__attribute__((amdgpu_flat_work_group_size(256, 256)))
__attribute__((amdgpu_waves_per_eu(2)))
void fused_mxfp4_gemm(
    const __hip_bfloat16* __restrict__ A_bf16,
    const char* __restrict__ B_shuffle,
    const unsigned char* __restrict__ B_scale_sh,
    __hip_bfloat16* __restrict__ C,
    int M, int N, int K
) {
    const int tid = threadIdx.x;
    const int wave_id = tid >> 6;
    const int lane = tid & 63;
    const int lane16 = lane & 15;
    const int group4 = lane >> 4;

    const int tile_m = blockIdx.y * BM;
    const int tile_n = blockIdx.x * BN;
    if (tile_m >= M) return;

    const int wave_n_base = tile_n + wave_id * 32;
    const int K_div_32 = K / 32;
    const int total_k_tiles = (K / 2) / 32;

    const int gnt0 = (wave_n_base) / 16;
    const int gnt1 = (wave_n_base + 16) / 16;

    f32x4_t acc00 = {};
    f32x4_t acc01 = {};
    f32x4_t acc10 = {};
    f32x4_t acc11 = {};

    // LDS layout:
    // A_lds: double-buffered, 2 * BM * A_LDS_ROW = 2 * 32 * 132 = 8448 bytes
    // A_scale_lds: double-buffered, 2 * BM * (BK/32) = 2 * 32 * 8 = 512 bytes
    __shared__ char A_lds[2][BM * A_LDS_ROW];
    __shared__ unsigned char A_scale_lds[2][BM * (BK / 32)];  // 32 * 8 = 256 per buf

    // Prefetch first A tile: quant bf16 -> FP4 in LDS
    coop_quant_A(A_bf16, A_lds[0], A_scale_lds[0], tile_m, 0, M, K, tid);
    __builtin_amdgcn_s_barrier();

    int buf = 0;

    for (int k = 0; k < K; k += BK) {
        int next_k = k + BK;

        // Prefetch next A tile (double buffer)
        if (next_k < K) {
            coop_quant_A(A_bf16, A_lds[1 - buf], A_scale_lds[1 - buf],
                         tile_m, next_k, M, K, tid);
        }

        int k_tile_base = (k / 2) / 32;

        #pragma unroll
        for (int ki = 0; ki < 2; ki++) {  // BK/128 = 2
            int k_abs = k + ki * 128;
            int k_scale_base = k_abs / 32;  // 4 scales per 128 FP4

            // Load A from LDS
            fp4x64_t a0 = load_a_from_lds(A_lds[buf], 0, ki, lane16, group4);
            fp4x64_t a1 = load_a_from_lds(A_lds[buf], 1, ki, lane16, group4);

            // Load B directly from global
            fp4x64_t b0 = load_b_direct(B_shuffle, gnt0, k_tile_base + ki * 2,
                                         total_k_tiles, group4, lane16);
            fp4x64_t b1 = load_b_direct(B_shuffle, gnt1, k_tile_base + ki * 2,
                                         total_k_tiles, group4, lane16);

            // A scales from LDS (computed during quant)
            // For 16x16x128 MFMA: lane16 = M-row within sub-tile, group4 = K-quarter
            // A_scale_lds[row][k_scale_base - k/32 + group4] since scale indices are relative to BK chunk
            // k_scale_base = (k + ki*128)/32, within BK chunk: ki*4 + group4
            int a_scale_offset_0 = lane16 * 8 + ki * 4 + group4;        // mt=0, row=lane16
            int a_scale_offset_1 = (16 + lane16) * 8 + ki * 4 + group4; // mt=1, row=16+lane16

            unsigned char sa0 = A_scale_lds[buf][a_scale_offset_0];
            unsigned char sa1 = A_scale_lds[buf][a_scale_offset_1];

            // Clamp sa for out-of-bounds M rows
            if (tile_m + lane16 >= M) sa0 = 127;
            if (tile_m + 16 + lane16 >= M) sa1 = 127;

            // B scales from global (inline unshuffle)
            int b_col0 = wave_n_base + lane16;
            int b_col1 = wave_n_base + 16 + lane16;

            unsigned char sb0 = 127, sb1 = 127;
            if (b_col0 < N)
                sb0 = load_b_scale_unshuffled(B_scale_sh, b_col0, k_scale_base + group4, N, K_div_32);
            if (b_col1 < N)
                sb1 = load_b_scale_unshuffled(B_scale_sh, b_col1, k_scale_base + group4, N, K_div_32);

            // 4 MFMAs: [mt][nt]
            acc00 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
                a0, b0, acc00, 4, 4, 0, (unsigned int)sa0, 0, (unsigned int)sb0);
            acc01 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
                a0, b1, acc01, 4, 4, 0, (unsigned int)sa0, 0, (unsigned int)sb1);
            acc10 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
                a1, b0, acc10, 4, 4, 0, (unsigned int)sa1, 0, (unsigned int)sb0);
            acc11 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
                a1, b1, acc11, 4, 4, 0, (unsigned int)sa1, 0, (unsigned int)sb1);
        }

        if (next_k < K) {
            __builtin_amdgcn_s_barrier();
            buf = 1 - buf;
        }
    }

    // Write back results
    auto write_tile = [&](f32x4_t& acc_reg, int m_base, int n_col) {
        if (n_col >= N) return;
        #pragma unroll
        for (int i = 0; i < 4; i++) {
            int row = m_base + group4 * 4 + i;
            if (row < M)
                C[row * N + n_col] = __float2bfloat16(acc_reg[i]);
        }
    };

    int n_col0 = wave_n_base + lane16;
    int n_col1 = wave_n_base + 16 + lane16;

    write_tile(acc00, tile_m,      n_col0);
    write_tile(acc01, tile_m,      n_col1);
    write_tile(acc10, tile_m + 16, n_col0);
    write_tile(acc11, tile_m + 16, n_col1);
}

torch::Tensor fused_gemm(
    torch::Tensor A_bf16,
    torch::Tensor B_shuffle,
    torch::Tensor B_scale_sh,
    int M, int N, int K, int sn
) {
    auto C = torch::empty({M, N},
        torch::TensorOptions().dtype(torch::kBFloat16).device(A_bf16.device()));

    dim3 grid((N + BN - 1) / BN, (M + BM - 1) / BM);
    dim3 block(256);

    auto a_ptr = reinterpret_cast<const __hip_bfloat16*>(A_bf16.data_ptr());
    auto b_ptr = reinterpret_cast<const char*>(B_shuffle.data_ptr<uint8_t>());
    auto bs_ptr = reinterpret_cast<const unsigned char*>(B_scale_sh.data_ptr<uint8_t>());
    auto c_ptr = reinterpret_cast<__hip_bfloat16*>(C.data_ptr());

    fused_mxfp4_gemm<<<grid, block>>>(a_ptr, b_ptr, bs_ptr, c_ptr, M, N, K);
    return C;
}
"""

CPP_SRC = """
torch::Tensor fused_gemm(
    torch::Tensor A_bf16,
    torch::Tensor B_shuffle,
    torch::Tensor B_scale_sh,
    int M, int N, int K, int sn
);
"""

module = load_inline(
    name='mxfp4_fused_v21b',
    cpp_sources=[CPP_SRC],
    cuda_sources=[HIP_SRC],
    functions=['fused_gemm'],
    verbose=True,
    extra_cuda_cflags=["--offload-arch=gfx950", "-std=c++20", "-O3",
                       "-mllvm", "-amdgpu-early-inline-all=true"],
)


import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle


def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data
    m, k = A.shape
    n = B_q.shape[0]

    # Dispatch: fused kernel for small K, aiter for large K
    # Fused kernel wins on K=512 (11.6 us vs 18.8 us)
    # aiter wins on K>=1536 (22.9-33 us vs 26-97 us)
    if k <= 1024:
        B_sh_u8 = B_shuffle.view(torch.uint8)
        B_sc_u8 = B_scale_sh.view(torch.uint8)
        sn = ((k // 32 + 7) // 8) * 8
        return module.fused_gemm(A, B_sh_u8, B_sc_u8, m, n, k, sn)
    else:
        A_q, A_scale = dynamic_mxfp4_quant(A)
        A_scale_sh = e8m0_shuffle(A_scale)
        return aiter.gemm_a4w4(
            A_q.view(dtypes.fp4x2), B_shuffle,
            A_scale_sh.view(dtypes.fp8_e8m0), B_scale_sh,
            dtype=dtypes.bf16, bpreshuffle=True,
        )
scrolls · 439 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