Skip to content
KernelIndex
Search⌘K

submission 563502

Danishlynx · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v59_flydsl_stage2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-563502?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
175.3µs
#351 of 782
2026-03-15

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:4ef6febf1bbafa6cde20bd4fc65e64f6e913d8ec978e0f058ddf16f7d37f6be9
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 v59 — flydsl stage2 replacement.

Kernel source

submission_v59_flydsl_stage2.py211 lines
"""
MoE MXFP4 v59 — flydsl stage2 replacement.

Strategy:
1. Standard warmup (preshuffle_off, fast ~102s)
2. Pre-compile flydsl stage2 kernels (~5s total)
3. Wrap get_2stage_cfgs to replace stage2 with flydsl in ALL metadata
4. Benchmark's first call triggers preshuffle_on JIT (stage1 needs it)
   but stage2 uses flydsl instead of CK
5. Block_m overrides from v49 included

flydsl stage2 uses MLIR-compiled kernels that may be faster than CK.
"""
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 (from v49)
_block_m_overrides = {
    (128, 9, 33, 512): 32,
    (128, 9, 33, 2048): 32,
    (512, 9, 33, 2048): 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 stage2 wrapper =====
_flydsl_ready = False

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):
    """Replace CK stage2 with flydsl_moe_stage2."""
    from aiter.ops.flydsl.moe_kernels import flydsl_moe_stage2
    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=256,
        a_dtype="fp4",
        b_dtype="fp4",
        out_dtype="bf16",
        mode="atomic",
        w2_scale=w2_scale,
        a2_scale=a2_scale,
        sorted_weights=sorted_weights,
    )

# ===== Wrap get_2stage_cfgs to inject flydsl stage2 =====
_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

# ===== Pre-compile flydsl stage2 kernels =====
def _compile_flydsl():
    global _flydsl_ready
    try:
        from aiter.ops.flydsl.moe_kernels import _get_compiled_stage2
        import time
        print("Compiling flydsl stage2 kernels...", file=sys.stderr)

        shapes = [
            (257, 256),   # E=257, d_expert_pad=256
            (33, 512),    # E=33, d_expert_pad=512
            (33, 2048),   # E=33, d_expert_pad=2048
        ]
        for E, inter_dim in shapes:
            t0 = time.time()
            _get_compiled_stage2(
                model_dim=7168,
                inter_dim=inter_dim,
                experts=E,
                topk=9,
                tile_m=32,
                tile_n=128,
                tile_k=256,
                doweight=True,
                a_dtype="fp4",
                b_dtype="fp4",
                out_dtype="bf16",
                accumulate=True,
            )
            t1 = time.time()
            print(f"  flydsl E={E} inter={inter_dim}: {t1-t0:.1f}s", file=sys.stderr)

        _flydsl_ready = True
        print("flydsl stage2 compilation complete", file=sys.stderr)
    except Exception as e:
        print(f"flydsl compilation FAILED: {e}", file=sys.stderr)
        import traceback
        traceback.print_exc(file=sys.stderr)

