Skip to content
KernelIndex
Search⌘K

submission 710593

kida023_89704 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:c99eb1d613262a9dec96724c7d7b84fa4bfd964346fbd3748735eff6ddca308f
license declaredunknown
license concludedunknown
authorskida023_89704
imported2026-08-26

Techniques

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

fp4static_assert(sizeof(KernelArgs) == 384, "Unexpected FP4 kernarg size");
num-warps = 1num_warps = 1
shared-memory__shared__ __align__(16) uint8_t s_a[2][TILE_M * (TILE_K / 2)];
split-ksplitk_enabled: int
stages = 1num_stages = 1
tile-k = 256constexpr int TILE_K = 256;
tile-m = 16static_assert(TILE_M == 16 || TILE_M == 32, "Only TILE_M=16/32 is supported");
tile-n = 128constexpr int TILE_N = 128;

Kernel source

submission_fused.py2104 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

from __future__ import annotations

import csv
import os
from dataclasses import dataclass
from functools import lru_cache
from pathlib import Path

os.environ.setdefault("CXX", "clang++")
os.environ.setdefault("HSA_XNACK", "0")
os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")

import torch
from torch.utils.cpp_extension import load_inline

from task import input_t, output_t


_ARCH = "gfx950"
_F4GEMM_SUBDIR = "f4gemm"
_CUSTOM_SMALLM_MAX_M = 32
_BENCHMARK_QUANT_WORKSPACE_MAX = 4
_BENCHMARK_OUT_CACHE_MAX = 4
_CUSTOM_N_TILE = 128
_FAST_FUSED_ENABLE = True

_SUPPORTED_SHAPES = {
    (4, 2880, 512),
    (8, 2112, 7168),
    (16, 2112, 7168),
    (16, 3072, 1536),
    (32, 2880, 512),
    (32, 4096, 512),
    (64, 3072, 1536),
    (64, 7168, 2048),
    (256, 2880, 512),
    (256, 3072, 1536),
}

_SMALLK_BENCHMARK_SHAPES = {
    (4, 2880, 512),
    (32, 2880, 512),
    (32, 4096, 512),
}

_COMBINED_HIP_QUANT_GEMM_SHAPES = {
    (4, 2880, 512),
    (16, 2112, 7168),
    (32, 2880, 512),
    (32, 4096, 512),
    (64, 7168, 2048),
    (256, 3072, 1536),
}

# Optional per-shape override:
# (m, n, k): (kernel_name, co_name_with_subdir, log2_k_split)
_KERNEL_OVERRIDE_BY_SHAPE: dict[tuple[int, int, int], tuple[str, str, int]] = {
    (4, 2880, 512): (
        "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
        "f4gemm/f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co",
        0,
    ),
    (8, 2112, 7168): (
        "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
        "f4gemm/f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co",
        0,
    ),
    (16, 2112, 7168): (
        "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
        "f4gemm/f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co",
        0,
    ),
    (16, 3072, 1536): (
        "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x768E",
        "f4gemm/f4gemm_bf16_per1x32Fp4_BpreShuffle_32x768.co",
        0,
    ),
    (32, 2880, 512): (
        "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
        "f4gemm/f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co",
        0,
    ),
    (32, 4096, 512): (
        "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
        "f4gemm/f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co",
        0,
    ),
    (64, 3072, 1536): (
        "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
        "f4gemm/f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co",
        0,
    ),
    (64, 7168, 2048): (
        "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
        "f4gemm/f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co",
        0,
    ),
    (256, 2880, 512): (
        "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
        "f4gemm/f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co",
        0,
    ),
    (256, 3072, 1536): (
        "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
        "f4gemm/f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co",
        0,
    ),
}


@dataclass(frozen=True)
class _KernelCfg:
    tile_m: int
    tile_n: int
    splitk_enabled: int
    bpreshuffle: int
    kernel_name: str
    co_name: str


@dataclass(frozen=True)
class _TunedKernel:
    cu_num: int
    m: int
    n: int
    k: int
    split_k: int
    kernel_name: str


@dataclass
class _QuantWorkspaceEntry:
    device_type: str
    device_index: int
    m: int
    n: int
    q_rows: int
    q: torch.Tensor
    scale_sh: torch.Tensor


@dataclass
class _OutCacheEntry:
    device_type: str
    device_index: int
    dtype: torch.dtype
    padded_m: int
    n: int
    out: torch.Tensor


_B_SHUFFLE_CACHE: dict[tuple[int, int, int, int, int], tuple[torch.Tensor, torch.Tensor, int]] = {}
_BENCHMARK_QUANT_WORKSPACES: list[_QuantWorkspaceEntry] = []
_BENCHMARK_OUT_CACHE: list[_OutCacheEntry] = []


def _is_benchmark_fastpath_shape(shape: tuple[int, int, int]) -> bool:
    return shape in _SUPPORTED_SHAPES


def _is_smallk_benchmark_shape(shape: tuple[int, int, int]) -> bool:
    return shape in _SMALLK_BENCHMARK_SHAPES


def _load_inline_with_trace(tag: str, **kwargs):
    print(f"[submission-ext] tag={tag} stage=load-start", flush=True)
    module = load_inline(**kwargs)
    print(f"[submission-ext] tag={tag} stage=load-done", flush=True)
    return module

_FAST_FUSED_CPP_SRC = r"""
#include <torch/extension.h>

void launch_fast_quant_mxfp4(
    torch::Tensor a_bf16,
    torch::Tensor q_out,
    torch::Tensor scale_sh_out,
    int real_m,
    int k);

void launch_fast_quant_and_f4gemm(
    torch::Tensor a_bf16,
    torch::Tensor q_out,
    torch::Tensor scale_sh_out,
    torch::Tensor b_shuffle,
    torch::Tensor b_scale_sh,
    torch::Tensor out,
    std::string co_path,
    std::string kernel_name,
    int tile_m,
    int tile_n,
    int log2_k_split,
    int real_m,
    int k);

void launch_fast_fused_mxfp4_gemm(
    torch::Tensor a_bf16,
    torch::Tensor b_shuffle_u8,
    torch::Tensor b_scale_sh_u8,
    torch::Tensor out,
    int padded_n,
    int real_n);
"""

_FAST_FUSED_CUDA_SRC = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <hip/hip_bfloat16.h>
#include <c10/hip/HIPFunctions.h>

#include <cmath>
#include <cstdint>
#include <mutex>
#include <stdexcept>
#include <string>
#include <unordered_map>

