Skip to content
KernelIndex
Search⌘K

submission 713996

hq_struggling · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

hq_submission_load_inline.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-713996?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
21.3µs
#735 of 1143
2026-04-03

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:5b04d2aec5785110db2200e01e644d49d8c7cb4e46df4fb4290881723a377b2f
license declaredunknown
license concludedunknown
authorshq_struggling
imported2026-08-26

Techniques

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

fp4const int k_blocks = K >> 6; // K / 64,也就是每个 row tile 里的 64-fp4 block 数
shared-memory__shared__ uint4 s_aq[2][TM][SPLITS];
split-k__global__ __launch_bounds__(512) void mxfp4_mfma_16x128x128_splitk_kernel(
vector-width = uint4static __device__ __forceinline__ uint4 quantize_32_bf16_to_fp4(
warp-specializationconst bool a_producer = tid < TM * SPLITS;

Kernel source

hq_submission_load_inline.py1966 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

from task import input_t, output_t

import os
os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
os.environ.setdefault("CXX", "clang++")
os.environ.setdefault("MAX_JOBS", "8")
import torch
from torch.utils.cpp_extension import load_inline
import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import (
    gemm_afp4wfp4_preshuffled_weight_scales,
)
from aiter.utility.fp4_utils import e8m0_shuffle

_THIS_DIR = os.path.dirname(os.path.abspath(__file__))
AITER_INCLUDE_DIR = None
USE_NATIVE_MFMA = os.environ.get("MXFP4_USE_NATIVE_MFMA") == "1"
USE_TRITON_PRESHUFFLE = os.environ.get("MXFP4_USE_TRITON_PRESHUFFLE", "1") == "1"
_module = None


def _resolve_aiter_include_dirs() -> list[str]:
    global AITER_INCLUDE_DIR
    if AITER_INCLUDE_DIR is None:
        for _root in (
            os.path.abspath(os.path.join(_THIS_DIR, "..")),
            os.path.abspath(os.path.join(_THIS_DIR, "..", "..")),
            "/home/runner",
            "/home/ubuntu/data/amd_202602",
        ):
            _include_dir = os.path.join(_root, "aiter", "csrc", "include")
            if os.path.isdir(_include_dir):
                AITER_INCLUDE_DIR = _include_dir
                break
        if AITER_INCLUDE_DIR is None:
            raise FileNotFoundError("Unable to locate aiter/csrc/include for load_inline")
    return [AITER_INCLUDE_DIR, os.path.join(AITER_INCLUDE_DIR, "opus")]


# C++ 侧只暴露一个最小入口,真正的实现放在下面的 HIP 源码里。
CPP_SRC = r"""
extern "C" torch::Tensor mxfp4_gemm_kernel(
    torch::Tensor A,
    torch::Tensor B_shuffle,
    torch::Tensor B_scale_sh
);
extern "C" torch::Tensor mxfp4_gemm_prequant(
    torch::Tensor A_q,
    torch::Tensor A_scale,
    torch::Tensor B_shuffle,
    torch::Tensor B_scale
);
extern "C" torch::Tensor mxfp4_pack_a_debug(
    torch::Tensor A
);
"""


# HIP 源码使用 gfx950 的 scaled MFMA:
# - B 数据继续按 16B 粒度搬运,避免逐 nibble 标量解包
# - A 在 kernel prologue 中按 1x32 动态量化成 MXFP4
# - B scale 直接保持 preshuffled E8M0 布局,在 device 侧按索引读取
# - quant_kernels.cu 里的 scale shuffle 只在“落内存”时需要;这里 A scale
#   直接留在寄存器喂给 MFMA,所以保留同构的索引定义,但不物化成 shuffled tensor
HIP_SRC = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include "opus.hpp"
#include <cstdint>
#include <limits>

using namespace opus;

// 基础参数检查,避免把错误类型或非连续张量送进 kernel。
#define CHECK_GPU(x) TORCH_CHECK(x.is_cuda(), #x " must be a GPU tensor")
#define CHECK_CONTIGUOUS(x) TORCH_CHECK(x.is_contiguous(), #x " must be contiguous")
#define CHECK_UINT8(x) TORCH_CHECK(x.scalar_type() == torch::kUInt8, #x " must be uint8")

// B_scale_sh 的物理布局。
static __device__ __forceinline__ int fp4_scale_shuffle_id(int scaleN_pad, int x, int y) {
    return (x / 32 * scaleN_pad) * 32 + (y / 8) * 256 + (y % 4) * 64 + (x % 16) * 4 +
           (y % 8) / 4 * 2 + (x % 32) / 16;
}

static __device__ __forceinline__ unsigned char load_preshuffled_scale(
    const uint8_t* base,
    int scaleNPad,
    int row,
    int group)
{
    return base[fp4_scale_shuffle_id(scaleNPad, row, group)];
}

static __device__ __forceinline__ unsigned char mxfp4_scale_byte(float amax) {
    // Match torch_dynamic_mxfp4_quant()/fp4_utils.dynamic_mxfp4_quant():
    // round amax via (bits + 0x200000) & 0xFF800000, then encode amax * 0.25
    // as E8M0.
    uint32_t bits = __builtin_bit_cast(uint32_t, amax);
    uint32_t rounded = (bits + 0x200000u) & 0xFF800000u;
    uint32_t exponent = (rounded >> 23) & 0xFFu;
    if (exponent == 0xFFu) {
        return static_cast<unsigned char>(0xFFu);
    }
    return static_cast<unsigned char>(exponent > 2u ? exponent - 2u : 0u);
}

static __device__ __forceinline__ float mxfp4_scale_f32(unsigned char scale_byte) {
    if (scale_byte == 0) {
        return __builtin_bit_cast(float, 0x00400000u);
    }
    if (scale_byte == 0xFFu) {
        return __builtin_bit_cast(float, 0x7F800001u);
    }
    return __builtin_bit_cast(float, static_cast<uint32_t>(scale_byte) << 23);
}

static __device__ __forceinline__ unsigned char float_to_fp4_nibble(float x_scaled) {
    constexpr uint32_t SIGN_MASK = 0x80000000u;
    constexpr float FP4_MAX_NORMAL = 6.0f;
    constexpr float FP4_MIN_NORMAL = 1.0f;
    constexpr int32_t EXP_BIAS_FP32 = 127;
    constexpr int32_t EXP_BIAS_FP4 = 1;
    constexpr int32_t MBITS_F32 = 23;
    constexpr int32_t MBITS_FP4 = 1;
    constexpr uint32_t FP4_SIGN_BIT = 0x8u;
    constexpr uint32_t DENORM_MASK_INT =
        ((EXP_BIAS_FP32 - EXP_BIAS_FP4) + (MBITS_F32 - MBITS_FP4) + 1) << MBITS_F32;
    constexpr float DENORM_MASK_FLOAT = __builtin_bit_cast(float, DENORM_MASK_INT);
    constexpr int32_t NORMAL_BIAS =
        ((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_F32) + (1 << 21) - 1;

    uint32_t bits = __builtin_bit_cast(uint32_t, x_scaled);
    uint32_t sign = bits & SIGN_MASK;
    bits ^= sign;
    float x_abs = __builtin_bit_cast(float, bits);

    uint8_t fp4_value = 0;
    if (x_abs >= FP4_MAX_NORMAL) {
        fp4_value = 0x7u;
    } else if (x_abs < FP4_MIN_NORMAL) {
        float denorm = x_abs + DENORM_MASK_FLOAT;
        uint32_t denorm_bits = __builtin_bit_cast(uint32_t, denorm) - DENORM_MASK_INT;
        fp4_value = static_cast<uint8_t>(denorm_bits);
    } else {
        uint32_t normal_bits = bits;
        uint32_t mant_odd = (normal_bits >> (MBITS_F32 - MBITS_FP4)) & 1u;
        normal_bits += static_cast<uint32_t>(NORMAL_BIAS);
        normal_bits += mant_odd;
        normal_bits >>= (MBITS_F32 - MBITS_FP4);
        fp4_value = static_cast<uint8_t>(normal_bits);
    }

    uint8_t sign_lp = static_cast<uint8_t>((sign >> 28) & FP4_SIGN_BIT);
    return static_cast<unsigned char>(fp4_value | sign_lp);
}

static __device__ __forceinline__ unsigned char bf16_to_fp4_packed_byte(
    const bf16x2_t& src,
    float scale_f32)
{
    const float quant_scale = 1.0f / scale_f32;
    const unsigned char lo = float_to_fp4_nibble(static_cast<float>(src[0]) * quant_scale);
    const unsigned char hi = float_to_fp4_nibble(static_cast<float>(src[1]) * quant_scale);
    return static_cast<unsigned char>((lo & 0xFu) | ((hi & 0xFu) << 4));
}

static __device__ __forceinline__ int preshuffled_mxfp4_block_offset(
    int n,
    int k0,
    int K,
    int split)
{
    // B_shuffle 的物理布局
    // -------------------------
    // 逻辑张量:
    //   B_shuffle[n, k_byte],形状为 [N, K/2]
    //
    // 打包事实:
    //   - 1 个字节存 2 个 fp4 值
    //   - 16 个字节 = 32 个 fp4 值
    //   - 64 个 fp4 值 = 32 个字节
    //
    // tile 顺序:
    //   [n_block = n / 16][k_block = k / 64][half = 0/1][row_in_16][byte_in_16]
    //
    // 一个物理 subtile 是 16 行 x 16 字节 = 256 字节:
    //
    //   第 0 行  -> 16 个 packed 字节
    //   第 1 行  -> 16 个 packed 字节
    //   ...
    //   第15 行  -> 16 个 packed 字节
    //
    // `split` 用来选择当前 K step 里读取哪一个 32-fp4 chunk。
    // 下面的 `k0` 以 fp4 元素为单位,所以这里要换算成字节偏移。
    const int n_block = n >> 4;  // 每个 tile 覆盖 16 行
    const int n_in = n & 15;
    const int k_blocks = K >> 6;  // K / 64,也就是每个 row tile 里的 64-fp4 block 数
    const int kb = (k0 >> 6) + (split >> 1);  // 当前 K block,按 split 分组修正
    const int c = split & 1;  // 选择 64-fp4 block 里的左/右半个 tile
    const int block = (n_block * k_blocks + kb) * 2 + c;
    return (block * 16 + n_in) * 16;  // 每个 half tile 是 16 行 x 16 字节
}

static __device__ __forceinline__ const uint8_t* preshuffled_mxfp4_block_ptr(
    const uint8_t* base,
    int n,
    int k0,
    int K,
    int split)
{
    return base + preshuffled_mxfp4_block_offset(n, k0, K, split);
}

static __device__ __forceinline__ const uint8_t* row_major_mxfp4_block_ptr(
    const uint8_t* base,
    int row,
    int k0,
    int K,
    int split)
{
    return base + row * (K >> 1) + (k0 >> 1) + split * 16;
}

static __device__ __forceinline__ void store_bf16x4_scatter(
    uint16_t* __restrict__ dst,
    int base_row,
    int col,
    int stride_n,
    int row_limit,
    const fp32x4_t& acc)
{
    #pragma unroll
    for (int i = 0; i < 4; ++i) {
        const int row = base_row + i;
        if (row < row_limit) {
            dst[row * stride_n + col] = fp32_to_bf16_rtn_raw(acc[i]);
        }
    }
}

static __device__ __forceinline__ void store_fp32x4_scatter(
    float* __restrict__ dst,
    int base_row,
    int col,
    int stride_n,
    int row_limit,
    const fp32x4_t& acc)
{
    #pragma unroll
    for (int i = 0; i < 4; ++i) {
        const int row = base_row + i;
        if (row < row_limit) {
            dst[row * stride_n + col] = acc[i];
        }
    }
}

static __device__ __forceinline__ fp32x4_t mxfp4_mma_16x16x128(
    const i32x8_t& a,
    const i32x8_t& b,
    const fp32x4_t& c,
    int block_sel,
    int scale_a,
    int scale_b)
{
    switch (block_sel) {
        case 0:
            return __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
                a, b, c, 4, 4, 0, scale_a, 0, scale_b);
        case 1:
            return __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
                a, b, c, 4, 4, 1, scale_a, 1, scale_b);
        case 2:
            return __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
                a, b, c, 4, 4, 2, scale_a, 2, scale_b);
        default:
            return __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
                a, b, c, 4, 4, 3, scale_a, 3, scale_b);
    }
}

