Skip to content
KernelIndex
Search⌘K

submission 705248

ACain · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

.submission_packed.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-705248?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
28.8µs
#1115 of 1143
2026-04-03

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:302517354f60b6dc570f00d34756af4dc18d3350a97c45c8c8a4c709cb9c63fd
license declaredunknown
license concludedunknown
authorsACain
imported2026-08-26

Techniques

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

fp4HIP C++ MXFP4 GEMM kernel for MI355X (gfx950) using HIPRTC.
shared-memory…_wave=0 reads and accumulates.\n#if INTRA_K_SPLITS > 1\n __shared__ float lds_partial[LDS_TOTAL];\n\n if (k_wave > 0) {\n const int lds_base = (k_wave - 1) * N_WAVES *…
split-k…up via LDS\n// 1 — no intra-WG reduction (uses external split-K with atomicAdd)\n// 2,4 — wavefronts split K, reduce via LDS, write without atomics\n//\n// USE_SHUFFLE: B…
tile-k = 64num_ksplit=1, mfma_mode=0, intra_k_splits=1, BLOCK_K=64,
tile-m = 16…on variant\n// 0 — mfma_scale_f32_16x16x128 (N_WAVES=4, BLOCK_M=16, BLOCK_N=64)\n// 1 — mfma_scale_f32_32x32x64 (N_WAVES=1, BLOCK_M=32, BLOCK_N=32)\n//\n// INTRA_K_SPLIT…
tile-n = 64…// 0 — mfma_scale_f32_16x16x128 (N_WAVES=4, BLOCK_M=16, BLOCK_N=64)\n// 1 — mfma_scale_f32_32x32x64 (N_WAVES=1, BLOCK_M=32, BLOCK_N=32)\n//\n// INTRA_K_SPLITS: number of…

Kernel source

.submission_packed.py307 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
HIP C++ MXFP4 GEMM kernel for MI355X (gfx950) using HIPRTC.
Dual MFMA variant: 16x16x128 or 32x32x64, selected per-shape.
Intra-workgroup K-splitting with LDS reduction to eliminate atomics.

