Skip to content
KernelIndex
Search⌘K

submission 573469

Aniket Sadashiva · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-573469?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
#199 of 782
2026-03-17

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:39df790f3594326fcd2fb11056b5cbb653eb91036658ba38379d24571cc2dd26
license declaredunknown
license concludedunknown
authorsAniket Sadashiva
imported2026-08-15

Techniques

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

split-k- E=257, bs<512: force CKTile split-k path

Kernel source

submission.py230 lines
"""Stable default submission.

This keeps the best known shape-specialized routing from the local experiment log:
- E=257, bs<512: force CKTile split-k path
- E=33, d=512, bs<=128: force CKTile split-k path
- E=33, large shapes: keep the tuned CSV-backed CK 2-stage kernels
"""
import functools
import os

from task import input_t, output_t

_CSV_PATH = "/home/runner/aiter/aiter/configs/model_configs/e33_fp4_tuned_fmoe.csv"
_CSV_HEADER = (
    "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 = (
    "ActivationType.Silu,torch.bfloat16,torch.float4_e2m1fn_x2,"
    "torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0"
)
_K1_512 = (
    "moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_"
    "Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
_K2_512 = (
    "moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_"
    "Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
)
_K1_2048 = (
    "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_"
    "Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
_K2_2048 = (
    "moe_ck2stages_gemm2_256x128x128x128_1x4_MulABScaleExpertWeightShuffled_v3_"
    "Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
)
_CSV_ROWS = [
    f"256,512,7168,512,33,9,{_COMMON},32,0,0,{_K1_512},0.0%,0,{_K2_512},0.0%,129.79,0,781.78,2884.18,",
    f"256,512,7168,2048,33,9,{_COMMON},128,0,0,{_K1_2048},0.0%,0,{_K2_2048},0.0%,275.08,0,1475.47,5323.27,",
]

try:
    with open(_CSV_PATH, "w") as f:
        f.write(_CSV_HEADER + "\n")
        for row in _CSV_ROWS:
            f.write(row + "\n")
except Exception:
    pass

os.environ["AITER_USE_OPUS_MOE_SORTING"] = "0"
os.environ["AITER_USE_NT"] = "0"

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

_LAST_RUNTIME_STATE = None
_ORIGINAL_GET_2STAGE_CFGS = fused_moe_mod.get_2stage_cfgs


def _make_cktile_metadata(hidden_pad, intermediate_pad, use_g1u1, activation, split_k=2):
    return fused_moe_mod.MOEMetadata(
        functools.partial(
            fused_moe_mod.cktile_moe_stage1,
            n_pad_zeros=intermediate_pad // 64 * 64 * (2 if use_g1u1 else 1),
            k_pad_zeros=hidden_pad // 128 * 128,
            activation=activation,
            split_k=split_k,
        ),
        functools.partial(
            fused_moe_mod.cktile_moe_stage2,
            n_pad_zeros=hidden_pad // 64 * 64,
            k_pad_zeros=intermediate_pad // 128 * 128,
            activation=activation,
        ),
        16,
        split_k,
        False,
        False,
        True,
    )


@functools.lru_cache(maxsize=2048)
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 (
        expert == 257
        and token < 512
        and model_dim == 7168
        and inter_dim == 256
        and topk == 9
        and dtype == dtypes.bf16
        and q_dtype_a == dtypes.fp4x2
        and q_dtype_w == dtypes.fp4x2
        and q_type == QuantType.per_1x32
        and use_g1u1
        and not doweight_stage1
        and is_shuffled
    ):
        return _make_cktile_metadata(hidden_pad, intermediate_pad, use_g1u1, activation, split_k=2)

    if (
        expert == 33
        and token <= 128
        and model_dim == 7168
        and inter_dim == 512
        and topk == 9
        and dtype == dtypes.bf16
        and q_dtype_a == dtypes.fp4x2
        and q_dtype_w == dtypes.fp4x2
        and q_type == QuantType.per_1x32
        and use_g1u1
        and not doweight_stage1
        and is_shuffled
    ):
        return _make_cktile_metadata(hidden_pad, intermediate_pad, use_g1u1, activation, split_k=2)

    return _ORIGINAL_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,
    )


fused_moe_mod.get_2stage_cfgs = _patched_get_2stage_cfgs


def _set_runtime_state(use_opus, use_nt, ksplit=None):
    global _LAST_RUNTIME_STATE

    state = (use_opus, use_nt, ksplit)
    if state == _LAST_RUNTIME_STATE:
        return

    fused_moe_mod._USE_OPUS_MOE_SORTING = use_opus
    os.environ["AITER_USE_NT"] = "1" if use_nt else "0"

    if ksplit in (None, 0):
        os.environ.pop("AITER_KSPLIT", None)
    else:
        os.environ["AITER_KSPLIT"] = str(ksplit)

    fused_moe_mod.use_nt.cache_clear()
    fused_moe_mod.get_ksplit.cache_clear()
    fused_moe_mod.get_2stage_cfgs.cache_clear()
    _ORIGINAL_GET_2STAGE_CFGS.cache_clear()
    _LAST_RUNTIME_STATE = state


def _policy(config):
    total_experts = config["n_routed_experts"] + config["n_shared_experts"]

    if total_experts == 257:
        return False, False, None

    if config["d_expert"] == 512 and config["bs"] <= 128:
        return False, True, None

    return True, True, None


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

    use_opus, use_nt, ksplit = _policy(config)
    _set_runtime_state(use_opus=use_opus, use_nt=use_nt, ksplit=ksplit)

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

    return 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,
        hidden_pad=hidden_pad,
        intermediate_pad=intermediate_pad,
    )
scrolls · 230 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 572094.