static __device__ __forceinline__ fp32x16_t mxfp4_mma_32x32x64(
    const i32x8_t& a,
    const i32x8_t& b,
    const fp32x16_t& c,
    int block_sel,
    int scale_a,
    int scale_b)
{
    if (block_sel == 0) {
        return __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
            a, b, c, 4, 4, 0, scale_a, 0, scale_b);
    }
    return __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
        a, b, c, 4, 4, 1, scale_a, 1, scale_b);
}

static __device__ __forceinline__ uint4 quantize_32_bf16_to_fp4(
    const bf16_t* src,
    unsigned char* scale_out)
{
    bf16x8_t v0 = *reinterpret_cast<const bf16x8_t*>(src + 0);
    bf16x8_t v1 = *reinterpret_cast<const bf16x8_t*>(src + 8);
    bf16x8_t v2 = *reinterpret_cast<const bf16x8_t*>(src + 16);
    bf16x8_t v3 = *reinterpret_cast<const bf16x8_t*>(src + 24);

    float amax = 1.0e-10f;
    #pragma unroll
    for (int i = 0; i < 8; ++i) {
        float x = static_cast<float>(v0[i]);
        x = x < 0.0f ? -x : x;
        amax = amax > x ? amax : x;
    }
    #pragma unroll
    for (int i = 0; i < 8; ++i) {
        float x = static_cast<float>(v1[i]);
        x = x < 0.0f ? -x : x;
        amax = amax > x ? amax : x;
    }
    #pragma unroll
    for (int i = 0; i < 8; ++i) {
        float x = static_cast<float>(v2[i]);
        x = x < 0.0f ? -x : x;
        amax = amax > x ? amax : x;
    }
    #pragma unroll
    for (int i = 0; i < 8; ++i) {
        float x = static_cast<float>(v3[i]);
        x = x < 0.0f ? -x : x;
        amax = amax > x ? amax : x;
    }

    const unsigned char scale_byte = mxfp4_scale_byte(amax);
    const float scale_f32 = mxfp4_scale_f32(scale_byte);

    union {
        uint4 v;
        unsigned char b[16];
    } out{};

    out.b[0] = bf16_to_fp4_packed_byte(bf16x2_t{v0[0], v0[1]}, scale_f32);
    out.b[1] = bf16_to_fp4_packed_byte(bf16x2_t{v0[2], v0[3]}, scale_f32);
    out.b[2] = bf16_to_fp4_packed_byte(bf16x2_t{v0[4], v0[5]}, scale_f32);
    out.b[3] = bf16_to_fp4_packed_byte(bf16x2_t{v0[6], v0[7]}, scale_f32);
    out.b[4] = bf16_to_fp4_packed_byte(bf16x2_t{v1[0], v1[1]}, scale_f32);
    out.b[5] = bf16_to_fp4_packed_byte(bf16x2_t{v1[2], v1[3]}, scale_f32);
    out.b[6] = bf16_to_fp4_packed_byte(bf16x2_t{v1[4], v1[5]}, scale_f32);
    out.b[7] = bf16_to_fp4_packed_byte(bf16x2_t{v1[6], v1[7]}, scale_f32);
    out.b[8] = bf16_to_fp4_packed_byte(bf16x2_t{v2[0], v2[1]}, scale_f32);
    out.b[9] = bf16_to_fp4_packed_byte(bf16x2_t{v2[2], v2[3]}, scale_f32);
    out.b[10] = bf16_to_fp4_packed_byte(bf16x2_t{v2[4], v2[5]}, scale_f32);
    out.b[11] = bf16_to_fp4_packed_byte(bf16x2_t{v2[6], v2[7]}, scale_f32);
    out.b[12] = bf16_to_fp4_packed_byte(bf16x2_t{v3[0], v3[1]}, scale_f32);
    out.b[13] = bf16_to_fp4_packed_byte(bf16x2_t{v3[2], v3[3]}, scale_f32);
    out.b[14] = bf16_to_fp4_packed_byte(bf16x2_t{v3[4], v3[5]}, scale_f32);
    out.b[15] = bf16_to_fp4_packed_byte(bf16x2_t{v3[6], v3[7]}, scale_f32);

    *scale_out = scale_byte;
    return out.v;
}

__global__ void mxfp4_pack_a_debug_kernel(
    const bf16_t* __restrict__ A,
    uint8_t* __restrict__ A_q,
    int M,
    int K)
{
    const int group_idx = static_cast<int>(blockIdx.x * blockDim.x + threadIdx.x);
    const int groups_per_row = K / 32;
    const int total_groups = M * groups_per_row;
    if (group_idx >= total_groups) {
        return;
    }

    const int row = group_idx / groups_per_row;
    const int group = group_idx % groups_per_row;
    unsigned char scale_byte = 127;
    union {
        uint4 v;
        unsigned char b[16];
    } packed{};
    packed.v = quantize_32_bf16_to_fp4(A + row * K + group * 32, &scale_byte);
    const int out_offset = row * (K / 2) + group * 16;
    #pragma unroll
    for (int i = 0; i < 16; ++i) {
        A_q[out_offset + i] = packed.b[i];
    }
}

template<int TM, int TN>
__global__ __launch_bounds__(64) void mxfp4_mfma_16x16x128_kernel(
    const bf16_t* __restrict__ A,
    const uint8_t* __restrict__ B_q,
    const uint8_t* __restrict__ B_scale_sh,
    uint16_t* __restrict__ C,
    int M,
    int N,
    int K,
    int scaleNPad)
{
    // 16x16x128 原生路径。
    //
    // wave -> tile 映射
    //   每个 wave 64 个线程
    //   lane   = 0..63
    //   lane16 = lane % 16
    //   group4 = lane / 16
    //
    //   group4 选择一个 128-fp4 K step 里的 4 个 32-fp4 切片:
    //     group4 0 -> k0 +  0..31
    //     group4 1 -> k0 + 32..63
    //     group4 2 -> k0 + 64..95
    //     group4 3 -> k0 + 96..127
    //
    // 数据流
    //   global B_shuffle -> 16B 向量加载 -> b_frag(VGPR)-> MFMA
    //   global A         -> BF16->MXFP4 prologue -> a_frag(VGPR)-> MFMA
    //   acc 一直保留在 FP32 寄存器里,直到最后写回
    //
    // 这个 kernel 里没有显式的 shared memory staging buffer。
    // 这里的 “tile” 只是逻辑访问模式,不是 LDS 缓冲区。
    //
    // 逻辑张量:
    // - A: [M, K] bf16,row-major,在 prologue 里量化成当前 32-fp4 chunk
    // - B_q: [N, K/2] packed fp4 字节,已经做过 bpreshuffle,适合 16x16 tile 读取
    // - B_scale: [N, K/32] uint8 E8M0 block scale
    // - C: [M, N] bf16 输出
    constexpr int KSTEP = 128;
    const int lane = static_cast<int>(__builtin_amdgcn_workitem_id_x());
    const int lane16 = lane & 15;
    const int group4 = lane >> 4;
    const int m0 = blockIdx.y * TM;
    const int n0 = blockIdx.x * TN;
    const int row = m0 + lane16;
    const int col = n0 + lane16;
    const bool a_active = row < M;
    const bool b_active = col < N;
    const int row_stride_a = K;
    // acc 是单个 lane 的 4 个输出所对应的寄存器态 FP32 累加器。
    fp32x4_t acc{0.0f, 0.0f, 0.0f, 0.0f};
    // 16x16 tile 内部的布局:
    // - lane16 选择 [0, 15] 范围内的列
    // - group4 选择输出 tile 中的 4 行带
    //   (group4 = 0..3,每一带分别覆盖 [0..3]、[4..7]、[8..11]、[12..15])
    for (int k0 = 0; k0 < K; k0 += KSTEP) {
        union {
            i32x8_t v;
            uint4 lo;
            unsigned char bytes[32];
        } a_frag{};
        union {
            i32x8_t v;
            uint4 lo;
            unsigned char bytes[32];
        } b_frag{};

        int scale_a = 127;
        int scale_b = 127;

        if (a_active) {
            const bf16_t* a_src = A + row * row_stride_a + k0 + group4 * 32;
            unsigned char scale_byte = 127;
            a_frag.lo = quantize_32_bf16_to_fp4(a_src, &scale_byte);
            scale_a = static_cast<int>(scale_byte);
        }
        if (b_active) {
            // B_q 在 bpreshuffle 之后仍然是逻辑上的 [N, K/2]。
            // helper 会把 (column, k0, group4) 映射到当前的 16 行 x 16 字节 subtile。
            const uint8_t* b_src = preshuffled_mxfp4_block_ptr(B_q, col, k0, K, group4);
            b_frag.lo = *reinterpret_cast<const uint4*>(b_src);
            scale_b = static_cast<int>(load_preshuffled_scale(
                B_scale_sh, scaleNPad, col, (k0 >> 5) + group4));
        }

        acc = mxfp4_mma_16x16x128(a_frag.v, b_frag.v, acc, group4, scale_a, scale_b);
    }

    const int out_row_base = m0 + group4 * 4;
    if (b_active) {
        // 每个 lane 针对一个列写回 4 个 bf16 输出:
        // 行是 out_row_base + [0..3],列是 n0 + lane16。
        store_bf16x4_scatter(C, out_row_base, col, N, M, acc);
    }
}