namespace {

using i32x8_t = int __attribute__((ext_vector_type(8)));
using fp32x4_t = float __attribute__((ext_vector_type(4)));

struct p3 {
    uint32_t x;
    uint32_t y;
    uint32_t z;
};

struct p2 {
    uint32_t x;
    uint32_t y;
};

struct __attribute__((packed)) KernelArgs {
    void* ptr_D;
    p2 _p0;
    void* ptr_C;
    p2 _p1;
    void* ptr_A;
    p2 _p2;
    void* ptr_B;
    p2 _p3;
    float alpha;
    p3 _p4;
    float beta;
    p3 _p5;
    uint32_t stride_D0;
    p3 _p6;
    uint32_t stride_D1;
    p3 _p7;
    uint32_t stride_C0;
    p3 _p8;
    uint32_t stride_C1;
    p3 _p9;
    uint32_t stride_A0;
    p3 _p10;
    uint32_t stride_A1;
    p3 _p11;
    uint32_t stride_B0;
    p3 _p12;
    uint32_t stride_B1;
    p3 _p13;
    uint32_t Mdim;
    p3 _p14;
    uint32_t Ndim;
    p3 _p15;
    uint32_t Kdim;
    p3 _p16;
    void* ptr_ScaleA;
    p2 _p17;
    void* ptr_ScaleB;
    p2 _p18;
    uint32_t stride_ScaleA0;
    p3 _p19;
    uint32_t stride_ScaleA1;
    p3 _p20;
    uint32_t stride_ScaleB0;
    p3 _p21;
    uint32_t stride_ScaleB1;
    p3 _p22;
    int32_t log2_k_split;
    p3 _p23;
};

static_assert(sizeof(KernelArgs) == 384, "Unexpected FP4 kernarg size");

struct CachedKernel {
    hipModule_t module = nullptr;
    hipFunction_t func = nullptr;
};

std::unordered_map<std::string, CachedKernel>& fast_kernel_cache() {
    static std::unordered_map<std::string, CachedKernel> cache;
    return cache;
}

std::mutex& fast_kernel_cache_mutex() {
    static std::mutex mu;
    return mu;
}

void hip_check(hipError_t err, const char* call_name) {
    if (err == hipSuccess) {
        return;
    }
    throw std::runtime_error(std::string(call_name) + " failed: " + hipGetErrorString(err));
}

CachedKernel& get_fast_kernel(const std::string& co_path, const std::string& kernel_name) {
    const std::string key = co_path + "|" + kernel_name;
    std::lock_guard<std::mutex> guard(fast_kernel_cache_mutex());
    auto& cache = fast_kernel_cache();
    auto it = cache.find(key);
    if (it != cache.end()) {
        return it->second;
    }

    CachedKernel entry;
    hip_check(hipModuleLoad(&entry.module, co_path.c_str()), "hipModuleLoad");
    hip_check(hipModuleGetFunction(&entry.func, entry.module, kernel_name.c_str()), "hipModuleGetFunction");
    auto [new_it, _inserted] = cache.emplace(key, entry);
    return new_it->second;
}

void pick_quant_launch_config(int real_m, int k, int* threads, int* blocks) {
    const int total_tasks = real_m * (k / 32);
    if (total_tasks <= 64) {
        *threads = 64;
        *blocks = 1;
        return;
    }
    if (total_tasks <= 1024) {
        *threads = 128;
        *blocks = max(1, min(8, (total_tasks + *threads - 1) / *threads));
        return;
    }
    *threads = 256;
    *blocks = max(1, min(120, (total_tasks + *threads - 1) / *threads));
}

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

__device__ __forceinline__ uint8_t amax_to_scale_e8m0(float amax) {
    const uint32_t rounded_bits = (__float_as_uint(amax) + 0x200000u) & 0xFF800000u;
    int exponent = static_cast<int>((rounded_bits >> 23) & 0xFFu) - 127;
    exponent -= 2;
    exponent = max(-127, min(127, exponent));
    return static_cast<uint8_t>(exponent + 127);
}

__device__ __forceinline__ uint8_t f32_to_fp4_bits(float x) {
    constexpr int MBITS_F32 = 23;
    constexpr int EBITS_F32 = 8;
    constexpr int ebits = 2;
    constexpr int mbits = 1;
    constexpr uint8_t max_int = static_cast<uint8_t>((1 << (ebits + mbits)) - 1);
    constexpr uint8_t sign_mask = static_cast<uint8_t>(1 << (ebits + mbits));
    constexpr int exp_bias = (1 << (ebits - 1)) - 1;
    constexpr int f32_exp_bias = (1 << (EBITS_F32 - 1)) - 1;
    constexpr int magic_adder = (1 << (MBITS_F32 - mbits - 1)) - 1;
    constexpr float max_normal = 6.0f;
    constexpr float min_normal = 1.0f;
    constexpr int denorm_exp = (f32_exp_bias - exp_bias) + (MBITS_F32 - mbits) + 1;
    constexpr int denorm_mask_int = denorm_exp << MBITS_F32;
    const float denorm_mask_float = __uint_as_float(static_cast<uint32_t>(denorm_mask_int));

    const uint32_t raw = __float_as_uint(x);
    const uint32_t sign = raw & 0x80000000u;
    const uint32_t mag_bits = raw ^ sign;
    const float mag = __uint_as_float(mag_bits);

    const bool saturate_mask = mag >= max_normal;
    const bool denormal_mask = (!saturate_mask) && (mag < min_normal);
    const bool normal_mask = (!saturate_mask) && (!denormal_mask);

    uint8_t out = max_int;
    if (denormal_mask) {
        int denormal_x = __float_as_int(mag + denorm_mask_float);
        denormal_x -= denorm_mask_int;
        out = static_cast<uint8_t>(denormal_x);
    } else if (normal_mask) {
        int normal_x = static_cast<int>(mag_bits);
        const int mant_odd = (normal_x >> (MBITS_F32 - mbits)) & 1;
        const int val_to_add = ((exp_bias - f32_exp_bias) << MBITS_F32) + magic_adder;
        normal_x += val_to_add;
        normal_x += mant_odd;
        normal_x >>= (MBITS_F32 - mbits);
        out = static_cast<uint8_t>(normal_x);
    }

    const uint8_t sign_lp = static_cast<uint8_t>((sign >> (MBITS_F32 + EBITS_F32 - mbits - ebits)) & sign_mask);
    return static_cast<uint8_t>(out | sign_lp);
}

template <int ScaleASel, int ScaleBSel>
__device__ __forceinline__ fp32x4_t mfma_scale_16x16x128(
    i32x8_t a_reg, i32x8_t b_reg, fp32x4_t c_reg, int scale_a_word, int scale_b_word) {
#if defined(__gfx950__)
    return __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
        a_reg, b_reg, c_reg, 4, 4, ScaleASel, scale_a_word, ScaleBSel, scale_b_word);
#else
    return c_reg;
#endif
}

__device__ __forceinline__ void sched_barrier() {
#if defined(__gfx950__)
    __builtin_amdgcn_sched_barrier(0);
#endif
}

__device__ __forceinline__ float reduce_max4(float x) {
    x = fmaxf(x, __shfl_xor(x, 1, 4));
    x = fmaxf(x, __shfl_xor(x, 2, 4));
    return x;
}

__device__ __forceinline__ int swizzle_xor16(int row, int col, int k_blocks16) {
    return col ^ ((row % k_blocks16) * 16);
}

__device__ __forceinline__ i32x8_t pack_i64x4_to_i32x8(uint64_t x0, uint64_t x1, uint64_t x2, uint64_t x3) {
    struct Pack {
        uint64_t u64[4];
    };
    const Pack p{{x0, x1, x2, x3}};
    return __builtin_bit_cast(i32x8_t, p);
}

template <int TILE_M>
__device__ __forceinline__ void quantize_a_tile_to_stage(
    const hip_bfloat16* __restrict__ a_bf16, int m, int a_stride, int tile_m_base, int tile_k_base, int tx,
    uint8_t* __restrict__ s_a_stage, uint8_t* __restrict__ s_scale_stage) {
    constexpr int TILE_K = 256;
    constexpr int SCALE_GROUPS = TILE_K / 32;
    constexpr int K_BLOCKS16 = TILE_K / 32;
    constexpr int BLOCK_THREADS = 256;
    constexpr int TASKS_PER_ROW = TILE_K / 8;
    constexpr int TASKS_TOTAL = TILE_M * TASKS_PER_ROW;
    static_assert(TASKS_TOTAL % BLOCK_THREADS == 0);
    constexpr int TASKS_PER_THREAD = TASKS_TOTAL / BLOCK_THREADS;

    #pragma unroll
    for (int i = 0; i < TASKS_PER_THREAD; ++i) {
        const int task_id = i * BLOCK_THREADS + tx;
        const int row_local = task_id / TASKS_PER_ROW;
        const int col_task = task_id % TASKS_PER_ROW;
        const int block32_idx = col_task / 4;
        const int in_block_task = col_task % 4;
        const bool is_scale_leader = (in_block_task == 0);
        const int src_row = tile_m_base + row_local;

        float vals[8];
        float max_abs = 0.0f;
        #pragma unroll
        for (int j = 0; j < 8; ++j) {
            float v = 0.0f;
            if (src_row < m) {
                const int src_col = tile_k_base + col_task * 8 + j;
                v = static_cast<float>(a_bf16[src_row * a_stride + src_col]);
            }
            vals[j] = v;
            max_abs = fmaxf(max_abs, fabsf(v));
        }

        const float group_amax = reduce_max4(max_abs);
        const uint8_t scale_byte = (src_row < m) ? amax_to_scale_e8m0(group_amax) : static_cast<uint8_t>(127);
        const float inv_scale = (src_row < m && scale_byte != 0) ? (1.0f / decode_e8m0(scale_byte)) : 0.0f;

        uint32_t packed = 0;
        #pragma unroll
        for (int j = 0; j < 8; ++j) {
            const uint8_t nibble = f32_to_fp4_bits(vals[j] * inv_scale);
            packed |= static_cast<uint32_t>(nibble) << (j * 4);
        }

        const int col_local_bytes = col_task * 4;
        const int col_swz_bytes = swizzle_xor16(row_local, col_local_bytes, K_BLOCKS16);
        reinterpret_cast<uint32_t*>(s_a_stage + row_local * (TILE_K / 2) + col_swz_bytes)[0] = packed;
        if (is_scale_leader) {
            s_scale_stage[row_local * SCALE_GROUPS + block32_idx] = scale_byte;
        }
    }
}

