Skip to content
KernelIndex
Search⌘K

submission 588558

uoedinburgh8089 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:265af3d5395193cbf7514d6dc46190794c36539d3f28312a4cbad3889f4d1f4a
license declaredunknown
license concludedunknown
authorsuoedinburgh8089
imported2026-08-15

Kernel source

submission.py245 lines
import os
import tempfile

import torch

_fp4x2_str = "torch.float4_e2m1fn_x2"
_bf16_str = "torch.bfloat16"

_HDR = "cu_num,token,model_dim,inter_dim,expert,topk,act_type,dtype,q_dtype_a,q_dtype_w,q_type,use_g1u1,doweight_stage1,block_m,ksplit,us1,kernelName1,err1,us2,kernelName2,err2,us,run_1stage,tflops,bw,_tag"
_COMMON = f"ActivationType.Silu,{_bf16_str},{_fp4x2_str},{_fp4x2_str},QuantType.per_1x32,1,0"
_K = "x"

_KN1_SM = "moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_KN2_SM = "moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"

_ROWS = [
    f"256,16,7168,256,257,9,{_COMMON},16,2,0,{_K},0,0,{_K},0,0,0,0,",
    f"256,128,7168,256,257,9,{_COMMON},16,2,0,{_K},0,0,{_K},0,0,0,0,",
    f"256,512,7168,256,257,9,{_COMMON},32,0,0,{_KN1_SM},0,0,{_KN2_SM},0,0,0,0,",
    f"256,16,7168,512,33,9,{_COMMON},16,2,0,{_K},0,0,{_K},0,0,0,0,",
    f"256,128,7168,512,33,9,{_COMMON},16,2,0,{_K},0,0,{_K},0,0,0,0,",
]

_csv_content = _HDR + "\n" + "\n".join(_ROWS) + "\n"
_cfg_file = tempfile.NamedTemporaryFile(mode="w", suffix=".csv", delete=False, prefix="fmoe_")
_cfg_file.write(_csv_content)
_cfg_file.close()
os.environ["AITER_CONFIG_FMOE"] = _cfg_file.name

import aiter
from aiter import ActivationType, QuantType, dtypes
from aiter.fused_moe import get_inter_dim, get_padded_M, get_2stage_cfgs
from aiter.ops.triton.quant.fused_mxfp4_quant import (
    fused_dynamic_mxfp4_quant_moe_sort,
)
from task import input_t, output_t

_cache = {}