template<int TM, int TN>
__global__ __launch_bounds__(64) void mxfp4_mfma_32x32x64_kernel(
    const bf16_t* __restrict__ A,
    const uint8_t* __restrict__ B_q,
    const uint8_t* __restrict__ B_scale_sh,
    uint16_t* __restrict__ C,
    int M,
    int N,
    int K,
    int scaleNPad)
{
    // 32x32x64 原生路径。
    //
    // wave -> tile 映射
    //   lane32 = lane % 32
    //   group  = lane / 32   // 0 或 1
    //
    //   group 0 -> 这个 K step 的前 32-fp4 切片
    //   group 1 -> 这个 K step 的后 32-fp4 切片
    //
    // 和 16x16 kernel 一样,这里是直接 global load -> VGPR fragment -> MFMA。
    // 不会显式构造 shared memory tile。
    //
    // 逻辑张量:
    // - A: [M, K] bf16,row-major,在 prologue 里量化成当前 32-fp4 chunk
    // - B_q: [N, K/2] packed fp4 字节,已经 bpreshuffle
    // - B_scale: [N, K/32] uint8 E8M0 block scale
    // - C: [M, N] bf16 输出
    constexpr int KSTEP = 64;
    const int lane = static_cast<int>(__builtin_amdgcn_workitem_id_x());
    const int lane32 = lane & 31;
    const int group = lane >> 5;
    const int m0 = blockIdx.y * TM;
    const int n0 = blockIdx.x * TN;
    const int row = m0 + lane32;
    const int col = n0 + lane32;
    const bool a_active = row < M;
    const bool b_active = col < N;
    const int row_stride_a = K;
    // acc 打包了最后要散写回去的 4x4 FP32 输出分块。
    fp32x16_t acc{0.0f};
    for (int k0 = 0; k0 < K; k0 += KSTEP) {
        union {
            i32x8_t v;
            uint4 lo;
            unsigned char bytes[32];
        } a_frag{};
        union {
            i32x8_t v;
            uint4 lo;
            unsigned char bytes[32];
        } b_frag{};

        int scale_a = 127;
        int scale_b = 127;

        if (a_active) {
            const bf16_t* a_src = A + row * row_stride_a + k0 + group * 32;
            unsigned char scale_byte = 127;
            a_frag.lo = quantize_32_bf16_to_fp4(a_src, &scale_byte);
            scale_a = static_cast<int>(scale_byte);
        }
        if (b_active) {
            // B_q 已经做过 bpreshuffle,所以 helper 会落到当前的 16x16 字节 subtile。
            const uint8_t* b_src = preshuffled_mxfp4_block_ptr(B_q, col, k0, K, group);
            b_frag.lo = *reinterpret_cast<const uint4*>(b_src);
            scale_b = static_cast<int>(load_preshuffled_scale(
                B_scale_sh, scaleNPad, col, (k0 >> 5) + group));
        }

        acc = mxfp4_mma_32x32x64(a_frag.v, b_frag.v, acc, group, scale_a, scale_b);
    }

    if (b_active) {
        // 行布局和 32x32 MFMA kernel 保持一致。
        store_bf16x4_scatter(
            C, m0 + group * 4 + 0, col, N, M, fp32x4_t{acc[0], acc[1], acc[2], acc[3]});
        store_bf16x4_scatter(
            C, m0 + group * 4 + 8, col, N, M, fp32x4_t{acc[4], acc[5], acc[6], acc[7]});
        store_bf16x4_scatter(
            C, m0 + group * 4 + 16, col, N, M, fp32x4_t{acc[8], acc[9], acc[10], acc[11]});
        store_bf16x4_scatter(
            C, m0 + group * 4 + 24, col, N, M, fp32x4_t{acc[12], acc[13], acc[14], acc[15]});
    }
}

__global__ __launch_bounds__(256) void mxfp4_mfma_16x64x128_kernel(
    const bf16_t* __restrict__ A,
    const uint8_t* __restrict__ B_q,
    const uint8_t* __restrict__ B_scale_sh,
    uint16_t* __restrict__ C,
    int M,
    int N,
    int K,
    int scaleNPad)
{
    constexpr int TM = 16;
    constexpr int TN = 64;
    constexpr int KSTEP = 128;
    constexpr int SPLITS = KSTEP / 32;

    __shared__ uint4 s_aq[2][TM][SPLITS];
    __shared__ unsigned char s_ascale[2][TM][SPLITS];
    __shared__ uint4 s_bq[2][TN][SPLITS];
    __shared__ unsigned char s_bscale[2][TN][SPLITS];

    const int tid = static_cast<int>(threadIdx.x);
    const int wave = tid >> 6;
    const int lane = tid & 63;
    const int lane16 = lane & 15;
    const int group4 = lane >> 4;
    const int m0 = blockIdx.y * TM;
    const int n0 = blockIdx.x * TN;
    const int row = m0 + lane16;
    const int col = n0 + wave * 16 + lane16;
    const bool a_active = row < M;
    const bool b_active = col < N;
    const int row_stride_a = K;
    const int a_q_group = tid / TM;
    const int a_q_row = tid - a_q_group * TM;
    const bool a_producer = tid < TM * SPLITS;
    const int b_split = tid / TN;
    const int b_col_local = tid - b_split * TN;
    const int b_col = n0 + b_col_local;
    fp32x4_t acc{0.0f, 0.0f, 0.0f, 0.0f};

    int read_buf = 0;
    if (a_producer) {
        const int global_row = m0 + a_q_row;
        unsigned char scale_byte = 127;
        uint4 packed{};
        if (global_row < M) {
            packed = quantize_32_bf16_to_fp4(
                A + global_row * row_stride_a + a_q_group * 32,
                &scale_byte);
        }
        s_aq[read_buf][a_q_row][a_q_group] = packed;
        s_ascale[read_buf][a_q_row][a_q_group] = scale_byte;
    }
    uint4 b_packed{};
    unsigned char b_scale = 127;
    if (b_col < N) {
        const uint8_t* b_src = preshuffled_mxfp4_block_ptr(B_q, b_col, 0, K, b_split);
        b_packed = *reinterpret_cast<const uint4*>(b_src);
        b_scale = load_preshuffled_scale(B_scale_sh, scaleNPad, b_col, b_split);
    }
    s_bq[read_buf][b_col_local][b_split] = b_packed;
    s_bscale[read_buf][b_col_local][b_split] = b_scale;
    __syncthreads();

    for (int k0 = 0; k0 < K; k0 += KSTEP) {
        union {
            i32x8_t v;
            uint4 lo;
            unsigned char bytes[32];
        } a_frag{};
        union {
            i32x8_t v;
            uint4 lo;
            unsigned char bytes[32];
        } b_frag{};

        int scale_a = 127;
        int scale_b = 127;

        if (a_active) {
            a_frag.lo = s_aq[read_buf][lane16][group4];
            scale_a = static_cast<int>(s_ascale[read_buf][lane16][group4]);
        }
        if (b_active) {
            b_frag.lo = s_bq[read_buf][col - n0][group4];
            scale_b = static_cast<int>(s_bscale[read_buf][col - n0][group4]);
        }

        const int next_k0 = k0 + KSTEP;
        uint4 a_next_packed{};
        unsigned char a_next_scale = 127;
        uint4 b_next_packed{};
        unsigned char b_next_scale = 127;
        if (next_k0 < K) {
            if (a_producer) {
                const int global_row = m0 + a_q_row;
                if (global_row < M) {
                    a_next_packed = quantize_32_bf16_to_fp4(
                        A + global_row * row_stride_a + next_k0 + a_q_group * 32,
                        &a_next_scale);
                }
            }
            if (b_col < N) {
                const uint8_t* b_src = preshuffled_mxfp4_block_ptr(B_q, b_col, next_k0, K, b_split);
                b_next_packed = *reinterpret_cast<const uint4*>(b_src);
                b_next_scale = load_preshuffled_scale(
                    B_scale_sh, scaleNPad, b_col, (next_k0 >> 5) + b_split);
            }
        }

        acc = mxfp4_mma_16x16x128(a_frag.v, b_frag.v, acc, group4, scale_a, scale_b);
        if (next_k0 < K) {
            const int write_buf = read_buf ^ 1;
            if (a_producer) {
                s_aq[write_buf][a_q_row][a_q_group] = a_next_packed;
                s_ascale[write_buf][a_q_row][a_q_group] = a_next_scale;
            }
            s_bq[write_buf][b_col_local][b_split] = b_next_packed;
            s_bscale[write_buf][b_col_local][b_split] = b_next_scale;
            __syncthreads();
            read_buf = write_buf;
        }
    }

    const int out_row_base = m0 + group4 * 4;
    if (b_active) {
        store_bf16x4_scatter(C, out_row_base, col, N, M, acc);
    }
}

