Skip to content
KernelIndex
Search⌘K

submission 544703

div22 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

solution_new_21.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-544703?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
13.8µs
#476 of 1143
2026-03-13

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:0ff26efd1a5091b68bc47ad72cd03fc179ffc4696349d7c4cdb67f7d92932f07
license declaredunknown
license concludedunknown
authorsdiv22
imported2026-08-15

Techniques

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

fp4MXFP4 GEMM v21 — gfx950 (MI355X) optimized.
split-k3. Improved splitK: only split K=7168 with small M (from v20)
tile-m = 0const int full_m = (BM >= 16) ? M / BM : 0;

Kernel source

solution_new_21.py533 lines
"""
MXFP4 GEMM v21 — gfx950 (MI355X) optimized.

Changes from v20:
  1. Hardware FP4 conversion via v_cvt_scalef32_pk_fp4_bf16
     - Replaces software f32_to_fp4_e2m1 with single instruction
     - Fixed scale semantics: pass bit_cast<float>(inv_exp << 23)
       where inv_exp = amax_exp - 2, so hardware applies 2^(129-amax_exp)
  2. Template GEMM on CKT for constexpr loop unrolling (from v20)
  3. Improved splitK: only split K=7168 with small M (from v20)
"""
import os
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"

from typing import Tuple
import torch
from torch.utils.cpp_extension import load_inline
import uuid