template <int TILE_M>
__global__ __launch_bounds__(256) void fast_fused_mxfp4_mfma_kernel(
    const hip_bfloat16* __restrict__ a_bf16,
    const uint8_t* __restrict__ b_shuffle,
    const uint8_t* __restrict__ b_scale_sh,
    hip_bfloat16* __restrict__ out,
    int m,
    int n_padded,
    int real_n,
    int k,
    int a_stride,
    int out_stride) {
    constexpr int TILE_K = 256;
    constexpr int TILE_N = 128;
    constexpr int SCALE_BLOCK_K = 32;
    constexpr int SCALE_GROUPS = TILE_K / SCALE_BLOCK_K;
    constexpr int K_BLOCKS16 = TILE_K / 32;
    constexpr int ROW_GROUPS = TILE_M / 16;
    constexpr int ACCUMULATORS_PER_WAVE = ROW_GROUPS * 2;
    static_assert(TILE_M == 16 || TILE_M == 32, "Only TILE_M=16/32 is supported");

    __shared__ __align__(16) uint8_t s_a[2][TILE_M * (TILE_K / 2)];
    __shared__ __align__(4) uint8_t s_scale[2][TILE_M * SCALE_GROUPS];
    __shared__ __align__(16) hip_bfloat16 s_out[TILE_M][TILE_N];

    const int tx = static_cast<int>(threadIdx.x);
    const int lane_id = tx & 63;
    const int wave_id = tx >> 6;
    const int lane_div_16 = lane_id >> 4;
    const int lane_mod_16 = lane_id & 15;

    const int tile_m_base = static_cast<int>(blockIdx.y) * TILE_M;
    const int tile_n_base = static_cast<int>(blockIdx.x) * TILE_N;

    fp32x4_t acc[ACCUMULATORS_PER_WAVE];
    #pragma unroll
    for (int i = 0; i < ACCUMULATORS_PER_WAVE; ++i) {
        acc[i] = fp32x4_t{0.0f, 0.0f, 0.0f, 0.0f};
    }

    const int k0_blocks = k / 128;
    const int scale_k_tiles = k / 256;
    const int num_k_tiles = k / TILE_K;
    const int row_a_lds = lane_mod_16;
    const int col_offset_base_bytes = lane_div_16 * 16;

    quantize_a_tile_to_stage<TILE_M>(a_bf16, m, a_stride, tile_m_base, 0, tx, s_a[0], s_scale[0]);
    __syncthreads();

    for (int tile_idx = 0; tile_idx < num_k_tiles; ++tile_idx) {
        const int curr = tile_idx & 1;
        const int next = curr ^ 1;
        const int tile_k_base = tile_idx * TILE_K;

        if (tile_idx + 1 < num_k_tiles) {
            quantize_a_tile_to_stage<TILE_M>(
                a_bf16, m, a_stride, tile_m_base, (tile_idx + 1) * TILE_K, tx, s_a[next], s_scale[next]);
        }

        const int row0 = lane_mod_16;
        const int blk0 = lane_div_16;
        const int blk1 = lane_div_16 + 4;
        const uint32_t a_scale_row0_blk0 = static_cast<uint32_t>(s_scale[curr][row0 * SCALE_GROUPS + blk0]);
        const uint32_t a_scale_row0_blk1 = static_cast<uint32_t>(s_scale[curr][row0 * SCALE_GROUPS + blk1]);
        uint32_t a_scale_word = a_scale_row0_blk0 | (a_scale_row0_blk1 << 16);
        if constexpr (TILE_M == 32) {
            const int row1 = row0 + 16;
            a_scale_word |= static_cast<uint32_t>(s_scale[curr][row1 * SCALE_GROUPS + blk0]) << 8;
            a_scale_word |= static_cast<uint32_t>(s_scale[curr][row1 * SCALE_GROUPS + blk1]) << 24;
        } else {
            a_scale_word |= 0x7F00u | 0x7F000000u;
        }

        const int n_pack = (tile_n_base + wave_id * 32) / 32;
        const int k_pack = tile_k_base / 256;
        
        uint32_t b_scale_word = 0x7F7F7F7F;
        if (n_pack < (real_n / 32)) {
            const int b_scale_word_idx = (((n_pack * scale_k_tiles + k_pack) * 4 + lane_div_16) * 16 + lane_mod_16);
            b_scale_word = reinterpret_cast<const uint32_t*>(b_scale_sh)[b_scale_word_idx];
        }

        #pragma unroll
        for (int k_half = 0; k_half < 2; ++k_half) {
            const int col_base_bytes = col_offset_base_bytes + k_half * 64;

            const int a_row0_col = swizzle_xor16(row_a_lds, col_base_bytes, K_BLOCKS16);
            const uint64_t a00 = reinterpret_cast<const uint64_t*>(s_a[curr] + row_a_lds * (TILE_K / 2) + a_row0_col)[0];
            const uint64_t a01 = reinterpret_cast<const uint64_t*>(s_a[curr] + row_a_lds * (TILE_K / 2) + a_row0_col)[1];
            const i32x8_t a_vec0 = pack_i64x4_to_i32x8(a00, a01, 0ull, 0ull);
            i32x8_t a_vec1 = i32x8_t{};
            if constexpr (TILE_M == 32) {
                const int a_row1_col = swizzle_xor16(row_a_lds + 16, col_base_bytes, K_BLOCKS16);
                const uint64_t a10 = reinterpret_cast<const uint64_t*>(s_a[curr] + (row_a_lds + 16) * (TILE_K / 2) + a_row1_col)[0];
                const uint64_t a11 = reinterpret_cast<const uint64_t*>(s_a[curr] + (row_a_lds + 16) * (TILE_K / 2) + a_row1_col)[1];
                a_vec1 = pack_i64x4_to_i32x8(a10, a11, 0ull, 0ull);
            }

            const int n_blk0 = (tile_n_base + wave_id * 32 + lane_mod_16) / 16;
            const int n_blk1 = (tile_n_base + wave_id * 32 + 16 + lane_mod_16) / 16;
            const int n_intra = lane_mod_16;
            const int k0 = tile_k_base / 128 + k_half;

            uint64_t b00 = 0;
            uint64_t b01 = 0;
            if (n_blk0 < (real_n / 16)) {
                const int b_idx0 = ((((n_blk0 * k0_blocks + k0) * 4 + lane_div_16) * 16 + n_intra) * 16);
                b00 = reinterpret_cast<const uint64_t*>(b_shuffle + b_idx0)[0];
                b01 = reinterpret_cast<const uint64_t*>(b_shuffle + b_idx0)[1];
            }
            const i32x8_t b_vec0 = pack_i64x4_to_i32x8(b00, b01, 0ull, 0ull);

            uint64_t b10 = 0;
            uint64_t b11 = 0;
            if (n_blk1 < (real_n / 16)) {
                const int b_idx1 = ((((n_blk1 * k0_blocks + k0) * 4 + lane_div_16) * 16 + n_intra) * 16);
                b10 = reinterpret_cast<const uint64_t*>(b_shuffle + b_idx1)[0];
                b11 = reinterpret_cast<const uint64_t*>(b_shuffle + b_idx1)[1];
            }
            const i32x8_t b_vec1 = pack_i64x4_to_i32x8(b10, b11, 0ull, 0ull);

            sched_barrier();
            if (k_half == 0) {
                acc[0] = mfma_scale_16x16x128<0, 0>(
                    a_vec0, b_vec0, acc[0], static_cast<int>(a_scale_word), static_cast<int>(b_scale_word));
                acc[1] = mfma_scale_16x16x128<0, 1>(
                    a_vec0, b_vec1, acc[1], static_cast<int>(a_scale_word), static_cast<int>(b_scale_word));
                if constexpr (TILE_M == 32) {
                    acc[2] = mfma_scale_16x16x128<1, 0>(
                        a_vec1, b_vec0, acc[2], static_cast<int>(a_scale_word), static_cast<int>(b_scale_word));
                    acc[3] = mfma_scale_16x16x128<1, 1>(
                        a_vec1, b_vec1, acc[3], static_cast<int>(a_scale_word), static_cast<int>(b_scale_word));
                }
            } else {
                acc[0] = mfma_scale_16x16x128<2, 2>(
                    a_vec0, b_vec0, acc[0], static_cast<int>(a_scale_word), static_cast<int>(b_scale_word));
                acc[1] = mfma_scale_16x16x128<2, 3>(
                    a_vec0, b_vec1, acc[1], static_cast<int>(a_scale_word), static_cast<int>(b_scale_word));
                if constexpr (TILE_M == 32) {
                    acc[2] = mfma_scale_16x16x128<3, 2>(
                        a_vec1, b_vec0, acc[2], static_cast<int>(a_scale_word), static_cast<int>(b_scale_word));
                    acc[3] = mfma_scale_16x16x128<3, 3>(
                        a_vec1, b_vec1, acc[3], static_cast<int>(a_scale_word), static_cast<int>(b_scale_word));
                }
            }
            sched_barrier();
        }
        if (tile_idx + 1 < num_k_tiles) {
            __syncthreads();
        }
    }

    const int col_base = tile_n_base + wave_id * 32 + lane_mod_16;
    const int col_local = wave_id * 32 + lane_mod_16;
    #pragma unroll
    for (int row_group = 0; row_group < ROW_GROUPS; ++row_group) {
        #pragma unroll
        for (int ii = 0; ii < 4; ++ii) {
            const int row_in_tile = row_group * 16 + lane_div_16 * 4 + ii;
            s_out[row_in_tile][col_local] = static_cast<hip_bfloat16>(acc[row_group * 2 + 0][ii]);
            s_out[row_in_tile][col_local + 16] = static_cast<hip_bfloat16>(acc[row_group * 2 + 1][ii]);
        }
    }
    __syncthreads();

    #pragma unroll
    for (int row_group = 0; row_group < ROW_GROUPS; ++row_group) {
        #pragma unroll
        for (int ii = 0; ii < 4; ++ii) {
            const int row_in_tile = row_group * 16 + lane_div_16 * 4 + ii;
            const int row = tile_m_base + row_in_tile;
            if (row >= m) continue;
            out[static_cast<int64_t>(row) * out_stride + col_base] = s_out[row_in_tile][col_local];
            out[static_cast<int64_t>(row) * out_stride + col_base + 16] = s_out[row_in_tile][col_local + 16];
        }
    }
}

}  // namespace

void launch_fast_fused_mxfp4_gemm(
    torch::Tensor a_bf16,
    torch::Tensor b_shuffle_u8,
    torch::Tensor b_scale_sh_u8,
    torch::Tensor out,
    int padded_n,
    int real_n) {
    
    const int m = static_cast<int>(a_bf16.size(0));
    const int k = static_cast<int>(a_bf16.size(1));
    const bool use_m16_kernel = m <= 16;
    const int tile_m = use_m16_kernel ? 16 : 32;
    const dim3 grid(static_cast<unsigned int>(padded_n / 128), static_cast<unsigned int>((m + tile_m - 1) / tile_m), 1);
    const dim3 block(256, 1, 1);

    if (use_m16_kernel) {
        hipLaunchKernelGGL(
            HIP_KERNEL_NAME((fast_fused_mxfp4_mfma_kernel<16>)),
            grid, block, 0, 0,
            reinterpret_cast<const hip_bfloat16*>(a_bf16.data_ptr()),
            reinterpret_cast<const uint8_t*>(b_shuffle_u8.data_ptr()),
            reinterpret_cast<const uint8_t*>(b_scale_sh_u8.data_ptr()),
            reinterpret_cast<hip_bfloat16*>(out.data_ptr()),
            m, padded_n, real_n, k,
            static_cast<int>(a_bf16.stride(0)), static_cast<int>(out.stride(0)));
    } else {
        hipLaunchKernelGGL(
            HIP_KERNEL_NAME((fast_fused_mxfp4_mfma_kernel<32>)),
            grid, block, 0, 0,
            reinterpret_cast<const hip_bfloat16*>(a_bf16.data_ptr()),
            reinterpret_cast<const uint8_t*>(b_shuffle_u8.data_ptr()),
            reinterpret_cast<const uint8_t*>(b_scale_sh_u8.data_ptr()),
            reinterpret_cast<hip_bfloat16*>(out.data_ptr()),
            m, padded_n, real_n, k,
            static_cast<int>(a_bf16.stride(0)), static_cast<int>(out.stride(0)));
    }

    const hipError_t err = hipGetLastError();
    TORCH_CHECK(err == hipSuccess, "fast_fused_mxfp4_mfma_kernel launch failed: ", hipGetErrorString(err));
}