- """
- submission__v23: submission__v21, but disable OPUS sorting on the E=257 family.
+ """Stable default submission.
- The new E=257 bs=16/128 CKTile split-k path is already a large win. This branch
- tests whether plain sorting beats OPUS once that path is active.
+ This keeps the best known shape-specialized routing from the local experiment log:
+ - E=257, bs<512: force CKTile split-k path
+ - E=33, d=512, bs<=128: force CKTile split-k path
+ - E=33, large shapes: keep the tuned CSV-backed CK 2-stage kernels
"""
import functools
import os
from task import input_t, output_t
- _csv_path = "/home/runner/aiter/aiter/configs/model_configs/e33_fp4_tuned_fmoe.csv"
- _csv_header = (
+ _CSV_PATH = "/home/runner/aiter/aiter/configs/model_configs/e33_fp4_tuned_fmoe.csv"
+ _CSV_HEADER = (
"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 = (
+ _COMMON = (
"ActivationType.Silu,torch.bfloat16,torch.float4_e2m1fn_x2,"
"torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0"
)
- _k1_512_med = (
+ _K1_512 = (
"moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_"
"Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
- _k2_small = (
+ _K2_512 = (
"moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_"
"Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
)
- _k1_2048 = (
+ _K1_2048 = (
"moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_"
"Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
- _k2_2048 = (
+ _K2_2048 = (
"moe_ck2stages_gemm2_256x128x128x128_1x4_MulABScaleExpertWeightShuffled_v3_"
"Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
)
- _rows = [
- (
- f"256,512,7168,512,33,9,{_common},32,0,0,{_k1_512_med},0.0%,0,"
- f"{_k2_small},0.0%,129.79,0,781.78,2884.18,"
- ),
- (
- f"256,512,7168,2048,33,9,{_common},128,0,0,{_k1_2048},0.0%,0,"
- f"{_k2_2048},0.0%,275.08,0,1475.47,5323.27,"
- ),
+ _CSV_ROWS = [
+ f"256,512,7168,512,33,9,{_COMMON},32,0,0,{_K1_512},0.0%,0,{_K2_512},0.0%,129.79,0,781.78,2884.18,",
+ f"256,512,7168,2048,33,9,{_COMMON},128,0,0,{_K1_2048},0.0%,0,{_K2_2048},0.0%,275.08,0,1475.47,5323.27,",
]
try:
- with open(_csv_path, "w") as f:
- f.write(_csv_header + "\n")
- for row in _rows:
+ with open(_CSV_PATH, "w") as f:
+ f.write(_CSV_HEADER + "\n")
+ for row in _CSV_ROWS:
f.write(row + "\n")
except Exception:
pass
⋯ 4 unchanged lines
from aiter import ActivationType, QuantType, dtypes
import aiter.fused_moe as fused_moe_mod
-
_LAST_RUNTIME_STATE = None
_ORIGINAL_GET_2STAGE_CFGS = fused_moe_mod.get_2stage_cfgs
- def _make_cktile_metadata(
- hidden_pad: int,
- intermediate_pad: int,
- use_g1u1: bool,
- activation,
- ):
+ def _make_cktile_metadata(hidden_pad, intermediate_pad, use_g1u1, activation, split_k=2):
return fused_moe_mod.MOEMetadata(
functools.partial(
fused_moe_mod.cktile_moe_stage1,
n_pad_zeros=intermediate_pad // 64 * 64 * (2 if use_g1u1 else 1),
k_pad_zeros=hidden_pad // 128 * 128,
activation=activation,
- split_k=2,
+ split_k=split_k,
),
functools.partial(
fused_moe_mod.cktile_moe_stage2,
⋯ 2 unchanged lines
activation=activation,
),
16,
- 2,
+ split_k,
False,
False,
True,
⋯ 32 unchanged lines
and not doweight_stage1
and is_shuffled
):
- return _make_cktile_metadata(
- hidden_pad=hidden_pad,
- intermediate_pad=intermediate_pad,
- use_g1u1=use_g1u1,
- activation=activation,
- )
+ return _make_cktile_metadata(hidden_pad, intermediate_pad, use_g1u1, activation, split_k=2)
+
+ if (
+ expert == 33
+ and token <= 128
+ and model_dim == 7168
+ and inter_dim == 512
+ and topk == 9
+ and dtype == dtypes.bf16
+ and q_dtype_a == dtypes.fp4x2
+ and q_dtype_w == dtypes.fp4x2
+ and q_type == QuantType.per_1x32
+ and use_g1u1
+ and not doweight_stage1
+ and is_shuffled
+ ):
+ return _make_cktile_metadata(hidden_pad, intermediate_pad, use_g1u1, activation, split_k=2)
+
return _ORIGINAL_GET_2STAGE_CFGS(
token,
model_dim,
⋯ 16 unchanged lines
fused_moe_mod.get_2stage_cfgs = _patched_get_2stage_cfgs
- def _set_runtime_state(use_opus: bool, use_nt: bool, ksplit: int | None = None) -> None:
+ def _set_runtime_state(use_opus, use_nt, ksplit=None):
global _LAST_RUNTIME_STATE
state = (use_opus, use_nt, ksplit)
⋯ 15 unchanged lines
_LAST_RUNTIME_STATE = state
- def _policy(config: dict) -> tuple[bool, bool, int | None]:
+ def _policy(config):
total_experts = config["n_routed_experts"] + config["n_shared_experts"]
if total_experts == 257:
return False, False, None
if config["d_expert"] == 512 and config["bs"] <= 128:
- return False, True, 2
+ return False, True, None
return True, True, None
scrolls · 166 diff lines total

Best evidence level for this revision: reported

JSON