Skip to content
KernelIndex
Search⌘K

submission 649634

XiaomingFun233 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_amd_mxfp4_mm.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-649634?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
31.5µs
#1120 of 1143
2026-03-27

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:82b8c28ecb0842294c733c27649106a7ef76fb35c25f9c52a078832cd77b9af7
license declaredunknown
license concludedunknown
authorsXiaomingFun233
imported2026-08-26

Techniques

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

fp4TORCH_CHECK(a_fp4_u8.dim() == 2 && b_u8.dim() == 2, "A/B FP4 must be 2D");
shared-memory__device__ __forceinline__ int smem_swizzled_dword(int logical_row, int logical_dword) {
split-k_SPLITK_ENV = "MXFP4_MM_LOG2_SPLITK"
stages = 3constexpr int PREFETCH_STAGES = 3;
tile-k = 32constexpr int BLOCK_K = 32;
tile-m = 64constexpr int GEMM_BLOCK_M = 64;
tile-n = 64constexpr int GEMM_BLOCK_N = 64;
vector-width = uint4constexpr int LDS_VEC_BYTES = sizeof(uint4);

Kernel source

submission_amd_mxfp4_mm.py2695 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

import os
import threading
from dataclasses import dataclass
from typing import Dict, Tuple

import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline


_B_LAYOUT_ENV = "MXFP4_MM_B_LAYOUT"  # raw | shuffle | auto
_B_SHUFFLE_INNER_ENV = "MXFP4_MM_B_SHUFFLE_INNER_MODE"  # 0 | 1
_SPLITK_ENV = "MXFP4_MM_LOG2_SPLITK"
_NATIVE_MFMA_ENV = "MXFP4_MM_NATIVE_MFMA"  # never | auto | force (force also enables native build)
_FUSED_A_QUANT_ENV = "MXFP4_MM_FUSED_A_QUANT"  # 0 | 1
_DEFAULT_B_LAYOUT = "shuffle"
_DEFAULT_B_SHUFFLE_INNER_MODE = 1  # Matches aiter.ops.shuffle.shuffle_weight() tile flattening.
_DEFAULT_NATIVE_MFMA = "auto"
_DEFAULT_FUSED_A_QUANT = "0"


_HIP_LOCK = threading.Lock()
_HIP_MODULE = None
_HIP_BUILD_ERROR = None
_HIP_NATIVE_BUILD_ENABLED = False
_HIP_BUILD_REPORT_EMITTED = False
_RUNTIME_REPORT_LOCK = threading.Lock()
_RUNTIME_REPORTED_SHAPES: set[Tuple[int, int, int]] = set()

_B_SCALE_RAW_LOCK = threading.Lock()
_B_SCALE_RAW_CACHE: Dict[Tuple[int, ...], torch.Tensor] = {}
_B_SCALE_RAW_CACHE_MAX = 8

_WORKSPACE_LOCK = threading.Lock()
_WORKSPACE_CACHE: Dict[Tuple[int, ...], torch.Tensor] = {}
_WORKSPACE_CACHE_MAX = 16


_STATIC_SPLITK_LOG2: Dict[Tuple[int, int, int], int] = {
    (4, 2880, 512): 2,
    (16, 2112, 7168): 3,
    (32, 4096, 512): 2,
    (32, 2880, 512): 2,
    (64, 7168, 2048): 1,
    (256, 3072, 1536): 0,
}

_RANKED_SHAPE_KEYS = frozenset(_STATIC_SPLITK_LOG2)


@dataclass(frozen=True)
class _LaunchPolicy:
    log2_k_split: int
    use_native_prequant: int
    use_native_fused: int
    is_ranked_shape: bool


CPP_WRAPPER = r"""
#include <cstdint>
#include <vector>
std::vector<torch::Tensor> hip_quant_mxfp4(torch::Tensor x);
torch::Tensor hip_gemm_mxfp4_prequant(
    torch::Tensor a_fp4_u8,
    torch::Tensor b_u8,
    torch::Tensor a_scale_u8,
    torch::Tensor b_scale_u8,
    int64_t layout_mode,
    int64_t log2_k_split,
    torch::Tensor workspace,
    int64_t workspace_stride,
    int64_t b_shuffle_inner_mode,
    int64_t native_mfma_mode);
torch::Tensor hip_gemm_mxfp4(
    torch::Tensor a_bf16,
    torch::Tensor b_u8,
    torch::Tensor b_scale_u8,
    int64_t layout_mode,
    int64_t log2_k_split,
    torch::Tensor workspace,
    int64_t workspace_stride,
    int64_t b_shuffle_inner_mode,
    int64_t native_mfma_mode);
"""


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

#include <algorithm>
#include <cmath>
#include <cstdint>
#include <vector>

#ifndef MXFP4_ENABLE_NATIVE_FP4_MFMA
#define MXFP4_ENABLE_NATIVE_FP4_MFMA 0
#endif

namespace {

constexpr int BLOCK_K = 32;
constexpr int NATIVE_MFMA_K_BLOCKS = 4;
constexpr int PAD_M_KERNEL = 64;
constexpr int PAD_M_SCALE = 256;
constexpr int PAD_SCALE_N = 8;

constexpr int GEMM_BLOCK_M = 64;
constexpr int GEMM_BLOCK_N = 64;
constexpr int GEMM_THREADS = 256;

constexpr int WAVE_SIZE = 64;
constexpr int WAVE_TILE_M = 32;
constexpr int WAVE_TILE_N = 32;
constexpr int MFMA_TILE_M = 16;
constexpr int MFMA_TILE_N = 16;
constexpr int PACKED_BLOCK_K = BLOCK_K / 2;
constexpr int SMEM_DWORDS = PACKED_BLOCK_K / static_cast<int>(sizeof(uint32_t));
constexpr int LDS_PAD_DWORDS = 1;
constexpr int SMEM_STRIDE_DWORDS = SMEM_DWORDS + LDS_PAD_DWORDS;
constexpr int SMEM_STRIDE = SMEM_STRIDE_DWORDS * static_cast<int>(sizeof(uint32_t));
constexpr int SMEM_SWIZZLE_MASK = SMEM_DWORDS - 1;
constexpr int PREFETCH_STAGES = 3;
constexpr int LDS_VEC_BYTES = sizeof(uint4);
constexpr int A_TILE_VEC_LOADS = (GEMM_BLOCK_M * PACKED_BLOCK_K) / LDS_VEC_BYTES;
constexpr int B_TILE_VEC_LOADS = (GEMM_BLOCK_N * PACKED_BLOCK_K) / LDS_VEC_BYTES;

__device__ __forceinline__ float bf16_to_f32(uint16_t x) {
    return __uint_as_float(static_cast<uint32_t>(x) << 16);
}

__device__ __forceinline__ uint8_t quantize_e2m1(float x_scaled) {
    // Match aiter dynamic_mxfp4_quant conversion path.
    uint32_t qx = __float_as_uint(x_scaled);
    uint32_t s = qx & 0x80000000u;
    uint32_t e = (qx >> 23) & 0xFFu;
    uint32_t m = qx & 0x7FFFFFu;

    if (e < 127u) {
        uint32_t adjusted_exponents = 127u - (e + 1u);
        uint32_t denorm_m = 0x400000u | (m >> 1);
        m = (adjusted_exponents >= 32u) ? 0u : (denorm_m >> adjusted_exponents);
    }

    e = ((e > 126u) ? e : 126u) - 126u;
    uint32_t e2m1_tmp = ((((e << 2) | (m >> 21)) + 1u) >> 1);
    if (e2m1_tmp > 0x7u) {
        e2m1_tmp = 0x7u;
    }

    return static_cast<uint8_t>((s >> 28) | e2m1_tmp);
}

__device__ __forceinline__ uint8_t amax_to_e8m0_scale(float amax) {
    // Match aiter dynamic_mxfp4_quant rounding to power-of-two scale.
    uint32_t bits = __float_as_uint(amax);
    bits = (bits + 0x200000u) & 0xFF800000u;
    int exp_unbiased = static_cast<int>((bits >> 23) & 0xFF) - 127;
    int scale_unbiased = exp_unbiased - 2;
    if (scale_unbiased < -127) {
        scale_unbiased = -127;
    }
    if (scale_unbiased > 127) {
        scale_unbiased = 127;
    }
    return static_cast<uint8_t>(scale_unbiased + 127);
}

__device__ __forceinline__ int64_t shuffled_scale_offset(
    int64_t row,
    int64_t col,
    int64_t scale_n_pad) {
    int64_t bs_offs_0 = row / 32;
    int64_t bs_offs_1 = row % 32;
    int64_t bs_offs_2 = bs_offs_1 % 16;
    bs_offs_1 = bs_offs_1 / 16;

    int64_t bs_offs_3 = col / 8;
    int64_t bs_offs_4 = col % 8;
    int64_t bs_offs_5 = bs_offs_4 % 4;
    bs_offs_4 = bs_offs_4 / 4;

    return bs_offs_1 + bs_offs_4 * 2 + bs_offs_2 * 4 + bs_offs_5 * 64 +
           bs_offs_3 * 256 + bs_offs_0 * 32 * scale_n_pad;
}

__device__ __forceinline__ int64_t get_b_scale_shuffled_offset(
    int64_t n,
    int64_t k_blk,
    int64_t scale_k_pad) {
    int64_t n_outer = n / 32;
    int64_t n_inner = n % 32;
    int64_t k_outer = k_blk / 8;
    int64_t k_inner = k_blk % 8;

    int64_t n_16 = n_inner % 16;
    int64_t n_2 = n_inner / 16;
    int64_t k_4 = k_inner % 4;
    int64_t k_2 = k_inner / 4;

    return n_2 + (k_2 * 2) + (n_16 * 4) + (k_4 * 64) + (k_outer * 256) +
           (n_outer * 32 * scale_k_pad);
}

__device__ __forceinline__ int64_t get_b_fp4_shuffled_offset(
    int64_t n,
    int64_t k_fp4,
    int64_t k_fp4_pad,
    int64_t inner_mode) {
    int64_t n_blk = n / 16;
    int64_t k_blk = k_fp4 / 16;
    int64_t n_in = n % 16;
    int64_t k_in = k_fp4 % 16;

    int64_t blk_stride = k_fp4_pad / 16;
    int64_t blk_offset = (n_blk * blk_stride + k_blk) * 256;
    int64_t inner_offset = (inner_mode == 0) ? (k_in * 16 + n_in) : (n_in * 16 + k_in);
    return blk_offset + inner_offset;
}

__device__ __forceinline__ float e8m0_to_f32_fast(uint8_t e) {
    if (e == 0) {
        return __uint_as_float(0x00400000u);
    }
    if (e == 0xFF) {
        return __uint_as_float(0x7F800001u);
    }
    return __uint_as_float(static_cast<uint32_t>(e) << 23);
}

using floatx4 = __attribute__((__vector_size__(4 * sizeof(float)))) float;
using bit16x4 = __attribute__((__vector_size__(4 * sizeof(uint16_t)))) uint16_t;
using bit16x8 = __attribute__((__vector_size__(8 * sizeof(uint16_t)))) uint16_t;
using int32x8 = __attribute__((__vector_size__(8 * sizeof(int32_t)))) int32_t;
struct B16x8 {
    bit16x4 xy[2];
};

__device__ __forceinline__ floatx4 gcn_mfma16x16x32_bf16(
    const B16x8& a,
    const B16x8& b,
    const floatx4& c) {
#if defined(__gfx950__)
    bit16x8 ta = __builtin_shufflevector(a.xy[0], a.xy[1], 0, 1, 2, 3, 4, 5, 6, 7);
    bit16x8 tb = __builtin_shufflevector(b.xy[0], b.xy[1], 0, 1, 2, 3, 4, 5, 6, 7);
    return __builtin_amdgcn_mfma_f32_16x16x32_bf16(ta, tb, c, 0, 0, 0);
#else
    return c;
#endif
}

__device__ __forceinline__ int32x8 zero_int32x8() {
    return {0, 0, 0, 0, 0, 0, 0, 0};
}

template <int ScaleIdxA, int ScaleIdxB>
__device__ __forceinline__ floatx4 gcn_mfma16x16x128_f4_scale(
    const int32x8& a,
    const int32x8& b,
    const floatx4& c,
    int packed_scale_a,
    int packed_scale_b) {
#if defined(__gfx950__) && MXFP4_ENABLE_NATIVE_FP4_MFMA
    // LLVM AMDGPU docs: the last four operands are
    //   scale_idx_a, scale_values_a, scale_idx_b, scale_values_b
    // where scale_values_* is a per-lane VGPR value and scale_idx_* is the
    // wave-uniform 2-bit byte selector. The selected byte across all 64 lanes
    // forms the 64 scale entries consumed by one scaled MFMA instruction.
    // ROCm/clang also requires scale_idx_* to be compile-time immediates.
    // For FP4 x FP4, both operand format codes are 4.
    static_assert(0 <= ScaleIdxA && ScaleIdxA < 4, "ScaleIdxA must be in [0, 3]");
    static_assert(0 <= ScaleIdxB && ScaleIdxB < 4, "ScaleIdxB must be in [0, 3]");
    return __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
        a, b, c, 4, 4, ScaleIdxA, packed_scale_a, ScaleIdxB, packed_scale_b);
#else
    return c;
#endif
}

__device__ __forceinline__ int pack_e8m0x4(
    uint8_t s0,
    uint8_t s1,
    uint8_t s2,
    uint8_t s3) {
    return static_cast<int>(
        static_cast<uint32_t>(s0) |
        (static_cast<uint32_t>(s1) << 8) |
        (static_cast<uint32_t>(s2) << 16) |
        (static_cast<uint32_t>(s3) << 24));
}

__device__ __forceinline__ int splat_e8m0(uint8_t s) {
    uint32_t v = static_cast<uint32_t>(s);
    return static_cast<int>(v | (v << 8) | (v << 16) | (v << 24));
}

__device__ __forceinline__ int smem_swizzled_dword(int logical_row, int logical_dword) {
    return logical_dword ^ (logical_row & SMEM_SWIZZLE_MASK);
}

__device__ __forceinline__ int smem_swizzled_byte_offset(int logical_row, int logical_byte) {
    int logical_dword = logical_byte >> 2;
    int byte_in_dword = logical_byte & 3;
    int physical_dword = smem_swizzled_dword(logical_row, logical_dword);
    return logical_row * SMEM_STRIDE + physical_dword * static_cast<int>(sizeof(uint32_t)) +
           byte_in_dword;
}

__device__ __forceinline__ void store_uint4_to_smem_swizzled(
    uint8_t* smem_dst,
    int logical_row,
    const uint4& v) {
    uint32_t* row_ptr = reinterpret_cast<uint32_t*>(smem_dst + logical_row * SMEM_STRIDE);
    row_ptr[smem_swizzled_dword(logical_row, 0)] = v.x;
    row_ptr[smem_swizzled_dword(logical_row, 1)] = v.y;
    row_ptr[smem_swizzled_dword(logical_row, 2)] = v.z;
    row_ptr[smem_swizzled_dword(logical_row, 3)] = v.w;
}

__device__ __forceinline__ void store_u8_to_smem_swizzled(
    uint8_t* smem_dst,
    int logical_row,
    int logical_byte,
    uint8_t value) {
    smem_dst[smem_swizzled_byte_offset(logical_row, logical_byte)] = value;
}

__device__ __forceinline__ uint32_t load_u32_from_smem_swizzled(
    const uint8_t* smem_src,
    int logical_row,
    int logical_byte) {
    const uint32_t* row_ptr =
        reinterpret_cast<const uint32_t*>(smem_src + logical_row * SMEM_STRIDE);
    return row_ptr[smem_swizzled_dword(logical_row, logical_byte >> 2)];
}

template <bool WRITE_BF16_OUT>
__device__ __forceinline__ void store_output_value(
    float* workspace,
    __hip_bfloat16* out,
    int64_t workspace_stride,
    int64_t n,
    int64_t split_idx,
    int64_t row,
    int64_t col,
    float value) {
    int64_t out_off = row * n + col;
    if constexpr (WRITE_BF16_OUT) {
        out[out_off] = static_cast<__hip_bfloat16>(value);
    } else {
        workspace[split_idx * workspace_stride + out_off] = value;
    }
}

template <bool SHUFFLED_B, int SHUFFLE_INNER_MODE>
__device__ __forceinline__ uint4 load_b_tile_vec(
    const uint8_t* b_u8,
    int tid,
    int64_t tile_n,
    int64_t n,
    int64_t k2,
    int64_t k_base) {
    uint4 v = {0u, 0u, 0u, 0u};

    if constexpr (!SHUFFLED_B) {
        int64_t global_col = tile_n + tid;
        if (global_col < n) {
            v = *reinterpret_cast<const uint4*>(b_u8 + global_col * k2 + k_base);
        }
    } else if constexpr (SHUFFLE_INNER_MODE == 1) {
        int64_t global_col = tile_n + tid;
        if (global_col < n) {
            int64_t boff = get_b_fp4_shuffled_offset(global_col, k_base, k2, 1);
            v = *reinterpret_cast<const uint4*>(b_u8 + boff);
        }
    } else {
        int col_group = tid / PACKED_BLOCK_K;
        int k_in = tid % PACKED_BLOCK_K;
        int local_col_base = col_group * 16;
        int64_t global_col_base = tile_n + local_col_base;
        uint8_t* v_bytes = reinterpret_cast<uint8_t*>(&v);

        if (global_col_base + 15 < n) {
            int64_t boff = get_b_fp4_shuffled_offset(global_col_base, k_base + k_in, k2, 0);
            v = *reinterpret_cast<const uint4*>(b_u8 + boff);
        } else {
            #pragma unroll
            for (int c = 0; c < 16; ++c) {
                int64_t global_col = global_col_base + c;
                if (global_col < n) {
                    int64_t boff = get_b_fp4_shuffled_offset(global_col, k_base + k_in, k2, 0);
                    v_bytes[c] = b_u8[boff];
                }
            }
        }
    }

    return v;
}

template <bool SHUFFLED_B, int SHUFFLE_INNER_MODE>
__device__ __forceinline__ void store_b_tile_vec(
    uint8_t* smem_dst,
    int tid,
    const uint4& v) {
    if constexpr (!SHUFFLED_B || SHUFFLE_INNER_MODE == 1) {
        store_uint4_to_smem_swizzled(smem_dst, tid, v);
    } else {
        int col_group = tid / PACKED_BLOCK_K;
        int k_in = tid % PACKED_BLOCK_K;
        int local_col_base = col_group * 16;
        const uint8_t* src = reinterpret_cast<const uint8_t*>(&v);
        #pragma unroll
        for (int c = 0; c < 16; ++c) {
            store_u8_to_smem_swizzled(smem_dst, local_col_base + c, k_in, src[c]);
        }
    }
}

template <bool SHUFFLED_B, int SHUFFLE_INNER_MODE>
__device__ __forceinline__ void clear_gemm_stage_smem(
    uint8_t* smem_a,
    uint8_t* smem_b,
    uint8_t* smem_a_scale,
    uint8_t* smem_b_scale,
    int tid) {
    const uint4 zero = {0u, 0u, 0u, 0u};
    if (tid < A_TILE_VEC_LOADS) {
        store_uint4_to_smem_swizzled(smem_a, tid, zero);
        smem_a_scale[tid] = 127;
    }
    if (tid < B_TILE_VEC_LOADS) {
        store_b_tile_vec<SHUFFLED_B, SHUFFLE_INNER_MODE>(smem_b, tid, zero);
    }
    if (tid < GEMM_BLOCK_N) {
        smem_b_scale[tid] = 127;
    }
}

template <bool SHUFFLED_B>
__device__ __forceinline__ uint8_t load_b_scale_code(
    const uint8_t* b_scale_u8,
    int64_t global_col,
    int64_t kb,
    int64_t n,
    int64_t k_blocks_pad) {
    if (global_col >= n) {
        return 127;
    }
    if constexpr (SHUFFLED_B) {
        int64_t soff = get_b_scale_shuffled_offset(global_col, kb, k_blocks_pad);
        return b_scale_u8[soff];
    } else {
        return b_scale_u8[global_col * k_blocks_pad + kb];
    }
}

__device__ __forceinline__ uint16_t fp4_e2m1_to_bf16_bits(uint8_t nib) {
    uint16_t sign = static_cast<uint16_t>(nib & 0x08u) << 12;
    uint16_t mag = static_cast<uint16_t>(nib & 0x07u);
    if (mag == 0) {
        return sign;
    }

    uint16_t exponent = static_cast<uint16_t>(126 + (mag >> 1));
    uint16_t mantissa = ((mag > 1) && (mag & 0x1u)) ? 0x40u : 0u;
    return static_cast<uint16_t>(sign | (exponent << 7) | mantissa);
}

__device__ __forceinline__ B16x8 unpack_fp4x8_to_bf16(uint32_t packed) {
    B16x8 reg;
    const uint8_t* bytes = reinterpret_cast<const uint8_t*>(&packed);
    #pragma unroll
    for (int i = 0; i < 8; ++i) {
        uint8_t byte = bytes[i >> 1];
        uint8_t nib = (i & 1) ? static_cast<uint8_t>(byte >> 4)
                              : static_cast<uint8_t>(byte & 0x0Fu);
        uint16_t bits = fp4_e2m1_to_bf16_bits(nib);
        if (i < 4) {
            reg.xy[0][i] = bits;
        } else {
            reg.xy[1][i - 4] = bits;
        }
    }
    return reg;
}

__device__ __forceinline__ void quantize_a_row_block_to_fp4(
    const __hip_bfloat16* a_bf16,
    int64_t global_row,
    int64_t k,
    int64_t k_block,
    int64_t m,
    uint4& out_fp4,
    uint8_t& out_scale) {
    out_fp4 = {0u, 0u, 0u, 0u};
    out_scale = 127;

    if (global_row >= m) {
        return;
    }

    const uint8_t* a_bytes = reinterpret_cast<const uint8_t*>(a_bf16);
    const uint8_t* row_ptr =
        a_bytes + (global_row * k + k_block * BLOCK_K) * static_cast<int64_t>(sizeof(__hip_bfloat16));

    uint4 chunks[4];
    #pragma unroll
    for (int vec_idx = 0; vec_idx < 4; ++vec_idx) {
        chunks[vec_idx] = reinterpret_cast<const uint4*>(row_ptr)[vec_idx];
    }

    float amax = 0.0f;
    #pragma unroll
    for (int vec_idx = 0; vec_idx < 4; ++vec_idx) {
        uint32_t words[4] = {chunks[vec_idx].x, chunks[vec_idx].y, chunks[vec_idx].z, chunks[vec_idx].w};
        #pragma unroll
        for (int word_idx = 0; word_idx < 4; ++word_idx) {
            uint32_t word = words[word_idx];
            float f0 = bf16_to_f32(static_cast<uint16_t>(word & 0xFFFFu));
            float f1 = bf16_to_f32(static_cast<uint16_t>((word >> 16) & 0xFFFFu));
            amax = fmaxf(amax, fabsf(f0));
            amax = fmaxf(amax, fabsf(f1));
        }
    }

    out_scale = amax_to_e8m0_scale(amax);
    float inv_scale = 1.0f / e8m0_to_f32_fast(out_scale);
    uint8_t* out_bytes = reinterpret_cast<uint8_t*>(&out_fp4);

    #pragma unroll
    for (int vec_idx = 0; vec_idx < 4; ++vec_idx) {
        uint32_t words[4] = {chunks[vec_idx].x, chunks[vec_idx].y, chunks[vec_idx].z, chunks[vec_idx].w};
        uint8_t vals[8];
        #pragma unroll
        for (int word_idx = 0; word_idx < 4; ++word_idx) {
            uint32_t word = words[word_idx];
            vals[word_idx * 2 + 0] =
                quantize_e2m1(bf16_to_f32(static_cast<uint16_t>(word & 0xFFFFu)) * inv_scale);
            vals[word_idx * 2 + 1] = quantize_e2m1(
                bf16_to_f32(static_cast<uint16_t>((word >> 16) & 0xFFFFu)) * inv_scale);
        }

        #pragma unroll
        for (int pair_idx = 0; pair_idx < 4; ++pair_idx) {
            uint8_t lo = vals[pair_idx * 2 + 0];
            uint8_t hi = vals[pair_idx * 2 + 1];
            out_bytes[vec_idx * 4 + pair_idx] = static_cast<uint8_t>((hi << 4) | lo);
        }
    }
}

template <bool SHUFFLED_B, int SHUFFLE_INNER_MODE>
__device__ __forceinline__ void load_gemm_stage_to_smem(
    const __hip_bfloat16* a_bf16,
    const uint8_t* b_u8,
    const uint8_t* b_scale_u8,
    uint8_t* smem_a,
    uint8_t* smem_b,
    uint8_t* smem_a_scale,
    uint8_t* smem_b_scale,
    int tid,
    int64_t tile_m,
    int64_t tile_n,
    int64_t m,
    int64_t n,
    int64_t k,
    int64_t k2,
    int64_t global_kb,
    int64_t k_blocks_pad) {
    const int64_t a_k_base = global_kb * PACKED_BLOCK_K;

    if (tid < A_TILE_VEC_LOADS) {
        int64_t global_row = tile_m + tid;
        uint4 a_vec = {0u, 0u, 0u, 0u};
        uint8_t a_scale = 127;
        quantize_a_row_block_to_fp4(a_bf16, global_row, k, global_kb, m, a_vec, a_scale);
        store_uint4_to_smem_swizzled(smem_a, tid, a_vec);
        smem_a_scale[tid] = a_scale;
    }

    if (tid < B_TILE_VEC_LOADS) {
        uint4 v = load_b_tile_vec<SHUFFLED_B, SHUFFLE_INNER_MODE>(b_u8, tid, tile_n, n, k2, a_k_base);
        store_b_tile_vec<SHUFFLED_B, SHUFFLE_INNER_MODE>(smem_b, tid, v);
    }

    if (tid < GEMM_BLOCK_N) {
        int64_t global_col = tile_n + tid;
        smem_b_scale[tid] =
            load_b_scale_code<SHUFFLED_B>(b_scale_u8, global_col, global_kb, n, k_blocks_pad);
    }
}

template <bool SHUFFLED_B, int SHUFFLE_INNER_MODE>
__device__ __forceinline__ void load_prequant_stage_to_smem(
    const uint8_t* a_fp4,
    const uint8_t* b_u8,
    const uint8_t* a_scale_u8,
    const uint8_t* b_scale_u8,
    uint8_t* smem_a,
    uint8_t* smem_b,
    uint8_t* smem_a_scale,
    uint8_t* smem_b_scale,
    int tid,
    int64_t tile_m,
    int64_t tile_n,
    int64_t m,
    int64_t n,
    int64_t k2,
    int64_t global_kb,
    int64_t k_blocks_pad) {
    const int64_t a_k_base = global_kb * PACKED_BLOCK_K;

    if (tid < A_TILE_VEC_LOADS) {
        uint4 a_vec = {0u, 0u, 0u, 0u};
        uint8_t a_scale = 127;
        int64_t global_row = tile_m + tid;
        if (global_row < m) {
            a_vec = *reinterpret_cast<const uint4*>(a_fp4 + global_row * k2 + a_k_base);
            a_scale = a_scale_u8[global_row * k_blocks_pad + global_kb];
        }
        store_uint4_to_smem_swizzled(smem_a, tid, a_vec);
        smem_a_scale[tid] = a_scale;
    }

    if (tid < B_TILE_VEC_LOADS) {
        uint4 v = load_b_tile_vec<SHUFFLED_B, SHUFFLE_INNER_MODE>(b_u8, tid, tile_n, n, k2, a_k_base);
        store_b_tile_vec<SHUFFLED_B, SHUFFLE_INNER_MODE>(smem_b, tid, v);
    }

    if (tid < GEMM_BLOCK_N) {
        int64_t global_col = tile_n + tid;
        smem_b_scale[tid] =
            load_b_scale_code<SHUFFLED_B>(b_scale_u8, global_col, global_kb, n, k_blocks_pad);
    }
}

__device__ __forceinline__ void accumulate_bf16_stage(
    const uint8_t* smem_a,
    const uint8_t* smem_b,
    const uint8_t* smem_a_scale,
    const uint8_t* smem_b_scale,
    int row_group,
    int local_a_row0,
    int local_a_row1,
    int local_b_col0,
    int local_b_col1,
    int local_row_block0,
    int local_row_block1,
    floatx4& c_acc_00,
    floatx4& c_acc_01,
    floatx4& c_acc_10,
    floatx4& c_acc_11) {
    int k_byte_base = row_group * 4;
    uint32_t a_pack_0 = load_u32_from_smem_swizzled(smem_a, local_a_row0, k_byte_base);
    uint32_t a_pack_1 = load_u32_from_smem_swizzled(smem_a, local_a_row1, k_byte_base);
    uint32_t b_pack_0 = load_u32_from_smem_swizzled(smem_b, local_b_col0, k_byte_base);
    uint32_t b_pack_1 = load_u32_from_smem_swizzled(smem_b, local_b_col1, k_byte_base);

    B16x8 a_reg_0 = unpack_fp4x8_to_bf16(a_pack_0);
    B16x8 a_reg_1 = unpack_fp4x8_to_bf16(a_pack_1);
    B16x8 b_reg_0 = unpack_fp4x8_to_bf16(b_pack_0);
    B16x8 b_reg_1 = unpack_fp4x8_to_bf16(b_pack_1);

    floatx4 temp_00 = {0.0f, 0.0f, 0.0f, 0.0f};
    floatx4 temp_01 = {0.0f, 0.0f, 0.0f, 0.0f};
    floatx4 temp_10 = {0.0f, 0.0f, 0.0f, 0.0f};
    floatx4 temp_11 = {0.0f, 0.0f, 0.0f, 0.0f};
    temp_00 = gcn_mfma16x16x32_bf16(a_reg_0, b_reg_0, temp_00);
    temp_01 = gcn_mfma16x16x32_bf16(a_reg_0, b_reg_1, temp_01);
    temp_10 = gcn_mfma16x16x32_bf16(a_reg_1, b_reg_0, temp_10);
    temp_11 = gcn_mfma16x16x32_bf16(a_reg_1, b_reg_1, temp_11);

    float a_scales_0[4];
    float a_scales_1[4];
    #pragma unroll
    for (int i = 0; i < 4; ++i) {
        a_scales_0[i] = e8m0_to_f32_fast(smem_a_scale[local_row_block0 + i]);
        a_scales_1[i] = e8m0_to_f32_fast(smem_a_scale[local_row_block1 + i]);
    }
    float b_scale_0 = e8m0_to_f32_fast(smem_b_scale[local_b_col0]);
    float b_scale_1 = e8m0_to_f32_fast(smem_b_scale[local_b_col1]);

    #pragma unroll
    for (int i = 0; i < 4; ++i) {
        c_acc_00[i] += temp_00[i] * (a_scales_0[i] * b_scale_0);
        c_acc_01[i] += temp_01[i] * (a_scales_0[i] * b_scale_1);
        c_acc_10[i] += temp_10[i] * (a_scales_1[i] * b_scale_0);
        c_acc_11[i] += temp_11[i] * (a_scales_1[i] * b_scale_1);
    }
}

__device__ __forceinline__ void accumulate_native_group(
    const uint8_t smem_a[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_M * SMEM_STRIDE],
    const uint8_t smem_b[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_N * SMEM_STRIDE],
    const uint8_t smem_a_scale[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_M],
    const uint8_t smem_b_scale[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_N],
    int row_group,
    int local_a_row0,
    int local_a_row1,
    int local_b_col0,
    int local_b_col1,
    floatx4& c_acc_00,
    floatx4& c_acc_01,
    floatx4& c_acc_10,
    floatx4& c_acc_11) {
    // Fragment contract for __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4:
    // each lane contributes exactly one 32-wide K block, selected by
    // `row_group = lane / 16`:
    //   row_group 0 -> K [0, 32)
    //   row_group 1 -> K [32, 64)
    //   row_group 2 -> K [64, 96)
    //   row_group 3 -> K [96, 128)
    // The lane's 16 packed FP4 bytes always live in the low 128 bits
    // (arg[0..3]) of the int32x8 fragment. The instruction assembles the full
    // 128-wide K dimension across all 64 lanes, so we must not shift the data
    // into stage-dependent dword slots inside the lane-local register.
    //
    // The scale operand follows the same distributed contract: each lane passes
    // the E8M0 code for its own 32-wide K block in byte0 of the int VGPR.
    // CK/Opus both use OPSEL=0 here, and CK notes that current backends read
    // byte0 regardless of OPSEL, so only the low byte is semantically relevant.
    const int stage_idx = row_group;

    int32x8 a_frag_0 = zero_int32x8();
    int32x8 a_frag_1 = zero_int32x8();
    int32x8 b_frag_0 = zero_int32x8();
    int32x8 b_frag_1 = zero_int32x8();

    int32_t* a0_words = reinterpret_cast<int32_t*>(&a_frag_0);
    int32_t* a1_words = reinterpret_cast<int32_t*>(&a_frag_1);
    int32_t* b0_words = reinterpret_cast<int32_t*>(&b_frag_0);
    int32_t* b1_words = reinterpret_cast<int32_t*>(&b_frag_1);

    #pragma unroll
    for (int d = 0; d < 4; ++d) {
        int byte_off = d * 4;
        a0_words[d] = static_cast<int32_t>(
            load_u32_from_smem_swizzled(smem_a[stage_idx], local_a_row0, byte_off));
        a1_words[d] = static_cast<int32_t>(
            load_u32_from_smem_swizzled(smem_a[stage_idx], local_a_row1, byte_off));
        b0_words[d] = static_cast<int32_t>(
            load_u32_from_smem_swizzled(smem_b[stage_idx], local_b_col0, byte_off));
        b1_words[d] = static_cast<int32_t>(
            load_u32_from_smem_swizzled(smem_b[stage_idx], local_b_col1, byte_off));
    }

    const int packed_a_scale_0 = static_cast<int>(static_cast<uint32_t>(smem_a_scale[stage_idx][local_a_row0]));
    const int packed_a_scale_1 = static_cast<int>(static_cast<uint32_t>(smem_a_scale[stage_idx][local_a_row1]));
    const int packed_b_scale_0 = static_cast<int>(static_cast<uint32_t>(smem_b_scale[stage_idx][local_b_col0]));
    const int packed_b_scale_1 = static_cast<int>(static_cast<uint32_t>(smem_b_scale[stage_idx][local_b_col1]));

    c_acc_00 = gcn_mfma16x16x128_f4_scale<0, 0>(
        a_frag_0, b_frag_0, c_acc_00, packed_a_scale_0, packed_b_scale_0);
    c_acc_01 = gcn_mfma16x16x128_f4_scale<0, 0>(
        a_frag_0, b_frag_1, c_acc_01, packed_a_scale_0, packed_b_scale_1);
    c_acc_10 = gcn_mfma16x16x128_f4_scale<0, 0>(
        a_frag_1, b_frag_0, c_acc_10, packed_a_scale_1, packed_b_scale_0);
    c_acc_11 = gcn_mfma16x16x128_f4_scale<0, 0>(
        a_frag_1, b_frag_1, c_acc_11, packed_a_scale_1, packed_b_scale_1);
}

__device__ __forceinline__ void accumulate_tail_stages(
    const uint8_t smem_a[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_M * SMEM_STRIDE],
    const uint8_t smem_b[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_N * SMEM_STRIDE],
    const uint8_t smem_a_scale[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_M],
    const uint8_t smem_b_scale[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_N],
    int tail_k_blocks,
    int row_group,
    int local_a_row0,
    int local_a_row1,
    int local_b_col0,
    int local_b_col1,
    int local_row_block0,
    int local_row_block1,
    floatx4& c_acc_00,
    floatx4& c_acc_01,
    floatx4& c_acc_10,
    floatx4& c_acc_11) {
    #pragma unroll
    for (int stage = 0; stage < NATIVE_MFMA_K_BLOCKS; ++stage) {
        if (stage < tail_k_blocks) {
            accumulate_bf16_stage(
                smem_a[stage],
                smem_b[stage],
                smem_a_scale[stage],
                smem_b_scale[stage],
                row_group,
                local_a_row0,
                local_a_row1,
                local_b_col0,
                local_b_col1,
                local_row_block0,
                local_row_block1,
                c_acc_00,
                c_acc_01,
                c_acc_10,
                c_acc_11);
        }
    }
}

__global__ void quant_mxfp4_kernel(
    const __hip_bfloat16* x,
    uint8_t* out_fp4,
    uint8_t* out_scale,
    int64_t m,
    int64_t m_pad,
    int64_t k,
    int64_t k_blocks_valid,
    int64_t k_blocks_pad) {
    int64_t linear = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
    int64_t total = m_pad * k_blocks_pad;
    if (linear >= total) {
        return;
    }

    int64_t row = linear / k_blocks_pad;
    int64_t kb = linear % k_blocks_pad;
    int64_t scale_off = shuffled_scale_offset(row, kb, k_blocks_pad);

    if (row >= m || kb >= k_blocks_valid) {
        out_scale[scale_off] = 127;

        if (row < m_pad && kb < k_blocks_valid) {
            int64_t out_base = row * (k / 2) + kb * (BLOCK_K / 2);
            #pragma unroll
            for (int i = 0; i < BLOCK_K / 2; ++i) {
                out_fp4[out_base + i] = 0;
            }
        }
        return;
    }

    int64_t in_base = row * k + kb * BLOCK_K;

    float vals[BLOCK_K];
    float amax = 0.0f;

    const uint8_t* x_bytes = reinterpret_cast<const uint8_t*>(x);
    const uint8_t* row_ptr = x_bytes + in_base * static_cast<int64_t>(sizeof(__hip_bfloat16));

    #pragma unroll
    for (int vec_idx = 0; vec_idx < 4; ++vec_idx) {
        uint4 v4 = *reinterpret_cast<const uint4*>(row_ptr + vec_idx * 16);
        uint32_t w0 = v4.x;
        uint32_t w1 = v4.y;
        uint32_t w2 = v4.z;
        uint32_t w3 = v4.w;

        uint16_t b0 = static_cast<uint16_t>(w0 & 0xFFFFu);
        uint16_t b1 = static_cast<uint16_t>((w0 >> 16) & 0xFFFFu);
        uint16_t b2 = static_cast<uint16_t>(w1 & 0xFFFFu);
        uint16_t b3 = static_cast<uint16_t>((w1 >> 16) & 0xFFFFu);
        uint16_t b4 = static_cast<uint16_t>(w2 & 0xFFFFu);
        uint16_t b5 = static_cast<uint16_t>((w2 >> 16) & 0xFFFFu);
        uint16_t b6 = static_cast<uint16_t>(w3 & 0xFFFFu);
        uint16_t b7 = static_cast<uint16_t>((w3 >> 16) & 0xFFFFu);

        int base = vec_idx * 8;
        float f0 = bf16_to_f32(b0);
        float f1 = bf16_to_f32(b1);
        float f2 = bf16_to_f32(b2);
        float f3 = bf16_to_f32(b3);
        float f4 = bf16_to_f32(b4);
        float f5 = bf16_to_f32(b5);
        float f6 = bf16_to_f32(b6);
        float f7 = bf16_to_f32(b7);

        vals[base + 0] = f0;
        vals[base + 1] = f1;
        vals[base + 2] = f2;
        vals[base + 3] = f3;
        vals[base + 4] = f4;
        vals[base + 5] = f5;
        vals[base + 6] = f6;
        vals[base + 7] = f7;

        amax = fmaxf(amax, fabsf(f0));
        amax = fmaxf(amax, fabsf(f1));
        amax = fmaxf(amax, fabsf(f2));
        amax = fmaxf(amax, fabsf(f3));
        amax = fmaxf(amax, fabsf(f4));
        amax = fmaxf(amax, fabsf(f5));
        amax = fmaxf(amax, fabsf(f6));
        amax = fmaxf(amax, fabsf(f7));
    }

    uint8_t scale_code = amax_to_e8m0_scale(amax);
    out_scale[scale_off] = scale_code;
    float scale = e8m0_to_f32_fast(scale_code);

    float inv_scale = 1.0f / scale;
    int64_t out_base = row * (k / 2) + kb * (BLOCK_K / 2);

    #pragma unroll
    for (int i = 0; i < BLOCK_K / 2; ++i) {
        uint8_t lo = quantize_e2m1(vals[2 * i] * inv_scale);
        uint8_t hi = quantize_e2m1(vals[2 * i + 1] * inv_scale);
        out_fp4[out_base + i] = static_cast<uint8_t>((hi << 4) | lo);
    }
}

template <bool SHUFFLED_B, int SHUFFLE_INNER_MODE, bool WRITE_BF16_OUT>
__global__ void gemm_mxfp4_prequant_kernel(
    const uint8_t* a_fp4,
    const uint8_t* b_u8,
    const uint8_t* a_scale_u8,
    const uint8_t* b_scale_u8,
    float* workspace,
    __hip_bfloat16* out,
    int64_t workspace_stride,
    int64_t m,
    int64_t n,
    int64_t k2,
    int64_t k_blocks_valid,
    int64_t k_blocks_pad,
    int64_t split_k) {
    const int tid = static_cast<int>(threadIdx.x);
    const int wave_id = tid / WAVE_SIZE;
    const int lane = tid & (WAVE_SIZE - 1);
    const int lane16 = lane & 15;
    const int row_group = lane >> 4;  // 0..3
    const int wave_row = wave_id >> 1;
    const int wave_col = wave_id & 1;
    int64_t tile_m = static_cast<int64_t>(blockIdx.y) * GEMM_BLOCK_M;
    int64_t tile_n = static_cast<int64_t>(blockIdx.x) * GEMM_BLOCK_N;
    const int local_wave_m = wave_row * WAVE_TILE_M;
    const int local_wave_n = wave_col * WAVE_TILE_N;
    const int local_row_block0 = local_wave_m + row_group * 4;
    const int local_row_block1 = local_row_block0 + MFMA_TILE_M;
    const int local_a_row0 = local_wave_m + lane16;
    const int local_a_row1 = local_a_row0 + MFMA_TILE_M;
    const int local_b_col0 = local_wave_n + lane16;
    const int local_b_col1 = local_b_col0 + MFMA_TILE_N;
    const int64_t col0 = tile_n + local_b_col0;
    const int64_t col1 = tile_n + local_b_col1;
    const int64_t split_idx = static_cast<int64_t>(blockIdx.z);

    floatx4 c_acc_00 = {0.0f, 0.0f, 0.0f, 0.0f};
    floatx4 c_acc_01 = {0.0f, 0.0f, 0.0f, 0.0f};
    floatx4 c_acc_10 = {0.0f, 0.0f, 0.0f, 0.0f};
    floatx4 c_acc_11 = {0.0f, 0.0f, 0.0f, 0.0f};

    __shared__ __align__(16) uint8_t smem_A[PREFETCH_STAGES][GEMM_BLOCK_M * SMEM_STRIDE];
    __shared__ __align__(16) uint8_t smem_B[PREFETCH_STAGES][GEMM_BLOCK_N * SMEM_STRIDE];
    __shared__ uint8_t smem_A_scale[PREFETCH_STAGES][GEMM_BLOCK_M];
    __shared__ uint8_t smem_B_scale[PREFETCH_STAGES][GEMM_BLOCK_N];

    int64_t kb_per_split = (k_blocks_valid + split_k - 1) / split_k;
    int64_t kb_start = static_cast<int64_t>(blockIdx.z) * kb_per_split;
    int64_t kb_end = kb_start + kb_per_split;
    if (kb_end > k_blocks_valid) {
        kb_end = k_blocks_valid;
    }

    if (kb_start >= kb_end) {
        if constexpr (!WRITE_BF16_OUT) {
            if (col0 < n) {
                #pragma unroll
                for (int i = 0; i < 4; ++i) {
                    int64_t row = tile_m + local_row_block0 + i;
                    if (row < m) {
                        store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col0, 0.0f);
                    }
                    row = tile_m + local_row_block1 + i;
                    if (row < m) {
                        store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col0, 0.0f);
                    }
                }
            }
            if (col1 < n) {
                #pragma unroll
                for (int i = 0; i < 4; ++i) {
                    int64_t row = tile_m + local_row_block0 + i;
                    if (row < m) {
                        store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col1, 0.0f);
                    }
                    row = tile_m + local_row_block1 + i;
                    if (row < m) {
                        store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col1, 0.0f);
                    }
                }
            }
        }
        return;
    }

    const int64_t tiles_in_split = kb_end - kb_start;

    const int64_t prologue_tiles = tiles_in_split < 2 ? tiles_in_split : 2;
    for (int64_t i = 0; i < prologue_tiles; ++i) {
        load_prequant_stage_to_smem<SHUFFLED_B, SHUFFLE_INNER_MODE>(
            a_fp4,
            b_u8,
            a_scale_u8,
            b_scale_u8,
            smem_A[static_cast<int>(i)],
            smem_B[static_cast<int>(i)],
            smem_A_scale[static_cast<int>(i)],
            smem_B_scale[static_cast<int>(i)],
            tid,
            tile_m,
            tile_n,
            m,
            n,
            k2,
            kb_start + i,
            k_blocks_pad);
    }
    __syncthreads();

    for (int64_t tile_idx = 0; tile_idx < tiles_in_split; ++tile_idx) {
        const int curr_buf = static_cast<int>(tile_idx % PREFETCH_STAGES);
        const int64_t prefetch_tile_idx = tile_idx + 2;
        const bool has_prefetch = prefetch_tile_idx < tiles_in_split;
        const int prefetch_buf = static_cast<int>(prefetch_tile_idx % PREFETCH_STAGES);
        const int64_t prefetch_kb = kb_start + prefetch_tile_idx;
        const int64_t prefetch_a_k_base = prefetch_kb * PACKED_BLOCK_K;

        uint4 next_a_v = {0u, 0u, 0u, 0u};
        uint4 next_b_v = {0u, 0u, 0u, 0u};
        uint8_t next_a_scale = 127;
        uint8_t next_b_scale = 127;

        if (has_prefetch && tid < A_TILE_VEC_LOADS) {
            int64_t global_row = tile_m + tid;
            if (global_row < m) {
                next_a_v = *reinterpret_cast<const uint4*>(a_fp4 + global_row * k2 + prefetch_a_k_base);
                next_a_scale = a_scale_u8[global_row * k_blocks_pad + prefetch_kb];
            }
        }

        if (has_prefetch && tid < B_TILE_VEC_LOADS) {
            next_b_v = load_b_tile_vec<SHUFFLED_B, SHUFFLE_INNER_MODE>(
                b_u8, tid, tile_n, n, k2, prefetch_a_k_base);
        }

        if (has_prefetch && tid < GEMM_BLOCK_N) {
            int64_t global_col = tile_n + tid;
            next_b_scale = load_b_scale_code<SHUFFLED_B>(
                b_scale_u8, global_col, prefetch_kb, n, k_blocks_pad);
        }

        accumulate_bf16_stage(
            smem_A[curr_buf],
            smem_B[curr_buf],
            smem_A_scale[curr_buf],
            smem_B_scale[curr_buf],
            row_group,
            local_a_row0,
            local_a_row1,
            local_b_col0,
            local_b_col1,
            local_row_block0,
            local_row_block1,
            c_acc_00,
            c_acc_01,
            c_acc_10,
            c_acc_11);

        if (has_prefetch) {
            if (tid < A_TILE_VEC_LOADS) {
                store_uint4_to_smem_swizzled(smem_A[prefetch_buf], tid, next_a_v);
                smem_A_scale[prefetch_buf][tid] = next_a_scale;
            }

            if (tid < B_TILE_VEC_LOADS) {
                store_b_tile_vec<SHUFFLED_B, SHUFFLE_INNER_MODE>(smem_B[prefetch_buf], tid, next_b_v);
            }

            if (tid < GEMM_BLOCK_N) {
                smem_B_scale[prefetch_buf][tid] = next_b_scale;
            }
        }
        __syncthreads();
    }

    if (col0 < n) {
        #pragma unroll
        for (int i = 0; i < 4; ++i) {
            int64_t row = tile_m + local_row_block0 + i;
            if (row < m) {
                store_output_value<WRITE_BF16_OUT>(
                    workspace, out, workspace_stride, n, split_idx, row, col0, c_acc_00[i]);
            }
            row = tile_m + local_row_block1 + i;
            if (row < m) {
                store_output_value<WRITE_BF16_OUT>(
                    workspace, out, workspace_stride, n, split_idx, row, col0, c_acc_10[i]);
            }
        }
    }
    if (col1 < n) {
        #pragma unroll
        for (int i = 0; i < 4; ++i) {
            int64_t row = tile_m + local_row_block0 + i;
            if (row < m) {
                store_output_value<WRITE_BF16_OUT>(
                    workspace, out, workspace_stride, n, split_idx, row, col1, c_acc_01[i]);
            }
            row = tile_m + local_row_block1 + i;
            if (row < m) {
                store_output_value<WRITE_BF16_OUT>(
                    workspace, out, workspace_stride, n, split_idx, row, col1, c_acc_11[i]);
            }
        }
    }
}