__global__ __launch_bounds__(256) void fast_quant_mxfp4_kernel(
    const hip_bfloat16* __restrict__ a_bf16,
    uint8_t* __restrict__ q_out,
    uint8_t* __restrict__ scale_sh_out,
    int real_m,
    int k,
    int a_stride,
    int q_stride,
    int scale_stride,
    int scale_n_valid,
    int scale_n_pad) {
    const int block32_per_row = k / 32;
    const int total_tasks = real_m * block32_per_row;
    const int tid = static_cast<int>(blockIdx.x) * static_cast<int>(blockDim.x) + static_cast<int>(threadIdx.x);
    const int stride = static_cast<int>(blockDim.x) * static_cast<int>(gridDim.x);

    for (int task = tid; task < total_tasks; task += stride) {
        const int row = task / block32_per_row;
        const int blk = task % block32_per_row;
        const int k_base = blk * 32;

        float vals[32];
        float max_abs = 0.0f;
        #pragma unroll
        for (int i = 0; i < 32; ++i) {
            const float v = static_cast<float>(a_bf16[static_cast<int64_t>(row) * a_stride + k_base + i]);
            vals[i] = v;
            max_abs = fmaxf(max_abs, fabsf(v));
        }

        const uint8_t scale_byte = amax_to_scale_e8m0(max_abs);
        const float inv_scale = scale_byte != 0 ? (1.0f / decode_e8m0(scale_byte)) : 0.0f;

        uint32_t packed[4] = {0u, 0u, 0u, 0u};
        #pragma unroll
        for (int pack_idx = 0; pack_idx < 4; ++pack_idx) {
            uint32_t out = 0u;
            #pragma unroll
            for (int j = 0; j < 8; ++j) {
                const uint8_t nibble = f32_to_fp4_bits(vals[pack_idx * 8 + j] * inv_scale);
                out |= static_cast<uint32_t>(nibble) << (j * 4);
            }
            packed[pack_idx] = out;
        }

        uint32_t* q_row_ptr = reinterpret_cast<uint32_t*>(q_out + static_cast<int64_t>(row) * q_stride + blk * 16);
        q_row_ptr[0] = packed[0];
        q_row_ptr[1] = packed[1];
        q_row_ptr[2] = packed[2];
        q_row_ptr[3] = packed[3];

        if (blk < scale_n_valid) {
            const int bs_offs_0 = row / 32;
            const int bs_offs_1 = (row % 32) / 16;
            const int bs_offs_2 = row % 16;
            const int bs_offs_3 = blk / 8;
            const int bs_offs_4 = (blk % 8) / 4;
            const int bs_offs_5 = blk % 4;
            const int flat = (
                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)
            );
            scale_sh_out[flat] = scale_byte;
        }
    }
}

void launch_fast_quant_mxfp4(
    torch::Tensor a_bf16,
    torch::Tensor q_out,
    torch::Tensor scale_sh_out,
    int real_m,
    int k) {
    TORCH_CHECK(a_bf16.is_cuda(), "a_bf16 must be a HIP tensor");
    TORCH_CHECK(q_out.is_cuda(), "q_out must be a HIP tensor");
    TORCH_CHECK(scale_sh_out.is_cuda(), "scale_sh_out must be a HIP tensor");

    const int scale_n_valid = (k + 31) / 32;

    int threads = 256;
    int blocks = 1;
    pick_quant_launch_config(real_m, k, &threads, &blocks);

    hipLaunchKernelGGL(
        HIP_KERNEL_NAME(fast_quant_mxfp4_kernel),
        dim3(static_cast<unsigned int>(blocks), 1, 1),
        dim3(static_cast<unsigned int>(threads), 1, 1),
        0,
        0,
        reinterpret_cast<const hip_bfloat16*>(a_bf16.data_ptr()),
        reinterpret_cast<uint8_t*>(q_out.data_ptr()),
        reinterpret_cast<uint8_t*>(scale_sh_out.data_ptr()),
        real_m,
        k,
        static_cast<int>(a_bf16.stride(0)),
        static_cast<int>(q_out.stride(0)),
        static_cast<int>(scale_sh_out.stride(0)),
        scale_n_valid,
        static_cast<int>(scale_sh_out.stride(0)));

    const hipError_t err = hipGetLastError();
    TORCH_CHECK(err == hipSuccess, "fast_quant_mxfp4_kernel launch failed: ", hipGetErrorString(err));
}

void launch_fast_quant_and_f4gemm(
    torch::Tensor a_bf16,
    torch::Tensor q_out,
    torch::Tensor scale_sh_out,
    torch::Tensor b_shuffle,
    torch::Tensor b_scale_sh,
    torch::Tensor out,
    std::string co_path,
    std::string kernel_name,
    int tile_m,
    int tile_n,
    int log2_k_split,
    int real_m,
    int k) {
    TORCH_CHECK(a_bf16.is_cuda(), "a_bf16 must be a HIP tensor");
    TORCH_CHECK(q_out.is_cuda(), "q_out must be a HIP tensor");
    TORCH_CHECK(scale_sh_out.is_cuda(), "scale_sh_out must be a HIP tensor");
    TORCH_CHECK(b_shuffle.is_cuda(), "b_shuffle must be a HIP tensor");
    TORCH_CHECK(b_scale_sh.is_cuda(), "b_scale_sh must be a HIP tensor");
    TORCH_CHECK(out.is_cuda(), "out must be a HIP tensor");

    const int scale_n_valid = (k + 31) / 32;

    int threads = 256;
    int blocks = 1;
    pick_quant_launch_config(real_m, k, &threads, &blocks);

    hipLaunchKernelGGL(
        HIP_KERNEL_NAME(fast_quant_mxfp4_kernel),
        dim3(static_cast<unsigned int>(blocks), 1, 1),
        dim3(static_cast<unsigned int>(threads), 1, 1),
        0,
        0,
        reinterpret_cast<const hip_bfloat16*>(a_bf16.data_ptr()),
        reinterpret_cast<uint8_t*>(q_out.data_ptr()),
        reinterpret_cast<uint8_t*>(scale_sh_out.data_ptr()),
        real_m,
        k,
        static_cast<int>(a_bf16.stride(0)),
        static_cast<int>(q_out.stride(0)),
        static_cast<int>(scale_sh_out.stride(0)),
        scale_n_valid,
        static_cast<int>(scale_sh_out.stride(0)));
    hip_check(hipGetLastError(), "fast_quant_mxfp4_kernel launch");

    auto& kernel = get_fast_kernel(co_path, kernel_name);

    KernelArgs args{};
    const int m = static_cast<int>(out.size(0));
    const int n = static_cast<int>(out.size(1));

    args.ptr_D = out.data_ptr();
    args.ptr_C = nullptr;
    args.ptr_A = q_out.data_ptr();
    args.ptr_B = b_shuffle.data_ptr();
    args.alpha = 1.0f;
    args.beta = 0.0f;
    args.stride_C0 = static_cast<uint32_t>(out.stride(0));
    args.stride_A0 = static_cast<uint32_t>(q_out.stride(0) * 2);
    args.stride_B0 = static_cast<uint32_t>(b_shuffle.stride(0) * 2);
    args.Mdim = static_cast<uint32_t>(m);
    args.Ndim = static_cast<uint32_t>(n);
    args.Kdim = static_cast<uint32_t>(k);
    args.ptr_ScaleA = scale_sh_out.data_ptr();
    args.ptr_ScaleB = b_scale_sh.data_ptr();
    args.stride_ScaleA0 = static_cast<uint32_t>(scale_sh_out.stride(0));
    args.stride_ScaleB0 = static_cast<uint32_t>(b_scale_sh.stride(0));
    args.log2_k_split = 0;

    int gdz = 1;
    if (log2_k_split > 0) {
        args.log2_k_split = log2_k_split;
        const int split_k = 1 << args.log2_k_split;
        TORCH_CHECK(k % split_k == 0, "K must be divisible by split-K factor");
        if (split_k > 1) {
            hip_check(
                hipMemsetAsync(out.data_ptr(), 0, out.numel() * out.element_size(), nullptr),
                "hipMemsetAsync");
        }
        const int k_per_tg = ((k / split_k + 255) / 256) * 256;
        gdz = (k + k_per_tg - 1) / k_per_tg;
    }

    size_t arg_size = sizeof(args);
    void* config[] = {
        HIP_LAUNCH_PARAM_BUFFER_POINTER,
        &args,
        HIP_LAUNCH_PARAM_BUFFER_SIZE,
        &arg_size,
        HIP_LAUNCH_PARAM_END,
    };

    hip_check(
        hipModuleLaunchKernel(
            kernel.func,
            static_cast<unsigned int>((n + tile_n - 1) / tile_n),
            static_cast<unsigned int>((m + tile_m - 1) / tile_m),
            static_cast<unsigned int>(gdz),
            256,
            1,
            1,
            0,
            nullptr,
            nullptr,
            reinterpret_cast<void**>(config)),
        "hipModuleLaunchKernel");
}
"""

_CPP_SRC = r"""
#include <torch/extension.h>
#include <c10/hip/HIPFunctions.h>
#include <hip/hip_runtime.h>

#include <cstdint>
#include <mutex>
#include <stdexcept>
#include <string>
#include <unordered_map>

