Skip to content
KernelIndex
Search⌘K

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
AMD MXFP4 GEMMsuite of 6 cases
AMD Instinct MI355X
28.1µs
#1111 of 1143
2026-03-16

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-specializationconstexpr 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