Skip to content
KernelIndex
Search⌘K

submission 691415

rosehulman. · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:9de3a2642f5add5846e8a1a977e3cd29217d5e399494085200384ae76fa3ac1d
license declaredunknown
license concludedunknown
authorsrosehulman.
imported2026-08-15

Kernel source

submission.py233 lines
"""
V86: Hybrid cktile/CK per shape. Always re-sort + direct stage calls.
  - Cktile (ksplit=2): E=257 bs=16/128 (skip quant = big win), E=33 bs=16/128
  - CK (ksplit=0): E=257 bs=512, E=33 bs=512 (all d_expert sizes)
  - Custom CSV for E=33 bs=512 CK shapes (tuned kernels)
  - Always re-sort for leaderboard correctness
"""
#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X
import torch
import os
import aiter
from aiter import ActivationType, QuantType, dtypes
from aiter.fused_moe import (
    fused_moe, get_2stage_cfgs, get_inter_dim, get_padded_M,
    fused_dynamic_mxfp4_quant_moe_sort,
)
from task import input_t, output_t

# Tuned CK kernel names for E=33 large-batch shapes 
CK_S1_64 = "moe_ck2stages_gemm1_256x64x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
FLYDSL_S2_64 = "flydsl_moe2_afp4_wfp4_bf16_t64x256x256_reduce"
FLYDSL_S2_128x256 = "flydsl_moe2_afp4_wfp4_bf16_t64x128x256_reduce"

_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"

def _row(cu, tok, mdim, idim, E, topk, bm, k1, k2, ks=0):
    return f"{cu},{tok},{mdim},{idim},{E},{topk},ActivationType.Silu,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,True,False,{bm},{ks},0,{k1},0,0,{k2},0,0,False,0,0"

def _build_custom_csv():
    rows = [_CSV_HEADER]
    rows.append(_row(256, 512, 7168, 512, 33, 9, 64, CK_S1_64, FLYDSL_S2_64))
    rows.append(_row(256, 512, 7168, 2048, 33, 9, 64, CK_S1_64, FLYDSL_S2_128x256))
    return "\n".join(rows) + "\n"

# Shapes that use CKTILE (ksplit=2, skip quant)
_CKTILE_SHAPES = {
    (16, 257),   # E=257 bs=16: 91 vs 145 (cktile wins big)
    (128, 257),  # E=257 bs=128: 175 vs 204 (cktile wins)
    (16, 33),    # E=33 bs=16: ~same
    (128, 33),   # E=33 bs=128: ~same
}
# All other shapes use CK (ksplit=0)

_cache = {}
_initialized = False
_sort_fn = None


