Skip to content
KernelIndex
Search⌘K

submission 651769

kfz · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

sub_v31_hip.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-651769?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.0µs
#401 of 1143
2026-03-28

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:159cf5fbcf42abd360737d9161260d361f7cbd9af94444bbc9ff260bb3fcd20b
license declaredunknown
license concludedunknown
authorskfz
imported2026-08-26

Techniques

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

shared-memoryextern __shared__ uint8_t smem[];
split-ktemplate<int WM, int WN, int GM, int GN, bool SPLIT_K, bool USE_HW_CVT, int PIPE_DEPTH>

Kernel source

sub_v31_hip.py542 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
# Variant 31: V29 pipeline depth + aggressive compiler flags + hardcoded best configs
# Compiler: -ffast-math -ffp-contract=fast -munsafe-fp-atomics
# Best configs from V29/V30 auto-tune (hardcoded, no warmup auto-tune overhead)

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
import sys


CUDA_SRC = r"""
#include <hip/hip_runtime.h>
#include <hip/hip_bf16.h>
#include <stdint.h>
#include <math.h>

#if defined(__gfx950__)
typedef int   i32x8_t  __attribute__((ext_vector_type(8)));
typedef float f32x4_t  __attribute__((ext_vector_type(4)));

__device__ __forceinline__ void async_global_load_lds_b16(
    const uint8_t* src, uint32_t m0_val)
{
    asm volatile(
        "\n\ts_mov_b32 m0, %1"
        "\n\tglobal_load_lds_dwordx4 %0, off"
        :
        : "v"(src), "s"(m0_val)
        : "memory"
    );
}

__device__ __forceinline__ void async_lds_fence() {
    asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
}

__device__ __forceinline__ unsigned int compute_e8m0(float amax) {
    union { float f; unsigned int u; } v;
    v.f = amax;
    v.u = (v.u + 0x200000u) & 0xFF800000u;
    if (v.u == 0u) return 127u;
    unsigned int biased_exp = (v.u >> 23) & 0xFFu;
    int e8m0 = (int)biased_exp - 2;
    return (unsigned int)max(0, min(254, e8m0));
}

__device__ __forceinline__ uint8_t float_to_fp4(float x) {
    union { float f; unsigned int u; } v;
    v.f = x;
    unsigned int sign = v.u & 0x80000000u;
    v.u ^= sign;
    float ax = v.f;
    uint8_t code;
    if (ax >= 6.0f) {
        code = 7u;
    } else if (ax >= 1.0f) {
        unsigned int mant_odd = (v.u >> 22) & 1u;
        v.u += 0xC11FFFFFu;
        v.u += mant_odd;
        code = (uint8_t)((v.u >> 22) & 0xFu);
    } else {
        union { float f; unsigned int u; } d;
        d.f = ax + 4194304.0f;
        code = (uint8_t)((d.u - 0x4A800000u) & 0xFFu);
    }
    return code | (uint8_t)(sign >> 28);
}
#endif

// ═══════════════════════════════════════════════════════════════
// Kernel with PIPE_DEPTH template parameter
// PIPE_DEPTH=2: standard double buffer (V26 behavior)
// PIPE_DEPTH=3: triple buffer
// PIPE_DEPTH=4: quad buffer (full prefetch for K_steps=4)
// ═══════════════════════════════════════════════════════════════
template<int WM, int WN, int GM, int GN, bool SPLIT_K, bool USE_HW_CVT, int PIPE_DEPTH>
__global__ void fp4_gemm_fused_kernel(
    const __hip_bfloat16* __restrict__ A_bf16,
    const uint8_t*  __restrict__ B_q,
    const uint8_t*  __restrict__ B_scale_sh,
    __hip_bfloat16* __restrict__ C_bf16,
    float*          __restrict__ C_f32,
    int M, int N, int K, int chunk_Ks
) {
#if defined(__gfx950__)
    constexpr int N_WARPS    = GM * GN;
    constexpr int BLOCK_MT   = WM * GM;
    constexpr int BLOCK_NT   = WN * GN;
    constexpr int BLOCK_M    = BLOCK_MT * 16;
    constexpr int BLOCK_N    = BLOCK_NT * 16;

    constexpr int B_ROW_BYTES  = 64;
    constexpr int B_TILE_BYTES = 16 * B_ROW_BYTES;
    constexpr int B_BUF        = BLOCK_NT * B_TILE_BYTES;

    constexpr int A_ROW_BYTES  = 256;
    constexpr int A_TILE_BYTES = 16 * A_ROW_BYTES;
    constexpr int A_BUF        = BLOCK_MT * A_TILE_BYTES;

    constexpr int S_CHUNK      = 256;
    constexpr int S_N_GROUPS   = (BLOCK_N + 31) / 32;
    constexpr int S_LDS_BUF    = 1024;

    // SMEM layout: PIPE_DEPTH buffers instead of 2
    constexpr int B_BASE = 0;
    constexpr int A_BASE = PIPE_DEPTH * B_BUF;
    constexpr int S_BASE = A_BASE + PIPE_DEPTH * A_BUF;

    extern __shared__ uint8_t smem[];

    const int tid      = threadIdx.x;
    const int warp_id  = tid / 64;
    const int lane     = tid % 64;
    const int lane16   = lane & 15;
    const int k_quarter = lane >> 4;

    const int warp_row = warp_id / GN;
    const int warp_col = warp_id % GN;

    const int m_block  = blockIdx.y * BLOCK_M;
    const int n_block  = blockIdx.x * BLOCK_N;
    const int m_warp   = m_block + warp_row * WM * 16;
    const int n_warp   = n_block + warp_col * WN * 16;
    const int n_block_i0 = n_block >> 5;

    const int half_K   = K >> 1;

    const int k_chunk   = SPLIT_K ? (int)blockIdx.z : 0;
    const int ks_base   = k_chunk * chunk_Ks;
    const int byte_base = ks_base * 16;
    const int K_steps   = chunk_Ks >> 2;

    f32x4_t acc[WM][WN];
    #pragma unroll
    for (int wm = 0; wm < WM; wm++)
        #pragma unroll
        for (int wn = 0; wn < WN; wn++)
            for (int i = 0; i < 4; i++)
                acc[wm][wn][i] = 0.0f;

    const int col_out = lane & 15;
    const int row_base_out = (lane >> 4) * 4;

    #define ASYNC_LOAD_B_TILES(BUF, KSTEP) do { \
        int _kbyte = byte_base + (KSTEP) * 64; \
        int _load_row = lane >> 2; \
        int _load_quarter = lane & 3; \
        for (int _task = warp_id; _task < BLOCK_NT; _task += N_WARPS) { \
            uint32_t _tile_lds = (uint32_t)(B_BASE + (BUF) * B_BUF + _task * B_TILE_BYTES); \
            int _g = n_block + _task * 16 + _load_row; \
            int _g_safe = _g < N ? _g : 0; \
            const uint8_t* _src = B_q + _g_safe * half_K + _kbyte + _load_quarter * 16; \
            uint32_t _m0 = __builtin_amdgcn_readfirstlane((int)_tile_lds); \
            async_global_load_lds_b16(_src, _m0); \
        } \
    } while(0)

    #define ASYNC_LOAD_A_TILES(BUF, KSTEP) do { \
        int _k_elem = ((int)ks_base + (KSTEP) * 4) * 32; \
        for (int _task = warp_id; _task < BLOCK_MT * 4; _task += N_WARPS) { \
            int _tile = _task >> 2; \
            int _sub  = _task & 3; \
            uint32_t _tile_lds = (uint32_t)(A_BASE + (BUF) * A_BUF \
                + _tile * A_TILE_BYTES + _sub * 1024); \
            int _chunk_base = _sub * 64 + lane; \
            int _row = _chunk_base >> 4; \
            int _col_chunk = _chunk_base & 15; \
            int _m_g = m_block + _tile * 16 + _row; \
            int _m_safe = (_m_g < M) ? _m_g : 0; \
            const uint8_t* _src = (const uint8_t*)(A_bf16 + (int64_t)_m_safe * K + _k_elem) \
                + _col_chunk * 16; \
            uint32_t _m0 = __builtin_amdgcn_readfirstlane((int)_tile_lds); \
            async_global_load_lds_b16(_src, _m0); \
        } \
    } while(0)

    #define ASYNC_LOAD_B_SCALE(BUF, KSTEP) do { \
        if (warp_id == 0) { \
            int _i3 = ((int)ks_base + (KSTEP) * 4) >> 3; \
            int _group = lane >> 4; \
            int _chunk_in_group = lane & 15; \
            int _i0 = n_block_i0 + _group; \
            int _i0_safe = (_group < S_N_GROUPS && _i0 * 32 < N) ? _i0 : 0; \
            const uint8_t* _src = B_scale_sh + _i0_safe * K \
                + _i3 * 256 + _chunk_in_group * 16; \
            uint32_t _s_lds = (uint32_t)(S_BASE + (BUF) * S_LDS_BUF); \
            uint32_t _m0 = __builtin_amdgcn_readfirstlane((int)_s_lds); \
            async_global_load_lds_b16(_src, _m0); \
        } \
    } while(0)

    #define GET_B_SCALE_LDS(BUF, N_COL, KS) \
        ((unsigned int)smem[S_BASE + (BUF) * S_LDS_BUF \
            + (((N_COL) >> 5) - n_block_i0) * S_CHUNK \
            + ((KS) & 3) * 64 \
            + ((N_COL) & 15) * 4 \
            + (((KS) >> 2) & 1) * 2 \
            + (((N_COL) >> 4) & 1)])

    i32x8_t a_frag[WM];
    int a_scale_packed[WM];

    #define QUANTIZE_A_FROM_LDS(BUF) do { \
        _Pragma("unroll") \
        for (int _wm = 0; _wm < WM; _wm++) { \
            int _wm_tile = warp_row * WM + _wm; \
            int _m_g = m_warp + _wm * 16 + lane16; \
            int _lds_off = A_BASE + (BUF) * A_BUF \
                + _wm_tile * A_TILE_BYTES + lane16 * A_ROW_BYTES + k_quarter * 64; \
            union { i32x8_t v; uint32_t u32[8]; uint8_t b[32]; } _abuf; \
            _abuf.v = i32x8_t{0,0,0,0,0,0,0,0}; \
            unsigned int _scale_byte = 127u; \
            if (_m_g < M) { \
                __hip_bfloat16 _av[32]; \
                *(uint4*)&_av[0]  = *(const uint4*)&smem[_lds_off]; \
                *(uint4*)&_av[8]  = *(const uint4*)&smem[_lds_off + 16]; \
                *(uint4*)&_av[16] = *(const uint4*)&smem[_lds_off + 32]; \
                *(uint4*)&_av[24] = *(const uint4*)&smem[_lds_off + 48]; \
                float _amax = 0.0f; \
                _Pragma("unroll") \
                for (int _i = 0; _i < 32; _i++) { \
                    float _v = __bfloat162float(_av[_i]); \
                    _amax = fmaxf(_amax, fabsf(_v)); \
                } \
                _scale_byte = compute_e8m0(_amax); \
                if constexpr (USE_HW_CVT) { \
                    union { float f; unsigned int u; } _sc; \
                    _sc.u = (unsigned int)_scale_byte << 23; \
                    float _sf = _sc.f; \
                    _Pragma("unroll") \
                    for (int _j = 0; _j < 4; _j++) { \
                        uint32_t _pk = 0; \
                        _pk = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(_pk, \
                            __bfloat162float(_av[8*_j+0]), __bfloat162float(_av[8*_j+1]), _sf, 0); \
                        _pk = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(_pk, \
                            __bfloat162float(_av[8*_j+2]), __bfloat162float(_av[8*_j+3]), _sf, 1); \
                        _pk = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(_pk, \
                            __bfloat162float(_av[8*_j+4]), __bfloat162float(_av[8*_j+5]), _sf, 2); \
                        _pk = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(_pk, \
                            __bfloat162float(_av[8*_j+6]), __bfloat162float(_av[8*_j+7]), _sf, 3); \
                        _abuf.u32[_j] = _pk; \
                    } \
                } else { \
                    union { float f; unsigned int u; } _inv; \
                    _inv.u = ((unsigned int)(254 - (int)_scale_byte)) << 23; \
                    float _inv_scale = _inv.f; \
                    _Pragma("unroll") \
                    for (int _i = 0; _i < 16; _i++) { \
                        float _f0 = __bfloat162float(_av[2*_i]) * _inv_scale; \
                        float _f1 = __bfloat162float(_av[2*_i+1]) * _inv_scale; \
                        _abuf.b[_i] = float_to_fp4(_f0) | (float_to_fp4(_f1) << 4); \
                    } \
                } \
            } \
            a_frag[_wm] = _abuf.v; \
            a_scale_packed[_wm] = (int)(_scale_byte \
                | (_scale_byte << 8) | (_scale_byte << 16) | (_scale_byte << 24)); \
        } \
    } while(0)

    // ═══════════════════════════════════════════════════════════
    // Prologue: issue loads for first min(PIPE_DEPTH, K_steps) ksteps
    // ═══════════════════════════════════════════════════════════
    #pragma unroll
    for (int p = 0; p < PIPE_DEPTH; p++) {
        if (p < K_steps) {
            ASYNC_LOAD_B_TILES(p, p);
            ASYNC_LOAD_A_TILES(p, p);
            ASYNC_LOAD_B_SCALE(p, p);
        }
    }
    async_lds_fence();
    __syncthreads();
    QUANTIZE_A_FROM_LDS(0);

    // ═══════════════════════════════════════════════════════════
    // Main loop
    // ═══════════════════════════════════════════════════════════
    // Pipeline: prologue loaded bufs 0..min(PIPE_DEPTH,K_steps)-1.
    // At each kstep, read from cur_buf, then prefetch kstep+PIPE_DEPTH
    // INTO cur_buf (pf_buf == cur_buf). Must read BEFORE prefetch!
    // Fence only needed when next_buf's data came from a prefetch
    // (kstep+1 >= PIPE_DEPTH), not from the already-fenced prologue.
    for (int kstep = 0; kstep < K_steps; kstep++) {
        int cur_buf = kstep % PIPE_DEPTH;
        bool has_next = (kstep + 1 < K_steps);

        // 1. Read B + B_scale from cur_buf, issue MFMAs
        int ks = ks_base + kstep * 4 + k_quarter;

        #pragma unroll
        for (int wn = 0; wn < WN; wn++) {
            int global_nt = warp_col * WN + wn;
            int n_col = n_warp + wn * 16 + lane16;

            union { i32x8_t v; uint8_t b[32]; } b_buf;
            b_buf.v = i32x8_t{0,0,0,0,0,0,0,0};
            if (n_col < N)
                *(uint4*)&b_buf.b[0] = *(const uint4*)&smem[B_BASE + cur_buf * B_BUF
                    + global_nt * B_TILE_BYTES + lane16 * B_ROW_BYTES + k_quarter * 16];

            unsigned int b_raw = 0x7Fu;
            if (n_col < N)
                b_raw = GET_B_SCALE_LDS(cur_buf, n_col, ks);
            int b_scale_packed = (int)(b_raw | (b_raw << 8) | (b_raw << 16) | (b_raw << 24));

            #pragma unroll
            for (int wm = 0; wm < WM; wm++) {
                acc[wm][wn] = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
                    a_frag[wm], b_buf.v,
                    acc[wm][wn],
                    4, 4,
                    0, a_scale_packed[wm],
                    0, b_scale_packed
                );
            }
        }

        // 2. Prepare for next iteration
        if (has_next) {
            int next_buf = (kstep + 1) % PIPE_DEPTH;

            // Fence if next_buf's data came from a prefetch (not prologue)
            if (kstep + 1 >= PIPE_DEPTH) {
                async_lds_fence();
            }

            // Sync: all warps done reading cur_buf + fence propagated
            __syncthreads();

            // Prefetch kstep+PIPE_DEPTH into cur_buf
            // (safe: all warps done reading cur_buf, cur_buf != next_buf)
            int pf_step = kstep + PIPE_DEPTH;
            if (pf_step < K_steps) {
                ASYNC_LOAD_B_TILES(cur_buf, pf_step);
                ASYNC_LOAD_A_TILES(cur_buf, pf_step);
                ASYNC_LOAD_B_SCALE(cur_buf, pf_step);
            }

            // Quant A from next_buf (ready: prologue-fenced or just-fenced)
            QUANTIZE_A_FROM_LDS(next_buf);
        } else {
            __syncthreads();
        }
    }

    #undef ASYNC_LOAD_B_TILES
    #undef ASYNC_LOAD_A_TILES
    #undef ASYNC_LOAD_B_SCALE
    #undef GET_B_SCALE_LDS
    #undef QUANTIZE_A_FROM_LDS

    // Store output
    #pragma unroll
    for (int wm = 0; wm < WM; wm++) {
        #pragma unroll
        for (int wn = 0; wn < WN; wn++) {
            int n_col = n_warp + wn * 16 + col_out;
            if (n_col < N) {
                for (int i = 0; i < 4; i++) {
                    int m_g = m_warp + wm * 16 + row_base_out + i;
                    if (m_g < M) {
                        if (SPLIT_K) {
                            atomicAdd(&C_f32[m_g * N + n_col], acc[wm][wn][i]);
                        } else {
                            C_bf16[m_g * N + n_col] = __float2bfloat16(acc[wm][wn][i]);
                        }
                    }
                }
            }
        }
    }
#endif
}

__global__ void f32_to_bf16_kernel(const float* __restrict__ in,
                                    __hip_bfloat16* __restrict__ out, int n) {
    int i = blockIdx.x * 256 + threadIdx.x;
    if (i < n) out[i] = __float2bfloat16(in[i]);
}

// ═══════════════════════════════════════════════════════════════
// Host dispatcher — tile configs × pipe depths
// ═══════════════════════════════════════════════════════════════
void fp4_gemm_fused(
    torch::Tensor A_bf16,
    int64_t B_q_ptr, int64_t B_scale_ptr,
    torch::Tensor C_bf16, torch::Tensor C_f32,
    int M, int N, int K,
    int wm, int wn, int gm, int gn, int k_split, int pipe_depth
) {
    const __hip_bfloat16* a = (const __hip_bfloat16*)A_bf16.data_ptr();
    const uint8_t* b  = (const uint8_t*)B_q_ptr;
    const uint8_t* bs = (const uint8_t*)B_scale_ptr;
    __hip_bfloat16* c_bf16 = (__hip_bfloat16*)C_bf16.data_ptr();
    float* c_f32 = k_split > 1 ? (float*)C_f32.data_ptr() : nullptr;

    int Ks = K >> 5;
    int chunk_Ks = Ks / k_split;
    int cfg = wm * 1000 + wn * 100 + gm * 10 + gn;

    // Encode cfg + pipe_depth
    int key = cfg * 10 + pipe_depth;

    #define LAUNCH_KERNEL(WM, WN, GM, GN, SPLITK, HWCVT, PD) do { \
        constexpr int BM = (WM)*(GM)*16, BN = (WN)*(GN)*16; \
        constexpr int BNT = (WN)*(GN); \
        constexpr int BMT = (WM)*(GM); \
        constexpr int SMEM = (PD) * BNT * 16 * 64 \
                           + (PD) * BMT * 16 * 256 \
                           + (PD) * 1024; \
        dim3 block((GM)*(GN)*64); \
        if (SPLITK) { \
            dim3 grid((N + BN-1) / BN, (M + BM-1) / BM, k_split); \
            fp4_gemm_fused_kernel<WM, WN, GM, GN, true, HWCVT, PD><<<grid, block, SMEM>>>( \
                a, b, bs, nullptr, c_f32, M, N, K, chunk_Ks); \
        } else { \
            dim3 grid((N + BN-1) / BN, (M + BM-1) / BM); \
            fp4_gemm_fused_kernel<WM, WN, GM, GN, false, HWCVT, PD><<<grid, block, SMEM>>>( \
                a, b, bs, c_bf16, nullptr, M, N, K, chunk_Ks); \
        } \
    } while(0)

    // Only 3 configs actually used by _select_config (reduces JIT compile time)
    switch (key) {
        case 11114: LAUNCH_KERNEL(1, 1, 1, 1, (k_split>1), true, 4); break;  // K<=512
        case 11144: LAUNCH_KERNEL(1, 1, 1, 4, (k_split>1), true, 4); break;  // K>512, M<=16
        case 12112: LAUNCH_KERNEL(1, 2, 1, 1, (k_split>1), true, 2); break;  // K>512, M>=32
    }
    #undef LAUNCH_KERNEL

    if (k_split > 1) {
        int total = M * N;
        f32_to_bf16_kernel<<<(total + 255) / 256, 256>>>(c_f32, c_bf16, total);
    }
}
"""

