Skip to content
KernelIndex
Search⌘K

submission 603031

kkosey · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:014c62f233165c5aedc6a56f94b739bda3a07a536e1ded56b358a28331a59e22
license declaredunknown
license concludedunknown
authorskkosey
imported2026-08-15

Kernel source

submission.py66 lines
#!POPCORN gpu MI355X
#!POPCORN leaderboard amd-moe-mxfp4
"""Selective ksplit: ksplit=2 for small batch (B0,B3) → CK Tile internal quant.
ksplit=0 for medium/large batch (B1,B2,B4,B5,B6) → standard CK 2stages.
B0: 128→94µs (-27%), B3: 88→65µs (-27%). Others unchanged."""
import os, sys, shutil

DSV3_CSV = "/home/runner/aiter/aiter/configs/model_configs/dsv3_fp4_tuned_fmoe.csv"

KN1_32_1w = "moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
KN1_32_4w = "moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
KN2_32_1w = "moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"

BASE = "ActivationType.Silu,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0"

entries = [
    # B0 (M=16, E=257, d=256): ksplit=2 → CK Tile (93.9µs vs 128µs)
    f"256,16,7168,256,257,9,{BASE},32,2,40.9,{KN1_32_1w},0.0%,27.8,{KN2_32_1w},1.3%,68.7,0,23.08,20597.69,",
    # B1 (M=128, E=257, d=256): ksplit=0 → standard CK (206µs), use existing DSV3 CSV
    # B2 (M=512, E=257, d=256): ksplit=0 → standard CK (243µs), use existing DSV3 CSV

    # B3 (M=16, E=33, d=512): ksplit=2 → CK Tile (64.6µs vs 88µs)
    f"256,16,7168,512,33,9,{BASE},32,2,50.0,{KN1_32_1w},0.0%,30.0,{KN2_32_1w},0.0%,80.0,0,10.0,5000.0,",
    # B4 (M=128, E=33, d=512): ksplit=0 → standard CK (110µs)
    f"256,128,7168,512,33,9,{BASE},32,0,50.0,{KN1_32_4w},0.0%,30.0,{KN2_32_1w},0.0%,80.0,0,10.0,5000.0,",
    # B5 (M=512, E=33, d=512): ksplit=0 → standard CK (179µs)
    f"256,512,7168,512,33,9,{BASE},32,0,50.0,{KN1_32_4w},0.0%,30.0,{KN2_32_1w},0.0%,80.0,0,10.0,5000.0,",
    # B6 NOT injected — default bm=128 + auto kernels (339µs)
]

with open(DSV3_CSV, 'a') as f:
    for e in entries:
        f.write(e + '\n')

if os.path.exists("/tmp/aiter_configs"):
    shutil.rmtree("/tmp/aiter_configs")

import torch
from task import input_t, output_t
from aiter import ActivationType, QuantType
from aiter.fused_moe import fused_moe

_ACT = int(ActivationType.Silu)
_QT = int(QuantType.per_1x32)
_jit_done = False


def custom_kernel(data: input_t) -> output_t:
    global _jit_done
    (hidden_states, _, _, _, _, w1, w2, w1s, w2s, tw, ti, cfg) = data
    hp = cfg["d_hidden_pad"] - cfg["d_hidden"]
    ip = cfg["d_expert_pad"] - cfg["d_expert"]

    if not _jit_done:
        _jit_done = True
        fused_moe(hidden_states, w1, w2, tw, ti,
                  activation=ActivationType.Silu, quant_type=QuantType.per_1x32,
                  w1_scale=w1s, w2_scale=w2s, hidden_pad=hp, intermediate_pad=ip)
        torch.cuda.synchronize()

    return torch.ops.aiter.fused_moe_(
        hidden_states, w1, w2, tw, ti,
        None, _ACT, _QT, False, w1s, w2s,
        None, None, -1, None, False, None,
        hp, ip, None, None)
scrolls · 66 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