Skip to content
KernelIndex
Search⌘K

submission 738819

chenxingqiang · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:10c12111c35dcf1c7c4318a3ed8c0f44ece6bed2a35a97a879a11128e63f2d27
license declaredunknown
license concludedunknown
authorschenxingqiang
imported2026-08-26

Techniques

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

fp4a_dtype="fp4",

Kernel source

submission.py251 lines
import torch
from task import input_t, output_t
import os

import aiter
import aiter.fused_moe as fused_moe_mod
from aiter import ActivationType, QuantType, dtypes
import functools
from aiter.ops.flydsl.utils import is_flydsl_available
from aiter.ops.flydsl.moe_kernels import flydsl_moe_stage1, flydsl_moe_stage2
from aiter.ops.shuffle import shuffle_scale_a16w4, shuffle_weight_a16w4
from aiter.ops.triton.quant.fused_mxfp4_quant import fused_dynamic_mxfp4_quant_moe_sort

_orig_get_2stage_cfgs = fused_moe_mod.get_2stage_cfgs
_SORT_CACHE = {}
_W2_LAYOUT_CACHE = {}
_CACHE_LIMIT = 8
_FLYDSL_DISABLED = False


def _cache_put(cache, key, value):
    if key in cache:
        cache[key] = value
        return
    if len(cache) >= _CACHE_LIMIT:
        cache.clear()
    cache[key] = value


