Skip to content
KernelIndex
Search⌘K

submission 666540

Jingze · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:31ef14c214dae0d9e19b647788533d77babef9d716583cb157d169c3468b2847
license declaredunknown
license concludedunknown
authorsJingze
imported2026-08-26

Techniques

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

fp4Converts given x (in fp32) to mxfp4 format.
num-warps = 1NUM_WARPS = 1
shared-memorysmem_a = gl.allocate_shared_memory(
split-kSPLITK_BLOCK_SIZE: gl.constexpr,
stages = 1NUM_STAGES = 1
tile-m = 64BLOCK_SIZE_M = 64
tile-n = 32BLOCK_SIZE_N = 32

Kernel source

submission_gluon.py2559 lines
from typing import Optional, Any
import functools
import json
import math
import os
import time


# os.environ.setdefault("TORCH_EXTENSIONS_DIR", "/tmp/torch_extensions")
# os.environ.setdefault("MAX_JOBS", "16")


import torch

import triton
import triton.language as tl

from triton.experimental import gluon
from triton.experimental.gluon import language as gl


MXFP4_GROUP_SIZE = 32


_HIP_QUANT_INLINE_MODULE = None
_HIP_QUANT_INLINE_LOAD_ERROR = None


def _inline_quant_backend() -> Optional[str]:
    if not torch.cuda.is_available():
        return None

    if torch.version.hip is not None:
        try:
            device_name = torch.cuda.get_device_name(torch.cuda.current_device()).lower()
        except Exception:
            return None

        if any(token in device_name for token in ("amd", "radeon", "instinct", "mi")):
            return "hip"

    if torch.version.cuda is not None:
        return "cuda"

    return None


def _dynamic_mxfp4_quant_inline_sources() -> tuple[str, str]:
    cpp_source = """
    #include <torch/extension.h>
    #include <vector>

    std::vector<torch::Tensor> dynamic_mxfp4_quantize(torch::Tensor x);
    """

    gpu_source = r"""
    #include <ATen/ATen.h>
    #include <ATen/cuda/CUDAContext.h>
    #include <c10/util/Exception.h>
    #include <c10/util/BFloat16.h>
    #include <cmath>
    #include <cstdint>
    #include <stdexcept>
    #include <vector>

    #if defined(__HIP_PLATFORM_AMD__) || defined(USE_ROCM)
    #include <hip/hip_runtime.h>
    #define MXFP4_GPU_BACKEND_HIP 1
    #else
    #include <cuda_runtime.h>
    #define MXFP4_GPU_BACKEND_CUDA 1
    #endif

    namespace {

    #define MX_CAT2(a, b) a##b
    #define MX_CAT3(a, b, c) a##b##c

    #if defined(MXFP4_GPU_BACKEND_CUDA)
    using mx_queue_t = MX_CAT2(cudaSt, ream_t);

    __device__ __forceinline__ float warp_shuffle_down(float value, int offset) {
        return __shfl_down_sync(0xFFFFFFFFu, value, offset, 32);
    }

    __device__ __forceinline__ float warp_shuffle(float value, int src_lane) {
        return __shfl_sync(0xFFFFFFFFu, value, src_lane, 32);
    }

    __device__ __forceinline__ uint32_t warp_shuffle_xor_u32(uint32_t value, int lane_mask) {
        return __shfl_xor_sync(0xFFFFFFFFu, value, lane_mask, 32);
    }

    __device__ __forceinline__ uint8_t warp_shuffle_xor(uint8_t value, int lane_mask) {
        return static_cast<uint8_t>(
            __shfl_xor_sync(0xFFFFFFFFu, static_cast<unsigned int>(value), lane_mask, 32)
        );
    }
    #else
    using mx_queue_t = MX_CAT2(hipSt, ream_t);

    __device__ __forceinline__ float warp_shuffle_down(float value, int offset) {
        return __shfl_down(value, offset, 32);
    }

    __device__ __forceinline__ float warp_shuffle(float value, int src_lane) {
        return __shfl(value, src_lane, 32);
    }

    __device__ __forceinline__ uint32_t warp_shuffle_xor_u32(uint32_t value, int lane_mask) {
        return static_cast<uint32_t>(__shfl_xor(value, lane_mask, 32));
    }

    __device__ __forceinline__ uint8_t warp_shuffle_xor(uint8_t value, int lane_mask) {
        return static_cast<uint8_t>(__shfl_xor(value, lane_mask, 32));
    }
    #endif

    __device__ __forceinline__ float bf16_to_float(const uint16_t v) {
        return __uint_as_float(static_cast<uint32_t>(v) << 16);
    }

    __device__ __forceinline__ int clamp_int(const int value, const int lo, const int hi) {
        return value < lo ? lo : (value > hi ? hi : value);
    }

    __device__ __forceinline__ uint32_t float_to_bits(float value) {
        return __float_as_uint(value);
    }

    __device__ __forceinline__ float bits_to_float(uint32_t value) {
        return __uint_as_float(value);
    }

    __device__ __forceinline__ float exact_pow2_neg_exp(int scale_unbiased) {
        return bits_to_float(static_cast<uint32_t>(127 - scale_unbiased) << 23);
    }

    __device__ __forceinline__ uint8_t quantize_mxfp4_bits(float value, float quant_scale) {
        constexpr uint32_t FP32_SIGN_MASK = 0x80000000u;
        constexpr uint32_t FP32_ONE_BITS = 0x3F800000u;
        constexpr uint32_t FP32_SIX_BITS = 0x40C00000u;
        constexpr int MBITS_F32 = 23;
        constexpr int MBITS_FP4 = 1;
        constexpr int EBITS_F32 = 8;
        constexpr int EBITS_FP4 = 2;
        constexpr int EXP_BIAS_FP32 = 127;
        constexpr int EXP_BIAS_FP4 = 1;
        constexpr int DENORM_EXP = (EXP_BIAS_FP32 - EXP_BIAS_FP4) + (MBITS_F32 - MBITS_FP4) + 1;
        constexpr uint32_t DENORM_MASK_INT = static_cast<uint32_t>(DENORM_EXP) << MBITS_F32;
        constexpr uint32_t VAL_TO_ADD =
            (static_cast<uint32_t>(EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_F32) + ((1u << 21) - 1u);
        const float DENORM_MASK_FLOAT = bits_to_float(DENORM_MASK_INT);

        const float scaled = value * quant_scale;
        const uint32_t scaled_bits = float_to_bits(scaled);
        const uint32_t sign = scaled_bits & FP32_SIGN_MASK;
        const uint32_t abs_bits = scaled_bits ^ sign;

        uint8_t e2m1 = 0x7;
        if (abs_bits < FP32_SIX_BITS) {
            if (abs_bits < FP32_ONE_BITS) {
                const float abs_value = bits_to_float(abs_bits);
                uint32_t denormal_x = float_to_bits(abs_value + DENORM_MASK_FLOAT);
                denormal_x -= DENORM_MASK_INT;
                e2m1 = static_cast<uint8_t>(denormal_x);
            } else {
                const uint32_t mant_odd = (abs_bits >> (MBITS_F32 - MBITS_FP4)) & 1u;
                const uint32_t normal_u32 = abs_bits + VAL_TO_ADD + mant_odd;
                e2m1 = static_cast<uint8_t>(normal_u32 >> (MBITS_F32 - MBITS_FP4));
            }
        }

        const uint8_t sign_lp = static_cast<uint8_t>(sign >> (MBITS_F32 + EBITS_F32 - MBITS_FP4 - EBITS_FP4));
        return static_cast<uint8_t>(e2m1 | sign_lp);
    }

    __global__ void dynamic_mxfp4_quantize_kernel(
        const uint16_t* __restrict__ x,
        uint8_t* __restrict__ x_fp4,
        uint8_t* __restrict__ scales,
        const int m,
        const int n,
        const int packed_n,
        const int quant_blocks
    ) {
        const int lane = threadIdx.x;
        const int warp_in_cta = threadIdx.y;
        const int block_n = blockIdx.x * blockDim.y + warp_in_cta;
        const int row = blockIdx.y;

        if (row >= m || block_n >= quant_blocks) {
            return;
        }

        const int base_n = block_n * 32;
        const float value = bf16_to_float(x[row * n + base_n + lane]);
        float abs_value = fabsf(value);

        #pragma unroll
        for (int offset = 16; offset > 0; offset >>= 1) {
            abs_value = fmaxf(abs_value, warp_shuffle_down(abs_value, offset));
        }

        const float amax = warp_shuffle(abs_value, 0);
        uint32_t amax_bits = float_to_bits(amax);
        amax_bits = (amax_bits + 0x200000u) & 0xFF800000u;
        int scale_unbiased = -127;
        if (amax_bits != 0) {
            scale_unbiased = static_cast<int>((amax_bits >> 23) & 0xFFu) - 127 - 2;
            scale_unbiased = clamp_int(scale_unbiased, -127, 127);
        }
        const float quant_scale = exact_pow2_neg_exp(scale_unbiased);

        const uint8_t q = quantize_mxfp4_bits(value, quant_scale);
        if ((lane & 1) == 0) {
            const float partner_value = bf16_to_float(x[row * n + base_n + lane + 1]);
            const uint8_t q_hi = quantize_mxfp4_bits(partner_value, quant_scale);
            const int packed_col = block_n * 16 + (lane >> 1);
            x_fp4[row * packed_n + packed_col] = static_cast<uint8_t>(q | (q_hi << 4));
        }

        if (lane == 0) {
            #if defined(MXFP4_GPU_BACKEND_HIP)
            const int stored_scale_unbiased = scale_unbiased;
            #else
            const int stored_scale_unbiased = scale_unbiased > 0 ? scale_unbiased : 0;
            #endif
            scales[row * quant_blocks + block_n] = static_cast<uint8_t>(stored_scale_unbiased + 127);
        }
    }

    }  // namespace

    std::vector<at::Tensor> dynamic_mxfp4_quantize(at::Tensor x) {
        const auto m = static_cast<int>(x.size(0));
        const auto n = static_cast<int>(x.size(1));

        const int packed_n = n / 2;
        const int quant_blocks = n / 32;
        const int groups_per_cta = quant_blocks >= 8 ? 8 : (quant_blocks >= 4 ? 4 : (quant_blocks >= 2 ? 2 : 1));

        auto options = at::TensorOptions().device(x.device()).dtype(at::kByte);
        auto x_fp4 = at::empty({m, packed_n}, options);
        auto scales = at::empty({m, quant_blocks}, options);

        dim3 block(32, groups_per_cta);
        dim3 grid((quant_blocks + groups_per_cta - 1) / groups_per_cta, m);

        #if defined(MXFP4_GPU_BACKEND_HIP)
        auto q = at::cuda::MX_CAT3(getCurrentHIPSt, ream, )();
        hipLaunchKernelGGL(
            dynamic_mxfp4_quantize_kernel,
            grid,
            block,
            0,
            static_cast<mx_queue_t>(q),
            reinterpret_cast<const uint16_t*>(x.data_ptr<c10::BFloat16>()),
            x_fp4.data_ptr<uint8_t>(),
            scales.data_ptr<uint8_t>(),
            m,
            n,
            packed_n,
            quant_blocks
        );

        const auto err = hipGetLastError();
        if (err != hipSuccess) {
            throw std::runtime_error(hipGetErrorString(err));
        }
        #else
        auto q = at::cuda::MX_CAT3(getCurrentCUDASt, ream, )();
        dynamic_mxfp4_quantize_kernel<<<grid, block, 0, static_cast<mx_queue_t>(q)>>>(
            reinterpret_cast<const uint16_t*>(x.data_ptr<c10::BFloat16>()),
            x_fp4.data_ptr<uint8_t>(),
            scales.data_ptr<uint8_t>(),
            m,
            n,
            packed_n,
            quant_blocks
        );

        const auto err = cudaGetLastError();
        if (err != cudaSuccess) {
            throw std::runtime_error(cudaGetErrorString(err));
        }
        #endif

        return {x_fp4, scales};
    }
    """
    return cpp_source, gpu_source


def _load_dynamic_mxfp4_quant_inline_module():
    global _HIP_QUANT_INLINE_MODULE, _HIP_QUANT_INLINE_LOAD_ERROR

    if _HIP_QUANT_INLINE_MODULE is not None:
        return _HIP_QUANT_INLINE_MODULE
    if _HIP_QUANT_INLINE_LOAD_ERROR is not None:
        raise RuntimeError(_HIP_QUANT_INLINE_LOAD_ERROR)

    backend = _inline_quant_backend()
    if backend is None:
        _HIP_QUANT_INLINE_LOAD_ERROR = "GPU inline quantization is unavailable in the current environment"
        raise RuntimeError(_HIP_QUANT_INLINE_LOAD_ERROR)

    try:
        from torch.utils.cpp_extension import load_inline

        cpp_source, gpu_source = _dynamic_mxfp4_quant_inline_sources()
        extra_cuda_cflags = ["-O3", "-std=c++17"]
        if backend == "cuda":
            extra_cuda_cflags.append("--use_fast_math")
        elif backend == "hip":
            extra_cuda_cflags.append("-ffast-math")

        _HIP_QUANT_INLINE_MODULE = load_inline(
            name=f"mxfp4_dynamic_quant_inline_v4_{backend}",
            cpp_sources=cpp_source,
            cuda_sources=gpu_source,
            functions=["dynamic_mxfp4_quantize"],
            extra_cflags=["-O3", "-std=c++17"],
            extra_cuda_cflags=extra_cuda_cflags,
            with_cuda=True,
            verbose=False,
        )
        return _HIP_QUANT_INLINE_MODULE
    except Exception as exc:
        _HIP_QUANT_INLINE_LOAD_ERROR = str(exc)
        raise RuntimeError(_HIP_QUANT_INLINE_LOAD_ERROR) from exc


def dynamic_mxfp4_quant_inline(
    x: torch.Tensor, scaling_mode: str = "even"
) -> tuple[torch.Tensor, torch.Tensor]:
    if scaling_mode != "even":
        raise NotImplementedError(f"Unsupported scaling_mode: {scaling_mode}")
    x_input = x.contiguous()
    if x_input.dtype != torch.bfloat16:
        x_input = x_input.to(torch.bfloat16)
    module = _load_dynamic_mxfp4_quant_inline_module()
    x_fp4, blockscale_e8m0 = module.dynamic_mxfp4_quantize(x_input)
    return x_fp4, blockscale_e8m0


_DEFAULT_GEMM_CONFIGS = {
    "M_LEQ_8": {
        "BLOCK_SIZE_M": 8,
        "BLOCK_SIZE_N": 64,
        "BLOCK_SIZE_K": 256,
        "GROUP_SIZE_M": 4,
        "num_warps": 2,
        "num_stages": 2,
        "waves_per_eu": 1,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": None,
        "NUM_KSPLIT": 1,
    },
    "M_LEQ_31": {
        "BLOCK_SIZE_M": 16,
        "BLOCK_SIZE_N": 64,
        "BLOCK_SIZE_K": 256,
        "GROUP_SIZE_M": 4,
        "num_warps": 2,
        "num_stages": 2,
        "waves_per_eu": 1,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": None,
        "NUM_KSPLIT": 1,
    },
    "M_LEQ_32": {
        "BLOCK_SIZE_M": 32,
        "BLOCK_SIZE_N": 64,
        "BLOCK_SIZE_K": 256,
        "GROUP_SIZE_M": 4,
        "num_warps": 2,
        "num_stages": 2,
        "waves_per_eu": 1,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": None,
        "NUM_KSPLIT": 1,
    },
    "M_LEQ_64": {
        "BLOCK_SIZE_M": 32,
        "BLOCK_SIZE_N": 64,
        "BLOCK_SIZE_K": 256,
        "GROUP_SIZE_M": 4,
        "num_warps": 2,
        "num_stages": 2,
        "waves_per_eu": 1,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": None,
        "NUM_KSPLIT": 1,
    },
    "M_LEQ_128": {
        "BLOCK_SIZE_M": 32,
        "BLOCK_SIZE_N": 64,
        "BLOCK_SIZE_K": 256,
        "GROUP_SIZE_M": 4,
        "num_warps": 2,
        "num_stages": 2,
        "waves_per_eu": 1,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": None,
        "NUM_KSPLIT": 1,
    },
    "M_LEQ_256": {
        "BLOCK_SIZE_M": 32,
        "BLOCK_SIZE_N": 64,
        "BLOCK_SIZE_K": 256,
        "GROUP_SIZE_M": 4,
        "num_warps": 2,
        "num_stages": 2,
        "waves_per_eu": 1,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": None,
        "NUM_KSPLIT": 1,
    },
    "any": {
        "BLOCK_SIZE_M": 32,
        "BLOCK_SIZE_N": 64,
        "BLOCK_SIZE_K": 256,
        "GROUP_SIZE_M": 4,
        "num_warps": 2,
        "num_stages": 2,
        "waves_per_eu": 1,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": None,
        "NUM_KSPLIT": 1,
    },
}


_BENCHMARK_GEMM_CONFIGS = {
    # (4, 2880, 512): {
    #     "BLOCK_SIZE_M": 4,
    #     "BLOCK_SIZE_N": 64,
    #     "BLOCK_SIZE_K": 512,
    #     "GROUP_SIZE_M": 1,
    #     "num_warps": 2,
    #     "num_stages": 2,
    #     "waves_per_eu": 3,
    #     "matrix_instr_nonkdim": 16,
    #     "cache_modifier": None,
    #     "NUM_KSPLIT": 1,
    # }, # 10.1 ± 0.02 µs
    (4, 2880, 512): {
        "BLOCK_SIZE_M": 4,
        "BLOCK_SIZE_N": 32,
        "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 1,
        "num_warps": 4,
        "num_stages": 2,
        "waves_per_eu": 3,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
        "NUM_KSPLIT": 1,
    },
    (16, 2112, 7168): {
        "BLOCK_SIZE_M": 16,
        "BLOCK_SIZE_N": 32,
        "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 1,
        "num_warps": 4,
        "num_stages": 2,
        "waves_per_eu": 2,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
        "NUM_KSPLIT": 7,
    },
    (32, 4096, 512): {
        "BLOCK_SIZE_M": 32,
        "BLOCK_SIZE_N": 32,
        "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 1,
        "num_warps": 4,
        "num_stages": 2,
        "waves_per_eu": 2,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
        "NUM_KSPLIT": 1,
    },
    (32, 2880, 512): {
        "BLOCK_SIZE_M": 32,
        "BLOCK_SIZE_N": 32,
        "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 1,
        "num_warps": 4,
        "num_stages": 2,
        "waves_per_eu": 2,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
        "NUM_KSPLIT": 1,
    },
    (64, 7168, 2048): {
        "BLOCK_SIZE_M": 32,
        "BLOCK_SIZE_N": 32,
        "BLOCK_SIZE_K": 1024,
        "GROUP_SIZE_M": 1,
        "num_warps": 2,
        "num_stages": 2,
        "waves_per_eu": 2,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
        "NUM_KSPLIT": 1,
    },
    (256, 3072, 1536): {
        "BLOCK_SIZE_M": 128,
        "BLOCK_SIZE_N": 32,
        "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 1,
        "num_warps": 4,
        "num_stages": 2,
        "waves_per_eu": 2,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
        "NUM_KSPLIT": 1,
    },
}


def _select_default_config_by_m(M: int):
    if M <= 8:
        return dict(_DEFAULT_GEMM_CONFIGS["M_LEQ_8"])
    if M <= 31:
        return dict(_DEFAULT_GEMM_CONFIGS["M_LEQ_31"])
    if M <= 32:
        return dict(_DEFAULT_GEMM_CONFIGS["M_LEQ_32"])
    if M <= 64:
        return dict(_DEFAULT_GEMM_CONFIGS["M_LEQ_64"])
    if M <= 128:
        return dict(_DEFAULT_GEMM_CONFIGS["M_LEQ_128"])
    if M <= 256:
        return dict(_DEFAULT_GEMM_CONFIGS["M_LEQ_256"])
    return dict(_DEFAULT_GEMM_CONFIGS["any"])


@triton.jit
def _mxfp4_quant_op(
    x,
    BLOCK_SIZE_N,
    BLOCK_SIZE_M,
    MXFP4_QUANT_BLOCK_SIZE,
):
    """
    Converts given x (in fp32) to mxfp4 format.
    x: [BLOCK_SIZE_M, BLOCK_SIZE_N], fp32

    """
    EXP_BIAS_FP32: tl.constexpr = 127
    EXP_BIAS_FP4: tl.constexpr = 1
    EBITS_F32: tl.constexpr = 8
    EBITS_FP4: tl.constexpr = 2
    MBITS_F32: tl.constexpr = 23
    MBITS_FP4: tl.constexpr = 1

    max_normal: tl.constexpr = 6
    min_normal: tl.constexpr = 1

    NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
    x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)
    # Calculate scale
    amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
    amax = amax.to(tl.int32, bitcast=True)
    amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
    amax = amax.to(tl.float32, bitcast=True)
    scale_e8m0_unbiased = tl.log2(amax).floor() - 2
    scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)

    # blockscale_e8m0
    bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127  # in fp32, we have 2&(e - 127)

    quant_scale = tl.exp2(-scale_e8m0_unbiased)

    # Compute quantized x
    qx = x * quant_scale

    # Convert quantized fp32 tensor to uint32 before converting to mxfp4 format
    # Note: MXFP4  S:1-bit, E:2-bit, M:1-bit
    #   Zeros: S000 -> +/-0
    #   Denormal Numbers: S001 -> +/- 0.5
    #   Normal Numbers:
    #           S010 -> +/- 1.0
    #           S011 -> +/- 1.5
    #           S100 -> +/- 2.0
    #           S101 -> +/- 3.0
    #           S110 -> +/- 4.0
    #           S111 -> +/- 6.0
    qx = qx.to(tl.uint32, bitcast=True)

    # Extract sign
    s = qx & 0x80000000
    # Set everything to positive, will add sign back at the end
    qx = qx ^ s

    qx_fp32 = qx.to(tl.float32, bitcast=True)
    saturate_mask = qx_fp32 >= max_normal
    denormal_mask = (not saturate_mask) & (qx_fp32 < min_normal)
    normal_mask = not (saturate_mask | denormal_mask)

    # Denormal numbers
    denorm_exp: tl.constexpr = (
        (EXP_BIAS_FP32 - EXP_BIAS_FP4) + (MBITS_F32 - MBITS_FP4) + 1
    )
    denorm_mask_int: tl.constexpr = denorm_exp << MBITS_F32
    denorm_mask_float: tl.constexpr = tl.cast(denorm_mask_int, tl.float32, bitcast=True)

    denormal_x = qx_fp32 + denorm_mask_float
    denormal_x = denormal_x.to(tl.uint32, bitcast=True)
    denormal_x -= denorm_mask_int
    denormal_x = denormal_x.to(tl.uint8)

    # Normal numbers
    normal_x = qx
    # resulting mantissa is odd
    mant_odd = (normal_x >> (MBITS_F32 - MBITS_FP4)) & 1
    # update exponent, rounding bias part 1
    val_to_add = ((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_F32) + (1 << 21) - 1
    normal_x += val_to_add
    # rounding bias part 2
    normal_x += mant_odd
    # take the bits!
    normal_x = normal_x >> (MBITS_F32 - MBITS_FP4)
    normal_x = normal_x.to(tl.uint8)

    # Merge results
    e2m1_value = tl.full(qx.type.get_block_shapes(), 0x7, dtype=tl.uint8)
    e2m1_value = tl.where(normal_mask, normal_x, e2m1_value)
    e2m1_value = tl.where(denormal_mask, denormal_x, e2m1_value)
    # add sign back
    sign_lp = s >> (MBITS_F32 + EBITS_F32 - MBITS_FP4 - EBITS_FP4)
    sign_lp = sign_lp.to(tl.uint8)
    e2m1_value = e2m1_value | sign_lp
    e2m1_value = tl.reshape(
        e2m1_value, [BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2]
    )
    evens, odds = tl.split(e2m1_value)
    x_fp4 = evens | (odds << 4)
    x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)

    return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)