__global__ __launch_bounds__(512) void mxfp4_mfma_16x128x128_kernel(
    const bf16_t* __restrict__ A,
    const uint8_t* __restrict__ B_q,
    const uint8_t* __restrict__ B_scale_sh,
    uint16_t* __restrict__ C,
    int M,
    int N,
    int K,
    int scaleNPad)
{
    constexpr int TM = 16;
    constexpr int TN = 128;
    constexpr int KSTEP = 128;
    constexpr int SPLITS = KSTEP / 32;

    __shared__ uint4 s_aq[2][TM][SPLITS];
    __shared__ unsigned char s_ascale[2][TM][SPLITS];
    __shared__ uint4 s_bq[2][TN][SPLITS];
    __shared__ unsigned char s_bscale[2][TN][SPLITS];

    const int tid = static_cast<int>(threadIdx.x);
    const int wave = tid >> 6;
    const int lane = tid & 63;
    const int lane16 = lane & 15;
    const int group4 = lane >> 4;
    const int m0 = blockIdx.y * TM;
    const int n0 = blockIdx.x * TN;
    const int row = m0 + lane16;
    const int col = n0 + wave * 16 + lane16;
    const bool a_active = row < M;
    const bool b_active = col < N;
    const int row_stride_a = K;
    const int a_q_group = tid / TM;
    const int a_q_row = tid - a_q_group * TM;
    const bool a_producer = tid < TM * SPLITS;
    const int b_split = tid / TN;
    const int b_col_local = tid - b_split * TN;
    const int b_col = n0 + b_col_local;
    fp32x4_t acc{0.0f, 0.0f, 0.0f, 0.0f};

    int read_buf = 0;
    if (a_producer) {
        const int global_row = m0 + a_q_row;
        unsigned char scale_byte = 127;
        uint4 packed{};
        if (global_row < M) {
            packed = quantize_32_bf16_to_fp4(
                A + global_row * row_stride_a + a_q_group * 32,
                &scale_byte);
        }
        s_aq[read_buf][a_q_row][a_q_group] = packed;
        s_ascale[read_buf][a_q_row][a_q_group] = scale_byte;
    }
    uint4 b_packed{};
    unsigned char b_scale = 127;
    if (b_col < N) {
        const uint8_t* b_src = preshuffled_mxfp4_block_ptr(B_q, b_col, 0, K, b_split);
        b_packed = *reinterpret_cast<const uint4*>(b_src);
        b_scale = load_preshuffled_scale(B_scale_sh, scaleNPad, b_col, b_split);
    }
    s_bq[read_buf][b_col_local][b_split] = b_packed;
    s_bscale[read_buf][b_col_local][b_split] = b_scale;
    __syncthreads();

    for (int k0 = 0; k0 < K; k0 += KSTEP) {
        union {
            i32x8_t v;
            uint4 lo;
            unsigned char bytes[32];
        } a_frag{};
        union {
            i32x8_t v;
            uint4 lo;
            unsigned char bytes[32];
        } b_frag{};

        int scale_a = 127;
        int scale_b = 127;

        if (a_active) {
            a_frag.lo = s_aq[read_buf][lane16][group4];
            scale_a = static_cast<int>(s_ascale[read_buf][lane16][group4]);
        }
        if (b_active) {
            b_frag.lo = s_bq[read_buf][col - n0][group4];
            scale_b = static_cast<int>(s_bscale[read_buf][col - n0][group4]);
        }

        const int next_k0 = k0 + KSTEP;
        uint4 a_next_packed{};
        unsigned char a_next_scale = 127;
        uint4 b_next_packed{};
        unsigned char b_next_scale = 127;
        if (next_k0 < K) {
            if (a_producer) {
                const int global_row = m0 + a_q_row;
                if (global_row < M) {
                    a_next_packed = quantize_32_bf16_to_fp4(
                        A + global_row * row_stride_a + next_k0 + a_q_group * 32,
                        &a_next_scale);
                }
            }
            if (b_col < N) {
                const uint8_t* b_src = preshuffled_mxfp4_block_ptr(B_q, b_col, next_k0, K, b_split);
                b_next_packed = *reinterpret_cast<const uint4*>(b_src);
                b_next_scale = load_preshuffled_scale(
                    B_scale_sh, scaleNPad, b_col, (next_k0 >> 5) + b_split);
            }
        }

        acc = mxfp4_mma_16x16x128(a_frag.v, b_frag.v, acc, group4, scale_a, scale_b);
        if (next_k0 < K) {
            const int write_buf = read_buf ^ 1;
            if (a_producer) {
                s_aq[write_buf][a_q_row][a_q_group] = a_next_packed;
                s_ascale[write_buf][a_q_row][a_q_group] = a_next_scale;
            }
            s_bq[write_buf][b_col_local][b_split] = b_next_packed;
            s_bscale[write_buf][b_col_local][b_split] = b_next_scale;
            __syncthreads();
            read_buf = write_buf;
        }
    }

    const int out_row_base = m0 + group4 * 4;
    if (b_active) {
        store_bf16x4_scatter(C, out_row_base, col, N, M, acc);
    }
}

__global__ __launch_bounds__(512) void mxfp4_mfma_16x128x128_splitk_kernel(
    const bf16_t* __restrict__ A,
    const uint8_t* __restrict__ B_q,
    const uint8_t* __restrict__ B_scale_sh,
    float* __restrict__ partial,
    int M,
    int N,
    int K,
    int scaleNPad,
    int split_k)
{
    constexpr int TM = 16;
    constexpr int TN = 128;
    constexpr int KSTEP = 128;
    constexpr int SPLITS = KSTEP / 32;

    __shared__ uint4 s_aq[2][TM][SPLITS];
    __shared__ unsigned char s_ascale[2][TM][SPLITS];
    __shared__ uint4 s_bq[2][TN][SPLITS];
    __shared__ unsigned char s_bscale[2][TN][SPLITS];

    const int tid = static_cast<int>(threadIdx.x);
    const int wave = tid >> 6;
    const int lane = tid & 63;
    const int lane16 = lane & 15;
    const int group4 = lane >> 4;
    const int split_idx = static_cast<int>(blockIdx.z);
    const int m0 = blockIdx.y * TM;
    const int n0 = blockIdx.x * TN;
    const int row = m0 + lane16;
    const int col = n0 + wave * 16 + lane16;
    const bool a_active = row < M;
    const bool b_active = col < N;
    const int row_stride_a = K;
    const int a_q_group = tid / TM;
    const int a_q_row = tid - a_q_group * TM;
    const bool a_producer = tid < TM * SPLITS;
    const int b_split = tid / TN;
    const int b_col_local = tid - b_split * TN;
    const int b_col = n0 + b_col_local;
    const int k_chunks = K / KSTEP;
    const int chunk_begin = (k_chunks * split_idx) / split_k;
    const int chunk_end = (k_chunks * (split_idx + 1)) / split_k;
    const int k_begin = chunk_begin * KSTEP;
    const int k_end = chunk_end * KSTEP;
    fp32x4_t acc{0.0f, 0.0f, 0.0f, 0.0f};
    if (k_begin >= k_end) {
        return;
    }

    int read_buf = 0;
    if (a_producer) {
        const int global_row = m0 + a_q_row;
        unsigned char scale_byte = 127;
        uint4 packed{};
        if (global_row < M) {
            packed = quantize_32_bf16_to_fp4(
                A + global_row * row_stride_a + k_begin + a_q_group * 32,
                &scale_byte);
        }
        s_aq[read_buf][a_q_row][a_q_group] = packed;
        s_ascale[read_buf][a_q_row][a_q_group] = scale_byte;
    }
    uint4 b_packed{};
    unsigned char b_scale = 127;
    if (b_col < N) {
        const uint8_t* b_src = preshuffled_mxfp4_block_ptr(B_q, b_col, k_begin, K, b_split);
        b_packed = *reinterpret_cast<const uint4*>(b_src);
        b_scale = load_preshuffled_scale(B_scale_sh, scaleNPad, b_col, (k_begin >> 5) + b_split);
    }
    s_bq[read_buf][b_col_local][b_split] = b_packed;
    s_bscale[read_buf][b_col_local][b_split] = b_scale;
    __syncthreads();

    for (int k0 = k_begin; k0 < k_end; k0 += KSTEP) {
        union {
            i32x8_t v;
            uint4 lo;
            unsigned char bytes[32];
        } a_frag{};
        union {
            i32x8_t v;
            uint4 lo;
            unsigned char bytes[32];
        } b_frag{};

        int scale_a = 127;
        int scale_b = 127;

        if (a_active) {
            a_frag.lo = s_aq[read_buf][lane16][group4];
            scale_a = static_cast<int>(s_ascale[read_buf][lane16][group4]);
        }
        if (b_active) {
            b_frag.lo = s_bq[read_buf][col - n0][group4];
            scale_b = static_cast<int>(s_bscale[read_buf][col - n0][group4]);
        }

        const int next_k0 = k0 + KSTEP;
        uint4 a_next_packed{};
        unsigned char a_next_scale = 127;
        uint4 b_next_packed{};
        unsigned char b_next_scale = 127;
        if (next_k0 < k_end) {
            if (a_producer) {
                const int global_row = m0 + a_q_row;
                if (global_row < M) {
                    a_next_packed = quantize_32_bf16_to_fp4(
                        A + global_row * row_stride_a + next_k0 + a_q_group * 32,
                        &a_next_scale);
                }
            }
            if (b_col < N) {
                const uint8_t* b_src = preshuffled_mxfp4_block_ptr(B_q, b_col, next_k0, K, b_split);
                b_next_packed = *reinterpret_cast<const uint4*>(b_src);
                b_next_scale = load_preshuffled_scale(
                    B_scale_sh, scaleNPad, b_col, (next_k0 >> 5) + b_split);
            }
        }

        acc = mxfp4_mma_16x16x128(a_frag.v, b_frag.v, acc, group4, scale_a, scale_b);
        if (next_k0 < k_end) {
            const int write_buf = read_buf ^ 1;
            if (a_producer) {
                s_aq[write_buf][a_q_row][a_q_group] = a_next_packed;
                s_ascale[write_buf][a_q_row][a_q_group] = a_next_scale;
            }
            s_bq[write_buf][b_col_local][b_split] = b_next_packed;
            s_bscale[write_buf][b_col_local][b_split] = b_next_scale;
            __syncthreads();
            read_buf = write_buf;
        }
    }

    const int out_row_base = m0 + group4 * 4;
    float* partial_base = partial + split_idx * M * N;
    if (b_active) {
        store_fp32x4_scatter(partial_base, out_row_base, col, N, M, acc);
    }
}