# ===== Standard warmup =====
_warmed = False
def _warmup():
    global _warmed
    if _warmed:
        return
    _warmed = True

    # Compile flydsl first (fast, ~5s total)
    _compile_flydsl()

    # Standard warmup (preshuffle_off, ~102s)
    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:
            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()
            print(f"Warmup OK: E={E} d_e={d_expert}", file=sys.stderr)
        except Exception as e:
            print(f"Warmup FAIL: E={E} d_e={d_expert}: {e}", file=sys.stderr)

    print("Warmup complete (v59 flydsl_stage2)", 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

    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 · 211 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 562885.

"""
- MoE MXFP4 v49 — Block_m tuning via monkey-patch.
+ MoE MXFP4 v59 — flydsl stage2 replacement.
- Default block_m values:
- M=16, E=257: 32 | M=128, E=257: 32 | M=512, E=257: 128
- M=16, E=33: 32 | M=128, E=33: 64 | M=512, E=33 d=512: 128 | M=512, E=33 d=2048: 128
+ Strategy:
+ 1. Standard warmup (preshuffle_off, fast ~102s)
+ 2. Pre-compile flydsl stage2 kernels (~5s total)
+ 3. Wrap get_2stage_cfgs to replace stage2 with flydsl in ALL metadata
+ 4. Benchmark's first call triggers preshuffle_on JIT (stage1 needs it)
+ but stage2 uses flydsl instead of CK
+ 5. Block_m overrides from v49 included
- The get_block_size_M heuristic selects block_m based on CU utilization.
- Let's try:
- - E=257 shapes: try block_m=64 for M=128 (default=32)
- - E=33 shapes: try block_m=32 for M=128 (default=64), block_m=64 for M=512 (default=128)
- These try smaller block_m for better CU occupancy or larger for fewer launches.
+ flydsl stage2 uses MLIR-compiled kernels that may be faster than CK.
"""
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
-
- # Monkey-patch get_block_size_M
import aiter.fused_moe as _fmoe
- import functools
- _orig_get_block_size_M = _fmoe.get_block_size_M.__wrapped__ if hasattr(_fmoe.get_block_size_M, '__wrapped__') else None
+ # Block_m overrides (from v49)
+ _block_m_overrides = {
+ (128, 9, 33, 512): 32,
+ (128, 9, 33, 2048): 32,
+ (512, 9, 33, 2048): 64,
+ }
- # Custom block_m selection
- _block_m_overrides = {}
-
- def _setup_overrides():
- """Set block_m overrides for specific shapes."""
- # Try different block_m for E=33 shapes
- # (token, topk, expert, inter_dim) -> block_m
- # Default E=33: M=16→32, M=128→64, M=512→128
- # Try: M=128→32 (smaller tiles, more CU parallelism)
- _block_m_overrides[(128, 9, 33, 512)] = 32
- _block_m_overrides[(128, 9, 33, 2048)] = 32
- # Try: M=512→64 for E=33 (more CU parallelism)
- _block_m_overrides[(512, 9, 33, 512)] = 64
- _block_m_overrides[(512, 9, 33, 2048)] = 64
-
- _setup_overrides()
-
@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:
- result = _block_m_overrides[key]
- print(f" block_m override: {key} -> {result}", file=sys.stderr)
- return result
- # Use original heuristic
+ return _block_m_overrides[key]
cu_num = _fmoe.get_cu_num()
tileN = 128
tgN = (inter_dim + tileN - 1) // tileN
⋯ 5 unchanged lines
rnd = (tg_num + cu_num - 1) // cu_num
empty = cu_num - tg_num % cu_num
tmp.append((rnd, empty, el))
- result = sorted(tmp, key=lambda x: x[:2])[0][-1]
- print(f" block_m default: {key} -> {result}", file=sys.stderr)
- return result
+ return sorted(tmp, key=lambda x: x[:2])[0][-1]
- # Apply monkey-patch
_fmoe.get_block_size_M = _custom_get_block_size_M
-
- # Also clear any cached get_2stage_cfgs
try:
_fmoe.get_2stage_cfgs.cache_clear()
except:
pass
- # Standard warmup
+ # ===== flydsl stage2 wrapper =====
+ _flydsl_ready = False
+
+ 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):
+ """Replace CK stage2 with flydsl_moe_stage2."""
+ from aiter.ops.flydsl.moe_kernels import flydsl_moe_stage2
+ 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=256,
+ a_dtype="fp4",
+ b_dtype="fp4",
+ out_dtype="bf16",
+ mode="atomic",
+ w2_scale=w2_scale,
+ a2_scale=a2_scale,
+ sorted_weights=sorted_weights,
+ )
+
+ # ===== Wrap get_2stage_cfgs to inject flydsl stage2 =====
+ _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
+
+ # ===== Pre-compile flydsl stage2 kernels =====
+ def _compile_flydsl():
+ global _flydsl_ready
+ try:
+ from aiter.ops.flydsl.moe_kernels import _get_compiled_stage2
+ import time
+ print("Compiling flydsl stage2 kernels...", file=sys.stderr)
+
+ shapes = [
+ (257, 256), # E=257, d_expert_pad=256
+ (33, 512), # E=33, d_expert_pad=512
+ (33, 2048), # E=33, d_expert_pad=2048
+ ]
+ for E, inter_dim in shapes:
+ t0 = time.time()
+ _get_compiled_stage2(
+ model_dim=7168,
+ inter_dim=inter_dim,
+ experts=E,
+ topk=9,
+ tile_m=32,
+ tile_n=128,
+ tile_k=256,
+ doweight=True,
+ a_dtype="fp4",
+ b_dtype="fp4",
+ out_dtype="bf16",
+ accumulate=True,
+ )
+ t1 = time.time()
+ print(f" flydsl E={E} inter={inter_dim}: {t1-t0:.1f}s", file=sys.stderr)
+
+ _flydsl_ready = True
+ print("flydsl stage2 compilation complete", file=sys.stderr)
+ except Exception as e:
+ print(f"flydsl compilation FAILED: {e}", file=sys.stderr)
+ import traceback
+ traceback.print_exc(file=sys.stderr)
+
+ # ===== Standard warmup =====
_warmed = False
def _warmup():
global _warmed
⋯ 1 unchanged lines
return
_warmed = True
+ # Compile flydsl first (fast, ~5s total)
+ _compile_flydsl()
+
+ # Standard warmup (preshuffle_off, ~102s)
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
- hidden_pad = d_hidden_pad - d_hidden
- intermediate_pad = d_expert_pad - d_expert
-
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")
⋯ 10 unchanged lines
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:
- 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=hidden_pad,
- intermediate_pad=intermediate_pad,
- )
+ 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()
- print(f" Warmup OK: E={E} d_e={d_expert}", file=sys.stderr)
+ print(f"Warmup OK: E={E} d_e={d_expert}", file=sys.stderr)
except Exception as e:
- print(f" Warmup FAIL: E={E} d_e={d_expert}: {e}", file=sys.stderr)
+ print(f"Warmup FAIL: E={E} d_e={d_expert}: {e}", file=sys.stderr)
- print("Warmup complete (v49 blockm_tune)", file=sys.stderr)
+ print("Warmup complete (v59 flydsl_stage2)", 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
+ (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
return fused_moe(
hidden_states,
scrolls · 257 diff lines total

Best evidence level for this revision: reported

JSON