Skip to content
KernelIndex
Search⌘K

submission 517676

_knarf_04 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission-hip.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-517676?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
17.2µs
#665 of 1143
2026-03-08

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:780edc2b295c095b485effd4b73169795cfac10ce5bcee1e64a4da44e0f226fa
license declaredunknown
license concludedunknown
authors_knarf_04
imported2026-08-26

Techniques

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

fp4MXFP4 GEMM — v11: v10 + pre-allocated buffers.
shared-memory__shared__ uint8_t A_lds[2][BM * BK_HALF_PAD];
tile-m = 16if (M <= 16) { WM_val = 1; BM = 16; BN = 64; }
tile-n = 128constexpr int BN = 128;
vector-width = int4int4 v0, v1, v2, v3;

Kernel source

submission-hip.py602 lines
"""
MXFP4 GEMM — v11: v10 + pre-allocated buffers.
- v10 LDS fixes (bank conflict padding, coalesced load, scale loading)
- Pre-allocate A_q, A_scale, ws in Python — eliminates torch::empty from hot path
"""
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


CPP_WRAPPER = """
void mxfp4_fused_gemm(
    torch::Tensor A,
    torch::Tensor B_shuf,
    torch::Tensor B_sc_shuf,
    torch::Tensor C,
    torch::Tensor A_q,
    torch::Tensor A_scale,
    torch::Tensor ws,
    int M, int N, int K
);
"""

