submission 590126
jIab-b · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1305 lines, June 9 Researcher Reciprocity License v1.0.
proto2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-590126?include=source"interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, fp32, fp8_e8m0, int32, 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:b3b112a9c1cdd6ef27aa2da2cb93146c4e7075cf0cf908a2a905f117eb58e443
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__ LdsState lds;stages = 3
constexpr int A_STAGES = 3;warp-specialization
constexpr int A_PRODUCER_WAVES = 4;Kernel source
proto2.py1305 lines
import os
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
os.environ["CXX"] = "clang++"
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
CPP_WRAPPER = """
torch::Tensor fused_moe(
torch::Tensor hidden_states,
torch::Tensor gate_up_weight,
torch::Tensor down_weight,
torch::Tensor gate_up_weight_scale,
torch::Tensor down_weight_scale,
torch::Tensor gate_up_weight_shuffled,
torch::Tensor down_weight_shuffled,
torch::Tensor gate_up_weight_scale_shuffled,
torch::Tensor down_weight_scale_shuffled,
torch::Tensor topk_weights,
torch::Tensor topk_ids,
int d_hidden,
int d_expert,
int d_hidden_pad,
int d_expert_pad,
int n_routed_experts,
int n_shared_experts,
int n_experts_per_token,
int total_top_k
);
"""
cuda_src = """
#include <torch/extension.h>
#include <c10/cuda/CUDAGuard.h>
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <vector>
#include <hip/amd_detail/amd_hip_fp8.h>
#include <hip/amd_detail/amd_hip_bf16.h>
// ============================================================
// Tunable schedule parameters
// ============================================================
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 CTA_M = MMA_M;
constexpr int FP4_TILE_BYTES = MMA_M * BLOCK / 2; // 1024
constexpr int SCALE_PACKED_WORDS = MMA_M; // 16 scale bytes, stored as 16 uint32
// --- Wave specialization config ---
constexpr int MMA_WAVES = 8;
constexpr int A_PRODUCER_WAVES = 4;
constexpr int B_PRODUCER_WAVES = 4;
constexpr int TOTAL_WAVES = MMA_WAVES + A_PRODUCER_WAVES + B_PRODUCER_WAVES;
// Each A-producer serves 2 MMA consumers, each B-producer serves 2 MMA consumers
constexpr int A_CONSUMERS_PER_PRODUCER = MMA_WAVES / A_PRODUCER_WAVES;
constexpr int B_CONSUMERS_PER_PRODUCER = MMA_WAVES / B_PRODUCER_WAVES;
// --- Pipeline depth ---
constexpr int A_STAGES = 3;
constexpr int B_STAGES = 3;
// --- Tiling ---
constexpr int FUSED_TILE_COLS = 256;
constexpr int INTERMEDIATE_TILE_COLS = FUSED_TILE_COLS / 2;
constexpr int LOCAL_TILES_S1 = FUSED_TILE_COLS / MMA_N; // 16 N-tiles for stage 1
constexpr int S2_N_TILES_PER_MMA = 1; // each MMA wave handles 1 N-tile at a time in stage 2
constexpr uint32_t INVALID_TOKEN = 0xFFFFFFFFu;
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 FragRegs {
u32x4 v4;
u32x8 v8;
uint64_t v2[4];
};
struct DispatchEntry {
int32_t token_id;
float weight;
};
// ============================================================
// LDS layout for producer-consumer pipeline
// ============================================================
struct alignas(16) AStage {
volatile uint8_t prod_to_mma;
volatile uint8_t mma_to_prod;
alignas(16) uint8_t a_data[FP4_TILE_BYTES];
alignas(4) uint32_t a_scales[SCALE_PACKED_WORDS];
};
struct alignas(16) BStage {
volatile uint8_t prod_to_mma;
volatile uint8_t mma_to_prod;
alignas(16) uint8_t b_data[FP4_TILE_BYTES];
alignas(4) uint32_t b_scales[SCALE_PACKED_WORDS];
};
struct LdsState {
// Dispatch metadata
uint32_t token_ids[CTA_M];
float topk_weights[CTA_M];
// Pipeline stages
AStage a_stage[MMA_WAVES][A_STAGES];
BStage b_stage[MMA_WAVES][B_STAGES];
// Stage 1 partial results & intermediate
float partial_f32[LOCAL_TILES_S1 * MMA_M * MMA_N];
uint16_t intermediate_bf16[(LOCAL_TILES_S1 / 2) * MMA_M * MMA_N];
// Stage 2: k-split reduction across MMA waves (split along N, so no partial sums needed)
// MMA waves split along N-tiles in stage 2, each writes its own output columns directly.
};
// ============================================================
// Kernel state
// ============================================================
struct KernelState {
int lane_id;
int wave_id;
int expert_id;
int token_tile_count;
int fused_tile_idx;
int d_hidden;
int d_hidden_pad;
int d_expert;
int d_expert_pad;
int hidden_row_stride_elems;
int output_row_base;
int k_offset;
// Wave role: 0=A-producer, 1=B-producer, 2=MMA-consumer, 3=idle
int role;
int consumer_id; // MMA consumer index (0..MMA_WAVES-1)
int prod_consumer_start; // first consumer this producer serves
int prod_consumer_count; // how many consumers this producer serves
// Stage 1 k-split for MMA consumers
int s1_k_tile_start;
int s1_k_tile_count;
// Stage 2 n-split for MMA consumers
int s2_n_tile_start;
int s2_n_tile_count;
bool valid;
const __hip_bfloat16* hidden_states;
const uint8_t* b_shuffle;
const uint8_t* b_scales;
const uint8_t* down_shuffle;
const uint8_t* down_scales;
float* final_out;
int final_out_row_stride_elems;
};
struct LaneArgs {
FragRegs a_regs = {}, b_regs = {};
uint32_t a_scale = 0, b_scale = 0;
f32x4 acc_local[LOCAL_TILES_S1] = {};
};
// ============================================================
// Utility functions (unchanged from proto.py)
// ============================================================
__device__ __forceinline__ int ceil_div_i32(int x, int y) {
return (x + y - 1) / y;
}
__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__ uint16_t f32_to_bf16_bits(float x) {
uint32_t bits = bitcast_f32_to_u32(x);
const uint32_t lsb = (bits >> 16) & 1u;
bits += 0x7FFFu + lsb;
return uint16_t(bits >> 16);
}
__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;
const uint32_t lo = v & 0xFFFFu;
return hi > lo ? hi : lo;
}
__device__ __forceinline__ void compute_mxfp4_scale_from_bf16max(
uint32_t max_bf16_abs,
uint32_t& scale_byte,
float& quant_scale
) {
const float max_f32 = bitcast_u32_to_f32(max_bf16_abs << 16);
const uint32_t raw_exp = (bitcast_f32_to_u32(max_f32) + 0x00200000u) >> 23;
scale_byte = raw_exp > 2u ? (raw_exp - 2u) : 0u;
quant_scale = bitcast_u32_to_f32(scale_byte << 23);
}
struct Fp4Thresholds {
uint16_t t[7];
};
__device__ __forceinline__ void compute_fp4_thresholds(float quant_scale, Fp4Thresholds& thr) {
thr.t[0] = f32_to_bf16_bits(quant_scale * 0.25f);
thr.t[1] = f32_to_bf16_bits(quant_scale * 0.75f);
thr.t[2] = f32_to_bf16_bits(quant_scale * 1.25f);
thr.t[3] = f32_to_bf16_bits(quant_scale * 1.75f);
thr.t[4] = f32_to_bf16_bits(quant_scale * 2.5f);
thr.t[5] = f32_to_bf16_bits(quant_scale * 3.5f);
thr.t[6] = f32_to_bf16_bits(quant_scale * 5.0f);
}
__device__ __forceinline__ uint8_t encode_fp4_nibble_sw_exact(uint16_t bits, const Fp4Thresholds& thr) {
const uint16_t absbits = bits & 0x7FFFu;
if (absbits == 0u) return 0u;
const uint8_t mag =
uint8_t(absbits >= thr.t[0]) +
uint8_t(absbits >= thr.t[1]) +
uint8_t(absbits >= thr.t[2]) +
uint8_t(absbits >= thr.t[3]) +
uint8_t(absbits >= thr.t[4]) +
uint8_t(absbits >= thr.t[5]) +
uint8_t(absbits >= thr.t[6]);
if (mag == 0u) return 0u;
return (bits & 0x8000u) ? uint8_t(mag | 0x8u) : mag;
}
__device__ __forceinline__ uint32_t pack_fp4_word_sw_exact(
const uint32_t* bf16_pairs,
const Fp4Thresholds& thr
) {
uint32_t packed = 0u;
#pragma unroll
for (int byte_idx = 0; byte_idx < 4; ++byte_idx) {
const uint32_t pair = bf16_pairs[byte_idx];
const uint8_t lo = encode_fp4_nibble_sw_exact(uint16_t(pair & 0xFFFFu), thr);
const uint8_t hi = encode_fp4_nibble_sw_exact(uint16_t(pair >> 16), thr);
packed |= uint32_t(lo | uint8_t(hi << 4)) << (byte_idx * 8);
}
return packed;
}
__device__ __forceinline__ void quantize_bf16_chunks_to_fp4(
const u32x4* src_chunks,
FragRegs& dst_regs,
uint32_t& scale_byte
) {
const uint32_t SIGN_MASK = 0x7FFF7FFFu;
uint32_t m0 = pk_max_u16(src_chunks[0][0] & SIGN_MASK, src_chunks[0][1] & SIGN_MASK);
uint32_t m1 = pk_max_u16(src_chunks[0][2] & SIGN_MASK, src_chunks[0][3] & SIGN_MASK);
uint32_t m2 = pk_max_u16(src_chunks[1][0] & SIGN_MASK, src_chunks[1][1] & SIGN_MASK);
uint32_t m3 = pk_max_u16(src_chunks[1][2] & SIGN_MASK, src_chunks[1][3] & SIGN_MASK);
uint32_t m4 = pk_max_u16(src_chunks[2][0] & SIGN_MASK, src_chunks[2][1] & SIGN_MASK);
uint32_t m5 = pk_max_u16(src_chunks[2][2] & SIGN_MASK, src_chunks[2][3] & SIGN_MASK);
uint32_t m6 = pk_max_u16(src_chunks[3][0] & SIGN_MASK, src_chunks[3][1] & SIGN_MASK);
uint32_t m7 = pk_max_u16(src_chunks[3][2] & SIGN_MASK, src_chunks[3][3] & SIGN_MASK);
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 quant_scale = 0.0f;
compute_mxfp4_scale_from_bf16max(reduce_pk_u16(m0), scale_byte, quant_scale);
dst_regs.v4 = {0u, 0u, 0u, 0u};
dst_regs.v2[2] = 0; dst_regs.v2[3] = 0;
if (scale_byte == 0u) return;
Fp4Thresholds thr{};
compute_fp4_thresholds(quant_scale, thr);
const uint32_t* raw = reinterpret_cast<const uint32_t*>(src_chunks);
#pragma unroll
for (int word_idx = 0; word_idx < 4; ++word_idx) {
dst_regs.v4[word_idx] = pack_fp4_word_sw_exact(&raw[word_idx * 4], thr);
}
}
__device__ __forceinline__ float silu_f32(float x) {
return x / (1.0f + __builtin_expf(-x));
}
__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;
}
// ============================================================
// Async global→LDS load intrinsic
// ============================================================
__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");
}
// Async LDS→register read
__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");
}
// ============================================================
// Pipeline init
// ============================================================
__device__ __forceinline__ void init_pipeline(LdsState& lds) {
const int tid = threadIdx.x;
const int total_threads = TOTAL_WAVES * WAVE_SIZE;
for (int idx = tid; idx < MMA_WAVES * A_STAGES; idx += total_threads) {
const int consumer = idx / A_STAGES;
const int stage = idx % A_STAGES;
lds.a_stage[consumer][stage].prod_to_mma = 1u; // not ready
lds.a_stage[consumer][stage].mma_to_prod = 0u; // available
}
for (int idx = tid; idx < MMA_WAVES * B_STAGES; idx += total_threads) {
const int consumer = idx / B_STAGES;
const int stage = idx % B_STAGES;
lds.b_stage[consumer][stage].prod_to_mma = 1u;
lds.b_stage[consumer][stage].mma_to_prod = 0u;
}
__syncthreads();
}
// ============================================================
// Wave role assignment
// ============================================================
__device__ __forceinline__ void split_range(
int total, int num_splits, int split_id,
int& start, int& count
) {
const int base = total / num_splits;
const int rem = total % num_splits;
count = base + (split_id < rem ? 1 : 0);
start = split_id * base + (split_id < rem ? split_id : rem);
}
__device__ __forceinline__ void init_wave_roles(KernelState& state, int s1_k_tiles, int s2_n_tiles) {
state.role = 3; // idle by default
state.consumer_id = -1;
state.prod_consumer_start = 0;
state.prod_consumer_count = 0;
state.s1_k_tile_start = 0;
state.s1_k_tile_count = 0;
state.s2_n_tile_start = 0;
state.s2_n_tile_count = 0;
const int active_mma = s1_k_tiles < MMA_WAVES ? s1_k_tiles : MMA_WAVES;
// Waves 0..MMA_WAVES-1: MMA consumers
if (state.wave_id < MMA_WAVES) {
if (state.wave_id < active_mma) {
state.role = 2;
state.consumer_id = state.wave_id;
// Stage 1: split K across MMA waves
split_range(s1_k_tiles, active_mma, state.wave_id,
state.s1_k_tile_start, state.s1_k_tile_count);
// Stage 2: split N across MMA waves
split_range(s2_n_tiles, active_mma, state.wave_id,
state.s2_n_tile_start, state.s2_n_tile_count);
}
return;
}
// Waves MMA_WAVES..MMA_WAVES+A_PRODUCER_WAVES-1: A producers
if (state.wave_id < MMA_WAVES + A_PRODUCER_WAVES) {
const int a_idx = state.wave_id - MMA_WAVES;
const int first = a_idx * A_CONSUMERS_PER_PRODUCER;
if (first < active_mma) {
state.role = 0;
state.prod_consumer_start = first;
state.prod_consumer_count =
(first + A_CONSUMERS_PER_PRODUCER <= active_mma)
? A_CONSUMERS_PER_PRODUCER
: (active_mma - first);
}
return;
}
// Waves MMA_WAVES+A_PRODUCER_WAVES..TOTAL_WAVES-1: B producers
if (state.wave_id < TOTAL_WAVES) {
const int b_idx = state.wave_id - MMA_WAVES - A_PRODUCER_WAVES;
const int first = b_idx * B_CONSUMERS_PER_PRODUCER;
if (first < active_mma) {
state.role = 1;
state.prod_consumer_start = first;
state.prod_consumer_count =
(first + B_CONSUMERS_PER_PRODUCER <= active_mma)
? B_CONSUMERS_PER_PRODUCER
: (active_mma - first);
}
}
}
// ============================================================
// Stage 1: A-producer — load bf16 from gmem, quantize, store to LDS
// ============================================================
__device__ __forceinline__ void a_prod_store_to_lds(
AStage& stage,
const KernelState& state,
const u32x4* src_chunks
) {
FragRegs regs{};
uint32_t scale_byte = 0;
quantize_bf16_chunks_to_fp4(src_chunks, regs, scale_byte);
// Store quantized fp4 data to LDS stage buffer
volatile uint32_t* a_words = reinterpret_cast<volatile uint32_t*>(stage.a_data);
const size_t lane_word = size_t(state.lane_id) * 4;
volatile u32x4* dst = reinterpret_cast<volatile u32x4*>(&a_words[lane_word]);
*dst = regs.v4;
reinterpret_cast<volatile uint8_t*>(stage.a_scales)[state.lane_id] = static_cast<uint8_t>(scale_byte);
asm volatile("s_waitcnt lgkmcnt(0)\\n\\t" ::: "memory");
}
__device__ __forceinline__ void s1_a_producer(
const KernelState& state,
LdsState& lds
) {
for (int ci = 0; ci < state.prod_consumer_count; ++ci) {
const int consumer_id = state.prod_consumer_start + ci;
int cons_k_start, cons_k_count;
const int s1_k_tiles = ceil_div_i32(state.d_hidden_pad, BLOCK);
const int active_mma = s1_k_tiles < MMA_WAVES ? s1_k_tiles : MMA_WAVES;
split_range(s1_k_tiles, active_mma, consumer_id, cons_k_start, cons_k_count);
for (int local_k = 0; local_k < cons_k_count; ++local_k) {
const int stage_id = local_k % A_STAGES;
const uint8_t parity = (local_k / A_STAGES) & 1;
// Wait for consumer to release this stage
while (lds.a_stage[consumer_id][stage_id].mma_to_prod != parity) {}
const int k_idx = cons_k_start + local_k;
// Load bf16 hidden_states from gmem → regs, quantize, store to LDS
const uint32_t row = state.lane_id & 15;
const uint32_t k_group = state.lane_id >> 4;
const bool lane_valid = row < uint32_t(state.token_tile_count)
&& lds.token_ids[row] != INVALID_TOKEN;
u32x4 src_chunks[4] = {};
if (lane_valid) {
const uint32_t token_id = lds.token_ids[row];
const uint32_t base_k = uint32_t(k_idx * BLOCK + int(k_group) * SCALE_GROUP);
const uint16_t* hidden_bits =
reinterpret_cast<const uint16_t*>(state.hidden_states) +
size_t(token_id) * size_t(state.hidden_row_stride_elems);
if (base_k + SCALE_GROUP <= uint32_t(state.d_hidden)) {
const u32x4* lane_chunks = reinterpret_cast<const u32x4*>(hidden_bits + base_k);
src_chunks[0] = lane_chunks[0];
src_chunks[1] = lane_chunks[1];
src_chunks[2] = lane_chunks[2];
src_chunks[3] = lane_chunks[3];
} else {
uint32_t* raw_ptr = reinterpret_cast<uint32_t*>(src_chunks);
#pragma unroll
for (int w = 0; w < 16; ++w) {
const uint32_t e0 = base_k + uint32_t(w * 2);
const uint32_t e1 = base_k + uint32_t(w * 2 + 1);
const uint32_t lo = e0 < uint32_t(state.d_hidden) ? uint32_t(hidden_bits[e0]) : 0u;
const uint32_t hi = e1 < uint32_t(state.d_hidden) ? (uint32_t(hidden_bits[e1]) << 16) : 0u;
raw_ptr[w] = lo | hi;
}
}
}
a_prod_store_to_lds(lds.a_stage[consumer_id][stage_id], state, src_chunks);
if (state.lane_id == 0) {
lds.a_stage[consumer_id][stage_id].prod_to_mma = parity;
}
}
}
}
// ============================================================
// Stage 1: B-producer — async DMA fp4 weights from gmem → LDS
// ============================================================
__device__ __forceinline__ void s1_b_producer(
const KernelState& state,
LdsState& lds
) {
for (int ci = 0; ci < state.prod_consumer_count; ++ci) {
const int consumer_id = state.prod_consumer_start + ci;
int cons_k_start, cons_k_count;
const int s1_k_tiles = ceil_div_i32(state.d_hidden_pad, BLOCK);
const int active_mma = s1_k_tiles < MMA_WAVES ? s1_k_tiles : MMA_WAVES;
split_range(s1_k_tiles, active_mma, consumer_id, cons_k_start, cons_k_count);
for (int local_k = 0; local_k < cons_k_count; ++local_k) {
for (int tile_idx = 0; tile_idx < LOCAL_TILES_S1; ++tile_idx) {
const int token_b = local_k * LOCAL_TILES_S1 + tile_idx;
const int stage_id = token_b % B_STAGES;
const uint8_t parity = (token_b / B_STAGES) & 1;
while (lds.b_stage[consumer_id][stage_id].mma_to_prod != parity) {}
const int k_idx = cons_k_start + local_k;
// Compute gate_up weight tile address
const int row_base =
(tile_idx < LOCAL_TILES_S1 / 2 ? 0 : state.d_expert_pad) +
state.fused_tile_idx * INTERMEDIATE_TILE_COLS +
(tile_idx % (LOCAL_TILES_S1 / 2)) * MMA_N;
const int rows_per_expert = 2 * state.d_expert_pad;
const int k_extent = state.d_hidden_pad;
const int k_base = k_idx * BLOCK;
const uint32_t row = state.lane_id & 15;
const uint32_t k_group = state.lane_id >> 4;
const int valid_rows_raw = rows_per_expert - row_base;
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);
uint32_t packed_scale = 0;
if (lane_valid) {
// Async DMA fp4 tile from gmem → LDS
const size_t expert_base =
size_t(state.expert_id) * size_t(rows_per_expert) * size_t(k_extent / 2);
const size_t tile_offset =
size_t(row_base) * size_t(k_extent / 2) +
size_t(k_base) * size_t(MMA_N / 2);
const size_t lane_offset = expert_base + tile_offset +
size_t(row) * 16u + size_t(k_group) * 256u;
const uint64_t b_src_addr =
reinterpret_cast<uint64_t>(state.b_shuffle + lane_offset);
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_data) -
reinterpret_cast<const char*>(&lds)));
global_load_lds_dwordx4_async(b_lds_off, b_src_addr);
// Load scale byte from gmem
const int global_row = row_base + int(row);
const uint32_t global_k_group = uint32_t(k_base / SCALE_GROUP + int(k_group));
const uint32_t flat_scale_row = uint32_t(state.expert_id * rows_per_expert + global_row);
const uint64_t bs_src =
reinterpret_cast<uint64_t>(state.b_scales) +
shuffled_scale_byte_offset(flat_scale_row, global_k_group, padded_scale_cols(k_extent));
asm volatile("global_load_ubyte %0, %1, off\\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].b_scales)[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 = parity;
}
}
}
}
}
// ============================================================
// MMA consumer: read A,B from LDS, issue MFMA
// ============================================================
__device__ __forceinline__ void mma_read_a(
const KernelState& state,
LdsState& lds,
LaneArgs& args,
int stage_id,
uint8_t parity
) {
while (lds.a_stage[state.consumer_id][stage_id].prod_to_mma != parity) {}
args.a_scale = reinterpret_cast<volatile uint8_t*>(
lds.a_stage[state.consumer_id][stage_id].a_scales)[state.lane_id];
const uint32_t a_lds_off = static_cast<uint32_t>(reinterpret_cast<uintptr_t>(
lds.a_stage[state.consumer_id][stage_id].a_data)) + uint32_t(state.lane_id) * 16u;
u32x4 a_tmp;
ds_read_b128_async(a_lds_off, a_tmp);
asm volatile("s_waitcnt lgkmcnt(0)\\n" ::: "memory");
args.a_regs.v4 = a_tmp;
}
__device__ __forceinline__ void mma_read_b(
const KernelState& state,
LdsState& lds,
LaneArgs& args,
int stage_id,
uint8_t parity,
bool release_a,
int a_stage_id,
uint8_t a_parity
) {
while (lds.b_stage[state.consumer_id][stage_id].prod_to_mma != parity) {}
args.b_scale = reinterpret_cast<volatile uint8_t*>(
lds.b_stage[state.consumer_id][stage_id].b_scales)[state.lane_id];
const uint32_t b_lds_off = static_cast<uint32_t>(reinterpret_cast<uintptr_t>(
lds.b_stage[state.consumer_id][stage_id].b_data)) + uint32_t(state.lane_id) * 16u;
u32x4 b_tmp;
ds_read_b128_async(b_lds_off, b_tmp);
asm volatile("s_waitcnt lgkmcnt(0)\\n" ::: "memory");
args.b_regs.v4 = b_tmp;
// Release B stage (and optionally A stage)
if (state.lane_id == 0) {
lds.b_stage[state.consumer_id][stage_id].mma_to_prod = parity ^ 1;
if (release_a) {
lds.a_stage[state.consumer_id][a_stage_id].mma_to_prod = a_parity ^ 1;
}
}
}
__device__ __forceinline__ void store_fragment_to_lds_f32(
float* partial_f32,
const f32x4& frag,
int lane_id,
int tile_idx
) {
const int lane_col = lane_id & (MMA_N - 1);
const int lane_row_base = (lane_id >> 4) * 4;
const size_t tile_base = size_t(tile_idx) * MMA_M * MMA_N;
#pragma unroll
for (int acc_idx = 0; acc_idx < 4; ++acc_idx) {
const int row = lane_row_base + acc_idx;
partial_f32[tile_base + size_t(row) * MMA_N + lane_col] = frag[acc_idx];
}
}
// ============================================================
// Stage 1 MMA consumer
// ============================================================
__device__ __forceinline__ void s1_mma_consumer(
const KernelState& state,
LdsState& lds
) {
LaneArgs args{};
#pragma unroll
for (int t = 0; t < LOCAL_TILES_S1; ++t) {
args.acc_local[t] = {0.0f, 0.0f, 0.0f, 0.0f};
}
for (int local_k = 0; local_k < state.s1_k_tile_count; ++local_k) {
const int a_stage_id = local_k % A_STAGES;
const uint8_t a_parity = (local_k / A_STAGES) & 1;
mma_read_a(state, lds, args, a_stage_id, a_parity);
for (int tile_idx = 0; tile_idx < LOCAL_TILES_S1; ++tile_idx) {
const int token_b = local_k * LOCAL_TILES_S1 + tile_idx;
const int b_stage_id = token_b % B_STAGES;
const uint8_t b_parity = (token_b / B_STAGES) & 1;
const bool release_a = (tile_idx == LOCAL_TILES_S1 - 1);
mma_read_b(state, lds, args, b_stage_id, b_parity,
release_a, a_stage_id, a_parity);
args.acc_local[tile_idx] =
__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
args.a_regs.v8, args.b_regs.v8, args.acc_local[tile_idx],
4, 4, 0, args.a_scale, 0, args.b_scale
);
}
}
// Store partial results to LDS
// If multiple MMA waves did k-splits, we need to reduce.
// For now store directly (single wave or last wave accumulates).
const int s1_k_tiles = ceil_div_i32(state.d_hidden_pad, BLOCK);
const int active_mma = s1_k_tiles < MMA_WAVES ? s1_k_tiles : MMA_WAVES;
if (active_mma == 1) {
// Only one MMA wave, write directly
#pragma unroll
for (int t = 0; t < LOCAL_TILES_S1; ++t) {
store_fragment_to_lds_f32(lds.partial_f32, args.acc_local[t], state.lane_id, t);
}
} else {
// K-split: accumulate via atomicAdd to partial_f32
#pragma unroll
for (int t = 0; t < LOCAL_TILES_S1; ++t) {
const int lane_col = state.lane_id & (MMA_N - 1);
const int lane_row_base = (state.lane_id >> 4) * 4;
const size_t tile_base = size_t(t) * MMA_M * MMA_N;
#pragma unroll
for (int acc_idx = 0; acc_idx < 4; ++acc_idx) {
const int row = lane_row_base + acc_idx;
atomicAdd(&lds.partial_f32[tile_base + size_t(row) * MMA_N + lane_col],
args.acc_local[t][acc_idx]);
}
}
}
}
// ============================================================
// SwiGLU activation (run by all waves cooperatively after barrier)
// ============================================================
__device__ __forceinline__ void apply_swiglu_to_lds(
LdsState& lds,
const KernelState& state,
int tid,
int total_threads
) {
constexpr int kElems = MMA_M * INTERMEDIATE_TILE_COLS;
for (int idx = tid; idx < kElems; idx += total_threads) {
const int row = idx / INTERMEDIATE_TILE_COLS;
const int col = idx % INTERMEDIATE_TILE_COLS;
if (row >= state.token_tile_count || lds.token_ids[row] == INVALID_TOKEN) continue;
const int tile_in_half = col / MMA_N;
const int col_in_tile = col % MMA_N;
const size_t gate_idx = size_t(tile_in_half) * MMA_M * MMA_N +
size_t(row) * MMA_N + col_in_tile;
const size_t up_idx = size_t(tile_in_half + LOCAL_TILES_S1 / 2) * MMA_M * MMA_N +
size_t(row) * MMA_N + col_in_tile;
const float gate = lds.partial_f32[gate_idx];
const float up = lds.partial_f32[up_idx];
lds.intermediate_bf16[gate_idx] = f32_to_bf16_bits(silu_f32(gate) * up);
}
}
// ============================================================
// Stage 2: A-producer — quantize intermediate from LDS → store back to A stage
// ============================================================
__device__ __forceinline__ void s2_a_producer(
const KernelState& state,
LdsState& lds
) {
// Stage 2 has exactly 1 k-tile (INTERMEDIATE_TILE_COLS == BLOCK),
// so each consumer needs exactly 1 A tile.
for (int ci = 0; ci < state.prod_consumer_count; ++ci) {
const int consumer_id = state.prod_consumer_start + ci;
const int stage_id = 0;
const uint8_t parity = 0;
while (lds.a_stage[consumer_id][stage_id].mma_to_prod != parity) {}
// Read intermediate bf16 from LDS, quantize, store to A stage
const uint32_t row = state.lane_id & 15;
const uint32_t k_group = state.lane_id >> 4;
const bool lane_valid = row < uint32_t(state.token_tile_count)
&& lds.token_ids[row] != INVALID_TOKEN;
u32x4 src_chunks[4] = {};
if (lane_valid) {
const uint32_t base_col = k_group * SCALE_GROUP;
if (base_col < uint32_t(INTERMEDIATE_TILE_COLS)) {
uint32_t bf16_pairs[16];
#pragma unroll
for (int w = 0; w < 16; ++w) {
const uint32_t col0 = base_col + uint32_t(w * 2);
const uint32_t col1 = base_col + uint32_t(w * 2 + 1);
uint32_t pair_word = 0u;
if (col0 < uint32_t(INTERMEDIATE_TILE_COLS)) {
const int t0 = col0 / MMA_N;
const int c0 = col0 % MMA_N;
pair_word |= uint32_t(lds.intermediate_bf16[
size_t(t0) * MMA_M * MMA_N + size_t(row) * MMA_N + c0]);
}
if (col1 < uint32_t(INTERMEDIATE_TILE_COLS)) {
const int t1 = col1 / MMA_N;
const int c1 = col1 % MMA_N;
pair_word |= uint32_t(lds.intermediate_bf16[
size_t(t1) * MMA_M * MMA_N + size_t(row) * MMA_N + c1]) << 16;
}
bf16_pairs[w] = pair_word;
}
const u32x4* chunks = reinterpret_cast<const u32x4*>(bf16_pairs);
src_chunks[0] = chunks[0];
src_chunks[1] = chunks[1];
src_chunks[2] = chunks[2];
src_chunks[3] = chunks[3];
}
}
a_prod_store_to_lds(lds.a_stage[consumer_id][stage_id], state, src_chunks);
if (state.lane_id == 0) {
lds.a_stage[consumer_id][stage_id].prod_to_mma = parity;
}
}
}
// ============================================================
// Stage 2: B-producer — async DMA down_weight fp4 from gmem → LDS
// ============================================================
__device__ __forceinline__ void s2_b_producer(
const KernelState& state,
LdsState& lds
) {
for (int ci = 0; ci < state.prod_consumer_count; ++ci) {
const int consumer_id = state.prod_consumer_start + ci;
// Get this consumer's N-tile range
const int s2_n_tiles = ceil_div_i32(state.d_hidden_pad, MMA_N);
const int s1_k_tiles = ceil_div_i32(state.d_hidden_pad, BLOCK);
const int active_mma = s1_k_tiles < MMA_WAVES ? s1_k_tiles : MMA_WAVES;
int cons_n_start, cons_n_count;
split_range(s2_n_tiles, active_mma, consumer_id, cons_n_start, cons_n_count);
for (int local_n = 0; local_n < cons_n_count; ++local_n) {
const int stage_id = local_n % B_STAGES;
const uint8_t parity = (local_n / B_STAGES) & 1;
while (lds.b_stage[consumer_id][stage_id].mma_to_prod != parity) {}
const int n_tile_idx = cons_n_start + local_n;
const int row_base = n_tile_idx * MMA_N;
const int rows_per_expert = state.d_hidden_pad;
const int k_extent = state.d_expert_pad;
const int k_base = state.k_offset;
const uint32_t row = state.lane_id & 15;
const uint32_t k_group = state.lane_id >> 4;
const int valid_rows_raw = rows_per_expert - row_base;
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);
uint32_t packed_scale = 0;
if (lane_valid) {
const size_t expert_base =
size_t(state.expert_id) * size_t(rows_per_expert) * size_t(k_extent / 2);
const size_t tile_offset =
size_t(row_base) * size_t(k_extent / 2) +
size_t(k_base) * size_t(MMA_N / 2);
const size_t lane_offset = expert_base + tile_offset +
size_t(row) * 16u + size_t(k_group) * 256u;
const uint64_t b_src_addr =
reinterpret_cast<uint64_t>(state.down_shuffle + lane_offset);
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_data) -
reinterpret_cast<const char*>(&lds)));
global_load_lds_dwordx4_async(b_lds_off, b_src_addr);
const int global_row = row_base + int(row);
const uint32_t global_k_group = uint32_t(k_base / SCALE_GROUP + int(k_group));
const uint32_t flat_scale_row = uint32_t(state.expert_id * rows_per_expert + global_row);
const uint64_t bs_src =
reinterpret_cast<uint64_t>(state.down_scales) +
shuffled_scale_byte_offset(flat_scale_row, global_k_group, padded_scale_cols(k_extent));
asm volatile("global_load_ubyte %0, %1, off\\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].b_scales)[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 = parity;
}
}
}
}
// ============================================================
// Stage 2 MMA consumer — N-split across MMA waves, no partial sums needed
// ============================================================
__device__ __forceinline__ void s2_mma_consumer(
const KernelState& state,
LdsState& lds
) {
LaneArgs args{};
// Read A tile (intermediate quantized) — only 1 k-tile
const int a_stage_id = 0;
const uint8_t a_parity = 0;
mma_read_a(state, lds, args, a_stage_id, a_parity);
for (int local_n = 0; local_n < state.s2_n_tile_count; ++local_n) {
const int b_stage_id = local_n % B_STAGES;
const uint8_t b_parity = (local_n / B_STAGES) & 1;
// Release A on last iteration only
const bool release_a = (local_n == state.s2_n_tile_count - 1);
mma_read_b(state, lds, args, b_stage_id, b_parity,
release_a, a_stage_id, a_parity);
f32x4 acc = {0.0f, 0.0f, 0.0f, 0.0f};
acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
args.a_regs.v8, args.b_regs.v8, acc,
4, 4, 0, args.a_scale, 0, args.b_scale
);
// Write results: multiply by topk_weight and atomicAdd to final_out
const int n_tile_idx = state.s2_n_tile_start + local_n;
const int lane_col = state.lane_id & (MMA_N - 1);
const int lane_row_base = (state.lane_id >> 4) * 4;
const int global_col = n_tile_idx * MMA_N + lane_col;
if (global_col < state.d_hidden) {
#pragma unroll
for (int acc_idx = 0; acc_idx < 4; ++acc_idx) {
const int local_row = lane_row_base + acc_idx;
if (local_row < state.token_tile_count &&
lds.token_ids[local_row] != INVALID_TOKEN) {
const uint32_t token_id = lds.token_ids[local_row];
const float weight = lds.topk_weights[local_row];
const float val = acc[acc_idx] * weight;
float* dst = state.final_out +
size_t(token_id) * size_t(state.final_out_row_stride_elems) +
size_t(global_col);
atomicAdd(dst, val);
}
}
}
}
}
// ============================================================
// Dispatch initialization (same as proto.py)
// ============================================================
__device__ __forceinline__ void init_kernel_dispatch(
KernelState& state,
LdsState& lds,
const int32_t* cta_to_expert,
const int32_t* cta_to_local_tile,
const int32_t* expert_batch_offsets,
const DispatchEntry* dispatch_buffer,
const __hip_bfloat16* hidden_states,
const uint8_t* b_shuffle,
const uint8_t* b_scales,
const uint8_t* down_shuffle,
const uint8_t* down_scales,
float* final_out,
int d_hidden,
int d_hidden_pad,
int d_expert,
int d_expert_pad,
int hidden_row_stride_elems,
int final_out_row_stride_elems
) {
const int tid = threadIdx.x;
state.lane_id = tid & (WAVE_SIZE - 1);
state.wave_id = tid / WAVE_SIZE;
state.expert_id = cta_to_expert[blockIdx.x];
state.fused_tile_idx = blockIdx.y;
state.d_hidden = d_hidden;
state.d_hidden_pad = d_hidden_pad;
state.d_expert = d_expert;
state.d_expert_pad = d_expert_pad;
state.hidden_row_stride_elems = hidden_row_stride_elems;
state.hidden_states = hidden_states;
state.b_shuffle = b_shuffle;
state.b_scales = b_scales;
state.down_shuffle = down_shuffle;
state.down_scales = down_scales;
state.final_out = final_out;
state.final_out_row_stride_elems = final_out_row_stride_elems;
state.k_offset = state.fused_tile_idx * INTERMEDIATE_TILE_COLS;
int local_tile = cta_to_local_tile[blockIdx.x];
int batch_idx = expert_batch_offsets[state.expert_id] + local_tile;
if (tid < CTA_M) {
const DispatchEntry& entry = dispatch_buffer[batch_idx * CTA_M + tid];
lds.token_ids[tid] = static_cast<uint32_t>(entry.token_id);
lds.topk_weights[tid] = entry.weight;
}
__syncthreads();
state.token_tile_count = 0;
for (int i = 0; i < CTA_M; ++i) {
if (lds.token_ids[i] != INVALID_TOKEN) state.token_tile_count++;
}
state.output_row_base = batch_idx * CTA_M;
state.valid = state.token_tile_count > 0;
}
// ============================================================
// Main kernel
// ============================================================
__global__ void fused_moe_kernel(
const __hip_bfloat16* hidden_states,
const uint8_t* gate_up_weight_shuffled,
const uint8_t* gate_up_weight_scale_shuffled,
const uint8_t* down_weight_shuffled,
const uint8_t* down_weight_scale_shuffled,
const DispatchEntry* dispatch_buffer,
const int32_t* cta_to_expert,
const int32_t* cta_to_local_tile,
const int32_t* expert_batch_offsets,
float* final_out,
int d_hidden,
int d_hidden_pad,
int d_expert,
int d_expert_pad,
int hidden_row_stride_elems,
int final_out_row_stride_elems
) {
__shared__ LdsState lds;
KernelState state;
init_kernel_dispatch(
state, lds,
cta_to_expert, cta_to_local_tile, expert_batch_offsets,
dispatch_buffer,
hidden_states, gate_up_weight_shuffled, gate_up_weight_scale_shuffled,
down_weight_shuffled, down_weight_scale_shuffled,
final_out,
d_hidden, d_hidden_pad, d_expert, d_expert_pad,
hidden_row_stride_elems, final_out_row_stride_elems
);
if (!state.valid) { return; }
const int s1_k_tiles = ceil_div_i32(d_hidden_pad, BLOCK);
const int s2_n_tiles = ceil_div_i32(d_hidden_pad, MMA_N);
init_wave_roles(state, s1_k_tiles, s2_n_tiles);
init_pipeline(lds);
// Need to zero partial_f32 if k-splitting
const int active_mma = s1_k_tiles < MMA_WAVES ? s1_k_tiles : MMA_WAVES;
if (active_mma > 1) {
const int tid = threadIdx.x;
const int total_threads = TOTAL_WAVES * WAVE_SIZE;
for (int i = tid; i < LOCAL_TILES_S1 * MMA_M * MMA_N; i += total_threads) {
lds.partial_f32[i] = 0.0f;
}
__syncthreads();
}
// ======== STAGE 1 ========
if (state.role == 0) { s1_a_producer(state, lds); }
if (state.role == 1) { s1_b_producer(state, lds); }
if (state.role == 2) { s1_mma_consumer(state, lds); }
__syncthreads();
// ======== SwiGLU ========
{
const int tid = threadIdx.x;
const int total_threads = TOTAL_WAVES * WAVE_SIZE;
apply_swiglu_to_lds(lds, state, tid, total_threads);
}
__syncthreads();
// Re-init pipeline for stage 2
init_pipeline(lds);
// ======== STAGE 2 ========
if (state.role == 0) { s2_a_producer(state, lds); }
if (state.role == 1) { s2_b_producer(state, lds); }
if (state.role == 2) { s2_mma_consumer(state, lds); }
}
// ============================================================
// Host dispatch (unchanged from proto.py)
// ============================================================
struct DispatchInfo {
torch::Tensor dispatch_buffer;
torch::Tensor cta_to_expert;
torch::Tensor cta_to_local_tile;
torch::Tensor expert_batch_offsets;
int total_ctas_m;
};
DispatchInfo build_dispatch(
const torch::Tensor& topk_ids,
const torch::Tensor& topk_weights,
int num_experts,
int total_top_k,
torch::Device device
) {
const int M = topk_ids.size(0);
auto ids_cpu = topk_ids.cpu().contiguous();
auto weights_cpu = topk_weights.cpu().contiguous();
const int32_t* ids_ptr = ids_cpu.data_ptr<int32_t>();
const float* weights_ptr = weights_cpu.data_ptr<float>();
std::vector<std::vector<std::pair<int32_t, float>>> per_expert(num_experts);
for (int tok = 0; tok < M; ++tok) {
for (int slot = 0; slot < total_top_k; ++slot) {
int eid = ids_ptr[tok * total_top_k + slot];
float w = weights_ptr[tok * total_top_k + slot];
if (eid >= 0 && eid < num_experts) {
per_expert[eid].emplace_back(tok, w);
}
}
}
std::vector<int32_t> expert_offsets(num_experts + 1, 0);
std::vector<int32_t> cta_expert_vec;
std::vector<int32_t> cta_tile_vec;
for (int e = 0; e < num_experts; ++e) {
int n_tokens = static_cast<int>(per_expert[e].size());
int n_batches = (n_tokens + CTA_M - 1) / CTA_M;
if (n_batches == 0) n_batches = 0;
expert_offsets[e + 1] = expert_offsets[e] + n_batches;
for (int t = 0; t < n_batches; ++t) {
cta_expert_vec.push_back(e);
cta_tile_vec.push_back(t);
}
}
int total_batches = expert_offsets[num_experts];
int total_ctas_m = static_cast<int>(cta_expert_vec.size());
auto dispatch_cpu = torch::full({total_batches * CTA_M, 2}, 0, torch::kInt32);
int32_t* dispatch_ptr = dispatch_cpu.data_ptr<int32_t>();
const int32_t INVALID = static_cast<int32_t>(0xFFFFFFFF);
for (int e = 0; e < num_experts; ++e) {
int base = expert_offsets[e] * CTA_M;
const auto& tokens = per_expert[e];
int n_batches = expert_offsets[e + 1] - expert_offsets[e];
for (int b = 0; b < n_batches; ++b) {
for (int s = 0; s < CTA_M; ++s) {
int idx = b * CTA_M + s;
int flat = (base + idx) * 2;
if (idx < static_cast<int>(tokens.size())) {
dispatch_ptr[flat + 0] = tokens[idx].first;
float w = tokens[idx].second;
int32_t w_bits;
std::memcpy(&w_bits, &w, sizeof(float));
dispatch_ptr[flat + 1] = w_bits;
} else {
dispatch_ptr[flat + 0] = INVALID;
dispatch_ptr[flat + 1] = 0;
}
}
}
}
auto offsets_cpu = torch::from_blob(expert_offsets.data(), {num_experts + 1}, torch::kInt32).clone();
auto cta_expert_cpu = torch::from_blob(cta_expert_vec.data(), {total_ctas_m}, torch::kInt32).clone();
auto cta_tile_cpu = torch::from_blob(cta_tile_vec.data(), {total_ctas_m}, torch::kInt32).clone();
DispatchInfo info;
info.dispatch_buffer = dispatch_cpu.to(device);
info.cta_to_expert = cta_expert_cpu.to(device);
info.cta_to_local_tile = cta_tile_cpu.to(device);
info.expert_batch_offsets = offsets_cpu.to(device);
info.total_ctas_m = total_ctas_m;
return info;
}
torch::Tensor fused_moe(
torch::Tensor hidden_states,
torch::Tensor gate_up_weight,
torch::Tensor down_weight,
torch::Tensor gate_up_weight_scale,
torch::Tensor down_weight_scale,
torch::Tensor gate_up_weight_shuffled,
torch::Tensor down_weight_shuffled,
torch::Tensor gate_up_weight_scale_shuffled,
torch::Tensor down_weight_scale_shuffled,
torch::Tensor topk_weights,
torch::Tensor topk_ids,
int d_hidden,
int d_expert,
int d_hidden_pad,
int d_expert_pad,
int n_routed_experts,
int n_shared_experts,
int n_experts_per_token,
int total_top_k
) {
const int batch_tokens = hidden_states.size(0);
const int num_experts = gate_up_weight_shuffled.size(0);
const int fused_groups_per_expert = (2 * d_expert_pad) / FUSED_TILE_COLS;
auto dispatch = build_dispatch(
topk_ids, topk_weights,
num_experts, total_top_k,
hidden_states.device()
);
auto final_out_fp32 = torch::zeros(
{batch_tokens, d_hidden},
torch::TensorOptions().device(hidden_states.device()).dtype(torch::kFloat32)
);
if (dispatch.total_ctas_m == 0) {
return final_out_fp32.to(torch::kBFloat16);
}
const dim3 grid(dispatch.total_ctas_m, fused_groups_per_expert);
const dim3 block(TOTAL_WAVES * WAVE_SIZE);
fused_moe_kernel<<<grid, block, 0, 0>>>(
reinterpret_cast<const __hip_bfloat16*>(hidden_states.data_ptr()),
reinterpret_cast<const uint8_t*>(gate_up_weight_shuffled.data_ptr()),
reinterpret_cast<const uint8_t*>(gate_up_weight_scale_shuffled.data_ptr()),
reinterpret_cast<const uint8_t*>(down_weight_shuffled.data_ptr()),
reinterpret_cast<const uint8_t*>(down_weight_scale_shuffled.data_ptr()),
reinterpret_cast<const DispatchEntry*>(dispatch.dispatch_buffer.data_ptr()),
dispatch.cta_to_expert.data_ptr<int32_t>(),
dispatch.cta_to_local_tile.data_ptr<int32_t>(),
dispatch.expert_batch_offsets.data_ptr<int32_t>(),
final_out_fp32.data_ptr<float>(),
d_hidden,
d_hidden_pad,
d_expert,
d_expert_pad,
hidden_states.stride(0),
final_out_fp32.stride(0)
);
return final_out_fp32.to(torch::kBFloat16);
}
"""
module = load_inline(
name="moe_proto2_prodcons",
cpp_sources=[CPP_WRAPPER],
cuda_sources=[cuda_src],
functions=["fused_moe"],
verbose=False,
extra_cuda_cflags=["--offload-arch=gfx950", "-std=c++20"],
)
def custom_kernel(data: input_t) -> output_t:
(
hidden_states,
gate_up_weight,
down_weight,
gate_up_weight_scale,
down_weight_scale,
gate_up_weight_shuffled,
down_weight_shuffled,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
config,
) = data
return module.fused_moe(
hidden_states,
gate_up_weight,
down_weight,
gate_up_weight_scale,
down_weight_scale,
gate_up_weight_shuffled,
down_weight_shuffled,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
config["d_hidden"],
config["d_expert"],
config["d_hidden_pad"],
config["d_expert_pad"],
config["n_routed_experts"],
config["n_shared_experts"],
config["n_experts_per_token"],
config["total_top_k"],
)
scrolls · 1305 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