def custom_kernel(data: input_t) -> output_t:
    global _initialized, _sort_fn

    (
        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]
    d_hidden = config["d_hidden"]
    d_hidden_pad = config["d_hidden_pad"]
    d_expert = config["d_expert"]
    d_expert_pad = config["d_expert_pad"]
    E = config["n_routed_experts"] + config["n_shared_experts"]
    topk = config["total_top_k"]
    hidden_pad = d_hidden_pad - d_hidden
    intermediate_pad = d_expert_pad - d_expert

    w1s = gate_up_weight_scale_shuffled.view(dtypes.fp8_e8m0)
    w2s = down_weight_scale_shuffled.view(dtypes.fp8_e8m0)

    if not _initialized:
        csv_path = "/tmp/custom_tuned_fmoe_v86.csv"
        with open(csv_path, 'w') as f:
            f.write(_build_custom_csv())
        aiter_root = os.path.dirname(os.path.abspath(aiter.__file__))
        default_csv = os.path.join(aiter_root, "configs", "tuned_fmoe.csv")
        dsv3_csv = os.path.join(aiter_root, "configs", "model_configs", "dsv3_fp4_tuned_fmoe.csv")
        paths = []
        if os.path.exists(default_csv):
            paths.append(default_csv)
        if os.path.exists(dsv3_csv):
            paths.append(dsv3_csv)
        paths.append(csv_path)
        os.environ["AITER_CONFIG_FMOE"] = ":".join(paths)
        os.environ["AITER_KSPLIT"] = "0"
        get_2stage_cfgs.cache_clear()
        _sort_fn = getattr(aiter, 'moe_sorting_opus_fwd', aiter.moe_sorting_fwd)
        _initialized = True

    shape_key = (M, E, d_expert)
    is_cktile = (M, E) in _CKTILE_SHAPES

    if shape_key not in _cache:
        # FIRST CALL: warm up + build cache
        # Set ksplit for this shape
        if is_cktile:
            os.environ["AITER_KSPLIT"] = "2"
            os.environ["AITER_BYPASS_TUNE_CONFIG"] = "1"
        else:
            os.environ["AITER_KSPLIT"] = "0"
            os.environ["AITER_BYPASS_TUNE_CONFIG"] = "0"
        get_2stage_cfgs.cache_clear()

        result = fused_moe(
            hidden_states, gate_up_weight_shuffled, down_weight_shuffled,
            topk_weights, topk_ids,
            activation=ActivationType.Silu, quant_type=QuantType.per_1x32,
            w1_scale=w1s, w2_scale=w2s,
            hidden_pad=hidden_pad, intermediate_pad=intermediate_pad,
        )

        E2, model_dim, inter_dim = get_inter_dim(
            gate_up_weight_shuffled.shape, down_weight_shuffled.shape
        )
        is_shuffled = getattr(gate_up_weight_shuffled, "is_shuffled", False)
        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, is_shuffled,
        )
        block_m = metadata.block_m

        device = hidden_states.device
        max_tok = M * topk + E * block_m
        max_mb = (max_tok + block_m - 1) // block_m

        sorted_ids = torch.empty(max_tok, dtype=torch.int32, device=device)
        sorted_weights = torch.empty(max_tok, dtype=torch.float32, device=device)
        sorted_expert_ids = torch.empty(max_mb, dtype=torch.int32, device=device)
        num_valid_ids = torch.empty(2, dtype=torch.int32, device=device)
        moe_buf = torch.zeros((M, d_hidden_pad), dtype=torch.bfloat16, device=device)

        entry = {
            'block_m': block_m,
            'metadata': metadata,
            'sorted_ids': sorted_ids,
            'sorted_weights': sorted_weights,
            'sorted_expert_ids': sorted_expert_ids,
            'num_valid_ids': num_valid_ids,
            'moe_buf': moe_buf,
            'is_cktile': is_cktile,
            'inter_dim': inter_dim,
        }

        if not is_cktile:
            # Pre-allocate a2 buffer for CK shapes
            entry['a2_buf'] = torch.empty(
                (M, topk, inter_dim),
                dtype=torch.bfloat16,
                device=device,
            )

        _cache[shape_key] = entry
        return result

    # SUBSEQUENT CALLS: always re-sort + direct stage calls
    c = _cache[shape_key]

    # Always re-sort
    _sort_fn(topk_ids, topk_weights, c['sorted_ids'], c['sorted_weights'],
             c['sorted_expert_ids'], c['num_valid_ids'], c['moe_buf'],
             E, c['block_m'], None, None, 0)

    if c['is_cktile']:
        # Cktile: no quant needed, stage1 creates own a2 buffer
        a2 = c['metadata'].stage1(
            hidden_states,
            gate_up_weight_shuffled, down_weight_shuffled,
            c['sorted_ids'], c['sorted_expert_ids'], c['num_valid_ids'],
            None,  # cktile ignores this
            topk,
            block_m=c['block_m'],
            a1_scale=None,
            w1_scale=w1s,
            sorted_weights=None,
        )
        c['metadata'].stage2(
            a2,
            gate_up_weight_shuffled, down_weight_shuffled,
            c['sorted_ids'], c['sorted_expert_ids'], c['num_valid_ids'],
            c['moe_buf'], topk,
            w2_scale=w2s,
            a2_scale=None,
            block_m=c['block_m'],
            sorted_weights=c['sorted_weights'],
        )
    else:
        # CK: quant input + stage1 + quant inter + stage2
        a1_fp4, a1_scale_sorted = fused_dynamic_mxfp4_quant_moe_sort(
            hidden_states,
            sorted_ids=c['sorted_ids'],
            num_valid_ids=c['num_valid_ids'],
            token_num=M,
            topk=1,
            block_size=c['block_m'],
        )
        c['metadata'].stage1(
            a1_fp4,
            gate_up_weight_shuffled, down_weight_shuffled,
            c['sorted_ids'], c['sorted_expert_ids'], c['num_valid_ids'],
            c['a2_buf'], topk,
            block_m=c['block_m'],
            a1_scale=a1_scale_sorted,
            w1_scale=w1s,
            sorted_weights=None,
        )
        a2_flat = c['a2_buf'].view(-1, c['inter_dim'])
        a2_fp4, a2_scale_sorted = fused_dynamic_mxfp4_quant_moe_sort(
            a2_flat,
            sorted_ids=c['sorted_ids'],
            num_valid_ids=c['num_valid_ids'],
            token_num=M,
            topk=topk,
            block_size=c['block_m'],
        )
        a2_fp4_3d = a2_fp4.view(M, topk, -1)
        c['metadata'].stage2(
            a2_fp4_3d,
            gate_up_weight_shuffled, down_weight_shuffled,
            c['sorted_ids'], c['sorted_expert_ids'], c['num_valid_ids'],
            c['moe_buf'], topk,
            w2_scale=w2s,
            a2_scale=a2_scale_sorted,
            block_m=c['block_m'],
            sorted_weights=c['sorted_weights'],
        )

    return c['moe_buf']