def _get_cached_sorting(topk_ids, topk_weights, num_experts, model_dim, block_m, out_dtype):
    device = topk_ids.device
    token_num, topk = topk_ids.shape
    max_num_tokens_padded = int(topk_ids.numel() + num_experts * block_m - topk)
    max_num_m_blocks = int((max_num_tokens_padded + block_m - 1) // block_m)
    key = (device, token_num, topk, num_experts, model_dim, block_m, out_dtype)

    bufs = _SORT_CACHE.get(key)
    if bufs is None:
        bufs = (
            torch.empty(max_num_tokens_padded, dtype=dtypes.i32, device=device),
            torch.empty(max_num_tokens_padded, dtype=dtypes.fp32, device=device),
            torch.empty(max_num_m_blocks, dtype=dtypes.i32, device=device),
            torch.empty(2, dtype=dtypes.i32, device=device),
            torch.empty((token_num, model_dim), dtype=out_dtype, device=device),
        )
        _cache_put(_SORT_CACHE, key, bufs)

    sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_out = bufs
    aiter.moe_sorting_opus_fwd(
        topk_ids,
        topk_weights,
        sorted_ids,
        sorted_weights,
        sorted_expert_ids,
        num_valid_ids,
        moe_out,
        num_experts,
        int(block_m),
        None,
        None,
        0,
    )
    return sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_out


def _get_cached_w2_layout(down_weight, down_weight_scale, num_experts):
    key = (
        down_weight.data_ptr(),
        down_weight_scale.data_ptr(),
        tuple(down_weight.shape),
        tuple(down_weight_scale.shape),
        down_weight.device,
    )
    cached = _W2_LAYOUT_CACHE.get(key)
    if cached is not None:
        return cached
    w2_a16w4 = shuffle_weight_a16w4(down_weight, 16, False)
    w2_scale_a16w4 = shuffle_scale_a16w4(
        down_weight_scale.view(num_experts, -1),
        num_experts,
        False,
    )
    _cache_put(_W2_LAYOUT_CACHE, key, (w2_a16w4, w2_scale_a16w4))
    return w2_a16w4, w2_scale_a16w4


def _flydsl_moe(
    hidden_states,
    gate_up_weight_shuffled,
    down_weight,
    gate_up_weight_scale_shuffled,
    down_weight_scale,
    topk_weights,
    topk_ids,
):
    token_num = hidden_states.shape[0]
    model_dim = hidden_states.shape[1]
    topk = topk_ids.shape[1]
    num_experts = gate_up_weight_shuffled.shape[0]
    inter_dim = gate_up_weight_shuffled.shape[1] // 2

    block_m = 128 if token_num >= 512 and inter_dim >= 1024 else 32
    sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_out = _get_cached_sorting(
        topk_ids,
        topk_weights,
        num_experts,
        model_dim,
        block_m,
        hidden_states.dtype,
    )
    w2_a16w4, w2_scale_a16w4 = _get_cached_w2_layout(down_weight, down_weight_scale, num_experts)

    a1_fp4, a1_scale = fused_dynamic_mxfp4_quant_moe_sort(
        hidden_states,
        sorted_ids=sorted_ids,
        num_valid_ids=num_valid_ids,
        token_num=token_num,
        topk=1,
        block_size=block_m,
    )
    stage1_out = flydsl_moe_stage1(
        a=a1_fp4,
        w1=gate_up_weight_shuffled,
        sorted_token_ids=sorted_ids,
        sorted_expert_ids=sorted_expert_ids,
        num_valid_ids=num_valid_ids,
        topk=topk,
        tile_m=block_m,
        tile_n=256,
        tile_k=256,
        a_dtype="fp4",
        b_dtype="fp4",
        out_dtype="bf16",
        act="silu",
        w1_scale=gate_up_weight_scale_shuffled,
        a1_scale=a1_scale,
        sorted_weights=None,
    )
    a2_fp4, a2_scale = fused_dynamic_mxfp4_quant_moe_sort(
        stage1_out.view(-1, inter_dim),
        sorted_ids=sorted_ids,
        num_valid_ids=num_valid_ids,
        token_num=token_num,
        topk=topk,
        block_size=block_m,
    )
    return flydsl_moe_stage2(
        inter_states=a2_fp4.view(token_num, topk, -1),
        w2=w2_a16w4,
        sorted_token_ids=sorted_ids,
        sorted_expert_ids=sorted_expert_ids,
        num_valid_ids=num_valid_ids,
        out=moe_out,
        topk=topk,
        tile_m=block_m,
        tile_n=256,
        tile_k=256,
        a_dtype="fp4",
        b_dtype="fp4",
        out_dtype="bf16",
        mode="atomic",
        w2_scale=w2_scale_a16w4,
        a2_scale=a2_scale,
        sorted_weights=sorted_weights,
    )


def _force_1stage_cfgs(*args, **kwargs):
    meta = _orig_get_2stage_cfgs(*args, **kwargs)
    meta.run_1stage = True

    act = args[10]
    q_type = args[8]
    dt = args[5]
    q_dt_a = args[6]
    q_dt_w = args[7]
    is_g1u1 = args[9]
    dow_s1 = args[11]

    meta.stage1 = functools.partial(
        fused_moe_mod.fused_moe_1stage,
        kernelName="",
        activation=act,
        quant_type=q_type,
    )
    return meta


_USE_1STAGE = os.environ.get("RK_FORCE_MOE_1STAGE", "0") == "1"
if _USE_1STAGE:
    fused_moe_mod.get_2stage_cfgs = _force_1stage_cfgs


def custom_kernel(data: input_t) -> output_t:
    global _FLYDSL_DISABLED
    (
        hidden_states,
        _gate_up_weight,  # raw, unused in baseline path
        down_weight,
        _gate_up_weight_scale,  # raw, unused in baseline path
        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"]

    use_flydsl = os.environ.get("RK_USE_FLYDSL_MOE", "0") == "1"
    if use_flydsl and (not _FLYDSL_DISABLED) and is_flydsl_available():
        try:
            return _flydsl_moe(
                hidden_states=hidden_states,
                gate_up_weight_shuffled=gate_up_weight_shuffled,
                down_weight=down_weight,
                gate_up_weight_scale_shuffled=gate_up_weight_scale_shuffled,
                down_weight_scale=down_weight_scale,
                topk_weights=topk_weights,
                topk_ids=topk_ids,
            )
        except Exception:
            # Disable FlyDSL fallback permanently in this process after first failure.
            _FLYDSL_DISABLED = True

    output = fused_moe_mod.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,
        block_size_M=None,
        moe_sorting_dispatch_policy=0,
        hidden_pad=hidden_pad,
        intermediate_pad=intermediate_pad,
    )

    return output
scrolls · 251 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