Skip to content
KernelIndex
Search⌘K

submission 569026

Danishlynx · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v99_tilem16_e257.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-569026?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
143.6µs
#122 of 782
2026-03-16

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:892aed5faecf8298ffbaa3f74b746ac459c41bcc4c5e0e65c37d2ce7968a28df
license declaredunknown
license concludedunknown
authorsDanishlynx
imported2026-08-15

Techniques

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

fp4MoE MXFP4 v99 — tile_m=16 ONLY for E=257 stage2 (proven -7us on bs=16).

Kernel source

submission_v99_tilem16_e257.py178 lines
"""
MoE MXFP4 v99 — tile_m=16 ONLY for E=257 stage2 (proven -7us on bs=16).
Keep tile_m=32 for E=33 (tile_m=16 causes issues on E=33).
Keep tile_k=128 for all (tile_k=256 proven worse for E=257).
"""
import os
os.environ["AITER_USE_NT"] = "1"

import sys
import functools
import torch
from task import input_t, output_t
from aiter import ActivationType, QuantType
from aiter.fused_moe import fused_moe
import aiter.fused_moe as _fmoe

_block_m_overrides = {
    (128, 9, 33, 512): 32,
    (128, 9, 33, 2048): 32,
    (512, 9, 33, 512): 64,
    (512, 9, 33, 2048): 64,
    (512, 9, 257, 256): 64,
}

@functools.lru_cache(maxsize=2048)
def _custom_get_block_size_M(token, topk, expert, inter_dim):
    key = (token, topk, expert, inter_dim)
    if key in _block_m_overrides:
        return _block_m_overrides[key]
    cu_num = _fmoe.get_cu_num()
    tileN = 128
    tgN = (inter_dim + tileN - 1) // tileN
    support_list = [32, 64, 128]
    tmp = []
    for el in support_list:
        max_num_tokens = token * topk + expert * el - topk
        tg_num = tgN * (max_num_tokens + el - 1) // el
        rnd = (tg_num + cu_num - 1) // cu_num
        empty = cu_num - tg_num % cu_num
        tmp.append((rnd, empty, el))
    return sorted(tmp, key=lambda x: x[:2])[0][-1]

_fmoe.get_block_size_M = _custom_get_block_size_M
try:
    _fmoe.get_2stage_cfgs.cache_clear()
except:
    pass


_flydsl_ready = False
_current_experts = [0]

def _flydsl_stage2_wrapper(inter_states, w1, w2, sorted_token_ids,
                            sorted_expert_ids, num_valid_ids, out, topk,
                            w2_scale=None, a2_scale=None, sorted_weights=None,
                            **_kwargs):
    from aiter.ops.flydsl.moe_kernels import flydsl_moe_stage2
    E = _current_experts[0]
    # tile_m=16 ONLY for E=257 shapes
    tile_m = 16 if E > 100 else 32
    flydsl_moe_stage2(
        inter_states=inter_states, w2=w2,
        sorted_token_ids=sorted_token_ids, sorted_expert_ids=sorted_expert_ids,
        num_valid_ids=num_valid_ids, out=out, topk=topk,
        tile_m=tile_m, tile_n=128, tile_k=128,
        a_dtype="fp4", b_dtype="fp4", out_dtype="bf16", mode="atomic",
        w2_scale=w2_scale, a2_scale=a2_scale, sorted_weights=sorted_weights,
    )


_orig_get_2stage_cfgs = _fmoe.get_2stage_cfgs

@functools.lru_cache(maxsize=2048)
def _patched_get_2stage_cfgs(*args, **kwargs):
    metadata = _orig_get_2stage_cfgs(*args, **kwargs)
    if _flydsl_ready and not metadata.run_1stage and metadata.stage2 is not None:
        metadata.stage2 = functools.partial(_flydsl_stage2_wrapper)
    return metadata

_fmoe.get_2stage_cfgs = _patched_get_2stage_cfgs


