Skip to content
KernelIndex
Search⌘K

submission 597664

Andrew Choi · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:978bbb7ef9003a47c23f9f52a2a3004f05cf5ea3bc51950d59ae2afb998df1e7
license declaredunknown
license concludedunknown
authorsAndrew Choi
imported2026-08-15

Techniques

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

fp4(7168, 512, 33, 9, 32, 128, 256, True, 'fp4', 'fp4', 'bf16'),
persistent-kernel3. HSA_NO_SCRATCH_RECLAIM=1 (persistent GPU scratch)
split-kkwargs['split_k'] = 1

Kernel source

submission.py224 lines
#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X

import os
import sys

os.environ["AITER_USE_OPUS_MOE_SORTING"] = "1"
os.environ["HIP_FORCE_DEV_KERNARG"] = "1"
os.environ["HSA_NO_SCRATCH_RECLAIM"] = "1"
os.environ["HSA_NO_SCRATCH_THREAD_LIMITER"] = "1"

import torch
import functools
from typing import Dict
from task import input_t, output_t

from aiter import ActivationType, QuantType, dtypes
import aiter
import aiter.fused_moe as _aiter_fused_moe_module

# Import flydsl
from aiter.ops.flydsl import flydsl_moe_stage1, flydsl_moe_stage2

# ---- Pre-compile flydsl stage2 kernels for all benchmark shapes ----
# This moves the FlyDSL MLIR compilation out of the first benchmark call,
# similar to how AITER pre-builds CK modules at import time.
try:
    from aiter.ops.flydsl import moe_kernels as _moe_kernels
    if hasattr(_moe_kernels, 'compile_flydsl_moe_stage2'):
        _precompile_configs = [
            # E=33/bs=512/d=512: tile_m=32 (forced override from heuristic block_m=64)
            (7168, 512, 33, 9, 32, 128, 256, True, 'fp4', 'fp4', 'bf16'),
            # E=33/bs=512/d=2048: tile_m=32 (forced override from heuristic block_m=128)
            (7168, 2048, 33, 9, 32, 128, 256, True, 'fp4', 'fp4', 'bf16'),
            # E=257/d=256: tile_m=32 (same config as CK CSV block_m=32)
            (7168, 256, 257, 9, 32, 128, 256, True, 'fp4', 'fp4', 'bf16'),
        ]
        for cfg in _precompile_configs:
            try:
                _moe_kernels.compile_flydsl_moe_stage2(*cfg)
            except Exception:
                pass  # silently skip if compilation fails for some shape
except Exception:
    pass  # flydsl not available or compilation error

# ---- Shape-dependent ksplit monkey-patch (for cktile path on small-batch E=33) ----
_original_get_ksplit = _aiter_fused_moe_module.get_ksplit

def _custom_get_ksplit(token, topk, expert, inter_dim, model_dim):
    if expert <= 64 and token <= 128:
        return 2
    return _original_get_ksplit(token, topk, expert, inter_dim, model_dim)

_aiter_fused_moe_module.get_ksplit = _custom_get_ksplit

# ---- cktile_moe_stage1 split_k=1 monkey-patch ----
if hasattr(_aiter_fused_moe_module, 'cktile_moe_stage1'):
    _original_cktile_moe_stage1 = _aiter_fused_moe_module.cktile_moe_stage1

    def _patched_cktile_moe_stage1(*args, **kwargs):
        kwargs['split_k'] = 1
        return _original_cktile_moe_stage1(*args, **kwargs)

    _aiter_fused_moe_module.cktile_moe_stage1 = _patched_cktile_moe_stage1

# ---- Replace ck_moe_stage2_fwd with flydsl_moe_stage2 for fp4x2 shapes ----
# FlyDSL generates optimized kernels for the down-projection GEMM.
# Available kernels: tile_m={32,64}, tile_n={128,256}, tile_k=256, mode={atomic,reduce}
# We use tile_m=32, tile_n=128, tile_k=256, mode=atomic for all shapes.