@triton.jit
def _dynamic_mxfp4_quant_kernel(
    x_ptr,
    x_fp4_ptr,
    bs_ptr,
    stride_x_m_in,
    stride_x_n_in,
    stride_x_fp4_m_in,
    stride_x_fp4_n_in,
    stride_bs_m_in,
    stride_bs_n_in,
    M,
    N,
    BLOCK_SIZE_M: tl.constexpr,
    BLOCK_SIZE_N: tl.constexpr,
    NUM_ITER: tl.constexpr,
    NUM_STAGES: tl.constexpr,
    MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
    SCALING_MODE: tl.constexpr,
    num_warps: tl.constexpr,
    waves_per_eu: tl.constexpr,
    num_stages: tl.constexpr,
):
    pid_m = tl.program_id(0)
    start_n = tl.program_id(1) * NUM_ITER
    # cast strides to int64, in case M*N > max int32
    stride_x_m = tl.cast(stride_x_m_in, tl.int64)
    stride_x_n = tl.cast(stride_x_n_in, tl.int64)
    stride_x_fp4_m = tl.cast(stride_x_fp4_m_in, tl.int64)
    stride_x_fp4_n = tl.cast(stride_x_fp4_n_in, tl.int64)
    stride_bs_m = tl.cast(stride_bs_m_in, tl.int64)
    stride_bs_n = tl.cast(stride_bs_n_in, tl.int64)

    NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE

    for pid_n in tl.range(start_n, min(start_n + NUM_ITER, N), num_stages=NUM_STAGES):
        x_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        x_offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
        x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n

        x = tl.load(x_ptr + x_offs, cache_modifier=".cg").to(tl.float32)

        out_tensor, bs_e8m0 = _mxfp4_quant_op(
            x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE
        )

        out_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        out_offs_n = pid_n * BLOCK_SIZE_N // 2 + tl.arange(0, BLOCK_SIZE_N // 2)
        out_offs = (
            out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n
        )

        tl.store(x_fp4_ptr + out_offs, out_tensor)

        bs_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        bs_offs_n = pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)
        bs_offs = bs_offs_m[:, None] * stride_bs_m + bs_offs_n[None, :] * stride_bs_n
        tl.store(bs_ptr + bs_offs, bs_e8m0)