HIP_KERNEL = r"""
#include <hip/hip_runtime.h>
#include <stdint.h>

using int4_v   = int   __attribute__((ext_vector_type(4)));
using float4_v = float __attribute__((ext_vector_type(4)));
using bf16x2   = __bf16 __attribute__((ext_vector_type(2)));

static constexpr int FP4_E2M1 = 4;

__device__ float4_v __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
    int4_v a, int4_v b, float4_v c,
    int cbsz, int blgp, int op_sel_a, int scale_a, int op_sel_b, int scale_b
) __asm("llvm.amdgcn.mfma.scale.f32.16x16x128.f8f6f4.v4i32.v4i32");

__device__ __forceinline__ int4_v load16(const uint8_t* __restrict__ p) {
    return *reinterpret_cast<const int4_v*>(p);
}

__device__ __forceinline__ uint16_t float_to_bf16(float f) {
    bf16x2 v;
    v[0] = static_cast<__bf16>(f);
    uint16_t r;
    __builtin_memcpy(&r, &v, sizeof(r));
    return r;
}

// Hardware BF16x2 -> FP4x2 conversion using v_cvt_scalef32_pk_fp4_bf16
// The instruction extracts biased exponent E from scale float, applies 2^(127-E) compression.
// To get correct compression factor 2^(129 - amax_exp), pass scale with biased_exp = amax_exp - 2.
__device__ __forceinline__ uint8_t hw_bf16x2_to_fp4x2(uint32_t bf16_pair, float scale) {
    uint32_t result;
    asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
                 : "=v"(result) : "v"(bf16_pair), "v"(scale));
    return static_cast<uint8_t>(result & 0xFFu);
}

// Quant kernel: BF16 [M, K] -> FP4x2 [M, K//2] + E8M0 scale [M, K//32]
// Uses hardware v_cvt_scalef32_pk_fp4_bf16 for FP4 conversion
__global__ void __launch_bounds__(128, 4)
mxfp4_quant(
    const __bf16* __restrict__ A_bf16,
    uint8_t*      __restrict__ A_fp4,
    uint8_t*      __restrict__ A_scale,
    int M, int K)
{
    const int KS    = K / 32;
    const int K2    = K / 2;
    const int group = blockIdx.x * 128 + threadIdx.x;
    const int row   = group / KS;
    const int kg    = group % KS;

    if (row >= M) return;

    const auto* src = A_bf16 + (long)row * K + kg * 32;

    float absMax = 1e-10f;
    #pragma unroll
    for (int i = 0; i < 32; ++i) {
        float v = __builtin_elementwise_abs(static_cast<float>(src[i]));
        absMax = (v > absMax) ? v : absMax;
    }

    uint32_t u32     = __builtin_bit_cast(uint32_t, absMax);
    const uint32_t amax_exp = ((u32 + 0x200000u) >> 23) & 0xFFu;
    const uint32_t inv_exp  = (amax_exp >= 2u) ? (amax_exp - 2u) : 0u;
    A_scale[(long)row * KS + kg] = static_cast<uint8_t>(inv_exp);

    // Hardware scale: biased_exp = inv_exp = amax_exp - 2
    // Instruction applies 2^(127 - inv_exp) = 2^(129 - amax_exp) for compression — correct!
    const float hw_scale = __builtin_bit_cast(float, static_cast<uint32_t>(inv_exp) << 23);

    const uint32_t* src_u32 = reinterpret_cast<const uint32_t*>(src);
    auto* dst = reinterpret_cast<uint8_t*>(A_fp4 + (long)row * K2 + kg * 16);
    #pragma unroll
    for (int i = 0; i < 16; ++i) {
        // Each uint32_t holds a pair of bf16 values (native layout)
        dst[i] = hw_bf16x2_to_fp4x2(src_u32[i], hw_scale);
    }
}

template<bool ALWAYS_VALID>
__device__ __forceinline__ int4_v load_or_zero(bool rt_valid, const uint8_t* p) {
    if constexpr (ALWAYS_VALID) return load16(p);
    else                        return rt_valid ? load16(p) : int4_v{0,0,0,0};
}

// Main GEMM kernel — templated on CKT (ktiles_per_split) for constexpr loop unrolling
// CKT=0 means runtime loop count
template<int BM, int BN, int NWARPS, bool SPLITK, bool A_VALID, bool B_VALID, int CKT = 0>
__global__ void __launch_bounds__(NWARPS * 64, 2)
mxfp4_gemm(
    const uint8_t* __restrict__ A,
    const uint8_t* __restrict__ As,
    const uint8_t* __restrict__ Bsh,
    const uint8_t* __restrict__ Bssh,
    float*         __restrict__ C_partial,
    uint16_t*      __restrict__ C_final,
    int M, int N, int K, int scaleN,
    int total_ktiles, int ktiles_per_split,
    int tile_off_x, int tile_off_y)
{
    static_assert(BN % 16 == 0);
    constexpr auto WAVES_M = (BM + 15) / 16;
    constexpr auto WAVES_N = BN / 16;
    static_assert(WAVES_M * WAVES_N == NWARPS);

    const auto ks_idx = blockIdx.z;
    const auto lane   = threadIdx.x % 64;
    const auto wave   = threadIdx.x / 64;
    const auto wave_m = wave / WAVES_N;
    const auto wave_n = wave % WAVES_N;

    const auto tile_m = (blockIdx.y + tile_off_y) * BM + wave_m * 16;
    const auto tile_n = (blockIdx.x + tile_off_x) * BN + wave_n * 16;

    if (tile_m >= M || tile_n >= N) return;

    const int ks_start = ks_idx * ktiles_per_split;
    const int ks_end   = min(ks_start + ktiles_per_split, total_ktiles);

    const auto lrow = lane % 16;
    const auto kgrp = lane / 16;

    const auto gm = tile_m + lrow;
    const auto gn = tile_n + lrow;
    const auto K2 = K / 2;
    const auto KS = K / 32;

    const auto a_rt = A_VALID | (gm < M);
    const auto b_rt = B_VALID | (gn < N);

    const auto n_tile        = tile_n / 16;
    const auto bsh_lane_base = Bsh + (long)n_tile * (K / 64) * 512 + (long)lrow * 16;

    const uint8_t* a_row  = nullptr;
    const uint8_t* as_row = nullptr;
    if constexpr (A_VALID) {
        a_row  = A  + (long)gm * K2;
        as_row = As + (long)gm * KS;
    } else {
        if (a_rt) { a_row  = A  + (long)gm * K2;
                    as_row = As + (long)gm * KS; }
    }

    int bssh_base = 0;
    if constexpr (B_VALID) {
        bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64 + (gn >> 5) * (32 * scaleN);
    } else {
        if (b_rt) bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64 + (gn >> 5) * (32 * scaleN);
    }

    const auto k_half_off = (kgrp & 1) * 256;
    const auto k_blk_base = kgrp >> 1;

    static constexpr long a_kt_stride   = 64L;
    static constexpr long bsh_kt_stride = 1024L;

    const uint8_t* a_ptr    = nullptr;
    const uint8_t* bsh_ptr  = nullptr;
    const uint8_t* bssh_ptr = nullptr;

    if constexpr (A_VALID) {
        a_ptr = a_row + (long)ks_start * a_kt_stride + kgrp * 16;
    } else {
        if (a_rt) a_ptr = a_row + (long)ks_start * a_kt_stride + kgrp * 16;
    }
    if constexpr (B_VALID) {
        bsh_ptr  = bsh_lane_base + (long)(ks_start * 2 + k_blk_base) * 512 + k_half_off;
        bssh_ptr = Bssh + bssh_base + (ks_start & 1) * 2 + (ks_start >> 1) * 256;
    } else {
        if (b_rt) {
            bsh_ptr  = bsh_lane_base + (long)(ks_start * 2 + k_blk_base) * 512 + k_half_off;
            bssh_ptr = Bssh + bssh_base + (ks_start & 1) * 2 + (ks_start >> 1) * 256;
        }
    }

    float4_v acc{0.f, 0.f, 0.f, 0.f};

    // bssh step pattern depends on ks_start parity; doesn't change across quads
    const auto bssh_step0 = (ks_start & 1) ? 254 : 2;
    const auto bssh_step1 = 256 - bssh_step0;

    // Macro for one MFMA iteration
    #define DO_MFMA(a_off, b_off, bssh_off, ks_val) \
    { \
        const auto av = load_or_zero<A_VALID>(a_rt, a_ptr + (a_off) * a_kt_stride); \
        const auto bv = load_or_zero<B_VALID>(b_rt, bsh_ptr + (b_off) * bsh_kt_stride); \
        const auto ks = (ks_val) * 4 + kgrp; \
        int sa, sb; \
        if constexpr (A_VALID) sa = static_cast<int>(as_row[ks]); \
        else                   sa = (a_rt & (ks < KS)) ? static_cast<int>(as_row[ks]) : 127; \
        if constexpr (B_VALID) sb = static_cast<int>(*(bssh_ptr + (bssh_off))); \
        else                   sb = (b_rt & (ks < KS)) ? static_cast<int>(*(bssh_ptr + (bssh_off))) : 127; \
        acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(av,bv,acc,FP4_E2M1,FP4_E2M1,0,sa,0,sb); \
    }

    // 4x unrolled main loop
    if constexpr (CKT > 0) {
        // Compile-time unrolled
        #pragma unroll
        for (int q = 0; q < (CKT / 4); ++q) {
            DO_MFMA(0, 0, 0, ks_start + q*4)
            DO_MFMA(1, 1, bssh_step0, ks_start + q*4 + 1)
            DO_MFMA(2, 2, bssh_step0 + bssh_step1, ks_start + q*4 + 2)
            DO_MFMA(3, 3, bssh_step0 + bssh_step1 + bssh_step0, ks_start + q*4 + 3)
            a_ptr    += 4 * a_kt_stride;
            bsh_ptr  += 4 * bsh_kt_stride;
            bssh_ptr += 512;
        }
        // Compile-time pair (parity same as bssh_step0 since quads advance by 4)
        if constexpr ((CKT % 4) >= 2) {
            DO_MFMA(0, 0, 0, ks_start + (CKT/4)*4)
            DO_MFMA(1, 1, bssh_step0, ks_start + (CKT/4)*4 + 1)
            a_ptr    += 2 * a_kt_stride;
            bsh_ptr  += 2 * bsh_kt_stride;
            bssh_ptr += 256;
        }
        // Compile-time single
        if constexpr ((CKT % 2) == 1) {
            DO_MFMA(0, 0, 0, ks_start + CKT - 1)
        }
    } else {
        // Runtime loop (fallback for unknown shapes)
        auto kt = ks_start;
        const int ks_count = ks_end - ks_start;
        const auto ks_end_quad = ks_start + (ks_count - (ks_count & 3));
        const auto ks_end_pair = ks_start + (ks_count - (ks_count & 1));

        for (; kt < ks_end_quad; kt += 4) {
            DO_MFMA(0, 0, 0, kt)
            DO_MFMA(1, 1, bssh_step0, kt + 1)
            DO_MFMA(2, 2, bssh_step0 + bssh_step1, kt + 2)
            DO_MFMA(3, 3, bssh_step0 + bssh_step1 + bssh_step0, kt + 3)
            a_ptr    += 4 * a_kt_stride;
            bsh_ptr  += 4 * bsh_kt_stride;
            bssh_ptr += 512;
        }
        const auto bssh_step0_trail = (kt & 1) ? 254 : 2;
        for (; kt < ks_end_pair; kt += 2) {
            DO_MFMA(0, 0, 0, kt)
            DO_MFMA(1, 1, bssh_step0_trail, kt + 1)
            a_ptr    += 2 * a_kt_stride;
            bsh_ptr  += 2 * bsh_kt_stride;
            bssh_ptr += 256;
        }
        if (kt < ks_end) {
            DO_MFMA(0, 0, 0, kt)
        }
    }
    #undef DO_MFMA

    const auto out_col      = tile_n + lrow;
    const auto out_row_base = tile_m + kgrp * 4;

    if constexpr (!B_VALID) { if (out_col >= N) return; }

    constexpr bool out_rows_always_valid = A_VALID && (BM >= 16);

    if constexpr (SPLITK) {
        auto c_out = C_partial + (long)ks_idx * M * N + (long)out_row_base * N + out_col;
        #pragma unroll
        for (int i = 0; i < 4; ++i) {
            if constexpr (out_rows_always_valid) c_out[i * N] = acc[i];
            else if (out_row_base + i < M)       c_out[i * N] = acc[i];
        }
    } else {
        auto c_out = C_final + (long)out_row_base * N + out_col;
        #pragma unroll
        for (int i = 0; i < 4; ++i) {
            if constexpr (out_rows_always_valid) c_out[i * N] = float_to_bf16(acc[i]);
            else if (out_row_base + i < M)       c_out[i * N] = float_to_bf16(acc[i]);
        }
    }
}

// Reduce kernel
__global__ void mxfp4_reduce(
    const float*  __restrict__ C_partial,
    uint16_t*     __restrict__ C_out,
    int M, int N, int NUM_KSPLIT)
{
    const auto col = blockIdx.x * 32 + threadIdx.x;
    const auto row = blockIdx.y * 16 + threadIdx.y;
    if (row >= M || col >= N) return;

    float sum = 0.f;
    const auto mn     = (long)row * N + col;
    const auto stride = (long)M * N;
    for (auto k = 0; k < NUM_KSPLIT; ++k)
        sum += C_partial[k * stride + mn];

    bf16x2 v;
    v[0] = static_cast<__bf16>(sum);
    uint16_t r;
    __builtin_memcpy(&r, &v, sizeof(r));
    C_out[mn] = r;
}

// Launch quant
extern "C" void launch_quant(
    const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale, int M, int K)
{
    const int KS       = K / 32;
    const int n_groups = M * KS;
    const dim3 block{128};
    const dim3 grid{static_cast<uint32_t>((n_groups + 127) / 128)};
    mxfp4_quant<<<grid, block>>>(A_bf16, A_fp4, A_scale, M, K);
}

// Templated GEMM launcher
template<int CKT>
void launch_gemm_ckt(
    const uint8_t* A, const uint8_t* As,
    const uint8_t* Bsh, const uint8_t* Bssh,
    float* C_partial, uint16_t* C_final,
    int M, int N, int K, int scaleN, int NUM_KSPLIT)
{
    const auto total_ktiles     = K / 128;
    const auto ktiles_per_split = (total_ktiles + NUM_KSPLIT - 1) / NUM_KSPLIT;
    const auto do_splitk        = NUM_KSPLIT > 1;

    auto launch = [&]<int BM, int BN, int NWARPS>() {
        static_assert(((BM + 15) / 16) * (BN / 16) == NWARPS);

        const int full_m  = (BM >= 16) ? M / BM : 0;
        const int full_n  = N / BN;
        const int total_m = (M + BM - 1) / BM;
        const int total_n = (N + BN - 1) / BN;
        const int edge_m  = total_m - full_m;
        const int edge_n  = total_n - full_n;

        const dim3 block{static_cast<uint32_t>(NWARPS * 64)};

        auto sub = [&]<bool AV, bool BV>(int gx, int gy, int ox, int oy) {
            if (gx <= 0 || gy <= 0) return;
            const dim3 grid{
                static_cast<uint32_t>(gx),
                static_cast<uint32_t>(gy),
                static_cast<uint32_t>(NUM_KSPLIT)
            };
            if (do_splitk)
                mxfp4_gemm<BM,BN,NWARPS,true,AV,BV,CKT><<<grid,block>>>(
                    A,As,Bsh,Bssh,C_partial,nullptr,
                    M,N,K,scaleN,total_ktiles,ktiles_per_split,ox,oy);
            else
                mxfp4_gemm<BM,BN,NWARPS,false,AV,BV,CKT><<<grid,block>>>(
                    A,As,Bsh,Bssh,nullptr,C_final,
                    M,N,K,scaleN,total_ktiles,ktiles_per_split,ox,oy);
        };

        sub.template operator()<true,  true >(full_n,  full_m,  0,      0);
        sub.template operator()<true,  false>(edge_n,  full_m,  full_n, 0);
        sub.template operator()<false, true >(full_n,  edge_m,  0,      full_m);
        sub.template operator()<false, false>(edge_n,  edge_m,  full_n, full_m);

        if (do_splitk) {
            const dim3 rblock{32, 16};
            const dim3 rgrid{
                static_cast<uint32_t>((N + 31) / 32),
                static_cast<uint32_t>((M + 15) / 16)
            };
            mxfp4_reduce<<<rgrid, rblock>>>(C_partial, C_final, M, N, NUM_KSPLIT);
        }
    };

    if      (M <=  8) launch.template operator()<  8, 32,  2>();
    else if (M <= 16) launch.template operator()< 16, 32,  2>();
    else if (M <= 32) launch.template operator()< 16, 32,  2>();
    else if (M <= 64) launch.template operator()< 32, 32,  4>();
    else if (M <=128) launch.template operator()< 32, 32,  4>();
    else              launch.template operator()< 64, 32,  8>();
}

// Main dispatch — routes to CKT-specialized launcher
extern "C" void launch_gemm(
    const uint8_t* A, const uint8_t* As,
    const uint8_t* Bsh, const uint8_t* Bssh,
    float* C_partial, uint16_t* C_final,
    int M, int N, int K, int scaleN, int NUM_KSPLIT)
{
    const int total_ktiles     = K / 128;
    const int ktiles_per_split = (total_ktiles + NUM_KSPLIT - 1) / NUM_KSPLIT;

    switch (ktiles_per_split) {
        case  4: launch_gemm_ckt< 4>(A,As,Bsh,Bssh,C_partial,C_final,M,N,K,scaleN,NUM_KSPLIT); break;
        case  8: launch_gemm_ckt< 8>(A,As,Bsh,Bssh,C_partial,C_final,M,N,K,scaleN,NUM_KSPLIT); break;
        case 12: launch_gemm_ckt<12>(A,As,Bsh,Bssh,C_partial,C_final,M,N,K,scaleN,NUM_KSPLIT); break;
        case 16: launch_gemm_ckt<16>(A,As,Bsh,Bssh,C_partial,C_final,M,N,K,scaleN,NUM_KSPLIT); break;
        default: launch_gemm_ckt< 0>(A,As,Bsh,Bssh,C_partial,C_final,M,N,K,scaleN,NUM_KSPLIT); break;
    }
}
"""


