Skip to content
KernelIndex
Search⌘K

submission 558039

Eurafat45 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:6dbd86f1a153879c249ae4fe4676c6c76a35c5fb35da910c41c026b55af6fdde
license declaredunknown
license concludedunknown
authorsEurafat45
imported2026-08-15

Techniques

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

fp4Submission template for DeepSeek-R1 MXFP4 MoE kernel.
split-ksplitk=0,
tile-m = 16BLOCK_SIZE_M=16,
tile-n = 4BLOCK_SIZE_N=4,

Kernel source

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

import importlib
import functools
import os
import time

# Must be set before importing aiter.fused_moe (read at import time).
os.environ.setdefault("AITER_USE_OPUS_MOE_SORTING", "1")

import torch
from task import input_t, output_t

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


_fused_moe_mod = importlib.import_module("aiter.fused_moe")
_orig_get_2stage_cfgs = _fused_moe_mod.get_2stage_cfgs
try:
    _fused_mxfp4_quant_moe_sort_fn = (
        importlib.import_module("aiter.ops.triton.quant.fused_mxfp4_quant")
        .fused_dynamic_mxfp4_quant_moe_sort
    )
except Exception:
    _fused_mxfp4_quant_moe_sort_fn = getattr(
        _fused_moe_mod, "fused_dynamic_mxfp4_quant_moe_sort", None
    )

_sort_workspace_cache = {}
_fastpath_workspace_cache = {}
_scale_sort_workspace_cache = {}
_route_bucket_cache = {}
_runtime_cfg_cache = {}
_runtime_state_cache = {}
_fastpath_disabled = False
_fastpath_error_printed = False
_fused_quant_sort_disabled = _fused_mxfp4_quant_moe_sort_fn is None
_fused_quant_sort_error_printed = False

_FUSED_QUANT_MOE_SORT_TOKEN_SWITCH = 1024

_RUNTIME_EXPLORE_SAMPLES = 0
_RUNTIME_RECHECK_INTERVAL = 0

