Skip to content
KernelIndex
Search⌘K

submission 661199

JIAQI PAN · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_hip_vG.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-661199?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
24.2µs
#944 of 1143
2026-03-29

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:e3ae6b7de46ccb4e2281c8988a9318853093bacd05c4f22afc61f6054c127e32
license declaredunknown
license concludedunknown
authorsJIAQI PAN
imported2026-08-26

Techniques

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

shared-memory__shared__ __half s_a[2][GEN_BLOCK_M * GEN_A_LDS_STRIDE];
tile-k = 32constexpr int BLOCK_K = 32;
tile-m = 8constexpr int GEN_BLOCK_M = 8;
tile-n = 32constexpr int GEN_BLOCK_N = 32;
vector-width = float2float2 p0 = __half22float2(__hmul2(a_row2[kk2 + 0], b_row2[kk2 + 0]));

Kernel source

submission_hip_vG.py646 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

"""
Shape-dispatched fused-A native HIP path.

Design:
- keeps fused A quantization
- keeps ping-pong LDS staging
- adds specialized wide-N kernels for the benchmark's small-M regimes
- keeps a general fallback kernel for non-matched shapes

Specialized dispatch:
- M == 4 and N % 64 == 0 -> 4x64 kernel
- M == 16 or M == 32 and N % 64 == 0 -> 8x64 kernel
- otherwise -> general 8x32 kernel
"""

from __future__ import annotations

import os

os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
os.environ.setdefault("CXX", "clang++")

import torch
from torch.utils.cpp_extension import load_inline

from task import input_t, output_t


CPP_SRC = """
void quant_gemm_mxfp4_vG(
    torch::Tensor a,
    torch::Tensor b_q,
    torch::Tensor b_scale_sh,
    torch::Tensor out
);
"""