CPP = r"""
#include <torch/extension.h>
#include <c10/core/DeviceGuard.h>

extern "C" void launch_quant(const __bf16*, uint8_t*, uint8_t*, int, int);
extern "C" void launch_gemm(const uint8_t*, const uint8_t*, const uint8_t*, const uint8_t*,
                             float*, uint16_t*, int, int, int, int, int);

static int get_num_ksplit(int M, int K) {
    // Only use splitK for K=7168+ (56+ ktiles) with small M
    // Matches AITER's tuned configs
    int total_ktiles = K / 128;
    if (total_ktiles < 28) return 1;  // K<=3456: never split
    if (M <= 8)  return 7;   // kps=8 (K=7168)
    if (M <= 16) return 14;  // kps=4 (K=7168)
    return 1;
}

struct Workspace {
    at::Tensor A_fp4;
    at::Tensor A_scale;
    at::Tensor C_partial;
    at::Tensor C;
    int64_t last_M = -1, last_K = -1, last_N = -1, last_ksplit = -1;

    void ensure(int M, int N, int K, int num_ksplit, const at::TensorOptions& opts) {
        if (M == last_M && K == last_K && N == last_N && num_ksplit == last_ksplit) return;
        int64_t KS = K / 32;
        A_fp4   = at::empty({(int64_t)M, (int64_t)(K / 2)}, opts.dtype(at::kByte));
        A_scale = at::empty({(int64_t)M, KS},                opts.dtype(at::kByte));
        C       = at::empty({(int64_t)M, (int64_t)N},        opts.dtype(at::kBFloat16));
        if (num_ksplit > 1)
            C_partial = at::empty({(int64_t)num_ksplit, (int64_t)M, (int64_t)N}, opts.dtype(at::kFloat));
        else
            C_partial = at::Tensor();
        last_M = M; last_K = K; last_N = N; last_ksplit = num_ksplit;
    }
};

static Workspace g_ws;

at::Tensor fwd(const at::Tensor& A,
               const at::Tensor& B_q,
               const at::Tensor& B_shuffle,
               const at::Tensor& B_scale_sh) {
    auto guard = at::DeviceGuard(A.device());

    at::Tensor A_bf16 = (A.scalar_type() == at::kBFloat16 && A.is_contiguous())
                        ? A : A.to(A.device(), at::kBFloat16, false, false,
                                   at::MemoryFormat::Contiguous);

    const int M = A_bf16.size(0);
    const int K = A_bf16.size(1);
    const int N = B_q.size(0);
    const int KS     = K / 32;
    const int scaleN = ((KS + 7) / 8) * 8;

    at::Tensor Bsh = B_shuffle.view(at::kByte);
    if (!Bsh.is_contiguous()) Bsh = Bsh.contiguous();
    at::Tensor Bssh = B_scale_sh.view(at::kByte);
    if (!Bssh.is_contiguous()) Bssh = Bssh.contiguous();

    const int num_ksplit = get_num_ksplit(M, K);

    g_ws.ensure(M, N, K, num_ksplit, A_bf16.options());

    launch_quant(
        reinterpret_cast<const __bf16*>(A_bf16.data_ptr<at::BFloat16>()),
        g_ws.A_fp4.data_ptr<uint8_t>(),
        g_ws.A_scale.data_ptr<uint8_t>(),
        M, K);

    float* c_partial_ptr = (num_ksplit > 1) ? g_ws.C_partial.data_ptr<float>() : nullptr;

    launch_gemm(
        g_ws.A_fp4.data_ptr<uint8_t>(), g_ws.A_scale.data_ptr<uint8_t>(),
        Bsh.data_ptr<uint8_t>(), Bssh.data_ptr<uint8_t>(),
        c_partial_ptr,
        reinterpret_cast<uint16_t*>(g_ws.C.data_ptr<at::BFloat16>()),
        M, N, K, scaleN, num_ksplit);

    return g_ws.C;
}
"""

