Skip to content
KernelIndex
Search⌘K

submission 675976

phoenixdna · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v41_hybrid_direct_bshuf.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-675976?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
11.6µs
#341 of 1143
2026-03-31

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:579bb3e619433df4d5287589c56c7af561d59e947897440ddae0d4fe7a172136
license declaredunknown
license concludedunknown
authorsphoenixdna
imported2026-08-26

Techniques

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

shared-memory__shared__ float s_abs[32 * WARP_SLOTS];
split-kf.write("cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio\n")

Kernel source

submission_v41_hybrid_direct_bshuf.py396 lines
import os
import tempfile

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


_RUNTIME = None
_NATIVE_MOD = None
_A_BUFFER_CACHE = {}

_CFG_ROWS = [
    (256, 4, 2880, 512, 21, 0, 5.05, "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"),
    (256, 16, 2112, 7168, 21, 1, 12.90, "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"),
    (256, 32, 4096, 512, 21, 0, 5.40, "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"),
    (256, 32, 2880, 512, 21, 0, 5.10, "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"),
]

_INLINE_CPP = r"""
#include <torch/extension.h>
void quant_mxfp4_a(torch::Tensor x, torch::Tensor x_fp4, torch::Tensor scale_sh);
"""

_INLINE_HIP = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <hip/amd_detail/amd_hip_bf16.h>
#include <cstdint>