def _compile_flydsl():
    global _flydsl_ready
    try:
        from aiter.ops.flydsl.moe_kernels import _get_compiled_stage2
        # E=257: tile_m=16
        _get_compiled_stage2(
            model_dim=7168, inter_dim=256, experts=257, topk=9,
            tile_m=16, tile_n=128, tile_k=128,
            doweight=True, a_dtype="fp4", b_dtype="fp4",
            out_dtype="bf16", accumulate=True,
        )
        # E=33: tile_m=32 (standard)
        for inter_dim in [512, 2048]:
            _get_compiled_stage2(
                model_dim=7168, inter_dim=inter_dim, experts=33, topk=9,
                tile_m=32, tile_n=128, tile_k=128,
                doweight=True, a_dtype="fp4", b_dtype="fp4",
                out_dtype="bf16", accumulate=True,
            )
        _flydsl_ready = True
        print("flydsl ready (v99 tm=16 for E257)", file=sys.stderr)
    except Exception as e:
        print(f"flydsl FAILED: {e}", file=sys.stderr)


_warmed = False
def _warmup():
    global _warmed
    if _warmed:
        return
    _warmed = True
    _compile_flydsl()
    configs = [
        (2, 256, 1, 7168, 256, 8),
        (2, 32, 1, 7168, 512, 8),
        (2, 32, 1, 7168, 2048, 8),
    ]
    for bs, n_routed, n_shared, d_hidden, d_expert, n_experts_per_token in configs:
        E = n_routed + n_shared
        total_topk = n_experts_per_token + n_shared
        d_hidden_pad = ((d_hidden + 255) // 256) * 256
        d_expert_pad = ((d_expert + 255) // 256) * 256
        h = torch.randn(bs, d_hidden, dtype=torch.bfloat16, device="cuda")
        w1 = torch.empty(E, 2 * d_expert_pad, d_hidden_pad // 2,
                         dtype=torch.float4_e2m1fn_x2, device="cuda")
        w2 = torch.empty(E, d_hidden_pad, d_expert_pad // 2,
                         dtype=torch.float4_e2m1fn_x2, device="cuda")
        w1_s = torch.empty(E, 2 * d_expert_pad, d_hidden_pad // 32,
                           dtype=torch.float8_e8m0fnu, device="cuda")
        w2_s = torch.empty(E, d_hidden_pad, d_expert_pad // 32,
                           dtype=torch.float8_e8m0fnu, device="cuda")
        topk_w = torch.ones(bs, total_topk, dtype=torch.float32, device="cuda")
        topk_i = torch.zeros(bs, total_topk, dtype=torch.int32, device="cuda")
        for t in range(bs):
            for k in range(n_experts_per_token):
                topk_i[t, k] = k % n_routed
            for k in range(n_shared):
                topk_i[t, n_experts_per_token + k] = n_routed + k
        try:
            _current_experts[0] = E
            fused_moe(h, w1, w2, topk_w, topk_i,
                      activation=ActivationType.Silu,
                      quant_type=QuantType.per_1x32,
                      w1_scale=w1_s, w2_scale=w2_s,
                      hidden_pad=d_hidden_pad - d_hidden,
                      intermediate_pad=d_expert_pad - d_expert)
            torch.cuda.synchronize()
        except Exception as e:
            print(f"Warmup FAIL: E={E}: {e}", file=sys.stderr)
    print("Warmup complete (v99)", file=sys.stderr)

_warmup()


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

    _current_experts[0] = config["n_routed_experts"] + config["n_shared_experts"]
    return 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=gate_up_weight_scale_shuffled,
        w2_scale=down_weight_scale_shuffled,
        hidden_pad=config["d_hidden_pad"] - config["d_hidden"],
        intermediate_pad=config["d_expert_pad"] - config["d_expert"],
    )
scrolls · 178 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 566515.

"""
- MoE MXFP4 v91 — Aggressive block_m tuning for ALL shapes + flydsl stage2 tk=128.
-
- Override block_m heuristic with empirically better values:
- - E=257 bs=512: block_m=64 (default=128, distributes better across 256 CUs)
- - E=33 bs=128 d=512: block_m=32 (from v75)
- - E=33 bs=128 d=2048: block_m=32 (from v75)
- - E=33 bs=512 d=2048: block_m=64 (from v75)
- - E=33 bs=512 d=512: block_m=64 (try instead of default 128)
+ MoE MXFP4 v99 — tile_m=16 ONLY for E=257 stage2 (proven -7us on bs=16).
+ Keep tile_m=32 for E=33 (tile_m=16 causes issues on E=33).
+ Keep tile_k=128 for all (tile_k=256 proven worse for E=257).
"""
import os
os.environ["AITER_USE_NT"] = "1"
⋯ 7 unchanged lines
import aiter.fused_moe as _fmoe
_block_m_overrides = {
- # E=33 shapes
(128, 9, 33, 512): 32,
(128, 9, 33, 2048): 32,
(512, 9, 33, 512): 64,
(512, 9, 33, 2048): 64,
- # E=257 shapes
(512, 9, 257, 256): 64,
}
⋯ 23 unchanged lines
_flydsl_ready = False
+ _current_experts = [0]
def _flydsl_stage2_wrapper(inter_states, w1, w2, sorted_token_ids,
sorted_expert_ids, num_valid_ids, out, topk,
w2_scale=None, a2_scale=None, sorted_weights=None,
**_kwargs):
from aiter.ops.flydsl.moe_kernels import flydsl_moe_stage2
+ E = _current_experts[0]
+ # tile_m=16 ONLY for E=257 shapes
+ tile_m = 16 if E > 100 else 32
flydsl_moe_stage2(
inter_states=inter_states, w2=w2,
sorted_token_ids=sorted_token_ids, sorted_expert_ids=sorted_expert_ids,
num_valid_ids=num_valid_ids, out=out, topk=topk,
- tile_m=32, tile_n=128, tile_k=128,
+ tile_m=tile_m, tile_n=128, tile_k=128,
a_dtype="fp4", b_dtype="fp4", out_dtype="bf16", mode="atomic",
w2_scale=w2_scale, a2_scale=a2_scale, sorted_weights=sorted_weights,
)
⋯ 15 unchanged lines
global _flydsl_ready
try:
from aiter.ops.flydsl.moe_kernels import _get_compiled_stage2
- for E, inter_dim in [(257, 256), (33, 512), (33, 2048)]:
+ # E=257: tile_m=16
+ _get_compiled_stage2(
+ model_dim=7168, inter_dim=256, experts=257, topk=9,
+ tile_m=16, tile_n=128, tile_k=128,
+ doweight=True, a_dtype="fp4", b_dtype="fp4",
+ out_dtype="bf16", accumulate=True,
+ )
+ # E=33: tile_m=32 (standard)
+ for inter_dim in [512, 2048]:
_get_compiled_stage2(
- model_dim=7168, inter_dim=inter_dim, experts=E, topk=9,
+ model_dim=7168, inter_dim=inter_dim, experts=33, topk=9,
tile_m=32, tile_n=128, tile_k=128,
doweight=True, a_dtype="fp4", b_dtype="fp4",
out_dtype="bf16", accumulate=True,
)
_flydsl_ready = True
- print("flydsl ready (v91)", file=sys.stderr)
+ print("flydsl ready (v99 tm=16 for E257)", file=sys.stderr)
except Exception as e:
print(f"flydsl FAILED: {e}", file=sys.stderr)
⋯ 32 unchanged lines
for k in range(n_shared):
topk_i[t, n_experts_per_token + k] = n_routed + k
try:
+ _current_experts[0] = E
fused_moe(h, w1, w2, topk_w, topk_i,
activation=ActivationType.Silu,
quant_type=QuantType.per_1x32,
⋯ 3 unchanged lines
torch.cuda.synchronize()
except Exception as e:
print(f"Warmup FAIL: E={E}: {e}", file=sys.stderr)
- print("Warmup complete (v91)", file=sys.stderr)
+ print("Warmup complete (v99)", file=sys.stderr)
_warmup()
⋯ 5 unchanged lines
gate_up_weight_scale_shuffled, down_weight_scale_shuffled,
topk_weights, topk_ids, config) = data
+ _current_experts[0] = config["n_routed_experts"] + config["n_shared_experts"]
return fused_moe(
hidden_states,
gate_up_weight_shuffled,
scrolls · 103 diff lines total

Best evidence level for this revision: reported

JSON