CUDA_SRC = r"""
#include <hip/hip_runtime.h>
#include <torch/extension.h>
#include <stdint.h>

#define WARP_SIZE 64

static __device__ __forceinline__ float bf16_to_f32(uint16_t x) {
    uint32_t u = (uint32_t)x << 16;
    float f;
    __builtin_memcpy(&f, &u, 4);
    return f;
}

static __device__ __forceinline__ uint16_t f32_to_bf16(float x) {
    uint32_t u;
    __builtin_memcpy(&u, &x, 4);
    u += 0x7FFF + ((u >> 16) & 1u);
    return (uint16_t)(u >> 16);
}

static __device__ __forceinline__ uint8_t float_to_fp4_e2m1(float x) {
    uint32_t u;
    __builtin_memcpy(&u, &x, 4);
    uint32_t sign = (u >> 28) & 0x8u;
    uint32_t e = (u >> 23) & 0xFFu;
    uint32_t m = u & 0x7FFFFFu;
    if (e == 0) return (uint8_t)sign;
    if (e < 127u) {
        uint32_t adj = 126u - e;
        m = (adj < 23u) ? ((0x400000u | (m >> 1u)) >> adj) : 0u;
        e = 0;
    } else {
        e = e - 126u;
    }
    uint32_t combined = (e << 2) | (m >> 21);
    uint32_t e2m1 = (combined + 1u) >> 1;
    if (e2m1 > 7u) e2m1 = 7u;
    return (uint8_t)(sign | e2m1);
}

typedef int   __attribute__((ext_vector_type(8)))  i32x8_t;
typedef float __attribute__((ext_vector_type(4)))  f32x4_t;

// ===================== Kernel 1: Quantize A =====================
__global__ void __launch_bounds__(256)
quant_a_kernel(
    const uint16_t* __restrict__ A,
    uint8_t* __restrict__ A_q,
    uint8_t* __restrict__ A_scale,
    int M, int K
) {
    int gidx = blockIdx.x * 256 + threadIdx.x;
    int k_groups = K / 32;
    int total = M * k_groups;
    if (gidx >= total) return;

    int row = gidx / k_groups;
    int kg  = gidx % k_groups;

    const uint16_t* ap = A + (long)row * K + kg * 32;
    int4 v0, v1, v2, v3;
    __builtin_memcpy(&v0, ap,      16);
    __builtin_memcpy(&v1, ap + 8,  16);
    __builtin_memcpy(&v2, ap + 16, 16);
    __builtin_memcpy(&v3, ap + 24, 16);
    const uint16_t* s0 = (const uint16_t*)&v0;
    const uint16_t* s1 = (const uint16_t*)&v1;
    const uint16_t* s2 = (const uint16_t*)&v2;
    const uint16_t* s3 = (const uint16_t*)&v3;

    float vals[32];
    float mx = 0.0f;
    #pragma unroll
    for (int i = 0; i < 8; i++) {
        vals[i]    = bf16_to_f32(s0[i]);
        vals[i+8]  = bf16_to_f32(s1[i]);
        vals[i+16] = bf16_to_f32(s2[i]);
        vals[i+24] = bf16_to_f32(s3[i]);
    }
    #pragma unroll
    for (int i = 0; i < 32; i++) mx = fmaxf(mx, fabsf(vals[i]));

    uint8_t sc;
    float qs;
    if (mx == 0.0f) { sc = 0; qs = 0.0f; }
    else {
        uint32_t mx_u;
        __builtin_memcpy(&mx_u, &mx, 4);
        mx_u = (mx_u + 0x200000u) & 0xFF800000u;
        float amr;
        __builtin_memcpy(&amr, &mx_u, 4);
        float su = __builtin_floorf(__builtin_log2f(amr)) - 2.0f;
        su = fmaxf(-127.0f, fminf(127.0f, su));
        sc = (uint8_t)((int)su + 127);
        qs = __builtin_exp2f(-su);
    }

    uint8_t fp4[16];
    #pragma unroll
    for (int i = 0; i < 16; i++) {
        uint8_t lo = float_to_fp4_e2m1(vals[2*i]   * qs);
        uint8_t hi = float_to_fp4_e2m1(vals[2*i+1] * qs);
        fp4[i] = (lo & 0xF) | ((hi & 0xF) << 4);
    }
    __builtin_memcpy(A_q + (long)row * (K / 2) + kg * 16, fp4, 16);
    A_scale[row * k_groups + kg] = sc;
}

// ===================== Kernel 2: Register-only GEMM =====================
template<int WM, int WN, bool DIRECT_BF16>
__global__ void __launch_bounds__(WM * WN * WARP_SIZE)
mxfp4_gemm_reg(
    const uint8_t*  __restrict__ A_q,
    const uint8_t*  __restrict__ A_scale,
    const uint8_t*  __restrict__ B_shuf,
    const uint8_t*  __restrict__ B_sc_shuf,
    void*           __restrict__ C_out,
    int M, int N, int K,
    int sB, int sSC,
    int k_per_split
) {
    const int tid = threadIdx.x;
    const int wid = tid / WARP_SIZE;
    const int lid = tid % WARP_SIZE;
    const int wm  = wid / WN;
    const int wn  = wid % WN;
    const int tile_m = blockIdx.y * (WM * 16) + wm * 16;
    const int tile_n = blockIdx.x * (WN * 16) + wn * 16;
    if (tile_m >= M || tile_n >= N) return;

    const int a_row = tile_m + (lid & 15);
    const int b_col = tile_n + (lid & 15);
    const int kg    = lid >> 4;
    const int b_rg  = b_col >> 4;
    const int b_i2  = b_col & 15;
    const int b_i0  = b_col >> 5;
    const int b_i1  = (b_col >> 4) & 1;
    const long b_base = (long)b_rg * (sB * 16);
    const bool a_ok = (a_row < M);
    const bool b_ok = (b_col < N);
    const int split_id = blockIdx.z;
    const int k_start  = split_id * k_per_split;
    const int k_end    = min(k_start + k_per_split, K);
    const int half_K = K / 2;
    const int sc_K   = K / 32;

    f32x4_t acc = {0.f, 0.f, 0.f, 0.f};
    for (int kb = k_start; kb < k_end; kb += 128) {
        i32x8_t a_frag = {0,0,0,0,0,0,0,0};
        int a_sv = 0;
        if (a_ok) {
            const uint8_t* ap = A_q + (long)a_row * half_K + (kb >> 1) + (kg << 4);
            int4 tmp; __builtin_memcpy(&tmp, ap, 16);
            a_frag[0] = tmp.x; a_frag[1] = tmp.y;
            a_frag[2] = tmp.z; a_frag[3] = tmp.w;
            a_sv = (int)A_scale[a_row * sc_K + (kb >> 5) + kg];
        }
        i32x8_t b_frag = {0,0,0,0,0,0,0,0};
        int b_sv = 0;
        if (b_ok) {
            int bk = (kb >> 1) + (kg << 4);
            int i3 = bk >> 5; int i4 = (bk >> 4) & 1;
            const uint8_t* bp = B_shuf + b_base + i3 * 512 + i4 * 256 + b_i2 * 16;
            int4 tmp; __builtin_memcpy(&tmp, bp, 16);
            b_frag[0] = tmp.x; b_frag[1] = tmp.y;
            b_frag[2] = tmp.z; b_frag[3] = tmp.w;
            int sg = (kb >> 5) + kg;
            b_sv = (int)B_sc_shuf[b_i0 * (sSC * 32) + (sg >> 3) * 256 + (sg & 3) * 64 + b_i2 * 4 + ((sg >> 2) & 1) * 2 + b_i1];
        }
#if defined(__gfx950__)
        acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
            a_frag, b_frag, acc, 4, 4, 0, a_sv, 0, b_sv);
#endif
    }
    if (b_ok) {
        if constexpr (DIRECT_BF16) {
            uint16_t* out = reinterpret_cast<uint16_t*>(C_out);
            #pragma unroll
            for (int r = 0; r < 4; r++) {
                int mr = tile_m + (lid >> 4) * 4 + r;
                if (mr < M) out[(long)mr * N + b_col] = f32_to_bf16(acc[r]);
            }
        } else {
            float* out = reinterpret_cast<float*>(C_out);
            long off = (long)split_id * M * N;
            #pragma unroll
            for (int r = 0; r < 4; r++) {
                int mr = tile_m + (lid >> 4) * 4 + r;
                if (mr < M) out[off + (long)mr * N + b_col] = acc[r];
            }
        }
    }
}

// ===================== Kernel 3: LDS GEMM with fixes =====================
// vs v8: (1) bank conflict padding, (2) coalesced A load, (3) scale loading bugfix
template<int BM, bool DIRECT_BF16>
__global__ void __launch_bounds__(256)
mxfp4_gemm_lds(
    const uint8_t*  __restrict__ A_q,
    const uint8_t*  __restrict__ A_scale,
    const uint8_t*  __restrict__ B_shuf,
    const uint8_t*  __restrict__ B_sc_shuf,
    void*           __restrict__ C_out,
    int M, int N, int K,
    int sB, int sSC,
    int k_per_split
) {
    constexpr int BN = 128;
    constexpr int BK_EXT = 512;
    constexpr int BK_HALF = BK_EXT / 2;        // 256 bytes per row (data)
    constexpr int BK_HALF_PAD = BK_HALF + 4;   // 260 bytes stride (bank conflict fix: 260/4=65, 65%32=1)
    constexpr int BK_SC = BK_EXT / 32;          // 16 scale groups per row
    constexpr int MXdl = BM / 16;
    constexpr int LOADS_PER_ROW = BK_HALF / 16; // 16

    // Double-buffered LDS with padded stride
    __shared__ uint8_t A_lds[2][BM * BK_HALF_PAD];
    __shared__ uint8_t A_sc_lds[2][BM * BK_SC];

    const int tid = threadIdx.x;
    const int wid = tid / WARP_SIZE;
    const int lid = tid % WARP_SIZE;

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

    const int split_id = blockIdx.z;
    const int k_start  = split_id * k_per_split;
    const int k_end    = min(k_start + k_per_split, K);
    const int half_K = K / 2;
    const int sc_K   = K / 32;

    // Each wave handles 32 N-cols, 2 MFMA tiles of 16 cols
    const int wave_n = tile_n + wid * 32;
    const int b_kg   = lid >> 4;

    int b_col[2];
    b_col[0] = wave_n + (lid & 15);
    b_col[1] = wave_n + 16 + (lid & 15);

    int b_i2[2], b_i0[2], b_i1[2];
    long b_base_addr[2];
    bool b_ok[2];
    #pragma unroll
    for (int nt = 0; nt < 2; nt++) {
        b_ok[nt]        = (b_col[nt] < N);
        b_i2[nt]        = b_col[nt] & 15;
        b_i0[nt]        = b_col[nt] >> 5;
        b_i1[nt]        = (b_col[nt] >> 4) & 1;
        b_base_addr[nt] = (long)(b_col[nt] >> 4) * (sB * 16);
    }

    // Accumulators
    f32x4_t acc[MXdl][2];
    #pragma unroll
    for (int mt = 0; mt < MXdl; mt++)
        for (int nt = 0; nt < 2; nt++)
            acc[mt][nt] = {0.f, 0.f, 0.f, 0.f};

    // === Cooperative A load with coalesced access + padded LDS ===
    constexpr int TOTAL_LOADS = BM * LOADS_PER_ROW;
    constexpr int TOTAL_SC = BM * BK_SC;

    auto load_a = [&](int buf, int sk) {
        int sk_half = sk >> 1;
        int sk_k32  = sk >> 5;

        // FIX 2: Coalesced loading — consecutive threads load consecutive 16B chunks
        for (int li = tid; li < TOTAL_LOADS; li += 256) {
            int row   = li / LOADS_PER_ROW;
            int col16 = li % LOADS_PER_ROW;
            int global_row = tile_m + row;
            int global_k_byte = sk_half + col16 * 16;
            int4 val = {0, 0, 0, 0};
            if (global_row < M && global_k_byte + 16 <= half_K) {
                __builtin_memcpy(&val, A_q + (long)global_row * half_K + global_k_byte, 16);
            }
            // FIX 1: Write to padded LDS stride (BK_HALF_PAD = 260, avoids bank conflicts)
            __builtin_memcpy(&A_lds[buf][row * BK_HALF_PAD + col16 * 16], &val, 16);
        }

        // FIX 3: Scale loading — loop covers ALL BM*BK_SC entries (was only 256 of 1024)
        for (int s = tid; s < TOTAL_SC; s += 256) {
            int row = s / BK_SC;
            int sg  = s % BK_SC;
            int global_row = tile_m + row;
            int global_sg  = sk_k32 + sg;
            uint8_t v = 0;
            if (global_row < M && global_sg < sc_K)
                v = A_scale[global_row * sc_K + global_sg];
            A_sc_lds[buf][s] = v;
        }
    };

    // === Prologue: load first super-K-step ===
    load_a(0, k_start);
    __syncthreads();

    int cur_buf = 0;
    for (int sk = k_start; sk < k_end; sk += BK_EXT) {
        int next_buf = cur_buf ^ 1;
        int sk_len = min(BK_EXT, k_end - sk);

        // Prefetch next super-K-step
        int next_sk = sk + BK_EXT;
        if (next_sk < k_end) {
            load_a(next_buf, next_sk);
        }

        // Inner K-loop: up to 4 iterations of 128 FP4 each
        for (int kk = 0; kk < sk_len; kk += 128) {
            int kb = sk + kk;

            // Load B from global (per-wave, 2 N-tiles)
            i32x8_t b_frag[2];
            int b_sv[2];
            #pragma unroll
            for (int nt = 0; nt < 2; nt++) {
                b_frag[nt] = {0,0,0,0,0,0,0,0};
                b_sv[nt] = 0;
                if (b_ok[nt]) {
                    int bk = (kb >> 1) + (b_kg << 4);
                    int i3 = bk >> 5;
                    int i4 = (bk >> 4) & 1;
                    const uint8_t* bp = B_shuf + b_base_addr[nt] + i3 * 512 + i4 * 256 + b_i2[nt] * 16;
                    int4 tmp;
                    __builtin_memcpy(&tmp, bp, 16);
                    b_frag[nt][0] = tmp.x; b_frag[nt][1] = tmp.y;
                    b_frag[nt][2] = tmp.z; b_frag[nt][3] = tmp.w;

                    int sg  = (kb >> 5) + b_kg;
                    int si3 = sg >> 3;
                    int si4 = (sg >> 2) & 1;
                    int si5 = sg & 3;
                    b_sv[nt] = (int)B_sc_shuf[b_i0[nt] * (sSC * 32) + si3 * 256 + si5 * 64 + b_i2[nt] * 4 + si4 * 2 + b_i1[nt]];
                }
            }

            // Compute MXdl × 2 MFMA tiles
            #pragma unroll
            for (int mt = 0; mt < MXdl; mt++) {
                // Read A from LDS (using padded stride BK_HALF_PAD)
                i32x8_t a_frag = {0,0,0,0,0,0,0,0};
                int a_sv = 0;
                {
                    int a_row_lds = mt * 16 + (lid & 15);
                    int a_k_off = (kk >> 1) + (lid >> 4) * 16;
                    if (tile_m + a_row_lds < M) {
                        const uint8_t* ap = &A_lds[cur_buf][a_row_lds * BK_HALF_PAD + a_k_off];
                        int4 tmp;
                        __builtin_memcpy(&tmp, ap, 16);
                        a_frag[0] = tmp.x; a_frag[1] = tmp.y;
                        a_frag[2] = tmp.z; a_frag[3] = tmp.w;
                        a_sv = (int)A_sc_lds[cur_buf][a_row_lds * BK_SC + (kk >> 5) + (lid >> 4)];
                    }
                }

                #pragma unroll
                for (int nt = 0; nt < 2; nt++) {
#if defined(__gfx950__)
                    acc[mt][nt] = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
                        a_frag, b_frag[nt], acc[mt][nt], 4, 4, 0, a_sv, 0, b_sv[nt]);
#endif
                }
            }
        }

        __syncthreads();
        cur_buf = next_buf;
    }

    // === Write output ===
    #pragma unroll
    for (int mt = 0; mt < MXdl; mt++) {
        #pragma unroll
        for (int nt = 0; nt < 2; nt++) {
            if (!b_ok[nt]) continue;
            if constexpr (DIRECT_BF16) {
                uint16_t* out = reinterpret_cast<uint16_t*>(C_out);
                #pragma unroll
                for (int r = 0; r < 4; r++) {
                    int mr = tile_m + mt * 16 + (lid >> 4) * 4 + r;
                    if (mr < M) out[(long)mr * N + b_col[nt]] = f32_to_bf16(acc[mt][nt][r]);
                }
            } else {
                float* out = reinterpret_cast<float*>(C_out);
                long off = (long)split_id * M * N;
                #pragma unroll
                for (int r = 0; r < 4; r++) {
                    int mr = tile_m + mt * 16 + (lid >> 4) * 4 + r;
                    if (mr < M) out[off + (long)mr * N + b_col[nt]] = acc[mt][nt][r];
                }
            }
        }
    }
}

// ===================== Kernel 4: Reduce + convert to bf16 =====================
__global__ void __launch_bounds__(256)
reduce_bf16(
    const float* __restrict__ C_f32,
    uint16_t* __restrict__ C,
    int M, int N, int P
) {
    int idx = blockIdx.x * 256 + threadIdx.x;
    if (idx >= M * N) return;
    float sum = C_f32[idx];
    long stride = (long)M * N;
    for (int p = 1; p < P; p++) sum += C_f32[p * stride + idx];
    uint32_t u;
    __builtin_memcpy(&u, &sum, 4);
    u += 0x7FFF + ((u >> 16) & 1u);
    C[idx] = (uint16_t)(u >> 16);
}

// ===================== Host dispatch =====================
void mxfp4_fused_gemm(
    torch::Tensor A, torch::Tensor B_shuf, torch::Tensor B_sc_shuf,
    torch::Tensor C,
    torch::Tensor A_q, torch::Tensor A_scale, torch::Tensor ws_buf,
    int M, int N, int K)
{
    auto* bs  = reinterpret_cast<const uint8_t*>(B_shuf.data_ptr());
    auto* bsc = reinterpret_cast<const uint8_t*>(B_sc_shuf.data_ptr());
    auto* c   = reinterpret_cast<uint16_t*>(C.data_ptr());
    int sB = K / 2, sSC = ((K / 32 + 7) / 8) * 8;

    // 1. Quant A (into pre-allocated buffers)
    auto* aq  = A_q.data_ptr<uint8_t>();
    auto* asc = A_scale.data_ptr<uint8_t>();
    {
        int total = M * (K / 32);
        int blocks = (total + 255) / 256;
        quant_a_kernel<<<blocks, 256>>>(
            reinterpret_cast<const uint16_t*>(A.data_ptr()),
            aq, asc, M, K);
    }

    // 2. GEMM — register-only for all shapes (higher occupancy, no LDS sync overhead)
    bool use_lds = false;

    if (!use_lds) {
        // === Register-only kernel (best for small M, small K) ===
        int WM_val, BM, BN;
        if (M <= 16) { WM_val = 1; BM = 16; BN = 64; }
        else         { WM_val = 2; BM = 32; BN = 64; }
        const int WN_val = 4;

        int grid_x = (N + BN - 1) / BN;
        int grid_y = (M + BM - 1) / BM;
        int grid_mn = grid_x * grid_y;

        int P = 1, max_splits = K / 128;
        while (grid_mn * P < 256 && P * 2 <= max_splits) P *= 2;
        int k_per_split = ((K / P + 127) / 128) * 128;

        if (P == 1) {
            dim3 grid(grid_x, grid_y, 1);
            if (M <= 16) mxfp4_gemm_reg<1,4,true><<<grid, 4*WARP_SIZE>>>(aq,asc,bs,bsc,(void*)c,M,N,K,sB,sSC,K);
            else         mxfp4_gemm_reg<2,4,true><<<grid, 8*WARP_SIZE>>>(aq,asc,bs,bsc,(void*)c,M,N,K,sB,sSC,K);
        } else {
            float* ws = ws_buf.data_ptr<float>();
            dim3 grid(grid_x, grid_y, P);
            if (M <= 16) mxfp4_gemm_reg<1,4,false><<<grid, 4*WARP_SIZE>>>(aq,asc,bs,bsc,(void*)ws,M,N,K,sB,sSC,k_per_split);
            else         mxfp4_gemm_reg<2,4,false><<<grid, 8*WARP_SIZE>>>(aq,asc,bs,bsc,(void*)ws,M,N,K,sB,sSC,k_per_split);
            int rblocks = (M * N + 255) / 256;
            reduce_bf16<<<rblocks, 256>>>(ws, c, M, N, P);
        }
    } else {
        // === LDS kernel (best for larger M or long K) ===
        int BM;
        if      (M <= 16) BM = 16;
        else if (M <= 32) BM = 32;
        else              BM = 64;

        int grid_x = (N + 127) / 128;
        int grid_y = (M + BM - 1) / BM;
        int grid_mn = grid_x * grid_y;

        // SplitK: target good wave occupancy
        int P = 1, max_splits = K / 128;
        int total_waves = grid_mn * 4;
        if (total_waves < 384) {
            while (grid_mn * P < 192 && P * 2 <= max_splits) P *= 2;
        }
        int k_per_split = ((K / P + 127) / 128) * 128;

        #define LAUNCH_LDS_K(BM_V, DIRECT, DST) do { \
            dim3 grid(grid_x, grid_y, (DIRECT) ? 1 : P); \
            mxfp4_gemm_lds<BM_V, DIRECT><<<grid, 256>>>( \
                aq, asc, bs, bsc, (void*)(DST), M, N, K, sB, sSC, \
                (DIRECT) ? K : k_per_split); \
        } while(0)

        if (P == 1) {
            if      (BM == 16) { LAUNCH_LDS_K(16, true, c); }
            else if (BM == 32) { LAUNCH_LDS_K(32, true, c); }
            else               { LAUNCH_LDS_K(64, true, c); }
        } else {
            float* ws = ws_buf.data_ptr<float>();
            if      (BM == 16) { LAUNCH_LDS_K(16, false, ws); }
            else if (BM == 32) { LAUNCH_LDS_K(32, false, ws); }
            else               { LAUNCH_LDS_K(64, false, ws); }
            int rblocks = (M * N + 255) / 256;
            reduce_bf16<<<rblocks, 256>>>(ws, c, M, N, P);
        }
        #undef LAUNCH_LDS_K
    }
}
"""

