submission 568320
jIab-b · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 887 lines, June 9 Researcher Reciprocity License v1.0.
static2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-568320?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:d3b7fef1bc45c73f5d55e470599ed5e8dbb0285401735cb0a7b9d6ca470d3959
license declaredunknown
license concludedunknown
authorsjIab-b
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
__shared__ LdsPipelineState lds;warp-specialization
constexpr int A_PRODUCER_WAVES = 4;Kernel source
static2.py887 lines
# This script provides a template for using load_inline to run a HIP kernel for
import os
os.environ['PYTORCH_ROCM_ARCH'] = 'gfx950'
from pathlib import Path
import re
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_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_SIZE = 64;
constexpr int NUM_WORKERS = 256;
constexpr int MAX_N_TILES_PER_CTA = 12;
constexpr int FP4_TILE_BYTES = MMA_M * BLOCK / 2;
constexpr int SCALE_PACKED_WORDS = MMA_M;
constexpr int MMA_WAVES = 4;
constexpr int A_PRODUCER_WAVES = 4;
constexpr int B_PRODUCER_WAVES = 8;
constexpr int TOTAL_WAVES = MMA_WAVES + A_PRODUCER_WAVES + B_PRODUCER_WAVES;
constexpr int A_PRODUCERS_PER_CONSUMER = A_PRODUCER_WAVES / MMA_WAVES;
constexpr int B_PRODUCERS_PER_CONSUMER = B_PRODUCER_WAVES / MMA_WAVES;
constexpr int A_STAGES_PER_CONSUMER = 3;
constexpr int B_STAGES_PER_CONSUMER = 8;
constexpr int MAX_LOCAL_TILES = MAX_N_TILES_PER_CTA;
static_assert(A_PRODUCERS_PER_CONSUMER == 1);
static_assert(B_PRODUCERS_PER_CONSUMER == 2);
enum StageParity : uint32_t {
STAGE_PARITY_0 = 0,
STAGE_PARITY_1 = 1,
};
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)));
// union for now for dumb clang reason
union FragRegs {
u32x4 v4;
u32x8 v8;
uint64_t v2[4];
};
struct alignas(16) AStage {
volatile uint8_t prod_to_mma;
volatile uint8_t mma_to_prod;
alignas(16) uint8_t a_base[FP4_TILE_BYTES];
alignas(4) uint32_t as_base[SCALE_PACKED_WORDS];
};
struct alignas(16) BStage {
volatile uint8_t prod_to_mma;
volatile uint8_t mma_to_prod;
alignas(16) uint8_t b_base[FP4_TILE_BYTES];
alignas(4) uint32_t bs_base[SCALE_PACKED_WORDS];
};
struct LdsPipelineState {
AStage a_stage[MMA_WAVES][A_STAGES_PER_CONSUMER];
BStage b_stage[MMA_WAVES][B_STAGES_PER_CONSUMER];
alignas(16) float reduce_accum[MAX_LOCAL_TILES][MMA_WAVES][WAVE_SIZE][4];
volatile uint8_t reduce_done[MAX_LOCAL_TILES][MMA_WAVES];
};
struct KernelState {
int m;
int n;
int k;
int lane_id;
int wave_id;
int n_tiles_per_cta;
int tile_m_count;
int tile_n_count;
int tile_count;
int tile_m0;
int tile_n0;
int k_tiles;
int active_mma_waves;
int role;
int consumer_id;
int split_id;
int a_prod_wave_id;
int b_prod_wave_id;
int k_tile_start;
int k_tile_count;
};
struct LaneArgs {
FragRegs a_regs = {}, b_regs = {};
uint32_t a_scale = 0, b_scale = 0;
uint8_t stage_a = 0, stage_b = 0;
uint8_t a_parity = 0, b_parity = 0;
f32x4 acc_local[MAX_LOCAL_TILES] = {};
};
__device__ __forceinline__ void global_load_lds_dwordx4_async(
uint32_t lds_byte_offset,
uint64_t src_addr
) {
asm volatile(
"s_mov_b32 m0, %0\\n\\t"
"global_load_lds_dwordx4 %1, off\\n\\t"
:
: "s"(lds_byte_offset), "v"(src_addr)
: "memory");
}
__device__ __forceinline__ void global_load_lds_dword_async(
uint32_t lds_byte_offset,
uint64_t src_addr
) {
asm volatile(
"s_mov_b32 m0, %0\\n\\t"
"global_load_lds_dword %1, off\\n\\t"
:
: "s"(lds_byte_offset), "v"(src_addr)
: "memory");
}
__device__ __forceinline__ void global_load_lds_ubyte_async(
uint32_t lds_byte_offset,
uint64_t src_addr
) {
asm volatile(
"s_mov_b32 m0, %0\\n\\t"
"global_load_lds_ubyte %1, off\\n\\t"
:
: "s"(lds_byte_offset), "v"(src_addr)
: "memory");
}
__device__ __forceinline__ void ds_read_b128_async(
uint32_t lds_byte_addr,
u32x4& out
) {
asm volatile("ds_read_b128 %0, %1\\n" : "=v"(out) : "v"(lds_byte_addr) : "memory");
}
__device__ __forceinline__ void ds_read_b64_tr_b4_pair_async(
uint32_t lds_byte_addr,
FragRegs& out
) {
asm volatile("ds_read_b64_tr_b4 %0, %1\\n" : "=v"(out.v2[0]) : "v"(lds_byte_addr) : "memory");
asm volatile("ds_read_b64_tr_b4 %0, %1 offset:128\\n" : "=v"(out.v2[1]) : "v"(lds_byte_addr) : "memory");
}
__device__ __forceinline__ void init_lds(LdsPipelineState& lds) {
const int tid = threadIdx.x +
blockDim.x * (threadIdx.y + blockDim.y * threadIdx.z);
const int threads_per_block = blockDim.x * blockDim.y * blockDim.z;
for (int idx = tid; idx < MMA_WAVES * A_STAGES_PER_CONSUMER; idx += threads_per_block) {
const int consumer = idx / A_STAGES_PER_CONSUMER;
const int stage = idx % A_STAGES_PER_CONSUMER;
lds.a_stage[consumer][stage].prod_to_mma = STAGE_PARITY_1;
lds.a_stage[consumer][stage].mma_to_prod = STAGE_PARITY_0;
}
for (int idx = tid; idx < MMA_WAVES * B_STAGES_PER_CONSUMER; idx += threads_per_block) {
const int consumer = idx / B_STAGES_PER_CONSUMER;
const int stage = idx % B_STAGES_PER_CONSUMER;
lds.b_stage[consumer][stage].prod_to_mma = STAGE_PARITY_1;
lds.b_stage[consumer][stage].mma_to_prod = STAGE_PARITY_0;
}
for (int idx = tid; idx < MAX_LOCAL_TILES * MMA_WAVES; idx += threads_per_block) {
reinterpret_cast<volatile uint8_t*>(lds.reduce_done)[idx] = 0u;
}
__syncthreads();
}
__device__ __forceinline__ int ceil_div_i32(int x, int y) {
return (x + y - 1) / y;
}
__device__ __forceinline__ void decode_a_stage(int k_idx, LaneArgs& lane_args) {
lane_args.stage_a = k_idx % A_STAGES_PER_CONSUMER;
lane_args.a_parity = (k_idx / A_STAGES_PER_CONSUMER) & 1;
}
__device__ __forceinline__ void decode_b_stage(
int tile_count,
int k_idx,
int tile_idx,
LaneArgs& lane_args
) {
const int token_b = k_idx * tile_count + tile_idx;
lane_args.stage_b = token_b % B_STAGES_PER_CONSUMER;
lane_args.b_parity = (token_b / B_STAGES_PER_CONSUMER) & 1;
}
__device__ __forceinline__ int owner_b_prod_wave_for_stage(int stage_b) {
return stage_b / (B_STAGES_PER_CONSUMER / B_PRODUCERS_PER_CONSUMER);
}
__device__ __forceinline__ void split_range_for_consumer(
const KernelState& state,
int consumer_id,
int& k_tile_start,
int& k_tile_count
) {
const int base_k = state.k_tiles / state.active_mma_waves;
const int rem_k = state.k_tiles % state.active_mma_waves;
k_tile_count = base_k + (consumer_id < rem_k ? 1 : 0);
k_tile_start = consumer_id * base_k + (consumer_id < rem_k ? consumer_id : rem_k);
}
__device__ __forceinline__ void init_wave_state(KernelState& state) {
state.k_tiles = ceil_div_i32(state.k, BLOCK);
state.active_mma_waves = state.k_tiles < MMA_WAVES ? state.k_tiles : MMA_WAVES;
state.role = 3;
state.consumer_id = -1;
state.split_id = -1;
state.a_prod_wave_id = -1;
state.b_prod_wave_id = -1;
state.k_tile_start = 0;
state.k_tile_count = 0;
if (state.wave_id < MMA_WAVES) {
if (state.wave_id < state.active_mma_waves) {
state.role = 2;
state.consumer_id = state.wave_id;
state.split_id = state.wave_id;
split_range_for_consumer(state, state.consumer_id, state.k_tile_start, state.k_tile_count);
}
return;
}
if (state.wave_id < MMA_WAVES + A_PRODUCER_WAVES) {
const int consumer_id = state.wave_id - MMA_WAVES;
if (consumer_id < state.active_mma_waves) {
state.role = 0;
state.consumer_id = consumer_id;
state.a_prod_wave_id = 0;
split_range_for_consumer(state, state.consumer_id, state.k_tile_start, state.k_tile_count);
}
return;
}
if (state.wave_id < TOTAL_WAVES) {
const int local_b_wave = state.wave_id - MMA_WAVES - A_PRODUCER_WAVES;
const int consumer_id = local_b_wave >> 1;
if (consumer_id < state.active_mma_waves) {
state.role = 1;
state.consumer_id = consumer_id;
state.b_prod_wave_id = local_b_wave & 1;
split_range_for_consumer(state, state.consumer_id, state.k_tile_start, state.k_tile_count);
}
}
}
__device__ __forceinline__ void decode_scheduled_tile(
const KernelState& state,
int local_tile_idx,
int& tile_m,
int& tile_n,
int& tile_bs
) {
tile_m = state.tile_m0;
tile_n = state.tile_n0 + local_tile_idx;
tile_bs = tile_n * MMA_N;
}
__device__ __forceinline__ void init_state(
KernelState& state,
int m,
int n,
int k,
int n_tiles_per_cta
) {
const int tid = threadIdx.x +
blockDim.x * (threadIdx.y + blockDim.y * threadIdx.z);
state.m = m;
state.n = n;
state.k = k;
state.lane_id = tid & (WAVE_SIZE - 1);
state.wave_id = tid / WAVE_SIZE;
state.n_tiles_per_cta = n_tiles_per_cta > 0 ? n_tiles_per_cta : 1;
state.tile_m_count = ceil_div_i32(m, MMA_M);
state.tile_n_count = ceil_div_i32(n, MMA_N);
state.tile_m0 = blockIdx.y;
state.tile_n0 = blockIdx.x * state.n_tiles_per_cta;
if (state.tile_m0 >= state.tile_m_count || state.tile_n0 >= state.tile_n_count) {
state.tile_count = 0;
state.tile_m0 = 0;
state.tile_n0 = 0;
return;
}
const int remaining_n_tiles = state.tile_n_count - state.tile_n0;
state.tile_count = remaining_n_tiles < state.n_tiles_per_cta ? remaining_n_tiles : state.n_tiles_per_cta;
init_wave_state(state);
}
__device__ __forceinline__ void mfma_scale_f32_16x16x128_fp4(
LaneArgs& lane_args,
int tile_idx
) {
asm volatile(
"v_mfma_scale_f32_16x16x128_f8f6f4 %0, %1, %2, %0, %3, %4 "
"op_sel:[0,0,0] op_sel_hi:[0,0,0] cbsz:4 blgp:4\\n"
: "+v"(lane_args.acc_local[tile_idx])
: "v"(lane_args.a_regs.v4), "v"(lane_args.b_regs.v4), "v"(lane_args.a_scale), "v"(lane_args.b_scale)
: "memory");
}
__device__ __forceinline__ uint32_t bitcast_f32_to_u32(float x) {
return __builtin_bit_cast(uint32_t, x);
}
__device__ __forceinline__ float bitcast_u32_to_f32(uint32_t x) {
return __builtin_bit_cast(float, x);
}
__device__ __forceinline__ float bf16_bits_to_f32(uint16_t bits) {
return bitcast_u32_to_f32(uint32_t(bits) << 16);
}
__device__ __forceinline__ void compute_mxfp4_scale(
float max_abs,
uint8_t& scale_byte,
float& quant_scale
) {
const uint32_t rounded_amax = (bitcast_f32_to_u32(max_abs) + 0x00200000u) & 0xFF800000u;
int scale_unbiased = int((rounded_amax >> 23) & 0xFFu) - 127 - 2;
scale_unbiased = scale_unbiased < -127 ? -127 : (scale_unbiased > 127 ? 127 : scale_unbiased);
scale_byte = scale_unbiased + 127;
quant_scale = bitcast_u32_to_f32(uint32_t(127 - scale_unbiased) << 23);
}
__device__ __forceinline__ uint32_t padded_scale_cols(int k) {
const uint32_t scale_cols = ceil_div_i32(k, SCALE_GROUP);
return (scale_cols + 7u) & ~7u;
}
__device__ __forceinline__ size_t shuffled_scale_byte_offset(
uint32_t row,
uint32_t k_group,
uint32_t padded_k_groups
) {
const uint32_t row_block = row >> 5;
const uint32_t row_half = (row >> 4) & 1u;
const uint32_t row_in_half = row & 15u;
const uint32_t group_block = k_group >> 3;
const uint32_t group_half = (k_group >> 2) & 1u;
const uint32_t group_in_block = k_group & 3u;
const size_t block_base =
(size_t(row_block) * (padded_k_groups >> 3) + group_block) * 256u;
const size_t within_block =
size_t(group_in_block) * 64u +
size_t(row_in_half) * 4u +
size_t(group_half) * 2u +
size_t(row_half);
return block_base + within_block;
}
__device__ __forceinline__ uint8_t pack_e2m1_from_f32(float x) {
constexpr uint32_t E8_BIAS = 127u;
constexpr uint32_t E2_BIAS = 1u;
constexpr uint32_t MBITS_F32 = 23u;
constexpr uint32_t MBITS_FP4 = 1u;
constexpr float MAX_NORMAL = 6.0f;
constexpr float MIN_NORMAL = 1.0f;
uint32_t qx = bitcast_f32_to_u32(x);
const uint32_t sign = qx & 0x80000000u;
qx ^= sign;
const float qx_f32 = bitcast_u32_to_f32(qx);
const bool saturate = qx_f32 >= MAX_NORMAL;
const bool denormal = (!saturate) && (qx_f32 < MIN_NORMAL);
uint8_t e2m1 = 0x7u;
if (denormal) {
constexpr uint32_t denorm_exp = (E8_BIAS - E2_BIAS) + (MBITS_F32 - MBITS_FP4) + 1u;
constexpr uint32_t denorm_mask_int = denorm_exp << MBITS_F32;
const float denorm_mask_float = bitcast_u32_to_f32(denorm_mask_int);
uint32_t denormal_x = bitcast_f32_to_u32(qx_f32 + denorm_mask_float);
denormal_x -= denorm_mask_int;
e2m1 = static_cast<uint8_t>(denormal_x);
} else if (!saturate) {
uint32_t normal_x = qx;
const uint32_t mant_odd = (normal_x >> (MBITS_F32 - MBITS_FP4)) & 1u;
const uint32_t val_to_add =
(uint32_t(E2_BIAS - E8_BIAS) << MBITS_F32) + (1u << 21) - 1u;
normal_x += val_to_add;
normal_x += mant_odd;
normal_x >>= (MBITS_F32 - MBITS_FP4);
e2m1 = static_cast<uint8_t>(normal_x);
}
const uint8_t sign_lp = static_cast<uint8_t>(sign >> (MBITS_F32 + 8u - MBITS_FP4 - 2u));
return static_cast<uint8_t>(e2m1 | sign_lp);
}
__device__ __forceinline__ uint32_t pack_fp4_word_bf16(
const uint32_t* bf16_pairs,
float quant_scale
) {
uint32_t packed = 0u;
#pragma unroll
for (int byte_idx = 0; byte_idx < 4; ++byte_idx) {
const uint32_t pair_word = bf16_pairs[byte_idx];
const float lo = bf16_bits_to_f32(static_cast<uint16_t>(pair_word)) * quant_scale;
const float hi = bf16_bits_to_f32(static_cast<uint16_t>(pair_word >> 16)) * quant_scale;
const uint8_t out_byte = pack_e2m1_from_f32(lo) | (pack_e2m1_from_f32(hi) << 4);
packed |= uint32_t(out_byte) << (byte_idx * 8);
}
return packed;
}
__device__ __forceinline__ void load_a_tile(
const char* a_offset,
const KernelState& state,
int row,
int k_block,
bool lane_valid,
u32x4* src_chunks
) {
if (lane_valid) {
const char* lane_src =
a_offset +
size_t(row * state.k + k_block * SCALE_GROUP) * sizeof(__hip_bfloat16);
const u32x4* lane_chunks = reinterpret_cast<const u32x4*>(lane_src);
src_chunks[0] = lane_chunks[0];
src_chunks[1] = lane_chunks[1];
src_chunks[2] = lane_chunks[2];
src_chunks[3] = lane_chunks[3];
}
asm volatile("s_waitcnt vmcnt(0) lgkmcnt(0)\\n\\t" ::: "memory");
}
__device__ __forceinline__ void scale_and_store_a_tile(
volatile uint32_t* a_words,
volatile uint8_t* as_bytes,
const KernelState& state,
const u32x4* src_chunks
) {
uint32_t bf16_pairs[16];
float max_abs = 0.0f;
#pragma unroll
for (int word_idx = 0; word_idx < 16; ++word_idx) {
const uint32_t pair_word = src_chunks[word_idx >> 2][word_idx & 3];
bf16_pairs[word_idx] = pair_word;
const float lo_abs = __builtin_fabsf(bf16_bits_to_f32(static_cast<uint16_t>(pair_word)));
const float hi_abs = __builtin_fabsf(bf16_bits_to_f32(static_cast<uint16_t>(pair_word >> 16)));
max_abs = lo_abs > max_abs ? lo_abs : max_abs;
max_abs = hi_abs > max_abs ? hi_abs : max_abs;
}
uint8_t scale_byte;
float quant_scale;
compute_mxfp4_scale(max_abs, scale_byte, quant_scale);
const uint32_t packed0 = pack_fp4_word_bf16(&bf16_pairs[0], quant_scale);
const uint32_t packed1 = pack_fp4_word_bf16(&bf16_pairs[4], quant_scale);
const uint32_t packed2 = pack_fp4_word_bf16(&bf16_pairs[8], quant_scale);
const uint32_t packed3 = pack_fp4_word_bf16(&bf16_pairs[12], quant_scale);
const size_t lane_word = size_t(state.lane_id) * 4;
a_words[lane_word + 0] = packed0;
a_words[lane_word + 1] = packed1;
a_words[lane_word + 2] = packed2;
a_words[lane_word + 3] = packed3;
as_bytes[state.lane_id] = scale_byte;
asm volatile("s_waitcnt lgkmcnt(0)\\n\\t" ::: "memory");
}
__device__ __forceinline__ void load_tile_lds_a(
const char* a_offset,
LdsPipelineState& lds,
const KernelState& state,
int consumer_id,
int stage_id,
uint8_t expected_parity
) {
while (lds.a_stage[consumer_id][stage_id].mma_to_prod != expected_parity) {}
volatile uint32_t* const a_words = reinterpret_cast<volatile uint32_t*>(lds.a_stage[consumer_id][stage_id].a_base);
volatile uint8_t* const as_bytes = reinterpret_cast<volatile uint8_t*>(lds.a_stage[consumer_id][stage_id].as_base);
int valid_rows = state.m - (state.tile_m0 * MMA_M);
valid_rows = valid_rows < 0 ? 0 : (valid_rows > MMA_M ? MMA_M : valid_rows);
const int row = state.lane_id & 15;
const int k_block = state.lane_id >> 4; // 1 row, 32 bf16 k vals per lane
const bool lane_valid = row < valid_rows;
u32x4 src_chunks[4] = {};
load_a_tile(a_offset, state, row, k_block, lane_valid, src_chunks);
scale_and_store_a_tile(a_words, as_bytes, state, src_chunks);
if (state.lane_id == 0) {
lds.a_stage[consumer_id][stage_id].prod_to_mma = expected_parity;
}
}
__device__ __forceinline__ void load_tile_lds_b(
const char* b_offset,
const uint8_t* b_scales,
LdsPipelineState& lds,
const KernelState& state,
int consumer_id,
int tile_n,
int k_idx,
int stage_id,
uint8_t expected_parity
) {
while (lds.b_stage[consumer_id][stage_id].mma_to_prod != expected_parity) {}
const uint32_t row = state.lane_id & 15;
const uint32_t k_group = state.lane_id >> 4;
const int valid_rows_raw = state.n - tile_n * MMA_N;
const int valid_rows = valid_rows_raw < 0 ? 0 : (valid_rows_raw > MMA_N ? MMA_N : valid_rows_raw);
const bool lane_valid = row < uint32_t(valid_rows);
const uint32_t off16 = row * 16 + k_group * 256;
const uint32_t global_row = tile_n * MMA_N + row;
const uint32_t global_k_group = k_idx * (BLOCK / SCALE_GROUP) + k_group;
const uint32_t scale_cols_padded = padded_scale_cols(state.k);
const uint64_t bs_src =
reinterpret_cast<uint64_t>(b_scales) +
shuffled_scale_byte_offset(global_row, global_k_group, scale_cols_padded);
volatile uint8_t* const b_stage_bytes = reinterpret_cast<volatile uint8_t*>(lds.b_stage[consumer_id][stage_id].b_base);
volatile uint8_t* const lane_b_dst = b_stage_bytes + size_t(state.lane_id) * 16;
volatile uint32_t* const lane_b_dst_words = reinterpret_cast<volatile uint32_t*>(const_cast<uint8_t*>(lane_b_dst));
// if (lane_valid) {
// const u32x4 b_vec = *reinterpret_cast<const u32x4*>(reinterpret_cast<const uint8_t*>(b_offset) + off16);
// lane_b_dst_words[0] = b_vec[0];
// lane_b_dst_words[1] = b_vec[1];
// lane_b_dst_words[2] = b_vec[2];
// lane_b_dst_words[3] = b_vec[3];
// } else {
// #pragma unroll
// for (int i = 0; i < 16; ++i) {
// lane_b_dst[i] = 0u;
// }
// }
// uint32_t packed_scale = 0;
// if (lane_valid) {
// asm volatile("flat_load_ubyte %0, %1\\n" : "=v"(packed_scale) : "v"(bs_src) : "memory");
// }
// asm volatile("s_waitcnt vmcnt(0) lgkmcnt(0)\\n\\t" ::: "memory");
// reinterpret_cast<volatile uint8_t*>(lds.b_stage[consumer_id][stage_id].bs_base)[state.lane_id] = static_cast<uint8_t>(packed_scale);
// asm volatile("s_waitcnt lgkmcnt(0)\\n\\t" ::: "memory");
uint32_t packed_scale = 0;
if (lane_valid) {
const uint64_t b_src_addr =
reinterpret_cast<uint64_t>(reinterpret_cast<const uint8_t*>(b_offset) + off16);
const uint32_t b_lds_off = __builtin_amdgcn_readfirstlane(static_cast<uint32_t>(
reinterpret_cast<const char*>(lds.b_stage[consumer_id][stage_id].b_base) -
reinterpret_cast<const char*>(&lds)));
global_load_lds_dwordx4_async(b_lds_off, b_src_addr);
asm volatile("flat_load_ubyte %0, %1\\n" : "=v"(packed_scale) : "v"(bs_src) : "memory");
}
asm volatile("s_waitcnt vmcnt(0) lgkmcnt(0)\\n\\t" ::: "memory");
reinterpret_cast<volatile uint8_t*>(lds.b_stage[consumer_id][stage_id].bs_base)[state.lane_id] = static_cast<uint8_t>(packed_scale);
asm volatile("s_waitcnt lgkmcnt(0)\\n\\t" ::: "memory");
if (state.lane_id == 0) {
lds.b_stage[consumer_id][stage_id].prod_to_mma = expected_parity;
}
}
__device__ __forceinline__ void producer_wave_a(
const __hip_bfloat16* a_full,
const KernelState& state,
LdsPipelineState& lds
) {
const char* a_base = reinterpret_cast<const char*>(a_full);
if (state.tile_count == 0 || state.k_tile_count <= 0) { return; }
const int m0 = state.tile_m0 * MMA_M;
const char* a_tile_base = a_base + size_t(m0) * size_t(state.k) * sizeof(__hip_bfloat16);
const size_t a_k_stride = size_t(BLOCK) * sizeof(__hip_bfloat16);
for (int local_k_idx = 0; local_k_idx < state.k_tile_count; ++local_k_idx) {
const uint8_t stage_a = local_k_idx % A_STAGES_PER_CONSUMER;
const uint8_t parity_a = (local_k_idx / A_STAGES_PER_CONSUMER) & 1;
const int k_idx = state.k_tile_start + local_k_idx;
const char* a_offset = a_tile_base + size_t(k_idx) * a_k_stride;
load_tile_lds_a(a_offset, lds, state, state.consumer_id, stage_a, parity_a);
}
}
__device__ __forceinline__ void producer_wave_b(
const uint8_t* b_shuffle,
const uint8_t* b_scales,
const KernelState& state,
LdsPipelineState& lds
) {
const char* b_base = reinterpret_cast<const char*>(b_shuffle);
if (state.tile_count == 0 || state.k_tile_count <= 0) { return; }
for (int local_k_idx = 0; local_k_idx < state.k_tile_count; ++local_k_idx) {
for (int tile_idx = 0; tile_idx < state.tile_count; ++tile_idx) {
const int abs_tile_idx = tile_idx;
int tile_m, tile_n, tile_bs;
decode_scheduled_tile(state, abs_tile_idx, tile_m, tile_n, tile_bs);
const int token_b = local_k_idx * state.tile_count + tile_idx;
const uint8_t stage_b = token_b % B_STAGES_PER_CONSUMER;
const uint8_t parity_b = (token_b / B_STAGES_PER_CONSUMER) & 1;
if (owner_b_prod_wave_for_stage(stage_b) != state.b_prod_wave_id) { continue; }
const int k_idx = state.k_tile_start + local_k_idx;
const size_t tile_offset =
size_t(tile_n) * MMA_N * (state.k / 2) +
size_t(k_idx) * MMA_N * (BLOCK / 2);
const char* b_offset = b_base + tile_offset;
load_tile_lds_b(b_offset, b_scales, lds, state, state.consumer_id, tile_n, k_idx, stage_b, parity_b);
}
}
}
__device__ __forceinline__ void store_k_split_results(
const KernelState& state,
LdsPipelineState& lds,
const LaneArgs& lane_args,
int tile_idx
) {
const int split = state.split_id;
lds.reduce_accum[tile_idx][split][state.lane_id][0] = lane_args.acc_local[tile_idx][0];
lds.reduce_accum[tile_idx][split][state.lane_id][1] = lane_args.acc_local[tile_idx][1];
lds.reduce_accum[tile_idx][split][state.lane_id][2] = lane_args.acc_local[tile_idx][2];
lds.reduce_accum[tile_idx][split][state.lane_id][3] = lane_args.acc_local[tile_idx][3];
asm volatile("s_waitcnt lgkmcnt(0)\\n\\t" ::: "memory");
if (state.lane_id == 0) {
lds.reduce_done[tile_idx][split] = 1u;
}
}
__device__ __forceinline__ void store_results(
const KernelState& state,
const LaneArgs& lane_args,
int acc_tile_idx,
int abs_tile_idx,
__hip_bfloat16* c
) {
const int tile_row = state.tile_m0 * MMA_M;
const int tile_col = (state.tile_n0 + abs_tile_idx) * MMA_N;
const int lane_col = state.lane_id & (MMA_N - 1);
const int lane_row_base = (state.lane_id >> 4) * 4;
#pragma unroll
for (int acc_idx = 0; acc_idx < 4; ++acc_idx) {
const int row = tile_row + lane_row_base + acc_idx;
const int col = tile_col + lane_col;
if (row < state.m && col < state.n) {
c[size_t(row) * state.n + col] = __hip_bfloat16(lane_args.acc_local[acc_tile_idx][acc_idx]);
}
}
}
__device__ __forceinline__ void reduce_and_store_results(
const KernelState& state,
LdsPipelineState& lds,
int tile_idx,
__hip_bfloat16* c
) {
while (true) {
bool ready = true;
#pragma unroll
for (int split = 0; split < MMA_WAVES; ++split) {
if (split < state.active_mma_waves && lds.reduce_done[tile_idx][split] == 0u) {
ready = false;
}
}
if (ready) { break; }
}
f32x4 acc = {0.0f, 0.0f, 0.0f, 0.0f};
#pragma unroll
for (int split = 0; split < MMA_WAVES; ++split) {
if (split < state.active_mma_waves) {
acc[0] += lds.reduce_accum[tile_idx][split][state.lane_id][0];
acc[1] += lds.reduce_accum[tile_idx][split][state.lane_id][1];
acc[2] += lds.reduce_accum[tile_idx][split][state.lane_id][2];
acc[3] += lds.reduce_accum[tile_idx][split][state.lane_id][3];
}
}
LaneArgs tmp{};
tmp.acc_local[0] = acc;
store_results(state, tmp, 0, tile_idx, c);
}
__device__ __forceinline__ void load_a_reg(
const KernelState& state,
LdsPipelineState& lds,
LaneArgs& lane_args
) {
while (lds.a_stage[state.consumer_id][lane_args.stage_a].prod_to_mma != lane_args.a_parity) {}
volatile uint64_t* const a_words =
reinterpret_cast<volatile uint64_t*>(lds.a_stage[state.consumer_id][lane_args.stage_a].a_base);
const size_t lane_qword = size_t(state.lane_id) * 2;
lane_args.a_regs.v2[0] = a_words[lane_qword + 0];
lane_args.a_regs.v2[1] = a_words[lane_qword + 1];
lane_args.a_regs.v2[2] = 0;
lane_args.a_regs.v2[3] = 0;
lane_args.a_scale = reinterpret_cast<volatile uint8_t*>(lds.a_stage[state.consumer_id][lane_args.stage_a].as_base)[state.lane_id];
asm volatile("s_waitcnt lgkmcnt(0)\\n" ::: "memory");
}
__device__ __forceinline__ void load_b_reg(
const KernelState& state,
LdsPipelineState& lds,
LaneArgs& lane_args,
bool release_a
) {
while (lds.b_stage[state.consumer_id][lane_args.stage_b].prod_to_mma != lane_args.b_parity) {}
volatile uint64_t* const b_words =
reinterpret_cast<volatile uint64_t*>(lds.b_stage[state.consumer_id][lane_args.stage_b].b_base);
const size_t lane_qword = size_t(state.lane_id) * 2;
lane_args.b_regs.v2[0] = b_words[lane_qword + 0];
lane_args.b_regs.v2[1] = b_words[lane_qword + 1];
lane_args.b_regs.v2[2] = 0;
lane_args.b_regs.v2[3] = 0;
lane_args.b_scale = reinterpret_cast<volatile uint8_t*>(lds.b_stage[state.consumer_id][lane_args.stage_b].bs_base)[state.lane_id];
asm volatile("s_waitcnt lgkmcnt(0)\\n" ::: "memory");
if (state.lane_id == 0) {
lds.b_stage[state.consumer_id][lane_args.stage_b].mma_to_prod = lane_args.b_parity ^ 1;
if (release_a) {
lds.a_stage[state.consumer_id][lane_args.stage_a].mma_to_prod = lane_args.a_parity ^ 1;
}
}
}
__device__ __forceinline__ void init_mma_args(const KernelState& state, LaneArgs& lane_args) {
const int local_tiles = state.tile_count < MAX_LOCAL_TILES ? state.tile_count : MAX_LOCAL_TILES;
#pragma unroll
for (int tile_idx = 0; tile_idx < local_tiles; ++tile_idx) {
lane_args.acc_local[tile_idx] = {0.0f, 0.0f, 0.0f, 0.0f};
}
}
__device__ __forceinline__ void mma_wave(
KernelState& state,
LdsPipelineState& lds,
__hip_bfloat16* c
) {
if (state.tile_count <= 0 || state.k_tile_count <= 0) { return; }
const int local_tiles = state.tile_count < MAX_LOCAL_TILES ? state.tile_count : MAX_LOCAL_TILES;
LaneArgs lane_args{};
init_mma_args(state, lane_args);
for (int local_k_idx = 0; local_k_idx < state.k_tile_count; ++local_k_idx) {
const bool final_k_tile = (local_k_idx == state.k_tile_count - 1);
decode_a_stage(local_k_idx, lane_args);
load_a_reg(state, lds, lane_args);
for (int tile_idx = 0; tile_idx < local_tiles; ++tile_idx) {
decode_b_stage(local_tiles, local_k_idx, tile_idx, lane_args);
const bool release_a = (tile_idx == local_tiles - 1);
load_b_reg(state, lds, lane_args, release_a);
lane_args.acc_local[tile_idx] =
__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
lane_args.a_regs.v8, lane_args.b_regs.v8, lane_args.acc_local[tile_idx],
4, 4, 0, lane_args.a_scale, 0 , lane_args.b_scale
);
// mfma_scale_f32_16x16x128_fp4(lane_args, tile_idx);
if (final_k_tile && state.active_mma_waves == 1) {
store_results(state, lane_args, tile_idx, tile_idx, c);
}
}
}
if (state.active_mma_waves > 1) {
for (int tile_idx = 0; tile_idx < local_tiles; ++tile_idx) {
store_k_split_results(state, lds, lane_args, tile_idx);
}
if (state.split_id == state.active_mma_waves - 1) {
for (int tile_idx = 0; tile_idx < local_tiles; ++tile_idx) {
reduce_and_store_results(state, lds, tile_idx, c);
}
}
}
}
__global__ void custom_kernel(const __hip_bfloat16* a_full, const uint8_t* b_shuffle,
const uint8_t* b_scales, __hip_bfloat16* c, int m, int n, int k, int n_tiles_per_cta) {
__shared__ LdsPipelineState lds;
KernelState state;
init_lds(lds);
init_state(state, m, n, k, n_tiles_per_cta); if (state.tile_count == 0) {return;}
if (state.role == 0) { producer_wave_a(a_full, state, lds); }
if (state.role == 1) {
producer_wave_b(b_shuffle, b_scales, state, lds);
}
if (state.role == 2) {
mma_wave(state, lds, c);
}
}
void fp8_mm(torch::Tensor a_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 int tile_m_count = (m + MMA_M - 1) / MMA_M; // launch ceil 256 ctas
const int tile_n_count = (n + MMA_N - 1) / MMA_N;
const int grid_n_target = max(1, NUM_WORKERS / max(1, tile_m_count));
const int n_tiles_raw = (tile_n_count + grid_n_target - 1) / grid_n_target;
const int n_tiles_per_cta = n_tiles_raw < 1 ? 1 : (n_tiles_raw > MAX_N_TILES_PER_CTA ? MAX_N_TILES_PER_CTA : n_tiles_raw);
const int grid_n = (tile_n_count + n_tiles_per_cta - 1) / n_tiles_per_cta;
const int num_wavefronts = TOTAL_WAVES;
const dim3 grid(grid_n, tile_m_count);
custom_kernel<<<grid, num_wavefronts * 64, 0, 0>>>
((const __hip_bfloat16*)a_full.data_ptr(), (const uint8_t*)b_shuffle.data_ptr(),
(const uint8_t*)b_scales.data_ptr(), (__hip_bfloat16*)c.data_ptr(),
m, n, k, n_tiles_per_cta);
}
"""
import os
os.environ["CXX"] = "clang++"
module = load_inline(
name='fp8_mm',
cpp_sources=[CPP_WRAPPER],
cuda_sources=[cuda_src],
functions=['fp8_mm'],
verbose=True,
extra_cuda_cflags=["--offload-arch=gfx950",
"-Rpass-analysis=kernel-resource-usage",
"--save-temps",
"-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)
flat = b_full.view(-1)
c = flat.narrow(0, 0, m * n).view(m, n)
module.fp8_mm(a_full, b_shuffle, b_scales, c)
return c
scrolls · 887 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