template <bool SHUFFLED_B, int SHUFFLE_INNER_MODE, bool WRITE_BF16_OUT>
__global__ void gemm_mxfp4_kernel(
    const __hip_bfloat16* a_bf16,
    const uint8_t* b_u8,
    const uint8_t* b_scale_u8,
    float* workspace,
    __hip_bfloat16* out,
    int64_t workspace_stride,
    int64_t m,
    int64_t n,
    int64_t k2,
    int64_t k_blocks_valid,
    int64_t k_blocks_pad,
    int64_t split_k) {
    const int tid = static_cast<int>(threadIdx.x);
    const int wave_id = tid / WAVE_SIZE;
    const int lane = tid & (WAVE_SIZE - 1);
    const int lane16 = lane & 15;
    const int row_group = lane >> 4;  // 0..3
    const int wave_row = wave_id >> 1;
    const int wave_col = wave_id & 1;
    int64_t tile_m = static_cast<int64_t>(blockIdx.y) * GEMM_BLOCK_M;
    int64_t tile_n = static_cast<int64_t>(blockIdx.x) * GEMM_BLOCK_N;
    const int local_wave_m = wave_row * WAVE_TILE_M;
    const int local_wave_n = wave_col * WAVE_TILE_N;
    const int local_row_block0 = local_wave_m + row_group * 4;
    const int local_row_block1 = local_row_block0 + MFMA_TILE_M;
    const int local_a_row0 = local_wave_m + lane16;
    const int local_a_row1 = local_a_row0 + MFMA_TILE_M;
    const int local_b_col0 = local_wave_n + lane16;
    const int local_b_col1 = local_b_col0 + MFMA_TILE_N;
    const int64_t col0 = tile_n + local_b_col0;
    const int64_t col1 = tile_n + local_b_col1;
    const int64_t split_idx = static_cast<int64_t>(blockIdx.z);

    floatx4 c_acc_00 = {0.0f, 0.0f, 0.0f, 0.0f};
    floatx4 c_acc_01 = {0.0f, 0.0f, 0.0f, 0.0f};
    floatx4 c_acc_10 = {0.0f, 0.0f, 0.0f, 0.0f};
    floatx4 c_acc_11 = {0.0f, 0.0f, 0.0f, 0.0f};

    __shared__ __align__(16) uint8_t smem_A[PREFETCH_STAGES][GEMM_BLOCK_M * SMEM_STRIDE];
    __shared__ __align__(16) uint8_t smem_B[PREFETCH_STAGES][GEMM_BLOCK_N * SMEM_STRIDE];
    __shared__ uint8_t smem_A_scale[PREFETCH_STAGES][GEMM_BLOCK_M];
    __shared__ uint8_t smem_B_scale[PREFETCH_STAGES][GEMM_BLOCK_N];

    int64_t kb_per_split = (k_blocks_valid + split_k - 1) / split_k;
    int64_t kb_start = static_cast<int64_t>(blockIdx.z) * kb_per_split;
    int64_t kb_end = kb_start + kb_per_split;
    if (kb_end > k_blocks_valid) {
        kb_end = k_blocks_valid;
    }

    if (kb_start >= kb_end) {
        if constexpr (!WRITE_BF16_OUT) {
            if (col0 < n) {
                #pragma unroll
                for (int i = 0; i < 4; ++i) {
                    int64_t row = tile_m + local_row_block0 + i;
                    if (row < m) {
                        store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col0, 0.0f);
                    }
                    row = tile_m + local_row_block1 + i;
                    if (row < m) {
                        store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col0, 0.0f);
                    }
                }
            }
            if (col1 < n) {
                #pragma unroll
                for (int i = 0; i < 4; ++i) {
                    int64_t row = tile_m + local_row_block0 + i;
                    if (row < m) {
                        store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col1, 0.0f);
                    }
                    row = tile_m + local_row_block1 + i;
                    if (row < m) {
                        store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col1, 0.0f);
                    }
                }
            }
        }
        return;
    }

    const int64_t tiles_in_split = kb_end - kb_start;
    const int64_t k = k2 * 2;

    const int64_t prologue_tiles = tiles_in_split < 2 ? tiles_in_split : 2;
    for (int64_t i = 0; i < prologue_tiles; ++i) {
        load_gemm_stage_to_smem<SHUFFLED_B, SHUFFLE_INNER_MODE>(
            a_bf16,
            b_u8,
            b_scale_u8,
            smem_A[static_cast<int>(i)],
            smem_B[static_cast<int>(i)],
            smem_A_scale[static_cast<int>(i)],
            smem_B_scale[static_cast<int>(i)],
            tid,
            tile_m,
            tile_n,
            m,
            n,
            k,
            k2,
            kb_start + i,
            k_blocks_pad);
    }
    __syncthreads();

    for (int64_t tile_idx = 0; tile_idx < tiles_in_split; ++tile_idx) {
        const int curr_buf = static_cast<int>(tile_idx % PREFETCH_STAGES);
        const int64_t prefetch_tile_idx = tile_idx + 2;
        const bool has_prefetch = prefetch_tile_idx < tiles_in_split;
        const int prefetch_buf = static_cast<int>(prefetch_tile_idx % PREFETCH_STAGES);
        const int64_t prefetch_kb = kb_start + prefetch_tile_idx;
        const int64_t prefetch_a_k_base = prefetch_kb * PACKED_BLOCK_K;

        uint4 next_a_v = {0u, 0u, 0u, 0u};
        uint4 next_b_v = {0u, 0u, 0u, 0u};
        uint8_t next_a_scale = 127;
        uint8_t next_b_scale = 127;

        if (has_prefetch && tid < A_TILE_VEC_LOADS) {
            int64_t global_row = tile_m + tid;
            quantize_a_row_block_to_fp4(a_bf16, global_row, k, prefetch_kb, m, next_a_v, next_a_scale);
        }

        if (has_prefetch && tid < B_TILE_VEC_LOADS) {
            next_b_v = load_b_tile_vec<SHUFFLED_B, SHUFFLE_INNER_MODE>(
                b_u8, tid, tile_n, n, k2, prefetch_a_k_base);
        }

        if (has_prefetch && tid < GEMM_BLOCK_N) {
            int64_t global_col = tile_n + tid;
            next_b_scale = load_b_scale_code<SHUFFLED_B>(
                b_scale_u8, global_col, prefetch_kb, n, k_blocks_pad);
        }

        accumulate_bf16_stage(
            smem_A[curr_buf],
            smem_B[curr_buf],
            smem_A_scale[curr_buf],
            smem_B_scale[curr_buf],
            row_group,
            local_a_row0,
            local_a_row1,
            local_b_col0,
            local_b_col1,
            local_row_block0,
            local_row_block1,
            c_acc_00,
            c_acc_01,
            c_acc_10,
            c_acc_11);

        if (has_prefetch) {
            if (tid < A_TILE_VEC_LOADS) {
                store_uint4_to_smem_swizzled(smem_A[prefetch_buf], tid, next_a_v);
                smem_A_scale[prefetch_buf][tid] = next_a_scale;
            }

            if (tid < B_TILE_VEC_LOADS) {
                store_b_tile_vec<SHUFFLED_B, SHUFFLE_INNER_MODE>(smem_B[prefetch_buf], tid, next_b_v);
            }

            if (tid < GEMM_BLOCK_N) {
                smem_B_scale[prefetch_buf][tid] = next_b_scale;
            }
        }
        __syncthreads();
    }

    if (col0 < n) {
        #pragma unroll
        for (int i = 0; i < 4; ++i) {
            int64_t row = tile_m + local_row_block0 + i;
            if (row < m) {
                store_output_value<WRITE_BF16_OUT>(
                    workspace, out, workspace_stride, n, split_idx, row, col0, c_acc_00[i]);
            }
            row = tile_m + local_row_block1 + i;
            if (row < m) {
                store_output_value<WRITE_BF16_OUT>(
                    workspace, out, workspace_stride, n, split_idx, row, col0, c_acc_10[i]);
            }
        }
    }
    if (col1 < n) {
        #pragma unroll
        for (int i = 0; i < 4; ++i) {
            int64_t row = tile_m + local_row_block0 + i;
            if (row < m) {
                store_output_value<WRITE_BF16_OUT>(
                    workspace, out, workspace_stride, n, split_idx, row, col1, c_acc_01[i]);
            }
            row = tile_m + local_row_block1 + i;
            if (row < m) {
                store_output_value<WRITE_BF16_OUT>(
                    workspace, out, workspace_stride, n, split_idx, row, col1, c_acc_11[i]);
            }
        }
    }
}

