Skip to content
KernelIndex
Search⌘K

submission 652012

DNAK-dnak · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 403 lines, June 9 Researcher Reciprocity License v1.0.

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-652012?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
180.5µs
#486 of 782
2026-03-28

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:a45032e6e8cfb1252662995debef1b30729fa0ad900b9b45a8af98be0a09db05
license declaredunknown
license concludedunknown
authorsDNAK-dnak
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

fp4MXFP4 MoE Submission for AMD MI355X

Kernel source

submission.py403 lines
"""
MXFP4 MoE Submission for AMD MI355X
GPU MODE Leaderboard 764: amd-moe-mxfp4

Strategy: Start with AITER baseline + incremental optimizations.
The key insight is that AITER's fused_moe is already very good —
we need to find the margins in sorting, dispatch, and shape-specific tuning.
"""

#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu AMD

from task import input_t, output_t
import torch
import torch.nn.functional as F

# ──────────────────────────────────────────────────────────────────────
# Approach 1: AITER baseline (identical to reference — sanity check)
# Submit this first to verify your setup works.
# ──────────────────────────────────────────────────────────────────────

import aiter
from aiter import ActivationType, QuantType
from aiter.fused_moe import fused_moe


def custom_kernel_v0_baseline(data: input_t) -> output_t:
    """Exact copy of reference — should match baseline times."""
    (
        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

    hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
    intermediate_pad = config["d_expert_pad"] - config["d_expert"]

    output = fused_moe(
        hidden_states,
        gate_up_weight_shuffled,
        down_weight_shuffled,
        topk_weights,
        topk_ids,
        expert_mask=None,
        activation=ActivationType.Silu,
        quant_type=QuantType.per_1x32,
        doweight_stage1=False,
        w1_scale=gate_up_weight_scale_shuffled,
        w2_scale=down_weight_scale_shuffled,
        a1_scale=None,
        a2_scale=None,
        hidden_pad=hidden_pad,
        intermediate_pad=intermediate_pad,
    )
    return output


# ──────────────────────────────────────────────────────────────────────
# Approach 2: Explore AITER's doweight_stage1 flag
# When doweight_stage1=True, routing weights are applied to inputs
# BEFORE Stage 1 GEMM rather than AFTER Stage 2. This changes the
# arithmetic slightly but may allow better fusion / less work in the
# reduction step. The 5% tolerance should absorb the numerical diff.
# ──────────────────────────────────────────────────────────────────────

def custom_kernel_v1_doweight(data: input_t) -> output_t:
    """Try applying router weights on input side (doweight_stage1=True)."""
    (
        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

    hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
    intermediate_pad = config["d_expert_pad"] - config["d_expert"]

    output = fused_moe(
        hidden_states,
        gate_up_weight_shuffled,
        down_weight_shuffled,
        topk_weights,
        topk_ids,
        expert_mask=None,
        activation=ActivationType.Silu,
        quant_type=QuantType.per_1x32,
        doweight_stage1=True,  # <-- Key change: weight on input
        w1_scale=gate_up_weight_scale_shuffled,
        w2_scale=down_weight_scale_shuffled,
        a1_scale=None,
        a2_scale=None,
        hidden_pad=hidden_pad,
        intermediate_pad=intermediate_pad,
    )
    return output


# ──────────────────────────────────────────────────────────────────────
# Approach 3: Pre-compute expert mask for sparse dispatch
# For bs=16 with E=257, most experts see 0 tokens. Building an
# expert_mask tells AITER to skip empty experts entirely.
# ──────────────────────────────────────────────────────────────────────

def _build_expert_mask(topk_ids: torch.Tensor, num_experts: int) -> torch.Tensor:
    """
    Build a boolean mask [E] indicating which experts have at least one token.
    This can help AITER skip launching work for empty experts.
    """
    flat_ids = topk_ids.view(-1)
    mask = torch.zeros(num_experts, dtype=torch.bool, device=topk_ids.device)
    mask.scatter_(0, flat_ids.long(), True)
    return mask


def custom_kernel_v2_expert_mask(data: input_t) -> output_t:
    """Use expert_mask to skip empty experts."""
    (
        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

    hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
    intermediate_pad = config["d_expert_pad"] - config["d_expert"]
    E = config["n_routed_experts"] + config["n_shared_experts"]

    expert_mask = _build_expert_mask(topk_ids, E)

    output = fused_moe(
        hidden_states,
        gate_up_weight_shuffled,
        down_weight_shuffled,
        topk_weights,
        topk_ids,
        expert_mask=expert_mask,
        activation=ActivationType.Silu,
        quant_type=QuantType.per_1x32,
        doweight_stage1=False,
        w1_scale=gate_up_weight_scale_shuffled,
        w2_scale=down_weight_scale_shuffled,
        a1_scale=None,
        a2_scale=None,
        hidden_pad=hidden_pad,
        intermediate_pad=intermediate_pad,
    )
    return output


# ──────────────────────────────────────────────────────────────────────
# Approach 4: Custom sorting + AITER lower-level API
# Instead of letting fused_moe handle sorting internally, we do it
# ourselves with optimized radix sort, then call the grouped GEMM
# stages directly.
# ──────────────────────────────────────────────────────────────────────

def custom_kernel_v3_custom_sort(data: input_t) -> output_t:
    """
    Custom token-expert sorting + direct AITER grouped GEMM.
    
    The idea: fused_moe's internal sorting may not be optimal for all
    shapes. We pre-sort with a counting sort (O(M*topk + E)) which is
    faster than the generic sort for small E or small M.
    """
    (
        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

    hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
    intermediate_pad = config["d_expert_pad"] - config["d_expert"]

    # Try using AITER's moe_sorting if available for better performance
    try:
        from aiter.fused_moe import moe_sorting, get_block_size_M
        
        M = hidden_states.shape[0]
        top_k = topk_ids.shape[1]
        E = config["n_routed_experts"] + config["n_shared_experts"]
        
        block_size_M = get_block_size_M(M)
        
        sorted_token_ids, sorted_weights, sorted_expert_ids, num_valid_ids = moe_sorting(
            topk_ids, topk_weights, E, block_size_M, 
        )
        
        # Now call fused_moe with pre-sorted data
        # (This depends on whether fused_moe accepts pre-sorted inputs —
        #  if not, fall back to the standard call)
    except (ImportError, AttributeError, TypeError):
        pass
    
    # Fallback: standard fused_moe call
    output = fused_moe(
        hidden_states,
        gate_up_weight_shuffled,
        down_weight_shuffled,
        topk_weights,
        topk_ids,
        expert_mask=None,
        activation=ActivationType.Silu,
        quant_type=QuantType.per_1x32,
        doweight_stage1=False,
        w1_scale=gate_up_weight_scale_shuffled,
        w2_scale=down_weight_scale_shuffled,
        a1_scale=None,
        a2_scale=None,
        hidden_pad=hidden_pad,
        intermediate_pad=intermediate_pad,
    )
    return output


# ──────────────────────────────────────────────────────────────────────
# Approach 5: Shape-specific dispatch
# Different shapes need different strategies. This dispatcher
# selects the best approach per benchmark configuration.
# ──────────────────────────────────────────────────────────────────────

def custom_kernel_v4_dispatch(data: input_t) -> output_t:
    """
    Shape-aware dispatcher.
    
    Key insight: The 7 benchmark shapes fall into 3 categories:
    1. Small batch (bs=16): Latency-bound, skip overhead
    2. Medium batch (bs=128): Balanced
    3. Large batch (bs=512): Throughput-bound, sorting matters
    
    And 2 expert-count regimes:
    A. Many experts (E=257): Most experts idle for small batches
    B. Few experts (E=33): All experts see many tokens
    """
    (
        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

    M = config["bs"]
    E = config["n_routed_experts"] + config["n_shared_experts"]
    d_expert = config["d_expert"]
    
    hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
    intermediate_pad = config["d_expert_pad"] - config["d_expert"]

    # For small batches with many experts: use expert_mask to skip empties
    if M <= 16 and E > 100:
        expert_mask = _build_expert_mask(topk_ids, E)
        output = fused_moe(
            hidden_states,
            gate_up_weight_shuffled,
            down_weight_shuffled,
            topk_weights,
            topk_ids,
            expert_mask=expert_mask,
            activation=ActivationType.Silu,
            quant_type=QuantType.per_1x32,
            doweight_stage1=False,
            w1_scale=gate_up_weight_scale_shuffled,
            w2_scale=down_weight_scale_shuffled,
            a1_scale=None,
            a2_scale=None,
            hidden_pad=hidden_pad,
            intermediate_pad=intermediate_pad,
        )
    else:
        # Default: standard AITER fused_moe
        output = fused_moe(
            hidden_states,
            gate_up_weight_shuffled,
            down_weight_shuffled,
            topk_weights,
            topk_ids,
            expert_mask=None,
            activation=ActivationType.Silu,
            quant_type=QuantType.per_1x32,
            doweight_stage1=False,
            w1_scale=gate_up_weight_scale_shuffled,
            w2_scale=down_weight_scale_shuffled,
            a1_scale=None,
            a2_scale=None,
            hidden_pad=hidden_pad,
            intermediate_pad=intermediate_pad,
        )

    return output


# ──────────────────────────────────────────────────────────────────────
# Approach 6 (ADVANCED): Triton MoE kernel for gfx950
# This is the nuclear option — write a custom Triton kernel that
# directly uses FP4 dot products on MI355X.
#
# NOTE: This is a SKELETON. You'd need to fill in the actual
# Triton kernel body, which requires testing on real MI355X hardware.
# The key challenge is getting the tl.dot to use MFMA FP4 instructions.
# ──────────────────────────────────────────────────────────────────────

"""
# UNCOMMENT AND DEVELOP ON ACTUAL MI355X HARDWARE

import triton
import triton.language as tl

@triton.jit
def moe_stage1_kernel(
    # Pointers
    hidden_ptr,          # [M, d_hidden] bf16
    gate_up_w_ptr,       # [E, 2*d_expert_pad, d_hidden_pad//2] fp4x2
    gate_up_scale_ptr,   # scales
    intermediate_ptr,    # [M*topk, d_expert] bf16 output
    sorted_token_ids_ptr,
    sorted_expert_ids_ptr,
    # Dimensions
    d_hidden: tl.constexpr,
    d_expert: tl.constexpr,
    d_hidden_pad: tl.constexpr,
    d_expert_pad: tl.constexpr,
    # Tile sizes
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
    # XCD config
    NUM_XCDS: tl.constexpr,
):
    # Program ID with XCD-aware remapping
    pid = tl.program_id(0)
    num_pid_m = tl.cdiv(TOKENS_PER_EXPERT, BLOCK_M)  # varies per expert
    num_pid_n = tl.cdiv(2 * d_expert_pad, BLOCK_N)
    
    # XCD-aware PID remapping for MI355X
    GRID_MN = num_pid_m * num_pid_n
    pids_per_xcd = (GRID_MN + NUM_XCDS - 1) // NUM_XCDS
    xcd_id = pid % NUM_XCDS
    local_pid = pid // NUM_XCDS
    remapped_pid = xcd_id * pids_per_xcd + local_pid
    
    pid_m = remapped_pid // num_pid_n
    pid_n = remapped_pid % num_pid_n
    
    # Load token indices for this expert group
    # ... (sorted by expert)
    
    # Accumulator
    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
    
    # K-loop: iterate over d_hidden in BLOCK_K chunks
    for k in range(0, d_hidden_pad, BLOCK_K):
        # Load activation block [BLOCK_M, BLOCK_K] in bf16
        # Quantize to FP4 on-the-fly (per-1x32 block scaling)
        # Load weight block [BLOCK_N, BLOCK_K] in fp4x2
        # Load scale block
        # acc += tl.dot(a_fp4, w_fp4, scale_a, scale_w)
        pass
    
    # Apply SwiGLU: split acc into gate and up, compute SiLU(gate) * up
    # gate_acc = acc[:, :d_expert]
    # up_acc = acc[:, d_expert:]
    # intermediate = tl.sigmoid(gate_acc) * gate_acc * up_acc  # SiLU = x * sigmoid(x)
    
    # Store intermediate result
    # ...
"""


# ──────────────────────────────────────────────────────────────────────
# ACTIVE SUBMISSION: Pick your best approach
# Start with v0 (baseline), verify it passes, then iterate.
# ──────────────────────────────────────────────────────────────────────

def custom_kernel(data: input_t) -> output_t:
    """
    Main entry point — called by the benchmark harness.
    Switch between approaches as you test on the leaderboard.
    """
    # Start here: verify baseline works
    return custom_kernel_v0_baseline(data)
    
    # Then try these one at a time:
    # return custom_kernel_v1_doweight(data)
    # return custom_kernel_v2_expert_mask(data)
    # return custom_kernel_v3_custom_sort(data)
    # return custom_kernel_v4_dispatch(data)
scrolls · 403 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