_ext = load_inline(
    name=f"g_{uuid.uuid4().hex[:8]}",
    cpp_sources=[CPP],
    cuda_sources=[HIP_KERNEL],
    functions=["fwd"],
    with_cuda=True,
    extra_cflags=["-O3", "-std=c++20"],
    extra_cuda_cflags=[
        "-O3",
        "--offload-arch=gfx950",
        "-ffast-math",
        "-munsafe-fp-atomics",
        "-std=c++20",
        "-mllvm", "-amdgpu-early-inline-all=true",
        "-mllvm", "-amdgpu-function-calls=false",
        "-mwavefrontsize64",
        "-mcumode",
        "-mllvm", "--amdgpu-kernarg-preload-count=16",
        "-mllvm", "-enable-post-misched=0",
        "-mllvm", "--lsr-drop-solution=1",
        "-mllvm", "-amdgpu-coerce-illegal-types=1",
        "-fgpu-flush-denormals-to-zero",
        "-fno-offload-uniform-block",
    ],
    extra_ldflags=["-lamdhip64"],
)


def custom_kernel(data: Tuple[torch.Tensor, ...]) -> torch.Tensor:
    """MXFP4 GEMM v21: hw FP4 conversion (fixed scale) + constexpr loop unrolling."""
    A, _, B_q, B_shuffle, B_scale_sh = data
    return _ext.fwd(A.cuda(), B_q, B_shuffle, B_scale_sh)
scrolls · 533 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 522107.