scrolls · 233 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 690372.

"""
- V82: Always re-sort + fused_moe_2stages for correctness in both benchmark and leaderboard.
- - First call: warm up with fused_moe, cache metadata + pre-allocate buffers
- - All subsequent calls: re-sort, then call fused_moe_2stages
- - Cktile for E=33 small-batch shapes via ksplit=2
- - Custom CSV tuning for E=33 large-batch shapes
- - Pre-allocated sort buffers + output buffer to save allocation overhead
+ V86: Hybrid cktile/CK per shape. Always re-sort + direct stage calls.
+ - Cktile (ksplit=2): E=257 bs=16/128 (skip quant = big win), E=33 bs=16/128
+ - CK (ksplit=0): E=257 bs=512, E=33 bs=512 (all d_expert sizes)
+ - Custom CSV for E=33 bs=512 CK shapes (tuned kernels)
+ - Always re-sort for leaderboard correctness
"""
#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X
⋯ 2 unchanged lines
import aiter
from aiter import ActivationType, QuantType, dtypes
from aiter.fused_moe import (
- fused_moe, fused_moe_2stages, get_2stage_cfgs, get_inter_dim, get_padded_M,
+ fused_moe, get_2stage_cfgs, get_inter_dim, get_padded_M,
+ fused_dynamic_mxfp4_quant_moe_sort,
)
from task import input_t, output_t
+ # Tuned CK kernel names for E=33 large-batch shapes
CK_S1_64 = "moe_ck2stages_gemm1_256x64x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
FLYDSL_S2_64 = "flydsl_moe2_afp4_wfp4_bf16_t64x256x256_reduce"
FLYDSL_S2_128x256 = "flydsl_moe2_afp4_wfp4_bf16_t64x128x256_reduce"
⋯ 9 unchanged lines
rows.append(_row(256, 512, 7168, 2048, 33, 9, 64, CK_S1_64, FLYDSL_S2_128x256))
return "\n".join(rows) + "\n"
+ # Shapes that use CKTILE (ksplit=2, skip quant)
_CKTILE_SHAPES = {
- (16, 33),
- (128, 33),
+ (16, 257), # E=257 bs=16: 91 vs 145 (cktile wins big)
+ (128, 257), # E=257 bs=128: 175 vs 204 (cktile wins)
+ (16, 33), # E=33 bs=16: ~same
+ (128, 33), # E=33 bs=128: ~same
}
+ # All other shapes use CK (ksplit=0)
_cache = {}
_initialized = False
⋯ 25 unchanged lines
w2s = down_weight_scale_shuffled.view(dtypes.fp8_e8m0)
if not _initialized:
- csv_path = "/tmp/custom_tuned_fmoe_v82.csv"
+ csv_path = "/tmp/custom_tuned_fmoe_v86.csv"
with open(csv_path, 'w') as f:
f.write(_build_custom_csv())
aiter_root = os.path.dirname(os.path.abspath(aiter.__file__))
⋯ 16 unchanged lines
if shape_key not in _cache:
# FIRST CALL: warm up + build cache
+ # Set ksplit for this shape
if is_cktile:
os.environ["AITER_KSPLIT"] = "2"
+ os.environ["AITER_BYPASS_TUNE_CONFIG"] = "1"
else:
os.environ["AITER_KSPLIT"] = "0"
+ os.environ["AITER_BYPASS_TUNE_CONFIG"] = "0"
+ get_2stage_cfgs.cache_clear()
result = fused_moe(
hidden_states, gate_up_weight_shuffled, down_weight_shuffled,
⋯ 3 unchanged lines
hidden_pad=hidden_pad, intermediate_pad=intermediate_pad,
)
- is_shuffled = getattr(gate_up_weight_shuffled, "is_shuffled", False)
E2, model_dim, inter_dim = get_inter_dim(
gate_up_weight_shuffled.shape, down_weight_shuffled.shape
)
+ is_shuffled = getattr(gate_up_weight_shuffled, "is_shuffled", False)
metadata = get_2stage_cfgs(
get_padded_M(M), model_dim, inter_dim, E, topk,
torch.bfloat16, dtypes.fp4x2, dtypes.fp4x2,
⋯ 12 unchanged lines
num_valid_ids = torch.empty(2, dtype=torch.int32, device=device)
moe_buf = torch.zeros((M, d_hidden_pad), dtype=torch.bfloat16, device=device)
- _cache[shape_key] = {
+ entry = {
'block_m': block_m,
+ 'metadata': metadata,
'sorted_ids': sorted_ids,
'sorted_weights': sorted_weights,
'sorted_expert_ids': sorted_expert_ids,
'num_valid_ids': num_valid_ids,
'moe_buf': moe_buf,
+ 'is_cktile': is_cktile,
+ 'inter_dim': inter_dim,
}
+
+ if not is_cktile:
+ # Pre-allocate a2 buffer for CK shapes
+ entry['a2_buf'] = torch.empty(
+ (M, topk, inter_dim),
+ dtype=torch.bfloat16,
+ device=device,
+ )
+
+ _cache[shape_key] = entry
return result
- # SUBSEQUENT CALLS: always re-sort + fused_moe_2stages
+ # SUBSEQUENT CALLS: always re-sort + direct stage calls
c = _cache[shape_key]
- # Re-sort (handles both benchmark same-data and leaderboard new-data correctly)
+ # Always re-sort
_sort_fn(topk_ids, topk_weights, c['sorted_ids'], c['sorted_weights'],
c['sorted_expert_ids'], c['num_valid_ids'], c['moe_buf'],
E, c['block_m'], None, None, 0)
- c['moe_buf'].zero_()
+ if c['is_cktile']:
+ # Cktile: no quant needed, stage1 creates own a2 buffer
+ a2 = c['metadata'].stage1(
+ hidden_states,
+ gate_up_weight_shuffled, down_weight_shuffled,
+ c['sorted_ids'], c['sorted_expert_ids'], c['num_valid_ids'],
+ None, # cktile ignores this
+ topk,
+ block_m=c['block_m'],
+ a1_scale=None,
+ w1_scale=w1s,
+ sorted_weights=None,
+ )
+ c['metadata'].stage2(
+ a2,
+ gate_up_weight_shuffled, down_weight_shuffled,
+ c['sorted_ids'], c['sorted_expert_ids'], c['num_valid_ids'],
+ c['moe_buf'], topk,
+ w2_scale=w2s,
+ a2_scale=None,
+ block_m=c['block_m'],
+ sorted_weights=c['sorted_weights'],
+ )
+ else:
+ # CK: quant input + stage1 + quant inter + stage2
+ a1_fp4, a1_scale_sorted = fused_dynamic_mxfp4_quant_moe_sort(
+ hidden_states,
+ sorted_ids=c['sorted_ids'],
+ num_valid_ids=c['num_valid_ids'],
+ token_num=M,
+ topk=1,
+ block_size=c['block_m'],
+ )
+ c['metadata'].stage1(
+ a1_fp4,
+ gate_up_weight_shuffled, down_weight_shuffled,
+ c['sorted_ids'], c['sorted_expert_ids'], c['num_valid_ids'],
+ c['a2_buf'], topk,
+ block_m=c['block_m'],
+ a1_scale=a1_scale_sorted,
+ w1_scale=w1s,
+ sorted_weights=None,
+ )
+ a2_flat = c['a2_buf'].view(-1, c['inter_dim'])
+ a2_fp4, a2_scale_sorted = fused_dynamic_mxfp4_quant_moe_sort(
+ a2_flat,
+ sorted_ids=c['sorted_ids'],
+ num_valid_ids=c['num_valid_ids'],
+ token_num=M,
+ topk=topk,
+ block_size=c['block_m'],
+ )
+ a2_fp4_3d = a2_fp4.view(M, topk, -1)
+ c['metadata'].stage2(
+ a2_fp4_3d,
+ gate_up_weight_shuffled, down_weight_shuffled,
+ c['sorted_ids'], c['sorted_expert_ids'], c['num_valid_ids'],
+ c['moe_buf'], topk,
+ w2_scale=w2s,
+ a2_scale=a2_scale_sorted,
+ block_m=c['block_m'],
+ sorted_weights=c['sorted_weights'],
+ )
- fused_moe_2stages(
- hidden_states,
- gate_up_weight_shuffled, down_weight_shuffled,
- topk, c['sorted_ids'], c['sorted_weights'],
- c['sorted_expert_ids'], c['num_valid_ids'], c['moe_buf'],
- True, c['block_m'],
- activation=ActivationType.Silu,
- quant_type=QuantType.per_1x32,
- doweight_stage1=False,
- q_dtype_a=dtypes.fp4x2,
- q_dtype_w=dtypes.fp4x2,
- w1_scale=w1s, w2_scale=w2s,
- a1_scale=None, a2_scale=None,
- hidden_pad=hidden_pad, intermediate_pad=intermediate_pad,
- )
-
return c['moe_buf']
scrolls · 202 diff lines total

Best evidence level for this revision: reported

JSON