def _init_config(M, topk, E, model_dim, inter_dim, w1, w2, hidden_pad, intermediate_pad, device):
    key = (M, E, inter_dim)
    if key in _cache:
        return _cache[key]

    metadata = get_2stage_cfgs(
        get_padded_M(M), model_dim, inter_dim, E, topk,
        torch.bfloat16, dtypes.fp4x2, dtypes.fp4x2,
        QuantType.per_1x32, True, ActivationType.Silu,
        False, hidden_pad, intermediate_pad, True,
    )

    block_m = int(metadata.block_m)
    ksplit = int(metadata.ksplit)

    s1_n_pad = intermediate_pad // 64 * 64 * 2
    s1_k_pad = hidden_pad // 128 * 128
    s2_n_pad = hidden_pad // 64 * 64
    s2_k_pad = intermediate_pad // 128 * 128

    _, n1, k1 = w1.shape
    _, k2, n2 = w2.shape
    D = n2 if k2 == k1 else n2 * 2
    if w1.dtype is torch.uint32:
        D = D * 8

    kn1 = ""
    kn2 = ""
    use_nt = False
    if hasattr(metadata.stage1, "keywords"):
        kn1 = metadata.stage1.keywords.get("kernelName", "")
        use_nt = metadata.stage1.keywords.get("use_non_temporal_load", False)
    if metadata.stage2 is not None and hasattr(metadata.stage2, "keywords"):
        kn2 = metadata.stage2.keywords.get("kernelName", "")

    max_num_tokens_padded = M * topk + E * block_m - topk
    max_num_m_blocks = (max_num_tokens_padded + block_m - 1) // block_m
    sorted_ids = torch.empty(max_num_tokens_padded, dtype=torch.int32, device=device)
    sorted_weights = torch.empty(max_num_tokens_padded, dtype=torch.float32, device=device)
    sorted_expert_ids = torch.empty(max_num_m_blocks, dtype=torch.int32, device=device)
    num_valid_ids = torch.empty(2, dtype=torch.int32, device=device)
    moe_buf = torch.empty((M, model_dim), dtype=torch.bfloat16, device=device)

    if ksplit >= 2:
        tmp_out = torch.empty((M, topk, w1.shape[1]), dtype=torch.bfloat16, device=device)
        a2_out = torch.empty((M, topk, D), dtype=torch.bfloat16, device=device)
    else:
        a2_out = torch.empty((M, topk, inter_dim), dtype=torch.bfloat16, device=device)
        tmp_out = None

    cfg = {
        "ksplit": ksplit,
        "block_m": block_m,
        "s1_n_pad": s1_n_pad,
        "s1_k_pad": s1_k_pad,
        "s2_n_pad": s2_n_pad,
        "s2_k_pad": s2_k_pad,
        "kn1": kn1,
        "kn2": kn2,
        "use_nt": use_nt,
        "sorted_ids": sorted_ids,
        "sorted_weights": sorted_weights,
        "sorted_expert_ids": sorted_expert_ids,
        "num_valid_ids": num_valid_ids,
        "moe_buf": moe_buf,
        "tmp_out": tmp_out,
        "a2_out": a2_out,
        "inter_dim": inter_dim,
    }

    _cache[key] = cfg
    return cfg


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

    M = hidden_states.shape[0]
    topk = topk_ids.shape[1]
    device = hidden_states.device

    hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
    intermediate_pad = config["d_expert_pad"] - config["d_expert"]

    w1 = gate_up_weight_shuffled
    w2 = down_weight_shuffled
    w1_scale = gate_up_weight_scale_shuffled
    w2_scale = down_weight_scale_shuffled
    w1.is_shuffled = True
    w2.is_shuffled = True

    E, model_dim, inter_dim = get_inter_dim(w1.shape, w2.shape)

    cfg = _init_config(
        M, topk, E, model_dim, inter_dim,
        w1, w2, hidden_pad, intermediate_pad, device,
    )

    block_m = cfg["block_m"]
    ksplit = cfg["ksplit"]
    sorted_ids = cfg["sorted_ids"]
    sorted_weights_buf = cfg["sorted_weights"]
    sorted_expert_ids = cfg["sorted_expert_ids"]
    num_valid_ids = cfg["num_valid_ids"]
    moe_buf = cfg["moe_buf"]

    aiter.moe_sorting_fwd(
        topk_ids, topk_weights,
        sorted_ids, sorted_weights_buf, sorted_expert_ids, num_valid_ids,
        moe_buf, E, block_m, None, None, 0,
    )

    w1_scale_e8m0 = w1_scale.view(dtypes.fp8_e8m0)
    w2_scale_e8m0 = w2_scale.view(dtypes.fp8_e8m0)

    if ksplit >= 2:
        tmp_out = cfg["tmp_out"]
        a2_out = cfg["a2_out"]

        tmp_out.zero_()
        aiter.moe_cktile2stages_gemm1(
            hidden_states, w1, tmp_out,
            sorted_ids, sorted_expert_ids, num_valid_ids, topk,
            cfg["s1_n_pad"], cfg["s1_k_pad"],
            None, None, w1_scale_e8m0, None,
            ActivationType.Silu, block_m, ksplit,
        )

        aiter.silu_and_mul(a2_out, tmp_out)

        aiter.moe_cktile2stages_gemm2(
            a2_out, w2, moe_buf,
            sorted_ids, sorted_expert_ids, num_valid_ids, topk,
            cfg["s2_n_pad"], cfg["s2_k_pad"],
            sorted_weights_buf, None, w2_scale_e8m0, None,
            ActivationType.Silu, block_m,
        )

        return moe_buf

    else:
        a2_out = cfg["a2_out"]

        a1, a1_scale = fused_dynamic_mxfp4_quant_moe_sort(
            hidden_states,
            sorted_ids=sorted_ids,
            num_valid_ids=num_valid_ids,
            token_num=M,
            topk=1,
            block_size=block_m,
        )

        aiter.ck_moe_stage1_fwd(
            a1, w1, w2,
            sorted_ids, sorted_expert_ids, num_valid_ids,
            a2_out, topk,
            cfg["kn1"],
            w1_scale_e8m0, a1_scale,
            block_m,
            None,
            QuantType.per_1x32,
            ActivationType.Silu,
            0,
            cfg["use_nt"],
            torch.bfloat16,
        )

        a2_q, a2_scale = fused_dynamic_mxfp4_quant_moe_sort(
            a2_out.view(-1, inter_dim),
            sorted_ids=sorted_ids,
            num_valid_ids=num_valid_ids,
            token_num=M,
            topk=topk,
            block_size=block_m,
        )
        a2_q = a2_q.view(M, topk, -1)

        aiter.ck_moe_stage2_fwd(
            a2_q, w1, w2,
            sorted_ids, sorted_expert_ids, num_valid_ids,
            moe_buf, topk,
            cfg["kn2"],
            w2_scale_e8m0, a2_scale,
            block_m,
            sorted_weights_buf,
            QuantType.per_1x32,
            ActivationType.Silu,
            cfg["use_nt"],
        )

        return moe_buf
scrolls · 245 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