_K_STAGE1_S32 = (
    "moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_"
    "MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
_K_STAGE1_W32 = (
    "moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_"
    "MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
_K_STAGE1_W64 = (
    "moe_ck2stages_gemm1_256x64x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_"
    "MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
_K_STAGE1_W128 = (
    "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_"
    "MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
_K_STAGE2_S32 = (
    "moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_"
    "Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
)
_K_STAGE2_W128 = (
    "moe_ck2stages_gemm2_256x128x128x128_1x4_MulABScaleExpertWeightShuffled_v3_"
    "Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
)
_K_STAGE2_FLYDSL_64 = "flydsl_moe2_afp4_wfp4_bf16_t64x256x256_reduce"
_K_STAGE2_FLYDSL_128 = "flydsl_moe2_afp4_wfp4_bf16_t128x256x256_reduce"

_FORCED_E257_CFG = {
    (128, 256): (
        32,
        _K_STAGE1_S32,
        _K_STAGE2_S32,
    ),
}

_FORCED_E33_CFG = {
    (16, 512): (
        32,
        _K_STAGE1_W32,
        _K_STAGE2_S32,
    ),
    (128, 512): (
        32,
        _K_STAGE1_W32,
        _K_STAGE2_S32,
    ),
    (512, 512): (
        32,
        _K_STAGE1_W32,
        _K_STAGE2_S32,
    ),
    (512, 2048): (
        128,
        _K_STAGE1_W128,
        _K_STAGE2_W128,
    ),
}


def _use_non_temporal_load(token: int, topk: int, expert: int) -> bool:
    use_nt_env = int(os.environ.get("AITER_USE_NT", "-1"))
    if use_nt_env != -1:
        return bool(use_nt_env)
    if int(expert) == 33:
        return True
    return (int(token) * int(topk) // int(expert)) < 64


def _workspace_moe_sorting(
    topk_ids,
    topk_weights,
    num_experts,
    model_dim,
    moebuf_dtype,
    block_size=32,
    expert_mask=None,
    num_local_tokens=None,
    dispatch_policy=0,
):
    device = topk_ids.device
    m, topk = topk_ids.shape
    block_size = int(block_size)
    num_experts = int(num_experts)
    model_dim = int(model_dim)

    max_num_tokens_padded = int(topk_ids.numel() + num_experts * block_size - topk)
    max_num_m_blocks = int((max_num_tokens_padded + block_size - 1) // block_size)

    key = (
        int(device.index if device.index is not None else -1),
        str(device.type),
        int(m),
        int(topk),
        num_experts,
        model_dim,
        str(moebuf_dtype),
        block_size,
    )
    ws = _sort_workspace_cache.get(key)
    if ws is None:
        ws = (
            torch.empty(max_num_tokens_padded, dtype=torch.int32, device=device),
            torch.empty(max_num_tokens_padded, dtype=torch.float32, device=device),
            torch.empty(max_num_m_blocks, dtype=torch.int32, device=device),
            torch.empty(2, dtype=torch.int32, device=device),
            torch.empty((m, model_dim), dtype=moebuf_dtype, device=device),
        )
        _sort_workspace_cache[key] = ws

    sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf = ws

    use_opus = os.environ.get("AITER_USE_OPUS_MOE_SORTING", "0") == "1"
    fwd_fn = (
        _fused_moe_mod.aiter.moe_sorting_opus_fwd
        if use_opus and hasattr(_fused_moe_mod.aiter, "moe_sorting_opus_fwd")
        else _fused_moe_mod.aiter.moe_sorting_fwd
    )
    fwd_fn(
        topk_ids,
        topk_weights,
        sorted_ids,
        sorted_weights,
        sorted_expert_ids,
        num_valid_ids,
        moe_buf,
        num_experts,
        block_size,
        expert_mask,
        num_local_tokens,
        dispatch_policy,
    )
    return sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf


def _cdiv(x: int, y: int) -> int:
    return (int(x) + int(y) - 1) // int(y)


def _alloc_fp8_e8m0(shape, device):
    return torch.empty(shape, dtype=torch.uint8, device=device).view(
        _fused_moe_mod.dtypes.fp8_e8m0
    )


def _get_fastpath_workspace(token_num, topk, model_dim, inter_dim, dtype, device):
    key = (
        int(device.index if device.index is not None else -1),
        str(device.type),
        int(token_num),
        int(topk),
        int(model_dim),
        int(inter_dim),
        str(dtype),
    )
    ws = _fastpath_workspace_cache.get(key)
    if ws is None:
        ws = {
            "a1_q": torch.empty(
                (token_num, model_dim // 2),
                dtype=_fused_moe_mod.dtypes.fp4x2,
                device=device,
            ),
            "a1_scale": _alloc_fp8_e8m0((token_num, model_dim // 32), device),
            "a2_inter": torch.empty(
                (token_num, topk, inter_dim),
                dtype=dtype,
                device=device,
            ),
            "a2_q": torch.empty(
                (token_num * topk, inter_dim // 2),
                dtype=_fused_moe_mod.dtypes.fp4x2,
                device=device,
            ),
            "a2_scale": _alloc_fp8_e8m0((token_num * topk, inter_dim // 32), device),
        }
        _fastpath_workspace_cache[key] = ws
    return ws


def _moe_mxfp4_sort_into(blockscale_e8m0, sorted_ids, num_valid_ids, token_num, tag):
    topk = 1
    src = blockscale_e8m0
    if len(src.shape) == 3:
        topk = int(src.shape[1])
        src = src.view(-1, src.shape[-1])

    m_i, n_i = src.shape
    m_o, n_o = int(sorted_ids.shape[0]), int(n_i)
    shape_u32 = (_cdiv(m_o, 32), _cdiv(n_o, 8), 4, 16)
    key = (
        int(src.device.index if src.device.index is not None else -1),
        str(src.device.type),
        int(m_o),
        int(n_o),
        str(tag),
    )
    out_u32 = _scale_sort_workspace_cache.get(key)
    if out_u32 is None:
        out_u32 = torch.empty(shape_u32, dtype=torch.uint32, device=src.device)
        _scale_sort_workspace_cache[key] = out_u32

    grid = (_cdiv(m_o, 32), _cdiv(n_i, 8))
    _fused_moe_mod.fp4_utils._moe_mxfp4_sort_kernel[grid](
        src.view(torch.uint8),
        sorted_ids,
        num_valid_ids,
        out_u32,
        *src.stride(),
        *out_u32.stride(),
        token_num=token_num,
        M_i=m_i,
        N_i=n_i,
        BLOCK_SIZE_M=16,
        BLOCK_SIZE_N=4,
        TOPK=topk,
    )
    return out_u32.view(_fused_moe_mod.dtypes.fp8_e8m0).view(-1, n_o)


def _try_fused_mxfp4_quant_moe_sort(
    x, sorted_ids, num_valid_ids, token_num, topk, block_m
):
    global _fused_quant_sort_disabled, _fused_quant_sort_error_printed
    if _fused_quant_sort_disabled or int(token_num) > _FUSED_QUANT_MOE_SORT_TOKEN_SWITCH:
        return None, None

    try:
        return _fused_mxfp4_quant_moe_sort_fn(
            x,
            sorted_ids=sorted_ids,
            num_valid_ids=num_valid_ids,
            token_num=int(token_num),
            topk=int(topk),
            block_size=int(block_m),
        )
    except Exception as e:
        _fused_quant_sort_disabled = True
        if not _fused_quant_sort_error_printed:
            print(f"[submission fused quant+sort disabled] {type(e).__name__}: {e}")
            _fused_quant_sort_error_printed = True
        return None, None


def _route_density_bucket(topk_ids: torch.Tensor, num_experts: int) -> str:
    key = (
        int(topk_ids.data_ptr()),
        int(topk_ids._version),
        int(topk_ids.shape[0]),
        int(topk_ids.shape[1]),
        int(num_experts),
    )
    cached = _route_bucket_cache.get(key)
    if cached is not None:
        return cached

    flat = topk_ids.reshape(-1).to(torch.int64)
    hist = torch.bincount(flat, minlength=int(num_experts))
    max_load = int(hist.max().item())
    active = int((hist > 0).sum().item())
    mean_load = float(flat.numel()) / float(max(num_experts, 1))
    max_ratio = max_load / (mean_load + 1e-6)
    active_ratio = float(active) / float(max(num_experts, 1))

    bucket = "skewed" if (max_ratio >= 3.0 or active_ratio <= 0.45) else "balanced"
    if len(_route_bucket_cache) > 128:
        _route_bucket_cache.clear()
    _route_bucket_cache[key] = bucket
    return bucket


def _stage2_callable(kernel_name, activation, q_type, use_non_temporal_load):
    if (
        isinstance(kernel_name, str)
        and kernel_name.startswith("flydsl_")
        and hasattr(_fused_moe_mod, "is_flydsl_available")
        and _fused_moe_mod.is_flydsl_available()
        and hasattr(_fused_moe_mod, "_flydsl_stage2_wrapper")
    ):
        return functools.partial(
            _fused_moe_mod._flydsl_stage2_wrapper,
            kernelName=kernel_name,
        )
    return functools.partial(
        _fused_moe_mod.aiter.ck_moe_stage2_fwd,
        kernelName=kernel_name,
        activation=activation,
        quant_type=q_type,
        use_non_temporal_load=use_non_temporal_load,
    )


def _build_metadata_from_cfg(cfg, token, expert, topk, activation, q_type, dtype):
    block_m, kernel1, kernel2 = cfg
    use_nt = _use_non_temporal_load(int(token), int(topk), int(expert))
    return _fused_moe_mod.MOEMetadata(
        functools.partial(
            _fused_moe_mod.ck_moe_stage1,
            kernelName=kernel1,
            activation=activation,
            quant_type=q_type,
            dtype=dtype,
            splitk=0,
            use_non_temporal_load=use_nt,
        ),
        _stage2_callable(kernel2, activation, q_type, use_nt),
        int(block_m),
        0,
        False,
    )


def _runtime_candidate_cfgs(token, inter_dim, expert, route_bucket):
    token = int(token)
    inter_dim = int(inter_dim)
    expert = int(expert)

    if expert == 257 and inter_dim == 256:
        if token >= 512:
            return [(32, _K_STAGE1_W32, _K_STAGE2_S32)]
        return [(32, _K_STAGE1_S32, _K_STAGE2_S32)]

    if expert == 33 and inter_dim == 512:
        if token >= 512:
            if (
                hasattr(_fused_moe_mod, "is_flydsl_available")
                and _fused_moe_mod.is_flydsl_available()
                and hasattr(_fused_moe_mod, "_flydsl_stage2_wrapper")
            ):
                return [(64, _K_STAGE1_W64, _K_STAGE2_FLYDSL_64)]
            return [(32, _K_STAGE1_W32, _K_STAGE2_S32)]
        return [(32, _K_STAGE1_S32, _K_STAGE2_S32)]

    if expert == 33 and inter_dim == 2048:
        return [(128, _K_STAGE1_W128, _K_STAGE2_W128)]

    return []


def _init_runtime_state(key, cands):
    if len(_runtime_state_cache) > 128:
        _runtime_state_cache.clear()
    state = {
        "cands": tuple(cands),
        "calls": 0,
        "stats": {cfg: [0, 0.0] for cfg in cands},
    }
    _runtime_state_cache[key] = state
    return state


def _pick_runtime_cfg(key, cands):
    state = _runtime_state_cache.get(key)
    if state is None or state.get("cands") != tuple(cands):
        state = _init_runtime_state(key, cands)

    state["calls"] += 1
    stats = state["stats"]

    for cfg in cands:
        if stats[cfg][0] < _RUNTIME_EXPLORE_SAMPLES:
            return cfg, True

    if (
        len(cands) > 1
        and _RUNTIME_RECHECK_INTERVAL > 0
        and (state["calls"] % _RUNTIME_RECHECK_INTERVAL == 0)
    ):
        cfg = min(cands, key=lambda c: stats[c][0])
        return cfg, True

    cfg = min(cands, key=lambda c: stats[c][1] / max(stats[c][0], 1))
    return cfg, False


def _update_runtime_stats(key, cfg, elapsed_us):
    if key is None or cfg is None:
        return
    state = _runtime_state_cache.get(key)
    if state is None:
        return
    entry = state["stats"].get(cfg)
    if entry is None:
        state["stats"][cfg] = [1, float(elapsed_us)]
        return
    entry[0] += 1
    entry[1] += float(elapsed_us)


def _run_fastpath_once(
    hidden_states,
    w1,
    w2,
    topk_weights,
    topk_ids,
    w1_scale,
    w2_scale,
    metadata,
):
    token_num, topk = topk_ids.shape
    expert, model_dim, inter_dim = _fused_moe_mod.get_inter_dim(w1.shape, w2.shape)
    expert = int(expert)
    model_dim = int(model_dim)
    inter_dim = int(inter_dim)
    block_m = int(metadata.block_m)
    dtype = hidden_states.dtype
    device = hidden_states.device

    sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf = (
        _workspace_moe_sorting(
            topk_ids=topk_ids,
            topk_weights=topk_weights,
            num_experts=expert,
            model_dim=model_dim,
            moebuf_dtype=dtype,
            block_size=block_m,
            expert_mask=None,
            num_local_tokens=None,
            dispatch_policy=0,
        )
    )

    ws = _get_fastpath_workspace(
        token_num=token_num,
        topk=topk,
        model_dim=model_dim,
        inter_dim=inter_dim,
        dtype=dtype,
        device=device,
    )

    a1_q, a1_scale_sorted = _try_fused_mxfp4_quant_moe_sort(
        hidden_states,
        sorted_ids=sorted_ids,
        num_valid_ids=num_valid_ids,
        token_num=token_num,
        topk=1,
        block_m=block_m,
    )
    if a1_q is None:
        _fused_moe_mod.aiter.dynamic_per_group_scaled_quant_fp4(
            ws["a1_q"],
            hidden_states,
            ws["a1_scale"],
            32,
            False,
            None,
            1,
        )
        a1_q = ws["a1_q"]
        a1_scale_sorted = _moe_mxfp4_sort_into(
            ws["a1_scale"],
            sorted_ids=sorted_ids,
            num_valid_ids=num_valid_ids,
            token_num=token_num,
            tag=("a1", token_num, model_dim, block_m),
        )

    a2_inter = metadata.stage1(
        a1_q,
        w1,
        w2,
        sorted_ids,
        sorted_expert_ids,
        num_valid_ids,
        ws["a2_inter"],
        topk,
        block_m=block_m,
        a1_scale=a1_scale_sorted,
        w1_scale=(
            w1_scale.view(_fused_moe_mod.dtypes.fp8_e8m0)
            if w1.dtype == _fused_moe_mod.dtypes.fp4x2
            else w1_scale
        ),
        sorted_weights=None,
    )

    a2_q, a2_scale_sorted = _try_fused_mxfp4_quant_moe_sort(
        a2_inter.view(-1, inter_dim),
        sorted_ids=sorted_ids,
        num_valid_ids=num_valid_ids,
        token_num=token_num,
        topk=topk,
        block_m=block_m,
    )
    if a2_q is None:
        _fused_moe_mod.aiter.dynamic_per_group_scaled_quant_fp4(
            ws["a2_q"],
            a2_inter.view(-1, inter_dim),
            ws["a2_scale"],
            32,
            False,
            None,
            int(topk),
        )
        a2_q = ws["a2_q"]
        a2_scale_sorted = _moe_mxfp4_sort_into(
            ws["a2_scale"].view(token_num, topk, -1),
            sorted_ids=sorted_ids,
            num_valid_ids=num_valid_ids,
            token_num=token_num,
            tag=("a2", token_num, topk, inter_dim, block_m),
        )

    metadata.stage2(
        a2_q.view(token_num, topk, -1),
        w1,
        w2,
        sorted_ids,
        sorted_expert_ids,
        num_valid_ids,
        moe_buf,
        topk,
        w2_scale=(
            w2_scale.view(_fused_moe_mod.dtypes.fp8_e8m0)
            if w2.dtype == _fused_moe_mod.dtypes.fp4x2
            else w2_scale
        ),
        a2_scale=a2_scale_sorted,
        block_m=block_m,
        sorted_weights=sorted_weights,
    )
    return moe_buf


def _select_runtime_metadata(
    token,
    model_dim,
    inter_dim,
    expert,
    topk,
    dtype,
    activation,
    q_type,
    route_bucket,
):
    key = (int(token), int(inter_dim), int(expert), str(route_bucket))
    cands = _runtime_candidate_cfgs(token, inter_dim, expert, route_bucket)
    if not cands:
        metadata = _patched_get_2stage_cfgs(
            token=token,
            model_dim=model_dim,
            inter_dim=inter_dim,
            expert=expert,
            topk=topk,
            dtype=dtype,
            q_dtype_a=_fused_moe_mod.dtypes.fp4x2,
            q_dtype_w=_fused_moe_mod.dtypes.fp4x2,
            q_type=q_type,
            use_g1u1=True,
            activation=activation,
            doweight_stage1=False,
            hidden_pad=0,
            intermediate_pad=0,
            is_shuffled=True,
        )
        return metadata, None, None, False

    if len(cands) == 1:
        _runtime_cfg_cache[key] = cands[0]
        metadata = _build_metadata_from_cfg(
            cands[0], token, expert, topk, activation, q_type, dtype
        )
        return metadata, key, cands[0], False

    cfg, need_timing = _pick_runtime_cfg(key, cands)
    _runtime_cfg_cache[key] = cfg
    metadata = _build_metadata_from_cfg(
        cfg, token, expert, topk, activation, q_type, dtype
    )
    return metadata, key, cfg, need_timing


def _fallback_fused_moe(
    hidden_states,
    gate_up_weight_shuffled,
    down_weight_shuffled,
    topk_weights,
    topk_ids,
    gate_up_weight_scale_shuffled,
    down_weight_scale_shuffled,
    hidden_pad,
    intermediate_pad,
):
    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=hidden_pad,
        intermediate_pad=intermediate_pad,
    )


def _patched_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 (
        int(model_dim) == 7168
        and int(expert) == 257
        and int(topk) == 9
        and activation == ActivationType.Silu
        and q_type == QuantType.per_1x32
        and dtype == torch.bfloat16
        and use_g1u1
        and not doweight_stage1
        and is_shuffled
    ):
        cfg = _FORCED_E257_CFG.get((int(token), int(inter_dim)))
        if cfg is not None:
            block_m, kernel1, kernel2 = cfg
            use_nt = _use_non_temporal_load(int(token), int(topk), int(expert))
            return _fused_moe_mod.MOEMetadata(
                functools.partial(
                    _fused_moe_mod.ck_moe_stage1,
                    kernelName=kernel1,
                    activation=activation,
                    quant_type=q_type,
                    dtype=dtype,
                    splitk=0,
                    use_non_temporal_load=use_nt,
                ),
                _stage2_callable(kernel2, activation, q_type, use_nt),
                int(block_m),
                0,
                False,
            )

    cfg = _FORCED_E33_CFG.get((int(token), int(inter_dim)))
    if (
        cfg is not None
        and int(model_dim) == 7168
        and int(expert) == 33
        and int(topk) == 9
        and activation == ActivationType.Silu
        and q_type == QuantType.per_1x32
        and dtype == torch.bfloat16
        and use_g1u1
        and not doweight_stage1
        and is_shuffled
    ):
        block_m, kernel1, kernel2 = cfg
        use_nt = _use_non_temporal_load(int(token), int(topk), int(expert))
        return _fused_moe_mod.MOEMetadata(
            functools.partial(
                _fused_moe_mod.ck_moe_stage1,
                kernelName=kernel1,
                activation=activation,
                quant_type=q_type,
                dtype=dtype,
                splitk=0,
                use_non_temporal_load=use_nt,
            ),
            _stage2_callable(kernel2, activation, q_type, use_nt),
            int(block_m),
            0,
            False,
        )

    return _orig_get_2stage_cfgs(
        token=token,
        model_dim=model_dim,
        inter_dim=inter_dim,
        expert=expert,
        topk=topk,
        dtype=dtype,
        q_dtype_a=q_dtype_a,
        q_dtype_w=q_dtype_w,
        q_type=q_type,
        use_g1u1=use_g1u1,
        activation=activation,
        doweight_stage1=doweight_stage1,
        hidden_pad=hidden_pad,
        intermediate_pad=intermediate_pad,
        is_shuffled=is_shuffled,
    )

# Keep only kernel-path override; do not cache cross-call intermediate/final results.
_fused_moe_mod.moe_sorting = _workspace_moe_sorting
_fused_moe_mod.get_2stage_cfgs = _patched_get_2stage_cfgs


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

    Input data tuple:
        hidden_states:                [M, d_hidden]                           bf16
        gate_up_weight:               [E, 2*d_expert_pad, d_hidden_pad//2]    fp4x2  (raw)
        down_weight:                  [E, d_hidden_pad, d_expert_pad//2]      fp4x2  (raw)
        gate_up_weight_scale:         [E, 2*d_expert_pad, scale_K]            e8m0   (raw)
        down_weight_scale:            [E, d_hidden_pad, scale_K]              e8m0   (raw)
        gate_up_weight_shuffled:      [E, 2*d_expert_pad, d_hidden_pad//2]    fp4x2  (shuffled)
        down_weight_shuffled:         [E, d_hidden_pad, d_expert_pad//2]      fp4x2  (shuffled)
        gate_up_weight_scale_shuffled:[padded, flat]                          e8m0   (shuffled)
        down_weight_scale_shuffled:   [padded, flat]                          e8m0   (shuffled)
        topk_weights:                 [M, total_top_k]                        float32
        topk_ids:                     [M, total_top_k]                        int32
        config:                       dict

    Returns:
        output: [M, d_hidden] bf16
    """
    (
        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"]

    global _fastpath_disabled
    if _fastpath_disabled:
        return _fallback_fused_moe(
            hidden_states,
            gate_up_weight_shuffled,
            down_weight_shuffled,
            topk_weights,
            topk_ids,
            gate_up_weight_scale_shuffled,
            down_weight_scale_shuffled,
            hidden_pad,
            intermediate_pad,
        )

    try:
        token_num = int(hidden_states.shape[0])
        token = int(_fused_moe_mod.get_padded_M(token_num))
        topk = int(topk_ids.shape[1])
        expert, model_dim, inter_dim = _fused_moe_mod.get_inter_dim(
            gate_up_weight_shuffled.shape, down_weight_shuffled.shape
        )
        expert = int(expert)
        model_dim = int(model_dim)
        inter_dim = int(inter_dim)
        use_fastpath = (expert == 257 and inter_dim == 256) or (
            expert == 33 and inter_dim == 512
        )
        if not use_fastpath:
            return _fallback_fused_moe(
                hidden_states,
                gate_up_weight_shuffled,
                down_weight_shuffled,
                topk_weights,
                topk_ids,
                gate_up_weight_scale_shuffled,
                down_weight_scale_shuffled,
                hidden_pad,
                intermediate_pad,
            )

        route_bucket = "balanced"
        run_args = (
            hidden_states,
            gate_up_weight_shuffled,
            down_weight_shuffled,
            topk_weights,
            topk_ids,
            gate_up_weight_scale_shuffled,
            down_weight_scale_shuffled,
        )
        metadata, runtime_key, runtime_cfg, need_timing = _select_runtime_metadata(
            token=token,
            model_dim=model_dim,
            inter_dim=inter_dim,
            expert=expert,
            topk=topk,
            dtype=hidden_states.dtype,
            activation=ActivationType.Silu,
            q_type=QuantType.per_1x32,
            route_bucket=route_bucket,
        )

        if need_timing:
            if hidden_states.is_cuda:
                torch.cuda.synchronize(hidden_states.device)
            t0 = time.perf_counter()
            out = _run_fastpath_once(*run_args, metadata=metadata)
            if hidden_states.is_cuda:
                torch.cuda.synchronize(hidden_states.device)
            _update_runtime_stats(
                runtime_key, runtime_cfg, (time.perf_counter() - t0) * 1e6
            )
            return out

        return _run_fastpath_once(*run_args, metadata=metadata)
    except Exception as e:
        global _fastpath_error_printed
        if not _fastpath_error_printed:
            print(f"[submission fastpath disabled] {type(e).__name__}: {e}")
            _fastpath_error_printed = True
        _fastpath_disabled = True
        return _fallback_fused_moe(
            hidden_states,
            gate_up_weight_shuffled,
            down_weight_shuffled,
            topk_weights,
            topk_ids,
            gate_up_weight_scale_shuffled,
            down_weight_scale_shuffled,
            hidden_pad,
            intermediate_pad,
        )
scrolls · 882 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