__global__ __launch_bounds__(256) void mxfp4_mfma_32x128x64_kernel(
    const bf16_t* __restrict__ A,
    const uint8_t* __restrict__ B_q,
    const uint8_t* __restrict__ B_scale_sh,
    uint16_t* __restrict__ C,
    int M,
    int N,
    int K,
    int scaleNPad)
{
    constexpr int TM = 32;
    constexpr int TN = 128;
    constexpr int KSTEP = 64;
    constexpr int SPLITS = KSTEP / 32;

    __shared__ uint4 s_aq[2][TM][SPLITS];
    __shared__ unsigned char s_ascale[2][TM][SPLITS];
    __shared__ uint4 s_bq[2][TN][SPLITS];
    __shared__ unsigned char s_bscale[2][TN][SPLITS];

    const int tid = static_cast<int>(threadIdx.x);
    const int wave = tid >> 6;
    const int lane = tid & 63;
    const int lane32 = lane & 31;
    const int group = lane >> 5;
    const int m0 = blockIdx.y * TM;
    const int n0 = blockIdx.x * TN;
    const int row = m0 + lane32;
    const int col = n0 + wave * 32 + lane32;
    const bool a_active = row < M;
    const bool b_active = col < N;
    const int row_stride_a = K;
    const int a_q_group = tid / TM;
    const int a_q_row = tid - a_q_group * TM;
    const bool a_producer = tid < TM * SPLITS;
    const int b_split = tid / TN;
    const int b_col_local = tid - b_split * TN;
    const int b_col = n0 + b_col_local;
    fp32x16_t acc{0.0f};

    int read_buf = 0;
    if (a_producer) {
        const int global_row = m0 + a_q_row;
        unsigned char scale_byte = 127;
        uint4 packed{};
        if (global_row < M) {
            packed = quantize_32_bf16_to_fp4(
                A + global_row * row_stride_a + a_q_group * 32,
                &scale_byte);
        }
        s_aq[read_buf][a_q_row][a_q_group] = packed;
        s_ascale[read_buf][a_q_row][a_q_group] = scale_byte;
    }
    uint4 b_packed{};
    unsigned char b_scale = 127;
    if (b_col < N) {
        const uint8_t* b_src = preshuffled_mxfp4_block_ptr(B_q, b_col, 0, K, b_split);
        b_packed = *reinterpret_cast<const uint4*>(b_src);
        b_scale = load_preshuffled_scale(B_scale_sh, scaleNPad, b_col, b_split);
    }
    s_bq[read_buf][b_col_local][b_split] = b_packed;
    s_bscale[read_buf][b_col_local][b_split] = b_scale;
    __syncthreads();

    for (int k0 = 0; k0 < K; k0 += KSTEP) {
        union {
            i32x8_t v;
            uint4 lo;
            unsigned char bytes[32];
        } a_frag{};
        union {
            i32x8_t v;
            uint4 lo;
            unsigned char bytes[32];
        } b_frag{};

        int scale_a = 127;
        int scale_b = 127;

        if (a_active) {
            a_frag.lo = s_aq[read_buf][lane32][group];
            scale_a = static_cast<int>(s_ascale[read_buf][lane32][group]);
        }
        if (b_active) {
            b_frag.lo = s_bq[read_buf][col - n0][group];
            scale_b = static_cast<int>(s_bscale[read_buf][col - n0][group]);
        }

        const int next_k0 = k0 + KSTEP;
        uint4 a_next_packed{};
        unsigned char a_next_scale = 127;
        uint4 b_next_packed{};
        unsigned char b_next_scale = 127;
        if (next_k0 < K) {
            if (a_producer) {
                const int global_row = m0 + a_q_row;
                if (global_row < M) {
                    a_next_packed = quantize_32_bf16_to_fp4(
                        A + global_row * row_stride_a + next_k0 + a_q_group * 32,
                        &a_next_scale);
                }
            }
            if (b_col < N) {
                const uint8_t* b_src = preshuffled_mxfp4_block_ptr(B_q, b_col, next_k0, K, b_split);
                b_next_packed = *reinterpret_cast<const uint4*>(b_src);
                b_next_scale = load_preshuffled_scale(
                    B_scale_sh, scaleNPad, b_col, (next_k0 >> 5) + b_split);
            }
        }

        acc = mxfp4_mma_32x32x64(a_frag.v, b_frag.v, acc, group, scale_a, scale_b);
        if (next_k0 < K) {
            const int write_buf = read_buf ^ 1;
            if (a_producer) {
                s_aq[write_buf][a_q_row][a_q_group] = a_next_packed;
                s_ascale[write_buf][a_q_row][a_q_group] = a_next_scale;
            }
            s_bq[write_buf][b_col_local][b_split] = b_next_packed;
            s_bscale[write_buf][b_col_local][b_split] = b_next_scale;
            __syncthreads();
            read_buf = write_buf;
        }
    }

    if (b_active) {
        store_bf16x4_scatter(
            C, m0 + group * 4 + 0, col, N, M, fp32x4_t{acc[0], acc[1], acc[2], acc[3]});
        store_bf16x4_scatter(
            C, m0 + group * 4 + 8, col, N, M, fp32x4_t{acc[4], acc[5], acc[6], acc[7]});
        store_bf16x4_scatter(
            C, m0 + group * 4 + 16, col, N, M, fp32x4_t{acc[8], acc[9], acc[10], acc[11]});
        store_bf16x4_scatter(
            C, m0 + group * 4 + 24, col, N, M, fp32x4_t{acc[12], acc[13], acc[14], acc[15]});
    }
}