namespace {

__device__ __forceinline__ uint32_t f32_as_u32(float x) {
    union {
        float f;
        uint32_t u;
    } v;
    v.f = x;
    return v.u;
}

__device__ __forceinline__ float u32_as_f32(uint32_t x) {
    union {
        float f;
        uint32_t u;
    } v;
    v.u = x;
    return v.f;
}

__device__ __forceinline__ float bf16_to_f32(uint16_t x) {
    __hip_bfloat16 v;
    *reinterpret_cast<uint16_t*>(&v) = x;
    return static_cast<float>(v);
}

__device__ __forceinline__ uint8_t quantize_e2m1(float x) {
    constexpr int EXP_BIAS_FP32 = 127;
    constexpr int EXP_BIAS_FP4 = 1;
    constexpr int MBITS_F32 = 23;
    constexpr int MBITS_FP4 = 1;
    constexpr uint8_t MAX_INT = 0x7;
    constexpr uint8_t SIGN_MASK = 0x8;
    constexpr uint32_t MAGIC_ADDER = (1u << 21) - 1u;
    constexpr float MAX_NORMAL = 6.0f;
    constexpr float MIN_NORMAL = 1.0f;
    constexpr uint32_t DENORM_MASK_INT =
        static_cast<uint32_t>(((EXP_BIAS_FP32 - EXP_BIAS_FP4) + (MBITS_F32 - MBITS_FP4) + 1) << MBITS_F32);
    constexpr int32_t VAL_TO_ADD =
        ((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_F32) + static_cast<int32_t>(MAGIC_ADDER);

    uint32_t ux = f32_as_u32(x);
    uint8_t sign_lp = static_cast<uint8_t>((ux >> 28) & SIGN_MASK);
    ux &= 0x7FFFFFFFu;
    float ax = u32_as_f32(ux);

    if (ax >= MAX_NORMAL) {
        return static_cast<uint8_t>(sign_lp | MAX_INT);
    }

    if (ax < MIN_NORMAL) {
        float denormal_f = ax + u32_as_f32(DENORM_MASK_INT);
        uint32_t denormal_u = f32_as_u32(denormal_f);
        uint8_t denormal_x = static_cast<uint8_t>(denormal_u - DENORM_MASK_INT);
        return static_cast<uint8_t>(sign_lp | (denormal_x & MAX_INT));
    }

    uint32_t mant_odd = (ux >> (MBITS_F32 - MBITS_FP4)) & 1u;
    uint32_t normal_x = ux + static_cast<uint32_t>(VAL_TO_ADD) + mant_odd;
    uint8_t e2m1 = static_cast<uint8_t>(normal_x >> (MBITS_F32 - MBITS_FP4));
    return static_cast<uint8_t>(sign_lp | (e2m1 & MAX_INT));
}

__device__ __forceinline__ int scale_shuffle_offset(int row, int col, int scale_n_pad) {
    int offs_0 = row / 32;
    int row_in_32 = row % 32;
    int offs_1 = row_in_32 / 16;
    int offs_2 = row_in_32 % 16;
    int offs_3 = col / 8;
    int col_in_8 = col % 8;
    int offs_4 = col_in_8 / 4;
    int offs_5 = col_in_8 % 4;
    return offs_1 + offs_4 * 2 + offs_2 * 4 + offs_5 * 64 + offs_3 * 256 + offs_0 * 32 * scale_n_pad;
}

template <int ROWS_PER_BLOCK, int GROUPS_PER_BLOCK>
__global__ void quant_mxfp4_a_kernel(
    const uint16_t* __restrict__ x,
    uint8_t* __restrict__ x_fp4,
    uint8_t* __restrict__ scale_sh,
    int m,
    int k,
    int scale_n_valid,
    int scale_n_pad
) {
    constexpr int WARP_SLOTS = ROWS_PER_BLOCK * GROUPS_PER_BLOCK;
    __shared__ float s_abs[32 * WARP_SLOTS];
    __shared__ uint8_t s_scale[WARP_SLOTS];

    constexpr int GROUP_SIZE = 32;
    constexpr int PAIRS_PER_GROUP = GROUP_SIZE / 2;

    int tid = static_cast<int>(threadIdx.x);
    int lane = tid & 31;
    int warp = tid >> 5;
    int row_in_block = warp / GROUPS_PER_BLOCK;
    int group_in_block = warp % GROUPS_PER_BLOCK;
    int warp_slot = row_in_block * GROUPS_PER_BLOCK + group_in_block;
    int row = static_cast<int>(blockIdx.y) * ROWS_PER_BLOCK + row_in_block;
    int scale_col = static_cast<int>(blockIdx.x) * GROUPS_PER_BLOCK + group_in_block;
    bool active_group = row < m && scale_col < scale_n_valid;

    float value = 0.0f;
    if (active_group) {
        int x_idx = row * k + scale_col * GROUP_SIZE + lane;
        value = bf16_to_f32(x[x_idx]);
    }

    int group_base = warp_slot * GROUP_SIZE;
    s_abs[group_base + lane] = active_group ? fabsf(value) : 0.0f;
    __syncthreads();

    if (lane < 16) {
        float a = s_abs[group_base + lane];
        float b = s_abs[group_base + lane + 16];
        s_abs[group_base + lane] = a > b ? a : b;
    }
    __syncthreads();
    if (lane < 8) {
        float a = s_abs[group_base + lane];
        float b = s_abs[group_base + lane + 8];
        s_abs[group_base + lane] = a > b ? a : b;
    }
    __syncthreads();
    if (lane < 4) {
        float a = s_abs[group_base + lane];
        float b = s_abs[group_base + lane + 4];
        s_abs[group_base + lane] = a > b ? a : b;
    }
    __syncthreads();
    if (lane < 2) {
        float a = s_abs[group_base + lane];
        float b = s_abs[group_base + lane + 2];
        s_abs[group_base + lane] = a > b ? a : b;
    }
    __syncthreads();
    if (lane == 0) {
        float a = s_abs[group_base];
        float b = s_abs[group_base + 1];
        float max_abs = a > b ? a : b;
        uint32_t max_bits = f32_as_u32(max_abs);
        max_bits = (max_bits + 0x00200000u) & 0xFF800000u;
        uint8_t scale_byte = 0;
        if (max_bits != 0) {
            scale_byte = static_cast<uint8_t>(((max_bits >> 23) & 0xFFu) - 2u);
        }
        s_scale[warp_slot] = scale_byte;
        if (active_group) {
            int scale_idx = scale_shuffle_offset(row, scale_col, scale_n_pad);
            scale_sh[scale_idx] = scale_byte;
        }
    }
    __syncthreads();

    if (lane < PAIRS_PER_GROUP && active_group) {
        uint8_t scale_byte = s_scale[warp_slot];
        float inv_scale = 0.0f;
        if (scale_byte != 0) {
            float scale_f = u32_as_f32(static_cast<uint32_t>(scale_byte) << 23);
            inv_scale = 1.0f / scale_f;
        }
        int pair_base = row * k + scale_col * GROUP_SIZE + lane * 2;
        float x0 = bf16_to_f32(x[pair_base]) * inv_scale;
        float x1 = bf16_to_f32(x[pair_base + 1]) * inv_scale;
        uint8_t q0 = quantize_e2m1(x0);
        uint8_t q1 = quantize_e2m1(x1);
        int out_idx = row * (k / 2) + scale_col * PAIRS_PER_GROUP + lane;
        x_fp4[out_idx] = static_cast<uint8_t>(q0 | (q1 << 4));
    }
}

}  // namespace

void quant_mxfp4_a(torch::Tensor x, torch::Tensor x_fp4, torch::Tensor scale_sh) {
    TORCH_CHECK(x.is_cuda(), "x must be a CUDA tensor");
    TORCH_CHECK(x_fp4.is_cuda(), "x_fp4 must be a CUDA tensor");
    TORCH_CHECK(scale_sh.is_cuda(), "scale_sh must be a CUDA tensor");
    TORCH_CHECK(x.scalar_type() == at::ScalarType::BFloat16, "x must be bf16");
    TORCH_CHECK(x_fp4.scalar_type() == at::ScalarType::Byte, "x_fp4 must be uint8");
    TORCH_CHECK(scale_sh.scalar_type() == at::ScalarType::Byte, "scale_sh must be uint8");
    TORCH_CHECK(x.dim() == 2, "x must be 2D");
    TORCH_CHECK(x.is_contiguous(), "x must be contiguous");
    TORCH_CHECK(x_fp4.is_contiguous(), "x_fp4 must be contiguous");
    TORCH_CHECK(scale_sh.is_contiguous(), "scale_sh must be contiguous");

    int m = static_cast<int>(x.size(0));
    int k = static_cast<int>(x.size(1));
    int scale_n_valid = k / 32;
    int scale_n_pad = static_cast<int>(scale_sh.size(1));

    TORCH_CHECK(k % 64 == 0, "k must be divisible by 64");
    TORCH_CHECK(x_fp4.size(0) == x.size(0), "x_fp4 row mismatch");
    TORCH_CHECK(x_fp4.size(1) == x.size(1) / 2, "x_fp4 col mismatch");

    if (k >= 1536) {
        dim3 grid((scale_n_valid + 1) / 2, (m + 3) / 4);
        dim3 block(256);
        hipLaunchKernelGGL(
            HIP_KERNEL_NAME(quant_mxfp4_a_kernel<4, 2>),
            grid,
            block,
            0,
            0,
            reinterpret_cast<const uint16_t*>(x.data_ptr<at::BFloat16>()),
            x_fp4.data_ptr<uint8_t>(),
            scale_sh.data_ptr<uint8_t>(),
            m,
            k,
            scale_n_valid,
            scale_n_pad
        );
    } else {
        dim3 grid((scale_n_valid + 1) / 2, m);
        dim3 block(64);
        hipLaunchKernelGGL(
            HIP_KERNEL_NAME(quant_mxfp4_a_kernel<1, 2>),
            grid,
            block,
            0,
            0,
            reinterpret_cast<const uint16_t*>(x.data_ptr<at::BFloat16>()),
            x_fp4.data_ptr<uint8_t>(),
            scale_sh.data_ptr<uint8_t>(),
            m,
            k,
            scale_n_valid,
            scale_n_pad
        );
    }
}
"""


def _get_native_mod():
    global _NATIVE_MOD
    if _NATIVE_MOD is None:
        _NATIVE_MOD = load_inline(
            name="mxfp4_mm_native_hip_quant_v41_hybrid_direct_bshuf",
            cpp_sources=[_INLINE_CPP],
            cuda_sources=[_INLINE_HIP],
            functions=["quant_mxfp4_a"],
            extra_cuda_cflags=["--offload-arch=gfx950", "-O3", "-std=c++20"],
            verbose=False,
        )
    return _NATIVE_MOD


def _ensure_runtime():
    global _RUNTIME
    if _RUNTIME is not None:
        return _RUNTIME

    cfg_path = os.path.join(tempfile.gettempdir(), "aiter_mxfp4_mm_submission.csv")
    if not os.path.exists(cfg_path):
        with open(cfg_path, "w", encoding="ascii", newline="") as f:
            f.write("cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio\n")
            for cu_num, m, n, k, kernel_id, split_k, us, kernel_name in _CFG_ROWS:
                f.write(
                    f"{cu_num},{m},{n},{k},{kernel_id},{split_k},{us},{kernel_name},0.0,0.0,0.0\n"
                )

    default_cfg = "/home/runner/aiter/aiter/configs/a4w4_blockscale_tuned_gemm.csv"
    os.environ["AITER_CONFIG_GEMM_A4W4"] = cfg_path + os.pathsep + default_cfg

    import aiter
    from aiter import dtypes

    _RUNTIME = (aiter, dtypes)
    return _RUNTIME


def _native_quant_a(A: torch.Tensor, dtypes):
    mod = _get_native_mod()
    m, k = A.shape
    scale_m_pad = ((m + 255) // 256) * 256
    scale_n_valid = k // 32
    scale_n_pad = ((scale_n_valid + 7) // 8) * 8

    cache_key = (A.device.index, m, k)
    cached = _A_BUFFER_CACHE.get(cache_key)
    if cached is None:
        A_q_u8 = torch.empty((m, k // 2), dtype=torch.uint8, device=A.device)
        A_scale_u8 = torch.zeros((scale_m_pad, scale_n_pad), dtype=torch.uint8, device=A.device)
        cached = (
            A_q_u8,
            A_scale_u8,
            A_q_u8.view(dtypes.fp4x2),
            A_scale_u8.view(dtypes.fp8_e8m0),
        )
        _A_BUFFER_CACHE[cache_key] = cached

    A_q_u8, A_scale_u8, A_q_view, A_scale_view = cached
    mod.quant_mxfp4_a(A, A_q_u8, A_scale_u8)
    return A_q_view, A_scale_view


def _prepare_triton_weight_from_task(B_shuffle: torch.Tensor, B_scale_sh: torch.Tensor):
    if B_shuffle.dtype != torch.uint8:
        b_shuf_u8 = B_shuffle.view(torch.uint8)
    else:
        b_shuf_u8 = B_shuffle

    if B_scale_sh.dtype != torch.uint8:
        b_scale_u8 = B_scale_sh.view(torch.uint8)
    else:
        b_scale_u8 = B_scale_sh

    n = b_shuf_u8.shape[0]
    k_pack = b_shuf_u8.shape[1]
    k_scale = k_pack // 16

    # The task already provides B_shuffle = shuffle_weight(B_q, (16, 16)).
    # Triton consumes the same bytes but with a different 2D view.
    w_triton = b_shuf_u8.contiguous().view(n // 16, k_pack * 16)

    # fp4_utils.e8m0_shuffle() and Triton's shuffle_scales() use the same
    # underlying byte ordering; only the final 2D view differs.
    w_scales_triton = b_scale_u8[:n, :k_scale].contiguous().view(n // 32, k_scale * 32)
    return w_triton, w_scales_triton


def custom_kernel(data: input_t) -> output_t:
    aiter, dtypes = _ensure_runtime()
    A, _, _, B_shuffle, B_scale_sh = data
    if not A.is_contiguous():
        A = A.contiguous()

    m, k = A.shape
    n = B_shuffle.shape[0]

    triton_shapes = {
        (4, 2880, 512),
        (16, 2112, 7168),
        (32, 4096, 512),
        (32, 2880, 512),
    }

    if (m, n, k) in triton_shapes:
        from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle

        w_triton, w_scales_triton = _prepare_triton_weight_from_task(
            B_shuffle.contiguous(), B_scale_sh.contiguous()
        )
        return gemm_a16wfp4_preshuffle(
            A,
            w_triton,
            w_scales_triton,
            prequant=True,
            dtype=torch.bfloat16,
        )

    A_q, A_scale_sh = _native_quant_a(A, dtypes)

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