def dynamic_mxfp4_quant(
    x: torch.Tensor, scaling_mode: str = "even"
) -> tuple[torch.Tensor, torch.Tensor]:
    """
    Quantize a tensor to MX FP4 format.

    Args:
        x: The input tensor, typically fp16 or bf16.
        scaling_mode: The method to calculate MX block scaling.
            - "even" (default): `even_round` in `quark.torch.quantization.utils`.
            - etc.
    Returns:
        A tuple of (x_fp4, blockscale_e8m0).
    """
    # Assume x is 2D-Tensor for now
    M, N = x.shape

    assert (N // 2) % 2 == 0

    # This is fixed by spec for MXFP4. Do not tune this.
    MXFP4_QUANT_BLOCK_SIZE = 32
    x_fp4 = torch.empty((M, N // 2), dtype=torch.uint8, device=x.device)
    blockscale_e8m0 = torch.empty(
        ((N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE, M),
        dtype=torch.uint8,
        device=x.device,
    ).T

    # for large N values
    if M <= 32:
        NUM_ITER = 1
        BLOCK_SIZE_M = triton.next_power_of_2(M)
        BLOCK_SIZE_N = 32
        NUM_WARPS = 1
        NUM_STAGES = 1
    else:
        NUM_ITER = 4
        BLOCK_SIZE_M = 64
        BLOCK_SIZE_N = 64
        NUM_WARPS = 4
        NUM_STAGES = 2

        if N <= 16384:
            BLOCK_SIZE_M = 32
            BLOCK_SIZE_N = 128

    # for small N values
    if N <= 1024:
        NUM_ITER = 1
        NUM_STAGES = 1
        NUM_WARPS = 4
        BLOCK_SIZE_N = min(256, triton.next_power_of_2(N))
        # BLOCK_SIZE_N needs to be multiple of 32
        BLOCK_SIZE_N = max(32, BLOCK_SIZE_N)
        BLOCK_SIZE_M = min(8, triton.next_power_of_2(M))

    grid = (
        triton.cdiv(M, BLOCK_SIZE_M),
        triton.cdiv(N, BLOCK_SIZE_N * NUM_ITER),
    )

    _dynamic_mxfp4_quant_kernel[grid](
        x,
        x_fp4,
        blockscale_e8m0,
        *x.stride(),
        *x_fp4.stride(),
        *blockscale_e8m0.stride(),
        M=M,
        N=N,
        MXFP4_QUANT_BLOCK_SIZE=MXFP4_QUANT_BLOCK_SIZE,
        SCALING_MODE=0,
        NUM_ITER=NUM_ITER,
        BLOCK_SIZE_M=BLOCK_SIZE_M,
        BLOCK_SIZE_N=BLOCK_SIZE_N,
        NUM_STAGES=NUM_STAGES,
        num_warps=NUM_WARPS,
        waves_per_eu=0,
        num_stages=1,
    )

    return (x_fp4, blockscale_e8m0)

@triton.jit
def pid_grid(pid: int, num_pid_m: int, num_pid_n: int, GROUP_SIZE_M: tl.constexpr = 1):
    """
    Maps 1D pid to 2D grid coords (pid_m, pid_n).

    Args:
        - pid: 1D pid
        - num_pid_m: grid m size
        - num_pid_n: grid n size
        - GROUP_SIZE_M: tl.constexpr: default is 1
    """
    if GROUP_SIZE_M == 1:
        pid_m = pid // num_pid_n
        pid_n = pid % num_pid_n
    else:
        num_pid_in_group = GROUP_SIZE_M * num_pid_n
        group_id = pid // num_pid_in_group
        first_pid_m = group_id * GROUP_SIZE_M
        group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M)
        tl.assume(group_size_m >= 0)
        pid_m = first_pid_m + (pid % group_size_m)
        pid_n = (pid % num_pid_in_group) // group_size_m

    return pid_m, pid_n


@triton.jit
def remap_xcd(pid, GRID_MN, NUM_XCDS: tl.constexpr = 8):
    ## pid remapping on xcds
    # Number of pids per XCD in the new arrangement
    pids_per_xcd = (GRID_MN + NUM_XCDS - 1) // NUM_XCDS
    # When GRID_MN cannot divide NUM_XCDS, some xcds will have
    # pids_per_xcd pids, the other will have pids_per_xcd - 1 pids.
    # We calculate the number of xcds that have pids_per_xcd pids as
    # tall_xcds
    tall_xcds = GRID_MN % NUM_XCDS
    tall_xcds = NUM_XCDS if tall_xcds == 0 else tall_xcds
    # Compute current XCD and local pid within the XCD
    xcd = pid % NUM_XCDS
    local_pid = pid // NUM_XCDS
    # Calculate new pid based on the new grouping
    # Note that we need to consider the following two cases:
    # 1. the current pid is on a tall xcd
    # 2. the current pid is on a short xcd
    if xcd < tall_xcds:
        pid = xcd * pids_per_xcd + local_pid
    else:
        pid = (
            tall_xcds * pids_per_xcd
            + (xcd - tall_xcds) * (pids_per_xcd - 1)
            + local_pid
        )

    return pid


@gluon.jit
def _gemm_afp4wfp4_kernel(
    a_ptr,
    b_ptr,
    c_ptr,
    a_scales_ptr,
    b_scales_ptr,
    M,
    N,
    K,
    stride_am,
    stride_ak,
    stride_bk,
    stride_bn,
    stride_ck,
    stride_cm,
    stride_cn,
    stride_asm,
    stride_ask,
    stride_bsn,
    stride_bsk,
    # Meta-parameters
    BLOCK_SIZE_M: gl.constexpr,
    BLOCK_SIZE_N: gl.constexpr,
    BLOCK_SIZE_K: gl.constexpr,
    GROUP_SIZE_M: gl.constexpr,
    NUM_KSPLIT: gl.constexpr,
    SPLITK_BLOCK_SIZE: gl.constexpr,
    num_warps: gl.constexpr,
    num_stages: gl.constexpr,
    waves_per_eu: gl.constexpr,
    matrix_instr_nonkdim: gl.constexpr,
    cache_modifier: gl.constexpr,
):
    """
    Kernel for computing the matmul C = A x B.
    A and B inputs are in the microscale fp4 (mxfp4) format.
    A_scales and B_scales are in e8m0 format.
    A has shape (M, K), B has shape (K, N) and C has shape (M, N)
    """
    GRID_MN = gl.cdiv(M, BLOCK_SIZE_M) * gl.cdiv(N, BLOCK_SIZE_N)

    # Grouped and XCD-remapped launch ordering improves L2 residency.
    pid_unified = gl.program_id(axis=0)
    pid_unified = remap_xcd(pid_unified, GRID_MN * NUM_KSPLIT, NUM_XCDS=8)

    pid_k = pid_unified % NUM_KSPLIT
    pid = pid_unified // NUM_KSPLIT
    num_pid_m = gl.cdiv(M, BLOCK_SIZE_M)
    num_pid_n = gl.cdiv(N, BLOCK_SIZE_N)

    if NUM_KSPLIT == 1:
        pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
    else:
        pid_m = pid // num_pid_n
        pid_n = pid % num_pid_n

    SCALE_GROUP_SIZE: gl.constexpr = 32
    BLOCK_K_PACKED: gl.constexpr = BLOCK_SIZE_K // 2
    BLOCK_K_SCALE: gl.constexpr = BLOCK_SIZE_K // SCALE_GROUP_SIZE
    NUM_BUFFERS: gl.constexpr = num_stages if num_stages > 1 else 2

    gl.static_assert(num_warps % 2 == 0)
    gl.static_assert(BLOCK_SIZE_K % SCALE_GROUP_SIZE == 0)

    wmma_layout: gl.constexpr = gl.amd.AMDWMMALayout(
        version=3,
        transposed=True,
        warps_per_cta=[2, num_warps // 2],
        instr_shape=[16, 16, 128],
    )
    wmma_packed_layout: gl.constexpr = gl.amd.AMDWMMALayout(
        version=3,
        transposed=True,
        warps_per_cta=[2, num_warps // 2],
        instr_shape=[16, 16, 64],
    )

    dot_a_layout: gl.constexpr = gl.DotOperandLayout(
        operand_index=0, parent=wmma_packed_layout, k_width=16
    )
    dot_b_layout: gl.constexpr = gl.DotOperandLayout(
        operand_index=1, parent=wmma_packed_layout, k_width=16
    )
    scale_a_layout: gl.constexpr = gl.amd.gfx1250.get_wmma_scale_layout(
        dot_a_layout, [BLOCK_SIZE_M, BLOCK_K_SCALE]
    )
    scale_b_layout: gl.constexpr = gl.amd.gfx1250.get_wmma_scale_layout(
        dot_b_layout, [BLOCK_SIZE_N, BLOCK_K_SCALE]
    )

    PAD_INTERVAL_A: gl.constexpr = 256 if BLOCK_K_PACKED <= 256 else BLOCK_K_PACKED
    PAD_INTERVAL_B: gl.constexpr = 256 if BLOCK_K_PACKED <= 256 else BLOCK_K_PACKED

    shared_layout_a: gl.constexpr = gl.PaddedSharedLayout.with_identity_for(
        [[PAD_INTERVAL_A, 16]], [BLOCK_SIZE_M, BLOCK_K_PACKED], [1, 0]
    )
    shared_layout_b: gl.constexpr = gl.PaddedSharedLayout.with_identity_for(
        [[PAD_INTERVAL_B, 16]], [BLOCK_K_PACKED, BLOCK_SIZE_N], [1, 0]
    )
    shared_layout_as: gl.constexpr = gl.PaddedSharedLayout.with_identity_for(
        [[256, 16]], [BLOCK_SIZE_M, BLOCK_K_SCALE], [1, 0]
    )
    shared_layout_bs: gl.constexpr = gl.PaddedSharedLayout.with_identity_for(
        [[256, 16]], [BLOCK_SIZE_N, BLOCK_K_SCALE], [1, 0]
    )

    split_k_start = pid_k * (SPLITK_BLOCK_SIZE // 2)
    split_k_start_scale = pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE)

    if split_k_start < K:
        valid_packed_k = K - split_k_start
        num_k_iter = gl.cdiv(valid_packed_k, BLOCK_K_PACKED)

        a_desc = gl.amd.gfx1250.tdm.make_tensor_descriptor(
            base=a_ptr + pid_m * BLOCK_SIZE_M * stride_am + split_k_start * stride_ak,
            shape=(M, K),
            strides=(stride_am, stride_ak),
            block_shape=(BLOCK_SIZE_M, BLOCK_K_PACKED),
            layout=shared_layout_a,
        )

        b_desc = gl.amd.gfx1250.tdm.make_tensor_descriptor(
            base=b_ptr + pid_n * BLOCK_SIZE_N * stride_bn + split_k_start * stride_bk,
            shape=(K, N),
            strides=(stride_bk, stride_bn),
            block_shape=(BLOCK_K_PACKED, BLOCK_SIZE_N),
            layout=shared_layout_b,
        )

        packed_scale_k = K // (SCALE_GROUP_SIZE // 2)
        a_scale_desc = gl.amd.gfx1250.tdm.make_tensor_descriptor(
            base=(
                a_scales_ptr
                + pid_m * BLOCK_SIZE_M * stride_asm
                + split_k_start_scale * stride_ask
            ),
            shape=(M, packed_scale_k),
            strides=(stride_asm, stride_ask),
            block_shape=(BLOCK_SIZE_M, BLOCK_K_SCALE),
            layout=shared_layout_as,
        )

        b_scale_desc = gl.amd.gfx1250.tdm.make_tensor_descriptor(
            base=(
                b_scales_ptr
                + pid_n * BLOCK_SIZE_N * stride_bsn
                + split_k_start_scale * stride_bsk
            ),
            shape=(N, packed_scale_k),
            strides=(stride_bsn, stride_bsk),
            block_shape=(BLOCK_SIZE_N, BLOCK_K_SCALE),
            layout=shared_layout_bs,
        )

        a_buffer = gl.allocate_shared_memory(
            a_desc.dtype, shape=[NUM_BUFFERS] + a_desc.block_shape, layout=a_desc.layout
        )
        b_buffer = gl.allocate_shared_memory(
            b_desc.dtype, shape=[NUM_BUFFERS] + b_desc.block_shape, layout=b_desc.layout
        )
        a_scale_buffer = gl.allocate_shared_memory(
            a_scale_desc.dtype,
            shape=[NUM_BUFFERS] + a_scale_desc.block_shape,
            layout=a_scale_desc.layout,
        )
        b_scale_buffer = gl.allocate_shared_memory(
            b_scale_desc.dtype,
            shape=[NUM_BUFFERS] + b_scale_desc.block_shape,
            layout=b_scale_desc.layout,
        )

        load_idx = 0
        wmma_idx = 0

        for _ in gl.static_range(NUM_BUFFERS - 1):
            if load_idx < num_k_iter:
                gl.amd.gfx1250.tdm.async_load(
                    a_desc, [0, load_idx * BLOCK_K_PACKED], a_buffer.index(load_idx)
                )
                gl.amd.gfx1250.tdm.async_load(
                    b_desc, [load_idx * BLOCK_K_PACKED, 0], b_buffer.index(load_idx)
                )
                gl.amd.gfx1250.tdm.async_load(
                    a_scale_desc,
                    [0, load_idx * BLOCK_K_SCALE],
                    a_scale_buffer.index(load_idx),
                )
                gl.amd.gfx1250.tdm.async_load(
                    b_scale_desc,
                    [0, load_idx * BLOCK_K_SCALE],
                    b_scale_buffer.index(load_idx),
                )
            load_idx += 1

        accumulator = gl.zeros(
            (BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=gl.float32, layout=wmma_layout
        )

        for _ in range(0, num_k_iter - (NUM_BUFFERS - 1)):
            gl.amd.gfx1250.tdm.async_load(
                a_desc,
                [0, load_idx * BLOCK_K_PACKED],
                a_buffer.index(load_idx % NUM_BUFFERS),
            )
            gl.amd.gfx1250.tdm.async_load(
                b_desc,
                [load_idx * BLOCK_K_PACKED, 0],
                b_buffer.index(load_idx % NUM_BUFFERS),
            )
            gl.amd.gfx1250.tdm.async_load(
                a_scale_desc,
                [0, load_idx * BLOCK_K_SCALE],
                a_scale_buffer.index(load_idx % NUM_BUFFERS),
            )
            gl.amd.gfx1250.tdm.async_load(
                b_scale_desc,
                [0, load_idx * BLOCK_K_SCALE],
                b_scale_buffer.index(load_idx % NUM_BUFFERS),
            )
            load_idx += 1

            gl.amd.gfx1250.tdm.async_wait((NUM_BUFFERS - 1) * 4)

            a = a_buffer.index(wmma_idx % NUM_BUFFERS).load(layout=dot_a_layout)
            b = b_buffer.index(wmma_idx % NUM_BUFFERS).load(layout=dot_b_layout)
            scale_a = a_scale_buffer.index(wmma_idx % NUM_BUFFERS).load(
                layout=scale_a_layout
            )
            scale_b = b_scale_buffer.index(wmma_idx % NUM_BUFFERS).load(
                layout=scale_b_layout
            )

            accumulator = gl.amd.gfx1250.wmma_scaled(
                a,
                scale_a,
                "e2m1",
                b,
                scale_b,
                "e2m1",
                accumulator,
            )
            wmma_idx += 1

        for i in gl.static_range(NUM_BUFFERS - 1):
            if wmma_idx < num_k_iter:
                gl.amd.gfx1250.tdm.async_wait((NUM_BUFFERS - 2 - i) * 4)

                a = a_buffer.index(wmma_idx % NUM_BUFFERS).load(layout=dot_a_layout)
                b = b_buffer.index(wmma_idx % NUM_BUFFERS).load(layout=dot_b_layout)
                scale_a = a_scale_buffer.index(wmma_idx % NUM_BUFFERS).load(
                    layout=scale_a_layout
                )
                scale_b = b_scale_buffer.index(wmma_idx % NUM_BUFFERS).load(
                    layout=scale_b_layout
                )

                accumulator = gl.amd.gfx1250.wmma_scaled(
                    a,
                    scale_a,
                    "e2m1",
                    b,
                    scale_b,
                    "e2m1",
                    accumulator,
                )
                wmma_idx += 1

        c = accumulator.to(c_ptr.type.element_ty)

        offs_cm = pid_m * BLOCK_SIZE_M + gl.arange(
            0, BLOCK_SIZE_M, layout=gl.SliceLayout(1, wmma_layout)
        )
        offs_cn = pid_n * BLOCK_SIZE_N + gl.arange(
            0, BLOCK_SIZE_N, layout=gl.SliceLayout(0, wmma_layout)
        )
        offs_c = (
            stride_cm * offs_cm[:, None]
            + stride_cn * offs_cn[None, :]
            + pid_k * stride_ck
        )
        c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
        gl.amd.gfx1250.buffer_store(c, c_ptr, offs_c, c_mask)


@triton.jit
def _gemm_afp4wfp4_triton_kernel(
    a_ptr,
    b_ptr,
    c_ptr,
    a_scales_ptr,
    b_scales_ptr,
    M: tl.constexpr,
    N: tl.constexpr,
    K: tl.constexpr,
    stride_am,
    stride_ak,
    stride_bk,
    stride_bn,
    stride_ck,
    stride_cm,
    stride_cn,
    stride_asm,
    stride_ask,
    stride_bsn,
    stride_bsk,
    BLOCK_SIZE_M: tl.constexpr,
    BLOCK_SIZE_N: tl.constexpr,
    BLOCK_SIZE_K: tl.constexpr,
    GROUP_SIZE_M: tl.constexpr,
    NUM_KSPLIT: tl.constexpr,
    SPLITK_BLOCK_SIZE: tl.constexpr,
    num_warps: tl.constexpr,
    num_stages: tl.constexpr,
    waves_per_eu: tl.constexpr,
    matrix_instr_nonkdim: tl.constexpr,
    cache_modifier: tl.constexpr,
):
    """Triton WMMA-scaled split-K kernel specialized for fixed benchmark shapes."""
    GRID_MN = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)

    pid_unified = tl.program_id(axis=0)
    pid_unified = remap_xcd(pid_unified, GRID_MN * NUM_KSPLIT, NUM_XCDS=8)

    pid_k = pid_unified % NUM_KSPLIT
    pid = pid_unified // NUM_KSPLIT
    num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
    num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)

    if NUM_KSPLIT == 1:
        pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
    else:
        pid_m = pid // num_pid_n
        pid_n = pid % num_pid_n

    SCALE_GROUP_SIZE: tl.constexpr = 32
    BLOCK_K_PACKED: tl.constexpr = BLOCK_SIZE_K // 2
    BLOCK_K_SCALE: tl.constexpr = BLOCK_SIZE_K // SCALE_GROUP_SIZE

    split_k_start = pid_k * (SPLITK_BLOCK_SIZE // 2)
    split_ks_start = pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE)

    if split_k_start < K:
        num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_K_PACKED)

        offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        offs_bn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
        offs_am_base = pid_m * BLOCK_SIZE_M
        offs_bn_base = pid_n * BLOCK_SIZE_N

        offs_cm = offs_am.to(tl.int64)
        offs_cn = offs_bn.to(tl.int64)

        a_desc = tl.make_tensor_descriptor(
            a_ptr,
            shape=[M, K],
            strides=[stride_am, stride_ak],
            block_shape=[BLOCK_SIZE_M, BLOCK_K_PACKED],
        )
        b_desc = tl.make_tensor_descriptor(
            b_ptr,
            shape=[N, K],
            strides=[stride_bn, stride_bk],
            block_shape=[BLOCK_SIZE_N, BLOCK_K_PACKED],
        )

        acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)

        for k_iter in tl.range(0, num_k_iter, num_stages=num_stages):
            k_base = split_k_start + k_iter * BLOCK_K_PACKED
            ks_base = split_ks_start + k_iter * BLOCK_K_SCALE

            a = a_desc.load([offs_am_base, k_base])
            b = b_desc.load([offs_bn_base, k_base]).trans(1, 0)
            offs_ks = ks_base + tl.arange(0, BLOCK_K_SCALE)
            a_scale_ptrs = (
                a_scales_ptr
                + offs_am[:, None] * stride_asm
                + offs_ks[None, :] * stride_ask
            )
            b_scale_ptrs = (
                b_scales_ptr
                + offs_bn[:, None] * stride_bsn
                + offs_ks[None, :] * stride_bsk
            )
            a_scales = tl.load(a_scale_ptrs)
            b_scales = tl.load(b_scale_ptrs)

            acc = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", acc)

        c = acc.to(c_ptr.type.element_ty)

        c_ptrs = (
            c_ptr
            + stride_cm * offs_cm[:, None]
            + stride_cn * offs_cn[None, :]
            + pid_k * stride_ck
        )
        tl.store(c_ptrs, c, cache_modifier=".wt")


@triton.jit
def _gemm_afp4wfp4_preshuffle_triton_kernel(
    a_ptr,
    b_ptr,
    c_ptr,
    a_scales_ptr,
    b_scales_ptr,
    M: tl.constexpr,
    N: tl.constexpr,
    K: tl.constexpr,
    stride_am,
    stride_ak,
    stride_bn,
    stride_bk,
    stride_ck,
    stride_cm,
    stride_cn,
    stride_asm,
    stride_ask,
    stride_bsn,
    stride_bsk,
    BLOCK_SIZE_M: tl.constexpr,
    BLOCK_SIZE_N: tl.constexpr,
    BLOCK_SIZE_K: tl.constexpr,
    GROUP_SIZE_M: tl.constexpr,
    NUM_KSPLIT: tl.constexpr,
    SPLITK_BLOCK_SIZE: tl.constexpr,
    A_SHUFFLED: tl.constexpr,
    B_SHUFFLED: tl.constexpr,
    ATOMIC_ADD: tl.constexpr,
    num_warps: tl.constexpr,
    num_stages: tl.constexpr,
    waves_per_eu: tl.constexpr,
    matrix_instr_nonkdim: tl.constexpr,
    cache_modifier: tl.constexpr,
):
    """Triton preshuffle kernel supporting shuffled A/B packed mxfp4 tensors."""
    GRID_MN = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)

    pid_unified = tl.program_id(axis=0)
    pid_unified = remap_xcd(pid_unified, GRID_MN * NUM_KSPLIT, NUM_XCDS=8)

    pid_k = pid_unified % NUM_KSPLIT
    pid = pid_unified // NUM_KSPLIT
    num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
    num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)

    if NUM_KSPLIT == 1:
        pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
    else:
        pid_m = pid // num_pid_n
        pid_n = pid % num_pid_n

    SCALE_GROUP_SIZE: tl.constexpr = 32
    BLOCK_K_PACKED: tl.constexpr = BLOCK_SIZE_K // 2
    BLOCK_K_SCALE: tl.constexpr = BLOCK_SIZE_K // SCALE_GROUP_SIZE
    packed_k = K
    packed_scale_k = K // (SCALE_GROUP_SIZE // 2)

    if A_SHUFFLED:
        tl.static_assert(BLOCK_SIZE_M % 16 == 0)
    if B_SHUFFLED:
        tl.static_assert(BLOCK_SIZE_N % 16 == 0)
        tl.static_assert(BLOCK_SIZE_K % 256 == 0)
    tl.static_assert(BLOCK_K_PACKED % 32 == 0)

    split_k_start = pid_k * (SPLITK_BLOCK_SIZE // 2)
    split_ks_start = pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE)

    if split_k_start < K:
        num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_K_PACKED)

        offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        offs_bn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
        offs_bn_packed = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % N
        offs_bsn = (pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)) % N
        offs_am_base = pid_m * BLOCK_SIZE_M

        offs_cm = offs_am.to(tl.int64)
        offs_cn = offs_bn.to(tl.int64)

        a_desc = tl.make_tensor_descriptor(
            a_ptr,
            shape=[M, K],
            strides=[stride_am, stride_ak],
            block_shape=[BLOCK_SIZE_M, BLOCK_K_PACKED],
        )
        acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)

        for k_iter in tl.range(0, num_k_iter, num_stages=num_stages):
            k_base = split_k_start + k_iter * BLOCK_K_PACKED
            ks_base = split_ks_start + k_iter * BLOCK_K_SCALE

            a = a_desc.load([offs_am_base, k_base])
            if A_SHUFFLED:
                a = (
                    a.reshape(
                        BLOCK_SIZE_M // 16,
                        BLOCK_K_PACKED // 32,
                        2,
                        16,
                        16,
                    )
                    .permute(0, 3, 1, 2, 4)
                    .reshape(BLOCK_SIZE_M, BLOCK_K_PACKED)
                )

            if B_SHUFFLED:
                offs_k_shuffle = (k_base * 16) + tl.arange(0, BLOCK_K_PACKED * 16)
                b_rows = offs_bn_packed[:, None] * 16 + (
                    offs_k_shuffle[None, :] // packed_k
                )
                b_cols = offs_k_shuffle[None, :] % packed_k
                b_ptrs = b_ptr + (b_rows * stride_bn + b_cols * stride_bk)
                b = (
                    tl.load(b_ptrs)
                    .reshape(
                        1,
                        BLOCK_SIZE_N // 16,
                        BLOCK_SIZE_K // 64,
                        2,
                        16,
                        16,
                    )
                    .permute(0, 1, 4, 2, 3, 5)
                    .reshape(BLOCK_SIZE_N, BLOCK_K_PACKED)
                )
            else:
                b_desc = tl.make_tensor_descriptor(
                    b_ptr,
                    shape=[N, K],
                    strides=[stride_bn, stride_bk],
                    block_shape=[BLOCK_SIZE_N, BLOCK_K_PACKED],
                )
                b = b_desc.load([pid_n * BLOCK_SIZE_N, k_base])
            b = b.trans(1, 0)

            offs_ks = ks_base + tl.arange(0, BLOCK_K_SCALE)
            a_scale_ptrs = (
                a_scales_ptr
                + offs_am[:, None] * stride_asm
                + offs_ks[None, :] * stride_ask
            )
            a_scales = tl.load(a_scale_ptrs)

            if B_SHUFFLED:
                offs_ks_shuffled = (ks_base * 32) + tl.arange(0, BLOCK_K_SCALE * 32)
                b_scale_rows = offs_bsn[:, None] * 32 + (
                    offs_ks_shuffled[None, :] // packed_scale_k
                )
                b_scale_cols = offs_ks_shuffled[None, :] % packed_scale_k
                b_scale_ptrs = (
                    b_scales_ptr
                    + b_scale_rows * stride_bsn
                    + b_scale_cols * stride_bsk
                )
                b_scales = (
                    tl.load(b_scale_ptrs)
                    .reshape(
                        BLOCK_SIZE_N // 32,
                        BLOCK_K_SCALE // 8,
                        4,
                        16,
                        2,
                        2,
                        1,
                    )
                    .permute(0, 5, 3, 1, 4, 2, 6)
                    .reshape(BLOCK_SIZE_N, BLOCK_K_SCALE)
                )
            else:
                b_scale_ptrs = (
                    b_scales_ptr
                    + offs_bn[:, None] * stride_bsn
                    + offs_ks[None, :] * stride_bsk
                )
                b_scales = tl.load(b_scale_ptrs)

            acc = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", acc)

        c_ptrs = (
            c_ptr
            + stride_cm * offs_cm[:, None]
            + stride_cn * offs_cn[None, :]
            + pid_k * stride_ck
        )
        if ATOMIC_ADD:
            tl.atomic_add(c_ptrs, acc, sem="relaxed")
        else:
            tl.store(c_ptrs, acc.to(c_ptr.type.element_ty), cache_modifier=".wt")