CUDA_SRC = r"""
#include <torch/extension.h>
#include <ATen/BFloat16.h>
#include <hip/amd_detail/amd_hip_bf16.h>
#include <hip/hip_fp16.h>
#include <hip/hip_runtime.h>

#include <cstdint>
#include <stdexcept>

namespace {

constexpr int SCALE_GROUP_SIZE = 32;
constexpr int BLOCK_K = 32;
constexpr int GEN_BLOCK_M = 8;
constexpr int GEN_BLOCK_N = 32;
constexpr int GEN_A_LDS_STRIDE = BLOCK_K + 2;
constexpr int GEN_B_LDS_STRIDE = BLOCK_K + 2;
constexpr int SPEC_BLOCK_N = 64;
constexpr int SPEC_A_LDS_STRIDE = BLOCK_K + 2;
constexpr int SPEC_B_LDS_STRIDE = BLOCK_K + 2;

__device__ __constant__ uint32_t FP4_TO_F32_BITS[16] = {
    0x00000000u, 0x3F000000u, 0x3F800000u, 0x3FC00000u,
    0x40000000u, 0x40400000u, 0x40800000u, 0x40C00000u,
    0x80000000u, 0xBF000000u, 0xBF800000u, 0xBFC00000u,
    0xC0000000u, 0xC0400000u, 0xC0800000u, 0xC0C00000u,
};

__device__ __forceinline__ uint16_t f32_to_bf16_rne(float x) {
    uint32_t bits = __float_as_uint(x);
    uint32_t lsb = (bits >> 16) & 1u;
    uint32_t bias = 0x7FFFu + lsb;
    return static_cast<uint16_t>((bits + bias) >> 16);
}

__device__ __forceinline__ float e8m0_bits_to_f32(uint8_t e8m0) {
    if (e8m0 == static_cast<uint8_t>(255)) {
        return NAN;
    }
    uint32_t bits = e8m0 == 0 ? 0x00400000u : (static_cast<uint32_t>(e8m0) << 23);
    return __uint_as_float(bits);
}

__device__ __forceinline__ float fp4_bits_to_f32(uint8_t q) {
    return __uint_as_float(FP4_TO_F32_BITS[q & 0xFu]);
}

__device__ __forceinline__ void unpack_fp4x8_half2(uint32_t packed, __half scale_h, __half2* dst) {
    __half2 scale2 = __halves2half2(scale_h, scale_h);
    #pragma unroll
    for (int i = 0; i < 4; ++i) {
        uint8_t lo = static_cast<uint8_t>((packed >> (8 * i)) & 0xFu);
        uint8_t hi = static_cast<uint8_t>((packed >> (8 * i + 4)) & 0xFu);
        __half2 vals = __floats2half2_rn(fp4_bits_to_f32(lo), fp4_bits_to_f32(hi));
        dst[i] = __hmul2(vals, scale2);
    }
}

__device__ __forceinline__ uint32_t wave_reduce_max_u32(uint32_t v) {
    uint32_t t = static_cast<uint32_t>(__shfl_down(static_cast<int>(v), 16, 32));
    v = v > t ? v : t;
    t = static_cast<uint32_t>(__shfl_down(static_cast<int>(v), 8, 32));
    v = v > t ? v : t;
    t = static_cast<uint32_t>(__shfl_down(static_cast<int>(v), 4, 32));
    v = v > t ? v : t;
    t = static_cast<uint32_t>(__shfl_down(static_cast<int>(v), 2, 32));
    v = v > t ? v : t;
    t = static_cast<uint32_t>(__shfl_down(static_cast<int>(v), 1, 32));
    v = v > t ? v : t;
    return v;
}

__device__ __forceinline__ uint8_t scale_bits_from_bf16_abs(uint16_t abs_bits) {
    uint32_t is_zero = abs_bits == 0u;
    uint32_t exp = (abs_bits >> 7) & 0xFFu;
    uint32_t mant = abs_bits & 0x7Fu;
    uint32_t is_subnormal = exp == 0u;

    uint32_t g = (mant >> 6) & 1u;
    uint32_t r = (mant >> 5) & 1u;
    uint32_t s = (mant & 0x1Fu) != 0u;
    uint32_t is_special = (exp == 0xFFu);
    uint32_t round_up = (1u - is_special) & (g & (r | s | 1u));
    uint32_t e8 = exp + round_up;
    int scale_bits = static_cast<int>(e8) - 2;
    if (scale_bits < 0) scale_bits = 0;
    if (scale_bits > 254) scale_bits = 254;
    uint32_t normal_val = static_cast<uint32_t>(scale_bits);
    uint32_t out =
        is_zero * 127u +
        (1u - is_zero) * ((1u - is_subnormal) * normal_val);
    return static_cast<uint8_t>(out);
}

__device__ __forceinline__ uint8_t bf16_to_fp4_bits_scaled(uint16_t x_bits, uint8_t scale_bits) {
    uint32_t sign = (x_bits >> 12) & 0x8u;
    uint32_t abs_bits = x_bits & 0x7FFFu;
    uint32_t exp = (abs_bits >> 7) & 0xFFu;
    uint32_t mant = abs_bits & 0x7Fu;

    uint32_t is_zero = abs_bits == 0u;
    uint32_t is_subnormal = exp == 0u;
    uint32_t is_special = exp == 0xFFu;
    uint32_t is_normal = (1u - is_zero) & (1u - is_subnormal) & (1u - is_special);

    int de = static_cast<int>(exp) - static_cast<int>(scale_bits);
    int d = de;
    if (d < -2) d = -2;
    if (d > 2) d = 2;

    uint32_t gt32 = mant > 32u;
    uint32_t ge64 = mant >= 64u;
    uint32_t ge96 = mant >= 96u;
    uint32_t nz = mant != 0u;

    uint32_t is_m2 = d == -2;
    uint32_t is_m1 = d == -1;
    uint32_t is_0 = d == 0;
    uint32_t is_p1 = d == 1;
    uint32_t is_p2 = d == 2;

    uint32_t mid_mag =
        is_m2 * nz +
        is_m1 * (1u + ge64) +
        is_0 * (2u + gt32 + ge96) +
        is_p1 * (4u + gt32 + ge96) +
        is_p2 * (6u + gt32);

    uint32_t ge3 = de >= 3;
    uint32_t mid_sel = (de > -3) & (de < 3);
    uint32_t mag = ge3 * 7u + mid_sel * mid_mag;
    uint32_t out_mag = is_special * 7u + is_normal * mag;
    return static_cast<uint8_t>(sign | out_mag);
}

template <int BLOCK_M_T, int A_STRIDE_T, bool CHECK_M>
__device__ __forceinline__ void load_a_tile(
    const uint16_t* __restrict__ a_bf16,
    __half* __restrict__ s_a_buf,
    int global_m,
    int local_m,
    int local_n,
    int M,
    int K,
    int kb
) {
    const __half zero_h = __float2half_rn(0.0f);
    const __half2 zero2 = __halves2half2(zero_h, zero_h);

    uint16_t a_bits0 = 0;
    uint16_t a_bits1 = 0;
    uint32_t local_abs = 0;
    bool row_valid = !CHECK_M || (global_m < M);

    if ((local_n & 1) == 0 && local_n < 32 && row_valid) {
        int gk = kb * BLOCK_K + local_n;
        const uint32_t* a_pair_ptr = reinterpret_cast<const uint32_t*>(
            a_bf16 + global_m * K + gk
        );
        uint32_t a_pair = *a_pair_ptr;
        a_bits0 = static_cast<uint16_t>(a_pair & 0xFFFFu);
        a_bits1 = static_cast<uint16_t>(a_pair >> 16);
        uint32_t abs0 = static_cast<uint32_t>(a_bits0 & 0x7FFFu);
        uint32_t abs1 = static_cast<uint32_t>(a_bits1 & 0x7FFFu);
        local_abs = abs0 > abs1 ? abs0 : abs1;
    }

    uint32_t amax_bits = wave_reduce_max_u32(local_abs);
    uint32_t scale_bits_local =
        local_n == 0 ? static_cast<uint32_t>(row_valid
            ? scale_bits_from_bf16_abs(static_cast<uint16_t>(amax_bits))
            : static_cast<uint8_t>(127))
        : 0u;
    uint32_t scale_bits_u32 = static_cast<uint32_t>(__shfl(static_cast<int>(scale_bits_local), 0, 32));

    if ((local_n & 1) == 0 && local_n < 32) {
        __half2 a_pair_h2 = zero2;
        if (row_valid) {
            __half scale_h = __float2half_rn(e8m0_bits_to_f32(static_cast<uint8_t>(scale_bits_u32)));
            __half2 scale2 = __halves2half2(scale_h, scale_h);
            uint8_t q0 = bf16_to_fp4_bits_scaled(a_bits0, static_cast<uint8_t>(scale_bits_u32));
            uint8_t q1 = bf16_to_fp4_bits_scaled(a_bits1, static_cast<uint8_t>(scale_bits_u32));
            __half2 vals = __floats2half2_rn(fp4_bits_to_f32(q0), fp4_bits_to_f32(q1));
            a_pair_h2 = __hmul2(vals, scale2);
        }
        reinterpret_cast<__half2*>(s_a_buf + local_m * A_STRIDE_T)[local_n >> 1] = a_pair_h2;
    }
}

template <int BLOCK_N_T, int B_STRIDE_T>
__device__ __forceinline__ void load_b_tile_loop(
    const uint8_t* __restrict__ b_q,
    const uint8_t* __restrict__ b_scale_sh,
    __half* __restrict__ s_b_buf,
    int linear_tid,
    int block_threads,
    int tile_n_base,
    int kb,
    int K,
    int k_scale,
    int b_q_stride
) {
    constexpr int B_TILE_LOADS = BLOCK_N_T * (BLOCK_K / 8);
    for (int load_idx = linear_tid; load_idx < B_TILE_LOADS; load_idx += block_threads) {
        int ln = load_idx / (BLOCK_K / 8);
        int chunk = load_idx % (BLOCK_K / 8);
        int gn = tile_n_base + ln;
        int gn_group = gn >> 5;
        int lane_row = gn & 31;
        int kb_group = kb >> 3;
        int kb_sub = kb & 7;
        int kb_c = kb_sub & 3;
        int kb_e = kb_sub >> 2;
        int blocks_k = k_scale >> 3;
        int b_scale_base =
            ((((gn_group * blocks_k + kb_group) * 4 + kb_c) << 4) << 2) + (kb_e << 1);
        int scale_src = b_scale_base + ((lane_row & 15) << 2) + (lane_row >> 4);
        __half scale_h = __float2half_rn(e8m0_bits_to_f32(b_scale_sh[scale_src]));
        int gk0 = kb * BLOCK_K + chunk * 8;
        __half2* dst = reinterpret_cast<__half2*>(s_b_buf + ln * B_STRIDE_T + chunk * 8);
        const uint32_t* packed_ptr = reinterpret_cast<const uint32_t*>(
            b_q + gn * b_q_stride + (gk0 >> 1)
        );
        unpack_fp4x8_half2(*packed_ptr, scale_h, dst);
    }
}

template <int A_STRIDE_T, int B_STRIDE_T>
__device__ __forceinline__ void compute_tile_1col(
    const __half* __restrict__ s_a_buf,
    const __half* __restrict__ s_b_buf,
    int local_m,
    int local_n,
    float& acc0,
    float& acc1,
    float& acc2,
    float& acc3
) {
    const __half2* a_row2 = reinterpret_cast<const __half2*>(s_a_buf + local_m * A_STRIDE_T);
    const __half2* b_row2 = reinterpret_cast<const __half2*>(s_b_buf + local_n * B_STRIDE_T);

    #pragma unroll
    for (int kk2 = 0; kk2 < BLOCK_K / 2; kk2 += 4) {
        float2 p0 = __half22float2(__hmul2(a_row2[kk2 + 0], b_row2[kk2 + 0]));
        float2 p1 = __half22float2(__hmul2(a_row2[kk2 + 1], b_row2[kk2 + 1]));
        float2 p2 = __half22float2(__hmul2(a_row2[kk2 + 2], b_row2[kk2 + 2]));
        float2 p3 = __half22float2(__hmul2(a_row2[kk2 + 3], b_row2[kk2 + 3]));
        acc0 += p0.x + p0.y;
        acc1 += p1.x + p1.y;
        acc2 += p2.x + p2.y;
        acc3 += p3.x + p3.y;
    }
}

template <int A_STRIDE_T, int B_STRIDE_T>
__device__ __forceinline__ void compute_tile_2col(
    const __half* __restrict__ s_a_buf,
    const __half* __restrict__ s_b_buf,
    int local_m,
    int local_n,
    float& acc0a,
    float& acc1a,
    float& acc2a,
    float& acc3a,
    float& acc0b,
    float& acc1b,
    float& acc2b,
    float& acc3b
) {
    const __half2* a_row2 = reinterpret_cast<const __half2*>(s_a_buf + local_m * A_STRIDE_T);
    const __half2* b_row2a = reinterpret_cast<const __half2*>(s_b_buf + local_n * B_STRIDE_T);
    const __half2* b_row2b = reinterpret_cast<const __half2*>(s_b_buf + (local_n + 32) * B_STRIDE_T);

    #pragma unroll
    for (int kk2 = 0; kk2 < BLOCK_K / 2; kk2 += 4) {
        float2 pa0 = __half22float2(__hmul2(a_row2[kk2 + 0], b_row2a[kk2 + 0]));
        float2 pa1 = __half22float2(__hmul2(a_row2[kk2 + 1], b_row2a[kk2 + 1]));
        float2 pa2 = __half22float2(__hmul2(a_row2[kk2 + 2], b_row2a[kk2 + 2]));
        float2 pa3 = __half22float2(__hmul2(a_row2[kk2 + 3], b_row2a[kk2 + 3]));
        acc0a += pa0.x + pa0.y;
        acc1a += pa1.x + pa1.y;
        acc2a += pa2.x + pa2.y;
        acc3a += pa3.x + pa3.y;

        float2 pb0 = __half22float2(__hmul2(a_row2[kk2 + 0], b_row2b[kk2 + 0]));
        float2 pb1 = __half22float2(__hmul2(a_row2[kk2 + 1], b_row2b[kk2 + 1]));
        float2 pb2 = __half22float2(__hmul2(a_row2[kk2 + 2], b_row2b[kk2 + 2]));
        float2 pb3 = __half22float2(__hmul2(a_row2[kk2 + 3], b_row2b[kk2 + 3]));
        acc0b += pb0.x + pb0.y;
        acc1b += pb1.x + pb1.y;
        acc2b += pb2.x + pb2.y;
        acc3b += pb3.x + pb3.y;
    }
}

template <bool CHECK_M>
__global__ __launch_bounds__(GEN_BLOCK_M * GEN_BLOCK_N, 2) void gemm_general_kernel(
    const uint16_t* __restrict__ a_bf16,
    const uint8_t* __restrict__ b_q,
    const uint8_t* __restrict__ b_scale_sh,
    uint16_t* __restrict__ out,
    int M,
    int N,
    int K,
    int k_scale,
    int row_offset
) {
    int local_n = threadIdx.x;
    int local_m = threadIdx.y;
    int linear_tid = local_m * GEN_BLOCK_N + local_n;
    int global_m = row_offset + blockIdx.y * GEN_BLOCK_M + local_m;
    int global_n = blockIdx.x * GEN_BLOCK_N + local_n;
    int tile_n_base = blockIdx.x * GEN_BLOCK_N;

    __shared__ __half s_a[2][GEN_BLOCK_M * GEN_A_LDS_STRIDE];
    __shared__ __half s_b[2][GEN_BLOCK_N * GEN_B_LDS_STRIDE];

    float acc0 = 0.0f;
    float acc1 = 0.0f;
    float acc2 = 0.0f;
    float acc3 = 0.0f;

    int k_blocks = K / BLOCK_K;
    int b_q_stride = K / 2;
    int block_threads = GEN_BLOCK_M * GEN_BLOCK_N;

    load_a_tile<GEN_BLOCK_M, GEN_A_LDS_STRIDE, CHECK_M>(
        a_bf16, s_a[0], global_m, local_m, local_n, M, K, 0
    );
    if (linear_tid < GEN_BLOCK_N * (BLOCK_K / 8)) {
        load_b_tile_loop<GEN_BLOCK_N, GEN_B_LDS_STRIDE>(
            b_q, b_scale_sh, s_b[0], linear_tid, block_threads, tile_n_base, 0, K, k_scale, b_q_stride
        );
    }
    __syncthreads();

    for (int kb = 0; kb < k_blocks - 1; ++kb) {
        int cur = kb & 1;
        int nxt = cur ^ 1;
        compute_tile_1col<GEN_A_LDS_STRIDE, GEN_B_LDS_STRIDE>(
            s_a[cur], s_b[cur], local_m, local_n, acc0, acc1, acc2, acc3
        );
        load_a_tile<GEN_BLOCK_M, GEN_A_LDS_STRIDE, CHECK_M>(
            a_bf16, s_a[nxt], global_m, local_m, local_n, M, K, kb + 1
        );
        if (linear_tid < GEN_BLOCK_N * (BLOCK_K / 8)) {
            load_b_tile_loop<GEN_BLOCK_N, GEN_B_LDS_STRIDE>(
                b_q, b_scale_sh, s_b[nxt], linear_tid, block_threads, tile_n_base, kb + 1, K, k_scale, b_q_stride
            );
        }
        __syncthreads();
    }

    compute_tile_1col<GEN_A_LDS_STRIDE, GEN_B_LDS_STRIDE>(
        s_a[(k_blocks - 1) & 1], s_b[(k_blocks - 1) & 1], local_m, local_n, acc0, acc1, acc2, acc3
    );

    if (global_n < N) {
        float acc = (acc0 + acc1) + (acc2 + acc3);
        if constexpr (CHECK_M) {
            if (global_m < M) {
                out[global_m * N + global_n] = f32_to_bf16_rne(acc);
            }
        } else {
            out[global_m * N + global_n] = f32_to_bf16_rne(acc);
        }
    }
}

template <int BLOCK_M_T, bool CHECK_M>
__global__ __launch_bounds__(BLOCK_M_T * 32, 2) void gemm_smallm_widen_kernel(
    const uint16_t* __restrict__ a_bf16,
    const uint8_t* __restrict__ b_q,
    const uint8_t* __restrict__ b_scale_sh,
    uint16_t* __restrict__ out,
    int M,
    int N,
    int K,
    int k_scale,
    int row_offset
) {
    int local_n = threadIdx.x;
    int local_m = threadIdx.y;
    int linear_tid = local_m * 32 + local_n;
    int global_m = row_offset + blockIdx.y * BLOCK_M_T + local_m;
    int tile_n_base = blockIdx.x * SPEC_BLOCK_N;
    int global_n0 = tile_n_base + local_n;
    int global_n1 = global_n0 + 32;

    __shared__ __half s_a[2][BLOCK_M_T * SPEC_A_LDS_STRIDE];
    __shared__ __half s_b[2][SPEC_BLOCK_N * SPEC_B_LDS_STRIDE];

    float acc0a = 0.0f;
    float acc1a = 0.0f;
    float acc2a = 0.0f;
    float acc3a = 0.0f;
    float acc0b = 0.0f;
    float acc1b = 0.0f;
    float acc2b = 0.0f;
    float acc3b = 0.0f;

    int k_blocks = K / BLOCK_K;
    int b_q_stride = K / 2;
    int block_threads = BLOCK_M_T * 32;

    load_a_tile<BLOCK_M_T, SPEC_A_LDS_STRIDE, CHECK_M>(
        a_bf16, s_a[0], global_m, local_m, local_n, M, K, 0
    );
    load_b_tile_loop<SPEC_BLOCK_N, SPEC_B_LDS_STRIDE>(
        b_q, b_scale_sh, s_b[0], linear_tid, block_threads, tile_n_base, 0, K, k_scale, b_q_stride
    );
    __syncthreads();

    for (int kb = 0; kb < k_blocks - 1; ++kb) {
        int cur = kb & 1;
        int nxt = cur ^ 1;
        compute_tile_2col<SPEC_A_LDS_STRIDE, SPEC_B_LDS_STRIDE>(
            s_a[cur], s_b[cur], local_m, local_n,
            acc0a, acc1a, acc2a, acc3a,
            acc0b, acc1b, acc2b, acc3b
        );
        load_a_tile<BLOCK_M_T, SPEC_A_LDS_STRIDE, CHECK_M>(
            a_bf16, s_a[nxt], global_m, local_m, local_n, M, K, kb + 1
        );
        load_b_tile_loop<SPEC_BLOCK_N, SPEC_B_LDS_STRIDE>(
            b_q, b_scale_sh, s_b[nxt], linear_tid, block_threads, tile_n_base, kb + 1, K, k_scale, b_q_stride
        );
        __syncthreads();
    }

    compute_tile_2col<SPEC_A_LDS_STRIDE, SPEC_B_LDS_STRIDE>(
        s_a[(k_blocks - 1) & 1], s_b[(k_blocks - 1) & 1], local_m, local_n,
        acc0a, acc1a, acc2a, acc3a,
        acc0b, acc1b, acc2b, acc3b
    );

    float acca = (acc0a + acc1a) + (acc2a + acc3a);
    float accb = (acc0b + acc1b) + (acc2b + acc3b);

    if constexpr (CHECK_M) {
        if (global_m < M) {
            out[global_m * N + global_n0] = f32_to_bf16_rne(acca);
            out[global_m * N + global_n1] = f32_to_bf16_rne(accb);
        }
    } else {
        out[global_m * N + global_n0] = f32_to_bf16_rne(acca);
        out[global_m * N + global_n1] = f32_to_bf16_rne(accb);
    }
}

}  // namespace

void quant_gemm_mxfp4_vG(
    torch::Tensor a,
    torch::Tensor b_q,
    torch::Tensor b_scale_sh,
    torch::Tensor out
) {
    TORCH_CHECK(a.is_cuda(), "a must be CUDA/HIP");
    TORCH_CHECK(b_q.is_cuda(), "b_q must be CUDA/HIP");
    TORCH_CHECK(b_scale_sh.is_cuda(), "b_scale_sh must be CUDA/HIP");
    TORCH_CHECK(out.is_cuda(), "out must be CUDA/HIP");

    TORCH_CHECK(a.dim() == 2, "a must be 2D");
    TORCH_CHECK(b_q.dim() == 2, "b_q must be 2D");
    TORCH_CHECK(b_scale_sh.dim() == 2, "b_scale_sh must be 2D");
    TORCH_CHECK(out.dim() == 2, "out must be 2D");

    TORCH_CHECK(a.scalar_type() == torch::kBFloat16, "a must be bf16");
    TORCH_CHECK(b_q.element_size() == 1, "b_q elements must be 1 byte");
    TORCH_CHECK(b_scale_sh.element_size() == 1, "b_scale_sh elements must be 1 byte");
    TORCH_CHECK(out.scalar_type() == torch::kBFloat16, "out must be bf16");

    TORCH_CHECK(a.is_contiguous(), "a must be contiguous");
    TORCH_CHECK(b_q.is_contiguous(), "b_q must be contiguous");
    TORCH_CHECK(b_scale_sh.is_contiguous(), "b_scale_sh must be contiguous");
    TORCH_CHECK(out.is_contiguous(), "out must be contiguous");

    int M = static_cast<int>(a.size(0));
    int K = static_cast<int>(a.size(1));
    int N = static_cast<int>(b_q.size(0));
    int k_scale = K / SCALE_GROUP_SIZE;

    TORCH_CHECK(K % 64 == 0, "K must be divisible by 64");
    TORCH_CHECK(b_q.size(1) == K / 2, "b_q shape mismatch");
    TORCH_CHECK(b_scale_sh.size(1) == k_scale, "b_scale_sh col mismatch");
    TORCH_CHECK(b_scale_sh.size(0) >= N, "b_scale_sh row mismatch");
    TORCH_CHECK(b_scale_sh.size(0) % 32 == 0, "b_scale_sh padded rows must be divisible by 32");
    TORCH_CHECK(out.size(0) == M && out.size(1) == N, "out shape mismatch");

    auto a_ptr = reinterpret_cast<const uint16_t*>(a.data_ptr<at::BFloat16>());
    auto out_ptr = reinterpret_cast<uint16_t*>(out.data_ptr<at::BFloat16>());
    auto b_q_ptr = reinterpret_cast<const uint8_t*>(b_q.data_ptr());
    auto b_scale_sh_ptr = reinterpret_cast<const uint8_t*>(b_scale_sh.data_ptr());

    if (N % 64 == 0 && M == 4) {
        dim3 grid(N / 64, 1);
        dim3 block(32, 4);
        gemm_smallm_widen_kernel<4, false><<<grid, block>>>(
            a_ptr, b_q_ptr, b_scale_sh_ptr, out_ptr, M, N, K, k_scale, 0
        );
    } else if (N % 64 == 0 && (M == 16 || M == 32)) {
        dim3 grid(N / 64, M / 8);
        dim3 block(32, 8);
        gemm_smallm_widen_kernel<8, false><<<grid, block>>>(
            a_ptr, b_q_ptr, b_scale_sh_ptr, out_ptr, M, N, K, k_scale, 0
        );
    } else {
        int full_m_tiles = M / GEN_BLOCK_M;
        int tail_m = M - full_m_tiles * GEN_BLOCK_M;
        dim3 block(GEN_BLOCK_N, GEN_BLOCK_M);

        if (full_m_tiles > 0) {
            dim3 full_grid((N + GEN_BLOCK_N - 1) / GEN_BLOCK_N, full_m_tiles);
            gemm_general_kernel<false><<<full_grid, block>>>(
                a_ptr, b_q_ptr, b_scale_sh_ptr, out_ptr, M, N, K, k_scale, 0
            );
        }

        if (tail_m > 0) {
            dim3 tail_grid((N + GEN_BLOCK_N - 1) / GEN_BLOCK_N, 1);
            gemm_general_kernel<true><<<tail_grid, block>>>(
                a_ptr, b_q_ptr, b_scale_sh_ptr, out_ptr, M, N, K, k_scale, full_m_tiles * GEN_BLOCK_M
            );
        }
    }

    hipError_t err = hipGetLastError();
    if (err != hipSuccess) {
        throw std::runtime_error(hipGetErrorString(err));
    }
}
"""


