Skip to content
KernelIndex
Search⌘K

submission 689780

mega-dmitriy · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v340.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-689780?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
141.3µs
#115 of 782
2026-04-01

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:01c7eb417319ea3ef36f148d0ca4d18c1128611879ef84a71bbe2d88a34c3e3d
license declaredunknown
license concludedunknown
authorsmega-dmitriy
imported2026-08-15

Kernel source

submission_v340.py199 lines
#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X
"""
V340: v262 + explicit cached num_local_tokens for E=33 EP shapes.

Hypothesis:
- fused_moe exposes num_local_tokens as an EP-mode hint and currently infers it
  internally on every call
- supplying a cached 1-element int32 tensor only for E=33 shapes may trim EP
  bookkeeping without changing the tuned kernel selection or numerics
"""
import os

os.environ["AITER_KSPLIT"] = "0"

import aiter.fused_moe as _fmoe
import torch
from task import input_t, output_t

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

_cfg_patched = False
_num_local_tokens_cache = {}


def _get_num_local_tokens(device, m_tokens):
    key = (device.type, device.index if device.index is not None else -1, m_tokens)
    cached = _num_local_tokens_cache.get(key)
    if cached is None:
        cached = torch.tensor([m_tokens], device=device, dtype=torch.int32)
        _num_local_tokens_cache[key] = cached
    return cached


def _maybe_patch_cfg_once(hs, w1s, w2s, hp, ip):
    global _cfg_patched
    if _cfg_patched:
        return
    if hasattr(_fmoe, "fused_dynamic_mxfp4_quant_moe_sort"):
        _of = _fmoe.fused_dynamic_mxfp4_quant_moe_sort
        _qf = _fmoe.get_quant(QuantType.per_1x32)
        from aiter.utility import fp4_utils as _fp4u

        def _cs(hidden_states, sorted_ids, num_valid_ids, token_num, topk, block_size):
            if token_num > 256:
                a1, a1s = _qf(hidden_states, scale=None, quant_dtype=dtypes.fp4x2, num_rows=None)
                a1s = _fp4u.moe_mxfp4_sort(
                    a1s,
                    sorted_ids=sorted_ids,
                    num_valid_ids=num_valid_ids,
                    token_num=token_num,
                    block_size=block_size,
                )
                return a1, a1s
            return _of(
                hidden_states,
                sorted_ids=sorted_ids,
                num_valid_ids=num_valid_ids,
                token_num=token_num,
                topk=topk,
                block_size=block_size,
            )

        _fmoe.fused_dynamic_mxfp4_quant_moe_sort = _cs

    E, md, inter = _fmoe.get_inter_dim(w1s.shape, w2s.shape)
    tok = _fmoe.get_padded_M(hs.shape[0])
    _fmoe.get_2stage_cfgs(
        tok,
        md,
        inter,
        E,
        9,
        hs.dtype,
        dtypes.fp4x2,
        w1s.dtype,
        QuantType.per_1x32,
        inter != w1s.shape[1],
        ActivationType.Silu,
        False,
        hp,
        ip,
        getattr(w1s, "is_shuffled", False),
    )
    cfg = _fmoe.cfg_2stages
    if cfg is None:
        _cfg_patched = True
        return
    from aiter.jit.utils.chip_info import get_cu_num

    cu = get_cu_num()
    c = (
        str(ActivationType.Silu),
        str(torch.bfloat16),
        str(dtypes.fp4x2),
        str(dtypes.fp4x2),
        str(QuantType.per_1x32),
        True,
        False,
    )

    k = (cu, 16, 7168, 256, 257, 9) + c
    if k in cfg:
        cfg[k]["ksplit"] = 2
        cfg[k]["kernelName2"] = "flydsl_moe2_afp4_wfp4_bf16_t32x128x256_atomic"

    k = (cu, 128, 7168, 256, 257, 9) + c
    if k in cfg:
        cfg[k]["ksplit"] = 4
        cfg[k]["kernelName2"] = "flydsl_moe2_afp4_wfp4_bf16_t32x128x256_atomic"

    k = (cu, 512, 7168, 256, 257, 9) + c
    if k in cfg:
        cfg[k]["ksplit"] = 1
        cfg[k]["kernelName2"] = "flydsl_moe2_afp4_wfp4_bf16_t32x128x256_atomic"

    for tok in (16, 128):
        cfg[(cu, tok, 7168, 512, 33, 9) + c] = {
            "block_m": 32,
            "ksplit": 2,
            "kernelName1": "moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",
            "kernelName2": "moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16",
            "run_1stage": False,
        }

    cfg[(cu, 512, 7168, 512, 33, 9) + c] = {
        "block_m": 64,
        "ksplit": 0,
        "kernelName1": "moe_ck2stages_gemm1_256x64x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",
        "kernelName2": "flydsl_moe2_afp4_wfp4_bf16_t32x128x256_atomic",
        "run_1stage": False,
    }
    cfg[(cu, 512, 7168, 2048, 33, 9) + c] = {
        "block_m": 64,
        "ksplit": 0,
        "kernelName1": "moe_ck2stages_gemm1_256x64x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",
        "kernelName2": "flydsl_moe2_afp4_wfp4_bf16_t32x128x256_atomic",
        "run_1stage": False,
    }

    _fmoe.get_2stage_cfgs.cache_clear()
    _cfg_patched = True


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
    hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
    intermediate_pad = config["d_expert_pad"] - config["d_expert"]
    _maybe_patch_cfg_once(
        hidden_states,
        gate_up_weight_shuffled,
        down_weight_shuffled,
        hidden_pad,
        intermediate_pad,
    )

    fused_kwargs = {
        "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,
    }

    total_experts = config["n_routed_experts"] + config["n_shared_experts"]
    if total_experts == 33:
        fused_kwargs["num_local_tokens"] = _get_num_local_tokens(
            hidden_states.device,
            hidden_states.shape[0],
        )

    return fused_moe(
        hidden_states,
        gate_up_weight_shuffled,
        down_weight_shuffled,
        topk_weights,
        topk_ids,
        **fused_kwargs,
    )
scrolls · 199 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