_module = load_inline(
    name="mxfp4_hip_v11",
    cpp_sources=[CPP_WRAPPER],
    cuda_sources=[CUDA_SRC],
    functions=["mxfp4_fused_gemm"],
    verbose=True,
    extra_cuda_cflags=["--offload-arch=gfx950", "-O3", "-std=c++17"],
)

# Pre-allocated buffer cache: (m, n, k) → (A_q, A_scale, ws, C)
_buf_cache = {}


def _get_buffers(m, n, k, device):
    key = (m, n, k)
    if key in _buf_cache:
        return _buf_cache[key]

    A_q = torch.empty((m, k // 2), dtype=torch.uint8, device=device)
    A_scale = torch.empty((m, k // 32), dtype=torch.uint8, device=device)
    C = torch.empty((m, n), dtype=torch.bfloat16, device=device)

    # Compute P to know ws size
    use_lds = (m >= 32 and k >= 1024) or m >= 64
    if not use_lds:
        BN = 64
        BM = 16 if m <= 16 else 32
        grid_mn = ((n + BN - 1) // BN) * ((m + BM - 1) // BM)
        P = 1
        max_splits = k // 128
        while grid_mn * P < 256 and P * 2 <= max_splits:
            P *= 2
    else:
        BM = 16 if m <= 16 else (32 if m <= 32 else 64)
        grid_mn = ((n + 127) // 128) * ((m + BM - 1) // BM)
        P = 1
        max_splits = k // 128
        total_waves = grid_mn * 4
        if total_waves < 384:
            while grid_mn * P < 192 and P * 2 <= max_splits:
                P *= 2

    if P > 1:
        ws = torch.empty((P, m, n), dtype=torch.float32, device=device)
    else:
        ws = torch.empty(1, dtype=torch.float32, device=device)  # dummy

    _buf_cache[key] = (A_q, A_scale, ws, C)
    return A_q, A_scale, ws, C


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

    A_q, A_scale, ws, C = _get_buffers(m, n, k, A.device)
    _module.mxfp4_fused_gemm(A, B_shuffle, B_scale_sh, C, A_q, A_scale, ws, m, n, k)
    return C
scrolls · 602 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