CPP_SRC = """
void fp4_gemm_fused(
    torch::Tensor A_bf16,
    int64_t B_q_ptr, int64_t B_scale_ptr,
    torch::Tensor C_bf16, torch::Tensor C_f32,
    int M, int N, int K,
    int wm, int wn, int gm, int gn, int k_split, int pipe_depth
);
"""

_module = load_inline(
    name="fp4_gemm_v31",
    cpp_sources=[CPP_SRC],
    cuda_sources=[CUDA_SRC],
    functions=["fp4_gemm_fused"],
    verbose=True,
    extra_cuda_cflags=[
        "--offload-arch=gfx950", "-std=c++20", "-O3",
        "-ffast-math",
        "-ffp-contract=fast",
        "-munsafe-fp-atomics",
    ],
)


def _select_config(M, N, K):
    """Hardcoded best (wm, wn, gm, gn, pipe_depth) from V29/V30 auto-tune."""
    if K <= 512:
        # Small K: (1,1,1,1) pipe=4 eliminates all fence stalls
        wm, wn, gm, gn = (1, 1, 1, 1)
        pipe_depth = 4
    elif M <= 16:
        # Small M, large K: (1,1,1,4) with k_split for parallelism
        wm, wn, gm, gn = (1, 1, 1, 4)
        pipe_depth = 4
    else:
        # Large shapes: (1,2,1,1) A-reuse
        wm, wn, gm, gn = (1, 2, 1, 1)
        pipe_depth = 2  # pipe=4 hurts on 256×3072×1536 (SMEM pressure)

    block_m = wm * gm * 16
    block_n = wn * gn * 16
    n_m_blocks = (M + block_m - 1) // block_m
    n_n_blocks = (N + block_n - 1) // block_n
    base_blocks = n_m_blocks * n_n_blocks
    Ks = K >> 5
    k_split = 1
    if K >= 1536 and base_blocks < 128:
        target_blocks = 256
        k_split = max(1, target_blocks // base_blocks)
        MIN_CHUNK_KS = 8
        k_split = min(k_split, max(1, Ks // MIN_CHUNK_KS))
        while k_split > 1:
            if Ks % k_split == 0 and (Ks // k_split) % 4 == 0:
                break
            k_split -= 1

    return wm, wn, gm, gn, k_split, pipe_depth


def _run_kernel(A, b_q_ptr, b_scale_ptr, M, N, K, wm, wn, gm, gn, k_split, pipe_depth):
    C = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
    if k_split > 1:
        C_f32 = torch.zeros((M, N), dtype=torch.float32, device=A.device)
    else:
        C_f32 = torch.empty(1, dtype=torch.float32, device=A.device)
    _module.fp4_gemm_fused(
        A, b_q_ptr, b_scale_ptr,
        C, C_f32, M, N, K, wm, wn, gm, gn, k_split, pipe_depth,
    )
    return C


# Minimal warmup — just JIT compile, no auto-tune overhead
def _warmup():
    from aiter.ops.triton.quant import dynamic_mxfp4_quant
    A = torch.randn((4, 512), dtype=torch.bfloat16, device="cuda")
    B = torch.randn((64, 512), dtype=torch.bfloat16, device="cuda")
    B_fp4, B_scale = dynamic_mxfp4_quant(B)
    for _ in range(3):
        _run_kernel(A, B_fp4.data_ptr(), B_scale.data_ptr(),
                   4, 64, 512, 1, 1, 1, 1, 1, 2)
    torch.cuda.synchronize()

_warmup()


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_q.shape[0]

    wm, wn, gm, gn, k_split, pipe_depth = _select_config(M, N, K)

    return _run_kernel(A, B_q.data_ptr(), B_scale_sh.data_ptr(),
                       M, N, K, wm, wn, gm, gn, k_split, pipe_depth)
scrolls · 542 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