try:
    _MODULE = load_inline(
        name="mxfp4_quant_gemm_vG",
        cpp_sources=[CPP_SRC],
        cuda_sources=[CUDA_SRC],
        functions=["quant_gemm_mxfp4_vG"],
        extra_cflags=["-O3"],
        extra_cuda_cflags=["-O3", "-std=c++20", "--offload-arch=gfx950"],
        verbose=False,
    )
except Exception:
    _MODULE = None


_OUT_CACHE: dict[tuple[int, int, int, str, int | None], torch.Tensor] = {}


def _get_output_buffer(
    m: int,
    n: int,
    k: int,
    device: torch.device,
) -> torch.Tensor:
    key = (m, n, k, device.type, device.index)
    cached = _OUT_CACHE.get(key)
    if cached is not None:
        return cached

    out = torch.empty((m, n), device=device, dtype=torch.bfloat16)
    _OUT_CACHE[key] = out
    return out


def _fallback_kernel(data: input_t) -> output_t:
    import aiter
    from aiter import dtypes
    from aiter.ops.triton.quant import dynamic_mxfp4_quant
    from aiter.utility.fp4_utils import e8m0_shuffle

    a, _b, _b_q, b_shuffle, b_scale_sh = data
    a_q, a_scale_sh = dynamic_mxfp4_quant(a.contiguous())
    a_q = a_q.view(dtypes.fp4x2)
    a_scale_sh = e8m0_shuffle(a_scale_sh).view(dtypes.fp8_e8m0)
    return aiter.gemm_a4w4(
        a_q,
        b_shuffle,
        a_scale_sh,
        b_scale_sh,
        dtype=dtypes.bf16,
        bpreshuffle=True,
    )


def custom_kernel(data: input_t) -> output_t:
    if _MODULE is None:
        return _fallback_kernel(data)

    a, _b, b_q, _b_shuffle, b_scale_sh = data
    if not a.is_contiguous():
        a = a.contiguous()
    if not b_q.is_contiguous():
        b_q = b_q.contiguous()
    if not b_scale_sh.is_contiguous():
        b_scale_sh = b_scale_sh.contiguous()

    out = _get_output_buffer(a.shape[0], b_q.shape[0], a.shape[1], a.device)
    _MODULE.quant_gemm_mxfp4_vG(a, b_q, b_scale_sh, out)
    return out
scrolls · 646 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