template <bool SHUFFLED_B, int SHUFFLE_INNER_MODE, bool WRITE_BF16_OUT>
__global__ void gemm_mxfp4_prequant_native_kernel(
    const uint8_t* a_fp4,
    const uint8_t* b_u8,
    const uint8_t* a_scale_u8,
    const uint8_t* b_scale_u8,
    float* workspace,
    __hip_bfloat16* out,
    int64_t workspace_stride,
    int64_t m,
    int64_t n,
    int64_t k2,
    int64_t k_blocks_valid,
    int64_t k_blocks_pad,
    int64_t split_k) {
    const int tid = static_cast<int>(threadIdx.x);
    const int wave_id = tid / WAVE_SIZE;
    const int lane = tid & (WAVE_SIZE - 1);
    const int lane16 = lane & 15;
    const int row_group = lane >> 4;
    const int wave_row = wave_id >> 1;
    const int wave_col = wave_id & 1;
    int64_t tile_m = static_cast<int64_t>(blockIdx.y) * GEMM_BLOCK_M;
    int64_t tile_n = static_cast<int64_t>(blockIdx.x) * GEMM_BLOCK_N;
    const int local_wave_m = wave_row * WAVE_TILE_M;
    const int local_wave_n = wave_col * WAVE_TILE_N;
    const int local_row_block0 = local_wave_m + row_group * 4;
    const int local_row_block1 = local_row_block0 + MFMA_TILE_M;
    const int local_a_row0 = local_wave_m + lane16;
    const int local_a_row1 = local_a_row0 + MFMA_TILE_M;
    const int local_b_col0 = local_wave_n + lane16;
    const int local_b_col1 = local_b_col0 + MFMA_TILE_N;
    const int64_t col0 = tile_n + local_b_col0;
    const int64_t col1 = tile_n + local_b_col1;
    const int64_t split_idx = static_cast<int64_t>(blockIdx.z);

    floatx4 c_acc_00 = {0.0f, 0.0f, 0.0f, 0.0f};
    floatx4 c_acc_01 = {0.0f, 0.0f, 0.0f, 0.0f};
    floatx4 c_acc_10 = {0.0f, 0.0f, 0.0f, 0.0f};
    floatx4 c_acc_11 = {0.0f, 0.0f, 0.0f, 0.0f};

    __shared__ __align__(16) uint8_t smem_A[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_M * SMEM_STRIDE];
    __shared__ __align__(16) uint8_t smem_B[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_N * SMEM_STRIDE];
    __shared__ uint8_t smem_A_scale[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_M];
    __shared__ uint8_t smem_B_scale[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_N];

    int64_t kb_per_split = (k_blocks_valid + split_k - 1) / split_k;
    int64_t kb_start = static_cast<int64_t>(blockIdx.z) * kb_per_split;
    int64_t kb_end = kb_start + kb_per_split;
    if (kb_end > k_blocks_valid) {
        kb_end = k_blocks_valid;
    }

    if (kb_start >= kb_end) {
        if constexpr (!WRITE_BF16_OUT) {
            if (col0 < n) {
                #pragma unroll
                for (int i = 0; i < 4; ++i) {
                    int64_t row = tile_m + local_row_block0 + i;
                    if (row < m) {
                        store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col0, 0.0f);
                    }
                    row = tile_m + local_row_block1 + i;
                    if (row < m) {
                        store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col0, 0.0f);
                    }
                }
            }
            if (col1 < n) {
                #pragma unroll
                for (int i = 0; i < 4; ++i) {
                    int64_t row = tile_m + local_row_block0 + i;
                    if (row < m) {
                        store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col1, 0.0f);
                    }
                    row = tile_m + local_row_block1 + i;
                    if (row < m) {
                        store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col1, 0.0f);
                    }
                }
            }
        }
        return;
    }

    const int64_t full_kb_end =
        kb_start + ((kb_end - kb_start) / NATIVE_MFMA_K_BLOCKS) * NATIVE_MFMA_K_BLOCKS;

    for (int64_t kb_group = kb_start; kb_group < full_kb_end; kb_group += NATIVE_MFMA_K_BLOCKS) {
        #pragma unroll
        for (int stage = 0; stage < NATIVE_MFMA_K_BLOCKS; ++stage) {
            load_prequant_stage_to_smem<SHUFFLED_B, SHUFFLE_INNER_MODE>(
                a_fp4,
                b_u8,
                a_scale_u8,
                b_scale_u8,
                smem_A[stage],
                smem_B[stage],
                smem_A_scale[stage],
                smem_B_scale[stage],
                tid,
                tile_m,
                tile_n,
                m,
                n,
                k2,
                kb_group + stage,
                k_blocks_pad);
        }
        __syncthreads();

#if defined(__gfx950__) && MXFP4_ENABLE_NATIVE_FP4_MFMA
        accumulate_native_group(
            smem_A,
            smem_B,
            smem_A_scale,
            smem_B_scale,
            row_group,
            local_a_row0,
            local_a_row1,
            local_b_col0,
            local_b_col1,
            c_acc_00,
            c_acc_01,
            c_acc_10,
            c_acc_11);
#else
        accumulate_tail_stages(
            smem_A,
            smem_B,
            smem_A_scale,
            smem_B_scale,
            NATIVE_MFMA_K_BLOCKS,
            row_group,
            local_a_row0,
            local_a_row1,
            local_b_col0,
            local_b_col1,
            local_row_block0,
            local_row_block1,
            c_acc_00,
            c_acc_01,
            c_acc_10,
            c_acc_11);
#endif
        __syncthreads();
    }

    const int tail_k_blocks = static_cast<int>(kb_end - full_kb_end);
    if (tail_k_blocks > 0) {
        #pragma unroll
        for (int stage = 0; stage < NATIVE_MFMA_K_BLOCKS; ++stage) {
            if (stage < tail_k_blocks) {
                load_prequant_stage_to_smem<SHUFFLED_B, SHUFFLE_INNER_MODE>(
                    a_fp4,
                    b_u8,
                    a_scale_u8,
                    b_scale_u8,
                    smem_A[stage],
                    smem_B[stage],
                    smem_A_scale[stage],
                    smem_B_scale[stage],
                    tid,
                    tile_m,
                    tile_n,
                    m,
                    n,
                    k2,
                    full_kb_end + stage,
                    k_blocks_pad);
            }
        }
        __syncthreads();

        accumulate_tail_stages(
            smem_A,
            smem_B,
            smem_A_scale,
            smem_B_scale,
            tail_k_blocks,
            row_group,
            local_a_row0,
            local_a_row1,
            local_b_col0,
            local_b_col1,
            local_row_block0,
            local_row_block1,
            c_acc_00,
            c_acc_01,
            c_acc_10,
            c_acc_11);
    }

    if (col0 < n) {
        #pragma unroll
        for (int i = 0; i < 4; ++i) {
            int64_t row = tile_m + local_row_block0 + i;
            if (row < m) {
                store_output_value<WRITE_BF16_OUT>(
                    workspace, out, workspace_stride, n, split_idx, row, col0, c_acc_00[i]);
            }
            row = tile_m + local_row_block1 + i;
            if (row < m) {
                store_output_value<WRITE_BF16_OUT>(
                    workspace, out, workspace_stride, n, split_idx, row, col0, c_acc_10[i]);
            }
        }
    }
    if (col1 < n) {
        #pragma unroll
        for (int i = 0; i < 4; ++i) {
            int64_t row = tile_m + local_row_block0 + i;
            if (row < m) {
                store_output_value<WRITE_BF16_OUT>(
                    workspace, out, workspace_stride, n, split_idx, row, col1, c_acc_01[i]);
            }
            row = tile_m + local_row_block1 + i;
            if (row < m) {
                store_output_value<WRITE_BF16_OUT>(
                    workspace, out, workspace_stride, n, split_idx, row, col1, c_acc_11[i]);
            }
        }
    }
}