"""
- FP4 quant + FP4 GEMM reference: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.
- Quant logic follows aiter op_tests/test_gemm_a4w4.py (get_triton_quant(QuantType.per_1x32)).
+ MXFP4 GEMM v21 — gfx950 (MI355X) optimized.
+
+ Changes from v20:
+ 1. Hardware FP4 conversion via v_cvt_scalef32_pk_fp4_bf16
+ - Replaces software f32_to_fp4_e2m1 with single instruction
+ - Fixed scale semantics: pass bit_cast<float>(inv_exp << 23)
+ where inv_exp = amax_exp - 2, so hardware applies 2^(129-amax_exp)
+ 2. Template GEMM on CKT for constexpr loop unrolling (from v20)
+ 3. Improved splitK: only split K=7168 with small M (from v20)
"""
- from task import input_t, output_t
+ import os
+ os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
+ from typing import Tuple
+ import torch
+ from torch.utils.cpp_extension import load_inline
+ import uuid
- def custom_kernel(data: input_t) -> output_t:
- """
- Reference: MXFP4 per-1x32 quant on A; B_shuffle, B_scale_sh from generate_input.
- gemm_a4w4 with bpreshuffle=True.
- """
- import aiter
- from aiter import QuantType, dtypes
- A, B, B_q, B_shuffle, B_scale_sh = data
- A = A.contiguous()
- B = B.contiguous()
- m, k = A.shape
- n, _ = B.shape
+ HIP_KERNEL = r"""
+ #include <hip/hip_runtime.h>
+ #include <stdint.h>
- quant_func = aiter.get_triton_quant(QuantType.per_1x32)
- A_q, A_scale_sh = quant_func(A, shuffle=True)
- out_gemm = aiter.gemm_a4w4(
- A_q,
- B_shuffle,
- A_scale_sh,
- B_scale_sh,
- dtype=dtypes.bf16,
- bpreshuffle=True,
- )
- return out_gemm
+ using int4_v = int __attribute__((ext_vector_type(4)));
+ using float4_v = float __attribute__((ext_vector_type(4)));
+ using bf16x2 = __bf16 __attribute__((ext_vector_type(2)));
+
+ static constexpr int FP4_E2M1 = 4;
+
+ __device__ float4_v __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
+ int4_v a, int4_v b, float4_v c,
+ int cbsz, int blgp, int op_sel_a, int scale_a, int op_sel_b, int scale_b
+ ) __asm("llvm.amdgcn.mfma.scale.f32.16x16x128.f8f6f4.v4i32.v4i32");
+
+ __device__ __forceinline__ int4_v load16(const uint8_t* __restrict__ p) {
+ return *reinterpret_cast<const int4_v*>(p);
+ }
+
+ __device__ __forceinline__ uint16_t float_to_bf16(float f) {
+ bf16x2 v;
+ v[0] = static_cast<__bf16>(f);
+ uint16_t r;
+ __builtin_memcpy(&r, &v, sizeof(r));
+ return r;
+ }
+
+ // Hardware BF16x2 -> FP4x2 conversion using v_cvt_scalef32_pk_fp4_bf16
+ // The instruction extracts biased exponent E from scale float, applies 2^(127-E) compression.
+ // To get correct compression factor 2^(129 - amax_exp), pass scale with biased_exp = amax_exp - 2.
+ __device__ __forceinline__ uint8_t hw_bf16x2_to_fp4x2(uint32_t bf16_pair, float scale) {
+ uint32_t result;
+ asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
+ : "=v"(result) : "v"(bf16_pair), "v"(scale));
+ return static_cast<uint8_t>(result & 0xFFu);
+ }
+
+ // Quant kernel: BF16 [M, K] -> FP4x2 [M, K//2] + E8M0 scale [M, K//32]
+ // Uses hardware v_cvt_scalef32_pk_fp4_bf16 for FP4 conversion
+ __global__ void __launch_bounds__(128, 4)
+ mxfp4_quant(
+ const __bf16* __restrict__ A_bf16,
+ uint8_t* __restrict__ A_fp4,
+ uint8_t* __restrict__ A_scale,
+ int M, int K)
+ {
+ const int KS = K / 32;
+ const int K2 = K / 2;
+ const int group = blockIdx.x * 128 + threadIdx.x;
+ const int row = group / KS;
+ const int kg = group % KS;
+
+ if (row >= M) return;
+
+ const auto* src = A_bf16 + (long)row * K + kg * 32;
+
+ float absMax = 1e-10f;
+ #pragma unroll
+ for (int i = 0; i < 32; ++i) {
+ float v = __builtin_elementwise_abs(static_cast<float>(src[i]));
+ absMax = (v > absMax) ? v : absMax;
+ }
+
+ uint32_t u32 = __builtin_bit_cast(uint32_t, absMax);
+ const uint32_t amax_exp = ((u32 + 0x200000u) >> 23) & 0xFFu;
+ const uint32_t inv_exp = (amax_exp >= 2u) ? (amax_exp - 2u) : 0u;
+ A_scale[(long)row * KS + kg] = static_cast<uint8_t>(inv_exp);
+
+ // Hardware scale: biased_exp = inv_exp = amax_exp - 2
+ // Instruction applies 2^(127 - inv_exp) = 2^(129 - amax_exp) for compression — correct!
+ const float hw_scale = __builtin_bit_cast(float, static_cast<uint32_t>(inv_exp) << 23);
+
+ const uint32_t* src_u32 = reinterpret_cast<const uint32_t*>(src);
+ auto* dst = reinterpret_cast<uint8_t*>(A_fp4 + (long)row * K2 + kg * 16);
+ #pragma unroll
+ for (int i = 0; i < 16; ++i) {
+ // Each uint32_t holds a pair of bf16 values (native layout)
+ dst[i] = hw_bf16x2_to_fp4x2(src_u32[i], hw_scale);
+ }
+ }
+
+ template<bool ALWAYS_VALID>
+ __device__ __forceinline__ int4_v load_or_zero(bool rt_valid, const uint8_t* p) {
+ if constexpr (ALWAYS_VALID) return load16(p);
+ else return rt_valid ? load16(p) : int4_v{0,0,0,0};
+ }
+
+ // Main GEMM kernel — templated on CKT (ktiles_per_split) for constexpr loop unrolling
+ // CKT=0 means runtime loop count
+ template<int BM, int BN, int NWARPS, bool SPLITK, bool A_VALID, bool B_VALID, int CKT = 0>
+ __global__ void __launch_bounds__(NWARPS * 64, 2)
+ mxfp4_gemm(
+ const uint8_t* __restrict__ A,
+ const uint8_t* __restrict__ As,
+ const uint8_t* __restrict__ Bsh,
+ const uint8_t* __restrict__ Bssh,
+ float* __restrict__ C_partial,
+ uint16_t* __restrict__ C_final,
+ int M, int N, int K, int scaleN,
+ int total_ktiles, int ktiles_per_split,
+ int tile_off_x, int tile_off_y)
+ {
+ static_assert(BN % 16 == 0);
+ constexpr auto WAVES_M = (BM + 15) / 16;
+ constexpr auto WAVES_N = BN / 16;
+ static_assert(WAVES_M * WAVES_N == NWARPS);
+
+ const auto ks_idx = blockIdx.z;
+ const auto lane = threadIdx.x % 64;
+ const auto wave = threadIdx.x / 64;
+ const auto wave_m = wave / WAVES_N;
+ const auto wave_n = wave % WAVES_N;
+
+ const auto tile_m = (blockIdx.y + tile_off_y) * BM + wave_m * 16;
+ const auto tile_n = (blockIdx.x + tile_off_x) * BN + wave_n * 16;
+
+ if (tile_m >= M || tile_n >= N) return;
+
+ const int ks_start = ks_idx * ktiles_per_split;
+ const int ks_end = min(ks_start + ktiles_per_split, total_ktiles);
+
+ const auto lrow = lane % 16;
+ const auto kgrp = lane / 16;
+
+ const auto gm = tile_m + lrow;
+ const auto gn = tile_n + lrow;
+ const auto K2 = K / 2;
+ const auto KS = K / 32;
+
+ const auto a_rt = A_VALID | (gm < M);
+ const auto b_rt = B_VALID | (gn < N);
+
+ const auto n_tile = tile_n / 16;
+ const auto bsh_lane_base = Bsh + (long)n_tile * (K / 64) * 512 + (long)lrow * 16;
+
+ const uint8_t* a_row = nullptr;
+ const uint8_t* as_row = nullptr;
+ if constexpr (A_VALID) {
+ a_row = A + (long)gm * K2;
+ as_row = As + (long)gm * KS;
+ } else {
+ if (a_rt) { a_row = A + (long)gm * K2;
+ as_row = As + (long)gm * KS; }
+ }
+
+ int bssh_base = 0;
+ if constexpr (B_VALID) {
+ bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64 + (gn >> 5) * (32 * scaleN);
+ } else {
+ if (b_rt) bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64 + (gn >> 5) * (32 * scaleN);
+ }
+
+ const auto k_half_off = (kgrp & 1) * 256;
+ const auto k_blk_base = kgrp >> 1;
+
+ static constexpr long a_kt_stride = 64L;
+ static constexpr long bsh_kt_stride = 1024L;
+
+ const uint8_t* a_ptr = nullptr;
+ const uint8_t* bsh_ptr = nullptr;
+ const uint8_t* bssh_ptr = nullptr;
+
+ if constexpr (A_VALID) {
+ a_ptr = a_row + (long)ks_start * a_kt_stride + kgrp * 16;
+ } else {
+ if (a_rt) a_ptr = a_row + (long)ks_start * a_kt_stride + kgrp * 16;
+ }
+ if constexpr (B_VALID) {
+ bsh_ptr = bsh_lane_base + (long)(ks_start * 2 + k_blk_base) * 512 + k_half_off;
+ bssh_ptr = Bssh + bssh_base + (ks_start & 1) * 2 + (ks_start >> 1) * 256;
+ } else {
+ if (b_rt) {
+ bsh_ptr = bsh_lane_base + (long)(ks_start * 2 + k_blk_base) * 512 + k_half_off;
+ bssh_ptr = Bssh + bssh_base + (ks_start & 1) * 2 + (ks_start >> 1) * 256;
+ }
+ }
+
+ float4_v acc{0.f, 0.f, 0.f, 0.f};
+
+ // bssh step pattern depends on ks_start parity; doesn't change across quads
+ const auto bssh_step0 = (ks_start & 1) ? 254 : 2;
+ const auto bssh_step1 = 256 - bssh_step0;
+
+ // Macro for one MFMA iteration
+ #define DO_MFMA(a_off, b_off, bssh_off, ks_val) \
+ { \
+ const auto av = load_or_zero<A_VALID>(a_rt, a_ptr + (a_off) * a_kt_stride); \
+ const auto bv = load_or_zero<B_VALID>(b_rt, bsh_ptr + (b_off) * bsh_kt_stride); \
+ const auto ks = (ks_val) * 4 + kgrp; \
+ int sa, sb; \
+ if constexpr (A_VALID) sa = static_cast<int>(as_row[ks]); \
+ else sa = (a_rt & (ks < KS)) ? static_cast<int>(as_row[ks]) : 127; \
+ if constexpr (B_VALID) sb = static_cast<int>(*(bssh_ptr + (bssh_off))); \
+ else sb = (b_rt & (ks < KS)) ? static_cast<int>(*(bssh_ptr + (bssh_off))) : 127; \
+ acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(av,bv,acc,FP4_E2M1,FP4_E2M1,0,sa,0,sb); \
+ }
+
+ // 4x unrolled main loop
+ if constexpr (CKT > 0) {
+ // Compile-time unrolled
+ #pragma unroll
+ for (int q = 0; q < (CKT / 4); ++q) {
+ DO_MFMA(0, 0, 0, ks_start + q*4)
+ DO_MFMA(1, 1, bssh_step0, ks_start + q*4 + 1)
+ DO_MFMA(2, 2, bssh_step0 + bssh_step1, ks_start + q*4 + 2)
+ DO_MFMA(3, 3, bssh_step0 + bssh_step1 + bssh_step0, ks_start + q*4 + 3)
+ a_ptr += 4 * a_kt_stride;
+ bsh_ptr += 4 * bsh_kt_stride;
+ bssh_ptr += 512;
+ }
+ // Compile-time pair (parity same as bssh_step0 since quads advance by 4)
+ if constexpr ((CKT % 4) >= 2) {
+ DO_MFMA(0, 0, 0, ks_start + (CKT/4)*4)
+ DO_MFMA(1, 1, bssh_step0, ks_start + (CKT/4)*4 + 1)
+ a_ptr += 2 * a_kt_stride;
+ bsh_ptr += 2 * bsh_kt_stride;
+ bssh_ptr += 256;
+ }
+ // Compile-time single
+ if constexpr ((CKT % 2) == 1) {
+ DO_MFMA(0, 0, 0, ks_start + CKT - 1)
+ }
+ } else {
+ // Runtime loop (fallback for unknown shapes)
+ auto kt = ks_start;
+ const int ks_count = ks_end - ks_start;
+ const auto ks_end_quad = ks_start + (ks_count - (ks_count & 3));
+ const auto ks_end_pair = ks_start + (ks_count - (ks_count & 1));
+
+ for (; kt < ks_end_quad; kt += 4) {
+ DO_MFMA(0, 0, 0, kt)
+ DO_MFMA(1, 1, bssh_step0, kt + 1)
+ DO_MFMA(2, 2, bssh_step0 + bssh_step1, kt + 2)
+ DO_MFMA(3, 3, bssh_step0 + bssh_step1 + bssh_step0, kt + 3)
+ a_ptr += 4 * a_kt_stride;
+ bsh_ptr += 4 * bsh_kt_stride;
+ bssh_ptr += 512;
+ }
+ const auto bssh_step0_trail = (kt & 1) ? 254 : 2;
+ for (; kt < ks_end_pair; kt += 2) {
+ DO_MFMA(0, 0, 0, kt)
+ DO_MFMA(1, 1, bssh_step0_trail, kt + 1)
+ a_ptr += 2 * a_kt_stride;
+ bsh_ptr += 2 * bsh_kt_stride;
+ bssh_ptr += 256;
+ }
+ if (kt < ks_end) {
+ DO_MFMA(0, 0, 0, kt)
+ }
+ }
+ #undef DO_MFMA
+
+ const auto out_col = tile_n + lrow;
+ const auto out_row_base = tile_m + kgrp * 4;
+
+ if constexpr (!B_VALID) { if (out_col >= N) return; }
+
+ constexpr bool out_rows_always_valid = A_VALID && (BM >= 16);
+
+ if constexpr (SPLITK) {
+ auto c_out = C_partial + (long)ks_idx * M * N + (long)out_row_base * N + out_col;
+ #pragma unroll
+ for (int i = 0; i < 4; ++i) {
+ if constexpr (out_rows_always_valid) c_out[i * N] = acc[i];
+ else if (out_row_base + i < M) c_out[i * N] = acc[i];
+ }
+ } else {
+ auto c_out = C_final + (long)out_row_base * N + out_col;
+ #pragma unroll
+ for (int i = 0; i < 4; ++i) {
+ if constexpr (out_rows_always_valid) c_out[i * N] = float_to_bf16(acc[i]);
+ else if (out_row_base + i < M) c_out[i * N] = float_to_bf16(acc[i]);
+ }
+ }
+ }
+
+ // Reduce kernel
+ __global__ void mxfp4_reduce(
+ const float* __restrict__ C_partial,
+ uint16_t* __restrict__ C_out,
+ int M, int N, int NUM_KSPLIT)
+ {
+ const auto col = blockIdx.x * 32 + threadIdx.x;
+ const auto row = blockIdx.y * 16 + threadIdx.y;
+ if (row >= M || col >= N) return;
+
+ float sum = 0.f;
+ const auto mn = (long)row * N + col;
+ const auto stride = (long)M * N;
+ for (auto k = 0; k < NUM_KSPLIT; ++k)
+ sum += C_partial[k * stride + mn];
+
+ bf16x2 v;
+ v[0] = static_cast<__bf16>(sum);
+ uint16_t r;
+ __builtin_memcpy(&r, &v, sizeof(r));
+ C_out[mn] = r;
+ }
+
+ // Launch quant
+ extern "C" void launch_quant(
+ const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale, int M, int K)
+ {
+ const int KS = K / 32;
+ const int n_groups = M * KS;
+ const dim3 block{128};
+ const dim3 grid{static_cast<uint32_t>((n_groups + 127) / 128)};
+ mxfp4_quant<<<grid, block>>>(A_bf16, A_fp4, A_scale, M, K);
+ }
+
+ // Templated GEMM launcher
+ template<int CKT>
+ void launch_gemm_ckt(
+ const uint8_t* A, const uint8_t* As,
+ const uint8_t* Bsh, const uint8_t* Bssh,
+ float* C_partial, uint16_t* C_final,
+ int M, int N, int K, int scaleN, int NUM_KSPLIT)
+ {
+ const auto total_ktiles = K / 128;
+ const auto ktiles_per_split = (total_ktiles + NUM_KSPLIT - 1) / NUM_KSPLIT;
+ const auto do_splitk = NUM_KSPLIT > 1;
+
+ auto launch = [&]<int BM, int BN, int NWARPS>() {
+ static_assert(((BM + 15) / 16) * (BN / 16) == NWARPS);
+
+ const int full_m = (BM >= 16) ? M / BM : 0;
+ const int full_n = N / BN;
+ const int total_m = (M + BM - 1) / BM;
+ const int total_n = (N + BN - 1) / BN;
+ const int edge_m = total_m - full_m;
+ const int edge_n = total_n - full_n;
+
+ const dim3 block{static_cast<uint32_t>(NWARPS * 64)};
+
+ auto sub = [&]<bool AV, bool BV>(int gx, int gy, int ox, int oy) {
+ if (gx <= 0 || gy <= 0) return;
+ const dim3 grid{
+ static_cast<uint32_t>(gx),
+ static_cast<uint32_t>(gy),
+ static_cast<uint32_t>(NUM_KSPLIT)
+ };
+ if (do_splitk)
+ mxfp4_gemm<BM,BN,NWARPS,true,AV,BV,CKT><<<grid,block>>>(
+ A,As,Bsh,Bssh,C_partial,nullptr,
+ M,N,K,scaleN,total_ktiles,ktiles_per_split,ox,oy);
+ else
+ mxfp4_gemm<BM,BN,NWARPS,false,AV,BV,CKT><<<grid,block>>>(
+ A,As,Bsh,Bssh,nullptr,C_final,
+ M,N,K,scaleN,total_ktiles,ktiles_per_split,ox,oy);
+ };
+
+ sub.template operator()<true, true >(full_n, full_m, 0, 0);
+ sub.template operator()<true, false>(edge_n, full_m, full_n, 0);
+ sub.template operator()<false, true >(full_n, edge_m, 0, full_m);
+ sub.template operator()<false, false>(edge_n, edge_m, full_n, full_m);
+
+ if (do_splitk) {
+ const dim3 rblock{32, 16};
+ const dim3 rgrid{
+ static_cast<uint32_t>((N + 31) / 32),
+ static_cast<uint32_t>((M + 15) / 16)
+ };
+ mxfp4_reduce<<<rgrid, rblock>>>(C_partial, C_final, M, N, NUM_KSPLIT);
+ }
+ };
+
+ if (M <= 8) launch.template operator()< 8, 32, 2>();
+ else if (M <= 16) launch.template operator()< 16, 32, 2>();
+ else if (M <= 32) launch.template operator()< 16, 32, 2>();
+ else if (M <= 64) launch.template operator()< 32, 32, 4>();
+ else if (M <=128) launch.template operator()< 32, 32, 4>();
+ else launch.template operator()< 64, 32, 8>();
+ }
+
+ // Main dispatch — routes to CKT-specialized launcher
+ extern "C" void launch_gemm(
+ const uint8_t* A, const uint8_t* As,
+ const uint8_t* Bsh, const uint8_t* Bssh,
+ float* C_partial, uint16_t* C_final,
+ int M, int N, int K, int scaleN, int NUM_KSPLIT)
+ {
+ const int total_ktiles = K / 128;
+ const int ktiles_per_split = (total_ktiles + NUM_KSPLIT - 1) / NUM_KSPLIT;
+
+ switch (ktiles_per_split) {
+ case 4: launch_gemm_ckt< 4>(A,As,Bsh,Bssh,C_partial,C_final,M,N,K,scaleN,NUM_KSPLIT); break;
+ case 8: launch_gemm_ckt< 8>(A,As,Bsh,Bssh,C_partial,C_final,M,N,K,scaleN,NUM_KSPLIT); break;
+ case 12: launch_gemm_ckt<12>(A,As,Bsh,Bssh,C_partial,C_final,M,N,K,scaleN,NUM_KSPLIT); break;
+ case 16: launch_gemm_ckt<16>(A,As,Bsh,Bssh,C_partial,C_final,M,N,K,scaleN,NUM_KSPLIT); break;
+ default: launch_gemm_ckt< 0>(A,As,Bsh,Bssh,C_partial,C_final,M,N,K,scaleN,NUM_KSPLIT); break;
+ }
+ }
+ """
+
+
+ CPP = r"""
+ #include <torch/extension.h>
+ #include <c10/core/DeviceGuard.h>
+
+ extern "C" void launch_quant(const __bf16*, uint8_t*, uint8_t*, int, int);
+ extern "C" void launch_gemm(const uint8_t*, const uint8_t*, const uint8_t*, const uint8_t*,
+ float*, uint16_t*, int, int, int, int, int);
+
+ static int get_num_ksplit(int M, int K) {
+ // Only use splitK for K=7168+ (56+ ktiles) with small M
+ // Matches AITER's tuned configs
+ int total_ktiles = K / 128;
+ if (total_ktiles < 28) return 1; // K<=3456: never split
+ if (M <= 8) return 7; // kps=8 (K=7168)
+ if (M <= 16) return 14; // kps=4 (K=7168)
+ return 1;
+ }
+
+ struct Workspace {
+ at::Tensor A_fp4;
+ at::Tensor A_scale;
+ at::Tensor C_partial;
+ at::Tensor C;
+ int64_t last_M = -1, last_K = -1, last_N = -1, last_ksplit = -1;
+
+ void ensure(int M, int N, int K, int num_ksplit, const at::TensorOptions& opts) {
+ if (M == last_M && K == last_K && N == last_N && num_ksplit == last_ksplit) return;
+ int64_t KS = K / 32;
+ A_fp4 = at::empty({(int64_t)M, (int64_t)(K / 2)}, opts.dtype(at::kByte));
+ A_scale = at::empty({(int64_t)M, KS}, opts.dtype(at::kByte));
+ C = at::empty({(int64_t)M, (int64_t)N}, opts.dtype(at::kBFloat16));
+ if (num_ksplit > 1)
+ C_partial = at::empty({(int64_t)num_ksplit, (int64_t)M, (int64_t)N}, opts.dtype(at::kFloat));
+ else
+ C_partial = at::Tensor();
+ last_M = M; last_K = K; last_N = N; last_ksplit = num_ksplit;
+ }
+ };
+
+ static Workspace g_ws;
+
+ at::Tensor fwd(const at::Tensor& A,
+ const at::Tensor& B_q,
+ const at::Tensor& B_shuffle,
+ const at::Tensor& B_scale_sh) {
+ auto guard = at::DeviceGuard(A.device());
+
+ at::Tensor A_bf16 = (A.scalar_type() == at::kBFloat16 && A.is_contiguous())
+ ? A : A.to(A.device(), at::kBFloat16, false, false,
+ at::MemoryFormat::Contiguous);
+
+ const int M = A_bf16.size(0);
+ const int K = A_bf16.size(1);
+ const int N = B_q.size(0);
+ const int KS = K / 32;
+ const int scaleN = ((KS + 7) / 8) * 8;
+
+ at::Tensor Bsh = B_shuffle.view(at::kByte);
+ if (!Bsh.is_contiguous()) Bsh = Bsh.contiguous();
+ at::Tensor Bssh = B_scale_sh.view(at::kByte);
+ if (!Bssh.is_contiguous()) Bssh = Bssh.contiguous();
+
+ const int num_ksplit = get_num_ksplit(M, K);
+
+ g_ws.ensure(M, N, K, num_ksplit, A_bf16.options());
+
+ launch_quant(
+ reinterpret_cast<const __bf16*>(A_bf16.data_ptr<at::BFloat16>()),
+ g_ws.A_fp4.data_ptr<uint8_t>(),
+ g_ws.A_scale.data_ptr<uint8_t>(),
+ M, K);
+
+ float* c_partial_ptr = (num_ksplit > 1) ? g_ws.C_partial.data_ptr<float>() : nullptr;
+
+ launch_gemm(
+ g_ws.A_fp4.data_ptr<uint8_t>(), g_ws.A_scale.data_ptr<uint8_t>(),
+ Bsh.data_ptr<uint8_t>(), Bssh.data_ptr<uint8_t>(),
+ c_partial_ptr,
+ reinterpret_cast<uint16_t*>(g_ws.C.data_ptr<at::BFloat16>()),
+ M, N, K, scaleN, num_ksplit);
+
+ return g_ws.C;
+ }
+ """
+
+ _ext = load_inline(
+ name=f"g_{uuid.uuid4().hex[:8]}",
+ cpp_sources=[CPP],
+ cuda_sources=[HIP_KERNEL],
+ functions=["fwd"],
+ with_cuda=True,
+ extra_cflags=["-O3", "-std=c++20"],
+ extra_cuda_cflags=[
+ "-O3",
+ "--offload-arch=gfx950",
+ "-ffast-math",
+ "-munsafe-fp-atomics",
+ "-std=c++20",
+ "-mllvm", "-amdgpu-early-inline-all=true",
+ "-mllvm", "-amdgpu-function-calls=false",
+ "-mwavefrontsize64",
+ "-mcumode",
+ "-mllvm", "--amdgpu-kernarg-preload-count=16",
+ "-mllvm", "-enable-post-misched=0",
+ "-mllvm", "--lsr-drop-solution=1",
+ "-mllvm", "-amdgpu-coerce-illegal-types=1",
+ "-fgpu-flush-denormals-to-zero",
+ "-fno-offload-uniform-block",
+ ],
+ extra_ldflags=["-lamdhip64"],
+ )
+
+
+ def custom_kernel(data: Tuple[torch.Tensor, ...]) -> torch.Tensor:
+ """MXFP4 GEMM v21: hw FP4 conversion (fixed scale) + constexpr loop unrolling."""
+ A, _, B_q, B_shuffle, B_scale_sh = data
+ return _ext.fwd(A.cuda(), B_q, B_shuffle, B_scale_sh)
scrolls · 558 diff lines total

Best evidence level for this revision: reported

JSON