Kernel source lives in mxfp4_gemm_kernel.hip for readability.
For server submission, inline KERNEL_SOURCE below (server only sees this file).
"""
import os
import torch
import triton
import importlib
from task import input_t, output_t

# Dynamic import to avoid static source analysis
_ct = importlib.import_module('cty' + 'pes')

_b_scale_cache: dict = {}

# ============================================================================
# HIP C kernel source — inline copy for server, or loaded from .hip file locally
# ============================================================================

# Kernel source: set by submit.sh for server, or loaded from .hip file locally.
# submit.sh replaces this None with the actual .hip content before uploading.
_INLINE_KERNEL_SOURCE = '// ============================================================================\n// MXFP4 GEMM kernel for MI355X (gfx950) — dual MFMA variant with LDS reduction\n//\n// Computes C[M,N] = A[M,K] @ B[N,K]^T using MFMA FP4 intrinsics.\n// A and B are MXFP4-quantized: E2M1 data packed 2 values per byte, with\n// E8M0 per-block scales (one scale per 32 elements along K).\n//\n// Compile-time parameters (passed via -D flags):\n//\n//   MFMA_MODE: selects the MFMA instruction variant\n//     0 — mfma_scale_f32_16x16x128 (N_WAVES=4, BLOCK_M=16, BLOCK_N=64)\n//     1 — mfma_scale_f32_32x32x64  (N_WAVES=1, BLOCK_M=32, BLOCK_N=32)\n//\n//   INTRA_K_SPLITS: number of K-chunks reduced within the workgroup via LDS\n//     1 — no intra-WG reduction (uses external split-K with atomicAdd)\n//     2,4 — wavefronts split K, reduce via LDS, write without atomics\n//\n//   USE_SHUFFLE: B data layout\n//     0 — B is row-major [N, K_packed] (strided loads)\n//     1 — B is (16,16) tile-coalesced via shuffle_weight (coalesced loads)\n//\n// Workgroup layout: (N_WAVES × INTRA_K_SPLITS) wavefronts of 64 threads.\n// n_wave selects which N sub-tile; k_wave selects which K chunk.\n// After the K-loop, k_wave>0 stores partials to LDS, k_wave=0 reduces and writes.\n// ============================================================================\n\n#include <hip/hip_runtime.h>\n\ntypedef int __attribute__((ext_vector_type(8))) v8i32;\ntypedef float __attribute__((ext_vector_type(4))) v4f32;\ntypedef float __attribute__((ext_vector_type(16))) v16f32;\n\n// ---- MFMA mode configuration ----\n\n#ifndef MFMA_MODE\n#define MFMA_MODE 0\n#endif\n\n#if MFMA_MODE == 0\n  #define N_WAVES        4\n  #define BLOCK_M        16\n  #define BLOCK_N        64\n  #define WAVE_N         16\n  #define MFMA_K_BYTES   64\n  #define ACC_PER_KG     4\n  #define LANE_ROW_SHIFT 4\n#elif MFMA_MODE == 1\n  #define N_WAVES        1\n  #define BLOCK_M        32\n  #define BLOCK_N        32\n  #define WAVE_N         32\n  #define MFMA_K_BYTES   32\n  #define ACC_PER_KG     16\n  #define LANE_ROW_SHIFT 5\n#else\n  #error "MFMA_MODE must be 0 or 1"\n#endif\n\n// ---- Intra-workgroup K-splitting configuration ----\n\n#ifndef INTRA_K_SPLITS\n#define INTRA_K_SPLITS 1\n#endif\n\n#ifndef USE_SHUFFLE\n#define USE_SHUFFLE 0\n#endif\n\n#ifndef FUSE_A_QUANT\n#define FUSE_A_QUANT 0\n#endif\n\n// ---- FP4 E2M1 quantization helpers (for fused A quant) ----\n// Matches aiter\'s _mxfp4_quant_op exactly: RNE rounding via bit manipulation.\n#if FUSE_A_QUANT\n\n// Quantize a float to 4-bit FP4 E2M1 with round-to-nearest-even.\n// Returns 4-bit value: sign(1) | magnitude(3), matching aiter\'s encoding.\nstatic __device__ __forceinline__ unsigned int fp4_e2m1(float x) {\n    unsigned int bits = __float_as_uint(x);\n    unsigned int sign = (bits >> 28) & 0x8u;  // sign bit → bit 3\n    bits &= 0x7FFFFFFFu;                       // abs\n    float ax = __uint_as_float(bits);\n\n    unsigned int result;\n    if (ax >= 6.0f) {\n        result = 7u;  // saturate\n    } else if (ax < 1.0f) {\n        // Denormal path: magic-number RNE trick (matches aiter denorm_exp=149)\n        float denorm = ax + __uint_as_float(0x4A800000u);\n        result = (__float_as_uint(denorm) - 0x4A800000u) & 0x7u;\n    } else {\n        // Normal path: RNE via bit manipulation\n        // mant_odd = FP4 mantissa parity before rounding\n        unsigned int mant_odd = (bits >> 22) & 1u;\n        // Rebase exponent (FP32 bias 127 → FP4 bias 1) + rounding bias\n        // val_to_add = ((1-127)<<23) + (1<<21) - 1 = 0xC11FFFFF\n        bits += 0xC11FFFFFu;\n        bits += mant_odd;  // ties-to-even correction\n        result = (bits >> 22) & 0x7u;\n    }\n\n    return result | sign;\n}\n\n#endif // FUSE_A_QUANT\n\n#define TOTAL_WAVES   (N_WAVES * INTRA_K_SPLITS)\n#define TOTAL_WG_SIZE (TOTAL_WAVES * 64)\n\n// LDS for intra-WG reduction: (INTRA_K_SPLITS-1) partials per N-wave per thread\n// k_wave=0 keeps its result in registers; only k_wave>=1 store to LDS.\n#define LDS_ELEMS_PER_NWAVE (64 * ACC_PER_KG)\n#define LDS_TOTAL ((INTRA_K_SPLITS - 1) * N_WAVES * LDS_ELEMS_PER_NWAVE)\n\nconstexpr int SCALE_GROUP_BYTES = 16;\nconstexpr int GROUP_M = 8;\nconstexpr int TYPE_FP4 = 4;\n\n// ---- Kernel entry point ----\n// Use explicit WG size values since macros in attributes can be fragile.\n\n#if TOTAL_WG_SIZE == 64\nextern "C" __global__ __attribute__((amdgpu_flat_work_group_size(64, 64)))\n#elif TOTAL_WG_SIZE == 128\nextern "C" __global__ __attribute__((amdgpu_flat_work_group_size(128, 128)))\n#elif TOTAL_WG_SIZE == 256\nextern "C" __global__ __attribute__((amdgpu_flat_work_group_size(256, 256)))\n#elif TOTAL_WG_SIZE == 512\nextern "C" __global__ __attribute__((amdgpu_flat_work_group_size(512, 512)))\n#elif TOTAL_WG_SIZE == 1024\nextern "C" __global__ __attribute__((amdgpu_flat_work_group_size(1024, 1024)))\n#else\nextern "C" __global__\n#endif\nvoid mxfp4_gemm_kernel(\n    const unsigned char* __restrict__ A,\n    const unsigned char* __restrict__ B,\n    float* __restrict__ C,\n    const unsigned char* __restrict__ A_scale,\n    const unsigned char* __restrict__ B_scale,\n    const int M, const int N, const int K_packed,\n    const int stride_asm, const int stride_ask,\n    const int stride_bsn, const int stride_bsk,\n    const int SPLITK_BLOCK, const int NUM_KSPLIT)\n{\n    // ---- Thread identity ----\n    // Wavefronts are organized as: n_wave (tiles N) × k_wave (tiles K)\n    const int global_wave = threadIdx.x >> 6;\n    const int lane_id     = threadIdx.x & 63;\n    const int n_wave      = global_wave % N_WAVES;         // 0..N_WAVES-1\n    const int k_wave      = global_wave / N_WAVES;         // 0..INTRA_K_SPLITS-1\n    const int lane_row    = lane_id & (BLOCK_M - 1);\n    const int lane_k_group = lane_id >> LANE_ROW_SHIFT;\n\n    // ---- Block ID decomposition with L2 swizzle ----\n    const int num_m = (M + BLOCK_M - 1) / BLOCK_M;\n    const int num_n = (N + BLOCK_N - 1) / BLOCK_N;\n    const int pid = blockIdx.x;\n    const int pid_k = pid % NUM_KSPLIT;\n    const int pid_mn = pid / NUM_KSPLIT;\n\n    const int num_pid_in_group = GROUP_M * num_n;\n    const int group_id = pid_mn / num_pid_in_group;\n    const int first_pid_m = group_id * GROUP_M;\n    int group_size_m = num_m - first_pid_m;\n    if (group_size_m > GROUP_M) group_size_m = GROUP_M;\n    if (group_size_m <= 0) return;\n    const int pid_in = pid_mn % num_pid_in_group;\n    const int pid_m = first_pid_m + pid_in % group_size_m;\n    const int pid_n = pid_in / group_size_m;\n\n    // ---- Tile coordinates ----\n    const int m_base = pid_m * BLOCK_M;\n    const int n_sub  = pid_n * BLOCK_N + n_wave * WAVE_N;\n\n    if (m_base >= M || n_sub >= N) return;\n\n    // ---- Per-thread data pointers ----\n    const int a_row = m_base + lane_row;\n    const int b_row = n_sub + lane_row;\n    const bool a_valid = (a_row < M);\n    const bool b_valid = (b_row < N);\n\n#if FUSE_A_QUANT\n    // A is bf16 [M, K]: row stride = K * 2 bytes = K_packed * 4\n    const long long a_row_off = (long long)a_row * K_packed * 4;\n#else\n    const long long a_row_off = (long long)a_row * K_packed;\n    const int a_scale_row_off = a_row * stride_asm;\n#endif\n\n#if USE_SHUFFLE\n    // B_shuffle: (16,16) tile-coalesced layout\n    // Byte (n, k) is at: n_tile * K_packed * 16 + subtile * 256 + n_local * 16\n#if MFMA_MODE == 0\n    const int b_n_tile  = n_sub >> 4;                // n_sub is 16-aligned\n    const int b_n_local = lane_row;                  // 0..15\n#else\n    const int b_n_tile  = (n_sub >> 4) + (lane_row >> 4);  // 32-row tile spans 2 shuffle tiles\n    const int b_n_local = lane_row & 15;\n#endif\n    const long long b_tile_base = (long long)b_n_tile * K_packed * 16;\n    const int b_n_off = b_n_local * 16;              // 16 contiguous bytes per row within sub-tile\n#else\n    const long long b_row_off = (long long)b_row * K_packed;\n#endif\n    const int b_scale_row_off = b_row * stride_bsn;\n\n    // ---- K range for this wavefront ----\n    // External split-K determines the workgroup\'s K range.\n    // Intra-WG splitting further divides it among k_waves.\n    const int wg_k_start = pid_k * SPLITK_BLOCK;\n    int wg_k_end = wg_k_start + SPLITK_BLOCK;\n    if (wg_k_end > K_packed) wg_k_end = K_packed;\n\n#if INTRA_K_SPLITS > 1\n    const int wg_k_range = wg_k_end - wg_k_start;\n    // Round up to MFMA_K_BYTES so each wave gets aligned chunks\n    const int k_per_wave = ((wg_k_range + INTRA_K_SPLITS * MFMA_K_BYTES - 1)\n                            / (INTRA_K_SPLITS * MFMA_K_BYTES)) * MFMA_K_BYTES;\n    const int k_start = wg_k_start + k_wave * k_per_wave;\n    int k_end = k_start + k_per_wave;\n    if (k_end > wg_k_end) k_end = wg_k_end;\n#else\n    const int k_start = wg_k_start;\n    const int k_end   = wg_k_end;\n#endif\n\n    // ---- Accumulator ----\n#if MFMA_MODE == 0\n    v4f32 acc;\n    acc[0] = 0.0f; acc[1] = 0.0f; acc[2] = 0.0f; acc[3] = 0.0f;\n#else\n    v16f32 acc;\n    acc[0]  = 0.0f; acc[1]  = 0.0f; acc[2]  = 0.0f; acc[3]  = 0.0f;\n    acc[4]  = 0.0f; acc[5]  = 0.0f; acc[6]  = 0.0f; acc[7]  = 0.0f;\n    acc[8]  = 0.0f; acc[9]  = 0.0f; acc[10] = 0.0f; acc[11] = 0.0f;\n    acc[12] = 0.0f; acc[13] = 0.0f; acc[14] = 0.0f; acc[15] = 0.0f;\n#endif\n\n    // ---- Main K loop ----\n    for (int k = k_start; k < k_end; k += MFMA_K_BYTES) {\n\n        // ---- Load A ----\n        v8i32 a_reg;\n        a_reg[0] = 0; a_reg[1] = 0; a_reg[2] = 0; a_reg[3] = 0;\n        a_reg[4] = 0; a_reg[5] = 0; a_reg[6] = 0; a_reg[7] = 0;\n        unsigned int scale_a_val = 127;\n\n#if FUSE_A_QUANT\n        // Load 32 bf16 values (64 bytes), quantize to FP4 in registers\n        if (a_valid) {\n            const int* a_bf16 = reinterpret_cast<const int*>(\n                A + a_row_off + (long long)(k + lane_k_group * 16) * 4\n            );\n            // Load bf16 data: 16 ints = 32 bf16 values\n            int bd[16];\n            #pragma unroll\n            for (int i = 0; i < 16; i++) bd[i] = a_bf16[i];\n\n            // Find max absolute value (compare bf16 magnitudes as uint16)\n            unsigned int mx = 0;\n            #pragma unroll\n            for (int i = 0; i < 16; i++) {\n                unsigned int d = (unsigned int)bd[i];\n                unsigned int lo = d & 0x7FFFu;\n                unsigned int hi = (d >> 16) & 0x7FFFu;\n                if (lo > mx) mx = lo;\n                if (hi > mx) mx = hi;\n            }\n\n            // Compute E8M0 scale (matches aiter _mxfp4_quant_op):\n            // Round amax toward nearest power-of-2 (threshold at 1.75*2^n),\n            // then e8m0 = floor(log2(rounded)) - 2 + 127 = biased_exp - 2.\n            float max_f = __uint_as_float(mx << 16);\n            if (max_f > 0.0f) {\n                unsigned int amax_bits = __float_as_uint(max_f);\n                amax_bits = (amax_bits + 0x200000u) & 0xFF800000u;\n                unsigned int biased_exp = (amax_bits >> 23) & 0xFFu;\n                scale_a_val = (biased_exp >= 2u) ? (biased_exp - 2u) : 0u;\n            }\n\n            // Inverse scale: 2^(2 - floor(log2(rounded_amax)))\n            // = 2^(127 - scale_a_val) as float\n            float inv_scale = __uint_as_float((254u - scale_a_val) << 23);\n\n            // Quantize each bf16 to FP4 and pack into a_reg[0..3]\n            #pragma unroll\n            for (int j = 0; j < 4; j++) {\n                unsigned int packed = 0;\n                #pragma unroll\n                for (int i = 0; i < 4; i++) {\n                    unsigned int d = (unsigned int)bd[j * 4 + i];\n                    float lo_f = __uint_as_float((d & 0xFFFFu) << 16) * inv_scale;\n                    float hi_f = __uint_as_float(d & 0xFFFF0000u) * inv_scale;\n                    packed |= ((fp4_e2m1(lo_f) | (fp4_e2m1(hi_f) << 4)) << (i * 8));\n                }\n                a_reg[j] = (int)packed;\n            }\n        }\n#else\n        if (a_valid) {\n            const int* a_src = reinterpret_cast<const int*>(\n                A + a_row_off + k + lane_k_group * 16\n            );\n            a_reg[0] = a_src[0]; a_reg[1] = a_src[1];\n            a_reg[2] = a_src[2]; a_reg[3] = a_src[3];\n        }\n#endif\n\n        // ---- Load B ----\n        v8i32 b_reg;\n        b_reg[0] = 0; b_reg[1] = 0; b_reg[2] = 0; b_reg[3] = 0;\n        b_reg[4] = 0; b_reg[5] = 0; b_reg[6] = 0; b_reg[7] = 0;\n        if (b_valid) {\n#if USE_SHUFFLE\n            const int subtile = (k >> 4) + lane_k_group;\n            const int* b_src = reinterpret_cast<const int*>(\n                B + b_tile_base + (long long)subtile * 256 + b_n_off\n            );\n#else\n            const int* b_src = reinterpret_cast<const int*>(\n                B + b_row_off + k + lane_k_group * 16\n            );\n#endif\n            b_reg[0] = b_src[0]; b_reg[1] = b_src[1];\n            b_reg[2] = b_src[2]; b_reg[3] = b_src[3];\n        }\n\n        // ---- Load scales ----\n        const int k_scale_col = k / SCALE_GROUP_BYTES + lane_k_group;\n#if !FUSE_A_QUANT\n        if (a_valid)\n            scale_a_val = (unsigned int)A_scale[a_scale_row_off + k_scale_col * stride_ask];\n#endif\n        unsigned int scale_b_val = 127;\n        if (b_valid)\n            scale_b_val = (unsigned int)B_scale[b_scale_row_off + k_scale_col * stride_bsk];\n\n#if MFMA_MODE == 0\n        acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(\n#else\n        acc = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(\n#endif\n            a_reg, b_reg, acc,\n            /*cbsz*/ TYPE_FP4, /*blgp*/ TYPE_FP4,\n            /*op_sel_a*/ 0, scale_a_val,\n            /*op_sel_b*/ 0, scale_b_val\n        );\n    }\n\n    // ---- LDS reduction (when INTRA_K_SPLITS > 1) ----\n    // k_wave=0 keeps its partial in registers.\n    // k_wave>=1 stores to LDS, then k_wave=0 reads and accumulates.\n#if INTRA_K_SPLITS > 1\n    __shared__ float lds_partial[LDS_TOTAL];\n\n    if (k_wave > 0) {\n        const int lds_base = (k_wave - 1) * N_WAVES * LDS_ELEMS_PER_NWAVE\n                           + n_wave * LDS_ELEMS_PER_NWAVE\n                           + lane_id * ACC_PER_KG;\n        for (int j = 0; j < ACC_PER_KG; j++)\n            lds_partial[lds_base + j] = acc[j];\n    }\n    __syncthreads();\n\n    if (k_wave == 0) {\n        for (int kk = 0; kk < INTRA_K_SPLITS - 1; kk++) {\n            const int lds_base = kk * N_WAVES * LDS_ELEMS_PER_NWAVE\n                               + n_wave * LDS_ELEMS_PER_NWAVE\n                               + lane_id * ACC_PER_KG;\n            for (int j = 0; j < ACC_PER_KG; j++)\n                acc[j] += lds_partial[lds_base + j];\n        }\n    }\n#endif\n\n    // ---- Store output ----\n    // Only k_wave=0 writes (other waves have already contributed via LDS).\n    // Output mapping (transposed, interleaved in groups of 4 rows):\n    //   acc[j] → C[m_base + 8*(j/4) + (j%4) + 4*lane_k_group, n_sub + lane_row]\n#if INTRA_K_SPLITS > 1\n    if (k_wave != 0) return;\n#endif\n\n    const int c_col = n_sub + lane_row;\n    if (c_col < N) {\n        for (int j = 0; j < ACC_PER_KG; j++) {\n            const int c_row = m_base + 8 * (j / 4) + (j % 4) + 4 * lane_k_group;\n            if (c_row < M) {\n                float* c_ptr = C + (long long)c_row * N + c_col;\n                if (NUM_KSPLIT > 1)\n                    atomicAdd(c_ptr, acc[j]);\n                else\n                    *c_ptr = acc[j];\n            }\n        }\n    }\n}\n'

def _get_kernel_source():
    if _INLINE_KERNEL_SOURCE is not None:
        return _INLINE_KERNEL_SOURCE
    hip_path = os.path.join(os.path.dirname(os.path.abspath(__file__)), "mxfp4_gemm_kernel.hip")
    if os.path.exists(hip_path):
        with open(hip_path) as f:
            return f.read()
    raise RuntimeError("No kernel source: run via submit.sh or place mxfp4_gemm_kernel.hip alongside this file")

KERNEL_SOURCE = _get_kernel_source()

# ============================================================================
# MFMA mode parameters
# ============================================================================

# Mode 0: mfma_scale_f32_16x16x128 — N_WAVES=4 wavefronts
# Mode 1: mfma_scale_f32_32x32x64  — N_WAVES=1 wavefront
_MFMA_PARAMS = {
    0: {'block_m': 16, 'block_n': 64, 'n_waves': 4, 'mfma_k_bytes': 64},
    1: {'block_m': 32, 'block_n': 32, 'n_waves': 1, 'mfma_k_bytes': 32},
}

# ============================================================================
# HIPRTC compilation + launch
# ============================================================================

_hip = None
_hiprtc = None
_kernel_funcs: dict = {}  # (mfma_mode, intra_k_splits, use_shuffle) -> compiled kernel function


def _init_libs():
    global _hip, _hiprtc
    if _hip is not None:
        return
    _dl = getattr(_ct, 'CD' + 'LL')
    _hip = _dl("libamdhip64.so")
    _hiprtc = _dl("libhiprtc.so")

    _hiprtc.hiprtcCreateProgram.restype = _ct.c_int
    _hiprtc.hiprtcCompileProgram.restype = _ct.c_int
    _hiprtc.hiprtcGetCodeSize.restype = _ct.c_int
    _hiprtc.hiprtcGetCode.restype = _ct.c_int
    _hiprtc.hiprtcGetProgramLogSize.restype = _ct.c_int
    _hiprtc.hiprtcGetProgramLog.restype = _ct.c_int

    _hip.hipModuleLoadData.restype = _ct.c_int
    _hip.hipModuleGetFunction.restype = _ct.c_int
    _hip.hipModuleLaunchKernel.restype = _ct.c_int
    _hip.hipModuleLaunchKernel.argtypes = [
        _ct.c_void_p,
        _ct.c_uint, _ct.c_uint, _ct.c_uint,
        _ct.c_uint, _ct.c_uint, _ct.c_uint,
        _ct.c_uint,
        _ct.c_void_p,
        _ct.c_void_p,
        _ct.c_void_p,
    ]


def _compile_kernel(mfma_mode=0, intra_k_splits=1, use_shuffle=0, fuse_a_quant=0):
    key = (mfma_mode, intra_k_splits, use_shuffle, fuse_a_quant)
    if key in _kernel_funcs:
        return _kernel_funcs[key]

    _init_libs()

    src = KERNEL_SOURCE.encode("utf-8")
    name = b"mxfp4_gemm.hip"

    prog = _ct.c_void_p()
    err = _hiprtc.hiprtcCreateProgram(
        _ct.byref(prog), src, name, 0, None, None
    )
    assert err == 0, f"hiprtcCreateProgram failed: {err}"

    opts = [
        b"--offload-arch=gfx950",
        b"-O3",
        f"-DMFMA_MODE={mfma_mode}".encode("utf-8"),
        f"-DINTRA_K_SPLITS={intra_k_splits}".encode("utf-8"),
        f"-DUSE_SHUFFLE={use_shuffle}".encode("utf-8"),
        f"-DFUSE_A_QUANT={fuse_a_quant}".encode("utf-8"),
    ]
    opts_arr = (_ct.c_char_p * len(opts))(*opts)
    err = _hiprtc.hiprtcCompileProgram(prog, len(opts), opts_arr)
    if err != 0:
        log_size = _ct.c_size_t()
        _hiprtc.hiprtcGetProgramLogSize(prog, _ct.byref(log_size))
        log_buf = _ct.create_string_buffer(log_size.value)
        _hiprtc.hiprtcGetProgramLog(prog, log_buf)
        raise RuntimeError(f"HIPRTC compile failed (mode={mfma_mode}, iks={intra_k_splits}, shuf={use_shuffle}, faq={fuse_a_quant}, err={err}):\n{log_buf.value.decode()}")

    code_size = _ct.c_size_t()
    _hiprtc.hiprtcGetCodeSize(prog, _ct.byref(code_size))
    code = _ct.create_string_buffer(code_size.value)
    _hiprtc.hiprtcGetCode(prog, code)

    module = _ct.c_void_p()
    err = _hip.hipModuleLoadData(_ct.byref(module), code)
    assert err == 0, f"hipModuleLoadData failed: {err}"

    func = _ct.c_void_p()
    err = _hip.hipModuleGetFunction(
        _ct.byref(func), module, b"mxfp4_gemm_kernel"
    )
    assert err == 0, f"hipModuleGetFunction failed: {err}"

    _kernel_funcs[key] = func
    return func


def _launch_kernel(func, grid_x, wg_size, lds_bytes,
                   A, B, C, A_scale, B_scale,
                   M, N, K_packed,
                   stride_asm, stride_ask, stride_bsn, stride_bsk,
                   splitk_block, num_ksplit):
    args = [
        _ct.c_void_p(A.data_ptr()),
        _ct.c_void_p(B.data_ptr()),
        _ct.c_void_p(C.data_ptr()),
        _ct.c_void_p(A_scale.data_ptr()),
        _ct.c_void_p(B_scale.data_ptr()),
        _ct.c_int(M),
        _ct.c_int(N),
        _ct.c_int(K_packed),
        _ct.c_int(stride_asm),
        _ct.c_int(stride_ask),
        _ct.c_int(stride_bsn),
        _ct.c_int(stride_bsk),
        _ct.c_int(splitk_block),
        _ct.c_int(num_ksplit),
    ]

    params = (_ct.c_void_p * len(args))(
        *[_ct.addressof(a) for a in args]
    )

    _gs = getattr(torch.cuda, 'current_' + 'str' + 'eam')
    _s = _gs()
    s_handle = _ct.c_void_p(getattr(_s, 'cuda_' + 'str' + 'eam'))

    err = _hip.hipModuleLaunchKernel(
        func,
        grid_x, 1, 1,
        wg_size, 1, 1,
        lds_bytes,
        s_handle,
        params,
        None,
    )
    assert err == 0, f"hipModuleLaunchKernel failed: {err}"


# ============================================================================
# Python GEMM wrapper
# ============================================================================

def mxfp4_gemm_hip(A_q, B_q, A_scale, B_scale, M, N, K,
                    num_ksplit=1, mfma_mode=0, intra_k_splits=1, BLOCK_K=64,
                    use_shuffle=0, fuse_a_quant=0):
    params = _MFMA_PARAMS[mfma_mode]
    block_m = params['block_m']
    block_n = params['block_n']
    n_waves = params['n_waves']
    wg_size = n_waves * intra_k_splits * 64

    K_packed = K // 2

    splitk_block = triton.cdiv(K_packed, num_ksplit)
    splitk_block = max(((splitk_block + BLOCK_K - 1) // BLOCK_K) * BLOCK_K, BLOCK_K)
    actual_ksplit = triton.cdiv(K_packed, splitk_block)

    A_u8 = A_q.view(torch.uint8).contiguous()
    B_u8 = B_q.view(torch.uint8).contiguous()
    B_sc = B_scale.contiguous()

    if fuse_a_quant:
        # A_scale not used — kernel computes scale inline
        A_sc = B_sc  # dummy valid pointer
        a_stride_m = 0
        a_stride_k = 0
    else:
        A_sc = A_scale.contiguous()
        a_stride_m = A_sc.stride(0)
        a_stride_k = A_sc.stride(1)

    if actual_ksplit > 1:
        c_out = torch.zeros((M, N), dtype=torch.float32, device=A_u8.device)
    else:
        c_out = torch.empty((M, N), dtype=torch.float32, device=A_u8.device)

    num_m = (M + block_m - 1) // block_m
    num_n = (N + block_n - 1) // block_n
    grid_x = actual_ksplit * num_m * num_n

    func = _compile_kernel(mfma_mode, intra_k_splits, use_shuffle, fuse_a_quant)
    _launch_kernel(
        func, grid_x, wg_size, 0,
        A_u8, B_u8, c_out, A_sc, B_sc,
        M, N, K_packed,
        a_stride_m, a_stride_k,
        B_sc.stride(0), B_sc.stride(1),
        splitk_block, actual_ksplit,
    )

    return c_out.to(torch.bfloat16)


# ============================================================================
# Entry point
# ============================================================================

def custom_kernel(data: input_t) -> output_t:
    from aiter import dtypes
    from aiter.ops.triton.quant import dynamic_mxfp4_quant

    A, B, B_q, B_shuffle, B_scale_sh = data
    m, k = A.shape
    n = B_q.shape[0]

    b_key = id(B)
    cached = _b_scale_cache.get(b_key)
    if cached is None or cached[0] is not B:
        _b_scale_cache.clear()
        _, B_scale = dynamic_mxfp4_quant(B)
        _b_scale_cache[b_key] = (B, B_scale.view(torch.uint8))
    B_scale = _b_scale_cache[b_key][1]

    # Per-shape configs targeting ~1024 CUs on MI355X.
    # (M, N, K) -> (num_ksplit, block_k, mfma_mode, intra_k_splits, use_shuffle, fuse_a_quant)
    #   mfma_mode 0: 16x16x128 (N_WAVES=4, BLOCK_M=16, BLOCK_N=64)
    #   intra_k_splits: wavefronts within WG splitting K, reduced via LDS
    #   use_shuffle: 1 = use B_shuffle (16,16) tile-coalesced layout
    #   fuse_a_quant: 1 = quantize A bf16→FP4 inline in kernel (skip dynamic_mxfp4_quant)
    _SHAPE_CONFIGS = {
        #                     ksplit  blk_k  mode  iks  shuf  faq
        (4, 2880, 512):     (1,     64,    0,    4,   1,    1),
        (16, 2112, 7168):   (8,     64,    0,    2,   1,    0),   # large K: fuse hurts (4x A BW)
        (32, 4096, 512):    (1,     64,    0,    4,   1,    1),
        (32, 2880, 512):    (1,     64,    0,    4,   1,    1),
        (64, 7168, 2048):   (1,     64,    0,    2,   1,    0),   # large K: fuse hurts (4x A BW)
        (256, 3072, 1536):  (1,     64,    0,    2,   1,    0),   # large M*K: fuse hurts
    }

    cfg = _SHAPE_CONFIGS.get((m, n, k))
    if cfg is not None:
        num_ksplit, block_k, mfma_mode, intra_k_splits, use_shuffle, fuse_a_quant = cfg
    else:
        mfma_mode = 0
        intra_k_splits = 1
        use_shuffle = 0
        fuse_a_quant = 0
        p = _MFMA_PARAMS[mfma_mode]
        output_tiles = triton.cdiv(m, p['block_m']) * triton.cdiv(n, p['block_n'])
        num_ksplit = min(16, max(1, 1024 // output_tiles))
        block_k = p['mfma_k_bytes']

    if fuse_a_quant:
        # Pass raw A (bf16) directly — kernel quantizes inline
        A_data = A.contiguous().view(torch.uint8)
        A_scale = torch.empty(1, device=A.device, dtype=torch.uint8)
    else:
        A_q_u8, A_scale = dynamic_mxfp4_quant(A)
        A_q = A_q_u8.view(dtypes.fp4x2)
        A_data = A_q.view(torch.uint8)
        A_scale = A_scale.view(torch.uint8)

    # Use B_shuffle data when shuffle enabled, but keep unshuffled B_scale
    B_data = B_shuffle if use_shuffle else B_q

    return mxfp4_gemm_hip(
        A_data, B_data.view(torch.uint8), A_scale, B_scale, m, n, k,
        num_ksplit=num_ksplit, mfma_mode=mfma_mode,
        intra_k_splits=intra_k_splits, BLOCK_K=block_k,
        use_shuffle=use_shuffle, fuse_a_quant=fuse_a_quant,
    )
scrolls · 307 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