Skip to content
KernelIndex
Search⌘K

submission 589135

Aniket Sadashiva · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v311.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-589135?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
148.4µs
#164 of 782
2026-03-19

Reported · How evidence levels are derived →

Source and license

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

Techniques

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

fp4"""v311: v252 + metadata-backed native BF16/FP4 E=33 bs=512 path.
split-k- CKTile split_k=2 for S1/S2/S4/S5 via get_2stage_cfgs patch

Kernel source

submission_v311.py425 lines
"""v311: v252 + metadata-backed native BF16/FP4 E=33 bs=512 path.

Key insight: per-shape OPUS switching (v249) may cause correctness issues from
cache_clear() interactions. Instead, set OPUS=1 globally before import.
OPUS sorting helps S6 (~181µs vs ~210µs) and is harmless for CKTile shapes.

- dsv3 CSV ksplit=2 for E=257 bs<=128
- CKTile split_k=2 for S1/S2/S4/S5 via get_2stage_cfgs patch
- E=33 CSV: S6 block_m=32, S7 block_m=128
- OPUS=1 globally (no switching)
- NT=1 globally (helps S3/S6/S7, harmless for CKTile shapes)
- block_m=64+NT for S7 via patch (overrides CSV)
- Cached sorting buffers
- Direct native BF16/FP4 2-stage metadata path for E=33 bs=512
"""
import functools
import os
import sys

import torch

from dataclasses import replace

# ── Step 1: Modify dsv3 CSV for E=257 ksplit=2 on bs<=128 ──
_dsv3_path = "/home/runner/aiter/aiter/configs/model_configs/dsv3_fp4_tuned_fmoe.csv"
try:
    with open(_dsv3_path, "r") as f:
        lines = f.readlines()
    header = lines[0].strip()
    modified_lines = [header + "\n"]
    for line in lines[1:]:
        stripped = line.strip()
        if not stripped:
            continue
        fields = stripped.split(",")
        try:
            token_val = int(fields[1])
            expert_val = int(fields[4])
        except (ValueError, IndexError):
            modified_lines.append(line)
            continue
        if expert_val == 257 and token_val <= 128:
            fields[14] = "2"  # ksplit=2
            modified_lines.append(",".join(fields) + "\n")
        else:
            modified_lines.append(line)
    with open(_dsv3_path, "w") as f:
        f.writelines(modified_lines)
except Exception as e:
    print(f"[v252] dsv3 error: {e}", file=sys.stderr)

# ── Step 2: E=33 CSV for S4/S5 ksplit=2, S6/S7 tuned configs ──
_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_small = (
    "moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_"
    "Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
_k2_small = (
    "moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_"
    "Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
)
_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 = [
    # S4 (bs=16, d=512): ksplit=2
    f"256,16,7168,512,33,9,{_common},32,2,0,{_k1_small},0.0%,0,{_k2_small},0.0%,47,0,67,7691,",
    # S5 (bs=128, d=512): ksplit=2
    f"256,128,7168,512,33,9,{_common},32,2,0,{_k1_small},0.0%,0,{_k2_small},0.0%,58,0,434,6263,",
    # S6 (bs=512, d=512): block_m=32, ksplit=0
    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,",
    # S7 (bs=512, d=2048): block_m=128 (default kernel), ksplit=0
    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:
    e33_path = "/home/runner/aiter/aiter/configs/model_configs/e33_fp4_tuned_fmoe.csv"
    with open(e33_path, "w") as f:
        f.write(_csv_header + "\n")
        for row in _csv_rows:
            f.write(row + "\n")
except Exception:
    pass

# ── Step 3: Global OPUS=1, NT=1 (no per-shape switching!) ──
os.environ["AITER_USE_OPUS_MOE_SORTING"] = "1"
os.environ["AITER_USE_NT"] = "1"

from task import input_t, output_t
from aiter import ActivationType, QuantType, dtypes
import aiter
import aiter.fused_moe as fused_moe_mod


# ═══════════════════════════════════════════════════════════════════════
# Cached sorting buffers (zero alloc after warmup)
# ═══════════════════════════════════════════════════════════════════════
_SORT_BUFS = {}
_DIRECT_E33_CK_BUFS = {}
_DIRECT_E33_CK_DISABLED = False


def _cached_moe_sorting_impl(
    topk_ids, topk_weights, num_experts, model_dim, moebuf_dtype,
    block_size, expert_mask, num_local_tokens, dispatch_policy, use_opus,
):
    device = topk_ids.device
    M, topk = topk_ids.shape
    key = (M, num_experts, block_size, model_dim)

    if key not in _SORT_BUFS:
        max_num_tokens_padded = int(M * topk + num_experts * block_size - topk)
        max_num_m_blocks = int(
            (max_num_tokens_padded + block_size - 1) // block_size
        )
        _SORT_BUFS[key] = (
            torch.empty(max_num_tokens_padded, dtype=dtypes.i32, device=device),
            torch.empty(max_num_tokens_padded, dtype=dtypes.fp32, device=device),
            torch.empty(max_num_m_blocks, dtype=dtypes.i32, device=device),
            torch.empty(2, dtype=dtypes.i32, device=device),
            torch.empty((M, model_dim), dtype=moebuf_dtype, device=device),
        )

    sid, sw, sei, nvi, mb = _SORT_BUFS[key]
    fwd = aiter.moe_sorting_opus_fwd if use_opus else aiter.moe_sorting_fwd
    fwd(
        topk_ids, topk_weights, sid, sw, sei, nvi, mb,
        num_experts, int(block_size), expert_mask, num_local_tokens,
        dispatch_policy,
    )
    return sid, sw, sei, nvi, mb


fused_moe_mod._moe_sorting_impl = _cached_moe_sorting_impl


def _get_direct_e33_ck_bufs(hidden_states, config, topk, block_m):
    token_num = hidden_states.shape[0]
    num_experts = config["n_routed_experts"] + config["n_shared_experts"]
    model_dim = config["d_hidden"]
    inter_dim = config["d_expert"]

    key = (
        hidden_states.device,
        hidden_states.dtype,
        token_num,
        topk,
        num_experts,
        model_dim,
        inter_dim,
        block_m,
    )
    bufs = _DIRECT_E33_CK_BUFS.get(key)
    if bufs is None:
        bufs = {
            "inter_buf": torch.empty(
                (token_num, topk, inter_dim),
                dtype=hidden_states.dtype,
                device=hidden_states.device,
            ),
            "out": torch.empty(
                (token_num, model_dim),
                dtype=hidden_states.dtype,
                device=hidden_states.device,
            ),
        }
        _DIRECT_E33_CK_BUFS[key] = bufs
    return bufs


def _direct_e33_bs512_ck(
    hidden_states,
    gate_up_weight_shuffled,
    down_weight_shuffled,
    gate_up_weight_scale_shuffled,
    down_weight_scale_shuffled,
    topk_weights,
    topk_ids,
    config,
):
    token_num = hidden_states.shape[0]
    topk = topk_ids.shape[1]
    model_dim = config["d_hidden"]
    inter_dim = config["d_expert"]
    hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
    intermediate_pad = config["d_expert_pad"] - config["d_expert"]
    total_experts = config["n_routed_experts"] + config["n_shared_experts"]

    policy_metadata = fused_moe_mod.get_2stage_cfgs(
        token_num,
        model_dim,
        inter_dim,
        total_experts,
        topk,
        hidden_states.dtype,
        dtypes.fp4x2,
        dtypes.fp4x2,
        QuantType.per_1x32,
        True,
        ActivationType.Silu,
        False,
        hidden_pad,
        intermediate_pad,
        True,
    )
    native_metadata = fused_moe_mod.get_2stage_cfgs(
        token_num,
        model_dim,
        inter_dim,
        total_experts,
        topk,
        hidden_states.dtype,
        hidden_states.dtype,
        dtypes.fp4x2,
        QuantType.per_1x32,
        True,
        ActivationType.Silu,
        False,
        hidden_pad,
        intermediate_pad,
        True,
    )
    if native_metadata.run_1stage:
        raise RuntimeError("unexpected 1-stage metadata for direct BF16/FP4 path")
    block_m = policy_metadata.block_m
    bufs = _get_direct_e33_ck_bufs(hidden_states, config, topk, block_m)

    sid, sw, sei, nvi, _ = _cached_moe_sorting_impl(
        topk_ids,
        topk_weights,
        total_experts,
        model_dim,
        hidden_states.dtype,
        block_m,
        None,
        None,
        0,
        True,
    )

    w1_scale = (
        gate_up_weight_scale_shuffled.view(dtypes.fp8_e8m0)
        if gate_up_weight_shuffled.dtype == dtypes.fp4x2
        else gate_up_weight_scale_shuffled
    )
    w2_scale = (
        down_weight_scale_shuffled.view(dtypes.fp8_e8m0)
        if down_weight_shuffled.dtype == dtypes.fp4x2
        else down_weight_scale_shuffled
    )

    bufs["inter_buf"].zero_()
    bufs["out"].zero_()

    native_metadata.stage1(
        hidden_states,
        gate_up_weight_shuffled,
        down_weight_shuffled,
        sid,
        sei,
        nvi,
        bufs["inter_buf"],
        topk,
        block_m=block_m,
        a1_scale=None,
        w1_scale=w1_scale,
        sorted_weights=None,
    )

    native_metadata.stage2(
        bufs["inter_buf"],
        gate_up_weight_shuffled,
        down_weight_shuffled,
        sid,
        sei,
        nvi,
        bufs["out"],
        topk,
        block_m=block_m,
        a2_scale=None,
        w2_scale=w2_scale,
        sorted_weights=sw,
    )

    return bufs["out"]


# ═══════════════════════════════════════════════════════════════════════
# Patch get_2stage_cfgs: CKTile for S1/S2/S4/S5, tuned S7
# ═══════════════════════════════════════════════════════════════════════
_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,
):
    common = (
        model_dim == 7168 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
    )

    if not common:
        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,
        )

    # E=257 bs<512 → CKTile split_k=2 (quant-skip)
    if expert == 257 and inter_dim == 256 and token < 512:
        return _make_cktile_metadata(
            hidden_pad, intermediate_pad, use_g1u1, activation, split_k=2,
        )

    # E=33 d=512 bs<=128 → CKTile split_k=2 (quant-skip)
    if expert == 33 and inter_dim == 512 and token <= 128:
        return _make_cktile_metadata(
            hidden_pad, intermediate_pad, use_g1u1, activation, split_k=2,
        )

    default = _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,
    )

    # S7 (E=33, d=2048, bs=512): block_m=64 + NT
    if expert == 33 and inter_dim == 2048 and token == 512:
        return replace(default, block_m=64, use_non_temporal_load=True)

    return default


fused_moe_mod.get_2stage_cfgs = _patched_get_2stage_cfgs


# ═══════════════════════════════════════════════════════════════════════
# Dispatch — simple, no state switching
# ═══════════════════════════════════════════════════════════════════════
def custom_kernel(data: input_t) -> output_t:
    global _DIRECT_E33_CK_DISABLED

    (
        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

    total_experts = config["n_routed_experts"] + config["n_shared_experts"]
    hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
    intermediate_pad = config["d_expert_pad"] - config["d_expert"]

    if total_experts == 33 and config["bs"] == 512 and not _DIRECT_E33_CK_DISABLED:
        try:
            return _direct_e33_bs512_ck(
                hidden_states,
                gate_up_weight_shuffled,
                down_weight_shuffled,
                gate_up_weight_scale_shuffled,
                down_weight_scale_shuffled,
                topk_weights,
                topk_ids,
                config,
            )
        except Exception as e:
            print(f"[v311] direct CK E33 path failed: {e}", file=sys.stderr)
            _DIRECT_E33_CK_DISABLED = True

    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 · 425 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 573469.

- """Stable default submission.
+ """v311: v252 + metadata-backed native BF16/FP4 E=33 bs=512 path.
- 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
+ Key insight: per-shape OPUS switching (v249) may cause correctness issues from
+ cache_clear() interactions. Instead, set OPUS=1 globally before import.
+ OPUS sorting helps S6 (~181µs vs ~210µs) and is harmless for CKTile shapes.
+
+ - dsv3 CSV ksplit=2 for E=257 bs<=128
+ - CKTile split_k=2 for S1/S2/S4/S5 via get_2stage_cfgs patch
+ - E=33 CSV: S6 block_m=32, S7 block_m=128
+ - OPUS=1 globally (no switching)
+ - NT=1 globally (helps S3/S6/S7, harmless for CKTile shapes)
+ - block_m=64+NT for S7 via patch (overrides CSV)
+ - Cached sorting buffers
+ - Direct native BF16/FP4 2-stage metadata path for E=33 bs=512
"""
import functools
import os
+ import sys
- from task import input_t, output_t
+ import torch
- _CSV_PATH = "/home/runner/aiter/aiter/configs/model_configs/e33_fp4_tuned_fmoe.csv"
- _CSV_HEADER = (
+ from dataclasses import replace
+
+ # ── Step 1: Modify dsv3 CSV for E=257 ksplit=2 on bs<=128 ──
+ _dsv3_path = "/home/runner/aiter/aiter/configs/model_configs/dsv3_fp4_tuned_fmoe.csv"
+ try:
+ with open(_dsv3_path, "r") as f:
+ lines = f.readlines()
+ header = lines[0].strip()
+ modified_lines = [header + "\n"]
+ for line in lines[1:]:
+ stripped = line.strip()
+ if not stripped:
+ continue
+ fields = stripped.split(",")
+ try:
+ token_val = int(fields[1])
+ expert_val = int(fields[4])
+ except (ValueError, IndexError):
+ modified_lines.append(line)
+ continue
+ if expert_val == 257 and token_val <= 128:
+ fields[14] = "2" # ksplit=2
+ modified_lines.append(",".join(fields) + "\n")
+ else:
+ modified_lines.append(line)
+ with open(_dsv3_path, "w") as f:
+ f.writelines(modified_lines)
+ except Exception as e:
+ print(f"[v252] dsv3 error: {e}", file=sys.stderr)
+
+ # ── Step 2: E=33 CSV for S4/S5 ksplit=2, S6/S7 tuned configs ──
+ _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 = (
+ _k1_small = (
+ "moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_"
+ "Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
+ )
+ _k2_small = (
+ "moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_"
+ "Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
+ )
+ _k1_512 = (
"moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_"
"Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
- _K2_512 = (
+ _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"
)
- _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,",
+ _csv_rows = [
+ # S4 (bs=16, d=512): ksplit=2
+ f"256,16,7168,512,33,9,{_common},32,2,0,{_k1_small},0.0%,0,{_k2_small},0.0%,47,0,67,7691,",
+ # S5 (bs=128, d=512): ksplit=2
+ f"256,128,7168,512,33,9,{_common},32,2,0,{_k1_small},0.0%,0,{_k2_small},0.0%,58,0,434,6263,",
+ # S6 (bs=512, d=512): block_m=32, ksplit=0
+ 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,",
+ # S7 (bs=512, d=2048): block_m=128 (default kernel), ksplit=0
+ 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:
+ e33_path = "/home/runner/aiter/aiter/configs/model_configs/e33_fp4_tuned_fmoe.csv"
+ with open(e33_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"
+ # ── Step 3: Global OPUS=1, NT=1 (no per-shape switching!) ──
+ os.environ["AITER_USE_OPUS_MOE_SORTING"] = "1"
+ os.environ["AITER_USE_NT"] = "1"
+ from task import input_t, output_t
from aiter import ActivationType, QuantType, dtypes
+ import aiter
import aiter.fused_moe as fused_moe_mod
- _LAST_RUNTIME_STATE = None
+
+ # ═══════════════════════════════════════════════════════════════════════
+ # Cached sorting buffers (zero alloc after warmup)
+ # ═══════════════════════════════════════════════════════════════════════
+ _SORT_BUFS = {}
+ _DIRECT_E33_CK_BUFS = {}
+ _DIRECT_E33_CK_DISABLED = False
+
+
+ def _cached_moe_sorting_impl(
+ topk_ids, topk_weights, num_experts, model_dim, moebuf_dtype,
+ block_size, expert_mask, num_local_tokens, dispatch_policy, use_opus,
+ ):
+ device = topk_ids.device
+ M, topk = topk_ids.shape
+ key = (M, num_experts, block_size, model_dim)
+
+ if key not in _SORT_BUFS:
+ max_num_tokens_padded = int(M * topk + num_experts * block_size - topk)
+ max_num_m_blocks = int(
+ (max_num_tokens_padded + block_size - 1) // block_size
+ )
+ _SORT_BUFS[key] = (
+ torch.empty(max_num_tokens_padded, dtype=dtypes.i32, device=device),
+ torch.empty(max_num_tokens_padded, dtype=dtypes.fp32, device=device),
+ torch.empty(max_num_m_blocks, dtype=dtypes.i32, device=device),
+ torch.empty(2, dtype=dtypes.i32, device=device),
+ torch.empty((M, model_dim), dtype=moebuf_dtype, device=device),
+ )
+
+ sid, sw, sei, nvi, mb = _SORT_BUFS[key]
+ fwd = aiter.moe_sorting_opus_fwd if use_opus else aiter.moe_sorting_fwd
+ fwd(
+ topk_ids, topk_weights, sid, sw, sei, nvi, mb,
+ num_experts, int(block_size), expert_mask, num_local_tokens,
+ dispatch_policy,
+ )
+ return sid, sw, sei, nvi, mb
+
+
+ fused_moe_mod._moe_sorting_impl = _cached_moe_sorting_impl
+
+
+ def _get_direct_e33_ck_bufs(hidden_states, config, topk, block_m):
+ token_num = hidden_states.shape[0]
+ num_experts = config["n_routed_experts"] + config["n_shared_experts"]
+ model_dim = config["d_hidden"]
+ inter_dim = config["d_expert"]
+
+ key = (
+ hidden_states.device,
+ hidden_states.dtype,
+ token_num,
+ topk,
+ num_experts,
+ model_dim,
+ inter_dim,
+ block_m,
+ )
+ bufs = _DIRECT_E33_CK_BUFS.get(key)
+ if bufs is None:
+ bufs = {
+ "inter_buf": torch.empty(
+ (token_num, topk, inter_dim),
+ dtype=hidden_states.dtype,
+ device=hidden_states.device,
+ ),
+ "out": torch.empty(
+ (token_num, model_dim),
+ dtype=hidden_states.dtype,
+ device=hidden_states.device,
+ ),
+ }
+ _DIRECT_E33_CK_BUFS[key] = bufs
+ return bufs
+
+
+ def _direct_e33_bs512_ck(
+ hidden_states,
+ gate_up_weight_shuffled,
+ down_weight_shuffled,
+ gate_up_weight_scale_shuffled,
+ down_weight_scale_shuffled,
+ topk_weights,
+ topk_ids,
+ config,
+ ):
+ token_num = hidden_states.shape[0]
+ topk = topk_ids.shape[1]
+ model_dim = config["d_hidden"]
+ inter_dim = config["d_expert"]
+ hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
+ intermediate_pad = config["d_expert_pad"] - config["d_expert"]
+ total_experts = config["n_routed_experts"] + config["n_shared_experts"]
+
+ policy_metadata = fused_moe_mod.get_2stage_cfgs(
+ token_num,
+ model_dim,
+ inter_dim,
+ total_experts,
+ topk,
+ hidden_states.dtype,
+ dtypes.fp4x2,
+ dtypes.fp4x2,
+ QuantType.per_1x32,
+ True,
+ ActivationType.Silu,
+ False,
+ hidden_pad,
+ intermediate_pad,
+ True,
+ )
+ native_metadata = fused_moe_mod.get_2stage_cfgs(
+ token_num,
+ model_dim,
+ inter_dim,
+ total_experts,
+ topk,
+ hidden_states.dtype,
+ hidden_states.dtype,
+ dtypes.fp4x2,
+ QuantType.per_1x32,
+ True,
+ ActivationType.Silu,
+ False,
+ hidden_pad,
+ intermediate_pad,
+ True,
+ )
+ if native_metadata.run_1stage:
+ raise RuntimeError("unexpected 1-stage metadata for direct BF16/FP4 path")
+ block_m = policy_metadata.block_m
+ bufs = _get_direct_e33_ck_bufs(hidden_states, config, topk, block_m)
+
+ sid, sw, sei, nvi, _ = _cached_moe_sorting_impl(
+ topk_ids,
+ topk_weights,
+ total_experts,
+ model_dim,
+ hidden_states.dtype,
+ block_m,
+ None,
+ None,
+ 0,
+ True,
+ )
+
+ w1_scale = (
+ gate_up_weight_scale_shuffled.view(dtypes.fp8_e8m0)
+ if gate_up_weight_shuffled.dtype == dtypes.fp4x2
+ else gate_up_weight_scale_shuffled
+ )
+ w2_scale = (
+ down_weight_scale_shuffled.view(dtypes.fp8_e8m0)
+ if down_weight_shuffled.dtype == dtypes.fp4x2
+ else down_weight_scale_shuffled
+ )
+
+ bufs["inter_buf"].zero_()
+ bufs["out"].zero_()
+
+ native_metadata.stage1(
+ hidden_states,
+ gate_up_weight_shuffled,
+ down_weight_shuffled,
+ sid,
+ sei,
+ nvi,
+ bufs["inter_buf"],
+ topk,
+ block_m=block_m,
+ a1_scale=None,
+ w1_scale=w1_scale,
+ sorted_weights=None,
+ )
+
+ native_metadata.stage2(
+ bufs["inter_buf"],
+ gate_up_weight_shuffled,
+ down_weight_shuffled,
+ sid,
+ sei,
+ nvi,
+ bufs["out"],
+ topk,
+ block_m=block_m,
+ a2_scale=None,
+ w2_scale=w2_scale,
+ sorted_weights=sw,
+ )
+
+ return bufs["out"]
+
+
+ # ═══════════════════════════════════════════════════════════════════════
+ # Patch get_2stage_cfgs: CKTile for S1/S2/S4/S5, tuned S7
+ # ═══════════════════════════════════════════════════════════════════════
_ORIGINAL_GET_2STAGE_CFGS = fused_moe_mod.get_2stage_cfgs
⋯ 12 unchanged lines
k_pad_zeros=intermediate_pad // 128 * 128,
activation=activation,
),
- 16,
- split_k,
- False,
- False,
- True,
+ 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,
+ 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,
+ common = (
+ model_dim == 7168 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
)
+ if not common:
+ 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
+ # E=257 bs<512 → CKTile split_k=2 (quant-skip)
+ if expert == 257 and inter_dim == 256 and token < 512:
+ return _make_cktile_metadata(
+ hidden_pad, intermediate_pad, use_g1u1, activation, split_k=2,
+ )
+ # E=33 d=512 bs<=128 → CKTile split_k=2 (quant-skip)
+ if expert == 33 and inter_dim == 512 and token <= 128:
+ return _make_cktile_metadata(
+ hidden_pad, intermediate_pad, use_g1u1, activation, split_k=2,
+ )
- def _set_runtime_state(use_opus, use_nt, ksplit=None):
- global _LAST_RUNTIME_STATE
+ default = _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,
+ )
- state = (use_opus, use_nt, ksplit)
- if state == _LAST_RUNTIME_STATE:
- return
+ # S7 (E=33, d=2048, bs=512): block_m=64 + NT
+ if expert == 33 and inter_dim == 2048 and token == 512:
+ return replace(default, block_m=64, use_non_temporal_load=True)
- fused_moe_mod._USE_OPUS_MOE_SORTING = use_opus
- os.environ["AITER_USE_NT"] = "1" if use_nt else "0"
+ return default
- 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
+ fused_moe_mod.get_2stage_cfgs = _patched_get_2stage_cfgs
- 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
-
-
+ # ═══════════════════════════════════════════════════════════════════════
+ # Dispatch — simple, no state switching
+ # ═══════════════════════════════════════════════════════════════════════
def custom_kernel(data: input_t) -> output_t:
+ global _DIRECT_E33_CK_DISABLED
+
(
- 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,
+ 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)
-
+ total_experts = config["n_routed_experts"] + config["n_shared_experts"]
hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
intermediate_pad = config["d_expert_pad"] - config["d_expert"]
+ if total_experts == 33 and config["bs"] == 512 and not _DIRECT_E33_CK_DISABLED:
+ try:
+ return _direct_e33_bs512_ck(
+ hidden_states,
+ gate_up_weight_shuffled,
+ down_weight_shuffled,
+ gate_up_weight_scale_shuffled,
+ down_weight_scale_shuffled,
+ topk_weights,
+ topk_ids,
+ config,
+ )
+ except Exception as e:
+ print(f"[v311] direct CK E33 path failed: {e}", file=sys.stderr)
+ _DIRECT_E33_CK_DISABLED = True
+
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,
+ 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,
+ a1_scale=None, a2_scale=None,
+ hidden_pad=hidden_pad, intermediate_pad=intermediate_pad,
)
scrolls · 561 diff lines total

Best evidence level for this revision: reported

JSON