template <bool SHUFFLED_B, int SHUFFLE_INNER_MODE, bool WRITE_BF16_OUT>
__global__ void gemm_mxfp4_native_kernel(
    const __hip_bfloat16* a_bf16,
    const uint8_t* b_u8,
    const uint8_t* b_scale_u8,
    float* workspace,
    __hip_bfloat16* out,
    int64_t workspace_stride,
    int64_t m,
    int64_t n,
    int64_t k2,
    int64_t k_blocks_valid,
    int64_t k_blocks_pad,
    int64_t split_k) {
    const int tid = static_cast<int>(threadIdx.x);
    const int wave_id = tid / WAVE_SIZE;
    const int lane = tid & (WAVE_SIZE - 1);
    const int lane16 = lane & 15;
    const int row_group = lane >> 4;
    const int wave_row = wave_id >> 1;
    const int wave_col = wave_id & 1;
    int64_t tile_m = static_cast<int64_t>(blockIdx.y) * GEMM_BLOCK_M;
    int64_t tile_n = static_cast<int64_t>(blockIdx.x) * GEMM_BLOCK_N;
    const int local_wave_m = wave_row * WAVE_TILE_M;
    const int local_wave_n = wave_col * WAVE_TILE_N;
    const int local_row_block0 = local_wave_m + row_group * 4;
    const int local_row_block1 = local_row_block0 + MFMA_TILE_M;
    const int local_a_row0 = local_wave_m + lane16;
    const int local_a_row1 = local_a_row0 + MFMA_TILE_M;
    const int local_b_col0 = local_wave_n + lane16;
    const int local_b_col1 = local_b_col0 + MFMA_TILE_N;
    const int64_t col0 = tile_n + local_b_col0;
    const int64_t col1 = tile_n + local_b_col1;
    const int64_t split_idx = static_cast<int64_t>(blockIdx.z);

    floatx4 c_acc_00 = {0.0f, 0.0f, 0.0f, 0.0f};
    floatx4 c_acc_01 = {0.0f, 0.0f, 0.0f, 0.0f};
    floatx4 c_acc_10 = {0.0f, 0.0f, 0.0f, 0.0f};
    floatx4 c_acc_11 = {0.0f, 0.0f, 0.0f, 0.0f};

    __shared__ __align__(16) uint8_t smem_A[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_M * SMEM_STRIDE];
    __shared__ __align__(16) uint8_t smem_B[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_N * SMEM_STRIDE];
    __shared__ uint8_t smem_A_scale[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_M];
    __shared__ uint8_t smem_B_scale[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_N];

    int64_t kb_per_split = (k_blocks_valid + split_k - 1) / split_k;
    int64_t kb_start = static_cast<int64_t>(blockIdx.z) * kb_per_split;
    int64_t kb_end = kb_start + kb_per_split;
    if (kb_end > k_blocks_valid) {
        kb_end = k_blocks_valid;
    }

    if (kb_start >= kb_end) {
        if constexpr (!WRITE_BF16_OUT) {
            if (col0 < n) {
                #pragma unroll
                for (int i = 0; i < 4; ++i) {
                    int64_t row = tile_m + local_row_block0 + i;
                    if (row < m) {
                        store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col0, 0.0f);
                    }
                    row = tile_m + local_row_block1 + i;
                    if (row < m) {
                        store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col0, 0.0f);
                    }
                }
            }
            if (col1 < n) {
                #pragma unroll
                for (int i = 0; i < 4; ++i) {
                    int64_t row = tile_m + local_row_block0 + i;
                    if (row < m) {
                        store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col1, 0.0f);
                    }
                    row = tile_m + local_row_block1 + i;
                    if (row < m) {
                        store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col1, 0.0f);
                    }
                }
            }
        }
        return;
    }

    const int64_t k = k2 * 2;
    const int64_t full_kb_end =
        kb_start + ((kb_end - kb_start) / NATIVE_MFMA_K_BLOCKS) * NATIVE_MFMA_K_BLOCKS;

    for (int64_t kb_group = kb_start; kb_group < full_kb_end; kb_group += NATIVE_MFMA_K_BLOCKS) {
        #pragma unroll
        for (int stage = 0; stage < NATIVE_MFMA_K_BLOCKS; ++stage) {
            load_gemm_stage_to_smem<SHUFFLED_B, SHUFFLE_INNER_MODE>(
                a_bf16,
                b_u8,
                b_scale_u8,
                smem_A[stage],
                smem_B[stage],
                smem_A_scale[stage],
                smem_B_scale[stage],
                tid,
                tile_m,
                tile_n,
                m,
                n,
                k,
                k2,
                kb_group + stage,
                k_blocks_pad);
        }
        __syncthreads();

        #if defined(__gfx950__) && MXFP4_ENABLE_NATIVE_FP4_MFMA
        accumulate_native_group(
            smem_A,
            smem_B,
            smem_A_scale,
            smem_B_scale,
            row_group,
            local_a_row0,
            local_a_row1,
            local_b_col0,
            local_b_col1,
            c_acc_00,
            c_acc_01,
            c_acc_10,
            c_acc_11);
        #else
        accumulate_tail_stages(
            smem_A,
            smem_B,
            smem_A_scale,
            smem_B_scale,
            NATIVE_MFMA_K_BLOCKS,
            row_group,
            local_a_row0,
            local_a_row1,
            local_b_col0,
            local_b_col1,
            local_row_block0,
            local_row_block1,
            c_acc_00,
            c_acc_01,
            c_acc_10,
            c_acc_11);
        #endif
        __syncthreads();
    }

    const int tail_k_blocks = static_cast<int>(kb_end - full_kb_end);
    if (tail_k_blocks > 0) {
        #pragma unroll
        for (int stage = 0; stage < NATIVE_MFMA_K_BLOCKS; ++stage) {
            if (stage < tail_k_blocks) {
                load_gemm_stage_to_smem<SHUFFLED_B, SHUFFLE_INNER_MODE>(
                    a_bf16,
                    b_u8,
                    b_scale_u8,
                    smem_A[stage],
                    smem_B[stage],
                    smem_A_scale[stage],
                    smem_B_scale[stage],
                    tid,
                    tile_m,
                    tile_n,
                    m,
                    n,
                    k,
                    k2,
                    full_kb_end + stage,
                    k_blocks_pad);
            }
        }
        __syncthreads();

        accumulate_tail_stages(
            smem_A,
            smem_B,
            smem_A_scale,
            smem_B_scale,
            tail_k_blocks,
            row_group,
            local_a_row0,
            local_a_row1,
            local_b_col0,
            local_b_col1,
            local_row_block0,
            local_row_block1,
            c_acc_00,
            c_acc_01,
            c_acc_10,
            c_acc_11);
    }

    if (col0 < n) {
        #pragma unroll
        for (int i = 0; i < 4; ++i) {
            int64_t row = tile_m + local_row_block0 + i;
            if (row < m) {
                store_output_value<WRITE_BF16_OUT>(
                    workspace, out, workspace_stride, n, split_idx, row, col0, c_acc_00[i]);
            }
            row = tile_m + local_row_block1 + i;
            if (row < m) {
                store_output_value<WRITE_BF16_OUT>(
                    workspace, out, workspace_stride, n, split_idx, row, col0, c_acc_10[i]);
            }
        }
    }
    if (col1 < n) {
        #pragma unroll
        for (int i = 0; i < 4; ++i) {
            int64_t row = tile_m + local_row_block0 + i;
            if (row < m) {
                store_output_value<WRITE_BF16_OUT>(
                    workspace, out, workspace_stride, n, split_idx, row, col1, c_acc_01[i]);
            }
            row = tile_m + local_row_block1 + i;
            if (row < m) {
                store_output_value<WRITE_BF16_OUT>(
                    workspace, out, workspace_stride, n, split_idx, row, col1, c_acc_11[i]);
            }
        }
    }
}