@gluon.jit
def _gemm_afp4wfp4_preshuffle_kernel(
    a_ptr,
    b_ptr,
    c_ptr,
    a_scales_ptr,
    b_scales_ptr,
    M,
    N,
    K,
    stride_am,
    stride_ak,
    stride_bn,
    stride_bk,
    stride_ck,
    stride_cm,
    stride_cn,
    stride_asm,
    stride_ask,
    stride_bsn,
    stride_bsk,
    # Meta-parameters
    BLOCK_SIZE_M: gl.constexpr,
    BLOCK_SIZE_N: gl.constexpr,
    BLOCK_SIZE_K: gl.constexpr,
    GROUP_SIZE_M: gl.constexpr,
    NUM_KSPLIT: gl.constexpr,
    SPLITK_BLOCK_SIZE: gl.constexpr,
    num_warps: gl.constexpr,
    num_stages: gl.constexpr,
    waves_per_eu: gl.constexpr,
    matrix_instr_nonkdim: gl.constexpr,
    cache_modifier: gl.constexpr,
):
    """
    Kernel for computing the matmul C = A x B.
    A and B inputs are in the microscale fp4 (mxfp4) format.
    A_scales and B_scales are in e8m0 format.
    A has shape (M, K), B and B_scales are loaded from preshuffled storage,
    and C has shape (M, N)
    """
    GRID_MN = gl.cdiv(M, BLOCK_SIZE_M) * gl.cdiv(N, BLOCK_SIZE_N)

    # -----------------------------------------------------------
    # Map program ids `pid` to the block of C it should compute.
    # This is done in a grouped ordering to promote L2 data reuse.
    pid_unified = gl.program_id(axis=0)
    # remap so that XCDs get continous chunks of pids (of CHUNK_SIZE).
    pid_unified = remap_xcd(pid_unified, GRID_MN * NUM_KSPLIT, NUM_XCDS=8)

    pid_k = pid_unified % NUM_KSPLIT
    pid = pid_unified // NUM_KSPLIT
    num_pid_m = gl.cdiv(M, BLOCK_SIZE_M)
    num_pid_n = gl.cdiv(N, BLOCK_SIZE_N)

    if NUM_KSPLIT == 1:
        pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
    else:
        pid_m = pid // num_pid_n
        pid_n = pid % num_pid_n

    # We assume 32 elements along K share the same scale.
    SCALE_GROUP_SIZE: gl.constexpr = 32

    blocked_mk: gl.constexpr = gl.BlockedLayout(
        size_per_thread=[1, 16],
        threads_per_warp=[8, 8],
        warps_per_cta=[num_warps, 1],
        order=[1, 0],
    )

    blocked_scales: gl.constexpr = gl.BlockedLayout(
        size_per_thread=[4, 1],
        threads_per_warp=[8, 8],
        warps_per_cta=[1, num_warps],
        order=[0, 1],
    )

    blocked_b_preshuffle: gl.constexpr = gl.BlockedLayout(
        size_per_thread=[1, 16],
        threads_per_warp=[8, 8],
        warps_per_cta=[1, num_warps],
        order=[1, 0],
    )

    blocked_shuffle_scales: gl.constexpr = gl.BlockedLayout(
        size_per_thread=[1, 4],
        threads_per_warp=[8, 8],
        warps_per_cta=[1, num_warps],
        order=[1, 0],
    )

    shared_a: gl.constexpr = gl.SwizzledSharedLayout(
        vec=16, per_phase=2, max_phase=8, order=[1, 0]
    )

    shared_b: gl.constexpr = gl.SwizzledSharedLayout(
        vec=16, per_phase=2, max_phase=8, order=[0, 1]
    )

    shared_scales: gl.constexpr = gl.SwizzledSharedLayout(
        vec=1, per_phase=1, max_phase=1, order=[0, 1]
    )

    mfma_layout: gl.constexpr = gl.amd.AMDMFMALayout(
        version=4,
        instr_shape=[32, 32, 32],
        transposed=True,
        warps_per_cta=[2, num_warps // 2],
    )

    dot_a_layout: gl.constexpr = gl.DotOperandLayout(
        operand_index=0, parent=mfma_layout, k_width=16
    )
    dot_b_layout: gl.constexpr = gl.DotOperandLayout(
        operand_index=1, parent=mfma_layout, k_width=16
    )
    scale_a_layout: gl.constexpr = gl.amd.cdna4.get_mfma_scale_layout(
        dot_a_layout, [BLOCK_SIZE_M, BLOCK_SIZE_K // SCALE_GROUP_SIZE]
    )
    scale_b_layout: gl.constexpr = gl.amd.cdna4.get_mfma_scale_layout(
        dot_b_layout, [BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE]
    )

    if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:

        num_k_iter = gl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
        packed_k = K
        packed_scale_k = K // (SCALE_GROUP_SIZE // 2)
        packed_bn = N // 16
        packed_bsn = N // 32

        # Create pointers for first block of A and B input matrices
        # A stays in packed [M, K/2] form, while B is read back from the
        # preshuffled storage used by the Triton submission path.
        offs_ak = gl.arange(0, BLOCK_SIZE_K // 2, layout=gl.SliceLayout(0, blocked_mk))
        offs_bk = gl.arange(
            0,
            (BLOCK_SIZE_K // 2) * 16,
            layout=gl.SliceLayout(0, blocked_b_preshuffle),
        )
        offs_ks_shuffle = gl.arange(
            0,
            (BLOCK_SIZE_K // SCALE_GROUP_SIZE) * 32,
            layout=gl.SliceLayout(0, blocked_shuffle_scales),
        )
        offs_ks = gl.arange(
            0,
            BLOCK_SIZE_K // SCALE_GROUP_SIZE,
            layout=gl.SliceLayout(0, blocked_scales),
        )
        offs_am = (
            pid_m * BLOCK_SIZE_M
            + gl.arange(0, BLOCK_SIZE_M, layout=gl.SliceLayout(1, blocked_mk))
        ) % M
        offs_bn_packed = (
            pid_n * (BLOCK_SIZE_N // 16)
            + gl.arange(
                0,
                BLOCK_SIZE_N // 16,
                layout=gl.SliceLayout(1, blocked_b_preshuffle),
            )
        ) % packed_bn
        offs_asm = (
            pid_m * BLOCK_SIZE_M
            + gl.arange(0, BLOCK_SIZE_M, layout=gl.SliceLayout(1, blocked_scales))
        ) % M
        offs_bsn_packed = (
            pid_n * (BLOCK_SIZE_N // 32)
            + gl.arange(
                0,
                BLOCK_SIZE_N // 32,
                layout=gl.SliceLayout(1, blocked_shuffle_scales),
            )
        ) % packed_bsn

        # Create shared memories
        smem_a = gl.allocate_shared_memory(
            a_ptr.type.element_ty, [BLOCK_SIZE_M, BLOCK_SIZE_K // 2], layout=shared_a
        )
        smem_b = gl.allocate_shared_memory(
            b_ptr.type.element_ty, [BLOCK_SIZE_K // 2, BLOCK_SIZE_N], layout=shared_b
        )

        smem_as = gl.allocate_shared_memory(
            a_scales_ptr.type.element_ty,
            [BLOCK_SIZE_M, BLOCK_SIZE_K // SCALE_GROUP_SIZE],
            layout=shared_scales,
        )
        smem_bs = gl.allocate_shared_memory(
            b_scales_ptr.type.element_ty,
            [BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE],
            layout=shared_scales,
        )

        accumulator = gl.zeros(
            (BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=gl.float32, layout=mfma_layout
        )

        # Load first blocks of A and B input matrices
        offs_ak_split = pid_k * (SPLITK_BLOCK_SIZE // 2) + offs_ak
        offs_a = offs_am[:, None] * stride_am + offs_ak_split[None, :] * stride_ak
        a = gl.amd.cdna4.buffer_load(
            ptr=a_ptr,
            offsets=offs_a,
        )

        # Create pointers for the first block of A and B scales.
        offs_ks_split = pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) + offs_ks
        offs_as = offs_asm[:, None] * stride_asm + offs_ks_split[None, :] * stride_ask
        a_scales = gl.amd.cdna4.buffer_load(
            ptr=a_scales_ptr,
            offsets=offs_as,
        )

        offs_bk_split = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_bk
        b_row_idx = offs_bn_packed[:, None] * 16 + (
            offs_bk_split[None, :] // packed_k
        )
        b_col_idx = offs_bk_split[None, :] % packed_k
        offs_b = b_row_idx * stride_bn + b_col_idx * stride_bk
        b = gl.amd.cdna4.buffer_load(
            ptr=b_ptr,
            offsets=offs_b,
            cache=cache_modifier,
        )

        # B scales are N x K even though B operand is K x N.
        offs_ks_split = (
            pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) * 32 + offs_ks_shuffle
        )
        b_scale_row_idx = offs_bsn_packed[:, None] * 32 + (
            offs_ks_split[None, :] // packed_scale_k
        )
        b_scale_col_idx = offs_ks_split[None, :] % packed_scale_k
        offs_bs = b_scale_row_idx * stride_bsn + b_scale_col_idx * stride_bsk
        b_scales = (
            gl.amd.cdna4.buffer_load(
                ptr=b_scales_ptr,
                offsets=offs_bs,
                cache=cache_modifier,
            )
            .reshape(
                BLOCK_SIZE_N // 32,
                BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8,
                4,
                16,
                2,
                2,
                1,
            )
            .permute(0, 5, 3, 1, 4, 2, 6)
            .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE)
        )

        # Reconstruct B from preshuffled storage into the layout consumed by LDS/MFMA.
        b = (
            b.reshape(
                1,
                BLOCK_SIZE_N // 16,
                BLOCK_SIZE_K // 64,
                2,
                16,
                16,
            )
            .permute(0, 1, 4, 2, 3, 5)
            .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
            .trans(1, 0)
        )

        smem_a.store(a)
        smem_as.store(a_scales)

        # num_stages:2
        for k in range(0, num_k_iter - 1):

            # Load next block of A.
            offs_ak_split = (
                pid_k * (SPLITK_BLOCK_SIZE // 2)
                + (k + 1) * (BLOCK_SIZE_K // 2)
                + offs_ak
            )
            offs_a = offs_am[:, None] * stride_am + offs_ak_split[None, :] * stride_ak
            a = gl.amd.cdna4.buffer_load(
                ptr=a_ptr,
                offsets=offs_a,
            )

            # LDS write current blocks of B and B scales.
            smem_b.store(b)
            smem_bs.store(b_scales)
            curr_a = smem_a.load(layout=dot_a_layout)
            curr_a_scales = smem_as.load(layout=scale_a_layout)

            # Load next block of A scales.
            offs_ks_split = (
                pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE)
                + (k + 1) * (BLOCK_SIZE_K // SCALE_GROUP_SIZE)
                + offs_ks
            )
            offs_as = (
                offs_asm[:, None] * stride_asm + offs_ks_split[None, :] * stride_ask
            )
            a_scales = gl.amd.cdna4.buffer_load(
                ptr=a_scales_ptr,
                offsets=offs_as,
            )

            curr_b_scales = smem_bs.load(layout=scale_b_layout)

            # Load next block of B from preshuffled storage.
            offs_bk_split = (
                pid_k * (SPLITK_BLOCK_SIZE // 2) * 16
                + (k + 1) * (BLOCK_SIZE_K // 2) * 16
                + offs_bk
            )
            b_row_idx = offs_bn_packed[:, None] * 16 + (
                offs_bk_split[None, :] // packed_k
            )
            b_col_idx = offs_bk_split[None, :] % packed_k
            offs_b = b_row_idx * stride_bn + b_col_idx * stride_bk

            b = gl.amd.cdna4.buffer_load(
                ptr=b_ptr,
                offsets=offs_b,
                cache=cache_modifier,
            )

            # Load next block of B scales from preshuffled storage.
            offs_ks_split = (
                pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) * 32
                + (k + 1) * BLOCK_SIZE_K
                + offs_ks_shuffle
            )
            b_scale_row_idx = offs_bsn_packed[:, None] * 32 + (
                offs_ks_split[None, :] // packed_scale_k
            )
            b_scale_col_idx = offs_ks_split[None, :] % packed_scale_k
            offs_bs = b_scale_row_idx * stride_bsn + b_scale_col_idx * stride_bsk
            b_scales = (
                gl.amd.cdna4.buffer_load(
                    ptr=b_scales_ptr,
                    offsets=offs_bs,
                    cache=cache_modifier,
                )
                .reshape(
                    BLOCK_SIZE_N // 32,
                    BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8,
                    4,
                    16,
                    2,
                    2,
                    1,
                )
                .permute(0, 5, 3, 1, 4, 2, 6)
                .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE)
            )

            # Read current block of B from LDS.
            curr_b = smem_b.load(layout=dot_b_layout)

            accumulator = gl.amd.cdna4.mfma_scaled(
                a=curr_a,
                a_scale=curr_a_scales,
                a_format="e2m1",
                b=curr_b,
                b_scale=curr_b_scales,
                b_format="e2m1",
                acc=accumulator,
            )

            # Reconstruct next block of B into the layout consumed by LDS/MFMA.
            b = (
                b.reshape(
                    1,
                    BLOCK_SIZE_N // 16,
                    BLOCK_SIZE_K // 64,
                    2,
                    16,
                    16,
                )
                .permute(0, 1, 4, 2, 3, 5)
                .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
                .trans(1, 0)
            )

            # LDS write next block of A and A scales.
            smem_a.store(a)
            smem_as.store(a_scales)

        # ======= Epilogue ========
        smem_b.store(b)
        smem_bs.store(b_scales)
        curr_a = smem_a.load(layout=dot_a_layout)
        curr_b = smem_b.load(layout=dot_b_layout)
        curr_a_scales = smem_as.load(layout=scale_a_layout)
        curr_b_scales = smem_bs.load(layout=scale_b_layout)

        accumulator = gl.amd.cdna4.mfma_scaled(
            a=curr_a,
            a_scale=curr_a_scales,
            a_format="e2m1",
            b=curr_b,
            b_scale=curr_b_scales,
            b_format="e2m1",
            acc=accumulator,
        )

        c = accumulator.to(c_ptr.type.element_ty)

        offs_cm = pid_m * BLOCK_SIZE_M + gl.arange(
            0, BLOCK_SIZE_M, layout=gl.SliceLayout(1, mfma_layout)
        )
        offs_cn = pid_n * BLOCK_SIZE_N + gl.arange(
            0, BLOCK_SIZE_N, layout=gl.SliceLayout(0, mfma_layout)
        )
        offs_c = (
            stride_cm * offs_cm[:, None]
            + stride_cn * offs_cn[None, :]
            + pid_k * stride_ck
        )
        c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)

        gl.amd.cdna4.buffer_store(c, c_ptr, offs_c, c_mask)


@gluon.jit
def _gemm_afp4wfp4_reduce_kernel(
    c_in_ptr,
    c_out_ptr,
    M,
    N,
    stride_c_in_k,
    stride_c_in_m,
    stride_c_in_n,
    stride_c_out_m,
    stride_c_out_n,
    BLOCK_SIZE_M: gl.constexpr,
    BLOCK_SIZE_N: gl.constexpr,
    ACTUAL_KSPLIT: gl.constexpr,
    MAX_KSPLIT: gl.constexpr,
):

    pid_m = gl.program_id(axis=0)
    pid_n = gl.program_id(axis=1)

    blocked_kmn: gl.constexpr = gl.BlockedLayout(
        size_per_thread=[1, 1, 4],
        threads_per_warp=[2, 2, 16],
        warps_per_cta=[1, 4, 1],
        order=[2, 0, 1],
    )

    blocked_mn: gl.constexpr = gl.BlockedLayout(
        size_per_thread=[1, 4],
        threads_per_warp=[4, 16],
        warps_per_cta=[4, 1],
        order=[1, 0],
    )

    offs_m = (
        pid_m * BLOCK_SIZE_M
        + gl.arange(
            0, BLOCK_SIZE_M, layout=gl.SliceLayout(0, gl.SliceLayout(2, blocked_kmn))
        )
    ) % M
    offs_n = (
        pid_n * BLOCK_SIZE_N
        + gl.arange(
            0, BLOCK_SIZE_N, layout=gl.SliceLayout(0, gl.SliceLayout(1, blocked_kmn))
        )
    ) % N
    offs_k = gl.arange(
        0, MAX_KSPLIT, layout=gl.SliceLayout(1, gl.SliceLayout(2, blocked_kmn))
    )
    c_in_ptrs = (
        c_in_ptr
        + (offs_k[:, None, None] * stride_c_in_k)
        + (offs_m[None, :, None] * stride_c_in_m)
        + (offs_n[None, None, :] * stride_c_in_n)
    )

    if ACTUAL_KSPLIT == MAX_KSPLIT:
        c = gl.load(c_in_ptrs)
    else:
        c = gl.load(c_in_ptrs, mask=offs_k[:, None, None] < ACTUAL_KSPLIT)
    c = gl.sum(c, axis=0)

    c = c.to(c_out_ptr.type.element_ty)
    offs_m = (
        pid_m * BLOCK_SIZE_M
        + gl.arange(0, BLOCK_SIZE_M, layout=gl.SliceLayout(1, blocked_mn))
    ) % M
    offs_n = (
        pid_n * BLOCK_SIZE_N
        + gl.arange(0, BLOCK_SIZE_N, layout=gl.SliceLayout(0, blocked_mn))
    ) % N
    c_out_ptrs = (
        c_out_ptr
        + (offs_m[:, None] * stride_c_out_m)
        + (offs_n[None, :] * stride_c_out_n)
    )
    c = gl.convert_layout(c, layout=blocked_mn, assert_trivial=False)
    gl.store(c_out_ptrs, c)


@triton.jit
def _gemm_afp4wfp4_reduce_triton_kernel(
    c_in_ptr,
    c_out_ptr,
    M,
    N,
    stride_c_in_k,
    stride_c_in_m,
    stride_c_in_n,
    stride_c_out_m,
    stride_c_out_n,
    BLOCK_SIZE_M: tl.constexpr,
    BLOCK_SIZE_N: tl.constexpr,
    ACTUAL_KSPLIT: tl.constexpr,
    MAX_KSPLIT: tl.constexpr,
):

    pid_m = tl.program_id(axis=0)
    pid_n = tl.program_id(axis=1)

    offs_m = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
    offs_n = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
    offs_k = tl.arange(0, MAX_KSPLIT)
    c_in_ptrs = (
        c_in_ptr
        + (offs_k[:, None, None] * stride_c_in_k)
        + (offs_m[None, :, None] * stride_c_in_m)
        + (offs_n[None, None, :] * stride_c_in_n)
    )

    if ACTUAL_KSPLIT == MAX_KSPLIT:
        c = tl.load(c_in_ptrs)
    else:
        c = tl.load(c_in_ptrs, mask=offs_k[:, None, None] < ACTUAL_KSPLIT)
    c = tl.sum(c, axis=0)

    c = c.to(c_out_ptr.type.element_ty)

    c_out_ptrs = (
        c_out_ptr
        + (offs_m[:, None] * stride_c_out_m)
        + (offs_n[None, :] * stride_c_out_n)
    )

    tl.store(c_out_ptrs, c)


@functools.lru_cache(maxsize=1024)
def _get_config(
    M: int,
    N: int,
    K: int,
):
    # K in this API is packed K/2, benchmark table uses full K.
    full_k = 2 * K
    key = (M, N, full_k)
    if key in _BENCHMARK_GEMM_CONFIGS:
        return dict(_BENCHMARK_GEMM_CONFIGS[key])

    return _select_default_config_by_m(M)


def gemm_afp4wfp4_preshuffle(
    x: torch.Tensor,
    w: torch.Tensor,
    x_scales: torch.Tensor,
    w_scales: torch.Tensor,
    dtype: Optional[torch.dtype] = torch.bfloat16,
    y: Optional[torch.Tensor] = None,
    config: Optional[dict] = None,
    skip_reduce: Optional[bool] = False,
    a_shuffled: Optional[bool] = False,
    b_shuffled: Optional[bool] = True,
) -> torch.Tensor:
    """
    Computes matrix multiplication Y = X @ W with FP4 activations and FP4 weights.

    Args:
        x (torch.Tensor): FP4 E2M1 input matrix with shape (M, K//2).
        w (torch.Tensor): FP4 E2M1 preshuffled weight storage.
        x_scales (torch.Tensor): E8M0 per-group scale for x with shape (M, K//32).
            One scale per 32 elements in K dimension.
        w_scales (torch.Tensor): E8M0 per-group scale for w with shape (N, K//32).
            One scale per 32 elements in K dimension.
        dtype (Optional[torch.dtype]): Output datatype (BF16 or FP16).
        y (Optional[torch.Tensor]): Pre-allocated output tensor with shape (M, N).
        config (Optional[dict]): Kernel tuning parameters (BLOCK_SIZE_M, BLOCK_SIZE_N,
            BLOCK_SIZE_K, GROUP_SIZE_M, NUM_KSPLIT, SPLITK_BLOCK_SIZE).
        skip_reduce (Optional[bool]): skip reduction, y becomes (SPK, M, N) where SPK is determined by config

    Returns:
        y (torch.Tensor): Output with shape (M, N) or (SPK, M, N).
    """
    atomic_add = False
    M, K = x.shape
    N, K = w.shape

    if config is None:
        config = _get_config(M, N, K)

    if config["BLOCK_SIZE_K"] >= K * 2:
        config["NUM_KSPLIT"] = 1

    if config["NUM_KSPLIT"] > 1:
        SPLITK_BLOCK_SIZE = (
            triton.cdiv(
                (2 * triton.cdiv(K, config["NUM_KSPLIT"])), config["BLOCK_SIZE_K"]
            )
            * config["BLOCK_SIZE_K"]
        )
    else:
        SPLITK_BLOCK_SIZE = 2 * K

    config["SPLITK_BLOCK_SIZE"] = SPLITK_BLOCK_SIZE

    if config["NUM_KSPLIT"] > 1 and not atomic_add:
        y_pp = torch.empty(
            (config["NUM_KSPLIT"], M, N), dtype=torch.float32, device=x.device
        )
    else:
        y_pp = None

    y_atomic = None
    if atomic_add:
        y_atomic = torch.zeros((M, N), dtype=torch.float32, device=x.device)

    if y is None and (config["NUM_KSPLIT"] == 1 or not skip_reduce):
        y = torch.empty((M, N), dtype=dtype, device=x.device)

    grid = lambda META: (  # noqa: E731
        (
            META["NUM_KSPLIT"]
            * triton.cdiv(M, META["BLOCK_SIZE_M"])
            * triton.cdiv(N, META["BLOCK_SIZE_N"])
        ),
    )

    def alloc_fn(size: int, align: int, sm: Optional[int]):
        return torch.empty(size, device=x.device, dtype=torch.int8)

    triton.set_allocator(alloc_fn)

    # _gemm_afp4wfp4_preshuffle_kernel[grid](
    _gemm_afp4wfp4_preshuffle_triton_kernel[grid](
        x,
        w,
        y_atomic if atomic_add else (y if y_pp is None else y_pp),
        x_scales,
        w_scales,
        M,
        N,
        K,
        x.stride(0),
        x.stride(1),
        w.stride(0),
        w.stride(1),
        0 if (atomic_add or y_pp is None) else y_pp.stride(0),
        (y_atomic.stride(0) if atomic_add else (y.stride(0) if y_pp is None else y_pp.stride(1))),
        (y_atomic.stride(1) if atomic_add else (y.stride(1) if y_pp is None else y_pp.stride(2))),
        x_scales.stride(0),
        x_scales.stride(1),
        w_scales.stride(0),
        w_scales.stride(1),
        A_SHUFFLED=a_shuffled,
        B_SHUFFLED=b_shuffled,
        ATOMIC_ADD=atomic_add,
        **config,
    )

    if config["NUM_KSPLIT"] > 1 and not atomic_add:
        if skip_reduce:
            return y_pp

        REDUCE_BLOCK_SIZE_M = 16
        REDUCE_BLOCK_SIZE_N = 64
        ACTUAL_KSPLIT = triton.cdiv(K, (config["SPLITK_BLOCK_SIZE"] // 2))

        grid_reduce = (
            triton.cdiv(M, REDUCE_BLOCK_SIZE_M),
            triton.cdiv(N, REDUCE_BLOCK_SIZE_N),
        )
        # _gemm_afp4wfp4_reduce_kernel[grid_reduce](
        _gemm_afp4wfp4_reduce_triton_kernel[grid_reduce](
            y_pp,
            y,
            M,
            N,
            y_pp.stride(0),
            y_pp.stride(1),
            y_pp.stride(2),
            y.stride(0),
            y.stride(1),
            REDUCE_BLOCK_SIZE_M,
            REDUCE_BLOCK_SIZE_N,
            ACTUAL_KSPLIT,
            triton.next_power_of_2(config["NUM_KSPLIT"]),
        )

    if atomic_add:
        y.copy_(y_atomic.to(y.dtype))

    return y


def gemm_afp4wfp4(
    x: torch.Tensor,
    w: torch.Tensor,
    x_scales: torch.Tensor,
    w_scales: torch.Tensor,
    dtype: Optional[torch.dtype] = torch.bfloat16,
    y: Optional[torch.Tensor] = None,
    config: Optional[dict] = None,
    skip_reduce: Optional[bool] = False,
) -> torch.Tensor:
    """
    Computes matrix multiplication Y = X @ W with FP4 activations and FP4 weights.

    This entrypoint expects the original non-preshuffled weight layout.
    """

    M, K = x.shape
    N, _ = w.shape

    if config is None:
        config = _get_config(M, N, K)

    if config["BLOCK_SIZE_K"] >= K * 2:
        config["NUM_KSPLIT"] = 1

    if config["NUM_KSPLIT"] > 1:
        SPLITK_BLOCK_SIZE = (
            triton.cdiv(
                (2 * triton.cdiv(K, config["NUM_KSPLIT"])), config["BLOCK_SIZE_K"]
            )
            * config["BLOCK_SIZE_K"]
        )
    else:
        SPLITK_BLOCK_SIZE = 2 * K

    config["SPLITK_BLOCK_SIZE"] = SPLITK_BLOCK_SIZE

    if config["NUM_KSPLIT"] > 1:
        y_pp = torch.empty(
            (config["NUM_KSPLIT"], M, N), dtype=torch.float32, device=x.device
        )
    else:
        y_pp = None

    if y is None and (config["NUM_KSPLIT"] == 1 or not skip_reduce):
        y = torch.empty((M, N), dtype=dtype, device=x.device)

    grid = lambda META: (  # noqa: E731
        (
            META["NUM_KSPLIT"]
            * triton.cdiv(M, META["BLOCK_SIZE_M"])
            * triton.cdiv(N, META["BLOCK_SIZE_N"])
        ),
    )

    def alloc_fn(size: int, align: int, sm: Optional[int]):
        return torch.empty(size, device=x.device, dtype=torch.int8)

    triton.set_allocator(alloc_fn)

    # _gemm_afp4wfp4_kernel[grid](
    _gemm_afp4wfp4_triton_kernel[grid](
        x,
        w,
        y if config["NUM_KSPLIT"] == 1 else y_pp,
        x_scales,
        w_scales,
        M,
        N,
        K,
        x.stride(0),
        x.stride(1),
        w.stride(1),
        w.stride(0),
        0 if config["NUM_KSPLIT"] == 1 else y_pp.stride(0),
        y.stride(0) if config["NUM_KSPLIT"] == 1 else y_pp.stride(1),
        y.stride(1) if config["NUM_KSPLIT"] == 1 else y_pp.stride(2),
        x_scales.stride(0),
        x_scales.stride(1),
        w_scales.stride(0),
        w_scales.stride(1),
        **config,
    )

    if config["NUM_KSPLIT"] > 1:
        if skip_reduce:
            return y_pp

        REDUCE_BLOCK_SIZE_M = 16
        REDUCE_BLOCK_SIZE_N = 64
        ACTUAL_KSPLIT = triton.cdiv(K, (config["SPLITK_BLOCK_SIZE"] // 2))

        grid_reduce = (
            triton.cdiv(M, REDUCE_BLOCK_SIZE_M),
            triton.cdiv(N, REDUCE_BLOCK_SIZE_N),
        )
        # _gemm_afp4wfp4_reduce_kernel[grid_reduce](
        _gemm_afp4wfp4_reduce_triton_kernel[grid_reduce](
            y_pp,
            y,
            M,
            N,
            y_pp.stride(0),
            y_pp.stride(1),
            y_pp.stride(2),
            y.stride(0),
            y.stride(1),
            REDUCE_BLOCK_SIZE_M,
            REDUCE_BLOCK_SIZE_N,
            ACTUAL_KSPLIT,
            triton.next_power_of_2(config["NUM_KSPLIT"]),
        )

    return y


def e8m0_shuffle(scale):
    if scale is None:
        return scale
    if scale.dtype == torch.float32:
        return scale
    assert scale.ndim == 2, "scale must be a 2D tensor"
    m, n = scale.shape
    scale_padded = torch.empty(
        (m + 255) // 256 * 256,
        (n + 7) // 8 * 8,
        dtype=scale.dtype,
        device=scale.device,
    )

    scale_padded[:m, :n] = scale
    scale = scale_padded
    sm, sn = scale.shape
    scale = scale.view(sm // 32, 2, 16, sn // 8, 2, 4)
    scale = scale.permute(0, 3, 5, 2, 4, 1).contiguous()
    scale = scale.view(sm, sn)
    return scale


def _quant_mxfp4(x, shuffle=True):
    # x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
    x_fp4, bs_e8m0 = dynamic_mxfp4_quant_inline(x)
    if shuffle:
        bs_e8m0 = e8m0_shuffle(bs_e8m0)
    return x_fp4, bs_e8m0


def _as_uint8_storage(x: torch.Tensor) -> torch.Tensor:
    if x.dtype == torch.uint8:
        return x
    return x.view(torch.uint8)

# import time
def custom_kernel(data: Any) -> Any:
    A, B, B_q, B_shuffle, B_scale_sh = data

    B_shuffle = _as_uint8_storage(B_shuffle)
    B_scale_sh = _as_uint8_storage(B_scale_sh)

    start = time.time()
    A_q, A_scale_sh = _quant_mxfp4(A, shuffle=False)
    end = time.time()
    print(f"Quantization time: {(end - start) * 1e6:.2f} us")

    start = time.time()
    out_gemm = gemm_afp4wfp4_preshuffle(
        A_q,
        B_shuffle,
        A_scale_sh,
        B_scale_sh,
        dtype=torch.bfloat16,
        a_shuffled=False,
        b_shuffled=True,
    )
    end = time.time()
    print(f"GEMM time: {(end - start) * 1e6:.2f} us")

    # A_q, A_scale_sh = _quant_mxfp4(A, shuffle=True)
    # B_q, B_scale_sh = _quant_mxfp4(B, shuffle=True)

    # out_gemm = gemm_afp4wfp4(
    #     A_q,
    #     B_q,
    #     A_scale_sh,
    #     B_scale_sh,
    #     dtype=torch.bfloat16,
    # )
    return out_gemm


# def shuffle_weight(x: torch.Tensor, layout=(16, 16), use_int4=False) -> torch.Tensor:
#     # Hardcode BLOCK_K and BLOCK_N
#     x_type = x.dtype
#     if hasattr(torch, "float4_e2m1fn_x2") and x_type == torch.float4_e2m1fn_x2:
#         x = x.view(torch.uint8)

#     IN, IK = layout
#     BK = IK * 2
#     K = 16 // x.element_size() if not use_int4 else 32
#     BN = IN
#     assert x.shape[-2] % BN == 0, f"{x.shape[-2]} % {BN} == {x.shape[-2] % BN }"
#     assert x.shape[-1] % BK == 0, f"{x.shape[-1]} % {BK} == {x.shape[-1] % BK }"

#     x_ = x
#     x_ = x_.view(-1, x.shape[-2] // BN, BN, x.shape[-1] // BK, BK // K, K)
#     x_ = x_.permute(0, 1, 3, 4, 2, 5)
#     x_ = x_.contiguous()
#     x_ = x_.view(*x.shape)
#     x_ = x_.view(x_type)
#     x_.is_shuffled = True
#     return x_


# def generate_input(m: int, n: int, k: int, seed: int) -> tuple[torch.Tensor, ...]:

#     assert k % 64 == 0, "k must be divisible by 64"
#     gen = torch.Generator(device="cuda")
#     gen.manual_seed(seed)
#     a = torch.randn((m, k), dtype=torch.bfloat16, device="cuda", generator=gen)
#     b = torch.randn((n, k), dtype=torch.bfloat16, device="cuda", generator=gen)
#     b_q, b_scale_sh = _quant_mxfp4(b, shuffle=True)
#     b_shuffle = _as_uint8_storage(shuffle_weight(b_q, layout=(16, 16)))
#     return a, b, b_q, b_shuffle, _as_uint8_storage(b_scale_sh)


# def mxfp4_to_f32(x):
#     if x.dtype == torch.float4_e2m1fn_x2:
#         x = x.view(torch.uint8)

#     # 2 because we pack fp4 in uint8.
#     x = x.repeat_interleave(2, dim=-1)
#     x[..., ::2] = x[..., ::2] & 0xF
#     x[..., 1::2] = x[..., 1::2] >> 4
#     mxfp4_list = [
#         0.0,
#         0.5,
#         1.0,
#         1.5,
#         2.0,
#         3.0,
#         4.0,
#         6.0,
#         -0.0,
#         -0.5,
#         -1.0,
#         -1.5,
#         -2.0,
#         -3.0,
#         -4.0,
#         -6.0,
#     ]
#     mxfp4_in_f32 = torch.tensor(mxfp4_list, dtype=torch.float32, device=x.device)
#     return mxfp4_in_f32[x.long()]



# def e8m0_to_f32(scale_e8m0_biased):
#     scale_e8m0_biased = scale_e8m0_biased.view(torch.uint8)
#     zero_case = scale_e8m0_biased == 0
#     nan_case = scale_e8m0_biased == 0xFF
#     scale_f32 = scale_e8m0_biased.to(torch.int32) << 23
#     scale_f32[zero_case] = 0x00400000
#     scale_f32[nan_case] = 0x7F800001
#     scale_f32 = scale_f32.view(torch.float32)
#     return scale_f32


# def run_torch_fp4_mm(
#     x: torch.Tensor,
#     w: torch.Tensor,
#     x_scales: torch.Tensor,
#     w_scales: torch.Tensor,
#     dtype: torch.dtype = torch.bfloat16,
# ) -> torch.Tensor:
#     """
#     PyTorch reference: dequant MXFP4 + E8M0 scale -> f32 -> mm -> dtype.
#     Same logic as aiter op_tests/test_gemm_a4w4.run_torch.
#     x: [m, k//2] fp4 packed, w: [n, k//2] fp4 packed
#     x_scales: [m, k//32] E8M0, w_scales: [n, k//32] E8M0
#     Returns: [m, n] in dtype
#     """

#     m, _ = x.shape
#     n, _ = w.shape
#     # fp4 packed -> f32
#     x_f32 = mxfp4_to_f32(x)
#     w_f32 = mxfp4_to_f32(w)
#     # E8M0 scale: [*, k//32] -> repeat 32 along k -> f32
#     x_scales = x_scales[:m].repeat_interleave(MXFP4_GROUP_SIZE, dim=1)
#     x_scales_f32 = e8m0_to_f32(x_scales)
#     x_f32 = x_f32 * x_scales_f32
#     w_scales = w_scales[:n].repeat_interleave(MXFP4_GROUP_SIZE, dim=1)
#     w_scales_f32 = e8m0_to_f32(w_scales)
#     w_f32 = w_f32 * w_scales_f32
#     return torch.mm(x_f32, w_f32.T).to(dtype)[:m, :n]


# def _calculate_stats(durations_ns: list[float]) -> dict[str, float]:
#     runs = len(durations_ns)
#     mean = sum(durations_ns) / runs
#     best = min(durations_ns)
#     worst = max(durations_ns)
#     if runs > 1:
#         variance = sum((x - mean) ** 2 for x in durations_ns) / (runs - 1)
#         std = math.sqrt(variance)
#         err = std / math.sqrt(runs)
#     else:
#         std = 0.0
#         err = 0.0
#     return {
#         "runs": float(runs),
#         "mean_ns": mean,
#         "std_ns": std,
#         "err_ns": err,
#         "best_ns": best,
#         "worst_ns": worst,
#     }


# def benchmark_custom_kernel(
#     benchmarks: list[dict[str, int]],
#     warmup: int = 0,
#     max_repeats: int = 1,
#     max_time_ns: float = 30e9,
# ) -> None:
#     if not torch.cuda.is_available():
#         print("CUDA/ROCm device not available, skip benchmark")
#         return

#     def _measure_latency_ns(fn, warmup_count: int, max_repeat_count: int) -> dict[str, float]:
#         for _ in range(warmup_count):
#             _ = fn()
#         torch.cuda.synchronize()

#         durations_ns: list[float] = []
#         bm_start = time.perf_counter_ns()
#         for i in range(max_repeat_count):
#             start_event = torch.cuda.Event(enable_timing=True)
#             end_event = torch.cuda.Event(enable_timing=True)
#             start_event.record()
#             _ = fn()
#             end_event.record()
#             torch.cuda.synchronize()
#             durations_ns.append(start_event.elapsed_time(end_event) * 1e6)

#             if i > 1:
#                 stats = _calculate_stats(durations_ns)
#                 total_bm_duration = time.perf_counter_ns() - bm_start
#                 if (
#                     stats["err_ns"] / max(stats["mean_ns"], 1.0) < 0.001
#                     or stats["mean_ns"] * len(durations_ns) > max_time_ns
#                     or total_bm_duration > 120e9
#                 ):
#                     break

#         return _calculate_stats(durations_ns)

#     print(f"benchmark-count: {len(benchmarks)}")
#     for idx, case in enumerate(benchmarks):
#         m, n, k, seed = case["m"], case["n"], case["k"], case["seed"]
#         spec = f"m:{m};n:{n};k:{k};seed:{seed}"
#         print(f"benchmark.{idx}.spec: {spec}")

#         data = generate_input(m, n, k, seed)
#         a, b, _b_q_shuffled_scale, _b_shuffle, _b_scale_sh = data
#         out = custom_kernel(data)
#         torch.cuda.synchronize()
#         if out.shape != (m, n):
#             print(f"benchmark.{idx}.status: fail")
#             print(f"benchmark.{idx}.error: shape mismatch got={tuple(out.shape)} expected={(m, n)}")
#             continue

#         a_q_ref, a_scale_ref = _quant_mxfp4(a, shuffle=False)
#         b_q_ref, b_scale_ref = _quant_mxfp4(b, shuffle=False)
#         out_ref = run_torch_fp4_mm(a_q_ref, b_q_ref, a_scale_ref, b_scale_ref, dtype=out.dtype)

#         out_f32 = out.float()
#         out_ref_f32 = out_ref.float()
#         abs_diff = (out_f32 - out_ref_f32).abs()
#         max_abs_diff = abs_diff.max().item()
#         mean_abs_diff = abs_diff.mean().item()
#         ref_abs_max = out_ref_f32.abs().max().item()
#         rel_max_diff = max_abs_diff / max(ref_abs_max, 1e-6)

#         custom_stats = _measure_latency_ns(lambda: custom_kernel(data), warmup, max_repeats)
#         # Torch reference is much slower, use fewer repeats to keep benchmark time practical.
#         ref_repeats = min(max_repeats, 20)
#         ref_stats = _measure_latency_ns(
#             lambda: run_torch_fp4_mm(a_q_ref, b_q_ref, a_scale_ref, b_scale_ref, dtype=out.dtype),
#             min(warmup, 5),
#             ref_repeats,
#         )

#         speedup = ref_stats["mean_ns"] / max(custom_stats["mean_ns"], 1.0)

#         print(f"benchmark.{idx}.correctness.max_abs_diff: {max_abs_diff:.6f}")
#         print(f"benchmark.{idx}.correctness.mean_abs_diff: {mean_abs_diff:.6f}")
#         print(f"benchmark.{idx}.correctness.rel_max_diff: {rel_max_diff:.6f}")
#         print(f"benchmark.{idx}.perf.custom_mean_us: {custom_stats['mean_ns'] / 1e3:.3f}")
#         print(f"benchmark.{idx}.perf.reference_mean_us: {ref_stats['mean_ns'] / 1e3:.3f}")
#         print(f"benchmark.{idx}.perf.speedup_vs_reference: {speedup:.2f}x")


# if __name__ == "__main__":
#     benchmarks = [
#         {"m": 4, "n": 2880, "k": 512, "seed": 4565},
#         {"m": 16, "n": 2112, "k": 7168, "seed": 15},
#         {"m": 32, "n": 4096, "k": 512, "seed": 457},
#         {"m": 32, "n": 2880, "k": 512, "seed": 54},
#         {"m": 64, "n": 7168, "k": 2048, "seed": 687},
#         {"m": 256, "n": 3072, "k": 1536, "seed": 7856},
#     ]
#     benchmark_custom_kernel(benchmarks)
scrolls · 2559 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