namespace {

struct p3 {
    uint32_t x;
    uint32_t y;
    uint32_t z;
};

struct p2 {
    uint32_t x;
    uint32_t y;
};

struct __attribute__((packed)) KernelArgs {
    void* ptr_D;
    p2 _p0;
    void* ptr_C;
    p2 _p1;
    void* ptr_A;
    p2 _p2;
    void* ptr_B;
    p2 _p3;
    float alpha;
    p3 _p4;
    float beta;
    p3 _p5;
    uint32_t stride_D0;
    p3 _p6;
    uint32_t stride_D1;
    p3 _p7;
    uint32_t stride_C0;
    p3 _p8;
    uint32_t stride_C1;
    p3 _p9;
    uint32_t stride_A0;
    p3 _p10;
    uint32_t stride_A1;
    p3 _p11;
    uint32_t stride_B0;
    p3 _p12;
    uint32_t stride_B1;
    p3 _p13;
    uint32_t Mdim;
    p3 _p14;
    uint32_t Ndim;
    p3 _p15;
    uint32_t Kdim;
    p3 _p16;
    void* ptr_ScaleA;
    p2 _p17;
    void* ptr_ScaleB;
    p2 _p18;
    uint32_t stride_ScaleA0;
    p3 _p19;
    uint32_t stride_ScaleA1;
    p3 _p20;
    uint32_t stride_ScaleB0;
    p3 _p21;
    uint32_t stride_ScaleB1;
    p3 _p22;
    int32_t log2_k_split;
    p3 _p23;
};

static_assert(sizeof(KernelArgs) == 384, "Unexpected FP4 kernarg size");

struct CachedKernel {
    hipModule_t module = nullptr;
    hipFunction_t func = nullptr;
};

std::unordered_map<std::string, CachedKernel>& kernel_cache() {
    static std::unordered_map<std::string, CachedKernel> cache;
    return cache;
}

std::mutex& kernel_cache_mutex() {
    static std::mutex mu;
    return mu;
}

void hip_check(hipError_t err, const char* call_name) {
    if (err == hipSuccess) {
        return;
    }
    throw std::runtime_error(std::string(call_name) + " failed: " + hipGetErrorString(err));
}

CachedKernel& get_kernel(const std::string& co_path, const std::string& kernel_name) {
    const std::string key = co_path + "|" + kernel_name;
    std::lock_guard<std::mutex> guard(kernel_cache_mutex());
    auto& cache = kernel_cache();
    auto it = cache.find(key);
    if (it != cache.end()) {
        return it->second;
    }

    CachedKernel entry;
    hip_check(hipModuleLoad(&entry.module, co_path.c_str()), "hipModuleLoad");
    hip_check(hipModuleGetFunction(&entry.func, entry.module, kernel_name.c_str()), "hipModuleGetFunction");
    auto [new_it, _inserted] = cache.emplace(key, entry);
    return new_it->second;
}

}  // namespace