__global__ __launch_bounds__(384) void mxfp4_mfma_32x192x64_kernel(
    const bf16_t* __restrict__ A,
    const uint8_t* __restrict__ B_q,
    const uint8_t* __restrict__ B_scale_sh,
    uint16_t* __restrict__ C,
    int M,
    int N,
    int K,
    int scaleNPad)
{
    constexpr int TM = 32;
    constexpr int TN = 192;
    constexpr int KSTEP = 64;
    constexpr int SPLITS = KSTEP / 32;

    __shared__ uint4 s_aq[2][TM][SPLITS];
    __shared__ unsigned char s_ascale[2][TM][SPLITS];
    __shared__ uint4 s_bq[2][TN][SPLITS];
    __shared__ unsigned char s_bscale[2][TN][SPLITS];

    const int tid = static_cast<int>(threadIdx.x);
    const int wave = tid >> 6;
    const int lane = tid & 63;
    const int lane32 = lane & 31;
    const int group = lane >> 5;
    const int m0 = blockIdx.y * TM;
    const int n0 = blockIdx.x * TN;
    const int row = m0 + lane32;
    const int col = n0 + wave * 32 + lane32;
    const bool a_active = row < M;
    const bool b_active = col < N;
    const int row_stride_a = K;
    const int a_q_group = tid / TM;
    const int a_q_row = tid - a_q_group * TM;
    const bool a_producer = tid < TM * SPLITS;
    const int b_split = tid / TN;
    const int b_col_local = tid - b_split * TN;
    const int b_col = n0 + b_col_local;
    fp32x16_t acc{0.0f};

    int read_buf = 0;
    if (a_producer) {
        const int global_row = m0 + a_q_row;
        unsigned char scale_byte = 127;
        uint4 packed{};
        if (global_row < M) {
            packed = quantize_32_bf16_to_fp4(
                A + global_row * row_stride_a + a_q_group * 32,
                &scale_byte);
        }
        s_aq[read_buf][a_q_row][a_q_group] = packed;
        s_ascale[read_buf][a_q_row][a_q_group] = scale_byte;
    }
    uint4 b_packed{};
    unsigned char b_scale = 127;
    if (b_col < N) {
        const uint8_t* b_src = preshuffled_mxfp4_block_ptr(B_q, b_col, 0, K, b_split);
        b_packed = *reinterpret_cast<const uint4*>(b_src);
        b_scale = load_preshuffled_scale(B_scale_sh, scaleNPad, b_col, b_split);
    }
    s_bq[read_buf][b_col_local][b_split] = b_packed;
    s_bscale[read_buf][b_col_local][b_split] = b_scale;
    __syncthreads();

    for (int k0 = 0; k0 < K; k0 += KSTEP) {
        union {
            i32x8_t v;
            uint4 lo;
            unsigned char bytes[32];
        } a_frag{};
        union {
            i32x8_t v;
            uint4 lo;
            unsigned char bytes[32];
        } b_frag{};

        int scale_a = 127;
        int scale_b = 127;

        if (a_active) {
            a_frag.lo = s_aq[read_buf][lane32][group];
            scale_a = static_cast<int>(s_ascale[read_buf][lane32][group]);
        }
        if (b_active) {
            b_frag.lo = s_bq[read_buf][col - n0][group];
            scale_b = static_cast<int>(s_bscale[read_buf][col - n0][group]);
        }

        const int next_k0 = k0 + KSTEP;
        uint4 a_next_packed{};
        unsigned char a_next_scale = 127;
        uint4 b_next_packed{};
        unsigned char b_next_scale = 127;
        if (next_k0 < K) {
            if (a_producer) {
                const int global_row = m0 + a_q_row;
                if (global_row < M) {
                    a_next_packed = quantize_32_bf16_to_fp4(
                        A + global_row * row_stride_a + next_k0 + a_q_group * 32,
                        &a_next_scale);
                }
            }
            if (b_col < N) {
                const uint8_t* b_src = preshuffled_mxfp4_block_ptr(B_q, b_col, next_k0, K, b_split);
                b_next_packed = *reinterpret_cast<const uint4*>(b_src);
                b_next_scale = load_preshuffled_scale(
                    B_scale_sh, scaleNPad, b_col, (next_k0 >> 5) + b_split);
            }
        }

        acc = mxfp4_mma_32x32x64(a_frag.v, b_frag.v, acc, group, scale_a, scale_b);
        if (next_k0 < K) {
            const int write_buf = read_buf ^ 1;
            if (a_producer) {
                s_aq[write_buf][a_q_row][a_q_group] = a_next_packed;
                s_ascale[write_buf][a_q_row][a_q_group] = a_next_scale;
            }
            s_bq[write_buf][b_col_local][b_split] = b_next_packed;
            s_bscale[write_buf][b_col_local][b_split] = b_next_scale;
            __syncthreads();
            read_buf = write_buf;
        }
    }

    if (b_active) {
        store_bf16x4_scatter(
            C, m0 + group * 4 + 0, col, N, M, fp32x4_t{acc[0], acc[1], acc[2], acc[3]});
        store_bf16x4_scatter(
            C, m0 + group * 4 + 8, col, N, M, fp32x4_t{acc[4], acc[5], acc[6], acc[7]});
        store_bf16x4_scatter(
            C, m0 + group * 4 + 16, col, N, M, fp32x4_t{acc[8], acc[9], acc[10], acc[11]});
        store_bf16x4_scatter(
            C, m0 + group * 4 + 24, col, N, M, fp32x4_t{acc[12], acc[13], acc[14], acc[15]});
    }
}

template<int TM, int TN>
__global__ __launch_bounds__(64) void mxfp4_mfma_16x16x128_prequant_kernel(
    const uint8_t* __restrict__ A_q,
    const uint8_t* __restrict__ A_scale,
    const uint8_t* __restrict__ B_q,
    const uint8_t* __restrict__ B_scale,
    uint16_t* __restrict__ C,
    int M,
    int N,
    int K)
{
    constexpr int KSTEP = 128;
    const int lane = static_cast<int>(__builtin_amdgcn_workitem_id_x());
    const int lane16 = lane & 15;
    const int group4 = lane >> 4;
    const int m0 = blockIdx.y * TM;
    const int n0 = blockIdx.x * TN;
    const int row = m0 + lane16;
    const int col = n0 + lane16;
    const bool a_active = row < M;
    const bool b_active = col < N;
    const int scale_stride = K / 32;
    const int scale_row_a = row * scale_stride;
    const int scale_row_b = col * scale_stride;

    fp32x4_t acc{0.0f, 0.0f, 0.0f, 0.0f};
    for (int k0 = 0; k0 < K; k0 += KSTEP) {
        union {
            i32x8_t v;
            uint4 lo;
            unsigned char bytes[32];
        } a_frag{};
        union {
            i32x8_t v;
            uint4 lo;
            unsigned char bytes[32];
        } b_frag{};

        int scale_a = 127;
        int scale_b = 127;

        if (a_active) {
            const uint8_t* a_src = row_major_mxfp4_block_ptr(A_q, row, k0, K, group4);
            a_frag.lo = *reinterpret_cast<const uint4*>(a_src);
            scale_a = static_cast<int>(A_scale[scale_row_a + (k0 >> 5) + group4]);
        }
        if (b_active) {
            const uint8_t* b_src = row_major_mxfp4_block_ptr(B_q, col, k0, K, group4);
            b_frag.lo = *reinterpret_cast<const uint4*>(b_src);
            scale_b = static_cast<int>(B_scale[scale_row_b + (k0 >> 5) + group4]);
        }

        acc = mxfp4_mma_16x16x128(a_frag.v, b_frag.v, acc, group4, scale_a, scale_b);
    }

    const int out_row_base = m0 + group4 * 4;
    if (b_active) {
        store_bf16x4_scatter(C, out_row_base, col, N, M, acc);
    }
}

template<int TM, int TN>
__global__ __launch_bounds__(64) void mxfp4_mfma_32x32x64_prequant_kernel(
    const uint8_t* __restrict__ A_q,
    const uint8_t* __restrict__ A_scale,
    const uint8_t* __restrict__ B_q,
    const uint8_t* __restrict__ B_scale,
    uint16_t* __restrict__ C,
    int M,
    int N,
    int K)
{
    constexpr int KSTEP = 64;
    const int lane = static_cast<int>(__builtin_amdgcn_workitem_id_x());
    const int lane32 = lane & 31;
    const int group = lane >> 5;
    const int m0 = blockIdx.y * TM;
    const int n0 = blockIdx.x * TN;
    const int row = m0 + lane32;
    const int col = n0 + lane32;
    const bool a_active = row < M;
    const bool b_active = col < N;
    const int scale_stride = K / 32;
    const int scale_row_a = row * scale_stride;
    const int scale_row_b = col * scale_stride;

    fp32x16_t acc{0.0f};
    for (int k0 = 0; k0 < K; k0 += KSTEP) {
        union {
            i32x8_t v;
            uint4 lo;
            unsigned char bytes[32];
        } a_frag{};
        union {
            i32x8_t v;
            uint4 lo;
            unsigned char bytes[32];
        } b_frag{};

        int scale_a = 127;
        int scale_b = 127;

        if (a_active) {
            const uint8_t* a_src = row_major_mxfp4_block_ptr(A_q, row, k0, K, group);
            a_frag.lo = *reinterpret_cast<const uint4*>(a_src);
            scale_a = static_cast<int>(A_scale[scale_row_a + (k0 >> 5) + group]);
        }
        if (b_active) {
            const uint8_t* b_src = row_major_mxfp4_block_ptr(B_q, col, k0, K, group);
            b_frag.lo = *reinterpret_cast<const uint4*>(b_src);
            scale_b = static_cast<int>(B_scale[scale_row_b + (k0 >> 5) + group]);
        }

        acc = mxfp4_mma_32x32x64(a_frag.v, b_frag.v, acc, group, scale_a, scale_b);
    }

    if (b_active) {
        store_bf16x4_scatter(
            C, m0 + group * 4 + 0, col, N, M, fp32x4_t{acc[0], acc[1], acc[2], acc[3]});
        store_bf16x4_scatter(
            C, m0 + group * 4 + 8, col, N, M, fp32x4_t{acc[4], acc[5], acc[6], acc[7]});
        store_bf16x4_scatter(
            C, m0 + group * 4 + 16, col, N, M, fp32x4_t{acc[8], acc[9], acc[10], acc[11]});
        store_bf16x4_scatter(
            C, m0 + group * 4 + 24, col, N, M, fp32x4_t{acc[12], acc[13], acc[14], acc[15]});
    }
}

__global__ void reduce_splitk_fp32_to_bf16_kernel(
    const float* __restrict__ partial,
    uint16_t* __restrict__ C,
    int split_k,
    int M,
    int N)
{
    const int idx = static_cast<int>(blockIdx.x * blockDim.x + threadIdx.x);
    const int total = M * N;
    if (idx >= total) {
        return;
    }
    float sum = 0.0f;
    #pragma unroll
    for (int s = 0; s < 4; ++s) {
        if (s < split_k) {
            sum += partial[s * total + idx];
        }
    }
    C[idx] = fp32_to_bf16_rtn_raw(sum);
}

enum class native_kernel_kind_t {
    k16x16x128,
    k16x64x128,
    k16x128x128,
    k16x128x128_splitk,
    k32x32x64,
    k32x128x64,
    k32x192x64,
};

struct native_kernel_plan_t {
    native_kernel_kind_t kind;
    int split_k;
};

constexpr int MI355X_CU_COUNT = 256;
constexpr int MI355X_LATENCY_WAVE_TARGET = MI355X_CU_COUNT * 2;

static inline int ceil_div_int(int x, int y) {
    return (x + y - 1) / y;
}

static inline int candidate_waves(
    int M,
    int N,
    int TM,
    int TN,
    int block_threads,
    int split_k = 1)
{
    return ceil_div_int(M, TM) * ceil_div_int(N, TN) * (block_threads / 64) * split_k;
}

static inline int candidate_waste_cols(int N, int TN) {
    return ceil_div_int(N, TN) * TN - N;
}

static inline int candidate_row_tiles(int M, int TM) {
    return ceil_div_int(M, TM);
}

static inline int candidate_stages(int K, int kstep, int split_k = 1) {
    const int chunks = K / kstep;
    return ceil_div_int(chunks, split_k);
}

