Skip to content
KernelIndex
Search⌘K

submission 728592

zainhaider5020 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-728592?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
151.9µs
#197 of 782
2026-04-04

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:ea60e88f785a2022b9656651674235d54fac7ac9b860d352e705870a9a07acbc
license declaredunknown
license concludedunknown
authorszainhaider5020
imported2026-08-15

Kernel source

submission.py103 lines
"""
v39: block_size_M=32 targeted only for E=33d512 bs=128.

Building on v38 benchmark findings:
  - block_m=32 HELPS E=33d512 bs=128 (K_gemm2=512): 114→105µs (-9µs)
    * K=512 > KPerBlock=256 → 2 k-iterations → ok pipelining
    * 3.78 waves vs 11.5 waves for gemm2 (3x fewer)
  - block_m=32 HURTS E=257 bs=128 (K_gemm2=256): 179→185µs (+6µs)
    * K=256 = KPerBlock=256 → 1 k-iteration → no pipelining/prefetch

Fix: block_size_M=32 only when d_expert==512 (E=33d512), not d_expert==256 (E=257).

Key format: (cu_num=256, token, hidden=7168, inter, n_experts, topk, act, dtype,
             q_dtype_a, q_dtype_w, q_type, use_g1u1=True, doweight_stage1=False)
"""

import torch
from task import input_t, output_t

import aiter.fused_moe as _fmoe_mod
from aiter import ActivationType, QuantType
from aiter.fused_moe import fused_moe, get_2stage_cfgs

# ── Kernel name constants ──────────────────────────────────────────────────────
# gemm1 kernels (all FP4X2 a4w4, silu, bf16 output)
_K1_32  = "moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_K1_32w = "moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_K1_64  = "moe_ck2stages_gemm1_256x64x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_K1_128 = "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"

# gemm2 kernels
_K2_32  = "moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
_K2_64  = "moe_ck2stages_gemm2_256x64x128x128_1x4_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
_K2_64s = "moe_ck2stages_gemm2_64x64x128x128_1x1_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
_K2_128 = "moe_ck2stages_gemm2_256x128x128x128_1x4_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"

# ── Config injection ───────────────────────────────────────────────────────────
_ACT = "ActivationType.Silu"
_DT  = "torch.bfloat16"
_QDA = "torch.float4_e2m1fn_x2"
_QDW = "torch.float4_e2m1fn_x2"
_QT  = "QuantType.per_1x32"


def _inject_all_configs():
    entries = {}

    # E=257 d=256 bs=16,128: cktile ksplit=2 (confirmed +36µs,+30µs savings in v35)
    for bs in (16, 128):
        entries[(256, bs, 7168, 256, 257, 9, _ACT, _DT, _QDA, _QDW, _QT, True, False)] = {
            "block_m": 16, "ksplit": 2, "kernelName1": _K1_32, "kernelName2": _K2_32, "run_1stage": False,
        }
    # E=257 d=256 bs=512: CK path (cktile regressed +22µs in v35)
    entries[(256, 512, 7168, 256, 257, 9, _ACT, _DT, _QDA, _QDW, _QT, True, False)] = {
        "block_m": 32, "ksplit": 0, "kernelName1": _K1_32, "kernelName2": _K2_32, "run_1stage": False,
    }

    # E=33 d=512, bs=16,128: cktile ksplit=2 (confirmed savings in v34)
    for bs in (16, 128):
        entries[(256, bs, 7168, 512, 33, 9, _ACT, _DT, _QDA, _QDW, _QT, True, False)] = {
            "block_m": 16, "ksplit": 2, "kernelName1": _K1_32, "kernelName2": _K2_32, "run_1stage": False,
        }
    # E=33 d=512, bs=512: bm=64 with 64-thread sparse gemm2 (-21µs vs bm=128)
    entries[(256, 512, 7168, 512, 33, 9, _ACT, _DT, _QDA, _QDW, _QT, True, False)] = {
        "block_m": 64, "ksplit": 0, "kernelName1": _K1_64, "kernelName2": _K2_64s, "run_1stage": False,
    }

    # E=33 d=2048, bs=512: bm=64, 256-thread gemm2
    entries[(256, 512, 7168, 2048, 33, 9, _ACT, _DT, _QDA, _QDW, _QT, True, False)] = {
        "block_m": 64, "ksplit": 0, "kernelName1": _K1_64, "kernelName2": _K2_64, "run_1stage": False,
    }

    if _fmoe_mod.cfg_2stages is None:
        _fmoe_mod.cfg_2stages = {}
    _fmoe_mod.cfg_2stages.update(entries)
    get_2stage_cfgs.cache_clear()


_inject_all_configs()


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
    _ = gate_up_weight, down_weight, gate_up_weight_scale, down_weight_scale
    # block_m=32 helps E=33d512 bs=128 (K_gemm2=512 > KPerBlock=256 → 2 k-iters)
    # block_m=32 hurts E=257 bs=128 (K_gemm2=256 = KPerBlock=256 → 1 k-iter, no pipeline)
    block_size_M = 32 if (hidden_states.shape[0] == 128 and config.get('d_expert', 0) == 512) else None
    return 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=config["d_hidden_pad"] - config["d_hidden"],
        intermediate_pad=config["d_expert_pad"] - config["d_expert"],
        block_size_M=block_size_M,
    )
scrolls · 103 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