void launch_f4gemm(
    torch::Tensor a_q,
    torch::Tensor b_shuffle,
    torch::Tensor a_scale_sh,
    torch::Tensor b_scale_sh,
    torch::Tensor out,
    std::string co_path,
    std::string kernel_name,
    int tile_m,
    int tile_n,
    int log2_k_split) {
    TORCH_CHECK(a_q.is_cuda(), "a_q must be a HIP tensor");
    TORCH_CHECK(b_shuffle.is_cuda(), "b_shuffle must be a HIP tensor");
    TORCH_CHECK(a_scale_sh.is_cuda(), "a_scale_sh must be a HIP tensor");
    TORCH_CHECK(b_scale_sh.is_cuda(), "b_scale_sh must be a HIP tensor");
    TORCH_CHECK(out.is_cuda(), "out must be a HIP tensor");

    auto& kernel = get_kernel(co_path, kernel_name);

    KernelArgs args{};
    const int m = static_cast<int>(out.size(0));
    const int n = static_cast<int>(out.size(1));
    const int k = static_cast<int>(a_q.size(1) * 2);

    // Match aiter's asm_gemm_a4w4 host-side argument population exactly.
    args.ptr_D = out.data_ptr();
    args.ptr_C = nullptr;
    args.ptr_A = a_q.data_ptr();
    args.ptr_B = b_shuffle.data_ptr();
    args.alpha = 1.0f;
    args.beta = 0.0f;
    args.stride_C0 = static_cast<uint32_t>(out.stride(0));
    args.stride_A0 = static_cast<uint32_t>(a_q.stride(0) * 2);
    args.stride_B0 = static_cast<uint32_t>(b_shuffle.stride(0) * 2);
    args.Mdim = static_cast<uint32_t>(m);
    args.Ndim = static_cast<uint32_t>(n);
    args.Kdim = static_cast<uint32_t>(k);
    args.ptr_ScaleA = a_scale_sh.data_ptr();
    args.ptr_ScaleB = b_scale_sh.data_ptr();
    args.stride_ScaleA0 = static_cast<uint32_t>(a_scale_sh.stride(0));
    args.stride_ScaleB0 = static_cast<uint32_t>(b_scale_sh.stride(0));
    args.log2_k_split = 0;

    int gdz = 1;
    if (log2_k_split > 0) {
        args.log2_k_split = log2_k_split;
        const int split_k = 1 << args.log2_k_split;
        TORCH_CHECK(k % split_k == 0, "K must be divisible by split-K factor");
        if (split_k > 1) {
            hip_check(
                hipMemsetAsync(out.data_ptr(), 0, out.numel() * out.element_size(), nullptr),
                "hipMemsetAsync");
        }
        const int k_per_tg = ((k / split_k + 255) / 256) * 256;
        gdz = (k + k_per_tg - 1) / k_per_tg;
    }

    size_t arg_size = sizeof(args);
    void* config[] = {
        HIP_LAUNCH_PARAM_BUFFER_POINTER,
        &args,
        HIP_LAUNCH_PARAM_BUFFER_SIZE,
        &arg_size,
        HIP_LAUNCH_PARAM_END,
    };

    hip_check(
        hipModuleLaunchKernel(
            kernel.func,
            static_cast<unsigned int>((n + tile_n - 1) / tile_n),
            static_cast<unsigned int>((m + tile_m - 1) / tile_m),
            static_cast<unsigned int>(gdz),
            256,
            1,
            1,
            0,
            nullptr,
            nullptr,
            reinterpret_cast<void**>(config)),
        "hipModuleLaunchKernel");
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("launch_f4gemm", &launch_f4gemm, "Launch gfx950 FP4 GEMM code object");
}
"""

@lru_cache(maxsize=1)
def _launcher():
    return _load_inline_with_trace(
        "launcher",
        name="mxfp4_gfx950_launcher_ext",
        cpp_sources=[_CPP_SRC],
        functions=None,
        extra_cflags=["-std=c++17", "-I/opt/rocm/include"],
        extra_ldflags=["-lamdhip64", "-L/opt/rocm/lib"],
        with_cuda=False,
        verbose=False,
    )

@lru_cache(maxsize=1)
def _fast_fused_kernel():
    return _load_inline_with_trace(
        "fast-fused",
        name="mxfp4_gfx950_fast_fused_ext",
        cpp_sources=[_FAST_FUSED_CPP_SRC],
        cuda_sources=[_FAST_FUSED_CUDA_SRC],
        functions=[
            "launch_fast_quant_mxfp4",
            "launch_fast_quant_and_f4gemm",
            "launch_fast_fused_mxfp4_gemm",
        ],
        extra_cflags=["-std=c++20", "-I/opt/rocm/include"],
        extra_cuda_cflags=_fused_cuda_cflags(),
        with_cuda=True,
        no_implicit_headers=True,
        verbose=False,
    )


@lru_cache(maxsize=1)
def _aiter_hsa_root() -> Path:
    env_root = os.environ.get("AITER_ASM_DIR")
    if env_root:
        root = Path(env_root)
        if root.exists():
            return root

    import aiter  # type: ignore

    pkg_file = Path(aiter.__file__).resolve()
    candidates = [
        pkg_file.parents[1] / "hsa",
        pkg_file.parents[0] / "hsa",
        Path("/home/runner/aiter/hsa"),
    ]
    for root in candidates:
        if (root / _ARCH / _F4GEMM_SUBDIR).exists():
            return root
    raise RuntimeError("Unable to locate aiter HSA directory for gfx950 FP4 kernels")


def _fused_cuda_cflags() -> list[str]:
    return [
        "-O3",
        "--offload-arch=gfx950",
        "-std=c++20",
        "-U__HIP_NO_HALF_OPERATORS__",
        "-U__HIP_NO_HALF_CONVERSIONS__",
    ]


@lru_cache(maxsize=1)
def _load_f4gemm_cfgs() -> tuple[_KernelCfg, ...]:
    csv_path = _aiter_hsa_root() / _ARCH / _F4GEMM_SUBDIR / "f4gemm_bf16_per1x32Fp4.csv"
    cfgs: list[_KernelCfg] = []
    with csv_path.open("r", encoding="utf-8") as f:
        for row in csv.DictReader(f):
            cfgs.append(
                _KernelCfg(
                    tile_m=int(row["tile_M"]),
                    tile_n=int(row["tile_N"]),
                    splitk_enabled=int(row["splitK"]),
                    bpreshuffle=int(row["bpreshuffle"]),
                    kernel_name=row["knl_name"],
                    co_name=f"{_F4GEMM_SUBDIR}/{row['co_name']}",
                )
            )
    return tuple(cfgs)


@lru_cache(maxsize=1)
def _aiter_config_root() -> Path:
    import aiter  # type: ignore

    pkg_file = Path(aiter.__file__).resolve()
    candidates = [
        pkg_file.parents[0] / "configs",
        pkg_file.parents[1] / "aiter" / "configs",
        Path("/home/runner/aiter/aiter/configs"),
    ]
    for root in candidates:
        if (root / "a4w4_blockscale_tuned_gemm.csv").exists():
            return root
    raise RuntimeError("Unable to locate aiter tuned GEMM config directory")


@lru_cache(maxsize=1)
def _load_tuned_a4w4_cfgs() -> tuple[_TunedKernel, ...]:
    csv_path = _aiter_config_root() / "a4w4_blockscale_tuned_gemm.csv"
    cfgs: list[_TunedKernel] = []
    with csv_path.open("r", encoding="utf-8") as f:
        for row in csv.DictReader(f):
            kernel_name = row["kernelName"]
            if not kernel_name.startswith("_ZN"):
                continue
            cfgs.append(
                _TunedKernel(
                    cu_num=int(row["cu_num"]),
                    m=int(row["M"]),
                    n=int(row["N"]),
                    k=int(row["K"]),
                    split_k=int(row["splitK"]),
                    kernel_name=kernel_name,
                )
            )
    return tuple(cfgs)


def _quant_mxfp4(x: torch.Tensor, *, shuffle_scale: bool) -> tuple[torch.Tensor, torch.Tensor]:
    from aiter import dtypes  # type: ignore
    from aiter.ops.triton.quant import dynamic_mxfp4_quant  # type: ignore
    from aiter.utility.fp4_utils import e8m0_shuffle  # type: ignore

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


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


def _pad_rows_u8(x: torch.Tensor, rows: int, fill: int) -> torch.Tensor:
    if int(x.shape[0]) >= rows:
        return x.contiguous()
    padded = torch.full((rows, *x.shape[1:]), fill, dtype=torch.uint8, device=x.device)
    padded[: int(x.shape[0])].copy_(x.contiguous())
    return padded


def _prepare_b_fast_inputs(
    b_shuffle: torch.Tensor,
    b_scale_sh: torch.Tensor,
    *,
    n: int,
) -> tuple[torch.Tensor, torch.Tensor, int]:
    cache_key = (
        int(b_shuffle.data_ptr()),
        int(b_scale_sh.data_ptr()),
        n,
        int(b_shuffle.shape[1]),
        int(b_scale_sh.shape[1]),
    )
    cached = _B_SHUFFLE_CACHE.get(cache_key)
    if cached is not None:
        return cached

    b_shuffle_u8_src = _as_u8_view(b_shuffle).contiguous()
    b_scale_sh_u8_src = _as_u8_view(b_scale_sh).contiguous()

    # b_shuffle is N-major and must be padded to the fast-path tile width.
    # b_scale_sh is already in preshuffled microscale layout and may have a
    # larger leading dimension than N because aiter pads it independently.
    n_padded = max(
        ((n + _CUSTOM_N_TILE - 1) // _CUSTOM_N_TILE) * _CUSTOM_N_TILE,
        int(b_shuffle_u8_src.shape[0]),
    )
    b_shuffle_u8 = _pad_rows_u8(b_shuffle_u8_src, n_padded, 0)
    b_scale_sh_u8 = b_scale_sh_u8_src

    result = (b_shuffle_u8, b_scale_sh_u8, n_padded)
    _B_SHUFFLE_CACHE[cache_key] = result
    return result


def _pad_a_quant_inputs(
    a_q: torch.Tensor,
    a_scale_sh: torch.Tensor,
    *,
    padded_m: int,
) -> tuple[torch.Tensor, torch.Tensor]:
    if int(a_q.shape[0]) >= padded_m:
        return a_q.contiguous(), a_scale_sh.contiguous()

    a_q_u8 = _as_u8_view(a_q).contiguous()
    a_q_pad = torch.zeros((padded_m, int(a_q_u8.shape[1])), dtype=torch.uint8, device=a_q.device)
    a_q_pad[: int(a_q.shape[0])].copy_(a_q_u8)

    if int(a_scale_sh.shape[0]) >= padded_m:
        return a_q_pad.view(a_q.dtype), a_scale_sh.contiguous()

    a_scale_u8 = _as_u8_view(a_scale_sh).contiguous()
    a_scale_pad = torch.full(
        (padded_m, int(a_scale_u8.shape[1])),
        127,
        dtype=torch.uint8,
        device=a_scale_sh.device,
    )
    a_scale_pad[: int(a_scale_sh.shape[0])].copy_(a_scale_u8)
    return a_q_pad.view(a_q.dtype), a_scale_pad.view(a_scale_sh.dtype)


@lru_cache(maxsize=1)
def _direct_shuffled_quant_kernel():
    import triton
    import triton.language as tl
    from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op  # type: ignore

    @triton.heuristics(
        {
            "EVEN_M_N": lambda args: args["M"] % args["BLOCK_SIZE_M"] == 0
            and args["N"] % (args["BLOCK_SIZE_N"] * args["NUM_ITER"]) == 0,
        }
    )
    @triton.jit
    def _kernel(
        x_ptr,
        x_fp4_ptr,
        scale_sh_ptr,
        stride_x_m_in,
        stride_x_n_in,
        stride_x_fp4_m_in,
        stride_x_fp4_n_in,
        M,
        N,
        NUM_ITER: tl.constexpr,
        NUM_STAGES: tl.constexpr,
        SCALE_N_VALID: tl.constexpr,
        SCALE_N_PAD: tl.constexpr,
        BLOCK_SIZE_M: tl.constexpr,
        BLOCK_SIZE_N: tl.constexpr,
        MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
        EVEN_M_N: tl.constexpr,
        SCALING_MODE: tl.constexpr,
    ):
        pid_m = tl.program_id(0)
        start_n = tl.program_id(1) * NUM_ITER

        stride_x_m = tl.cast(stride_x_m_in, tl.int64)
        stride_x_n = tl.cast(stride_x_n_in, tl.int64)
        stride_x_fp4_m = tl.cast(stride_x_fp4_m_in, tl.int64)
        stride_x_fp4_n = tl.cast(stride_x_fp4_n_in, tl.int64)

        num_quant_blocks: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE

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

            if EVEN_M_N:
                x = tl.load(x_ptr + x_offs, cache_modifier=".cg").to(tl.float32)
            else:
                x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
                x = tl.load(x_ptr + x_offs, mask=x_mask, cache_modifier=".cg").to(tl.float32)

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

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

            if EVEN_M_N:
                tl.store(x_fp4_ptr + out_offs, out_tensor)
            else:
                out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :]
                tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask)

            bs_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
            bs_offs_n = pid_n * num_quant_blocks + tl.arange(0, num_quant_blocks)

            bs_offs_0 = bs_offs_m[:, None] // 32
            bs_offs_1 = (bs_offs_m[:, None] % 32) // 16
            bs_offs_2 = bs_offs_m[:, None] % 16
            bs_offs_3 = bs_offs_n[None, :] // 8
            bs_offs_4 = (bs_offs_n[None, :] % 8) // 4
            bs_offs_5 = bs_offs_n[None, :] % 4

            bs_flat_offs = (
                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)
            )
            bs_mask = (bs_offs_m < M)[:, None] & (bs_offs_n < SCALE_N_VALID)[None, :]
            tl.store(scale_sh_ptr + bs_flat_offs, bs_e8m0, mask=bs_mask)

    return _kernel


def _quant_mxfp4_direct_shuffled(
    x: torch.Tensor,
    *,
    x_fp4_out: torch.Tensor | None = None,
    scale_sh_out: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
    import triton
    from aiter import dtypes  # type: ignore

    m, n = map(int, x.shape)
    scale_n_valid = (n + 31) // 32
    scale_n_pad = ((scale_n_valid + 7) // 8) * 8
    scale_m_pad = ((m + 255) // 256) * 256

    if x_fp4_out is None:
        x_fp4 = torch.empty((m, n // 2), dtype=torch.uint8, device=x.device)
    else:
        x_fp4 = x_fp4_out
        if int(x_fp4.shape[0]) != m:
            x_fp4.zero_()
    if scale_sh_out is None:
        scale_sh = torch.full((scale_m_pad, scale_n_pad), 127, dtype=torch.uint8, device=x.device)
    else:
        scale_sh = scale_sh_out
        scale_sh.fill_(127)

    if m <= _CUSTOM_SMALLM_MAX_M:
        num_iter = 1
        block_size_m = triton.next_power_of_2(m)
        block_size_n = 32
        num_warps = 1
        num_stages = 1
    else:
        num_iter = 4
        block_size_m = 64
        block_size_n = 64
        num_warps = 4
        num_stages = 2
        if n <= 16384:
            block_size_m = 32
            block_size_n = 128

    if n <= 1024:
        num_iter = 1
        num_stages = 1
        num_warps = 4
        block_size_n = min(256, triton.next_power_of_2(n))
        block_size_n = max(32, block_size_n)
        block_size_m = min(8, triton.next_power_of_2(m))

    grid = (
        triton.cdiv(m, block_size_m),
        triton.cdiv(n, block_size_n * num_iter),
    )

    _direct_shuffled_quant_kernel()[grid](
        x,
        x_fp4,
        scale_sh,
        *x.stride(),
        *x_fp4.stride(),
        M=m,
        N=n,
        NUM_ITER=num_iter,
        NUM_STAGES=num_stages,
        SCALE_N_VALID=scale_n_valid,
        SCALE_N_PAD=scale_n_pad,
        BLOCK_SIZE_M=block_size_m,
        BLOCK_SIZE_N=block_size_n,
        MXFP4_QUANT_BLOCK_SIZE=32,
        SCALING_MODE=0,
        num_warps=num_warps,
        num_stages=1,
        waves_per_eu=0,
    )
    return x_fp4.view(dtypes.fp4x2), scale_sh.view(dtypes.fp8_e8m0)


def _quant_mxfp4_direct_shuffled_k512_smallm(
    x: torch.Tensor,
    *,
    x_fp4_out: torch.Tensor | None = None,
    scale_sh_out: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
    import triton
    from aiter import dtypes  # type: ignore

    m, n = map(int, x.shape)
    if n != 512 or m > _CUSTOM_SMALLM_MAX_M:
        raise RuntimeError(f"k=512 small-shape direct quant does not support shape {(m, n)}")

    scale_n_pad = 16
    scale_m_pad = ((m + 255) // 256) * 256

    if x_fp4_out is None:
        x_fp4 = torch.empty((m, n // 2), dtype=torch.uint8, device=x.device)
    else:
        x_fp4 = x_fp4_out
        if int(x_fp4.shape[0]) != m:
            x_fp4.zero_()
    if scale_sh_out is None:
        scale_sh = torch.full((scale_m_pad, scale_n_pad), 127, dtype=torch.uint8, device=x.device)
    else:
        scale_sh = scale_sh_out
        scale_sh.fill_(127)

    block_size_m = min(8, triton.next_power_of_2(m))
    num_warps = 4
    block_size_n = 256
    num_iter = 2
    grid = (
        triton.cdiv(m, block_size_m),
        triton.cdiv(n, block_size_n * num_iter),
    )
    _direct_shuffled_quant_kernel()[grid](
        x,
        x_fp4,
        scale_sh,
        *x.stride(),
        *x_fp4.stride(),
        M=m,
        N=n,
        NUM_ITER=num_iter,
        NUM_STAGES=1,
        SCALE_N_VALID=16,
        SCALE_N_PAD=16,
        BLOCK_SIZE_M=block_size_m,
        BLOCK_SIZE_N=block_size_n,
        MXFP4_QUANT_BLOCK_SIZE=32,
        SCALING_MODE=0,
        num_warps=num_warps,
        num_stages=1,
        waves_per_eu=0,
    )
    return x_fp4.view(dtypes.fp4x2), scale_sh.view(dtypes.fp8_e8m0)


def _quant_mxfp4_direct_shuffled_largek_smallm(
    x: torch.Tensor,
    *,
    x_fp4_out: torch.Tensor | None = None,
    scale_sh_out: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
    from aiter import dtypes  # type: ignore

    m, n = map(int, x.shape)
    if (m, n) != (16, 7168):
        raise RuntimeError(f"large-k small-m direct quant does not support shape {(m, n)}")

    scale_n_valid = (n + 31) // 32
    scale_n_pad = ((scale_n_valid + 7) // 8) * 8
    scale_m_pad = ((m + 255) // 256) * 256

    if x_fp4_out is None:
        x_fp4 = torch.empty((m, n // 2), dtype=torch.uint8, device=x.device)
    else:
        x_fp4 = x_fp4_out
        if int(x_fp4.shape[0]) != m:
            x_fp4.zero_()
    if scale_sh_out is None:
        scale_sh = torch.full((scale_m_pad, scale_n_pad), 127, dtype=torch.uint8, device=x.device)
    else:
        scale_sh = scale_sh_out
        scale_sh.fill_(127)

    grid = (
        1,
        (n + 511) // 512,
    )
    _direct_shuffled_quant_kernel()[grid](
        x,
        x_fp4,
        scale_sh,
        *x.stride(),
        *x_fp4.stride(),
        M=m,
        N=n,
        NUM_ITER=2,
        NUM_STAGES=2,
        SCALE_N_VALID=scale_n_valid,
        SCALE_N_PAD=scale_n_pad,
        BLOCK_SIZE_M=16,
        BLOCK_SIZE_N=256,
        MXFP4_QUANT_BLOCK_SIZE=32,
        SCALING_MODE=0,
        num_warps=4,
        num_stages=1,
        waves_per_eu=0,
    )
    return x_fp4.view(dtypes.fp4x2), scale_sh.view(dtypes.fp8_e8m0)


def _quant_mxfp4_hip_direct_shuffled(
    x: torch.Tensor,
    *,
    x_fp4_out: torch.Tensor | None = None,
    scale_sh_out: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
    from aiter import dtypes  # type: ignore

    m, n = map(int, x.shape)
    scale_n_valid = (n + 31) // 32
    scale_n_pad = ((scale_n_valid + 7) // 8) * 8
    scale_m_pad = ((m + 255) // 256) * 256

    if x_fp4_out is None:
        x_fp4 = torch.empty((m, n // 2), dtype=torch.uint8, device=x.device)
    else:
        x_fp4 = x_fp4_out
        if int(x_fp4.shape[0]) != m:
            x_fp4.zero_()
    if scale_sh_out is None:
        scale_sh = torch.full((scale_m_pad, scale_n_pad), 127, dtype=torch.uint8, device=x.device)
    else:
        scale_sh = scale_sh_out
        scale_sh.fill_(127)

    _fast_fused_kernel().launch_fast_quant_mxfp4(x, x_fp4, scale_sh, int(m), int(n))
    return x_fp4.view(dtypes.fp4x2), scale_sh.view(dtypes.fp8_e8m0)


def _quant_mxfp4_aiter_hip(
    x: torch.Tensor,
    *,
    x_fp4_out: torch.Tensor | None = None,
    scale_sh_out: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
    from aiter import dtypes  # type: ignore
    from aiter.ops.quant import dynamic_per_group_scaled_quant_fp4  # type: ignore

    m, n = map(int, x.shape)
    scale_n_valid = (n + 31) // 32
    scale_n_pad = ((scale_n_valid + 7) // 8) * 8
    scale_m_pad = ((m + 255) // 256) * 256

    if x_fp4_out is None:
        x_fp4 = torch.empty((m, n // 2), dtype=torch.uint8, device=x.device)
    else:
        x_fp4 = x_fp4_out
        if int(x_fp4.shape[0]) != m:
            x_fp4.zero_()
    if scale_sh_out is None:
        scale_sh = torch.empty((scale_m_pad, scale_n_pad), dtype=torch.uint8, device=x.device)
    else:
        scale_sh = scale_sh_out

    dynamic_per_group_scaled_quant_fp4(
        x_fp4.view(dtypes.fp4x2),
        x,
        scale_sh.view(dtypes.fp8_e8m0),
        32,
        shuffle_scale=True,
    )
    return x_fp4.view(dtypes.fp4x2), scale_sh.view(dtypes.fp8_e8m0)


def _device_index(device: torch.device) -> int:
    return -1 if device.index is None else int(device.index)


def _lookup_benchmark_quant_workspace(
    *,
    device: torch.device,
    m: int,
    n: int,
    q_rows: int,
) -> tuple[torch.Tensor, torch.Tensor] | None:
    device_type = device.type
    device_index = _device_index(device)
    for idx, entry in enumerate(_BENCHMARK_QUANT_WORKSPACES):
        if (
            entry.device_type != device_type
            or entry.device_index != device_index
            or entry.m != m
            or entry.n != n
            or entry.q_rows != q_rows
        ):
            continue
        if idx:
            _BENCHMARK_QUANT_WORKSPACES.insert(0, _BENCHMARK_QUANT_WORKSPACES.pop(idx))
            entry = _BENCHMARK_QUANT_WORKSPACES[0]
        return entry.q, entry.scale_sh
    return None


def _get_benchmark_quant_workspace(
    *,
    device: torch.device,
    m: int,
    n: int,
    q_rows: int | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
    q_rows = m if q_rows is None else q_rows
    cached = _lookup_benchmark_quant_workspace(device=device, m=m, n=n, q_rows=q_rows)
    if cached is not None:
        return cached

    scale_n_valid = (n + 31) // 32
    scale_n_pad = ((scale_n_valid + 7) // 8) * 8
    scale_m_pad = ((max(m, q_rows) + 255) // 256) * 256
    q = torch.empty((q_rows, n // 2), dtype=torch.uint8, device=device)
    scale_sh = torch.empty((scale_m_pad, scale_n_pad), dtype=torch.uint8, device=device)
    _BENCHMARK_QUANT_WORKSPACES.insert(
        0,
        _QuantWorkspaceEntry(
            device_type=device.type,
            device_index=_device_index(device),
            m=m,
            n=n,
            q_rows=q_rows,
            q=q,
            scale_sh=scale_sh,
        ),
    )
    del _BENCHMARK_QUANT_WORKSPACES[_BENCHMARK_QUANT_WORKSPACE_MAX:]
    return q, scale_sh


def _compute_benchmark_quant(
    x: torch.Tensor,
    *,
    shape: tuple[int, int, int],
    q_out: torch.Tensor,
    scale_sh_out: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
    if shape in {(64, 7168, 2048), (256, 3072, 1536)}:
        return _quant_mxfp4_hip_direct_shuffled(
            x,
            x_fp4_out=q_out,
            scale_sh_out=scale_sh_out,
        )
    if shape == (16, 2112, 7168):
        return _quant_mxfp4_direct_shuffled(
            x,
            x_fp4_out=q_out,
            scale_sh_out=scale_sh_out,
        )
    if _is_smallk_benchmark_shape(shape):
        return _quant_mxfp4_direct_shuffled_k512_smallm(
            x,
            x_fp4_out=q_out,
            scale_sh_out=scale_sh_out,
        )
    return _quant_mxfp4_direct_shuffled(
        x,
        x_fp4_out=q_out,
        scale_sh_out=scale_sh_out,
    )


def _lookup_benchmark_out_cache(
    *,
    device: torch.device,
    dtype: torch.dtype,
    padded_m: int,
    n: int,
) -> torch.Tensor | None:
    device_type = device.type
    device_index = _device_index(device)
    for idx, entry in enumerate(_BENCHMARK_OUT_CACHE):
        if (
            entry.device_type != device_type
            or entry.device_index != device_index
            or entry.dtype != dtype
            or entry.padded_m != padded_m
            or entry.n != n
        ):
            continue
        if idx:
            _BENCHMARK_OUT_CACHE.insert(0, _BENCHMARK_OUT_CACHE.pop(idx))
            entry = _BENCHMARK_OUT_CACHE[0]
        return entry.out
    return None


def _get_cached_benchmark_out(
    *,
    device: torch.device,
    dtype: torch.dtype,
    padded_m: int,
    n: int,
) -> torch.Tensor:
    cached = _lookup_benchmark_out_cache(
        device=device,
        dtype=dtype,
        padded_m=padded_m,
        n=n,
    )
    if cached is not None:
        return cached
    out = torch.empty((padded_m, n), dtype=dtype, device=device)
    _BENCHMARK_OUT_CACHE.insert(
        0,
        _OutCacheEntry(
            device_type=device.type,
            device_index=_device_index(device),
            dtype=dtype,
            padded_m=padded_m,
            n=n,
            out=out,
        ),
    )
    del _BENCHMARK_OUT_CACHE[_BENCHMARK_OUT_CACHE_MAX:]
    return out


def _maybe_contiguous(x: torch.Tensor) -> torch.Tensor:
    return x if x.is_contiguous() else x.contiguous()


def _use_fast_fused_path(m: int, n: int, k: int) -> bool:
    if (m, n, k) == (16, 2112, 7168):
        return False
    return _FAST_FUSED_ENABLE and m <= 32 and k >= 1536 and k % 256 == 0 and n > 0


@lru_cache(maxsize=1)
def _aiter_cu_num() -> int:
    from aiter.jit.utils.chip_info import get_cu_num  # type: ignore

    return int(get_cu_num())


def _find_kernel_cfg(kernel_name: str) -> _KernelCfg:
    for cfg in _load_f4gemm_cfgs():
        if cfg.kernel_name == kernel_name:
            return cfg
    raise RuntimeError(f"Kernel {kernel_name} not found in FP4 config table")


def _lookup_tuned_kernel(m: int, n: int, k: int, padded_m: int, cu_num: int) -> tuple[_KernelCfg, int] | None:
    for candidate_m in (m, padded_m):
        for tuned in _load_tuned_a4w4_cfgs():
            if (tuned.cu_num, tuned.m, tuned.n, tuned.k) != (cu_num, candidate_m, n, k):
                continue
            return _find_kernel_cfg(tuned.kernel_name), tuned.split_k
    return None


def _pick_kernel(m: int, n: int, k: int, num_cu: int) -> tuple[_KernelCfg, int]:
    override = _KERNEL_OVERRIDE_BY_SHAPE.get((m, n, k))
    if override is not None:
        kernel_name, co_name, log2_k_split = override
        for cfg in _load_f4gemm_cfgs():
            if cfg.kernel_name == kernel_name and cfg.co_name == co_name:
                return cfg, log2_k_split
        raise RuntimeError(f"Override kernel not found in FP4 config table: {override}")

    padded_m = ((m + 31) // 32) * 32
    tuned = _lookup_tuned_kernel(m, n, k, padded_m, num_cu)
    if tuned is not None:
        return tuned

    empty_cu = num_cu
    best_round = 1 << 30
    best_eff = 1.0
    best_cfg: _KernelCfg | None = None
    for cfg in _load_f4gemm_cfgs():
        if cfg.bpreshuffle != 1:
            continue
        if cfg.tile_m == 128 and cfg.tile_n == 512 and n % cfg.tile_n != 0:
            continue
        tg_num_m = (padded_m + cfg.tile_m - 1) // cfg.tile_m
        tg_num_n = (n + cfg.tile_n - 1) // cfg.tile_n
        tg_num = tg_num_m * tg_num_n
        local_round = (tg_num + num_cu - 1) // num_cu
        local_eff = (cfg.tile_m * cfg.tile_n) / (cfg.tile_m + cfg.tile_n)
        is_earlier_round = local_round < best_round
        is_same_round = local_round == best_round
        has_sufficient_empty_cu = empty_cu > (local_round * num_cu - tg_num)
        has_better_efficiency = local_eff > best_eff
        if is_earlier_round or (is_same_round and (has_sufficient_empty_cu or has_better_efficiency)):
            best_round = local_round
            empty_cu = local_round * num_cu - tg_num
            best_eff = local_eff
            best_cfg = cfg
    if best_cfg is None:
        raise RuntimeError(f"No preshuffled FP4 kernel found for shape {(m, n, k)}")
    return best_cfg, 0


def _launch_f4gemm(
    *,
    a_q: torch.Tensor,
    b_shuffle: torch.Tensor,
    a_scale_sh: torch.Tensor,
    b_scale_sh: torch.Tensor,
    out: torch.Tensor,
    kernel_cfg: _KernelCfg,
    log2_k_split: int,
) -> None:
    co_path = _aiter_hsa_root() / _ARCH / kernel_cfg.co_name
    if not co_path.exists():
        raise RuntimeError(f"Missing code object: {co_path}")

    _launcher().launch_f4gemm(
        a_q,
        b_shuffle,
        a_scale_sh,
        b_scale_sh,
        out,
        str(co_path),
        kernel_cfg.kernel_name,
        int(kernel_cfg.tile_m),
        int(kernel_cfg.tile_n),
        int(log2_k_split),
    )


def _launch_f4gemm_official(
    *,
    a_q: torch.Tensor,
    b_shuffle: torch.Tensor,
    a_scale_sh: torch.Tensor,
    b_scale_sh: torch.Tensor,
    out: torch.Tensor,
    kernel_cfg: _KernelCfg,
    log2_k_split: int,
) -> None:
    from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm  # type: ignore

    gemm_a4w4_asm(
        a_q,
        b_shuffle,
        a_scale_sh,
        b_scale_sh,
        out,
        kernel_cfg.kernel_name,
        None,
        1.0,
        0.0,
        True,
        log2_k_split,
    )


def _launch_fused_fast(
    *,
    a: torch.Tensor,
    b_shuffle: torch.Tensor,
    b_scale_sh: torch.Tensor,
    out: torch.Tensor,
    n: int,
) -> None:
    a_contig = a if a.is_contiguous() else a.contiguous()
    b_shuffle_u8, b_scale_sh_u8, n_padded = _prepare_b_fast_inputs(b_shuffle, b_scale_sh, n=n)
    _fast_fused_kernel().launch_fast_fused_mxfp4_gemm(
        a_contig,
        b_shuffle_u8,
        b_scale_sh_u8,
        out,
        int(n_padded),
        int(n),
    )


def _launch_fast_quant_and_f4gemm(
    *,
    a: torch.Tensor,
    a_q: torch.Tensor,
    a_scale_sh: torch.Tensor,
    b_shuffle: torch.Tensor,
    b_scale_sh: torch.Tensor,
    out: torch.Tensor,
    kernel_cfg: _KernelCfg,
    log2_k_split: int,
) -> None:
    co_path = _aiter_hsa_root() / _ARCH / kernel_cfg.co_name
    if not co_path.exists():
        raise RuntimeError(f"Missing code object: {co_path}")

    _fast_fused_kernel().launch_fast_quant_and_f4gemm(
        a if a.is_contiguous() else a.contiguous(),
        _as_u8_view(a_q).contiguous(),
        _as_u8_view(a_scale_sh).contiguous(),
        b_shuffle,
        b_scale_sh,
        out,
        str(co_path),
        kernel_cfg.kernel_name,
        int(kernel_cfg.tile_m),
        int(kernel_cfg.tile_n),
        int(log2_k_split),
        int(a.shape[0]),
        int(a.shape[1]),
    )


def custom_kernel(data: input_t) -> output_t:
    a, _b, _b_q, b_shuffle, b_scale_sh = data
    a = _maybe_contiguous(a)

    m, k = map(int, a.shape)
    n = int(b_shuffle.shape[0])
    shape = (m, n, k)
    padded_m = ((m + 31) // 32) * 32

    if _is_benchmark_fastpath_shape(shape):
        q_rows = padded_m if shape in _COMBINED_HIP_QUANT_GEMM_SHAPES else m
        a_q_ws, a_scale_sh_ws = _get_benchmark_quant_workspace(
            device=a.device,
            m=m,
            n=k,
            q_rows=q_rows,
        )
        if shape in _COMBINED_HIP_QUANT_GEMM_SHAPES:
            a_q, a_scale_sh = a_q_ws, a_scale_sh_ws
        else:
            a_q, a_scale_sh = _compute_benchmark_quant(
                a,
                shape=shape,
                q_out=a_q_ws,
                scale_sh_out=a_scale_sh_ws,
            )
        out = _get_cached_benchmark_out(
            device=a.device,
            dtype=torch.bfloat16,
            padded_m=((m + 31) // 32) * 32,
            n=n,
        )
    else:
        a_q, a_scale_sh = _quant_mxfp4(a, shuffle_scale=True)
        out = torch.empty((padded_m, n), dtype=torch.bfloat16, device=a.device)

    kernel_cfg, log2_k_split = _pick_kernel(m, n, k, _aiter_cu_num())
    if shape == (16, 2112, 7168) and int(a_q.shape[0]) < padded_m:
        a_q, a_scale_sh = _pad_a_quant_inputs(a_q, a_scale_sh, padded_m=padded_m)

    if shape in _COMBINED_HIP_QUANT_GEMM_SHAPES:
        _launch_fast_quant_and_f4gemm(
            a=a,
            a_q=a_q,
            a_scale_sh=a_scale_sh,
            b_shuffle=b_shuffle,
            b_scale_sh=b_scale_sh,
            out=out,
            kernel_cfg=kernel_cfg,
            log2_k_split=log2_k_split,
        )
        return out[:m]

    use_custom_launcher = (
        _is_benchmark_fastpath_shape(shape)
        and int(a_q.shape[0]) == int(out.shape[0])
        and shape != (16, 2112, 7168)
    )
    if use_custom_launcher:
        _launch_f4gemm(
            a_q=a_q,
            b_shuffle=b_shuffle,
            a_scale_sh=a_scale_sh,
            b_scale_sh=b_scale_sh,
            out=out,
            kernel_cfg=kernel_cfg,
            log2_k_split=log2_k_split,
        )
    else:
        _launch_f4gemm_official(
            a_q=a_q,
            b_shuffle=b_shuffle,
            a_scale_sh=a_scale_sh,
            b_scale_sh=b_scale_sh,
            out=out,
            kernel_cfg=kernel_cfg,
            log2_k_split=log2_k_split,
        )
    return out[:m]
scrolls · 2104 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