__global__ void reduce_splitk_kernel(
    const float* workspace,
    __hip_bfloat16* out,
    int64_t workspace_stride,
    int64_t m,
    int64_t n,
    int64_t split_k) {
    int64_t row = static_cast<int64_t>(blockIdx.y) * blockDim.y + threadIdx.y;
    int64_t col = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
    if (row >= m || col >= n) {
        return;
    }

    float sum = 0.0f;
    int64_t base = row * n + col;
    for (int64_t s = 0; s < split_k; ++s) {
        sum += workspace[s * workspace_stride + base];
    }
    out[base] = static_cast<__hip_bfloat16>(sum);
}

inline void check_hip_error(const char* where) {
    hipError_t err = hipGetLastError();
    TORCH_CHECK(err == hipSuccess, where, " failed: ", hipGetErrorString(err));
}

template <bool SHUFFLED_B, int SHUFFLE_INNER_MODE, bool WRITE_BF16_OUT>
inline void launch_gemm_mxfp4_variant(
    dim3 grid,
    dim3 block,
    const __hip_bfloat16* a_bf16,
    const uint8_t* b_u8,
    const uint8_t* b_scale_u8,
    float* workspace,
    __hip_bfloat16* out,
    int64_t workspace_stride,
    int64_t m,
    int64_t n,
    int64_t k2,
    int64_t k_blocks_valid,
    int64_t k_blocks_pad,
    int64_t split_k) {
    hipLaunchKernelGGL(
        HIP_KERNEL_NAME(gemm_mxfp4_kernel<SHUFFLED_B, SHUFFLE_INNER_MODE, WRITE_BF16_OUT>),
        grid,
        block,
        0,
        0,
        a_bf16,
        b_u8,
        b_scale_u8,
        workspace,
        out,
        workspace_stride,
        m,
        n,
        k2,
        k_blocks_valid,
        k_blocks_pad,
        split_k);
}