static inline int choose_small_m_split_k(int M, int N, int K) {
    if (K < 4096 || N < 128) {
        return 1;
    }
    const int base_waves = candidate_waves(M, N, 16, 128, 512);
    if (base_waves >= MI355X_LATENCY_WAVE_TARGET) {
        return 1;
    }
    int split_k = 1;
    const int max_chunks = K / 128;
    while (split_k < 4 && split_k < max_chunks &&
           base_waves * split_k < MI355X_LATENCY_WAVE_TARGET) {
        split_k <<= 1;
    }
    return split_k;
}

static inline int64_t score_candidate(
    int M,
    int N,
    int K,
    int TM,
    int TN,
    int block_threads,
    int kstep,
    int resource_penalty,
    int split_k = 1)
{
    const int waves = candidate_waves(M, N, TM, TN, block_threads, split_k);
    const int stage_count = candidate_stages(K, kstep, split_k);
    const int waste_cols = candidate_waste_cols(N, TN);
    const int row_tiles = candidate_row_tiles(M, TM);
    const int wave_shortfall = waves < MI355X_CU_COUNT ? MI355X_CU_COUNT - waves : 0;
    int long_k_block_penalty = 0;
    if (row_tiles <= 2 && stage_count >= 24 && block_threads > 64) {
        long_k_block_penalty = ((block_threads / 64) - 1) * 256;
    }
    int64_t score = 0;
    score += static_cast<int64_t>(TN) * 8;
    score -= static_cast<int64_t>(stage_count) * 64;
    score -= static_cast<int64_t>(resource_penalty);
    score -= static_cast<int64_t>(waste_cols) * 2;
    score -= static_cast<int64_t>(wave_shortfall) * 32;
    score -= static_cast<int64_t>(long_k_block_penalty);
    score -= static_cast<int64_t>(split_k - 1) * 128;
    return score;
}

static inline native_kernel_plan_t choose_native_kernel_plan(int M, int N, int K) {
    native_kernel_plan_t best{native_kernel_kind_t::k16x16x128, 1};
    int64_t best_score = std::numeric_limits<int64_t>::min();
    auto consider = [&](native_kernel_kind_t kind,
                        int TM,
                        int TN,
                        int block_threads,
                        int kstep,
                        int resource_penalty,
                        int split_k = 1) {
        const int64_t score =
            score_candidate(M, N, K, TM, TN, block_threads, kstep, resource_penalty, split_k);
        if (score > best_score) {
            best = native_kernel_plan_t{kind, split_k};
            best_score = score;
        }
    };

    if (M >= 32) {
        consider(native_kernel_kind_t::k32x32x64, 32, 32, 64, 64, 96);
        if (N >= 128) {
            consider(native_kernel_kind_t::k32x128x64, 32, 128, 256, 64, 512);
        }
        if (N >= 192) {
            consider(native_kernel_kind_t::k32x192x64, 32, 192, 384, 64, 1088);
        }
    } else {
        consider(native_kernel_kind_t::k16x16x128, 16, 16, 64, 128, 96);
        if (N >= 64) {
            consider(native_kernel_kind_t::k16x64x128, 16, 64, 256, 128, 448);
        }
        if (N >= 128) {
            consider(native_kernel_kind_t::k16x128x128, 16, 128, 512, 128, 1024);
            const int split_k = choose_small_m_split_k(M, N, K);
            if (split_k > 1) {
                consider(native_kernel_kind_t::k16x128x128_splitk, 16, 128, 512, 128, 1024, split_k);
            }
        }
    }
    return best;
}

// 检查 HIP launch 是否成功。
static inline void check_hip() {
    hipError_t err = hipGetLastError();
    TORCH_CHECK(err == hipSuccess, "HIP kernel launch failed: ", hipGetErrorString(err));
}

extern "C" torch::Tensor mxfp4_gemm_kernel(
    torch::Tensor A,
    torch::Tensor B_shuffle,
    torch::Tensor B_scale_sh
) {
    CHECK_GPU(A);
    CHECK_GPU(B_shuffle);
    CHECK_GPU(B_scale_sh);
    CHECK_CONTIGUOUS(A);
    CHECK_CONTIGUOUS(B_shuffle);
    CHECK_CONTIGUOUS(B_scale_sh);
    TORCH_CHECK(A.scalar_type() == torch::kBFloat16, "A must be bf16");
    CHECK_UINT8(B_shuffle);
    CHECK_UINT8(B_scale_sh);
    TORCH_CHECK(A.dim() == 2 && B_shuffle.dim() == 2 && B_scale_sh.dim() == 2,
        "A, B_shuffle and B_scale_sh must be 2D");
    TORCH_CHECK(A.size(1) == B_shuffle.size(1) * 2, "A and B_shuffle must have same K");
    TORCH_CHECK(B_scale_sh.size(0) >= B_shuffle.size(0), "B_scale_sh rows must cover N");
    TORCH_CHECK(B_scale_sh.size(1) >= A.size(1) / 32, "B_scale_sh cols must cover K/32 groups");
    TORCH_CHECK(A.size(1) % 64 == 0, "K must be divisible by 64");

    int64_t M = A.size(0);
    int64_t N = B_shuffle.size(0);
    int64_t K = A.size(1);

    auto C = torch::empty({M, N}, A.options().dtype(torch::kBFloat16));

    const auto* A_ptr = reinterpret_cast<const bf16_t*>(A.data_ptr());
    const auto* B_shuffle_ptr = B_shuffle.data_ptr<uint8_t>();
    const auto* B_scale_sh_ptr = B_scale_sh.data_ptr<uint8_t>();
    auto* C_ptr = reinterpret_cast<uint16_t*>(C.data_ptr());
    const int scaleNPad = static_cast<int>(B_scale_sh.size(1));
    const int M_int = static_cast<int>(M);
    const int N_int = static_cast<int>(N);
    const int K_int = static_cast<int>(K);
    const native_kernel_plan_t plan = choose_native_kernel_plan(M_int, N_int, K_int);

    switch (plan.kind) {
        case native_kernel_kind_t::k32x192x64: {
            constexpr int TM = 32;
            constexpr int TN = 192;
            dim3 block(384, 1, 1);
            dim3 grid((N_int + TN - 1) / TN, (M_int + TM - 1) / TM, 1);
            hipLaunchKernelGGL(
                HIP_KERNEL_NAME(mxfp4_mfma_32x192x64_kernel),
                grid,
                block,
                0,
                0,
                A_ptr,
                B_shuffle_ptr,
                B_scale_sh_ptr,
                C_ptr,
                M_int,
                N_int,
                K_int,
                scaleNPad
            );
            break;
        }
        case native_kernel_kind_t::k32x128x64: {
            constexpr int TM = 32;
            constexpr int TN = 128;
            dim3 block(256, 1, 1);
            dim3 grid((N_int + TN - 1) / TN, (M_int + TM - 1) / TM, 1);
            hipLaunchKernelGGL(
                HIP_KERNEL_NAME(mxfp4_mfma_32x128x64_kernel),
                grid,
                block,
                0,
                0,
                A_ptr,
                B_shuffle_ptr,
                B_scale_sh_ptr,
                C_ptr,
                M_int,
                N_int,
                K_int,
                scaleNPad
            );
            break;
        }
        case native_kernel_kind_t::k32x32x64: {
            constexpr int TM = 32;
            constexpr int TN = 32;
            dim3 block(64, 1, 1);
            dim3 grid((N_int + TN - 1) / TN, (M_int + TM - 1) / TM, 1);
            hipLaunchKernelGGL(
                HIP_KERNEL_NAME(mxfp4_mfma_32x32x64_kernel<TM, TN>),
                grid,
                block,
                0,
                0,
                A_ptr,
                B_shuffle_ptr,
                B_scale_sh_ptr,
                C_ptr,
                M_int,
                N_int,
                K_int,
                scaleNPad
            );
            break;
        }
        case native_kernel_kind_t::k16x128x128_splitk: {
            constexpr int TM = 16;
            constexpr int TN = 128;
            auto partial = torch::empty({plan.split_k, M, N}, A.options().dtype(torch::kFloat32));
            auto* partial_ptr = partial.data_ptr<float>();
            dim3 block(512, 1, 1);
            dim3 grid((N_int + TN - 1) / TN, (M_int + TM - 1) / TM, plan.split_k);
            hipLaunchKernelGGL(
                HIP_KERNEL_NAME(mxfp4_mfma_16x128x128_splitk_kernel),
                grid,
                block,
                0,
                0,
                A_ptr,
                B_shuffle_ptr,
                B_scale_sh_ptr,
                partial_ptr,
                M_int,
                N_int,
                K_int,
                scaleNPad,
                plan.split_k
            );
            constexpr int REDUCE_THREADS = 256;
            const int total = M_int * N_int;
            dim3 reduce_block(REDUCE_THREADS, 1, 1);
            dim3 reduce_grid((total + REDUCE_THREADS - 1) / REDUCE_THREADS, 1, 1);
            hipLaunchKernelGGL(
                reduce_splitk_fp32_to_bf16_kernel,
                reduce_grid,
                reduce_block,
                0,
                0,
                partial_ptr,
                C_ptr,
                plan.split_k,
                M_int,
                N_int
            );
            break;
        }
        case native_kernel_kind_t::k16x128x128: {
            constexpr int TM = 16;
            constexpr int TN = 128;
            dim3 block(512, 1, 1);
            dim3 grid((N_int + TN - 1) / TN, (M_int + TM - 1) / TM, 1);
            hipLaunchKernelGGL(
                HIP_KERNEL_NAME(mxfp4_mfma_16x128x128_kernel),
                grid,
                block,
                0,
                0,
                A_ptr,
                B_shuffle_ptr,
                B_scale_sh_ptr,
                C_ptr,
                M_int,
                N_int,
                K_int,
                scaleNPad
            );
            break;
        }
        case native_kernel_kind_t::k16x64x128: {
            constexpr int TM = 16;
            constexpr int TN = 64;
            dim3 block(256, 1, 1);
            dim3 grid((N_int + TN - 1) / TN, (M_int + TM - 1) / TM, 1);
            hipLaunchKernelGGL(
                HIP_KERNEL_NAME(mxfp4_mfma_16x64x128_kernel),
                grid,
                block,
                0,
                0,
                A_ptr,
                B_shuffle_ptr,
                B_scale_sh_ptr,
                C_ptr,
                M_int,
                N_int,
                K_int,
                scaleNPad
            );
            break;
        }
        default: {
            constexpr int TM = 16;
            constexpr int TN = 16;
            dim3 block(64, 1, 1);
            dim3 grid((N_int + TN - 1) / TN, (M_int + TM - 1) / TM, 1);
            hipLaunchKernelGGL(
                HIP_KERNEL_NAME(mxfp4_mfma_16x16x128_kernel<TM, TN>),
                grid,
                block,
                0,
                0,
                A_ptr,
                B_shuffle_ptr,
                B_scale_sh_ptr,
                C_ptr,
                M_int,
                N_int,
                K_int,
                scaleNPad
            );
            break;
        }
    }
    check_hip();
    return C;
}

