Skip to content
KernelIndex
Search⌘K

submission 745403

Lemonade · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:edfc9abc4fb8c7c1808ccd458ce88e09a81db73ca9b6d2468031030b84a996cb
license declaredunknown
license concludedunknown
authorsLemonade
imported2026-08-15

Techniques

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

split-kk_pad_zeros=hidden_pad // 128 * 128, activation=activation, split_k=ksplit),

Kernel source

submission.py108 lines
import os
import torch
import functools
from task import input_t, output_t

os.environ.setdefault("AITER_LOG_TUNED_CONFIG", "0")

import aiter
from aiter import ActivationType, QuantType, dtypes
from aiter.fused_moe import (
    fused_moe, cktile_moe_stage1, cktile_moe_stage2, ck_moe_stage1,
    MOEMetadata,
    get_2stage_cfgs as _orig_get_2stage_cfgs,
)
import aiter.fused_moe as _fm

# ═══════════════════════════════════════════════════════
# Pre-computed kernel configs
# ═══════════════════════════════════════════════════════

_CK1_32 = "moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_CK1_32v2 = "moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_CK1_64 = "moe_ck2stages_gemm1_256x64x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_CK1_128 = "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_CK2_32v1 = "moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
_CK2_32v3 = "moe_ck2stages_gemm2_256x32x128x128_1x4_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
_CK2_64 = "moe_ck2stages_gemm2_256x64x128x128_1x4_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
_CK2_128 = "moe_ck2stages_gemm2_256x128x128x128_1x4_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"

_TUNED = {
    # E=33, d=512
    (512, 7168, 512, 33, 9):  {"bm": 32, "k1": _CK1_32v2, "k2": _CK2_32v1},
    (256, 7168, 512, 33, 9):  {"bm": 64, "k1": _CK1_64, "k2": _CK2_64},
    (128, 7168, 512, 33, 9):  {"bm": 64, "k1": _CK1_64, "k2": _CK2_64},
    # E=33, d=2048
    (512, 7168, 2048, 33, 9): {"bm": 64, "k1": _CK1_64, "k2": _CK2_64, "nt": True},
    (256, 7168, 2048, 33, 9): {"bm": 64, "k1": _CK1_64, "k2": _CK2_64},
    (128, 7168, 2048, 33, 9): {"bm": 64, "k1": _CK1_64, "k2": _CK2_64},
    (64, 7168, 2048, 33, 9):  {"bm": 32, "k1": _CK1_32, "k2": _CK2_32v3},
    (32, 7168, 2048, 33, 9):  {"bm": 32, "k1": _CK1_32, "k2": _CK2_32v3},
    (16, 7168, 2048, 33, 9):  {"bm": 32, "k1": _CK1_32, "k2": _CK2_32v3},
    # E=257, d=256
    (512, 7168, 256, 257, 9): {"bm": 32, "k1": _CK1_32, "k2": _CK2_32v1},
}


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,
):
    if not (q_type == QuantType.per_1x32 and dtype in [dtypes.bf16, dtypes.fp16]
            and q_dtype_w == dtypes.fp4x2 and activation == ActivationType.Silu and is_shuffled):
        return _orig_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)

    # Path 1: cktile bf16 for token<=256 (skips intermediate requant)
    if token <= 256 and model_dim == 7168 and expert in (33, 257):
        ksplit = 7 if token <= 64 else 4
        return MOEMetadata(
            functools.partial(cktile_moe_stage1,
                n_pad_zeros=intermediate_pad // 64 * 64 * 2,
                k_pad_zeros=hidden_pad // 128 * 128, activation=activation, split_k=ksplit),
            functools.partial(cktile_moe_stage2,
                n_pad_zeros=hidden_pad // 64 * 64,
                k_pad_zeros=intermediate_pad // 128 * 128, activation=activation),
            16, ksplit, False)

    # Path 2: tuned CK kernels for token=512
    key = (token, model_dim, inter_dim, expert, topk)
    if key in _TUNED:
        cfg = _TUNED[key]
        use_nt = cfg.get("nt", False)
        return MOEMetadata(
            functools.partial(ck_moe_stage1, kernelName=cfg["k1"],
                activation=activation, quant_type=q_type, dtype=dtype,
                splitk=0, use_non_temporal_load=use_nt),
            functools.partial(aiter.ck_moe_stage2_fwd, kernelName=cfg["k2"],
                activation=activation, quant_type=q_type, use_non_temporal_load=use_nt),
            cfg["bm"], 0, False)

    return _orig_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)


_fm.get_2stage_cfgs = functools.lru_cache(maxsize=2048)(_custom_get_2stage_cfgs)


def custom_kernel(data: input_t) -> output_t:
    hs = data[0]
    w1 = data[5]
    w2 = data[6]
    w1_scale = data[7]
    w2_scale = data[8]
    tw = data[9]
    ti = data[10]
    cfg = data[11]

    hp = cfg["d_hidden_pad"] - cfg["d_hidden"]
    ip = cfg["d_expert_pad"] - cfg["d_expert"]

    return fused_moe(hs, w1, w2, tw, ti, expert_mask=None,
        activation=ActivationType.Silu, quant_type=QuantType.per_1x32,
        doweight_stage1=False, w1_scale=w1_scale, w2_scale=w2_scale,
        a1_scale=None, a2_scale=None, hidden_pad=hp, intermediate_pad=ip)
scrolls · 108 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