Skip to content
KernelIndex
Search⌘K

submission 741515

HE · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-741515?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
8.45µs
#65 of 1143
2026-04-05

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:28b3a3408b227ca41f6112cbc3bc5743816f489224c977f7f219066cbbab61af
license declaredunknown
license concludedunknown
authorsHE
imported2026-08-15

Techniques

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

fp4FP4 quant + FP4 GEMM reference: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.
shared-memory__shared__ bf16x8 A_lds [2][mxc_tile_x*mxc_tile_y/8];
split-kvoid mm_impl_16x128_splitk(bf16 *C_ptr, const bf16 *A_ptr, const fp4x32 *B_ptr, const fp8e8m0 *Bs_ptr,

Kernel source

submission.py554 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

"""
FP4 quant + FP4 GEMM reference: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.
Quant logic follows aiter op_tests/test_gemm_a4w4.py (get_triton_quant(QuantType.per_1x32)).
"""
import os
from task import input_t, output_t
import torch
from torch.utils.cpp_extension import load_inline

kern_src = r"""
#include <hip/hip_runtime.h>

using fp4x2      = unsigned char;
using fp4x4      = fp4x2 __attribute__((ext_vector_type(2)));
using fp4x8      = fp4x2 __attribute__((ext_vector_type(4)));
using fp4x32     = fp4x2 __attribute__((ext_vector_type(16)));
using fp8e8m0    = unsigned char;
using fp8e8m0x4  = unsigned char __attribute__((ext_vector_type(4)));
using fp8e8m0x16 = unsigned char __attribute__((ext_vector_type(16)));
using fp16       = _Float16;
using fp16x2     = fp16 __attribute__((ext_vector_type(2)));
using fp16x32    = fp16 __attribute__((ext_vector_type(32)));
using bf16       = __bf16;
using bf16x2     = bf16 __attribute__((ext_vector_type(2)));
using bf16x8     = bf16 __attribute__((ext_vector_type(8)));
using bf16x32    = bf16 __attribute__((ext_vector_type(32)));
using fp32       = float;
using fp32x4     = fp32 __attribute__((ext_vector_type(4)));
using u8         = unsigned char;
using u16        = unsigned short;
using u32        = unsigned int;
using u32x4      = u32 __attribute__((ext_vector_type(4)));
using u32x8      = u32 __attribute__((ext_vector_type(8)));
using u64        = unsigned long;

#define alias __builtin_bit_cast
#define align(x, a) (((x) + (a) - 1) & ~((a) - 1))

union AliasedBlockBf16 {
    bf16x32 x32;
    bf16x8  x8[4];
    bf16x2  x2[16];
    bf16    x1[32];
    u32     u32[16];
};
static_assert(sizeof(AliasedBlockBf16) == 64);

union AliasedBlockFp4 {
    fp4x32 x32;
    fp4x8  x8[4];
    fp4x4  x4[8];
    fp4x2  x2[16];
    u32    u32[4];
    u32x4  u32x4;
};
static_assert(sizeof(AliasedBlockFp4) == 16);

template <int v>
struct C {
    static constexpr int value = v;
};
constexpr static auto C2 = C<2>{};
constexpr static auto C8 = C<8>{};

using i32x4 = int32_t __attribute__((ext_vector_type(4)));
using as3_uint32_ptr = uint32_t __attribute__((address_space(3))) *;

extern "C" __device__ void llvm_amdgcn_raw_buffer_load_lds(
    i32x4 rsrc, as3_uint32_ptr lds_ptr, int size, int voffset, int soffset, int offset, int aux)
        __asm("llvm.amdgcn.raw.buffer.load.lds");

struct buffer_resource {
    uint64_t ptr;
    uint32_t range;
    uint32_t config;
};

__device__ inline i32x4 make_srsrc(const void *ptr, uint32_t range_bytes) {
    buffer_resource rsrc = {reinterpret_cast<uint64_t>(ptr), range_bytes, 0x10000};
    return alias(i32x4, rsrc);
}

struct QuantizedBlock {
    AliasedBlockFp4 block;
    fp8e8m0         scale;
};

__device__ __forceinline__
QuantizedBlock quantize_block_mxfp4(AliasedBlockBf16 dat) {
    auto dat2 = dat;

    for (u32 i = 0; i < 16; ++i)
        dat2.u32[i] &= 0x7fff7fff;

    auto maxv = ({
        bf16x2 tmp0, tmp1, tmp2, tmp3, tmp4;
        asm(
            "v_pk_maximum3_f16  %16,  %0, %1,  %2\n\t"
            "v_pk_maximum3_f16  %17,  %3, %4,  %5\n\t"
            "v_pk_maximum3_f16  %18,  %6, %7,  %8\n\t"
            "v_pk_maximum3_f16  %19,  %9, %10, %11\n\t"
            "v_pk_maximum3_f16  %20, %12, %13, %14\n\t"
            "v_pk_maximum3_f16  %16, %16, %17, %15\n\t"
            "v_pk_maximum3_f16  %17, %18, %19, %20\n\t"
            "v_pk_max_f16       %16, %16, %17\n\t"
            "v_max_f16_sdwa     %16, %16, %16 dst_sel:DWORD dst_unused:UNUSED_PAD src0_sel:WORD_0 src1_sel:WORD_1\n\t"
            : "+v"(dat2.x2[ 0]), "+v"(dat2.x2[ 1]), "+v"(dat2.x2[ 2]), "+v"(dat2.x2[ 3]),
              "+v"(dat2.x2[ 4]), "+v"(dat2.x2[ 5]), "+v"(dat2.x2[ 6]), "+v"(dat2.x2[ 7]),
              "+v"(dat2.x2[ 8]), "+v"(dat2.x2[ 9]), "+v"(dat2.x2[10]), "+v"(dat2.x2[11]),
              "+v"(dat2.x2[12]), "+v"(dat2.x2[13]), "+v"(dat2.x2[14]), "+v"(dat2.x2[15]),
              "=v"(tmp0), "=v"(tmp1), "=v"(tmp2), "=v"(tmp3), "=v"(tmp4)
        );

        tmp0.x;
    });

    auto s      = alias(fp8e8m0, static_cast< u8>(((alias(u16, maxv) + 0x20) >> 7) - 2));
    auto s_fp32 = alias(fp32,    static_cast<u32>(s) << 23);

    AliasedBlockFp4 block = {};
    block.u32[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(block.u32[0], dat.x2[ 0], s_fp32, 0);
    block.u32[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(block.u32[0], dat.x2[ 1], s_fp32, 1);
    block.u32[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(block.u32[0], dat.x2[ 2], s_fp32, 2);
    block.u32[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(block.u32[0], dat.x2[ 3], s_fp32, 3);
    block.u32[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(block.u32[1], dat.x2[ 4], s_fp32, 0);
    block.u32[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(block.u32[1], dat.x2[ 5], s_fp32, 1);
    block.u32[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(block.u32[1], dat.x2[ 6], s_fp32, 2);
    block.u32[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(block.u32[1], dat.x2[ 7], s_fp32, 3);
    block.u32[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(block.u32[2], dat.x2[ 8], s_fp32, 0);
    block.u32[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(block.u32[2], dat.x2[ 9], s_fp32, 1);
    block.u32[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(block.u32[2], dat.x2[10], s_fp32, 2);
    block.u32[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(block.u32[2], dat.x2[11], s_fp32, 3);
    block.u32[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(block.u32[3], dat.x2[12], s_fp32, 0);
    block.u32[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(block.u32[3], dat.x2[13], s_fp32, 1);
    block.u32[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(block.u32[3], dat.x2[14], s_fp32, 2);
    block.u32[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(block.u32[3], dat.x2[15], s_fp32, 3);

    return { block, s };
}

constexpr int mxc_tile_x = 128, mxc_tile_y = 16, mxc_qblock_sz = 32;

__device__ __forceinline__
u32 swiz(u32 idx) {
    return ((idx >> 4 & 0b11) | (idx >> 3 & 0b1000));
};

__device__ __forceinline__
void mm_impl_16x128_basic(bf16 *C_ptr, const bf16 *A_ptr, const fp4x32 *B_ptr, const fp8e8m0 *Bs_ptr,
                          u32 tm, u32 tn, u32 tk, u32 nt, u32 m, u32 n, u32 k)
{
    auto bid = blockIdx.x, rbid = ((bid & 7) * (nt >> 3)) + (bid >> 3);
    auto tid = threadIdx.x, lid = tid & 63;

    auto C_tile_x = rbid >> __builtin_ctz(tm), C_tile_y = rbid & (tm - 1),
        C_x = C_tile_x * mxc_tile_y, C_y = C_tile_y * mxc_tile_y;

    if (C_tile_x >= tn)
        return;

    auto A_rsrc  = make_srsrc((void *)A_ptr,  m * k * sizeof(*A_ptr )),
         Bq_rsrc = make_srsrc((void *)B_ptr,  n * k * sizeof(*B_ptr ) / 32),
         Bs_rsrc = make_srsrc((void *)Bs_ptr, n * k * sizeof(*Bs_ptr) / mxc_qblock_sz);

    __shared__ bf16x8  A_lds [2][mxc_tile_x*mxc_tile_y/8];
    __shared__ fp4x32  Bq_lds[2][mxc_tile_x*mxc_tile_y/32];
    __shared__ fp8e8m0 Bs_lds[2][mxc_tile_x*mxc_tile_y/32*4];

    auto load_to_lds = [&](u32 it, u32 phase) __attribute__((always_inline)) {
        __builtin_amdgcn_s_setprio(1);

        auto A_voff = (lid >> 4) * k + (lid & 15) * 8,
             A_soff = C_y * k + it,
             A_incr = 4 * k;

        for (u32 i = 0; i < 4; ++i) {
            as3_uint32_ptr lds_ptr = (as3_uint32_ptr)(reinterpret_cast<uintptr_t>(A_lds[phase]));
            llvm_amdgcn_raw_buffer_load_lds(A_rsrc, lds_ptr + i * 64 * 4, sizeof(bf16x8),
                (A_voff ^ (swiz(lid | (i << 6)) * 8)) * sizeof(bf16), (A_soff + i * A_incr) * sizeof(bf16), 0, 0);
        }

        __builtin_amdgcn_sched_barrier(1);

        {
            as3_uint32_ptr lds_ptr = (as3_uint32_ptr)(reinterpret_cast<uintptr_t>(Bq_lds[phase]));
            llvm_amdgcn_raw_buffer_load_lds(Bq_rsrc, lds_ptr, sizeof(fp4x32),
                lid * sizeof(fp4x32), (C_tile_x * tk + it / mxc_tile_x) * 64 * sizeof(fp4x32), 0, 0);
        }

        {
            as3_uint32_ptr lds_ptr = (as3_uint32_ptr)(reinterpret_cast<uintptr_t>(Bs_lds[phase]));
            llvm_amdgcn_raw_buffer_load_lds(Bs_rsrc, lds_ptr, sizeof(fp8e8m0x4),
                lid * sizeof(fp8e8m0x4), ((C_tile_x & ~1) * tk + (it & ~(mxc_tile_x * 2 - 1)) / mxc_tile_x * 2) * 16 * sizeof(fp8e8m0x4), 0, 0);
        }
    };

    auto do_compute = [&]<int outstanding>(fp32x4 &C_dat, u32 it, u32 prev_phase, C<outstanding>) __attribute__((always_inline)) {
        __builtin_amdgcn_s_setprio(0);

        __builtin_amdgcn_s_waitcnt(0xf70 | outstanding);

        auto A_dat = AliasedBlockBf16{};
        for (u32 i = 0; i < 4; ++i) {
            auto idx = ((lid & 15) * 4 + (lid >> 4)) * 4 + i;
            A_dat.x8[i] = A_lds[prev_phase][idx ^ swiz(idx)];
        }

        auto [Aq_dat, As] = quantize_block_mxfp4(A_dat);

        __builtin_amdgcn_s_waitcnt(0xf70 | (outstanding - 2));

        auto Bs_dat = reinterpret_cast<u32 *>(Bs_lds[prev_phase])[lid];
        Bs_dat >>= ((C_tile_x & 1) + (it / mxc_tile_x & 1) * 2) * 8;

        auto Bq_dat = alias(AliasedBlockFp4, Bq_lds[prev_phase][lid]);

        C_dat = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
            u32x8{Aq_dat.u32[0], Aq_dat.u32[1], Aq_dat.u32[2], Aq_dat.u32[3]},
            u32x8{Bq_dat.u32[0], Bq_dat.u32[1], Bq_dat.u32[2], Bq_dat.u32[3]},
            C_dat, 4, 4, 0, As, 0, Bs_dat);
    };

    u32 it = 0, prev_it = it, phase = 0, prev_phase = phase;
    load_to_lds(it, phase);

    fp32x4 C_dat = {};
    for (it = mxc_tile_x, phase = 1; it < k; prev_it = it, it += mxc_tile_x, prev_phase = phase, phase ^= 1) {
        load_to_lds(it, phase);
        do_compute(C_dat, prev_it, prev_phase, C8);
    }

    do_compute(C_dat, prev_it, prev_phase, C2);

    asm volatile("" :: "v"(C_dat[0]), "v"(C_dat[1]), "v"(C_dat[2]), "v"(C_dat[3]));

    if (C_y + (lid >> 4) * 4 < m) {
        for (u32 i = 0; i < 4; ++i)
            __builtin_nontemporal_store(static_cast<bf16>(C_dat[i]),
                C_ptr + (C_y + (lid >> 4) * 4 + i) * n + C_x + (lid & 15));
    }
}

__device__ __forceinline__
void mm_impl_32x256_blocked(bf16 *C_ptr, const bf16 *A_ptr, const fp4x32 *B_ptr, const fp8e8m0 *Bs_ptr,
                            u32 tm, u32 tn, u32 tk, u32 nt, u32 m, u32 n, u32 k)
{
    auto bid = blockIdx.x, rbid = ((bid & 7) * (nt >> 5)) + (bid >> 3);
    auto tid = threadIdx.x, wid = tid >> 6, lid = tid & 63;

    wid = __builtin_amdgcn_readfirstlane(wid);

    auto C_tile_x = (rbid >> (__builtin_ctz(tm) - 1)) * 2, C_tile_y = (rbid & ((tm >> 1) - 1)) * 2,
        C_x = C_tile_x * mxc_tile_y, C_y = C_tile_y * mxc_tile_y;
    if (C_tile_x >= tn || C_tile_y >= tm)
        return;

    auto A_rsrc  = make_srsrc((void *)A_ptr,  m * k * sizeof(*A_ptr )),
         Bq_rsrc = make_srsrc((void *)B_ptr,  n * k * sizeof(*B_ptr ) / 32),
         Bs_rsrc = make_srsrc((void *)Bs_ptr, n * k * sizeof(*Bs_ptr) / mxc_qblock_sz);

    __shared__ bf16x8  A_lds[2][4][mxc_tile_x*mxc_tile_y/8];
    __shared__ fp4x32  Aq_lds  [4][mxc_tile_x*mxc_tile_y/32],     Bq_lds[2][4][mxc_tile_x*mxc_tile_y/32];
    __shared__ fp8e8m0 As_lds     [mxc_tile_x*mxc_tile_y/32 * 4], Bs_lds[2]   [mxc_tile_x*mxc_tile_y/32 * 4];

    auto load_to_lds = [&](u32 it, u32 phase) __attribute__((always_inline)) {
        __builtin_amdgcn_s_setprio(1);

        auto A_voff = (lid >> 4) * k + (lid & 15) * 8,
             A_soff = (C_y + (wid >> 1) * mxc_tile_y) * k + it + (wid & 1) * mxc_tile_x,
             A_incr = 4 * k;

        for (u32 i = 0; i < 4; ++i) {
            as3_uint32_ptr lds_ptr = (as3_uint32_ptr)(reinterpret_cast<uintptr_t>(A_lds[phase][wid]));
            llvm_amdgcn_raw_buffer_load_lds(A_rsrc, lds_ptr + i * 64 * 4, sizeof(bf16x8),
                (A_voff ^ (swiz(lid | (i << 6)) * 8)) * sizeof(bf16), (A_soff + i * A_incr) * sizeof(bf16), 0, 0);
        }

        __builtin_amdgcn_sched_barrier(1);

        {
            as3_uint32_ptr lds_ptr = (as3_uint32_ptr)(reinterpret_cast<uintptr_t>(Bq_lds[phase][wid]));
            llvm_amdgcn_raw_buffer_load_lds(Bq_rsrc, lds_ptr, sizeof(fp4x32),
                lid * sizeof(fp4x32), ((C_tile_x + (wid >> 1)) * tk + it / mxc_tile_x + (wid & 1)) * 64 * sizeof(fp4x32), 0, 0);
        }

        if (lid < 16) {
            as3_uint32_ptr lds_ptr = (as3_uint32_ptr)(reinterpret_cast<uintptr_t>(Bs_lds[phase]));
            llvm_amdgcn_raw_buffer_load_lds(Bs_rsrc, lds_ptr + wid * 4 * 4, sizeof(fp8e8m0x4),
                lid * sizeof(fp8e8m0x4), ((C_tile_x * tk) + it / mxc_tile_x * 2 + wid) * 16 * sizeof(fp8e8m0x4), 0, 0);
        }
    };

    auto do_compute = [&]<int outstanding>(fp32x4 &C_dat, u32 prev_phase, C<outstanding>) __attribute__((always_inline)) {
        __builtin_amdgcn_s_setprio(0);

        __builtin_amdgcn_s_waitcnt(0xf70 | outstanding);

        auto A_dat = AliasedBlockBf16{};
        for (u32 i = 0; i < 4; ++i) {
            auto idx = ((lid & 15) * 4 + (lid >> 4)) * 4 + i;
            A_dat.x8[i] = A_lds[prev_phase][wid][idx ^ swiz(idx)];
        }

        auto [Aq_dat, As] = quantize_block_mxfp4(A_dat);

        Aq_lds[wid][lid] = Aq_dat.x32;
        As_lds[lid *  4 + wid] = As;

        __builtin_amdgcn_s_waitcnt(0xf70 | (outstanding - 2));
        __builtin_amdgcn_s_barrier();

        auto Bs_dat = reinterpret_cast<u32 *>(Bs_lds[prev_phase])[lid];
        Bs_dat = __builtin_amdgcn_perm(Bs_dat, Bs_dat, 0x01030200) >> ((wid & 1) * 16);

        auto Bq_dat = alias(AliasedBlockFp4, Bq_lds[prev_phase][(wid & 1) * 3][lid]);

        C_dat = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
            u32x8{Aq_dat.u32[0], Aq_dat.u32[1], Aq_dat.u32[2], Aq_dat.u32[3]},
            u32x8{Bq_dat.u32[0], Bq_dat.u32[1], Bq_dat.u32[2], Bq_dat.u32[3]},
            C_dat, 4, 4, 0, As, 0, Bs_dat);

        __builtin_amdgcn_sched_barrier(0);
        __builtin_amdgcn_s_barrier();

        auto As_dat = reinterpret_cast<u32 *>(As_lds)[lid] >> ((wid ^ 1) * 8);

        Aq_dat = alias(AliasedBlockFp4, Aq_lds[wid ^ 1][lid]);
        Bq_dat = alias(AliasedBlockFp4, Bq_lds[prev_phase][(wid & 1) + 1][lid]);

        C_dat = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
            u32x8{Aq_dat.u32[0], Aq_dat.u32[1], Aq_dat.u32[2], Aq_dat.u32[3]},
            u32x8{Bq_dat.u32[0], Bq_dat.u32[1], Bq_dat.u32[2], Bq_dat.u32[3]},
            C_dat, 4, 4, 0, As_dat, 1, Bs_dat);

        __builtin_amdgcn_s_barrier();
    };

    load_to_lds(0, 0);

    u32 prev_phase = 0;
    fp32x4 C_dat = {};
    for (u32 it = 2*mxc_tile_x, phase = 1; it < k; it += 2*mxc_tile_x, prev_phase = phase, phase ^= 1) {
        load_to_lds(it, phase);
        do_compute(C_dat, prev_phase, C8);
    }

    do_compute(C_dat, prev_phase, C2);

    for (u32 i = 0; i < 4; ++i)
        __builtin_nontemporal_store(static_cast<bf16>(C_dat[i]),
            C_ptr + (C_y + (wid >> 1) * mxc_tile_y + (lid >> 4) * 4 + i) * n +
            C_x + (wid & 1) * mxc_tile_y + (lid & 15));
}

__device__ __forceinline__
void mm_impl_16x128_splitk(bf16 *C_ptr, const bf16 *A_ptr, const fp4x32 *B_ptr, const fp8e8m0 *Bs_ptr,
                           u32 tm, u32 tn, u32 tk, u32 nt, u32 m, u32 n, u32 k)
{
    auto bid = blockIdx.x, rbid = ((bid & 7) * (nt >> 3)) + (bid >> 3);
    auto tid = threadIdx.x, wid = tid >> 6, lid = tid & 63;

    wid = __builtin_amdgcn_readfirstlane(wid);

    auto C_tile_x = rbid >> __builtin_ctz(tm), C_tile_y = rbid & (tm - 1),
        C_x = C_tile_x * mxc_tile_y, C_y = C_tile_y * mxc_tile_y;

    if (C_tile_x >= tn)
        return;

    auto A_rsrc  = make_srsrc((void *)A_ptr,  m * k * sizeof(*A_ptr )),
         Bq_rsrc = make_srsrc((void *)B_ptr,  n * k * sizeof(*B_ptr ) / 32),
         Bs_rsrc = make_srsrc((void *)Bs_ptr, n * k * sizeof(*Bs_ptr) / mxc_qblock_sz);

    __shared__ bf16x8  A_lds [2][4][mxc_tile_x*mxc_tile_y/8];
    __shared__ fp4x32  Bq_lds[2][4][mxc_tile_x*mxc_tile_y/32];
    __shared__ fp8e8m0 Bs_lds[2][4][mxc_tile_x*mxc_tile_y/32*4];
    __shared__ fp32x4  C_lds    [4][mxc_tile_x*mxc_tile_y/8];

    auto load_to_lds = [&](u32 it, u32 phase) __attribute__((always_inline)) {
        __builtin_amdgcn_s_setprio(1);

        auto A_voff = (lid >> 4) * k + (lid & 15) * 8,
             A_soff = C_y * k + it,
             A_incr = 4 * k;

        for (u32 i = 0; i < 4; ++i) {
            as3_uint32_ptr lds_ptr = (as3_uint32_ptr)(reinterpret_cast<uintptr_t>(A_lds[phase][wid]));
            llvm_amdgcn_raw_buffer_load_lds(A_rsrc, lds_ptr + i * 64 * 4, sizeof(bf16x8),
                (A_voff ^ (swiz(lid | (i << 6)) * 8)) * sizeof(bf16), (A_soff + i * A_incr) * sizeof(bf16), 0, 0);
        }

        __builtin_amdgcn_sched_barrier(1);

        {
            as3_uint32_ptr lds_ptr = (as3_uint32_ptr)(reinterpret_cast<uintptr_t>(Bq_lds[phase][wid]));
            llvm_amdgcn_raw_buffer_load_lds(Bq_rsrc, lds_ptr, sizeof(fp4x32),
                lid * sizeof(fp4x32), (C_tile_x * tk + it / mxc_tile_x) * 64 * sizeof(fp4x32), 0, 0);
        }

        {
            as3_uint32_ptr lds_ptr = (as3_uint32_ptr)(reinterpret_cast<uintptr_t>(Bs_lds[phase][wid]));
            llvm_amdgcn_raw_buffer_load_lds(Bs_rsrc, lds_ptr, sizeof(fp8e8m0x4),
                lid * sizeof(fp8e8m0x4), ((C_tile_x & ~1) * tk + (it & ~(mxc_tile_x * 2 - 1)) / mxc_tile_x * 2) * 16 * sizeof(fp8e8m0x4), 0, 0);
        }
    };

    auto do_compute = [&]<int outstanding>(fp32x4 &C_dat, u32 it, u32 prev_phase, C<outstanding>) __attribute__((always_inline)) {
        __builtin_amdgcn_s_setprio(0);

        __builtin_amdgcn_s_waitcnt(0xf70 | outstanding);

        auto A_dat = AliasedBlockBf16{};
        for (u32 i = 0; i < 4; ++i) {
            auto idx = ((lid & 15) * 4 + (lid >> 4)) * 4 + i;
            A_dat.x8[i] = A_lds[prev_phase][wid][idx ^ swiz(idx)];
        }

        auto [Aq_dat, As] = quantize_block_mxfp4(A_dat);

        __builtin_amdgcn_s_waitcnt(0xf70 | (outstanding - 2));

        auto Bs_dat = reinterpret_cast<u32 *>(Bs_lds[prev_phase][wid])[lid];
        Bs_dat >>= ((C_tile_x & 1) + (it / mxc_tile_x & 1) * 2) * 8;

        auto Bq_dat = alias(AliasedBlockFp4, Bq_lds[prev_phase][wid][lid]);

        C_dat = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
            u32x8{Aq_dat.u32[0], Aq_dat.u32[1], Aq_dat.u32[2], Aq_dat.u32[3]},
            u32x8{Bq_dat.u32[0], Bq_dat.u32[1], Bq_dat.u32[2], Bq_dat.u32[3]},
            C_dat, 4, 4, 0, As, 0, Bs_dat);
    };

    u32 it = k / 4 * wid, end = k / 4 * (wid + 1);
    load_to_lds(it, 0);

    u32 prev_it = it, prev_phase = 0;
    fp32x4 C_dat = {};
    for (u32 phase = (it += mxc_tile_x, 1); it < end; prev_it = it, prev_phase = phase, it += mxc_tile_x, phase ^= 1) {
        load_to_lds(it, phase);
        do_compute(C_dat, prev_it, prev_phase, C8);
    }

    do_compute(C_dat, prev_it, prev_phase, C2);

    C_lds[wid][lid] = C_dat;
    __builtin_amdgcn_s_waitcnt(0xc07f);
    __builtin_amdgcn_s_barrier();

    if (wid == 0) {
        fp32x4 sum = C_lds[0][lid];
        sum += C_lds[1][lid];
        sum += C_lds[2][lid];
        sum += C_lds[3][lid];

        if (C_y + (lid >> 4) * 4 < m) {
            for (u32 i = 0; i < 4; ++i)
                __builtin_nontemporal_store(static_cast<bf16>(sum[i]),
                    C_ptr + (C_y + (lid >> 4) * 4 + i) * n + C_x + (lid & 15));
        }
    }
}

__global__ __launch_bounds__(64)
void kern_mm_basic(bf16 *C_ptr, const bf16 *A_ptr, const fp4x32 *B_ptr, const fp8e8m0 *Bs_ptr,
                   u32 tm, u32 tn, u32 tk, u32 nt, u32 m, u32 n, u32 k)
{
    __builtin_assume(m >= 4 && n >= 2112 && k >= 512 && (m & 3) == 0 && (n & 63) == 0 && (k & 511) == 0);
    __builtin_assume(tm >= 1 && tn >= 132 && tk >= 4 && (tm & (tm - 1)) == 0 && (tn & 3) == 0 && (tk & 3) == 0);
    mm_impl_16x128_basic(C_ptr, A_ptr, B_ptr, Bs_ptr, tm, tn, tk, nt, m, n, k);
}

__global__ __launch_bounds__(4*64, 2)
void kern_mm_blocked(bf16 *C_ptr, const bf16 *A_ptr, const fp4x32 *B_ptr, const fp8e8m0 *Bs_ptr,
                     u32 tm, u32 tn, u32 tk, u32 nt, u32 m, u32 n, u32 k)
{
    __builtin_assume(m >= 4 && n >= 2112 && k >= 512 && (m & 3) == 0 && (n & 63) == 0 && (k & 511) == 0);
    __builtin_assume(tm >= 1 && tn >= 132 && tk >= 4 && (tm & (tm - 1)) == 0 && (tn & 3) == 0 && (tk & 3) == 0);
    mm_impl_32x256_blocked(C_ptr, A_ptr, B_ptr, Bs_ptr, tm, tn, tk, nt, m, n, k);
}

__global__ __launch_bounds__(4*64, 2)
void kern_mm_splitk(bf16 *C_ptr, const bf16 *A_ptr, const fp4x32 *B_ptr, const fp8e8m0 *Bs_ptr,
                    u32 tm, u32 tn, u32 tk, u32 nt, u32 m, u32 n, u32 k)
{
    __builtin_assume(m >= 4 && n >= 2112 && k >= 512 && (m & 3) == 0 && (n & 63) == 0 && (k & 511) == 0);
    __builtin_assume(tm >= 1 && tn >= 132 && tk >= 4 && (tm & (tm - 1)) == 0 && (tn & 3) == 0 && (tk & 3) == 0);
    mm_impl_16x128_splitk(C_ptr, A_ptr, B_ptr, Bs_ptr, tm, tn, tk, nt, m, n, k);
}
"""

torch_src = r"""
void mm(torch::Tensor C, torch::Tensor A, torch::Tensor B_shuffle, torch::Tensor B_scale_sh) {
    auto m = C.size(0), n = C.size(1), k = A.size(1);
    auto tm = (m + mxc_tile_y - 1) / mxc_tile_y, tn = (n + mxc_tile_y - 1) / mxc_tile_y,
         tk = (k + mxc_tile_x - 1) / mxc_tile_x,
         nt = tm * tn;

    if (tm < 2 && tk <= 4) {
        nt = align(nt, 8);
        kern_mm_basic<<<nt, 64>>>(
            reinterpret_cast<      bf16    *>(C.data_ptr()),
            reinterpret_cast<const bf16    *>(A.data_ptr()),
            reinterpret_cast<const fp4x32  *>(B_shuffle.data_ptr()),
            reinterpret_cast<const fp8e8m0 *>(B_scale_sh.data_ptr()),
            tm, tn, tk, nt, m, n, k
        );
    } else if (tm < 2) {
        nt = align(nt, 8);
        kern_mm_splitk<<<nt, 4*64>>>(
            reinterpret_cast<      bf16    *>(C.data_ptr()),
            reinterpret_cast<const bf16    *>(A.data_ptr()),
            reinterpret_cast<const fp4x32  *>(B_shuffle.data_ptr()),
            reinterpret_cast<const fp8e8m0 *>(B_scale_sh.data_ptr()),
            tm, tn, tk, nt, m, n, k
        );
    } else {
        nt = align(nt, 32);
        kern_mm_blocked<<<nt >> 2, 4*64>>>(
            reinterpret_cast<      bf16    *>(C.data_ptr()),
            reinterpret_cast<const bf16    *>(A.data_ptr()),
            reinterpret_cast<const fp4x32  *>(B_shuffle.data_ptr()),
            reinterpret_cast<const fp8e8m0 *>(B_scale_sh.data_ptr()),
            tm, tn, tk, nt, m, n, k
        );
    }
}
"""

os.environ["TORCH_DONT_CHECK_COMPILER_ABI"] = "1"
os.environ["CXX"] = "/opt/rocm/bin/amdclang++"
module = load_inline(
    name="mm",
    cpp_sources=[kern_src + torch_src],
    functions=["mm"],
    verbose=True,
    extra_cflags=[
        "-xhip", "--offload-arch=gfx950",
        "-std=gnu++23", "-g", "-O3", "-ffast-math", "-funroll-loops",
        "-mllvm=--amdgpu-kernarg-preload-count=16",
    ],
)

def custom_kernel(data: input_t) -> output_t:
    A, B, _, B_shuffle, B_scale_sh = data
    A = A.contiguous()
    C = torch.empty((A.shape[0], B.shape[0]), dtype=torch.bfloat16, device="cuda")

    module.mm(C, A, B_shuffle, B_scale_sh)
    return C

scrolls · 554 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