Skip to content
KernelIndex
Search⌘K

submission 727931

Navid Khazaee · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

mm-submission_v64.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-727931?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
13.5µs
#447 of 1143
2026-04-04

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:8f77273f31b58dc65685605cffc3fb4cba34ec5df0e0eaf202cb58bcb8cf919c
license declaredunknown
license concludedunknown
authorsNavid Khazaee
imported2026-08-26

Techniques

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

fp4v64: Hardware FP4 quant + ASM GEMM.

Kernel source

mm-submission_v64.py136 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
v64: Hardware FP4 quant + ASM GEMM.

Uses __builtin_amdgcn_cvt_scalef32_pk_fp4_f32 for quantization.
1 HW instruction per 2 values vs 20+ software instructions.
Single module, no fallback, minimal compilation.
"""
import os
os.environ['PYTORCH_ROCM_ARCH'] = 'gfx950'
os.environ['CXX'] = 'clang++'

import torch
import sys
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
import aiter
from aiter import dtypes

_ = aiter.gemm_a4w4
_ga = torch.ops.aiter.gemm_a4w4_asm
_KN = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"

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

__device__ __forceinline__ void compute_scale(float mx, float& sf, uint8_t& e8) {
    if (mx == 0.0f) { sf = 1.0f; e8 = 0; return; }
    uint32_t ab; memcpy(&ab, &mx, 4);
    uint32_t adj = (ab + 0x200000u) & 0xFF800000u;
    int eu = (int)((adj >> 23) & 0xFFu) - 129;
    if (eu < -127) eu = -127; if (eu > 127) eu = 127;
    e8 = (uint8_t)(eu + 127);
    uint32_t sb = (uint32_t)(eu + 127) << 23;
    memcpy(&sf, &sb, 4);
}

__global__ void hw_qk(
    const __hip_bfloat16* __restrict__ A,
    uint8_t* __restrict__ fp4,
    uint8_t* __restrict__ sc,
    int M, int K, int sn8
) {
    int r = blockIdx.x;
    if (r >= M) return;
    int ng = K / 32;
    const __hip_bfloat16* rp = A + r * K;
    uint8_t* fr = fp4 + r * (K/2);
    int i0=r/32, i1=(r/16)%2, i2=r%16;
    int bo = i0*(sn8*256) + i2*4 + i1;

    for (int kg = threadIdx.x; kg < ng; kg += blockDim.x) {
        int b = kg * 32;
        float mx = 0.0f, v[32];

        #pragma unroll
        for (int i = 0; i < 32; i++) {
            float x = __bfloat162float(rp[b+i]);
            v[i] = x;
            mx = fmaxf(mx, fabsf(x));
        }

        float sf; uint8_t e8;
        compute_scale(mx, sf, e8);

        uint32_t w0 = 0;
        w0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(w0, v[0],  v[1],  sf, 0);
        w0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(w0, v[2],  v[3],  sf, 1);
        w0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(w0, v[4],  v[5],  sf, 2);
        w0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(w0, v[6],  v[7],  sf, 3);
        uint32_t w1 = 0;
        w1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(w1, v[8],  v[9],  sf, 0);
        w1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(w1, v[10], v[11], sf, 1);
        w1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(w1, v[12], v[13], sf, 2);
        w1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(w1, v[14], v[15], sf, 3);
        uint32_t w2 = 0;
        w2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(w2, v[16], v[17], sf, 0);
        w2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(w2, v[18], v[19], sf, 1);
        w2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(w2, v[20], v[21], sf, 2);
        w2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(w2, v[22], v[23], sf, 3);
        uint32_t w3 = 0;
        w3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(w3, v[24], v[25], sf, 0);
        w3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(w3, v[26], v[27], sf, 1);
        w3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(w3, v[28], v[29], sf, 2);
        w3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(w3, v[30], v[31], sf, 3);

        uint32_t* out4 = reinterpret_cast<uint32_t*>(fr + kg*16);
        out4[0] = w0; out4[1] = w1; out4[2] = w2; out4[3] = w3;
        sc[bo + kg/8*256 + (kg%4)*64 + ((kg/4)%2)*2] = e8;
    }
}

void qf(torch::Tensor A, torch::Tensor fp4, torch::Tensor sc, int sn8) {
    int M = A.size(0), K = A.size(1);
    int ng = K / 32, t = (ng < 256) ? ng : 256;
    hw_qk<<<M, t>>>(
        reinterpret_cast<const __hip_bfloat16*>(A.data_ptr()),
        fp4.data_ptr<uint8_t>(), sc.data_ptr<uint8_t>(),
        M, K, sn8);
}
"""

_qm = load_inline(name='q64',
    cpp_sources=["void qf(torch::Tensor,torch::Tensor,torch::Tensor,int);"],
    cuda_sources=[SRC], functions=['qf'],
    extra_cuda_cflags=["--offload-arch=gfx950","-std=c++20","-O3"],
    verbose=True)

class _S:
    __slots__=['fp4','fp4v','scb','scv','out','osl','sn8']
_st = {}
def _mk(m,n,k,d):
    s=_S();ng=k//32;sm=((m+255)//256)*256;sn=((ng+7)//8)*8
    s.sn8=sn//8
    s.fp4=torch.empty((m,k//2),dtype=torch.uint8,device=d)
    s.fp4v=s.fp4.view(dtypes.fp4x2)
    s.scb=torch.zeros((sm,sn),dtype=torch.uint8,device=d)
    s.scv=s.scb.view(dtypes.fp8_e8m0)
    mp=((m+31)//32)*32
    s.out=torch.empty((mp,n),dtype=torch.bfloat16,device=d);s.osl=s.out[:m]
    _st[(m,n,k)]=s;return s

def custom_kernel(data: input_t) -> output_t:
    A=data[0]; Bsh=data[3]; Bsc=data[4]
    m,k=A.shape; n=Bsh.shape[0]; key=(m,n,k)
    try: s=_st[key]
    except KeyError: s=_mk(m,n,k,A.device)
    _qm.qf(A, s.fp4, s.scb, s.sn8)
    _ga(s.fp4v, Bsh, s.scv, Bsc, s.out, _KN, None, 1.0, 0.0, True, 0)
    return s.osl
scrolls · 136 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