Skip to content
KernelIndex
Search⌘K

submission 610306

Zaber · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v66.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-610306?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.9µs
#539 of 1143
2026-03-22

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:b777950aac5507b3bb7f629a5b62cc7f0285e2795c00d9b69c5759e9c6c1b265
license declaredunknown
license concludedunknown
authorsZaber
imported2026-08-26

Techniques

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

split-kint m, k, n, split_k;

Kernel source

submission_v66.py251 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

import os
os.environ['PYTORCH_ROCM_ARCH'] = 'gfx950'
os.environ['CXX'] = 'clang++'

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

CPP_SRC = """
#include <torch/extension.h>
#include <c10/core/DispatchKeySet.h>
#include <ATen/core/dispatch/Dispatcher.h>
#include <tuple>
#include <string>

void init_kernels(std::string hsa_dir);

torch::Tensor monolithic_gemm(
    torch::Tensor A, torch::Tensor B_shuffle, torch::Tensor B_scale_sh, 
    int m, int k, int n);
"""

CUDA_SRC = """
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <hip/amd_detail/amd_hip_bf16.h>
#include <unordered_map>
#include <string>
#include <c10/core/DispatchKeySet.h>
#include <ATen/core/dispatch/Dispatcher.h>

struct CacheKey {
    int m, k, n, split_k;
    bool operator==(const CacheKey& o) const {
        return m == o.m && k == o.k && n == o.n && split_k == o.split_k;
    }
};

struct CacheKeyHash {
    size_t operator()(const CacheKey& k) const {
        return ((size_t)k.m << 48) ^ ((size_t)k.split_k << 32) ^ ((size_t)k.k << 16) ^ k.n;
    }
};

static std::unordered_map<CacheKey, std::tuple<torch::Tensor, torch::Tensor, torch::Tensor>, CacheKeyHash> g_tensor_cache;
static std::string cached_hsa_dir = "";

std::pair<const char*, int> get_dispatch(int m, int k, int n) {
    if (m <= 4 && n == 2880 && k == 512) return {"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E", 0};
    if (m <= 32 && n == 4096 && k == 512) return {"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E", 0};
    if (m <= 32 && n == 2880 && k == 512) return {"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E", 0};
    
    if (m <= 16 && n == 2112 && k == 7168) return {"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E", 2};
    if (m <= 64 && n == 7168 && k == 2048) return {"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E", 0};
    if (m <= 256 && n == 3072 && k == 1536) return {"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E", 1};
    
    if (m <= 8 && n == 2112 && k == 7168) return {"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E", 2}; 
    if (m <= 16 && n == 3072 && k == 1536) return {"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E", 1}; 
    if (m <= 64 && n == 3072 && k == 1536) return {"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E", 1}; 
    
    return {"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E", 0};
}

void init_kernels(std::string hsa_dir) {
    cached_hsa_dir = hsa_dir;
}

__global__ void prepare_a_kernel(const __nv_bfloat16* __restrict__ A, 
                                 uint8_t* __restrict__ A_q,
                                 uint8_t* __restrict__ bs_e8m0,
                                 int M, int K, int scaleN, int scaleM_pad,
                                 int stride_am, int stride_ak) {
    int global_tid = blockIdx.x * blockDim.x + threadIdx.x;
    int subgroup_id = global_tid / 4;
    int lane_id = threadIdx.x % 4;

    int total_subgroups = M * (K / 32); 
    if (subgroup_id >= total_subgroups) return;

    int m = subgroup_id / (K / 32);
    int n = subgroup_id % (K / 32);

    int base_idx = m * stride_am + n * 32 + lane_id * 8;
    
    union {
        ulonglong2 vec;
        uint16_t u16[8];
    } a_vec;
    a_vec.vec = *reinterpret_cast<const ulonglong2*>(&A[base_idx]);

    uint16_t thread_max = 0;
    for (int i = 0; i < 8; ++i) {
        uint16_t abs_val = a_vec.u16[i] & 0x7FFF;
        if (abs_val > thread_max) thread_max = abs_val;
    }

    for (int offset = 2; offset > 0; offset /= 2) {
        int other_i = __shfl_down((int)thread_max, offset, 64);
        uint16_t other = (uint16_t)other_i;
        if ((threadIdx.x % 4) + offset < 4) {
            if (other > thread_max) thread_max = other;
        }
    }
    
    int max_i = __shfl((int)thread_max, (threadIdx.x / 4) * 4, 64);
    uint16_t max_abs = (uint16_t)max_i;

    uint16_t amax_rounded = (max_abs + 0x0020) & 0xFF80;
    int exp = amax_rounded >> 7; 
    int scale_e8m0_unbiased = (exp == 0) ? -127 : exp - 129;
    scale_e8m0_unbiased = max(-127, min(127, scale_e8m0_unbiased));

    uint32_t quant_scale_u = (127 - scale_e8m0_unbiased) << 23;
    float quant_scale = __uint_as_float(quant_scale_u);

    uint32_t out_packed = 0;

    #pragma unroll
    for(int i = 0; i < 4; ++i) {
        __nv_bfloat16 a_val0, a_val1;
        *(uint16_t*)&a_val0 = a_vec.u16[i * 2];
        *(uint16_t*)&a_val1 = a_vec.u16[i * 2 + 1];
        
        float val0 = __bfloat162float(a_val0) * quant_scale;
        float val1 = __bfloat162float(a_val1) * quant_scale;
        
        uint32_t cvt = 0;
        cvt = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(cvt, val0, val1, 1.0f, 0);
        
        out_packed |= ((cvt & 0xFF) << (i * 8));
    }

    uint32_t* A_q_u32 = (uint32_t*)A_q;
    int aq_idx = m * (K / 8) + n * 4 + lane_id;
    A_q_u32[aq_idx] = out_packed;

    if (lane_id == 0) {
        uint8_t bs_val = (uint8_t)(scale_e8m0_unbiased + 127);
        int m_mod_32 = m % 32;
        int bs_0 = m / 32;
        int bs_1 = m_mod_32 / 16;
        int bs_2 = m_mod_32 % 16;
        int n_mod_8 = n % 8;
        int bs_3 = n / 8;
        int bs_4 = n_mod_8 / 4;
        int bs_5 = n_mod_8 % 4;
        int bs_offs = bs_1 + (bs_4 * 2) + (bs_2 * 4) + (bs_5 * 64) + (bs_3 * 256) + (bs_0 * 32 * scaleN);
        bs_e8m0[bs_offs] = bs_val;
    }
}

torch::Tensor monolithic_gemm(
    torch::Tensor A, torch::Tensor B_shuffle, torch::Tensor B_scale_sh, 
    int m, int k, int n) {
    
    auto dispatch = get_dispatch(m, k, n);
    const char* kname_ptr = dispatch.first;
    int split_k = dispatch.second;
    
    CacheKey key = {m, k, n, split_k};
    torch::Tensor x_fp4, bs_e8m0, out;
    int scaleN_valid = k / 32;
    int scaleN = ((scaleN_valid + 7) / 8) * 8;
    int scaleM_pad = ((m + 31) / 32) * 32;
    int scaleM_256 = ((m + 255) / 256) * 256;
    int padded_m = ((m + 31) / 32) * 32;

    if (g_tensor_cache.find(key) == g_tensor_cache.end()) {
        auto options_u8 = torch::TensorOptions().dtype(torch::kUInt8).device(A.device());
        auto options_bf16 = torch::TensorOptions().dtype(torch::kBFloat16).device(A.device());
        x_fp4 = torch::empty({m, k / 2}, options_u8);
        bs_e8m0 = torch::empty({scaleM_256, scaleN}, options_u8);
        bs_e8m0.fill_(127);
        if (split_k > 0) out = torch::zeros({padded_m, n}, options_bf16);
        else out = torch::empty({padded_m, n}, options_bf16);
        g_tensor_cache[key] = std::make_tuple(x_fp4, bs_e8m0, out);
    } else {
        auto& tuple_val = g_tensor_cache[key];
        x_fp4 = std::get<0>(tuple_val);
        bs_e8m0 = std::get<1>(tuple_val);
        out = std::get<2>(tuple_val);
        if (split_k > 0) out.zero_();
    }
    
    int total_subgroups = m * (k / 32); 
    int threads = 256;
    int blocks = (total_subgroups * 4 + threads - 1) / threads;
    if (blocks > 0) {
        prepare_a_kernel<<<blocks, threads>>>(
            reinterpret_cast<const __nv_bfloat16*>(A.data_ptr()),
            x_fp4.data_ptr<uint8_t>(), bs_e8m0.data_ptr<uint8_t>(),
            m, k, scaleN, scaleM_pad, A.stride(0), A.stride(1)
        );
    }
    
    // Bypass ATen dtype enforcement dynamically using the B_shuffle and B_scale types 
    at::Tensor x_view = at::from_blob(x_fp4.data_ptr(), {m, k / 2}, B_shuffle.options());
    at::Tensor bs_view = at::from_blob(bs_e8m0.data_ptr(), {scaleM_256, scaleN}, B_scale_sh.options());

    static auto op = c10::Dispatcher::singleton().findSchemaOrThrow("aiter::gemm_a4w4_asm", "");
    
    torch::jit::Stack stack;
    stack.push_back(x_view);
    stack.push_back(B_shuffle);
    stack.push_back(bs_view);
    stack.push_back(B_scale_sh);
    stack.push_back(out);
    stack.push_back(std::string(kname_ptr));
    stack.push_back(c10::IValue()); // bias=None
    stack.push_back(1.0); // alpha=1.0
    stack.push_back(0.0); // beta=0.0
    stack.push_back(true); // bpreshuffle=True
    stack.push_back(c10::optional<int64_t>(split_k)); // log2_k_split
    
    op.callBoxed(&stack);

    return out;
}
"""

_module_cache = None

def get_module():
    global _module_cache
    if _module_cache is None:
        _module_cache = load_inline(
            name='prepare_a_module_v66',
            cpp_sources=[CPP_SRC],
            cuda_sources=[CUDA_SRC],
            functions=['monolithic_gemm', 'init_kernels'],
            verbose=False,
            extra_cuda_cflags=["--offload-arch=gfx950", "-std=c++20", "-O3"],
        )
        import aiter
        hsa_dir = os.path.normpath(os.path.join(os.path.dirname(aiter.__file__), "..", "hsa", "gfx950", "f4gemm"))
        _module_cache.init_kernels(hsa_dir)
        
    return _module_cache

def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data
    m, k = A.shape
    n = B.shape[0]

    module = get_module()
    out = module.monolithic_gemm(A, B_shuffle, B_scale_sh, m, k, n)

    return out[:m]
scrolls · 251 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