Skip to content
KernelIndex
Search⌘K

submission 609571

nicholaswilde_08140 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-609571?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
15.5µs
#629 of 1143
2026-03-22

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:f291514ff502140ff95ab8c5e563aff47410b241d0a1eee157f5850e9190962f
license declaredunknown
license concludedunknown
authorsnicholaswilde_08140
imported2026-08-26

Techniques

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

fp4int fp4;

Kernel source

submission.py117 lines
import torch
from task import input_t, output_t

CPP_WRAPPER = """
void quant_bf16_mxfp4(torch::Tensor x, torch::Tensor x_q, torch::Tensor x_s, int m, int n);
"""

CUDA_SRC = """
#include <hip/hip_runtime.h>
#include <hip/amd_detail/amd_hip_bf16.h>

__global__ void kernel_quant_bf16_mxfp4(__hip_bfloat16 *x, char *x_q, char *x_s, int m, int n) {

    if (blockIdx.x * 512 + threadIdx.x * 8 >= m * n) return;

    __hip_bfloat16 elements[8];
    *(float4 *)elements = *((float4*)(x + blockIdx.x * 512 + threadIdx.x * 8));

    float absmax = 0.0f;
    for (int i = 0; i < 8; i++) {
        float element = __bfloat162float(elements[i]);
        absmax = fmaxf(absmax, fabsf(element));
    }
    for (int off = 1; off < 4; off <<= 1) {
        absmax = fmaxf(absmax, __shfl_down(absmax, off));
    }
    absmax = __shfl(absmax, (threadIdx.x / 4) * 4);

    float scale;
    unsigned int amax_bits = (__float_as_uint(absmax) + 0x200000u) & 0x7F800000u;
    scale = __uint_as_float(amax_bits) / 4.0f;

    int exp = (amax_bits >> 23) - 2;
    if (threadIdx.x % 4 == 0) {
        // x_s[blockIdx.x * 16 + threadIdx.x / 4] = (char)exp;
        int this_m = blockIdx.x * 512 / n;
        int this_n = (blockIdx.x * 512 + threadIdx.x * 8) % n / 32;
        x_s[this_m / 32 * 32 * (n / 32) + this_n / 8 * 256 + (this_n % 8) % 4 * 64 + (this_m % 32) % 16 * 4 + (this_n % 8) / 4 * 2 + (this_m % 32) / 16] = (char)exp;
    }

    int fp4x8 = 0;
    for (int i = 0; i < 8; i++) {
        float element = __bfloat162float(elements[i]);
        int sign = (element >= 0) ? 0x0 : 0x8;
        float q = element / scale;
        int fp4;
        if (fabs(q) > 5.0) fp4 = sign | 0x7;
        else if (fabs(q) >= 3.5) fp4 = sign | 0x6;
        else if (fabs(q) > 2.5) fp4 = sign | 0x5;
        else if (fabs(q) >= 1.75) fp4 = sign | 0x4;
        else if (fabs(q) > 1.25) fp4 = sign | 0x3;
        else if (fabs(q) >= 0.75) fp4 = sign | 0x2;
        else if (fabs(q) > 0.25) fp4 = sign | 0x1;
        else fp4 = 0x0;
        fp4x8 |= (fp4 << (i * 4));
    }
    *((int *)(x_q + blockIdx.x * 256 + threadIdx.x * 4)) = fp4x8;
}

void quant_bf16_mxfp4(torch::Tensor x, torch::Tensor x_q, torch::Tensor x_s, int m, int n) {
    __hip_bfloat16 *x_ptr = (__hip_bfloat16 *)x.data_ptr();
    char *x_q_ptr = (char *)x_q.data_ptr();
    char *x_s_ptr = (char *)x_s.data_ptr();

    int grid = (m * n + 511) / 512;
    hipLaunchKernelGGL(kernel_quant_bf16_mxfp4, dim3(grid), dim3(64), 0, 0, x_ptr, x_q_ptr, x_s_ptr, m, n);
}
"""

import os
os.environ['PYTORCH_ROCM_ARCH'] = 'gfx950'
os.environ["CXX"] = "clang++"
from torch.utils.cpp_extension import load_inline

module = load_inline(
    name='mxfp4_mm',
    cpp_sources=[CPP_WRAPPER],
    cuda_sources=[CUDA_SRC],
    functions=['quant_bf16_mxfp4'],
    verbose=True,
    extra_cuda_cflags=["--offload-arch=gfx950", "-std=c++20", "-O3"],
)

def custom_kernel(data: input_t) -> output_t:

    import aiter
    from aiter import QuantType, dtypes
    from aiter.ops.triton.quant import dynamic_mxfp4_quant
    from aiter.utility.fp4_utils import e8m0_shuffle

    def _quant_mxfp4(x, shuffle=True):
        x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
        if shuffle:
            bs_e8m0 = e8m0_shuffle(bs_e8m0)
        return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)

    A, B, B_q, B_shuffle, B_scale_sh = data
    A = A.contiguous()
    B = B.contiguous()
    m, k = A.shape
    n, _ = B.shape

    # A_q, A_scale_sh = _quant_mxfp4(A, shuffle=True)

    A_q = torch.empty((m, k // 2), dtype=dtypes.fp4x2, device=A.device)
    A_scale_sh = torch.empty((256, k // 32), dtype=dtypes.fp8_e8m0, device=A.device)
    module.quant_bf16_mxfp4(A, A_q, A_scale_sh, m, k)

    out_gemm = aiter.gemm_a4w4(
        A_q,
        B_shuffle,
        A_scale_sh,
        B_scale_sh,
        dtype=dtypes.bf16,
        bpreshuffle=True,
    )
    return out_gemm
scrolls · 117 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