Skip to content
KernelIndex
Search⌘K

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
AMD MXFP4 MoEsuite of 7 cases
AMD Instinct MI355X
1.12ms
#779 of 782
2026-03-19

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 = 3constexpr int A_STAGES = 3;
warp-specializationconstexpr 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