def _flydsl_stage2_replacement(
    inter_states,    # a2: (token_num, topk, inter_dim) -- fp4x2 quantized
    w1,              # gate_up_weight (unused in stage2)
    w2,              # down_weight shuffled
    sorted_token_ids,
    sorted_expert_ids,
    num_valid_ids,
    out,             # moe_out: (M, model_dim) bf16
    topk,
    w2_scale=None,
    a2_scale=None,
    block_m=32,
    _tile_m_override=None,
    sorted_weights=None,
    **kwargs,
):
    """Wrapper to call flydsl_moe_stage2 with ck_moe_stage2_fwd's call signature."""
    # Use override tile_m if provided, otherwise use block_m (from CK heuristic)
    tile_m = _tile_m_override if _tile_m_override is not None else block_m
    # Clamp to flydsl's supported values {32, 64}
    if tile_m > 64:
        tile_m = 64
    elif tile_m < 32:
        tile_m = 32
    flydsl_moe_stage2(
        inter_states=inter_states,
        w2=w2,
        sorted_token_ids=sorted_token_ids,
        sorted_expert_ids=sorted_expert_ids,
        num_valid_ids=num_valid_ids,
        out=out,
        topk=topk,
        tile_m=tile_m,
        tile_n=128,   # default
        tile_k=256,   # default
        a_dtype='fp4',
        b_dtype='fp4',
        out_dtype='bf16',
        mode='atomic',
        w2_scale=w2_scale,
        a2_scale=a2_scale,
        sorted_weights=sorted_weights,
    )

# Monkey-patch get_2stage_cfgs to replace stage2 with flydsl for standard 2-stage path
_original_get_2stage_cfgs = _aiter_fused_moe_module.get_2stage_cfgs.__wrapped__ \
    if hasattr(_aiter_fused_moe_module.get_2stage_cfgs, '__wrapped__') \
    else _aiter_fused_moe_module.get_2stage_cfgs

@functools.lru_cache(maxsize=2048)
def _custom_get_2stage_cfgs(
    token, model_dim, inter_dim, expert, topk,
    dtype, q_dtype_a, q_dtype_w, q_type, use_g1u1,
    activation, doweight_stage1, hidden_pad, intermediate_pad, is_shuffled=True,
):
    # Get original metadata
    metadata = _original_get_2stage_cfgs(
        token, model_dim, inter_dim, expert, topk,
        dtype, q_dtype_a, q_dtype_w, q_type, use_g1u1,
        activation, doweight_stage1, hidden_pad, intermediate_pad, is_shuffled,
    )

    # For the standard 2-stage path (not cktile, not 1-stage), replace stage2 with flydsl
    # for ALL expert counts on fp4x2 per_1x32 workloads.
    # E=33/bs=512: proven -20% (d=512), -15% (d=2048) with tile_m=32
    # E=257: flydsl stage2 replaces CSV-tuned CK stage2 with tile_m=32
    if (
        not metadata.run_1stage
        and metadata.ksplit <= 1  # not cktile path
        and q_type == QuantType.per_1x32
        and q_dtype_w == dtypes.fp4x2
        and metadata.stage2 is not None
    ):
        # Force tile_m=32 for all shapes:
        # - E=33/bs=512: overrides heuristic block_m=64/128 for better CU utilization
        # - E=257: matches CSV block_m=32, flydsl replaces CK stage2
        tile_m_override = 32

        # Replace stage2 with flydsl
        from aiter.fused_moe import MOEMetadata
        return MOEMetadata(
            metadata.stage1,
            functools.partial(
                _flydsl_stage2_replacement,
                block_m=metadata.block_m,
                _tile_m_override=tile_m_override,
            ),
            metadata.block_m,
            metadata.ksplit,
            metadata.run_1stage,
            metadata.has_bias,
            metadata.use_non_temporal_load,
        )

    return metadata

_aiter_fused_moe_module.get_2stage_cfgs = _custom_get_2stage_cfgs

from aiter.fused_moe import fused_moe


def custom_kernel(data: input_t) -> output_t:
    """
    Optimized DeepSeek-R1 MXFP4 MoE kernel.

    Key optimizations:
    1. AITER_USE_OPUS_MOE_SORTING=1 (faster sorting kernel)
    2. HIP_FORCE_DEV_KERNARG=1 (device-side kernel args)
    3. HSA_NO_SCRATCH_RECLAIM=1 (persistent GPU scratch)
    4. HSA_NO_SCRATCH_THREAD_LIMITER=1 (full wave occupancy)
    5. Shape-dependent ksplit (cktile for E<=64/bs<=128)
    6. cktile_moe_stage1 split_k=1
    7. flydsl_moe_stage2 for ALL standard 2-stage shapes (E=33 and E=257)
    8. tile_m=32 override for flydsl stage2
    9. Pre-compile flydsl kernels at import time (including E=257)
    """
    (
        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
scrolls · 224 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