submission 673831
John Hahn · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1435 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-673831?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:ed7ac14219cd39186e5175a2a19b62f0773982a79aae9581b28dff4fb273f7e1
license declaredunknown
license concludedunknown
authorsJohn Hahn
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
extern __shared__ uint8_t lds[];split-k
mxfp4_gemm_splitk_fused(vector-width = uint4
const uint4* src4 = reinterpret_cast<const uint4*>(src + offset);Kernel source
submission.py1435 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""v343: 4-way split (narrow+hex together)."""
from task import input_t, output_t
import torch
import ctypes
import os
from concurrent.futures import ThreadPoolExecutor
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
HIP_SOURCE = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <hip/hip_fp16.h>
#include <hip/hip_bfloat16.h>
#include <cstdint>
using mfma_acc_t = __attribute__((ext_vector_type(4))) float;
using mfma_input_t = __attribute__((ext_vector_type(8))) uint32_t;
using u32x4 = __attribute__((ext_vector_type(4))) uint32_t;
#define MXFP4_BLOCK_SIZE 32
__device__ __forceinline__ mfma_acc_t
mfma_fp4_16x16(mfma_input_t a, mfma_input_t b, mfma_acc_t c,
uint32_t sa, uint32_t sb) {
#if defined(__gfx950__)
return __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
a, b, c, 4, 4, 0, sa, 0, sb);
#else
return c;
#endif
}
#ifdef UNIT_HEX
// Inline ASM: 6 MFMAs for hex kernel in single block (no compiler waits)
__device__ __forceinline__ void
asm_mfma_hex(u32x4 af, uint32_t sa,
u32x4 bf0, u32x4 bf1, u32x4 bf2, u32x4 bf3, u32x4 bf4, u32x4 bf5,
uint32_t sb0, uint32_t sb1, uint32_t sb2, uint32_t sb3, uint32_t sb4, uint32_t sb5,
mfma_acc_t& acc0, mfma_acc_t& acc1, mfma_acc_t& acc2,
mfma_acc_t& acc3, mfma_acc_t& acc4, mfma_acc_t& acc5) {
#if defined(__gfx950__)
asm volatile(
"s_setprio 1\n\t"
"v_mfma_scale_f32_16x16x128_f8f6f4 %[a0], %[af], %[b0], %[a0], %[sa], %[s0] cbsz:4 blgp:4\n\t"
"v_mfma_scale_f32_16x16x128_f8f6f4 %[a1], %[af], %[b1], %[a1], %[sa], %[s1] cbsz:4 blgp:4\n\t"
"v_mfma_scale_f32_16x16x128_f8f6f4 %[a2], %[af], %[b2], %[a2], %[sa], %[s2] cbsz:4 blgp:4\n\t"
"v_mfma_scale_f32_16x16x128_f8f6f4 %[a3], %[af], %[b3], %[a3], %[sa], %[s3] cbsz:4 blgp:4\n\t"
"v_mfma_scale_f32_16x16x128_f8f6f4 %[a4], %[af], %[b4], %[a4], %[sa], %[s4] cbsz:4 blgp:4\n\t"
"v_mfma_scale_f32_16x16x128_f8f6f4 %[a5], %[af], %[b5], %[a5], %[sa], %[s5] cbsz:4 blgp:4\n\t"
"s_setprio 0\n\t"
: [a0] "+v"(acc0), [a1] "+v"(acc1), [a2] "+v"(acc2),
[a3] "+v"(acc3), [a4] "+v"(acc4), [a5] "+v"(acc5)
: [af] "v"(af), [sa] "v"(sa),
[b0] "v"(bf0), [b1] "v"(bf1), [b2] "v"(bf2),
[b3] "v"(bf3), [b4] "v"(bf4), [b5] "v"(bf5),
[s0] "v"(sb0), [s1] "v"(sb1), [s2] "v"(sb2),
[s3] "v"(sb3), [s4] "v"(sb4), [s5] "v"(sb5)
: "memory"
);
#endif
}
// Inline ASM: 6 MFMAs interleaved with 6 dwordx4 NT B loads for hex
// MFMA captures inputs at issue (4 cycles), so bf VGPRs safe to reuse as load dests
__device__ __forceinline__ void
asm_mfma_hex_loadb(
u32x4 af, uint32_t sa,
u32x4& bf0, u32x4& bf1, u32x4& bf2, u32x4& bf3, u32x4& bf4, u32x4& bf5,
uint32_t sb0, uint32_t sb1, uint32_t sb2, uint32_t sb3, uint32_t sb4, uint32_t sb5,
mfma_acc_t& acc0, mfma_acc_t& acc1, mfma_acc_t& acc2,
mfma_acc_t& acc3, mfma_acc_t& acc4, mfma_acc_t& acc5,
const void* p0, const void* p1, const void* p2,
const void* p3, const void* p4, const void* p5) {
#if defined(__gfx950__)
asm volatile(
"s_setprio 1\n\t"
"v_mfma_scale_f32_16x16x128_f8f6f4 %[a0], %[af], %[b0], %[a0], %[sa], %[s0] cbsz:4 blgp:4\n\t"
"global_load_dwordx4 %[b0], %[p0], off nt\n\t"
"v_mfma_scale_f32_16x16x128_f8f6f4 %[a1], %[af], %[b1], %[a1], %[sa], %[s1] cbsz:4 blgp:4\n\t"
"global_load_dwordx4 %[b1], %[p1], off nt\n\t"
"v_mfma_scale_f32_16x16x128_f8f6f4 %[a2], %[af], %[b2], %[a2], %[sa], %[s2] cbsz:4 blgp:4\n\t"
"global_load_dwordx4 %[b2], %[p2], off nt\n\t"
"v_mfma_scale_f32_16x16x128_f8f6f4 %[a3], %[af], %[b3], %[a3], %[sa], %[s3] cbsz:4 blgp:4\n\t"
"global_load_dwordx4 %[b3], %[p3], off nt\n\t"
"v_mfma_scale_f32_16x16x128_f8f6f4 %[a4], %[af], %[b4], %[a4], %[sa], %[s4] cbsz:4 blgp:4\n\t"
"global_load_dwordx4 %[b4], %[p4], off nt\n\t"
"v_mfma_scale_f32_16x16x128_f8f6f4 %[a5], %[af], %[b5], %[a5], %[sa], %[s5] cbsz:4 blgp:4\n\t"
"global_load_dwordx4 %[b5], %[p5], off nt\n\t"
"s_setprio 0\n\t"
: [a0] "+v"(acc0), [a1] "+v"(acc1), [a2] "+v"(acc2),
[a3] "+v"(acc3), [a4] "+v"(acc4), [a5] "+v"(acc5),
[b0] "+v"(bf0), [b1] "+v"(bf1), [b2] "+v"(bf2),
[b3] "+v"(bf3), [b4] "+v"(bf4), [b5] "+v"(bf5)
: [af] "v"(af), [sa] "v"(sa),
[s0] "v"(sb0), [s1] "v"(sb1), [s2] "v"(sb2),
[s3] "v"(sb3), [s4] "v"(sb4), [s5] "v"(sb5),
[p0] "v"(p0), [p1] "v"(p1), [p2] "v"(p2),
[p3] "v"(p3), [p4] "v"(p4), [p5] "v"(p5)
: "memory"
);
#endif
}
#endif // UNIT_HEX
#if defined(UNIT_QUAD)
// Inline ASM: 4 MFMAs for quad kernel (no compiler waits)
__device__ __forceinline__ void
asm_mfma_quad(u32x4 af, uint32_t sa,
u32x4 bf0, u32x4 bf1, u32x4 bf2, u32x4 bf3,
uint32_t sb0, uint32_t sb1, uint32_t sb2, uint32_t sb3,
mfma_acc_t& acc0, mfma_acc_t& acc1, mfma_acc_t& acc2, mfma_acc_t& acc3) {
#if defined(__gfx950__)
asm volatile(
"s_setprio 1\n\t"
"v_mfma_scale_f32_16x16x128_f8f6f4 %[a0], %[af], %[b0], %[a0], %[sa], %[s0] cbsz:4 blgp:4\n\t"
"v_mfma_scale_f32_16x16x128_f8f6f4 %[a1], %[af], %[b1], %[a1], %[sa], %[s1] cbsz:4 blgp:4\n\t"
"v_mfma_scale_f32_16x16x128_f8f6f4 %[a2], %[af], %[b2], %[a2], %[sa], %[s2] cbsz:4 blgp:4\n\t"
"v_mfma_scale_f32_16x16x128_f8f6f4 %[a3], %[af], %[b3], %[a3], %[sa], %[s3] cbsz:4 blgp:4\n\t"
"s_setprio 0\n\t"
: [a0] "+v"(acc0), [a1] "+v"(acc1), [a2] "+v"(acc2), [a3] "+v"(acc3)
: [af] "v"(af), [sa] "v"(sa),
[b0] "v"(bf0), [b1] "v"(bf1), [b2] "v"(bf2), [b3] "v"(bf3),
[s0] "v"(sb0), [s1] "v"(sb1), [s2] "v"(sb2), [s3] "v"(sb3)
: "memory"
);
#endif
}
// Inline ASM: 4 MFMAs interleaved with 4 dwordx4 NT B loads
// MFMA captures inputs at issue (4 cycles), so bf VGPRs safe to reuse as load dests
__device__ __forceinline__ void
asm_mfma_quad_loadb(
u32x4 af, uint32_t sa,
u32x4& bf0, u32x4& bf1, u32x4& bf2, u32x4& bf3,
uint32_t sb0, uint32_t sb1, uint32_t sb2, uint32_t sb3,
mfma_acc_t& acc0, mfma_acc_t& acc1, mfma_acc_t& acc2, mfma_acc_t& acc3,
const void* p0, const void* p1, const void* p2, const void* p3) {
#if defined(__gfx950__)
asm volatile(
"s_setprio 1\n\t"
"v_mfma_scale_f32_16x16x128_f8f6f4 %[a0], %[af], %[b0], %[a0], %[sa], %[s0] cbsz:4 blgp:4\n\t"
"global_load_dwordx4 %[b0], %[p0], off nt\n\t"
"v_mfma_scale_f32_16x16x128_f8f6f4 %[a1], %[af], %[b1], %[a1], %[sa], %[s1] cbsz:4 blgp:4\n\t"
"global_load_dwordx4 %[b1], %[p1], off nt\n\t"
"v_mfma_scale_f32_16x16x128_f8f6f4 %[a2], %[af], %[b2], %[a2], %[sa], %[s2] cbsz:4 blgp:4\n\t"
"global_load_dwordx4 %[b2], %[p2], off nt\n\t"
"v_mfma_scale_f32_16x16x128_f8f6f4 %[a3], %[af], %[b3], %[a3], %[sa], %[s3] cbsz:4 blgp:4\n\t"
"global_load_dwordx4 %[b3], %[p3], off nt\n\t"
"s_setprio 0\n\t"
: [a0] "+v"(acc0), [a1] "+v"(acc1), [a2] "+v"(acc2), [a3] "+v"(acc3),
[b0] "+v"(bf0), [b1] "+v"(bf1), [b2] "+v"(bf2), [b3] "+v"(bf3)
: [af] "v"(af), [sa] "v"(sa),
[s0] "v"(sb0), [s1] "v"(sb1), [s2] "v"(sb2), [s3] "v"(sb3),
[p0] "v"(p0), [p1] "v"(p1), [p2] "v"(p2), [p3] "v"(p3)
: "memory"
);
#endif
}
#endif // UNIT_QUAD (quad ASM functions)
// Single dwordx4 NT load into u32x4 via ASM
__device__ __forceinline__ void
load_b_nt_u32x4(const uint8_t* __restrict__ B_shuffle, size_t base, int lid,
u32x4& bf) {
#if defined(__gfx950__)
const void* ptr = (const void*)(B_shuffle + base + (size_t)lid * 16);
asm volatile("global_load_dwordx4 %0, %1, off nt"
: "=v"(bf) : "v"(ptr) : "memory");
#endif
}
__device__ __forceinline__ uint32_t
read_bscale_sh(const uint8_t* __restrict__ B_scale_sh,
int nts, int pos, int grp, int sd4, int sp1d8) {
int row = nts + pos, col = sd4 * 4 + grp;
int r0 = row >> 5, r1 = (row >> 4) & 1, c0 = col >> 3, c1 = (col >> 2) & 1;
return (uint32_t)B_scale_sh[r0 * sp1d8 * 256 + c0 * 256 + grp * 64 + pos * 4 + c1 * 2 + r1];
}
__device__ __forceinline__ void
load_b_nt(const uint8_t* __restrict__ B_shuffle, size_t base, int lid,
mfma_input_t& bf) {
const uint32_t* bp = reinterpret_cast<const uint32_t*>(
B_shuffle + base + (size_t)lid * 16);
bf[0] = __builtin_nontemporal_load(bp + 0);
bf[1] = __builtin_nontemporal_load(bp + 1);
bf[2] = __builtin_nontemporal_load(bp + 2);
bf[3] = __builtin_nontemporal_load(bp + 3);
}
// Packed single-pass quant: uint4 loads + v_pk_max_u16 absmax (from v104)
__device__ __forceinline__ void
quant_bf16x32_to_fp4(const uint16_t* __restrict__ src, int offset,
uint8_t* __restrict__ dst, uint32_t& scale_out) {
const uint4* src4 = reinterpret_cast<const uint4*>(src + offset);
uint4 d[4];
#pragma unroll
for (int j = 0; j < 4; j++) d[j] = src4[j];
#if defined(__gfx950__)
uint32_t pk_max = 0;
#pragma unroll
for (int j = 0; j < 4; j++) {
uint32_t a0 = d[j].x & 0x7FFF7FFFu;
uint32_t a1 = d[j].y & 0x7FFF7FFFu;
uint32_t a2 = d[j].z & 0x7FFF7FFFu;
uint32_t a3 = d[j].w & 0x7FFF7FFFu;
asm volatile("v_pk_max_u16 %0, %1, %2" : "=v"(pk_max) : "v"(a0), "v"(pk_max));
asm volatile("v_pk_max_u16 %0, %1, %2" : "=v"(pk_max) : "v"(a1), "v"(pk_max));
asm volatile("v_pk_max_u16 %0, %1, %2" : "=v"(pk_max) : "v"(a2), "v"(pk_max));
asm volatile("v_pk_max_u16 %0, %1, %2" : "=v"(pk_max) : "v"(a3), "v"(pk_max));
}
uint32_t lo16 = pk_max & 0xFFFFu;
uint32_t hi16 = pk_max >> 16;
uint32_t max16;
asm volatile("v_max_u32 %0, %1, %2" : "=v"(max16) : "v"(lo16), "v"(hi16));
float absmax = __uint_as_float(max16 << 16);
#else
float absmax = 0.0f;
#pragma unroll
for (int j = 0; j < 4; j++) {
float v0, v1;
v0 = __uint_as_float((d[j].x & 0xFFFFu) << 16);
v1 = __uint_as_float(d[j].x & 0xFFFF0000u);
absmax = fmaxf(absmax, fmaxf(fabsf(v0), fabsf(v1)));
v0 = __uint_as_float((d[j].y & 0xFFFFu) << 16);
v1 = __uint_as_float(d[j].y & 0xFFFF0000u);
absmax = fmaxf(absmax, fmaxf(fabsf(v0), fabsf(v1)));
v0 = __uint_as_float((d[j].z & 0xFFFFu) << 16);
v1 = __uint_as_float(d[j].z & 0xFFFF0000u);
absmax = fmaxf(absmax, fmaxf(fabsf(v0), fabsf(v1)));
v0 = __uint_as_float((d[j].w & 0xFFFFu) << 16);
v1 = __uint_as_float(d[j].w & 0xFFFF0000u);
absmax = fmaxf(absmax, fmaxf(fabsf(v0), fabsf(v1)));
}
#endif
uint32_t bits = __float_as_uint(absmax);
uint32_t rounded = (bits + 0x200000u) & 0xFF800000u;
uint32_t E = (rounded >> 23) & 0xFFu;
uint32_t e8m0 = (E > 2u) ? (E - 2u) : 0u;
scale_out = e8m0;
#if defined(__gfx950__)
float scale_f = (e8m0 > 0u) ? __uint_as_float(e8m0 << 23) : 1.0f;
#pragma unroll
for (int j = 0; j < 4; j++) {
uint32_t pk0, pk1, pk2, pk3;
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
: "=v"(pk0) : "v"(d[j].x), "v"(scale_f));
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
: "=v"(pk1) : "v"(d[j].y), "v"(scale_f));
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
: "=v"(pk2) : "v"(d[j].z), "v"(scale_f));
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
: "=v"(pk3) : "v"(d[j].w), "v"(scale_f));
uint32_t lo, hi, packed;
asm volatile("v_perm_b32 %0, %1, %2, %3" : "=v"(lo) : "v"(pk1), "v"(pk0), "v"(0x0C0C0400u));
asm volatile("v_perm_b32 %0, %1, %2, %3" : "=v"(hi) : "v"(pk3), "v"(pk2), "v"(0x04000C0Cu));
asm volatile("v_or_b32 %0, %1, %2" : "=v"(packed) : "v"(lo), "v"(hi));
reinterpret_cast<uint32_t*>(dst)[j] = packed;
}
#else
(void)dst;
#endif
}
// Compute-only quant: takes pre-loaded d[4], produces fp4 + scale (for deep pipeline)
__device__ __forceinline__ void
quant_compute_fp4(uint4 d[4], uint8_t* __restrict__ dst, uint32_t& scale_out) {
#if defined(__gfx950__)
uint32_t pk_max = 0;
#pragma unroll
for (int j = 0; j < 4; j++) {
uint32_t a0 = d[j].x & 0x7FFF7FFFu;
uint32_t a1 = d[j].y & 0x7FFF7FFFu;
uint32_t a2 = d[j].z & 0x7FFF7FFFu;
uint32_t a3 = d[j].w & 0x7FFF7FFFu;
asm volatile("v_pk_max_u16 %0, %1, %2" : "=v"(pk_max) : "v"(a0), "v"(pk_max));
asm volatile("v_pk_max_u16 %0, %1, %2" : "=v"(pk_max) : "v"(a1), "v"(pk_max));
asm volatile("v_pk_max_u16 %0, %1, %2" : "=v"(pk_max) : "v"(a2), "v"(pk_max));
asm volatile("v_pk_max_u16 %0, %1, %2" : "=v"(pk_max) : "v"(a3), "v"(pk_max));
}
uint32_t lo16 = pk_max & 0xFFFFu;
uint32_t hi16 = pk_max >> 16;
uint32_t max16;
asm volatile("v_max_u32 %0, %1, %2" : "=v"(max16) : "v"(lo16), "v"(hi16));
float absmax = __uint_as_float(max16 << 16);
uint32_t bits = __float_as_uint(absmax);
uint32_t rounded = (bits + 0x200000u) & 0xFF800000u;
uint32_t E = (rounded >> 23) & 0xFFu;
uint32_t e8m0 = (E > 2u) ? (E - 2u) : 0u;
scale_out = e8m0;
float scale_f = (e8m0 > 0u) ? __uint_as_float(e8m0 << 23) : 1.0f;
#pragma unroll
for (int j = 0; j < 4; j++) {
uint32_t pk0, pk1, pk2, pk3;
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
: "=v"(pk0) : "v"(d[j].x), "v"(scale_f));
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
: "=v"(pk1) : "v"(d[j].y), "v"(scale_f));
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
: "=v"(pk2) : "v"(d[j].z), "v"(scale_f));
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
: "=v"(pk3) : "v"(d[j].w), "v"(scale_f));
uint32_t lo, hi, packed;
asm volatile("v_perm_b32 %0, %1, %2, %3" : "=v"(lo) : "v"(pk1), "v"(pk0), "v"(0x0C0C0400u));
asm volatile("v_perm_b32 %0, %1, %2, %3" : "=v"(hi) : "v"(pk3), "v"(pk2), "v"(0x04000C0Cu));
asm volatile("v_or_b32 %0, %1, %2" : "=v"(packed) : "v"(lo), "v"(hi));
reinterpret_cast<uint32_t*>(dst)[j] = packed;
}
#else
scale_out = 0x7F;
(void)dst;
#endif
}
// Compute-only quant returning u32x4 directly (no stack intermediary)
__device__ __forceinline__ void
quant_compute_u32x4(uint4 d[4], u32x4& af_out, uint32_t& scale_out) {
#if defined(__gfx950__)
uint32_t pk_max = 0;
#pragma unroll
for (int j = 0; j < 4; j++) {
uint32_t a0 = d[j].x & 0x7FFF7FFFu;
uint32_t a1 = d[j].y & 0x7FFF7FFFu;
uint32_t a2 = d[j].z & 0x7FFF7FFFu;
uint32_t a3 = d[j].w & 0x7FFF7FFFu;
asm volatile("v_pk_max_u16 %0, %1, %2" : "=v"(pk_max) : "v"(a0), "v"(pk_max));
asm volatile("v_pk_max_u16 %0, %1, %2" : "=v"(pk_max) : "v"(a1), "v"(pk_max));
asm volatile("v_pk_max_u16 %0, %1, %2" : "=v"(pk_max) : "v"(a2), "v"(pk_max));
asm volatile("v_pk_max_u16 %0, %1, %2" : "=v"(pk_max) : "v"(a3), "v"(pk_max));
}
uint32_t lo16 = pk_max & 0xFFFFu;
uint32_t hi16 = pk_max >> 16;
uint32_t max16;
asm volatile("v_max_u32 %0, %1, %2" : "=v"(max16) : "v"(lo16), "v"(hi16));
float absmax = __uint_as_float(max16 << 16);
uint32_t bits = __float_as_uint(absmax);
uint32_t rounded = (bits + 0x200000u) & 0xFF800000u;
uint32_t E = (rounded >> 23) & 0xFFu;
uint32_t e8m0 = (E > 2u) ? (E - 2u) : 0u;
scale_out = e8m0;
float scale_f = (e8m0 > 0u) ? __uint_as_float(e8m0 << 23) : 1.0f;
#pragma unroll
for (int j = 0; j < 4; j++) {
uint32_t pk0, pk1, pk2, pk3;
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
: "=v"(pk0) : "v"(d[j].x), "v"(scale_f));
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
: "=v"(pk1) : "v"(d[j].y), "v"(scale_f));
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
: "=v"(pk2) : "v"(d[j].z), "v"(scale_f));
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
: "=v"(pk3) : "v"(d[j].w), "v"(scale_f));
uint32_t lo, hi, packed;
asm volatile("v_perm_b32 %0, %1, %2, %3" : "=v"(lo) : "v"(pk1), "v"(pk0), "v"(0x0C0C0400u));
asm volatile("v_perm_b32 %0, %1, %2, %3" : "=v"(hi) : "v"(pk3), "v"(pk2), "v"(0x04000C0Cu));
asm volatile("v_or_b32 %0, %1, %2" : "=v"(packed) : "v"(lo), "v"(hi));
af_out[j] = packed;
}
#else
af_out = {0,0,0,0};
scale_out = 0x7F;
#endif
}
// Full quant from bf16 source, returns u32x4 directly
__device__ __forceinline__ void
quant_bf16x32_to_u32x4(const uint16_t* __restrict__ src, int offset,
u32x4& af_out, uint32_t& scale_out) {
const uint4* src4 = reinterpret_cast<const uint4*>(src + offset);
uint4 d[4];
#pragma unroll
for (int j = 0; j < 4; j++) d[j] = src4[j];
quant_compute_u32x4(d, af_out, scale_out);
}
// Single MFMA with u32x4 inputs (for narrow kernel)
#if defined(UNIT_NARROW)
__device__ __forceinline__ void
asm_mfma_single(u32x4 af, uint32_t sa, u32x4 bf, uint32_t sb,
mfma_acc_t& acc) {
#if defined(__gfx950__)
asm volatile(
"s_setprio 1\n\t"
"v_mfma_scale_f32_16x16x128_f8f6f4 %[a0], %[af], %[b0], %[a0], %[sa], %[s0] cbsz:4 blgp:4\n\t"
"s_setprio 0\n\t"
: [a0] "+v"(acc)
: [af] "v"(af), [sa] "v"(sa),
[b0] "v"(bf), [s0] "v"(sb)
: "memory"
);
#endif
}
#endif // UNIT_NARROW
// Inline ASM: 2 MFMAs (no compiler waits) — shared by wide + splitk
#if defined(UNIT_WIDE) || defined(UNIT_SPLITK)
__device__ __forceinline__ void
asm_mfma_pair(u32x4 af, uint32_t sa,
u32x4 bf0, u32x4 bf1,
uint32_t sb0, uint32_t sb1,
mfma_acc_t& acc0, mfma_acc_t& acc1) {
#if defined(__gfx950__)
asm volatile(
"s_setprio 1\n\t"
"v_mfma_scale_f32_16x16x128_f8f6f4 %[a0], %[af], %[b0], %[a0], %[sa], %[s0] cbsz:4 blgp:4\n\t"
"v_mfma_scale_f32_16x16x128_f8f6f4 %[a1], %[af], %[b1], %[a1], %[sa], %[s1] cbsz:4 blgp:4\n\t"
"s_setprio 0\n\t"
: [a0] "+v"(acc0), [a1] "+v"(acc1)
: [af] "v"(af), [sa] "v"(sa),
[b0] "v"(bf0), [b1] "v"(bf1),
[s0] "v"(sb0), [s1] "v"(sb1)
: "memory"
);
#endif
}
#endif // UNIT_WIDE || UNIT_SPLITK
#ifdef UNIT_NARROW
// ---- Shape 1: Narrow 16-wide fused (M=4) ----
__global__ void __launch_bounds__(256, 2)
mxfp4_fused_narrow(
const uint16_t* __restrict__ A_bf16,
const uint8_t* __restrict__ B_shuffle,
const uint8_t* __restrict__ B_scale_sh,
hip_bfloat16* __restrict__ output,
int M, int N, int K, int sp1d8
) {
constexpr int NW = 4, T = 16, MB = 64;
const int Kp = K / 2;
const int KPW = Kp / NW, KI = KPW / MB, KB = Kp / MB;
const int tid = threadIdx.x, wid = tid / 64, lid = tid % 64;
const int ntx = blockIdx.x, nts = ntx * T, mts = blockIdx.y * T;
if (nts >= N) return;
const int grp = lid / 16, pos = lid % 16, mr = mts + pos;
const bool valid = (mr < M);
mfma_acc_t acc = {0,0,0,0};
if (KI >= 2) {
int ki = 0;
int kb0 = wid*KI+ki, keo0 = wid*(K/NW)+ki*128;
int sd40 = keo0/(MXFP4_BLOCK_SIZE*4);
uint32_t sa0=0x7F;
u32x4 af0={0,0,0,0};
// Issue B load first, overlap with A quant
u32x4 bf0={0,0,0,0};
load_b_nt_u32x4(B_shuffle, (size_t)ntx*KB*1024+(size_t)kb0*1024, lid, bf0);
if (valid) {
quant_bf16x32_to_u32x4(A_bf16, mr*K+keo0+grp*32, af0, sa0);
}
asm volatile("s_waitcnt vmcnt(0)":::"memory");
uint32_t sb0 = read_bscale_sh(B_scale_sh, nts, pos, grp, sd40, sp1d8);
for (ki = 1; ki < KI; ki++) {
int kb1=wid*KI+ki, keo1=wid*(K/NW)+ki*128;
int sd41=keo1/(MXFP4_BLOCK_SIZE*4);
asm_mfma_single(af0, sa0, bf0, sb0, acc);
// Issue B load first, overlap with A quant
load_b_nt_u32x4(B_shuffle, (size_t)ntx*KB*1024+(size_t)kb1*1024, lid, bf0);
af0={0,0,0,0}; sa0=0x7F;
if (valid) {
quant_bf16x32_to_u32x4(A_bf16, mr*K+keo1+grp*32, af0, sa0);
}
asm volatile("s_waitcnt vmcnt(0)":::"memory");
sb0 = read_bscale_sh(B_scale_sh, nts, pos, grp, sd41, sp1d8);
}
asm_mfma_single(af0, sa0, bf0, sb0, acc);
} else if (KI == 1) {
int kb=wid, keo=wid*(K/NW), sd4=keo/(MXFP4_BLOCK_SIZE*4);
uint32_t sa=0x7F;
u32x4 af={0,0,0,0};
// Issue B load first, overlap with A quant
u32x4 bf={0,0,0,0};
load_b_nt_u32x4(B_shuffle, (size_t)ntx*KB*1024+(size_t)kb*1024, lid, bf);
if (valid) {
quant_bf16x32_to_u32x4(A_bf16, mr*K+keo+grp*32, af, sa);
}
asm volatile("s_waitcnt vmcnt(0)":::"memory");
uint32_t sb = read_bscale_sh(B_scale_sh, nts, pos, grp, sd4, sp1d8);
asm_mfma_single(af, sa, bf, sb, acc);
}
extern __shared__ uint8_t lds[];
float* lr = reinterpret_cast<float*>(lds);
#pragma unroll
for (int i=0;i<4;i++) lr[wid*T*T+(grp*4+i)*T+pos]=acc[i];
__syncthreads();
{
int ml=wid*4+grp, mg=mts+ml, ng=nts+pos;
if (mg<M&&ng<N) {
float s=0;
#pragma unroll
for (int w=0;w<NW;w++) s+=lr[w*T*T+ml*T+pos];
output[mg*N+ng]=hip_bfloat16(s);
}
}
}
#endif // UNIT_NARROW
#ifdef UNIT_WIDE
// ---- Shapes 3,4: Wide 32-N fused (M=32) ----
__global__ void __launch_bounds__(256, 2)
mxfp4_fused_wide(
const uint16_t* __restrict__ A_bf16,
const uint8_t* __restrict__ B_shuffle,
const uint8_t* __restrict__ B_scale_sh,
hip_bfloat16* __restrict__ output,
int M, int N, int K, int sp1d8
) {
constexpr int NW = 4, T = 16, MB = 64;
const int Kp = K / 2;
const int KPW = Kp / NW, KI = KPW / MB, KB = Kp / MB;
const int tid = threadIdx.x, wid = tid / 64, lid = tid % 64;
const int ntx0 = blockIdx.x * 2, ntx1 = ntx0 + 1;
const int nts0 = ntx0 * T, nts1 = ntx1 * T;
const int mts = blockIdx.y * T;
if (nts0 >= N) return;
const bool has_sub1 = (nts1 < N);
const int grp = lid / 16, pos = lid % 16, mr = mts + pos;
const bool valid = (mr < M);
mfma_acc_t acc0 = {0,0,0,0}, acc1 = {0,0,0,0};
if (KI >= 2) {
int ki = 0;
int kb0 = wid*KI+ki, keo0 = wid*(K/NW)+ki*128;
int sd40 = keo0/(MXFP4_BLOCK_SIZE*4);
uint32_t sa0=0x7F;
u32x4 af0={0,0,0,0};
if (valid) {
quant_bf16x32_to_u32x4(A_bf16, mr*K+keo0+grp*32, af0, sa0);
}
u32x4 bfa={0,0,0,0}, bfb={0,0,0,0};
load_b_nt_u32x4(B_shuffle, (size_t)ntx0*KB*1024+(size_t)kb0*1024, lid, bfa);
uint32_t sba = read_bscale_sh(B_scale_sh, nts0, pos, grp, sd40, sp1d8);
uint32_t sbb = 0x7F;
if (has_sub1) {
load_b_nt_u32x4(B_shuffle, (size_t)ntx1*KB*1024+(size_t)kb0*1024, lid, bfb);
sbb = read_bscale_sh(B_scale_sh, nts1, pos, grp, sd40, sp1d8);
}
for (ki = 1; ki < KI; ki++) {
int kb1=wid*KI+ki, keo1=wid*(K/NW)+ki*128;
int sd41=keo1/(MXFP4_BLOCK_SIZE*4);
asm_mfma_pair(af0, sa0, bfa, bfb, sba, sbb, acc0, acc1);
af0={0,0,0,0}; sa0=0x7F;
if (valid) {
quant_bf16x32_to_u32x4(A_bf16, mr*K+keo1+grp*32, af0, sa0);
}
load_b_nt_u32x4(B_shuffle, (size_t)ntx0*KB*1024+(size_t)kb1*1024, lid, bfa);
if (has_sub1) load_b_nt_u32x4(B_shuffle, (size_t)ntx1*KB*1024+(size_t)kb1*1024, lid, bfb);
asm volatile("s_waitcnt vmcnt(0)":::"memory");
sba = read_bscale_sh(B_scale_sh, nts0, pos, grp, sd41, sp1d8);
if (has_sub1) sbb = read_bscale_sh(B_scale_sh, nts1, pos, grp, sd41, sp1d8);
}
asm_mfma_pair(af0, sa0, bfa, bfb, sba, sbb, acc0, acc1);
} else if (KI == 1) {
int kb=wid, keo=wid*(K/NW), sd4=keo/(MXFP4_BLOCK_SIZE*4);
uint32_t sa=0x7F;
u32x4 af={0,0,0,0};
// Issue both B loads upfront using u32x4 (1 vmcnt each, overlap with quant)
u32x4 bf0={0,0,0,0}, bf1={0,0,0,0};
load_b_nt_u32x4(B_shuffle, (size_t)ntx0*KB*1024+(size_t)kb*1024, lid, bf0);
if (has_sub1) load_b_nt_u32x4(B_shuffle, (size_t)ntx1*KB*1024+(size_t)kb*1024, lid, bf1);
if (valid) {
quant_bf16x32_to_u32x4(A_bf16, mr*K+keo+grp*32, af, sa);
}
asm volatile("s_waitcnt vmcnt(0)":::"memory");
uint32_t sb0 = read_bscale_sh(B_scale_sh, nts0, pos, grp, sd4, sp1d8);
uint32_t sb1 = 0x7F;
if (has_sub1) sb1 = read_bscale_sh(B_scale_sh, nts1, pos, grp, sd4, sp1d8);
// ASM MFMA pair: both N-tiles in one block
asm_mfma_pair(af, sa, bf0, bf1, sb0, sb1, acc0, acc1);
}
extern __shared__ uint8_t lds[];
float* lr = reinterpret_cast<float*>(lds);
#pragma unroll
for (int i=0;i<4;i++) lr[wid*T*T+(grp*4+i)*T+pos]=acc0[i];
#pragma unroll
for (int i=0;i<4;i++) lr[NW*T*T+wid*T*T+(grp*4+i)*T+pos]=acc1[i];
__syncthreads();
{
int ml=wid*4+grp, mg=mts+ml, ng=nts0+pos;
if (mg<M&&ng<N) {
float s=0;
#pragma unroll
for (int w=0;w<NW;w++) s+=lr[w*T*T+ml*T+pos];
output[mg*N+ng]=hip_bfloat16(s);
}
if (has_sub1) {
ng=nts1+pos;
if (mg<M&&ng<N) {
float s=0;
#pragma unroll
for (int w=0;w<NW;w++) s+=lr[NW*T*T+w*T*T+ml*T+pos];
output[mg*N+ng]=hip_bfloat16(s);
}
}
}
}
#endif // UNIT_WIDE
#ifdef UNIT_QUAD
// ---- Shape 5: Quad 64-N fused (M=64, 4 N-tiles per CTA) with deep pipeline ----
__global__ void __launch_bounds__(256, 2)
mxfp4_fused_quad(
const uint16_t* __restrict__ A_bf16,
const uint8_t* __restrict__ B_shuffle,
const uint8_t* __restrict__ B_scale_sh,
hip_bfloat16* __restrict__ output,
int M, int N, int K, int sp1d8
) {
constexpr int NW = 4, T = 16, MB = 64;
const int Kp = K / 2;
const int KPW = Kp / NW, KI = KPW / MB, KB = Kp / MB;
const int tid = threadIdx.x, wid = tid / 64, lid = tid % 64;
const int ntx0 = blockIdx.x * 4;
const int nts0 = ntx0 * T, nts1 = (ntx0+1) * T, nts2 = (ntx0+2) * T, nts3 = (ntx0+3) * T;
const int mts = blockIdx.y * T;
if (nts0 >= N) return;
const bool has1 = (nts1 < N), has2 = (nts2 < N), has3 = (nts3 < N);
const int grp = lid / 16, pos = lid % 16, mr = mts + pos;
const bool valid = (mr < M);
mfma_acc_t acc0 = {0,0,0,0}, acc1 = {0,0,0,0}, acc2 = {0,0,0,0}, acc3 = {0,0,0,0};
if (KI >= 2 && has3) {
// === FAST PATH: interleaved MFMA+B loads (all 4 tiles present) ===
int ki = 0;
int kb0 = wid*KI+ki, keo0 = wid*(K/NW)+ki*128;
int sd40 = keo0/(MXFP4_BLOCK_SIZE*4);
uint32_t sa0=0x7F;
u32x4 af0={0,0,0,0};
// Prologue: Load A[0]
uint4 ad[4] = {};
if (valid) {
const uint4* a_src = reinterpret_cast<const uint4*>(A_bf16 + mr*K+keo0+grp*32);
#pragma unroll
for (int j = 0; j < 4; j++) ad[j] = a_src[j];
}
// Load B[0] for all 4 tiles via ASM
u32x4 bf0_u={0,0,0,0}, bf1_u={0,0,0,0}, bf2_u={0,0,0,0}, bf3_u={0,0,0,0};
load_b_nt_u32x4(B_shuffle, (size_t)ntx0*KB*1024+(size_t)kb0*1024, lid, bf0_u);
load_b_nt_u32x4(B_shuffle, (size_t)(ntx0+1)*KB*1024+(size_t)kb0*1024, lid, bf1_u);
load_b_nt_u32x4(B_shuffle, (size_t)(ntx0+2)*KB*1024+(size_t)kb0*1024, lid, bf2_u);
load_b_nt_u32x4(B_shuffle, (size_t)(ntx0+3)*KB*1024+(size_t)kb0*1024, lid, bf3_u);
// Wait for A (oldest, vmcnt(4) = 4 B loads still in flight)
asm volatile("s_waitcnt vmcnt(4)":::"memory");
if (valid) {
quant_compute_u32x4(ad, af0, sa0);
}
asm volatile("s_waitcnt vmcnt(0)":::"memory");
uint32_t sba = read_bscale_sh(B_scale_sh, nts0, pos, grp, sd40, sp1d8);
uint32_t sbb = read_bscale_sh(B_scale_sh, nts1, pos, grp, sd40, sp1d8);
uint32_t sbc = read_bscale_sh(B_scale_sh, nts2, pos, grp, sd40, sp1d8);
uint32_t sbd = read_bscale_sh(B_scale_sh, nts3, pos, grp, sd40, sp1d8);
// Pre-load A[1] for double-buffer
uint4 ad_next[4] = {};
if (valid && KI > 2) {
int keo_next = wid*(K/NW)+1*128;
const uint4* a_src_next = reinterpret_cast<const uint4*>(A_bf16 + mr*K+keo_next+grp*32);
#pragma unroll
for (int j = 0; j < 4; j++) ad_next[j] = a_src_next[j];
}
for (ki = 1; ki < KI; ki++) {
int kb1=wid*KI+ki, keo1=wid*(K/NW)+ki*128;
int sd41=keo1/(MXFP4_BLOCK_SIZE*4);
// Compute B pointers for next iteration loads
const void* p0 = (const void*)(B_shuffle + (size_t)ntx0*KB*1024 + (size_t)kb1*1024 + (size_t)lid*16);
const void* p1 = (const void*)(B_shuffle + (size_t)(ntx0+1)*KB*1024 + (size_t)kb1*1024 + (size_t)lid*16);
const void* p2 = (const void*)(B_shuffle + (size_t)(ntx0+2)*KB*1024 + (size_t)kb1*1024 + (size_t)lid*16);
const void* p3 = (const void*)(B_shuffle + (size_t)(ntx0+3)*KB*1024 + (size_t)kb1*1024 + (size_t)lid*16);
// Interleaved MFMA(old data) + B loads(new addresses)
asm_mfma_quad_loadb(af0, sa0, bf0_u, bf1_u, bf2_u, bf3_u,
sba, sbb, sbc, sbd,
acc0, acc1, acc2, acc3,
p0, p1, p2, p3);
// Wait for pre-loaded A[ki]
asm volatile("s_waitcnt vmcnt(4)":::"memory");
ad[0] = ad_next[0]; ad[1] = ad_next[1]; ad[2] = ad_next[2]; ad[3] = ad_next[3];
// Quant A[ki] (VALU only, overlaps with B loads in flight)
af0 = {0,0,0,0}; sa0 = 0x7F;
if (valid) quant_compute_u32x4(ad, af0, sa0);
// Issue Bscale reads early (overlap with remaining B loads)
sba = read_bscale_sh(B_scale_sh, nts0, pos, grp, sd41, sp1d8);
sbb = read_bscale_sh(B_scale_sh, nts1, pos, grp, sd41, sp1d8);
sbc = read_bscale_sh(B_scale_sh, nts2, pos, grp, sd41, sp1d8);
sbd = read_bscale_sh(B_scale_sh, nts3, pos, grp, sd41, sp1d8);
// Pre-load A[ki+1] (issued AFTER bscale so vmcnt(4) retains only A loads)
ad_next[0] = {}; ad_next[1] = {}; ad_next[2] = {}; ad_next[3] = {};
if (valid && ki + 1 < KI) {
int keo_next = wid*(K/NW) + (ki+1)*128;
const uint4* a_src_next = reinterpret_cast<const uint4*>(A_bf16 + mr*K+keo_next+grp*32);
#pragma unroll
for (int j = 0; j < 4; j++) ad_next[j] = a_src_next[j];
}
// Wait for B[ki] + Bscale (4 B + 4 bscale drained, 4 A newest remain)
if (ki + 1 < KI) {
asm volatile("s_waitcnt vmcnt(4)":::"memory");
} else {
asm volatile("s_waitcnt vmcnt(0)":::"memory");
}
}
// Epilogue: final MFMAs (no B loads)
asm_mfma_quad(af0, sa0, bf0_u, bf1_u, bf2_u, bf3_u,
sba, sbb, sbc, sbd, acc0, acc1, acc2, acc3);
} else if (KI >= 2) {
// Fallback for edge tiles (has3=false)
int ki = 0;
int kb0 = wid*KI+ki, keo0 = wid*(K/NW)+ki*128;
int sd40 = keo0/(MXFP4_BLOCK_SIZE*4);
uint8_t ap[16]; uint32_t sa0=0x7F;
mfma_input_t af0={0,0,0,0,0,0,0,0};
uint4 ad[4] = {};
if (valid) {
const uint4* a_src = reinterpret_cast<const uint4*>(A_bf16 + mr*K+keo0+grp*32);
#pragma unroll
for (int j = 0; j < 4; j++) ad[j] = a_src[j];
}
mfma_input_t bfa={0,0,0,0,0,0,0,0}, bfb={0,0,0,0,0,0,0,0};
mfma_input_t bfc={0,0,0,0,0,0,0,0}, bfd={0,0,0,0,0,0,0,0};
load_b_nt(B_shuffle, (size_t)ntx0*KB*1024+(size_t)kb0*1024, lid, bfa);
if (has1) load_b_nt(B_shuffle, (size_t)(ntx0+1)*KB*1024+(size_t)kb0*1024, lid, bfb);
if (has2) load_b_nt(B_shuffle, (size_t)(ntx0+2)*KB*1024+(size_t)kb0*1024, lid, bfc);
asm volatile("s_waitcnt vmcnt(0)":::"memory");
if (valid) {
quant_compute_fp4(ad, ap, sa0);
uint4 c=*reinterpret_cast<const uint4*>(ap);
af0[0]=c.x;af0[1]=c.y;af0[2]=c.z;af0[3]=c.w;
}
uint32_t sba = read_bscale_sh(B_scale_sh, nts0, pos, grp, sd40, sp1d8);
uint32_t sbb = 0x7F, sbc = 0x7F;
if (has1) sbb = read_bscale_sh(B_scale_sh, nts1, pos, grp, sd40, sp1d8);
if (has2) sbc = read_bscale_sh(B_scale_sh, nts2, pos, grp, sd40, sp1d8);
for (ki = 1; ki < KI; ki++) {
int kb1=wid*KI+ki, keo1=wid*(K/NW)+ki*128;
int sd41=keo1/(MXFP4_BLOCK_SIZE*4);
ad[0] = {}; ad[1] = {}; ad[2] = {}; ad[3] = {};
if (valid) {
const uint4* a_src = reinterpret_cast<const uint4*>(A_bf16 + mr*K+keo1+grp*32);
#pragma unroll
for (int j = 0; j < 4; j++) ad[j] = a_src[j];
}
{
u32x4 af_u4 = {af0[0], af0[1], af0[2], af0[3]};
u32x4 bfu0 = {bfa[0], bfa[1], bfa[2], bfa[3]};
u32x4 bfu1 = {bfb[0], bfb[1], bfb[2], bfb[3]};
u32x4 bfu2 = {bfc[0], bfc[1], bfc[2], bfc[3]};
u32x4 zero4 = {0,0,0,0};
asm_mfma_quad(af_u4, sa0, bfu0, bfu1, bfu2, zero4,
sba, sbb, sbc, 0x7F, acc0, acc1, acc2, acc3);
}
load_b_nt(B_shuffle, (size_t)ntx0*KB*1024+(size_t)kb1*1024, lid, bfa);
if (has1) load_b_nt(B_shuffle, (size_t)(ntx0+1)*KB*1024+(size_t)kb1*1024, lid, bfb);
if (has2) load_b_nt(B_shuffle, (size_t)(ntx0+2)*KB*1024+(size_t)kb1*1024, lid, bfc);
asm volatile("s_waitcnt vmcnt(0)":::"memory");
af0={0,0,0,0,0,0,0,0}; sa0=0x7F;
if (valid) {
quant_compute_fp4(ad, ap, sa0);
uint4 c=*reinterpret_cast<const uint4*>(ap);
af0[0]=c.x;af0[1]=c.y;af0[2]=c.z;af0[3]=c.w;
}
sba = read_bscale_sh(B_scale_sh, nts0, pos, grp, sd41, sp1d8);
if (has1) sbb = read_bscale_sh(B_scale_sh, nts1, pos, grp, sd41, sp1d8);
if (has2) sbc = read_bscale_sh(B_scale_sh, nts2, pos, grp, sd41, sp1d8);
}
{
u32x4 af_u4 = {af0[0], af0[1], af0[2], af0[3]};
u32x4 bfu0 = {bfa[0], bfa[1], bfa[2], bfa[3]};
u32x4 bfu1 = {bfb[0], bfb[1], bfb[2], bfb[3]};
u32x4 bfu2 = {bfc[0], bfc[1], bfc[2], bfc[3]};
u32x4 zero4 = {0,0,0,0};
asm_mfma_quad(af_u4, sa0, bfu0, bfu1, bfu2, zero4,
sba, sbb, sbc, 0x7F, acc0, acc1, acc2, acc3);
}
} else if (KI == 1) {
int kb=wid, keo=wid*(K/NW), sd4=keo/(MXFP4_BLOCK_SIZE*4);
uint8_t ap[16]; uint32_t sa=0x7F;
mfma_input_t af={0,0,0,0,0,0,0,0};
if (valid) {
quant_bf16x32_to_fp4(A_bf16, mr*K+keo+grp*32, ap, sa);
uint4 c=*reinterpret_cast<const uint4*>(ap);
af[0]=c.x;af[1]=c.y;af[2]=c.z;af[3]=c.w;
}
mfma_input_t bf0={0,0,0,0,0,0,0,0};
load_b_nt(B_shuffle, (size_t)ntx0*KB*1024+(size_t)kb*1024, lid, bf0);
uint32_t sb0 = read_bscale_sh(B_scale_sh, nts0, pos, grp, sd4, sp1d8);
acc0 = mfma_fp4_16x16(af, bf0, acc0, sa, sb0);
if (has1) {
mfma_input_t bf1={0,0,0,0,0,0,0,0};
load_b_nt(B_shuffle, (size_t)(ntx0+1)*KB*1024+(size_t)kb*1024, lid, bf1);
uint32_t sb1 = read_bscale_sh(B_scale_sh, nts1, pos, grp, sd4, sp1d8);
acc1 = mfma_fp4_16x16(af, bf1, acc1, sa, sb1);
}
if (has2) {
mfma_input_t bf2={0,0,0,0,0,0,0,0};
load_b_nt(B_shuffle, (size_t)(ntx0+2)*KB*1024+(size_t)kb*1024, lid, bf2);
uint32_t sb2 = read_bscale_sh(B_scale_sh, nts2, pos, grp, sd4, sp1d8);
acc2 = mfma_fp4_16x16(af, bf2, acc2, sa, sb2);
}
if (has3) {
mfma_input_t bf3={0,0,0,0,0,0,0,0};
load_b_nt(B_shuffle, (size_t)(ntx0+3)*KB*1024+(size_t)kb*1024, lid, bf3);
uint32_t sb3 = read_bscale_sh(B_scale_sh, nts3, pos, grp, sd4, sp1d8);
acc3 = mfma_fp4_16x16(af, bf3, acc3, sa, sb3);
}
}
// Warp reduction via LDS: 4 N-tiles x 4 warps
extern __shared__ uint8_t lds_quad[];
float* lr = reinterpret_cast<float*>(lds_quad);
constexpr int NW2 = 4, T2 = 16;
#pragma unroll
for (int i=0;i<4;i++) lr[wid*T2*T2+(grp*4+i)*T2+pos]=acc0[i];
#pragma unroll
for (int i=0;i<4;i++) lr[NW2*T2*T2+wid*T2*T2+(grp*4+i)*T2+pos]=acc1[i];
#pragma unroll
for (int i=0;i<4;i++) lr[2*NW2*T2*T2+wid*T2*T2+(grp*4+i)*T2+pos]=acc2[i];
#pragma unroll
for (int i=0;i<4;i++) lr[3*NW2*T2*T2+wid*T2*T2+(grp*4+i)*T2+pos]=acc3[i];
__syncthreads();
// Parallel output: warp w handles tile w
{
const int nts_w[4] = {nts0, nts1, nts2, nts3};
const bool has_w[4] = {true, has1, has2, has3};
if (has_w[wid]) {
#pragma unroll
for (int i=0;i<4;i++) {
int ml=grp*4+i,mg=mts+ml;
if (mg<M) {
int ng=nts_w[wid]+pos;
if (ng<N) { float s=0; for (int w=0;w<NW2;w++) s+=lr[wid*NW2*T2*T2+w*T2*T2+ml*T2+pos]; output[mg*N+ng]=hip_bfloat16(s); }
}
}
}
}
}
#endif // UNIT_QUAD
#ifdef UNIT_HEX
// ---- Shape 6: Hex 96-N fused (M>=128, 6 N-tiles per CTA) with deep pipeline ----
__global__ void __launch_bounds__(256, 2)
mxfp4_fused_hex(
const uint16_t* __restrict__ A_bf16,
const uint8_t* __restrict__ B_shuffle,
const uint8_t* __restrict__ B_scale_sh,
hip_bfloat16* __restrict__ output,
int M, int N, int K, int sp1d8
) {
constexpr int NW = 4, T = 16, MB = 64;
const int Kp = K / 2;
const int KPW = Kp / NW, KI = KPW / MB, KB = Kp / MB;
const int tid = threadIdx.x, wid = tid / 64, lid = tid % 64;
const int ntx0 = blockIdx.x * 6;
const int nts[6] = {ntx0*T, (ntx0+1)*T, (ntx0+2)*T, (ntx0+3)*T, (ntx0+4)*T, (ntx0+5)*T};
const int mts = blockIdx.y * T;
if (nts[0] >= N) return;
const bool has[6] = {true, (nts[1]<N), (nts[2]<N), (nts[3]<N), (nts[4]<N), (nts[5]<N)};
const int grp = lid / 16, pos = lid % 16, mr = mts + pos;
const bool valid = (mr < M);
mfma_acc_t acc[6] = {{0,0,0,0},{0,0,0,0},{0,0,0,0},{0,0,0,0},{0,0,0,0},{0,0,0,0}};
if (KI >= 2) {
int ki = 0;
int kb0 = wid*KI+ki, keo0 = wid*(K/NW)+ki*128;
int sd40 = keo0/(MXFP4_BLOCK_SIZE*4);
uint32_t sa0=0x7F;
u32x4 af0={0,0,0,0};
// Prologue: Load A data into registers (oldest in vmcnt)
uint4 ad[4] = {};
if (valid) {
const uint4* a_src = reinterpret_cast<const uint4*>(A_bf16 + mr*K+keo0+grp*32);
#pragma unroll
for (int j = 0; j < 4; j++) ad[j] = a_src[j];
}
// Issue B loads for all 6 tiles via u32x4 ASM (1 vmcnt each)
u32x4 bf[6] = {{0,0,0,0},{0,0,0,0},{0,0,0,0},{0,0,0,0},{0,0,0,0},{0,0,0,0}};
load_b_nt_u32x4(B_shuffle, (size_t)ntx0*KB*1024+(size_t)kb0*1024, lid, bf[0]);
#pragma unroll
for (int t = 1; t < 6; t++)
if (has[t]) load_b_nt_u32x4(B_shuffle, (size_t)(ntx0+t)*KB*1024+(size_t)kb0*1024, lid, bf[t]);
// Wait for A only (oldest), B loads overlap with quant ALU
if (has[5]) {
asm volatile("s_waitcnt vmcnt(6)":::"memory");
} else {
asm volatile("s_waitcnt vmcnt(0)":::"memory");
}
if (valid) quant_compute_u32x4(ad, af0, sa0);
asm volatile("s_waitcnt vmcnt(0)":::"memory");
uint32_t sb[6];
sb[0] = read_bscale_sh(B_scale_sh, nts[0], pos, grp, sd40, sp1d8);
#pragma unroll
for (int t = 1; t < 6; t++) {
sb[t] = 0x7F;
if (has[t]) sb[t] = read_bscale_sh(B_scale_sh, nts[t], pos, grp, sd40, sp1d8);
}
// Pre-load A[1] for double-buffer
uint4 ad_next[4] = {};
if (valid && KI > 2) {
int keo_next = wid*(K/NW)+1*128;
const uint4* a_src_next = reinterpret_cast<const uint4*>(A_bf16 + mr*K+keo_next+grp*32);
#pragma unroll
for (int j = 0; j < 4; j++) ad_next[j] = a_src_next[j];
}
for (ki = 1; ki < KI; ki++) {
int kb1=wid*KI+ki, keo1=wid*(K/NW)+ki*128;
int sd41=keo1/(MXFP4_BLOCK_SIZE*4);
// Compute B pointers for interleaved loads
const void* p[6];
p[0] = (const void*)(B_shuffle + (size_t)ntx0*KB*1024 + (size_t)kb1*1024 + (size_t)lid*16);
#pragma unroll
for (int t = 1; t < 6; t++)
p[t] = (const void*)(B_shuffle + (size_t)(ntx0+t)*KB*1024 + (size_t)kb1*1024 + (size_t)lid*16);
// Interleaved MFMA(old data) + B loads(new addresses)
asm_mfma_hex_loadb(af0, sa0,
bf[0], bf[1], bf[2], bf[3], bf[4], bf[5],
sb[0], sb[1], sb[2], sb[3], sb[4], sb[5],
acc[0], acc[1], acc[2], acc[3], acc[4], acc[5],
p[0], p[1], p[2], p[3], p[4], p[5]);
// Wait for pre-loaded A[ki] (4 oldest vmcnt entries behind 6 B loads)
asm volatile("s_waitcnt vmcnt(6)":::"memory");
ad[0] = ad_next[0]; ad[1] = ad_next[1]; ad[2] = ad_next[2]; ad[3] = ad_next[3];
// Pre-load A[ki+1] if not last iteration
ad_next[0] = {}; ad_next[1] = {}; ad_next[2] = {}; ad_next[3] = {};
if (valid && ki + 1 < KI) {
int keo_next = wid*(K/NW) + (ki+1)*128;
const uint4* a_src_next = reinterpret_cast<const uint4*>(A_bf16 + mr*K+keo_next+grp*32);
#pragma unroll
for (int j = 0; j < 4; j++) ad_next[j] = a_src_next[j];
}
// Quant A[ki]
af0 = {0,0,0,0}; sa0 = 0x7F;
if (valid) quant_compute_u32x4(ad, af0, sa0);
// Wait for B[ki]
if (ki + 1 < KI) {
asm volatile("s_waitcnt vmcnt(4)":::"memory");
} else {
asm volatile("s_waitcnt vmcnt(0)":::"memory");
}
sb[0] = read_bscale_sh(B_scale_sh, nts[0], pos, grp, sd41, sp1d8);
#pragma unroll
for (int t = 1; t < 6; t++)
if (has[t]) sb[t] = read_bscale_sh(B_scale_sh, nts[t], pos, grp, sd41, sp1d8);
}
// Epilogue: final MFMAs (no B loads)
asm_mfma_hex(af0, sa0,
bf[0], bf[1], bf[2], bf[3], bf[4], bf[5],
sb[0], sb[1], sb[2], sb[3], sb[4], sb[5],
acc[0], acc[1], acc[2], acc[3], acc[4], acc[5]);
} else if (KI == 1) {
int kb=wid, keo=wid*(K/NW), sd4=keo/(MXFP4_BLOCK_SIZE*4);
uint8_t ap[16]; uint32_t sa=0x7F;
mfma_input_t af={0,0,0,0,0,0,0,0};
if (valid) {
quant_bf16x32_to_fp4(A_bf16, mr*K+keo+grp*32, ap, sa);
uint4 c=*reinterpret_cast<const uint4*>(ap);
af[0]=c.x;af[1]=c.y;af[2]=c.z;af[3]=c.w;
}
#pragma unroll
for (int t = 0; t < 6; t++) {
if (has[t]) {
mfma_input_t bft={0,0,0,0,0,0,0,0};
load_b_nt(B_shuffle, (size_t)(ntx0+t)*KB*1024+(size_t)kb*1024, lid, bft);
uint32_t sbt = read_bscale_sh(B_scale_sh, nts[t], pos, grp, sd4, sp1d8);
acc[t] = mfma_fp4_16x16(af, bft, acc[t], sa, sbt);
}
}
}
// Warp reduction via LDS: 6 N-tiles x 4 warps
extern __shared__ uint8_t lds_hex[];
float* lr = reinterpret_cast<float*>(lds_hex);
constexpr int NW2 = 4, T2 = 16;
#pragma unroll
for (int t = 0; t < 6; t++) {
#pragma unroll
for (int i=0;i<4;i++) lr[t*NW2*T2*T2+wid*T2*T2+(grp*4+i)*T2+pos]=acc[t][i];
}
__syncthreads();
// Parallel reduction: warp 0 handles tiles 0,1; warp 1 handles 2,3; warp 2 handles 4,5
if (wid <= 2) {
int t0 = wid * 2, t1 = t0 + 1;
#pragma unroll
for (int i=0;i<4;i++) {
int ml=grp*4+i,mg=mts+ml;
if (mg<M) {
int ng0=nts[t0]+pos;
if (has[t0] && ng0<N) { float s=0; for (int w=0;w<NW2;w++) s+=lr[t0*NW2*T2*T2+w*T2*T2+ml*T2+pos]; output[mg*N+ng0]=hip_bfloat16(s); }
int ng1=nts[t1]+pos;
if (has[t1] && ng1<N) { float s=0; for (int w=0;w<NW2;w++) s+=lr[t1*NW2*T2*T2+w*T2*T2+ml*T2+pos]; output[mg*N+ng1]=hip_bfloat16(s); }
}
}
}
}
#endif // UNIT_HEX
#ifdef UNIT_SPLITK
// ---- Shape 2: Fused quant + wide split-K GEMM with 2-stage deep pipeline (from v104) ----
__global__ void __launch_bounds__(256, 2)
mxfp4_gemm_splitk_fused(
const uint16_t* __restrict__ A_bf16,
const uint8_t* __restrict__ B_shuffle,
const uint8_t* __restrict__ B_scale_sh,
float* __restrict__ partial_out,
int M, int N, int K, int sp1d8, int ksplit
) {
constexpr int NW = 4, T = 16, MB = 64;
const int kslice = blockIdx.z;
const int Kp = K / 2;
const int KB = Kp / MB;
const int K_slice = K / ksplit;
const int Kp_slice = Kp / ksplit;
const int KPW = Kp_slice / NW, KI = KPW / MB;
const int ke_base = kslice * K_slice;
const int kb_base = kslice * (Kp_slice / MB);
const int tid = threadIdx.x, wid = tid / 64, lid = tid % 64;
const int ntx0 = blockIdx.x * 2, ntx1 = ntx0 + 1;
const int nts0 = ntx0 * T, nts1 = ntx1 * T;
const int mts = blockIdx.y * T;
if (nts0 >= N) return;
const bool has_sub1 = (nts1 < N);
const int grp = lid / 16, pos = lid % 16, mr = mts + pos;
const bool valid = (mr < M);
mfma_acc_t acc0 = {0,0,0,0}, acc1 = {0,0,0,0};
if (KI == 2) {
// ============ 2-STAGE DEEP PIPELINE for KI==2 ============
int kb0 = kb_base + wid*2;
int kb1 = kb0 + 1;
int keo0 = ke_base + wid*(K_slice/NW);
int keo1 = keo0 + 128;
int sd40 = keo0/(MXFP4_BLOCK_SIZE*4);
int sd41 = keo1/(MXFP4_BLOCK_SIZE*4);
// Load A data for BOTH iterations upfront (oldest in vmcnt)
uint4 ad0[4] = {}, ad1[4] = {};
if (valid) {
const uint4* a_src0 = reinterpret_cast<const uint4*>(A_bf16 + mr*K + keo0 + grp*32);
#pragma unroll
for (int j = 0; j < 4; j++) ad0[j] = a_src0[j];
const uint4* a_src1 = reinterpret_cast<const uint4*>(A_bf16 + mr*K + keo1 + grp*32);
#pragma unroll
for (int j = 0; j < 4; j++) ad1[j] = a_src1[j];
}
// Pre-read all Bscale values upfront (avoids compiler vmcnt interference later)
uint32_t sba0 = read_bscale_sh(B_scale_sh, nts0, pos, grp, sd40, sp1d8);
uint32_t sbb0 = 0x7F;
if (has_sub1) sbb0 = read_bscale_sh(B_scale_sh, nts1, pos, grp, sd40, sp1d8);
uint32_t sba1 = read_bscale_sh(B_scale_sh, nts0, pos, grp, sd41, sp1d8);
uint32_t sbb1 = 0x7F;
if (has_sub1) sbb1 = read_bscale_sh(B_scale_sh, nts1, pos, grp, sd41, sp1d8);
// Issue B loads for BOTH iterations via u32x4 (1 vmcnt each, 4 total or 2 if no sub1)
u32x4 bfa0={0,0,0,0}, bfb0={0,0,0,0};
u32x4 bfa1={0,0,0,0}, bfb1={0,0,0,0};
load_b_nt_u32x4(B_shuffle, (size_t)ntx0*KB*1024+(size_t)kb0*1024, lid, bfa0);
if (has_sub1) load_b_nt_u32x4(B_shuffle, (size_t)ntx1*KB*1024+(size_t)kb0*1024, lid, bfb0);
load_b_nt_u32x4(B_shuffle, (size_t)ntx0*KB*1024+(size_t)kb1*1024, lid, bfa1);
if (has_sub1) load_b_nt_u32x4(B_shuffle, (size_t)ntx1*KB*1024+(size_t)kb1*1024, lid, bfb1);
// Wait for A loads only (oldest) — B loads remain in flight
if (has_sub1) {
asm volatile("s_waitcnt vmcnt(4)":::"memory");
} else {
asm volatile("s_waitcnt vmcnt(2)":::"memory");
}
// Compute quant iter 0 directly to u32x4 (no stack intermediary)
uint32_t sa0 = 0x7F;
u32x4 af0 = {0,0,0,0};
if (valid) {
quant_compute_u32x4(ad0, af0, sa0);
}
// Wait for iter 0 B loads (2 if has_sub1, 1 if not)
if (has_sub1) {
asm volatile("s_waitcnt vmcnt(2)":::"memory");
} else {
asm volatile("s_waitcnt vmcnt(1)":::"memory");
}
// ASM MFMA pair iter 0 (Bscale already in registers from pre-read)
asm_mfma_pair(af0, sa0, bfa0, bfb0, sba0, sbb0, acc0, acc1);
// Wait for iter 1 B loads (should be ready by now)
asm volatile("s_waitcnt vmcnt(0)":::"memory");
// Compute quant iter 1 directly to u32x4 (no stack intermediary)
uint32_t sa1 = 0x7F;
u32x4 af1 = {0,0,0,0};
if (valid) {
quant_compute_u32x4(ad1, af1, sa1);
}
// ASM MFMA pair iter 1 (Bscale already in registers from pre-read)
asm_mfma_pair(af1, sa1, bfa1, bfb1, sba1, sbb1, acc0, acc1);
} else if (KI >= 3) {
// General case
int ki = 0;
int kb0 = kb_base+wid*KI+ki;
int keo0 = ke_base+wid*(K_slice/NW)+ki*128;
int sd40 = keo0/(MXFP4_BLOCK_SIZE*4);
uint8_t ap[16]; uint32_t sa0=0x7F;
mfma_input_t af0={0,0,0,0,0,0,0,0};
if (valid) {
quant_bf16x32_to_fp4(A_bf16, mr*K+keo0+grp*32, ap, sa0);
uint4 c=*reinterpret_cast<const uint4*>(ap);
af0[0]=c.x;af0[1]=c.y;af0[2]=c.z;af0[3]=c.w;
}
mfma_input_t bfa={0,0,0,0,0,0,0,0}, bfb={0,0,0,0,0,0,0,0};
load_b_nt(B_shuffle, (size_t)ntx0*KB*1024+(size_t)kb0*1024, lid, bfa);
uint32_t sba = read_bscale_sh(B_scale_sh, nts0, pos, grp, sd40, sp1d8);
uint32_t sbb = 0x7F;
if (has_sub1) {
load_b_nt(B_shuffle, (size_t)ntx1*KB*1024+(size_t)kb0*1024, lid, bfb);
sbb = read_bscale_sh(B_scale_sh, nts1, pos, grp, sd40, sp1d8);
}
for (ki = 1; ki < KI; ki++) {
int kb1=kb_base+wid*KI+ki;
int keo1=ke_base+wid*(K_slice/NW)+ki*128;
int sd41=keo1/(MXFP4_BLOCK_SIZE*4);
#if defined(__gfx950__)
asm volatile("s_setprio 1" ::: "memory");
#endif
acc0 = mfma_fp4_16x16(af0, bfa, acc0, sa0, sba);
if (has_sub1) acc1 = mfma_fp4_16x16(af0, bfb, acc1, sa0, sbb);
#if defined(__gfx950__)
asm volatile("s_setprio 0" ::: "memory");
#endif
af0={0,0,0,0,0,0,0,0}; sa0=0x7F;
if (valid) {
quant_bf16x32_to_fp4(A_bf16, mr*K+keo1+grp*32, ap, sa0);
uint4 c=*reinterpret_cast<const uint4*>(ap);
af0[0]=c.x;af0[1]=c.y;af0[2]=c.z;af0[3]=c.w;
}
load_b_nt(B_shuffle, (size_t)ntx0*KB*1024+(size_t)kb1*1024, lid, bfa);
if (has_sub1) load_b_nt(B_shuffle, (size_t)ntx1*KB*1024+(size_t)kb1*1024, lid, bfb);
asm volatile("s_waitcnt vmcnt(0)":::"memory");
sba = read_bscale_sh(B_scale_sh, nts0, pos, grp, sd41, sp1d8);
if (has_sub1) sbb = read_bscale_sh(B_scale_sh, nts1, pos, grp, sd41, sp1d8);
}
#if defined(__gfx950__)
asm volatile("s_setprio 1" ::: "memory");
#endif
acc0 = mfma_fp4_16x16(af0, bfa, acc0, sa0, sba);
if (has_sub1) acc1 = mfma_fp4_16x16(af0, bfb, acc1, sa0, sbb);
#if defined(__gfx950__)
asm volatile("s_setprio 0" ::: "memory");
#endif
} else if (KI == 1) {
int kb=kb_base+wid;
int keo=ke_base+wid*(K_slice/NW);
int sd4=keo/(MXFP4_BLOCK_SIZE*4);
uint8_t ap[16]; uint32_t sa=0x7F;
mfma_input_t af={0,0,0,0,0,0,0,0};
if (valid) {
quant_bf16x32_to_fp4(A_bf16, mr*K+keo+grp*32, ap, sa);
uint4 c=*reinterpret_cast<const uint4*>(ap);
af[0]=c.x;af[1]=c.y;af[2]=c.z;af[3]=c.w;
}
mfma_input_t bf0={0,0,0,0,0,0,0,0};
load_b_nt(B_shuffle, (size_t)ntx0*KB*1024+(size_t)kb*1024, lid, bf0);
uint32_t sb0 = read_bscale_sh(B_scale_sh, nts0, pos, grp, sd4, sp1d8);
#if defined(__gfx950__)
asm volatile("s_setprio 1" ::: "memory");
#endif
acc0 = mfma_fp4_16x16(af, bf0, acc0, sa, sb0);
#if defined(__gfx950__)
asm volatile("s_setprio 0" ::: "memory");
#endif
if (has_sub1) {
mfma_input_t bf1={0,0,0,0,0,0,0,0};
load_b_nt(B_shuffle, (size_t)ntx1*KB*1024+(size_t)kb*1024, lid, bf1);
uint32_t sb1 = read_bscale_sh(B_scale_sh, nts1, pos, grp, sd4, sp1d8);
#if defined(__gfx950__)
asm volatile("s_setprio 1" ::: "memory");
#endif
acc1 = mfma_fp4_16x16(af, bf1, acc1, sa, sb1);
#if defined(__gfx950__)
asm volatile("s_setprio 0" ::: "memory");
#endif
}
}
extern __shared__ uint8_t lds[];
float* lr = reinterpret_cast<float*>(lds);
#pragma unroll
for (int i=0;i<4;i++) lr[wid*T*T+(grp*4+i)*T+pos]=acc0[i];
#pragma unroll
for (int i=0;i<4;i++) lr[NW*T*T+wid*T*T+(grp*4+i)*T+pos]=acc1[i];
__syncthreads();
{
int ml=wid*4+grp, mg=mts+ml, ng=nts0+pos;
if (mg<M&&ng<N) {
float s=0;
#pragma unroll
for (int w=0;w<NW;w++) s+=lr[w*T*T+ml*T+pos];
partial_out[kslice*M*N+mg*N+ng]=s;
}
if (has_sub1) {
ng=nts1+pos;
if (mg<M&&ng<N) {
float s=0;
#pragma unroll
for (int w=0;w<NW;w++) s+=lr[NW*T*T+w*T*T+ml*T+pos];
partial_out[kslice*M*N+mg*N+ng]=s;
}
}
}
}
__global__ void __launch_bounds__(256)
reduce_splitk(
const float* __restrict__ partial_in,
hip_bfloat16* __restrict__ output,
int M, int N, int ksplit
) {
const int idx = blockIdx.x * 256 + threadIdx.x;
const int total = M * N;
if (idx >= total) return;
float s = 0.0f;
#pragma unroll 7
for (int k = 0; k < ksplit; k++) s += partial_in[k * total + idx];
output[idx] = hip_bfloat16(s);
}
#endif // UNIT_SPLITK
#ifdef UNIT_NARROW
extern "C" void ct_fused_narrow(void* A, void* B, void* Bs, void* out,
int M, int N, int K, int sp1d8) {
int nt=(N+15)/16, mt=(M+15)/16;
hipLaunchKernelGGL(mxfp4_fused_narrow, dim3(nt,mt), dim3(256), 4*16*16*4, 0,
(const uint16_t*)A, (const uint8_t*)B, (const uint8_t*)Bs,
(hip_bfloat16*)out, M, N, K, sp1d8);
}
#endif // UNIT_NARROW
#ifdef UNIT_WIDE
extern "C" void ct_fused_wide(void* A, void* B, void* Bs, void* out,
int M, int N, int K, int sp1d8) {
int nt=(N+31)/32, mt=(M+15)/16;
hipLaunchKernelGGL(mxfp4_fused_wide, dim3(nt,mt), dim3(256), 2*4*16*16*4, 0,
(const uint16_t*)A, (const uint8_t*)B, (const uint8_t*)Bs,
(hip_bfloat16*)out, M, N, K, sp1d8);
}
#endif // UNIT_WIDE
#ifdef UNIT_QUAD
extern "C" void ct_fused_quad(void* A, void* B, void* Bs, void* out,
int M, int N, int K, int sp1d8) {
int nt=(N+63)/64, mt=(M+15)/16;
hipLaunchKernelGGL(mxfp4_fused_quad, dim3(nt,mt), dim3(256), 4*4*16*16*4, 0,
(const uint16_t*)A, (const uint8_t*)B, (const uint8_t*)Bs,
(hip_bfloat16*)out, M, N, K, sp1d8);
}
#endif // UNIT_QUAD
#ifdef UNIT_SPLITK
extern "C" void ct_splitk_fused(void* A, void* B, void* Bs,
void* partial, void* out,
int M, int N, int K, int sp1d8, int ksplit) {
int nt=(N+31)/32, mt=(M+15)/16;
hipLaunchKernelGGL(mxfp4_gemm_splitk_fused, dim3(nt,mt,ksplit), dim3(256), 2*4*16*16*4, 0,
(const uint16_t*)A, (const uint8_t*)B,
(const uint8_t*)Bs, (float*)partial, M, N, K, sp1d8, ksplit);
int total = M * N;
int nblk = (total + 255) / 256;
hipLaunchKernelGGL(reduce_splitk, dim3(nblk), dim3(256), 0, 0,
(const float*)partial, (hip_bfloat16*)out, M, N, ksplit);
}
#endif // UNIT_SPLITK
#ifdef UNIT_HEX
extern "C" void ct_fused_hex(void* A, void* B, void* Bs, void* out,
int M, int N, int K, int sp1d8) {
int nt=(N+95)/96, mt=(M+15)/16;
hipLaunchKernelGGL(mxfp4_fused_hex, dim3(nt,mt), dim3(256), 6*4*16*16*4, 0,
(const uint16_t*)A, (const uint8_t*)B, (const uint8_t*)Bs,
(hip_bfloat16*)out, M, N, K, sp1d8);
}
#endif // UNIT_HEX
"""
CPP_NARROW = r"""
#include <torch/extension.h>
extern "C" void ct_fused_narrow(void*, void*, void*, void*, int, int, int, int);
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {}
"""
CPP_NARROW_HEX = r"""
#include <torch/extension.h>
extern "C" void ct_fused_narrow(void*, void*, void*, void*, int, int, int, int);
extern "C" void ct_fused_hex(void*, void*, void*, void*, int, int, int, int);
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {}
"""
CPP_WIDE = r"""
#include <torch/extension.h>
extern "C" void ct_fused_wide(void*, void*, void*, void*, int, int, int, int);
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {}
"""
CPP_QUAD = r"""
#include <torch/extension.h>
extern "C" void ct_fused_quad(void*, void*, void*, void*, int, int, int, int);
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {}
"""
CPP_SPLITK = r"""
#include <torch/extension.h>
extern "C" void ct_splitk_fused(void*, void*, void*, void*, void*, int, int, int, int, int);
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {}
"""
CPP_HEX = r"""
#include <torch/extension.h>
extern "C" void ct_fused_hex(void*, void*, void*, void*, int, int, int, int);
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {}
"""
_lib_narrow = None
_lib_wide = None
_lib_quad = None
_lib_splitk = None
_lib_hex = None
_output_cache = {}
_partial_cache = {}
_b_cache = {}
VP = ctypes.c_void_p
CI = ctypes.c_int
_BASE_FLAGS = [
"-O3", "--offload-arch=gfx950",
"-mllvm", "-amdgpu-early-inline-all=true",
"-mllvm", "--amdgpu-kernarg-preload-count=16",
"-DHIP_HCC_COMPAT_MODE=1",
]
def _compile_unit(name, cpp_src, defines, extra_flags=None):
from torch.utils.cpp_extension import load_inline
if isinstance(defines, str):
defines = [defines]
flags = _BASE_FLAGS + [f"-D{d}" for d in defines]
if extra_flags:
flags += extra_flags
mod = load_inline(
name=name, cpp_sources=cpp_src, cuda_sources=HIP_SOURCE,
extra_cuda_cflags=flags,
verbose=False,
)
return ctypes.CDLL(mod.__file__)
def _get_modules():
global _lib_narrow, _lib_wide, _lib_quad, _lib_splitk, _lib_hex
if _lib_narrow is not None:
return
with ThreadPoolExecutor(max_workers=4) as pool:
f_narrow_hex = pool.submit(_compile_unit, "mxfp4_nh_v765", CPP_NARROW_HEX,
["UNIT_NARROW", "UNIT_HEX"],
["-mllvm", "-amdgpu-function-calls=false"])
f_wide = pool.submit(_compile_unit, "mxfp4_w_v765", CPP_WIDE, "UNIT_WIDE")
f_quad = pool.submit(_compile_unit, "mxfp4_q_v765", CPP_QUAD, "UNIT_QUAD",
["-mllvm", "-amdgpu-function-calls=false"])
f_splitk = pool.submit(_compile_unit, "mxfp4_s_v765", CPP_SPLITK, "UNIT_SPLITK")
lib_nh = f_narrow_hex.result()
_lib_narrow = lib_nh
_lib_hex = lib_nh
_lib_wide = f_wide.result()
_lib_quad = f_quad.result()
_lib_splitk = f_splitk.result()
_lib_narrow.ct_fused_narrow.argtypes = [VP, VP, VP, VP, CI, CI, CI, CI]
_lib_narrow.ct_fused_narrow.restype = None
_lib_wide.ct_fused_wide.argtypes = [VP, VP, VP, VP, CI, CI, CI, CI]
_lib_wide.ct_fused_wide.restype = None
_lib_quad.ct_fused_quad.argtypes = [VP, VP, VP, VP, CI, CI, CI, CI]
_lib_quad.ct_fused_quad.restype = None
_lib_splitk.ct_splitk_fused.argtypes = [VP, VP, VP, VP, VP, CI, CI, CI, CI, CI]
_lib_splitk.ct_splitk_fused.restype = None
_lib_hex.ct_fused_hex.argtypes = [VP, VP, VP, VP, CI, CI, CI, CI]
_lib_hex.ct_fused_hex.restype = None
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]
_get_modules()
spr = K // 32
sp1d8 = (spr + 7) // 8
ok = (M, N)
if ok not in _output_cache:
_output_cache[ok] = torch.empty(M, N, dtype=torch.bfloat16, device=A.device)
out = _output_cache[ok]
a_ptr = A.data_ptr()
b_ptr = B_shuffle.data_ptr()
bs_ptr = B_scale_sh.data_ptr()
o_ptr = out.data_ptr()
if M >= 128:
_lib_hex.ct_fused_hex(a_ptr, b_ptr, bs_ptr, o_ptr, M, N, K, sp1d8)
elif M >= 64:
_lib_quad.ct_fused_quad(a_ptr, b_ptr, bs_ptr, o_ptr, M, N, K, sp1d8)
elif M <= 8:
_lib_narrow.ct_fused_narrow(a_ptr, b_ptr, bs_ptr, o_ptr, M, N, K, sp1d8)
elif K > 4096:
ksplit = 7 if K == 7168 else 2
pk = (ksplit, M, N)
if pk not in _partial_cache:
_partial_cache[pk] = torch.empty(ksplit, M, N, dtype=torch.float32, device=A.device)
partial = _partial_cache[pk]
_lib_splitk.ct_splitk_fused(a_ptr, b_ptr, bs_ptr,
partial.data_ptr(), o_ptr,
M, N, K, sp1d8, ksplit)
else:
_lib_wide.ct_fused_wide(a_ptr, b_ptr, bs_ptr, o_ptr, M, N, K, sp1d8)
return out
scrolls · 1435 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