template <bool SHUFFLED_B, int SHUFFLE_INNER_MODE, bool WRITE_BF16_OUT>
inline void launch_gemm_mxfp4_prequant_variant(
    dim3 grid,
    dim3 block,
    const uint8_t* a_fp4,
    const uint8_t* b_u8,
    const uint8_t* a_scale_u8,
    const uint8_t* b_scale_u8,
    float* workspace,
    __hip_bfloat16* out,
    int64_t workspace_stride,
    int64_t m,
    int64_t n,
    int64_t k2,
    int64_t k_blocks_valid,
    int64_t k_blocks_pad,
    int64_t split_k) {
    hipLaunchKernelGGL(
        HIP_KERNEL_NAME(gemm_mxfp4_prequant_kernel<SHUFFLED_B, SHUFFLE_INNER_MODE, WRITE_BF16_OUT>),
        grid,
        block,
        0,
        0,
        a_fp4,
        b_u8,
        a_scale_u8,
        b_scale_u8,
        workspace,
        out,
        workspace_stride,
        m,
        n,
        k2,
        k_blocks_valid,
        k_blocks_pad,
        split_k);
}

template <bool SHUFFLED_B, int SHUFFLE_INNER_MODE, bool WRITE_BF16_OUT>
inline void launch_gemm_mxfp4_prequant_native_variant(
    dim3 grid,
    dim3 block,
    const uint8_t* a_fp4,
    const uint8_t* b_u8,
    const uint8_t* a_scale_u8,
    const uint8_t* b_scale_u8,
    float* workspace,
    __hip_bfloat16* out,
    int64_t workspace_stride,
    int64_t m,
    int64_t n,
    int64_t k2,
    int64_t k_blocks_valid,
    int64_t k_blocks_pad,
    int64_t split_k) {
    hipLaunchKernelGGL(
        HIP_KERNEL_NAME(gemm_mxfp4_prequant_native_kernel<SHUFFLED_B, SHUFFLE_INNER_MODE, WRITE_BF16_OUT>),
        grid,
        block,
        0,
        0,
        a_fp4,
        b_u8,
        a_scale_u8,
        b_scale_u8,
        workspace,
        out,
        workspace_stride,
        m,
        n,
        k2,
        k_blocks_valid,
        k_blocks_pad,
        split_k);
}

template <bool SHUFFLED_B, int SHUFFLE_INNER_MODE, bool WRITE_BF16_OUT>
inline void launch_gemm_mxfp4_native_variant(
    dim3 grid,
    dim3 block,
    const __hip_bfloat16* a_bf16,
    const uint8_t* b_u8,
    const uint8_t* b_scale_u8,
    float* workspace,
    __hip_bfloat16* out,
    int64_t workspace_stride,
    int64_t m,
    int64_t n,
    int64_t k2,
    int64_t k_blocks_valid,
    int64_t k_blocks_pad,
    int64_t split_k) {
    hipLaunchKernelGGL(
        HIP_KERNEL_NAME(gemm_mxfp4_native_kernel<SHUFFLED_B, SHUFFLE_INNER_MODE, WRITE_BF16_OUT>),
        grid,
        block,
        0,
        0,
        a_bf16,
        b_u8,
        b_scale_u8,
        workspace,
        out,
        workspace_stride,
        m,
        n,
        k2,
        k_blocks_valid,
        k_blocks_pad,
        split_k);
}

}  // namespace

std::vector<torch::Tensor> hip_quant_mxfp4(torch::Tensor x) {
    TORCH_CHECK(x.is_cuda(), "x must be CUDA tensor");
    TORCH_CHECK(x.scalar_type() == torch::kBFloat16, "x must be bfloat16");
    TORCH_CHECK(x.dim() == 2, "x must be 2D [M, K]");

    auto x_contig = x.contiguous();
    int64_t m = x_contig.size(0);
    int64_t k = x_contig.size(1);
    TORCH_CHECK(k % 64 == 0, "k must be divisible by 64");

    int64_t k_blocks_valid = k / BLOCK_K;
    int64_t k_blocks_pad = ((k_blocks_valid + (PAD_SCALE_N - 1)) / PAD_SCALE_N) * PAD_SCALE_N;
    int64_t m_pad_kernel = ((m + (PAD_M_KERNEL - 1)) / PAD_M_KERNEL) * PAD_M_KERNEL;
    int64_t m_pad_scale = ((m + (PAD_M_SCALE - 1)) / PAD_M_SCALE) * PAD_M_SCALE;

    auto u8_opts = x_contig.options().dtype(torch::kUInt8);
    auto out_fp4 = torch::empty({m_pad_kernel, k / 2}, u8_opts);
    auto out_scale = torch::full({m_pad_scale, k_blocks_pad}, 127, u8_opts);

    int64_t total = m_pad_kernel * k_blocks_pad;
    int threads = 256;
    int blocks = static_cast<int>((total + threads - 1) / threads);
    if (blocks > 0) {
        hipLaunchKernelGGL(
            quant_mxfp4_kernel,
            dim3(blocks),
            dim3(threads),
            0,
            0,
            reinterpret_cast<const __hip_bfloat16*>(x_contig.data_ptr()),
            reinterpret_cast<uint8_t*>(out_fp4.data_ptr()),
            reinterpret_cast<uint8_t*>(out_scale.data_ptr()),
            m,
            m_pad_kernel,
            k,
            k_blocks_valid,
            k_blocks_pad);
        check_hip_error("quant_mxfp4_kernel");
    }
    return {out_fp4, out_scale};
}

