submission 721028
jiab_85281 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 704 lines, June 9 Researcher Reciprocity License v1.0.
mm_best.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-721028?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:8b2fce1267afcbca99623003a09138d0d75fc084c6ec54b04c77232e250b2699
license declaredunknown
license concludedunknown
authorsjiab_85281
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
__shared__ Lds lds;split-k
namespace splitk {Kernel source
mm_best.py704 lines
# Optimized MXFP4 GEMM: shape6 with 12 waves (all active), shape5 original 16 waves
import os
os.environ['PYTORCH_ROCM_ARCH'] = 'gfx950'
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
CPP_WRAPPER = """
void fp8_mm(torch::Tensor a_full, torch::Tensor b_full, torch::Tensor b_shuffle, torch::Tensor b_scales, torch::Tensor c);
"""
cuda_src = """
#include <torch/extension.h>
#include <cstddef>
#include <cstdint>
#include <hip/amd_detail/amd_hip_fp8.h>
#include <hip/amd_detail/amd_hip_bf16.h>
constexpr int BLOCK = 128;
constexpr int SCALE_GROUP = 32;
constexpr int MMA_M = 16;
constexpr int MMA_N = 16;
constexpr int WAVE = 64;
typedef uint32_t u32x4 __attribute__((ext_vector_type(4)));
typedef uint32_t u32x8 __attribute__((ext_vector_type(8)));
typedef float f32x4 __attribute__((ext_vector_type(4)));
typedef __bf16 bf16x2 __attribute__((ext_vector_type(2)));
union FragRegs { u32x4 v4; u32x8 v8; uint64_t v2[4]; };
__host__ __device__ __forceinline__ int cdiv(int x, int y) { return (x + y - 1) / y; }
__device__ __forceinline__ int lid() { return threadIdx.x & 63; }
__device__ __forceinline__ int wid() { return threadIdx.x / WAVE; }
// ============================================================
// MXFP4 quantization
// ============================================================
__device__ __forceinline__ uint32_t pk_max_u16(uint32_t a, uint32_t b) {
uint32_t r; asm("v_pk_max_u16 %0, %1, %2" : "=v"(r) : "v"(a), "v"(b)); return r;
}
__device__ __forceinline__ uint32_t reduce_pk_u16(uint32_t v) {
const uint32_t hi = v >> 16, lo = v & 0xFFFFu; return hi > lo ? hi : lo;
}
__device__ __forceinline__ void mxfp4_scale(uint32_t mx, uint8_t& sb, float& sf) {
float mf = __builtin_bit_cast(float, mx << 16);
uint32_t re = (__builtin_bit_cast(uint32_t, mf) + 0x00200000u) >> 23;
sb = (re > 2u) ? uint8_t(re - 2u) : uint8_t(0u);
sf = __builtin_bit_cast(float, uint32_t(sb) << 23);
}
__device__ __forceinline__ uint32_t pack_fp4(const uint32_t* bp, float sf) {
uint32_t p = 0u;
p = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(p, __builtin_bit_cast(bf16x2, bp[0]), sf, 0);
p = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(p, __builtin_bit_cast(bf16x2, bp[1]), sf, 1);
p = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(p, __builtin_bit_cast(bf16x2, bp[2]), sf, 2);
p = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(p, __builtin_bit_cast(bf16x2, bp[3]), sf, 3);
return p;
}
__device__ __forceinline__ void quant_bf16x32(const u32x4* src, u32x4& packed, uint8_t& sb) {
constexpr uint32_t SM = 0x7FFF7FFFu;
uint32_t m0 = pk_max_u16(src[0][0]&SM, src[0][1]&SM);
uint32_t m1 = pk_max_u16(src[0][2]&SM, src[0][3]&SM);
uint32_t m2 = pk_max_u16(src[1][0]&SM, src[1][1]&SM);
uint32_t m3 = pk_max_u16(src[1][2]&SM, src[1][3]&SM);
uint32_t m4 = pk_max_u16(src[2][0]&SM, src[2][1]&SM);
uint32_t m5 = pk_max_u16(src[2][2]&SM, src[2][3]&SM);
uint32_t m6 = pk_max_u16(src[3][0]&SM, src[3][1]&SM);
uint32_t m7 = pk_max_u16(src[3][2]&SM, src[3][3]&SM);
m0 = pk_max_u16(m0,m1); m2 = pk_max_u16(m2,m3);
m4 = pk_max_u16(m4,m5); m6 = pk_max_u16(m6,m7);
m0 = pk_max_u16(m0,m2); m4 = pk_max_u16(m4,m6);
m0 = pk_max_u16(m0,m4);
float sf; mxfp4_scale(reduce_pk_u16(m0), sb, sf);
const uint32_t* raw = reinterpret_cast<const uint32_t*>(src);
packed = { pack_fp4(&raw[0], sf), pack_fp4(&raw[4], sf),
pack_fp4(&raw[8], sf), pack_fp4(&raw[12], sf) };
}
// ============================================================
// Scale addressing
// ============================================================
__device__ __forceinline__ uint32_t padded_scale_cols(int k) {
return (uint32_t(cdiv(k, SCALE_GROUP)) + 7u) & ~7u;
}
__device__ __forceinline__ size_t shuffled_scale_offset(uint32_t row, uint32_t kg, uint32_t pkg) {
uint32_t rb = row >> 5, rh = (row >> 4) & 1u, ri = row & 15u;
uint32_t gb = kg >> 3, gh = (kg >> 2) & 1u, gi = kg & 3u;
return (size_t(rb) * (pkg >> 3) + gb) * 256u +
size_t(gi) * 64u + size_t(ri) * 4u + size_t(gh) * 2u + size_t(rh);
}
// ============================================================
// Helper: write output tile
// ============================================================
__device__ __forceinline__ void write_c(
__hip_bfloat16* c, f32x4 sum, int m0, int tn, int m, int n, int lane
) {
int col = tn * MMA_N + (lane & 15);
int rb = m0 + (lane >> 4) * 4;
uint32_t pk01, pk23;
asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(pk01) : "v"(sum[0]), "v"(sum[1]));
asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(pk23) : "v"(sum[2]), "v"(sum[3]));
const uint16_t* v01 = reinterpret_cast<const uint16_t*>(&pk01);
const uint16_t* v23 = reinterpret_cast<const uint16_t*>(&pk23);
#pragma unroll
for (int i = 0; i < 4; ++i) {
int r = rb + i;
if (r < m && col < n) {
uint16_t val = (i < 2) ? v01[i] : v23[i - 2];
reinterpret_cast<uint16_t*>(c)[size_t(r) * n + col] = val;
}
}
}
// ============================================================
// Shape 2: M=16, N=2112, K=7168 (original 16 waves, 4 iters)
// ============================================================
namespace shape2 {
constexpr int WAVES = 16;
constexpr int KT = 56;
struct Lds {
alignas(16) float reduce[WAVES][WAVE][4];
};
__global__ __launch_bounds__(WAVES * WAVE)
void kernel(
const __hip_bfloat16* a, const uint8_t* bsh, const uint8_t* bsc,
__hip_bfloat16* c, int m, int n, int k
) {
__shared__ Lds lds;
const int w = wid();
const int lane = lid();
const int tn = int(blockIdx.x);
const int tm = int(blockIdx.y);
const int m0 = tm * MMA_M;
const int row = lane & 15, kgrp = lane >> 4;
const uint32_t pkg = padded_scale_cols(k);
f32x4 acc = {0.f, 0.f, 0.f, 0.f};
for (int kk = 0; kk < 4; ++kk) {
int ki = w + kk * WAVES;
if (ki >= KT) break;
int vrows = m - m0;
vrows = vrows < 0 ? 0 : (vrows > MMA_M ? MMA_M : vrows);
u32x4 chunks[4] = {};
if (row < vrows) {
const char* base = reinterpret_cast<const char*>(a) +
(size_t(m0) * k + size_t(ki) * BLOCK) * sizeof(__hip_bfloat16);
const char* lsrc = base + size_t(row * k + kgrp * SCALE_GROUP) * sizeof(__hip_bfloat16);
const u32x4* p = reinterpret_cast<const u32x4*>(lsrc);
chunks[0] = p[0]; chunks[1] = p[1]; chunks[2] = p[2]; chunks[3] = p[3];
}
u32x4 a_packed; uint8_t a_sb;
quant_bf16x32(chunks, a_packed, a_sb);
uint32_t a_sc = a_sb;
u32x4 b_data = {};
uint32_t b_sc = 0;
int vraw = n - tn * MMA_N;
int valid = vraw < 0 ? 0 : (vraw > MMA_N ? MMA_N : vraw);
if (row < uint32_t(valid)) {
size_t toff = size_t(tn) * MMA_N * (k / 2) + size_t(ki) * MMA_N * (BLOCK / 2);
uint32_t off16 = row * 16 + kgrp * 256;
const u32x4* bptr = reinterpret_cast<const u32x4*>(bsh + toff + off16);
b_data = *bptr;
uint32_t grow = tn * MMA_N + row;
uint32_t gkg = ki * (BLOCK / SCALE_GROUP) + kgrp;
b_sc = reinterpret_cast<const uint8_t*>(bsc)[
shuffled_scale_offset(grow, gkg, pkg)];
}
FragRegs ar, br;
ar.v4 = a_packed;
br.v4 = b_data;
acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
ar.v8, br.v8, acc, 4, 4, 0, a_sc, 0, b_sc);
}
lds.reduce[w][lane][0] = acc[0];
lds.reduce[w][lane][1] = acc[1];
lds.reduce[w][lane][2] = acc[2];
lds.reduce[w][lane][3] = acc[3];
__syncthreads();
if (w == 0) {
f32x4 sum = {0.f, 0.f, 0.f, 0.f};
#pragma unroll
for (int s = 0; s < WAVES; ++s) {
sum[0] += lds.reduce[s][lane][0];
sum[1] += lds.reduce[s][lane][1];
sum[2] += lds.reduce[s][lane][2];
sum[3] += lds.reduce[s][lane][3];
}
write_c(c, sum, m0, tn, m, n, lane);
}
}
} // shape2
// ============================================================
// Shape 5: M=64, N=7168, K=2048 (original 16 waves, 4 N-tiles)
// ============================================================
namespace shape5 {
constexpr int WAVES = 16;
constexpr int N_TILES_PER_CTA = 4;
constexpr int S5_N = 7168;
constexpr int S5_K = 2048;
constexpr int S5_PKG = 64;
struct Lds {
alignas(16) float reduce[N_TILES_PER_CTA][WAVES][WAVE][4];
};
__global__ __launch_bounds__(WAVES * WAVE)
void kernel(
const __hip_bfloat16* a, const uint8_t* bsh, const uint8_t* bsc,
__hip_bfloat16* c, int m, int n, int k
) {
__shared__ Lds lds;
const int w = wid();
const int lane = lid();
const int tm = int(blockIdx.y);
const int ng = int(blockIdx.x);
const int m0 = tm * MMA_M;
const int row = lane & 15, kgrp = lane >> 4;
const int ki = w;
const char* base = reinterpret_cast<const char*>(a) +
(size_t(m0) * S5_K + size_t(ki) * BLOCK) * sizeof(__hip_bfloat16);
const char* lsrc = base + size_t(row * S5_K + kgrp * SCALE_GROUP) * sizeof(__hip_bfloat16);
const u32x4* p = reinterpret_cast<const u32x4*>(lsrc);
u32x4 chunks[4];
chunks[0] = p[0]; chunks[1] = p[1]; chunks[2] = p[2]; chunks[3] = p[3];
u32x4 a_packed; uint8_t a_sb;
quant_bf16x32(chunks, a_packed, a_sb);
uint32_t a_sc = a_sb;
FragRegs ar; ar.v4 = a_packed;
#pragma unroll
for (int nt = 0; nt < N_TILES_PER_CTA; ++nt) {
int tn = ng * N_TILES_PER_CTA + nt;
size_t toff = size_t(tn) * MMA_N * (S5_K / 2) + size_t(ki) * MMA_N * (BLOCK / 2);
uint32_t off16 = row * 16 + kgrp * 256;
const u32x4* bptr = reinterpret_cast<const u32x4*>(bsh + toff + off16);
u32x4 b_data = *bptr;
uint32_t grow = tn * MMA_N + row;
uint32_t gkg = ki * (BLOCK / SCALE_GROUP) + kgrp;
uint32_t b_sc = reinterpret_cast<const uint8_t*>(bsc)[
shuffled_scale_offset(grow, gkg, S5_PKG)];
f32x4 acc = {0.f, 0.f, 0.f, 0.f};
FragRegs br; br.v4 = b_data;
acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
ar.v8, br.v8, acc, 4, 4, 0, a_sc, 0, b_sc);
lds.reduce[nt][w][lane][0] = acc[0];
lds.reduce[nt][w][lane][1] = acc[1];
lds.reduce[nt][w][lane][2] = acc[2];
lds.reduce[nt][w][lane][3] = acc[3];
}
__syncthreads();
if (w < N_TILES_PER_CTA) {
int tn = ng * N_TILES_PER_CTA + w;
f32x4 sum = {0.f, 0.f, 0.f, 0.f};
#pragma unroll
for (int s = 0; s < WAVES; ++s) {
sum[0] += lds.reduce[w][s][lane][0];
sum[1] += lds.reduce[w][s][lane][1];
sum[2] += lds.reduce[w][s][lane][2];
sum[3] += lds.reduce[w][s][lane][3];
}
int col = tn * MMA_N + (lane & 15);
int rb = m0 + (lane >> 4) * 4;
uint32_t pk01, pk23;
asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(pk01) : "v"(sum[0]), "v"(sum[1]));
asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(pk23) : "v"(sum[2]), "v"(sum[3]));
const uint16_t* v01 = reinterpret_cast<const uint16_t*>(&pk01);
const uint16_t* v23 = reinterpret_cast<const uint16_t*>(&pk23);
reinterpret_cast<uint16_t*>(c)[size_t(rb) * S5_N + col] = v01[0];
reinterpret_cast<uint16_t*>(c)[size_t(rb + 1) * S5_N + col] = v01[1];
reinterpret_cast<uint16_t*>(c)[size_t(rb + 2) * S5_N + col] = v23[0];
reinterpret_cast<uint16_t*>(c)[size_t(rb + 3) * S5_N + col] = v23[1];
}
}
} // shape5
// ============================================================
// Shape 6: M=256, N=3072, K=1536
// 12 waves (all active), 6 n-tiles in 2 batches of 3
// ============================================================
namespace shape6 {
constexpr int WAVES = 12;
constexpr int K_TILES = 12;
constexpr int N_TILES_PER_CTA = 6;
constexpr int BATCH = 3;
constexpr int S6_N = 3072;
constexpr int S6_K = 1536;
constexpr int S6_PKG = 48;
struct Lds {
alignas(16) float reduce[BATCH][K_TILES][WAVE][4];
};
__global__ __launch_bounds__(WAVES * WAVE)
void kernel(
const __hip_bfloat16* a, const uint8_t* bsh, const uint8_t* bsc,
__hip_bfloat16* c, int m, int n, int k
) {
__shared__ Lds lds;
const int w = wid();
const int lane = lid();
const int tm = int(blockIdx.y);
const int ng = int(blockIdx.x);
const int m0 = tm * MMA_M;
const int row = lane & 15, kgrp = lane >> 4;
const int ki = w;
// All 12 waves active
const char* base = reinterpret_cast<const char*>(a) +
(size_t(m0) * S6_K + size_t(ki) * BLOCK) * sizeof(__hip_bfloat16);
const char* lsrc = base + size_t(row * S6_K + kgrp * SCALE_GROUP) * sizeof(__hip_bfloat16);
const u32x4* p = reinterpret_cast<const u32x4*>(lsrc);
u32x4 chunks[4];
chunks[0] = p[0]; chunks[1] = p[1]; chunks[2] = p[2]; chunks[3] = p[3];
u32x4 a_packed; uint8_t a_sb;
quant_bf16x32(chunks, a_packed, a_sb);
uint32_t a_sc = a_sb;
FragRegs ar; ar.v4 = a_packed;
#pragma unroll
for (int batch = 0; batch < 2; ++batch) {
#pragma unroll
for (int bi = 0; bi < BATCH; ++bi) {
int nt = batch * BATCH + bi;
int tn = ng * N_TILES_PER_CTA + nt;
size_t toff = size_t(tn) * MMA_N * (S6_K / 2) + size_t(ki) * MMA_N * (BLOCK / 2);
uint32_t off16 = row * 16 + kgrp * 256;
const u32x4* bptr = reinterpret_cast<const u32x4*>(bsh + toff + off16);
u32x4 b_data = *bptr;
uint32_t grow = tn * MMA_N + row;
uint32_t gkg = ki * (BLOCK / SCALE_GROUP) + kgrp;
uint32_t b_sc = reinterpret_cast<const uint8_t*>(bsc)[
shuffled_scale_offset(grow, gkg, S6_PKG)];
FragRegs br; br.v4 = b_data;
f32x4 acc = {0.f, 0.f, 0.f, 0.f};
acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
ar.v8, br.v8, acc, 4, 4, 0, a_sc, 0, b_sc);
lds.reduce[bi][w][lane][0] = acc[0];
lds.reduce[bi][w][lane][1] = acc[1];
lds.reduce[bi][w][lane][2] = acc[2];
lds.reduce[bi][w][lane][3] = acc[3];
}
__syncthreads();
if (w < BATCH) {
int nt = batch * BATCH + w;
int tn = ng * N_TILES_PER_CTA + nt;
f32x4 sum = {0.f, 0.f, 0.f, 0.f};
#pragma unroll
for (int s = 0; s < K_TILES; ++s) {
sum[0] += lds.reduce[w][s][lane][0];
sum[1] += lds.reduce[w][s][lane][1];
sum[2] += lds.reduce[w][s][lane][2];
sum[3] += lds.reduce[w][s][lane][3];
}
int col = tn * MMA_N + (lane & 15);
int rb = m0 + (lane >> 4) * 4;
uint32_t pk01, pk23;
asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(pk01) : "v"(sum[0]), "v"(sum[1]));
asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(pk23) : "v"(sum[2]), "v"(sum[3]));
const uint16_t* v01 = reinterpret_cast<const uint16_t*>(&pk01);
const uint16_t* v23 = reinterpret_cast<const uint16_t*>(&pk23);
reinterpret_cast<uint16_t*>(c)[size_t(rb) * S6_N + col] = v01[0];
reinterpret_cast<uint16_t*>(c)[size_t(rb + 1) * S6_N + col] = v01[1];
reinterpret_cast<uint16_t*>(c)[size_t(rb + 2) * S6_N + col] = v23[0];
reinterpret_cast<uint16_t*>(c)[size_t(rb + 3) * S6_N + col] = v23[1];
}
__syncthreads();
}
}
} // shape6
// ============================================================
// Generic nosplit (shapes with kt<=4)
// ============================================================
namespace nosplit {
struct Lds {
alignas(16) float reduce[4][WAVE][4];
};
__global__ void kernel(
const __hip_bfloat16* a, const uint8_t* bsh, const uint8_t* bsc,
__hip_bfloat16* c, int m, int n, int k
) {
__shared__ Lds lds;
const int w = wid();
const int lane = lid();
const int tm = int(blockIdx.y);
const int tn = int(blockIdx.x);
if (tm >= cdiv(m, MMA_M) || tn >= cdiv(n, MMA_N)) return;
const int ki = w;
const int m0 = tm * MMA_M;
const int row = lane & 15, kgrp = lane >> 4;
int vrows = m - m0; vrows = vrows < 0 ? 0 : (vrows > MMA_M ? MMA_M : vrows);
u32x4 chunks[4] = {};
if (row < vrows) {
const char* base = reinterpret_cast<const char*>(a) +
(size_t(m0) * k + size_t(ki) * BLOCK) * sizeof(__hip_bfloat16);
const char* lsrc = base + size_t(row * k + kgrp * SCALE_GROUP) * sizeof(__hip_bfloat16);
const u32x4* p = reinterpret_cast<const u32x4*>(lsrc);
chunks[0] = p[0]; chunks[1] = p[1]; chunks[2] = p[2]; chunks[3] = p[3];
}
u32x4 a_packed; uint8_t a_sb;
quant_bf16x32(chunks, a_packed, a_sb);
uint32_t a_sc = a_sb;
u32x4 b_data = {};
uint32_t b_sc = 0;
int vraw = n - tn * MMA_N;
int valid = vraw < 0 ? 0 : (vraw > MMA_N ? MMA_N : vraw);
if (row < uint32_t(valid)) {
size_t toff = size_t(tn) * MMA_N * (k / 2) + size_t(ki) * MMA_N * (BLOCK / 2);
uint32_t off16 = row * 16 + kgrp * 256;
const u32x4* bptr = reinterpret_cast<const u32x4*>(bsh + toff + off16);
b_data = *bptr;
uint32_t grow = tn * MMA_N + row;
uint32_t gkg = ki * (BLOCK / SCALE_GROUP) + kgrp;
b_sc = reinterpret_cast<const uint8_t*>(bsc)[
shuffled_scale_offset(grow, gkg, padded_scale_cols(k))];
}
f32x4 acc = {0.f, 0.f, 0.f, 0.f};
FragRegs ar, br; ar.v4 = a_packed; br.v4 = b_data;
acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
ar.v8, br.v8, acc, 4, 4, 0, a_sc, 0, b_sc);
lds.reduce[w][lane][0] = acc[0];
lds.reduce[w][lane][1] = acc[1];
lds.reduce[w][lane][2] = acc[2];
lds.reduce[w][lane][3] = acc[3];
__syncthreads();
if (w == 0) {
f32x4 sum = {0.f, 0.f, 0.f, 0.f};
#pragma unroll
for (int s = 0; s < 4; ++s) {
sum[0] += lds.reduce[s][lane][0];
sum[1] += lds.reduce[s][lane][1];
sum[2] += lds.reduce[s][lane][2];
sum[3] += lds.reduce[s][lane][3];
}
write_c(c, sum, tm * MMA_M, tn, m, n, lane);
}
}
} // nosplit
// ============================================================
// Generic split-K (fallback)
// ============================================================
namespace splitk {
constexpr int MAX_WAVES = 16;
struct Lds {
alignas(16) float reduce[MAX_WAVES][WAVE][4];
};
__global__ void kernel(
const __hip_bfloat16* a, const uint8_t* bsh, const uint8_t* bsc,
float* workspace, __hip_bfloat16* c_out,
int m, int n, int k, int k_split
) {
__shared__ Lds lds;
const int waves_per_cta = int(blockDim.x) / WAVE;
const int w = wid();
const int lane = lid();
const int tm = int(blockIdx.y);
const int tn = int(blockIdx.x);
const int kz = int(blockIdx.z);
if (tm >= cdiv(m, MMA_M) || tn >= cdiv(n, MMA_N)) return;
const int ki = kz * waves_per_cta + w;
const int k_tiles = cdiv(k, BLOCK);
const int row = lane & 15, kgrp = lane >> 4;
f32x4 acc = {0.f, 0.f, 0.f, 0.f};
if (ki < k_tiles) {
const int m0 = tm * MMA_M;
int vrows = m - m0; vrows = vrows < 0 ? 0 : (vrows > MMA_M ? MMA_M : vrows);
u32x4 chunks[4] = {};
if (row < vrows) {
const char* base = reinterpret_cast<const char*>(a) +
(size_t(m0) * k + size_t(ki) * BLOCK) * sizeof(__hip_bfloat16);
const char* lsrc = base + size_t(row * k + kgrp * SCALE_GROUP) * sizeof(__hip_bfloat16);
const u32x4* p = reinterpret_cast<const u32x4*>(lsrc);
chunks[0] = p[0]; chunks[1] = p[1]; chunks[2] = p[2]; chunks[3] = p[3];
}
u32x4 a_packed; uint8_t a_sb;
quant_bf16x32(chunks, a_packed, a_sb);
uint32_t a_sc = a_sb;
u32x4 b_data = {};
uint32_t b_sc = 0;
int vraw = n - tn * MMA_N;
int valid = vraw < 0 ? 0 : (vraw > MMA_N ? MMA_N : vraw);
if (row < uint32_t(valid)) {
size_t toff = size_t(tn) * MMA_N * (k / 2) + size_t(ki) * MMA_N * (BLOCK / 2);
uint32_t off16 = row * 16 + kgrp * 256;
const u32x4* bptr = reinterpret_cast<const u32x4*>(bsh + toff + off16);
b_data = *bptr;
uint32_t grow = tn * MMA_N + row;
uint32_t gkg = ki * (BLOCK / SCALE_GROUP) + kgrp;
b_sc = reinterpret_cast<const uint8_t*>(bsc)[
shuffled_scale_offset(grow, gkg, padded_scale_cols(k))];
}
FragRegs ar, br; ar.v4 = a_packed; br.v4 = b_data;
acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
ar.v8, br.v8, acc, 4, 4, 0, a_sc, 0, b_sc);
}
lds.reduce[w][lane][0] = acc[0];
lds.reduce[w][lane][1] = acc[1];
lds.reduce[w][lane][2] = acc[2];
lds.reduce[w][lane][3] = acc[3];
__syncthreads();
if (w == 0) {
f32x4 sum = {0.f, 0.f, 0.f, 0.f};
for (int s = 0; s < waves_per_cta; ++s) {
sum[0] += lds.reduce[s][lane][0];
sum[1] += lds.reduce[s][lane][1];
sum[2] += lds.reduce[s][lane][2];
sum[3] += lds.reduce[s][lane][3];
}
int col = tn * MMA_N + (lane & 15);
int rb = tm * MMA_M + (lane >> 4) * 4;
if (k_split == 1) {
uint32_t pk01, pk23;
asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(pk01) : "v"(sum[0]), "v"(sum[1]));
asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(pk23) : "v"(sum[2]), "v"(sum[3]));
const uint16_t* v01 = reinterpret_cast<const uint16_t*>(&pk01);
const uint16_t* v23 = reinterpret_cast<const uint16_t*>(&pk23);
#pragma unroll
for (int i = 0; i < 4; ++i) {
int r = rb + i;
if (r < m && col < n) {
uint16_t val = (i < 2) ? v01[i] : v23[i - 2];
reinterpret_cast<uint16_t*>(c_out)[size_t(r) * n + col] = val;
}
}
} else {
#pragma unroll
for (int i = 0; i < 4; ++i) {
int r = rb + i;
if (r < m && col < n)
atomicAdd(&workspace[size_t(r) * n + col], sum[i]);
}
size_t mn = size_t(m) * n;
int tile_id = tm * cdiv(n, MMA_N) + tn;
int* counters = reinterpret_cast<int*>(workspace + mn);
int done = 0;
if (lane == 0)
done = atomicAdd(&counters[tile_id], 1);
done = __builtin_amdgcn_readfirstlane(done);
if (done == k_split - 1) {
#pragma unroll
for (int i = 0; i < 4; ++i) {
int r = rb + i;
if (r < m && col < n) {
float v = reinterpret_cast<volatile float*>(workspace)[size_t(r) * n + col];
uint32_t pk;
asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(pk) : "v"(v), "v"(0.f));
reinterpret_cast<uint16_t*>(c_out)[size_t(r) * n + col] =
uint16_t(pk & 0xFFFFu);
}
}
}
}
}
}
} // splitk
// ============================================================
// Host dispatch
// ============================================================
void fp8_mm(torch::Tensor a_full, torch::Tensor b_full,
torch::Tensor b_shuffle, torch::Tensor b_scales, torch::Tensor c) {
int m = a_full.size(0);
int n = c.size(1);
int k = a_full.size(1);
const auto* ap = reinterpret_cast<const __hip_bfloat16*>(a_full.data_ptr());
const auto* bsh = reinterpret_cast<const uint8_t*>(b_shuffle.data_ptr());
const auto* bsc = reinterpret_cast<const uint8_t*>(b_scales.data_ptr());
auto* cp = reinterpret_cast<__hip_bfloat16*>(c.data_ptr());
int tm = cdiv(m, MMA_M);
int tn = cdiv(n, MMA_N);
int kt = cdiv(k, BLOCK);
// Shape 2: N=2112, K=7168 (any M)
if (n == 2112 && k == 7168) {
shape2::kernel<<<dim3(tn, tm), dim3(shape2::WAVES * WAVE)>>>(
ap, bsh, bsc, cp, m, n, k);
return;
}
// Shape 5: M=64, N=7168, K=2048
if (m == 64 && n == 7168 && k == 2048) {
int ng = cdiv(tn, shape5::N_TILES_PER_CTA);
shape5::kernel<<<dim3(ng, tm), dim3(shape5::WAVES * WAVE)>>>(ap, bsh, bsc, cp, m, n, k);
return;
}
// Shape 6: M=256, N=3072, K=1536
if (m == 256 && n == 3072 && k == 1536) {
int ng = cdiv(tn, shape6::N_TILES_PER_CTA);
shape6::kernel<<<dim3(ng, tm), dim3(shape6::WAVES * WAVE)>>>(ap, bsh, bsc, cp, m, n, k);
return;
}
if (kt <= 4) {
nosplit::kernel<<<dim3(tn, tm), dim3(4 * WAVE)>>>(ap, bsh, bsc, cp, m, n, k);
} else {
int waves_per_cta = kt;
int k_split = 1;
while (waves_per_cta > 16) {
k_split++;
waves_per_cta = cdiv(kt, k_split);
}
float* wp = reinterpret_cast<float*>(
reinterpret_cast<__hip_bfloat16*>(b_full.data_ptr()) + m * n);
splitk::kernel<<<dim3(tn, tm, k_split), dim3(waves_per_cta * WAVE)>>>(
ap, bsh, bsc, wp, cp, m, n, k, k_split);
}
}
"""
module = load_inline(
name='fp8_mm_v2',
cpp_sources=[CPP_WRAPPER],
cuda_sources=[cuda_src],
functions=['fp8_mm'],
verbose=True,
extra_cuda_cflags=["--offload-arch=gfx950",
"-O3",
"-std=c++20"],
)
import torch
def custom_kernel(data: input_t) -> output_t:
a_full, b_full, b_fp4, b_shuffle, b_scales = data
m = a_full.size(0)
n = b_fp4.size(0)
k = a_full.size(1)
flat = b_full.view(-1)
c = flat.narrow(0, 0, m * n).view(m, n)
kt = (k + 127) // 128
is_shape2 = (n == 2112 and k == 7168)
is_shape5 = (m == 64 and n == 7168 and k == 2048)
is_shape6 = (m == 256 and n == 3072 and k == 1536)
# Zero workspace for generic splitk path
if kt > 4 and not is_shape2 and not is_shape5 and not is_shape6:
waves_per_cta = kt
k_split = 1
while waves_per_cta > 16:
k_split += 1
waves_per_cta = (kt + k_split - 1) // k_split
if k_split > 1:
ws_bf16 = m * n * 2
tm = (m + 15) // 16
tn = (n + 15) // 16
cnt_bf16 = ((tm * tn * 2) + 1) & ~1
flat.narrow(0, m * n, ws_bf16 + cnt_bf16).zero_()
module.fp8_mm(a_full, b_full, b_shuffle, b_scales, c)
return c
scrolls · 704 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