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
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.
fp4
FP4 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-k
void 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