torch::Tensor hip_gemm_mxfp4_prequant(
    torch::Tensor a_fp4_u8,
    torch::Tensor b_u8,
    torch::Tensor a_scale_u8,
    torch::Tensor b_scale_u8,
    int64_t layout_mode,
    int64_t log2_k_split,
    torch::Tensor workspace,
    int64_t workspace_stride,
    int64_t b_shuffle_inner_mode,
    int64_t native_mfma_mode) {
    TORCH_CHECK(a_fp4_u8.is_cuda(), "a_fp4_u8 must be CUDA");
    TORCH_CHECK(b_u8.is_cuda(), "b_u8 must be CUDA");
    TORCH_CHECK(a_scale_u8.is_cuda(), "a_scale_u8 must be CUDA");
    TORCH_CHECK(b_scale_u8.is_cuda(), "b_scale_u8 must be CUDA");
    TORCH_CHECK(a_fp4_u8.scalar_type() == torch::kUInt8, "a_fp4_u8 must be uint8");
    TORCH_CHECK(b_u8.scalar_type() == torch::kUInt8, "b_u8 must be uint8");
    TORCH_CHECK(a_scale_u8.scalar_type() == torch::kUInt8, "a_scale_u8 must be uint8");
    TORCH_CHECK(b_scale_u8.scalar_type() == torch::kUInt8, "b_scale_u8 must be uint8");
    TORCH_CHECK(a_fp4_u8.dim() == 2 && b_u8.dim() == 2, "A/B FP4 must be 2D");
    TORCH_CHECK(a_scale_u8.dim() == 2 && b_scale_u8.dim() == 2, "A/B scale must be 2D");
    TORCH_CHECK(layout_mode == 0 || layout_mode == 1, "layout_mode must be 0 or 1");
    TORCH_CHECK(layout_mode == 0 || b_shuffle_inner_mode == 0 || b_shuffle_inner_mode == 1,
                "b_shuffle_inner_mode must be 0 or 1 for shuffled B");
    TORCH_CHECK(native_mfma_mode == 0 || native_mfma_mode == 1,
                "native_mfma_mode must be 0 or 1");

    auto a_fp4 = a_fp4_u8.contiguous();
    auto b_fp4 = b_u8.contiguous();
    auto a_scale = a_scale_u8.contiguous();
    auto b_scale = b_scale_u8.contiguous();

    int64_t m = a_fp4.size(0);
    int64_t n = b_fp4.size(0);
    int64_t k2 = a_fp4.size(1);
    TORCH_CHECK(b_fp4.size(1) == k2, "A/B K/2 mismatch");

    int64_t k = k2 * 2;
    TORCH_CHECK(k % BLOCK_K == 0, "K must be divisible by 32");
    int64_t k_blocks_valid = k / BLOCK_K;
    int64_t k_blocks_pad = a_scale.size(1);
    TORCH_CHECK(b_scale.size(1) == k_blocks_pad, "A/B scale padded K-block mismatch");
    TORCH_CHECK(a_scale.size(0) >= m, "a_scale rows must cover m");
    TORCH_CHECK(b_scale.size(0) >= n, "b_scale rows must cover n");

    int64_t split_k = 1;
    if (log2_k_split > 0) {
        split_k = static_cast<int64_t>(1) << log2_k_split;
    }
    if (split_k < 1) {
        split_k = 1;
    }
    const bool direct_output = (split_k == 1);

    int64_t mn = m * n;
    auto ws = workspace;
    if (!direct_output) {
        if (workspace_stride <= 0) {
            workspace_stride = mn;
        }
        TORCH_CHECK(workspace_stride >= mn, "workspace_stride must be >= m*n");

        auto ws_opts = a_fp4.options().dtype(torch::kFloat);
        int64_t need = split_k * workspace_stride;
        if (!ws.defined() || !ws.is_cuda() || ws.scalar_type() != torch::kFloat || ws.numel() < need) {
            ws = torch::empty({split_k, workspace_stride}, ws_opts);
        } else {
            ws = ws.contiguous().view({split_k, workspace_stride});
        }
    }

    auto out = torch::empty({m, n}, a_fp4.options().dtype(torch::kBFloat16));

    dim3 block(GEMM_THREADS);
    dim3 grid(
        static_cast<unsigned int>((n + GEMM_BLOCK_N - 1) / GEMM_BLOCK_N),
        static_cast<unsigned int>((m + GEMM_BLOCK_M - 1) / GEMM_BLOCK_M),
        static_cast<unsigned int>(split_k));
    const uint8_t* a_ptr = reinterpret_cast<const uint8_t*>(a_fp4.data_ptr());
    const uint8_t* b_ptr = reinterpret_cast<const uint8_t*>(b_fp4.data_ptr());
    const uint8_t* a_scale_ptr = reinterpret_cast<const uint8_t*>(a_scale.data_ptr());
    const uint8_t* b_scale_ptr = reinterpret_cast<const uint8_t*>(b_scale.data_ptr());
    float* ws_ptr = direct_output ? nullptr : reinterpret_cast<float*>(ws.data_ptr());
    auto* out_ptr = reinterpret_cast<__hip_bfloat16*>(out.data_ptr());
    constexpr bool native_build_enabled = static_cast<bool>(MXFP4_ENABLE_NATIVE_FP4_MFMA);
    const bool use_native_mfma =
        native_build_enabled && (native_mfma_mode != 0) && (k_blocks_valid >= NATIVE_MFMA_K_BLOCKS);

    if (direct_output) {
        if (layout_mode == 0) {
            if (use_native_mfma) {
                launch_gemm_mxfp4_prequant_native_variant<false, 0, true>(
                    grid, block, a_ptr, b_ptr, a_scale_ptr, b_scale_ptr, ws_ptr, out_ptr,
                    workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
            } else {
                launch_gemm_mxfp4_prequant_variant<false, 0, true>(
                    grid, block, a_ptr, b_ptr, a_scale_ptr, b_scale_ptr, ws_ptr, out_ptr,
                    workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
            }
        } else if (b_shuffle_inner_mode == 0) {
            if (use_native_mfma) {
                launch_gemm_mxfp4_prequant_native_variant<true, 0, true>(
                    grid, block, a_ptr, b_ptr, a_scale_ptr, b_scale_ptr, ws_ptr, out_ptr,
                    workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
            } else {
                launch_gemm_mxfp4_prequant_variant<true, 0, true>(
                    grid, block, a_ptr, b_ptr, a_scale_ptr, b_scale_ptr, ws_ptr, out_ptr,
                    workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
            }
        } else {
            if (use_native_mfma) {
                launch_gemm_mxfp4_prequant_native_variant<true, 1, true>(
                    grid, block, a_ptr, b_ptr, a_scale_ptr, b_scale_ptr, ws_ptr, out_ptr,
                    workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
            } else {
                launch_gemm_mxfp4_prequant_variant<true, 1, true>(
                    grid, block, a_ptr, b_ptr, a_scale_ptr, b_scale_ptr, ws_ptr, out_ptr,
                    workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
            }
        }
    } else {
        if (layout_mode == 0) {
            if (use_native_mfma) {
                launch_gemm_mxfp4_prequant_native_variant<false, 0, false>(
                    grid, block, a_ptr, b_ptr, a_scale_ptr, b_scale_ptr, ws_ptr, out_ptr,
                    workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
            } else {
                launch_gemm_mxfp4_prequant_variant<false, 0, false>(
                    grid, block, a_ptr, b_ptr, a_scale_ptr, b_scale_ptr, ws_ptr, out_ptr,
                    workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
            }
        } else if (b_shuffle_inner_mode == 0) {
            if (use_native_mfma) {
                launch_gemm_mxfp4_prequant_native_variant<true, 0, false>(
                    grid, block, a_ptr, b_ptr, a_scale_ptr, b_scale_ptr, ws_ptr, out_ptr,
                    workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
            } else {
                launch_gemm_mxfp4_prequant_variant<true, 0, false>(
                    grid, block, a_ptr, b_ptr, a_scale_ptr, b_scale_ptr, ws_ptr, out_ptr,
                    workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
            }
        } else {
            if (use_native_mfma) {
                launch_gemm_mxfp4_prequant_native_variant<true, 1, false>(
                    grid, block, a_ptr, b_ptr, a_scale_ptr, b_scale_ptr, ws_ptr, out_ptr,
                    workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
            } else {
                launch_gemm_mxfp4_prequant_variant<true, 1, false>(
                    grid, block, a_ptr, b_ptr, a_scale_ptr, b_scale_ptr, ws_ptr, out_ptr,
                    workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
            }
        }
    }
    check_hip_error(use_native_mfma ? "gemm_mxfp4_prequant_native_kernel"
                                    : "gemm_mxfp4_prequant_kernel");

    if (!direct_output) {
        dim3 rblock(16, 16);
        dim3 rgrid(
            static_cast<unsigned int>((n + 15) / 16),
            static_cast<unsigned int>((m + 15) / 16));
        hipLaunchKernelGGL(
            reduce_splitk_kernel,
            rgrid,
            rblock,
            0,
            0,
            reinterpret_cast<const float*>(ws.data_ptr()),
            reinterpret_cast<__hip_bfloat16*>(out.data_ptr()),
            workspace_stride,
            m,
            n,
            split_k);
        check_hip_error("reduce_splitk_kernel");
    }
    return out;
}

torch::Tensor hip_gemm_mxfp4(
    torch::Tensor a_bf16,
    torch::Tensor b_u8,
    torch::Tensor b_scale_u8,
    int64_t layout_mode,
    int64_t log2_k_split,
    torch::Tensor workspace,
    int64_t workspace_stride,
    int64_t b_shuffle_inner_mode,
    int64_t native_mfma_mode) {
    TORCH_CHECK(a_bf16.is_cuda(), "a_bf16 must be CUDA");
    TORCH_CHECK(b_u8.is_cuda(), "b_u8 must be CUDA");
    TORCH_CHECK(b_scale_u8.is_cuda(), "b_scale_u8 must be CUDA");
    TORCH_CHECK(a_bf16.scalar_type() == torch::kBFloat16, "a_bf16 must be bfloat16");
    TORCH_CHECK(b_u8.scalar_type() == torch::kUInt8, "b_u8 must be uint8");
    TORCH_CHECK(b_scale_u8.scalar_type() == torch::kUInt8, "b_scale_u8 must be uint8");
    TORCH_CHECK(a_bf16.dim() == 2 && b_u8.dim() == 2, "A BF16 and B FP4 must be 2D");
    TORCH_CHECK(b_scale_u8.dim() == 2, "B scale must be 2D");
    TORCH_CHECK(layout_mode == 0 || layout_mode == 1, "layout_mode must be 0 or 1");
    TORCH_CHECK(layout_mode == 0 || b_shuffle_inner_mode == 0 || b_shuffle_inner_mode == 1,
                "b_shuffle_inner_mode must be 0 or 1 for shuffled B");

    auto a = a_bf16.contiguous();
    auto b_fp4 = b_u8.contiguous();
    auto b_scale = b_scale_u8.contiguous();

    int64_t m = a.size(0);
    int64_t n = b_fp4.size(0);
    int64_t k = a.size(1);
    int64_t k2 = b_fp4.size(1);
    TORCH_CHECK(k == k2 * 2, "A/B K mismatch");

    TORCH_CHECK(k % BLOCK_K == 0, "K must be divisible by 32");
    int64_t k_blocks_valid = k / BLOCK_K;
    int64_t k_blocks_pad = b_scale.size(1);
    TORCH_CHECK(k_blocks_pad >= k_blocks_valid, "B scale padded K-block mismatch");
    TORCH_CHECK(b_scale.size(0) >= n, "b_scale rows must cover n");
    TORCH_CHECK(native_mfma_mode == 0 || native_mfma_mode == 1,
                "native_mfma_mode must be 0 or 1");

    int64_t split_k = 1;
    if (log2_k_split > 0) {
        split_k = static_cast<int64_t>(1) << log2_k_split;
    }
    if (split_k < 1) {
        split_k = 1;
    }
    const bool direct_output = (split_k == 1);

    int64_t mn = m * n;
    auto ws = workspace;
    if (!direct_output) {
        if (workspace_stride <= 0) {
            workspace_stride = mn;
        }
        TORCH_CHECK(workspace_stride >= mn, "workspace_stride must be >= m*n");

        auto ws_opts = a.options().dtype(torch::kFloat);
        int64_t need = split_k * workspace_stride;
        if (!ws.defined() || !ws.is_cuda() || ws.scalar_type() != torch::kFloat || ws.numel() < need) {
            ws = torch::empty({split_k, workspace_stride}, ws_opts);
        } else {
            ws = ws.contiguous().view({split_k, workspace_stride});
        }
    }

    auto out = torch::empty({m, n}, a.options().dtype(torch::kBFloat16));

    dim3 block(GEMM_THREADS);
    dim3 grid(
        static_cast<unsigned int>((n + GEMM_BLOCK_N - 1) / GEMM_BLOCK_N),
        static_cast<unsigned int>((m + GEMM_BLOCK_M - 1) / GEMM_BLOCK_M),
        static_cast<unsigned int>(split_k));
    const __hip_bfloat16* a_ptr = reinterpret_cast<const __hip_bfloat16*>(a.data_ptr());
    const uint8_t* b_ptr = reinterpret_cast<const uint8_t*>(b_fp4.data_ptr());
    const uint8_t* b_scale_ptr = reinterpret_cast<const uint8_t*>(b_scale.data_ptr());
    float* ws_ptr = direct_output ? nullptr : reinterpret_cast<float*>(ws.data_ptr());
    auto* out_ptr = reinterpret_cast<__hip_bfloat16*>(out.data_ptr());
    constexpr bool native_build_enabled = static_cast<bool>(MXFP4_ENABLE_NATIVE_FP4_MFMA);
    const bool use_native_mfma =
        native_build_enabled && (native_mfma_mode != 0) && (k_blocks_valid >= NATIVE_MFMA_K_BLOCKS);

    if (direct_output) {
        if (layout_mode == 0) {
            if (use_native_mfma) {
                launch_gemm_mxfp4_native_variant<false, 0, true>(
                    grid, block, a_ptr, b_ptr, b_scale_ptr, ws_ptr, out_ptr,
                    workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
            } else {
                launch_gemm_mxfp4_variant<false, 0, true>(
                    grid, block, a_ptr, b_ptr, b_scale_ptr, ws_ptr, out_ptr,
                    workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
            }
        } else if (b_shuffle_inner_mode == 0) {
            if (use_native_mfma) {
                launch_gemm_mxfp4_native_variant<true, 0, true>(
                    grid, block, a_ptr, b_ptr, b_scale_ptr, ws_ptr, out_ptr,
                    workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
            } else {
                launch_gemm_mxfp4_variant<true, 0, true>(
                    grid, block, a_ptr, b_ptr, b_scale_ptr, ws_ptr, out_ptr,
                    workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
            }
        } else {
            if (use_native_mfma) {
                launch_gemm_mxfp4_native_variant<true, 1, true>(
                    grid, block, a_ptr, b_ptr, b_scale_ptr, ws_ptr, out_ptr,
                    workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
            } else {
                launch_gemm_mxfp4_variant<true, 1, true>(
                    grid, block, a_ptr, b_ptr, b_scale_ptr, ws_ptr, out_ptr,
                    workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
            }
        }
    } else {
        if (layout_mode == 0) {
            if (use_native_mfma) {
                launch_gemm_mxfp4_native_variant<false, 0, false>(
                    grid, block, a_ptr, b_ptr, b_scale_ptr, ws_ptr, out_ptr,
                    workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
            } else {
                launch_gemm_mxfp4_variant<false, 0, false>(
                    grid, block, a_ptr, b_ptr, b_scale_ptr, ws_ptr, out_ptr,
                    workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
            }
        } else if (b_shuffle_inner_mode == 0) {
            if (use_native_mfma) {
                launch_gemm_mxfp4_native_variant<true, 0, false>(
                    grid, block, a_ptr, b_ptr, b_scale_ptr, ws_ptr, out_ptr,
                    workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
            } else {
                launch_gemm_mxfp4_variant<true, 0, false>(
                    grid, block, a_ptr, b_ptr, b_scale_ptr, ws_ptr, out_ptr,
                    workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
            }
        } else {
            if (use_native_mfma) {
                launch_gemm_mxfp4_native_variant<true, 1, false>(
                    grid, block, a_ptr, b_ptr, b_scale_ptr, ws_ptr, out_ptr,
                    workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
            } else {
                launch_gemm_mxfp4_variant<true, 1, false>(
                    grid, block, a_ptr, b_ptr, b_scale_ptr, ws_ptr, out_ptr,
                    workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
            }
        }
    }
    check_hip_error(use_native_mfma ? "gemm_mxfp4_native_kernel" : "gemm_mxfp4_kernel");

    if (!direct_output) {
        dim3 rblock(16, 16);
        dim3 rgrid(
            static_cast<unsigned int>((n + 15) / 16),
            static_cast<unsigned int>((m + 15) / 16));
        hipLaunchKernelGGL(
            reduce_splitk_kernel,
            rgrid,
            rblock,
            0,
            0,
            reinterpret_cast<const float*>(ws.data_ptr()),
            reinterpret_cast<__hip_bfloat16*>(out.data_ptr()),
            workspace_stride,
            m,
            n,
            split_k);
        check_hip_error("reduce_splitk_kernel");
    }
    return out;
}
"""


def _sanitize_b_layout(mode: str) -> str:
    mode = (mode or _DEFAULT_B_LAYOUT).strip().lower()
    if mode in {"raw", "shuffle", "auto"}:
        return mode
    return _DEFAULT_B_LAYOUT


def _get_b_layout() -> str:
    mode = _sanitize_b_layout(os.getenv(_B_LAYOUT_ENV, _DEFAULT_B_LAYOUT))
    if mode == "auto":
        return "shuffle"
    return mode


def _get_fused_a_quant() -> bool:
    raw = os.getenv(_FUSED_A_QUANT_ENV, _DEFAULT_FUSED_A_QUANT)
    return str(raw).strip().lower() in {"1", "true", "yes", "on"}


def _sanitize_native_mfma(mode: str) -> str:
    mode = (mode or _DEFAULT_NATIVE_MFMA).strip().lower()
    if mode in {"never", "auto", "force"}:
        return mode
    return _DEFAULT_NATIVE_MFMA


def _get_native_mfma_mode() -> str:
    return _sanitize_native_mfma(os.getenv(_NATIVE_MFMA_ENV, _DEFAULT_NATIVE_MFMA))


def _sanitize_b_shuffle_inner_mode(mode: str | None) -> int:
    if mode is None:
        return _DEFAULT_B_SHUFFLE_INNER_MODE
    try:
        value = int(str(mode).strip())
    except ValueError:
        return _DEFAULT_B_SHUFFLE_INNER_MODE
    return value if value in {0, 1} else _DEFAULT_B_SHUFFLE_INNER_MODE


def _get_b_shuffle_inner_mode() -> int:
    return _sanitize_b_shuffle_inner_mode(os.getenv(_B_SHUFFLE_INNER_ENV))


def _get_splitk_override() -> int | None:
    raw = os.getenv(_SPLITK_ENV)
    if raw is None:
        return None
    try:
        return max(0, int(raw))
    except ValueError:
        return None


def _e8m0_unshuffle(scale_sh: torch.Tensor) -> torch.Tensor:
    if scale_sh.ndim != 2:
        raise RuntimeError(f"scale_sh must be 2D, got {tuple(scale_sh.shape)}")
    sm, sn = scale_sh.shape
    if sm % 32 != 0 or sn % 8 != 0:
        raise RuntimeError(f"scale_sh shape must be divisible by (32,8), got {(sm, sn)}")

    s = scale_sh.view(torch.uint8)
    s = s.view(sm // 32, sn // 8, 4, 16, 2, 2)
    s = s.permute(0, 5, 3, 1, 4, 2).contiguous()
    s = s.view(sm, sn)
    return s.view(scale_sh.dtype)


def _get_b_scale_raw_cached(b_scale_sh: torch.Tensor) -> torch.Tensor:
    dev = int(b_scale_sh.device.index) if b_scale_sh.device.index is not None else -1
    key = (
        int(b_scale_sh.data_ptr()),
        int(b_scale_sh.shape[0]),
        int(b_scale_sh.shape[1]),
        dev,
    )
    cached = _B_SCALE_RAW_CACHE.get(key)
    if cached is not None:
        return cached

    with _B_SCALE_RAW_LOCK:
        cached = _B_SCALE_RAW_CACHE.get(key)
        if cached is not None:
            return cached
        raw = _e8m0_unshuffle(b_scale_sh).contiguous()
        if len(_B_SCALE_RAW_CACHE) >= _B_SCALE_RAW_CACHE_MAX:
            _B_SCALE_RAW_CACHE.pop(next(iter(_B_SCALE_RAW_CACHE)))
        _B_SCALE_RAW_CACHE[key] = raw
        return raw


def _get_workspace(device: torch.device, m: int, n: int, split_k: int) -> torch.Tensor:
    dev = int(device.index) if device.index is not None else -1
    key = (dev, int(m), int(n), int(split_k))
    cached = _WORKSPACE_CACHE.get(key)
    need = split_k * m * n
    if cached is not None and cached.numel() >= need:
        return cached

    with _WORKSPACE_LOCK:
        cached = _WORKSPACE_CACHE.get(key)
        if cached is not None and cached.numel() >= need:
            return cached
        ws = torch.empty((split_k, m * n), dtype=torch.float32, device=device)
        if len(_WORKSPACE_CACHE) >= _WORKSPACE_CACHE_MAX:
            _WORKSPACE_CACHE.pop(next(iter(_WORKSPACE_CACHE)))
        _WORKSPACE_CACHE[key] = ws
        return ws


def _pick_splitk_log2(m: int, n: int, k: int) -> int:
    override = _get_splitk_override()
    if override is not None:
        return override

    key = (int(m), int(n), int(k))
    if key in _STATIC_SPLITK_LOG2:
        return _STATIC_SPLITK_LOG2[key]

    if m <= 16 and k >= 4096:
        return 3
    if m <= 32:
        return 2
    if m <= 64 and k >= 2048:
        return 1
    return 0


def _resolve_launch_policy(m: int, n: int, k: int, native_mfma_mode: str) -> _LaunchPolicy:
    shape_key = (int(m), int(n), int(k))
    is_ranked_shape = shape_key in _RANKED_SHAPE_KEYS

    use_native_prequant = 0
    use_native_fused = 0
    if native_mfma_mode == "force":
        use_native_prequant = 1
        use_native_fused = 1
    elif native_mfma_mode == "auto" and is_ranked_shape:
        # Ranked benchmark shapes now have a correctness-validated native
        # prequant path on gfx950. Keep fused-A native guarded behind `force`
        # until that variant gets its own correctness sign-off.
        use_native_prequant = 1

    return _LaunchPolicy(
        log2_k_split=_pick_splitk_log2(*shape_key),
        use_native_prequant=use_native_prequant,
        use_native_fused=use_native_fused,
        is_ranked_shape=is_ranked_shape,
    )


def _get_hip_module():
    global _HIP_MODULE, _HIP_BUILD_ERROR, _HIP_NATIVE_BUILD_ENABLED, _HIP_BUILD_REPORT_EMITTED
    if _HIP_MODULE is not None:
        return _HIP_MODULE
    if _HIP_BUILD_ERROR is not None:
        raise RuntimeError(f"HIP inline build failed previously: {_HIP_BUILD_ERROR}")

    with _HIP_LOCK:
        if _HIP_MODULE is not None:
            return _HIP_MODULE
        if _HIP_BUILD_ERROR is not None:
            raise RuntimeError(f"HIP inline build failed previously: {_HIP_BUILD_ERROR}")

        native_mode = _get_native_mfma_mode()
        build_attempts = [0]
        if native_mode != "never":
            build_attempts = [1, 0]

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

        last_error = None
        for enable_native_fp4_mfma in build_attempts:
            try:
                _HIP_MODULE = load_inline(
                    name=f"mxfp4_mm_inline_quant_gemm_v13_nat{enable_native_fp4_mfma}",
                    cpp_sources=[CPP_WRAPPER],
                    cuda_sources=[HIP_SRC],
                    functions=["hip_quant_mxfp4", "hip_gemm_mxfp4_prequant", "hip_gemm_mxfp4"],
                    verbose=False,
                    extra_cuda_cflags=[
                        "-O3",
                        "-std=c++20",
                        "--offload-arch=gfx950",
                        f"-DMXFP4_ENABLE_NATIVE_FP4_MFMA={enable_native_fp4_mfma}",
                    ],
                )
                _HIP_NATIVE_BUILD_ENABLED = bool(enable_native_fp4_mfma)
                _HIP_BUILD_ERROR = None
                if not _HIP_BUILD_REPORT_EMITTED:
                    print(
                        f"[mxfp4_mm] hip_inline_build requested={native_mode} "
                        f"built=nat{1 if _HIP_NATIVE_BUILD_ENABLED else 0}",
                        flush=True,
                    )
                    _HIP_BUILD_REPORT_EMITTED = True
                break
            except Exception as e:  # pragma: no cover - runtime dependent
                last_error = e
                _HIP_MODULE = None
                if enable_native_fp4_mfma == 1 and native_mode != "force":
                    if not _HIP_BUILD_REPORT_EMITTED:
                        print(
                            "[mxfp4_mm] hip_inline_build requested="
                            f"{native_mode} native_build_failed -> retry nat0",
                            flush=True,
                        )
                        _HIP_BUILD_REPORT_EMITTED = True
                    continue
                _HIP_BUILD_ERROR = e
                raise RuntimeError(f"HIP inline build failed: {e}") from e

        if _HIP_MODULE is None:
            _HIP_BUILD_ERROR = last_error
            raise RuntimeError(f"HIP inline build failed: {last_error}")

    return _HIP_MODULE


def _quant_triton_mxfp4(x: torch.Tensor, shuffle: bool = True):
    from aiter import dtypes
    from aiter.ops.triton.quant import dynamic_mxfp4_quant
    from aiter.utility.fp4_utils import e8m0_shuffle

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


def _detect_b_shuffle_inner_mode(
    module,
    a_bf16: torch.Tensor,
    b_shuffle: torch.Tensor,
    b_scale_sh: torch.Tensor,
) -> int:
    del module, a_bf16, b_shuffle, b_scale_sh
    # Keep one fixed B_shuffle inner layout in the steady-state path. The
    # pre-shuffled MXFP4 tensor produced by aiter.ops.shuffle.shuffle_weight()
    # flattens 16x16 byte tiles in N-major inner order, which matches mode 1.
    return _get_b_shuffle_inner_mode()


def _run_hip_full_pipeline(
    a_bf16: torch.Tensor,
    b_q: torch.Tensor,
    b_shuffle: torch.Tensor,
    b_scale_sh: torch.Tensor,
) -> torch.Tensor:
    module = _get_hip_module()
    b_layout = _get_b_layout()
    native_mfma_mode = _get_native_mfma_mode()
    if not _HIP_NATIVE_BUILD_ENABLED and native_mfma_mode != "never":
        native_mfma_mode = "never"

    m = int(a_bf16.shape[0])
    k = int(a_bf16.shape[1])
    n = int(b_q.shape[0])
    launch_policy = _resolve_launch_policy(m, n, k, native_mfma_mode)
    shape_key = (m, n, k)
    with _RUNTIME_REPORT_LOCK:
        if shape_key not in _RUNTIME_REPORTED_SHAPES:
            print(
                "[mxfp4_mm] launch "
                f"shape={shape_key} requested={_get_native_mfma_mode()} "
                f"effective={native_mfma_mode} build=nat{1 if _HIP_NATIVE_BUILD_ENABLED else 0} "
                f"prequant_native={launch_policy.use_native_prequant} "
                f"fused_native={launch_policy.use_native_fused} "
                f"splitk_log2={launch_policy.log2_k_split}",
                flush=True,
            )
            _RUNTIME_REPORTED_SHAPES.add(shape_key)

    if b_layout == "shuffle":
        b_u8 = b_shuffle.view(torch.uint8).contiguous()
        b_scale_u8 = b_scale_sh.view(torch.uint8).contiguous()
        layout_mode = 1
        b_inner_mode = _detect_b_shuffle_inner_mode(module, a_bf16, b_shuffle, b_scale_sh)
    else:
        b_u8 = b_q.view(torch.uint8).contiguous()
        b_scale_u8 = _get_b_scale_raw_cached(b_scale_sh).view(torch.uint8).contiguous()
        layout_mode = 0
        b_inner_mode = 0

    log2_k_split = launch_policy.log2_k_split
    split_k = 1 << log2_k_split
    if split_k > 1:
        ws = _get_workspace(a_bf16.device, m, n, split_k)
        workspace_stride = int(m * n)
    else:
        ws = torch.empty((0,), dtype=torch.float32, device=a_bf16.device)
        workspace_stride = 0

    if _get_fused_a_quant():
        return module.hip_gemm_mxfp4(
            a_bf16,
            b_u8,
            b_scale_u8,
            int(layout_mode),
            int(log2_k_split),
            ws,
            int(workspace_stride),
            int(b_inner_mode),
            int(launch_policy.use_native_fused),
        )

    a_q, a_scale_raw = _quant_triton_mxfp4(a_bf16, shuffle=False)
    a_fp4_u8 = a_q.view(torch.uint8).contiguous()
    a_scale_u8 = a_scale_raw.view(torch.uint8).contiguous()
    return module.hip_gemm_mxfp4_prequant(
        a_fp4_u8,
        b_u8,
        a_scale_u8,
        b_scale_u8,
        int(layout_mode),
        int(log2_k_split),
        ws,
        int(workspace_stride),
        int(b_inner_mode),
        int(launch_policy.use_native_prequant),
    )


def custom_kernel(data: input_t) -> output_t:
    """
    Default HIP path:
      bf16 A -> Triton per-1x32 quant -> HIP split-k GEMM -> bf16 C.
      Default validation mode now forces the prequant native FP4 scale-MFMA
      path on gfx950-capable builds, so nat1 compile/runtime issues are not
      hidden by the nat0 fallback.

    Experimental fused HIP path:
      set MXFP4_MM_FUSED_A_QUANT=1 to quantize A inside the GEMM prolog.
    Native path controls:
      MXFP4_MM_NATIVE_MFMA=auto uses native only for the ranked benchmark shapes.
      MXFP4_MM_NATIVE_MFMA=force tries native on every eligible shape, including fused A experiments.
      MXFP4_MM_NATIVE_MFMA=never keeps the BF16-unpack fallback.
      MXFP4_MM_B_SHUFFLE_INNER_MODE lets debug runs override the fixed shuffled-B inner layout.
    """
    A, B, B_q, B_shuffle, B_scale_sh = data
    del B

    A = A.contiguous()
    B_q = B_q.contiguous()
    B_shuffle = B_shuffle.contiguous()
    B_scale_sh = B_scale_sh.contiguous()

    return _run_hip_full_pipeline(
        a_bf16=A,
        b_q=B_q,
        b_shuffle=B_shuffle,
        b_scale_sh=B_scale_sh,
    )
scrolls · 2695 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