extern "C" torch::Tensor mxfp4_gemm_prequant(
    torch::Tensor A_q,
    torch::Tensor A_scale,
    torch::Tensor B_shuffle,
    torch::Tensor B_scale
) {
    CHECK_GPU(A_q);
    CHECK_GPU(A_scale);
    CHECK_GPU(B_shuffle);
    CHECK_GPU(B_scale);
    CHECK_CONTIGUOUS(A_q);
    CHECK_CONTIGUOUS(A_scale);
    CHECK_CONTIGUOUS(B_shuffle);
    CHECK_CONTIGUOUS(B_scale);
    CHECK_UINT8(A_q);
    CHECK_UINT8(A_scale);
    CHECK_UINT8(B_shuffle);
    CHECK_UINT8(B_scale);
    TORCH_CHECK(
        A_q.dim() == 2 && A_scale.dim() == 2 && B_shuffle.dim() == 2 && B_scale.dim() == 2,
        "A_q, A_scale, B_shuffle and B_scale must be 2D");
    TORCH_CHECK(A_q.size(0) == A_scale.size(0), "A_q and A_scale must have same M");
    TORCH_CHECK(A_q.size(1) * 2 == B_shuffle.size(1) * 2, "A_q and B_shuffle must have same K");
    TORCH_CHECK(A_scale.size(1) * 32 == A_q.size(1) * 2, "A_scale must have K/32 groups");
    TORCH_CHECK(B_scale.size(0) == B_shuffle.size(0), "B_shuffle and B_scale must have same N");
    TORCH_CHECK(B_scale.size(1) * 32 == A_q.size(1) * 2, "B_scale must have K/32 groups");
    TORCH_CHECK((A_q.size(1) * 2) % 64 == 0, "K must be divisible by 64");

    int64_t M = A_q.size(0);
    int64_t N = B_shuffle.size(0);
    int64_t K = A_q.size(1) * 2;

    auto C = torch::empty({M, N}, A_q.options().dtype(torch::kBFloat16));

    const auto* A_q_ptr = A_q.data_ptr<uint8_t>();
    const auto* A_scale_ptr = A_scale.data_ptr<uint8_t>();
    const auto* B_shuffle_ptr = B_shuffle.data_ptr<uint8_t>();
    const auto* B_scale_ptr = B_scale.data_ptr<uint8_t>();
    auto* C_ptr = reinterpret_cast<uint16_t*>(C.data_ptr());

    if (M >= 32) {
        constexpr int TM = 32;
        constexpr int TN = 32;
        dim3 block(64, 1, 1);
        dim3 grid((N + TN - 1) / TN, (M + TM - 1) / TM, 1);
        hipLaunchKernelGGL(
            HIP_KERNEL_NAME(mxfp4_mfma_32x32x64_prequant_kernel<TM, TN>),
            grid,
            block,
            0,
            0,
            A_q_ptr,
            A_scale_ptr,
            B_shuffle_ptr,
            B_scale_ptr,
            C_ptr,
            static_cast<int>(M),
            static_cast<int>(N),
            static_cast<int>(K)
        );
    } else {
        constexpr int TM = 16;
        constexpr int TN = 16;
        dim3 block(64, 1, 1);
        dim3 grid((N + TN - 1) / TN, (M + TM - 1) / TM, 1);
        hipLaunchKernelGGL(
            HIP_KERNEL_NAME(mxfp4_mfma_16x16x128_prequant_kernel<TM, TN>),
            grid,
            block,
            0,
            0,
            A_q_ptr,
            A_scale_ptr,
            B_shuffle_ptr,
            B_scale_ptr,
            C_ptr,
            static_cast<int>(M),
            static_cast<int>(N),
            static_cast<int>(K)
        );
    }
    check_hip();
    return C;
}

extern "C" torch::Tensor mxfp4_pack_a_debug(
    torch::Tensor A
) {
    CHECK_GPU(A);
    CHECK_CONTIGUOUS(A);
    TORCH_CHECK(A.scalar_type() == torch::kBFloat16, "A must be bf16");
    TORCH_CHECK(A.dim() == 2, "A must be 2D");
    TORCH_CHECK(A.size(1) % 64 == 0, "K must be divisible by 64");

    int64_t M = A.size(0);
    int64_t K = A.size(1);
    auto A_q = torch::empty({M, K / 2}, A.options().dtype(torch::kUInt8));

    const auto* A_ptr = reinterpret_cast<const bf16_t*>(A.data_ptr());
    auto* A_q_ptr = A_q.data_ptr<uint8_t>();
    const int total_groups = static_cast<int>(M * (K / 32));
    constexpr int THREADS = 256;
    dim3 block(THREADS, 1, 1);
    dim3 grid((total_groups + THREADS - 1) / THREADS, 1, 1);
    hipLaunchKernelGGL(
        mxfp4_pack_a_debug_kernel,
        grid,
        block,
        0,
        0,
        A_ptr,
        A_q_ptr,
        static_cast<int>(M),
        static_cast<int>(K)
    );
    check_hip();
    return A_q;
}
"""

def _get_module():
    global _module
    if _module is None:
        _module = load_inline(
            name="mxfp4_gemm_release6",
            cpp_sources=[CPP_SRC],
            cuda_sources=[HIP_SRC],
            functions=["mxfp4_gemm_kernel"],
            verbose=False,
            extra_cflags=["-std=c++20", "-O3"],
            extra_cuda_cflags=["--offload-arch=gfx950", "-std=c++20", "-O3"],
            extra_include_paths=_resolve_aiter_include_dirs(),
        )
    return _module


def _view_preshuffled_weight_u8(weight_sh: torch.Tensor) -> torch.Tensor:
    rows, cols = weight_sh.shape
    assert rows % 16 == 0, "B_shuffle rows must be divisible by 16"
    return weight_sh.view(rows // 16, cols * 16)


def _view_preshuffled_scales_u8(scale_sh: torch.Tensor) -> torch.Tensor:
    scale_u8 = scale_sh.view(torch.uint8)
    rows, cols = scale_u8.shape
    assert rows % 32 == 0, "scale rows must be padded to a multiple of 32"
    return scale_u8.view(rows // 32, cols * 32)


def _quantize_a_for_preshuffle_gemm(A: torch.Tensor):
    A_q_u8, A_scale_u8 = dynamic_mxfp4_quant(A)
    A_scale_u8 = A_scale_u8.view(torch.uint8).contiguous()
    A_scale_triton = None
    if A.shape[0] >= 32:
        A_scale_triton = _view_preshuffled_scales_u8(e8m0_shuffle(A_scale_u8))
    return A_q_u8, A_scale_u8, A_scale_triton


def custom_kernel(data: input_t) -> output_t:
    """
    默认走 aiter 的量化 + ASM GEMM 热路径。
    原生 inline MFMA 仅保留为实验分支。
    布局说明:
    - A 是稠密 bf16 [M, K]
    - B_shuffle 是已经做过 bpreshuffle 的 packed fp4 [N, K/2]
    - B_scale_sh 保持 shuffle 后的 E8M0 布局
    - 默认路径使用参考的 Triton quant + e8m0 shuffle,再交给 aiter ASM GEMM
    """
    A, _, _, B_shuffle, B_scale_sh = data
    A = A.contiguous()
    B_shuffle = B_shuffle.contiguous()
    B_scale_sh = B_scale_sh.contiguous()

    if not USE_NATIVE_MFMA and USE_TRITON_PRESHUFFLE:
        M, K = A.shape
        A_q_u8, A_scale_u8, A_scale_triton = _quantize_a_for_preshuffle_gemm(A)
        B_triton = _view_preshuffled_weight_u8(B_shuffle.view(torch.uint8))
        B_scale_triton = _view_preshuffled_scales_u8(B_scale_sh)

        if M < 32:
            A_scale_triton = A_scale_u8

        return gemm_afp4wfp4_preshuffled_weight_scales(
            A_q_u8,
            B_triton,
            A_scale_triton,
            B_scale_triton,
            torch.bfloat16,
        )

    if not USE_NATIVE_MFMA:
        A_q_u8, A_scale_u8 = dynamic_mxfp4_quant(A)
        A_q = A_q_u8.view(dtypes.fp4x2)
        A_scale_sh = e8m0_shuffle(A_scale_u8).view(dtypes.fp8_e8m0)
        return aiter.gemm_a4w4(
            A_q,
            B_shuffle,
            A_scale_sh,
            B_scale_sh,
            dtype=dtypes.bf16,
            bpreshuffle=True,
        )

    return _get_module().mxfp4_gemm_kernel(
        A,
        B_shuffle.view(torch.uint8),
        B_scale_sh.view(torch.uint8),
    )
scrolls · 1966 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