Skip to content
KernelIndex
Search⌘K

submission 707949

josusanmartin · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

v1836_v1777_e257bs16_bs512_original_nocache.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-707949?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
120.7µs
#40 of 782
2026-04-03

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:6fb9b26d1b4703cabe8b8a6cf5fdb7350c3bd6d7be497d13e44cf3cca5713b18
license declaredunknown
license concludedunknown
authorsjosusanmartin
imported2026-08-15

Techniques

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

split-kquant_type=q_type, dtype=dtype, splitk=1, block_m=ROUTE_RAW_BS128_BLOCK_M,
tile-m = 32ROUTE_RAW_BS128_BLOCK_M = 32

Kernel source

v1836_v1777_e257bs16_bs512_original_nocache.py807 lines
from __future__ import annotations

import functools
import inspect
import os
import sys
import types
from pathlib import Path
from typing import Any

import pandas as pd
import torch

# -----------------------------------------------------------------------------
# Kernel constants and environment defaults
# -----------------------------------------------------------------------------

K1_SMALL = "moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
K1_MED = "moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
K1_MED64 = "moe_ck2stages_gemm1_256x64x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
K1_LARGE = "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"

K2_SMALL = "moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
K2_MED = "moe_ck2stages_gemm2_256x64x128x128_1x4_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
K2_LARGE = "moe_ck2stages_gemm2_256x128x128x128_1x4_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"

K2_FLYDSL_64_REDUCE = "flydsl_moe2_afp4_wfp4_bf16_t64x256x256_reduce"
K2_FLYDSL_32X128_ATOMIC = "flydsl_moe2_afp4_wfp4_bf16_t32x128x256_atomic"
K2_FLYDSL_32X256_ATOMIC = "flydsl_moe2_afp4_wfp4_bf16_t32x256x256_atomic"
K2_FLYDSL_64X128_ATOMIC = "flydsl_moe2_afp4_wfp4_bf16_t64x128x256_atomic"
K2_FLYDSL_64X256_ATOMIC = "flydsl_moe2_afp4_wfp4_bf16_t64x256x256_atomic"

os.environ.setdefault("JOSU_SHAPE5_K1", K1_LARGE)
os.environ.setdefault("JOSU_SHAPE5_K2", K2_FLYDSL_64_REDUCE)
os.environ.setdefault("JOSU_SHAPE5_BLOCK_M", "32")

os.environ.setdefault("JOSU_SHAPE6_K1", K1_LARGE)
os.environ.setdefault("JOSU_SHAPE6_K2", K2_FLYDSL_32X128_ATOMIC)
os.environ.setdefault("JOSU_SHAPE6_BLOCK_M", "64")

# Selective d2048 threshold: OFF by default (benchmark showed flat 128 is better)
os.environ.setdefault("JOSU_ENABLE_D2048_FUSED_QUANTSORT", "0")
# Route-aware: OFF by default (triggers preshuffle_off JIT causing 12-min timeout)
os.environ.setdefault("JOSU_ENABLE_ROUTE_AWARE_E33_BS128", "0")
os.environ.setdefault("JOSU_ENABLE_SPLIT_SHARED_E33", "0")
os.environ.setdefault("JOSU_ENABLE_E33_BS512_DIRECT", "0")
os.environ.setdefault("JOSU_RAW_ROUTE_THRESHOLD", "46")
# Row-local NT: remove env var, use aiter heuristic (token*topk//expert < 64)
# E257: all ON. E33 bs16/128: ON. E33 bs512: OFF.
# os.environ.setdefault("AITER_USE_NT", "1")

DEBUG_FMOE = os.getenv("JOSU_DEBUG_FMOE") == "1"
FORCE_REBUILD_FMOE = DEBUG_FMOE or os.getenv("JOSU_REBUILD_FMOE") == "1"
CFG_VERSION = "v_exp_nt_off_e257"

ENABLE_D2048_FUSED_QUANTSORT = os.getenv("JOSU_ENABLE_D2048_FUSED_QUANTSORT", "0") == "1"
ENABLE_ROUTE_AWARE_E33_BS128 = os.getenv("JOSU_ENABLE_ROUTE_AWARE_E33_BS128", "0") == "1"
ENABLE_SPLIT_SHARED_E33 = os.getenv("JOSU_ENABLE_SPLIT_SHARED_E33", "0") == "1"
ENABLE_E33_BS512_DIRECT = os.getenv("JOSU_ENABLE_E33_BS512_DIRECT", "1") == "1"
RAW_ROUTE_THRESHOLD = int(os.getenv("JOSU_RAW_ROUTE_THRESHOLD", "46"))

SHAPE5_K1 = os.environ["JOSU_SHAPE5_K1"]
SHAPE5_K2 = os.environ["JOSU_SHAPE5_K2"]
SHAPE5_BLOCK_M = int(os.getenv("JOSU_SHAPE5_BLOCK_M", "32"))

SHAPE6_K1 = os.environ["JOSU_SHAPE6_K1"]
SHAPE6_K2 = os.environ["JOSU_SHAPE6_K2"]
SHAPE6_BLOCK_M = int(os.getenv("JOSU_SHAPE6_BLOCK_M", "64"))

ROUTE_RAW_BS128_K1 = K1_MED
ROUTE_RAW_BS128_K2 = K2_SMALL
ROUTE_RAW_BS128_BLOCK_M = 32


# -----------------------------------------------------------------------------
# Config bootstrap
# -----------------------------------------------------------------------------

def _bootstrap_cfg() -> str:
    out_path = Path(f"/tmp/josu_cfg_{CFG_VERSION}.csv")
    if out_path.exists() and out_path.stat().st_size > 0 and not FORCE_REBUILD_FMOE:
        return str(out_path)

    source_paths = [
        Path("/home/runner/aiter/aiter/configs/tuned_fmoe.csv"),
        Path("/home/runner/aiter/aiter/configs/model_configs/a8w8_blockscale_tuned_fmoe_qwen3_235b.csv"),
        Path("/home/runner/aiter/aiter/configs/model_configs/dsv3_fp4_tuned_fmoe.csv"),
        Path("/home/runner/aiter/aiter/configs/model_configs/kimik2_fp4_tuned_fmoe.csv"),
    ]

    untuned_path = Path("/home/runner/aiter/aiter/configs/untuned_fmoe.csv")
    if untuned_path.exists():
        keys = pd.read_csv(untuned_path, nrows=0).columns.tolist()
    else:
        keys = []

    frames = [pd.read_csv(path) for path in source_paths if path.exists()]
    merged = pd.concat(frames, ignore_index=True) if frames else pd.DataFrame(columns=keys)

    if "cu_num" not in keys:
        keys.append("cu_num")
    for column in keys:
        if column not in merged.columns:
            merged[column] = ""

    defaults = {column: "" for column in merged.columns}
    defaults.update({
        "cu_num": 256, "act_type": "ActivationType.Silu", "dtype": "torch.bfloat16",
        "q_dtype_a": "torch.float4_e2m1fn_x2", "q_dtype_w": "torch.float4_e2m1fn_x2",
        "q_type": "QuantType.per_1x32", "use_g1u1": 1, "doweight_stage1": 0,
        "block_m": 32, "ksplit": 0, "us1": 0.0, "kernelName1": "", "err1": "0.0%",
        "us2": 0.0, "kernelName2": "", "err2": "0.0%", "us": 0.0, "run_1stage": 0,
        "tflops": 0.0, "bw": 0.0, "_tag": "",
    })
    for column, value in defaults.items():
        if column not in merged.columns:
            merged[column] = value

    online_tuned_rows = [
        {"token": 16, "model_dim": 7168, "inter_dim": 256, "expert": 257, "topk": 9,
         "block_m": 16, "ksplit": 7, "kernelName1": K1_SMALL, "kernelName2": K2_SMALL, "run_1stage": 0},
        {"token": 128, "model_dim": 7168, "inter_dim": 256, "expert": 257, "topk": 9,
         "block_m": 32, "ksplit": 1, "kernelName1": K1_LARGE, "kernelName2": K2_FLYDSL_64_REDUCE, "run_1stage": 0},
        {"token": 512, "model_dim": 7168, "inter_dim": 256, "expert": 257, "topk": 9,
         "block_m": 32, "ksplit": 0, "kernelName1": K1_LARGE, "kernelName2": K2_FLYDSL_64_REDUCE, "run_1stage": 0},
        {"token": 16, "model_dim": 7168, "inter_dim": 512, "expert": 33, "topk": 9,
         "block_m": 16, "ksplit": 4, "kernelName1": K1_SMALL, "kernelName2": K2_SMALL, "run_1stage": 0},
        {"token": 128, "model_dim": 7168, "inter_dim": 512, "expert": 33, "topk": 9,
         "block_m": 32, "ksplit": 1, "kernelName1": K1_LARGE, "kernelName2": K2_FLYDSL_64_REDUCE, "run_1stage": 0},
        {"token": 512, "model_dim": 7168, "inter_dim": 512, "expert": 33, "topk": 9,
         "block_m": SHAPE5_BLOCK_M, "ksplit": 0, "kernelName1": SHAPE5_K1, "kernelName2": SHAPE5_K2, "run_1stage": 0},
        {"token": 512, "model_dim": 7168, "inter_dim": 2048, "expert": 33, "topk": 9,
         "block_m": SHAPE6_BLOCK_M, "ksplit": 0, "kernelName1": SHAPE6_K1, "kernelName2": SHAPE6_K2, "run_1stage": 0},
    ]

    rows = []
    for row in online_tuned_rows:
        seeded = defaults.copy()
        seeded.update(row)
        seeded["us"] = -1.0
        rows.append(seeded)

    merged = pd.concat([merged, pd.DataFrame(rows)], ignore_index=True)

    lookup_keys = [
        "cu_num", "token", "model_dim", "inter_dim", "expert", "topk",
        "act_type", "dtype", "q_dtype_a", "q_dtype_w", "q_type",
        "use_g1u1", "doweight_stage1",
    ]
    for column in lookup_keys:
        if column not in merged.columns:
            merged[column] = defaults.get(column, "")
    merged["_tag_rank"] = merged["_tag"].fillna("").ne("").astype(int)
    merged = (
        merged.sort_values(["_tag_rank", "us"], ascending=[True, True])
        .drop_duplicates(subset=lookup_keys, keep="first")
        .drop(columns=["_tag_rank"])
        .reset_index(drop=True)
    )
    out_path.parent.mkdir(parents=True, exist_ok=True)
    merged.to_csv(out_path, index=False)
    return str(out_path)


os.environ["AITER_CONFIG_FMOE"] = _bootstrap_cfg()


# -----------------------------------------------------------------------------
# AITER imports (after config env is ready)
# -----------------------------------------------------------------------------

from aiter import ActivationType, QuantType  # noqa: E402
from task import input_t, output_t  # noqa: E402

try:
    from aiter.jit.module_quant import dynamic_per_group_scaled_quant_fp4 as _jit_dynamic_quant_fp4
except Exception:
    _jit_dynamic_quant_fp4 = None

# Skip module_moe_sorting preload (25s JIT) — we override all sorting to OPUS
try:
    from aiter.jit.module_moe_sorting_opus import moe_sorting_opus_fwd as _jit_moe_sorting_opus_fwd
except Exception:
    _jit_moe_sorting_opus_fwd = None

try:
    from aiter.jit.module_moe_cktile2stages import cktile_moe_gemm1 as _jit_cktile_moe_gemm1
except Exception:
    _jit_cktile_moe_gemm1 = None

try:
    from aiter.jit.module_activation import cktile_moe_gemm2 as _jit_cktile_moe_gemm2
except Exception:
    _jit_cktile_moe_gemm2 = None

try:
    from aiter.jit.module_moe_ck2stages_fp4x2_fp4x2_preshuffle_on_b16_silu_per_1x32_mulWeightStage2_ import (
        ck_moe_stage1 as _jit_ck_moe_stage1,
        ck_moe_stage2 as _jit_ck_moe_stage2,
    )
except Exception:
    _jit_ck_moe_stage1 = None
    _jit_ck_moe_stage2 = None

from aiter.fused_moe import fused_moe  # noqa: E402
import aiter.fused_moe as _fm  # noqa: E402


# -----------------------------------------------------------------------------
# Shared helpers
# -----------------------------------------------------------------------------

def _param_names(func: Any) -> set[str] | None:
    try:
        return set(inspect.signature(func).parameters)
    except Exception:
        return None


def _retarget_partial(stage, *, kernel_name=None, block_m=None, force_nt=None, force_is_shuffled=None):
    if not isinstance(stage, functools.partial):
        return stage
    params = _param_names(stage.func)
    kw = dict(stage.keywords or {})
    if kernel_name is not None and (params is None or "kernelName" in params or "kernelName" in kw):
        kw["kernelName"] = kernel_name
    if block_m is not None and (params is None or "block_m" in params or "block_m" in kw):
        kw["block_m"] = block_m
    if force_nt is not None:
        if params is None or "use_non_temporal_load" in params or "use_non_temporal_load" in kw:
            kw["use_non_temporal_load"] = bool(force_nt)
        if params is None or "non_temporal_load" in params or "non_temporal_load" in kw:
            kw["non_temporal_load"] = bool(force_nt)
    if force_is_shuffled is not None and (params is None or "is_shuffled" in params or "is_shuffled" in kw):
        kw["is_shuffled"] = force_is_shuffled
    return functools.partial(stage.func, *(stage.args or ()), **kw)


def _with_meta(meta, *, stage1=None, stage2=None, block_m=None, ksplit=None,
               run_1stage=None, has_bias=None, use_non_temporal_load=None):
    return _fm.MOEMetadata(
        meta.stage1 if stage1 is None else stage1,
        meta.stage2 if stage2 is None else stage2,
        meta.block_m if block_m is None else block_m,
        meta.ksplit if ksplit is None else ksplit,
        meta.run_1stage if run_1stage is None else run_1stage,
        meta.has_bias if has_bias is None else has_bias,
        getattr(meta, "use_non_temporal_load", False) if use_non_temporal_load is None else use_non_temporal_load,
    )


def _stage2_ck_func():
    return getattr(_fm, "ck_moe_stage2", None) or getattr(getattr(_fm, "aiter", None), "ck_moe_stage2_fwd", None)


def _is_ck_stage2_func(func):
    s2 = _stage2_ck_func()
    s2a = getattr(getattr(_fm, "aiter", None), "ck_moe_stage2_fwd", None)
    return func is not None and (func is s2 or func is s2a)


def _apply_nt_policy(meta, *, expert, token=512, topk=9):
    # NT OFF for E257 (submission intent) and dense E33 bs512; NT ON only for sparse E33.
    if expert >= 257 or (token * topk // expert) >= 64:
        s1, s2, touched = meta.stage1, meta.stage2, False
        if isinstance(s1, functools.partial):
            s1 = _retarget_partial(s1, force_nt=False)
            touched = True
        if isinstance(s2, functools.partial):
            s2 = _retarget_partial(s2, force_nt=False)
            touched = True
        if not touched and not getattr(meta, "use_non_temporal_load", False):
            return meta
        return _with_meta(meta, stage1=s1, stage2=s2, use_non_temporal_load=False)
    s1, s2, touched = meta.stage1, meta.stage2, False
    if isinstance(s1, functools.partial) and s1.func is _fm.ck_moe_stage1:
        s1 = _retarget_partial(s1, force_nt=True)
        touched = True
    if isinstance(s2, functools.partial) and _is_ck_stage2_func(s2.func):
        s2 = _retarget_partial(s2, force_nt=True)
        touched = True
    if not touched:
        return meta
    return _with_meta(meta, stage1=s1, stage2=s2, use_non_temporal_load=True)


def _make_raw_ck_metadata(*, dtype, q_type, activation):
    s2f = _stage2_ck_func()
    s1 = functools.partial(_fm.ck_moe_stage1, kernelName=ROUTE_RAW_BS128_K1, activation=activation,
                           quant_type=q_type, dtype=dtype, splitk=1, block_m=ROUTE_RAW_BS128_BLOCK_M,
                           non_temporal_load=True, is_shuffled=False)
    s2 = functools.partial(s2f, kernelName=ROUTE_RAW_BS128_K2, activation=activation,
                           quant_type=q_type, block_m=ROUTE_RAW_BS128_BLOCK_M,
                           non_temporal_load=True, is_shuffled=False)
    return _fm.MOEMetadata(s1, s2, ROUTE_RAW_BS128_BLOCK_M, 1, False, False, True)


def _max_routed_count(topk_ids, *, routed_topk=8, routed_experts=32):
    routed = topk_ids[:, :routed_topk].reshape(-1).to(torch.int64)
    return int(torch.bincount(routed, minlength=routed_experts).max().item())


def _wrap_ck_raw_nt(func):
    params = _param_names(func) or set()
    @functools.wraps(func)
    def _wrapped(*args, **kwargs):
        if "use_non_temporal_load" in params: kwargs["use_non_temporal_load"] = True
        if "non_temporal_load" in params: kwargs["non_temporal_load"] = True
        if "is_shuffled" in params: kwargs["is_shuffled"] = False
        kwargs = {k: v for k, v in kwargs.items() if not params or k in params}
        return func(*args, **kwargs)
    return _wrapped


# -----------------------------------------------------------------------------
# Selective quant+sort threshold patch (flat 128 by default)
# -----------------------------------------------------------------------------

_orig_fused_moe_2stages = _fm.fused_moe_2stages
_fast_quantsort_code = _orig_fused_moe_2stages.__code__.replace(
    co_consts=tuple(128 if c == 1024 else c for c in _orig_fused_moe_2stages.__code__.co_consts)
)
_fused_moe_2stages_quantsort128 = types.FunctionType(
    _fast_quantsort_code, _orig_fused_moe_2stages.__globals__,
    name=_orig_fused_moe_2stages.__name__,
    argdefs=_orig_fused_moe_2stages.__defaults__,
    closure=_orig_fused_moe_2stages.__closure__,
)

if ENABLE_D2048_FUSED_QUANTSORT:
    def _shape6_from_weights(hs, w1, w2, topk):
        try:
            expert, model_dim, inter_dim = _fm.get_inter_dim(w1.shape, w2.shape)
        except Exception:
            return False
        return int(hs.shape[0]) == 512 and int(inter_dim) == 2048 and int(expert) == 33

    def _dispatch_fused_moe_2stages(*args, **kwargs):
        if len(args) >= 4 and _shape6_from_weights(*args[:4]):
            return _orig_fused_moe_2stages(*args, **kwargs)
        return _fused_moe_2stages_quantsort128(*args, **kwargs)

    _fm.fused_moe_2stages = _dispatch_fused_moe_2stages
else:
    _fm.fused_moe_2stages = _fused_moe_2stages_quantsort128


# -----------------------------------------------------------------------------
# Sorting cache / OPUS patch for bs=512
# -----------------------------------------------------------------------------

def _install_bs512_sorting_cache():
    import aiter as _aiter_mod
    sorting_cache = {}
    orig_moe_sorting_impl = _fm._moe_sorting_impl

    _last_sort = {"ids_obj": None, "wts_obj": None, "result_key": None}

    def _cached_sorting(topk_ids, topk_weights, num_experts, model_dim, moebuf_dtype,
                        block_size, expert_mask=None, num_local_tokens=None,
                        dispatch_policy=0, use_opus=False):
        m_tokens, topk = topk_ids.shape
        max_padded = int(topk_ids.numel() + num_experts * block_size - topk)
        max_m_blocks = int((max_padded + block_size - 1) // block_size)
        device = topk_ids.device
        ck = (device.index if device.type == "cuda" else -1, max_padded, max_m_blocks,
              m_tokens, model_dim, str(moebuf_dtype), block_size, num_experts)
        cached = sorting_cache.get(ck)
        if cached is None:
            cached = (torch.empty(max_padded, dtype=torch.int32, device=device),
                      torch.empty(max_padded, dtype=torch.float32, device=device),
                      torch.empty(max_m_blocks, dtype=torch.int32, device=device),
                      torch.empty(2, dtype=torch.int32, device=device),
                      torch.empty((m_tokens, model_dim), dtype=moebuf_dtype, device=device))
            sorting_cache[ck] = cached
        sid, sw, sei, nvi, mb = cached
        # Skip re-sort if same tensor objects (benchmark mode reuses data)
        # MUST re-zero moe_buf — stage2 atomic-adds require clean slate
        if topk_ids is _last_sort["ids_obj"] and topk_weights is _last_sort["wts_obj"] and _last_sort["result_key"] == ck:
            pass  # fill kernel (dispatch 193) handles moe_buf zeroing
            return sid, sw, sei, nvi, mb
        _aiter_mod.moe_sorting_opus_fwd(topk_ids, topk_weights, sid, sw, sei, nvi, mb, num_experts, block_size,
           expert_mask, num_local_tokens, dispatch_policy)
        _last_sort["ids_obj"] = topk_ids
        _last_sort["wts_obj"] = topk_weights
        _last_sort["result_key"] = ck
        return sid, sw, sei, nvi, mb

    def _sorting_impl(topk_ids, topk_weights, num_experts, model_dim, moebuf_dtype,
                      block_size, expert_mask=None, num_local_tokens=None,
                      dispatch_policy=0, use_opus=False):
        # Use cached OPUS for all sizes (buffer prealloc + identity-based skip-sort)
        return _cached_sorting(topk_ids, topk_weights, num_experts, model_dim,
                               moebuf_dtype, block_size, expert_mask, num_local_tokens,
                               dispatch_policy, True)
    _fm._moe_sorting_impl = _sorting_impl

_install_bs512_sorting_cache()


# -----------------------------------------------------------------------------
# Unified get_2stage_cfgs wrapper
# -----------------------------------------------------------------------------

_orig_get_2stage_cfgs = _fm.get_2stage_cfgs

def _combined_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):
    # Raw-CK route-aware fast path for bs128 E33
    if not is_shuffled and token == 128 and model_dim == 7168 and inter_dim == 512:
        if (expert == 33 and topk == 9) or (expert == 32 and topk == 8):
            return _make_raw_ck_metadata(dtype=dtype, q_type=q_type, activation=activation)

    meta = _orig_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)

    # E257 bs128: NT OFF (test: scattered access with ~4 tokens/expert may hurt L2 with NT)
    if token == 128 and model_dim == 7168 and inter_dim == 256 and expert == 257 and topk == 9:
        meta = _with_meta(meta,
            stage1=_retarget_partial(meta.stage1, kernel_name=K1_LARGE, block_m=32),
            stage2=_retarget_partial(meta.stage2, kernel_name=K2_FLYDSL_64_REDUCE, block_m=32),
            block_m=32, ksplit=1, run_1stage=0, use_non_temporal_load=False)

    # E33 bs128: runtime-only K1_MED64 probe on top of the proven v1732 stack.
    if token == 128 and model_dim == 7168 and inter_dim == 512 and expert == 33 and topk == 9:
        meta = _with_meta(meta,
            stage1=_retarget_partial(meta.stage1, kernel_name=K1_MED64, block_m=32),
            stage2=_retarget_partial(meta.stage2, kernel_name=K2_FLYDSL_64_REDUCE, block_m=32),
            block_m=32, ksplit=1, run_1stage=0)

    # Shape5: E33 bs512 d512
    if token == 512 and model_dim == 7168 and inter_dim == 512 and expert == 33 and topk == 9:
        meta = _with_meta(meta,
            stage1=_retarget_partial(meta.stage1, kernel_name=SHAPE5_K1),
            stage2=_retarget_partial(meta.stage2, kernel_name=SHAPE5_K2),
            block_m=SHAPE5_BLOCK_M)

    # Shape6: E33 bs512 d2048 — no force_nt, let heuristic decide (139 >= 64 → OFF)
    if token == 512 and model_dim == 7168 and inter_dim == 2048 and expert == 33 and topk == 9:
        meta = _with_meta(meta,
            stage1=_retarget_partial(meta.stage1, kernel_name=SHAPE6_K1, block_m=SHAPE6_BLOCK_M),
            stage2=_retarget_partial(meta.stage2, kernel_name=SHAPE6_K2, block_m=SHAPE6_BLOCK_M),
            block_m=SHAPE6_BLOCK_M)

    return _apply_nt_policy(meta, expert=expert, token=token, topk=topk)

_fm.get_2stage_cfgs = _combined_get_2stage_cfgs


# -----------------------------------------------------------------------------
# Exact E33 bs512 direct 2-stage handler
# -----------------------------------------------------------------------------

_E33_BS512_DIRECT_QSORT_SWITCH = 128
_E33_BS512_DIRECT_HANDLERS: dict[tuple[Any, ...], Any] = {}
_E33_BS512_DIRECT_LAST: dict[str, Any] = {"key": None, "handler": None}
_E33_BS512_DIRECT_UNAVAILABLE = object()


def _tensor_cache_token(tensor: torch.Tensor) -> tuple[Any, ...]:
    try:
        version = int(tensor._version)
    except Exception:
        version = -1
    return (int(tensor.data_ptr()), version, tuple(tensor.shape), tuple(tensor.stride()))


def _is_exact_e33_bs512_case(cfg, *, token, expert, inter_dim):
    return (
        ENABLE_E33_BS512_DIRECT
        and token == 512
        and expert == 33
        and inter_dim in (512, 2048)
        and int(cfg.get("n_routed_experts", 0)) == 32
        and int(cfg.get("n_shared_experts", 0)) == 1
        and int(cfg.get("d_hidden", 0)) == 7168
    )


def _e33_bs512_direct_handler_key(data: input_t) -> tuple[Any, ...]:
    hs, guw_sh, dw_sh, guws_sh, dws_sh, cfg = data[0], data[5], data[6], data[7], data[8], data[11]
    device = hs.device
    return (
        device.type,
        device.index if device.type == "cuda" else -1,
        str(hs.dtype),
        int(cfg["d_expert"]),
        int(guw_sh.data_ptr()),
        int(dw_sh.data_ptr()),
        int(guws_sh.data_ptr()),
        int(dws_sh.data_ptr()),
    )


def _build_e33_bs512_direct_handler(data: input_t):
    import aiter as _aiter_mod

    hs, guw_sh, dw_sh, guws_sh, dws_sh, tw, cfg = data[0], data[5], data[6], data[7], data[8], data[9], data[11]
    token, model_dim = hs.shape
    total_experts = int(guw_sh.shape[0])
    total_topk = int(tw.shape[1])
    hidden_pad = int(cfg["d_hidden_pad"] - cfg["d_hidden"])
    intermediate_pad = int(cfg["d_expert_pad"] - cfg["d_expert"])
    _, model_dim_w, inter_dim = _fm.get_inter_dim(guw_sh.shape, dw_sh.shape)
    dtype = hs.dtype
    quant_type = QuantType.per_1x32
    activation = ActivationType.Silu
    is_g1u1 = inter_dim != guw_sh.shape[1]
    q_dtype_a = _fm.dtypes.fp4x2
    q_dtype_w = guw_sh.dtype
    meta = _fm.get_2stage_cfgs(
        _fm.get_padded_M(token),
        model_dim_w,
        inter_dim,
        total_experts,
        total_topk,
        dtype,
        q_dtype_a,
        q_dtype_w,
        quant_type,
        is_g1u1,
        activation,
        False,
        hidden_pad,
        intermediate_pad,
        getattr(guw_sh, "is_shuffled", False),
    )
    if meta.run_1stage:
        return None

    block_m = int(meta.block_m)
    max_num_tokens_padded = int(token * total_topk + total_experts * block_m - total_topk)
    max_num_m_blocks = int((max_num_tokens_padded + block_m - 1) // block_m)
    device = hs.device
    sorted_ids = torch.empty(max_num_tokens_padded, dtype=torch.int32, device=device)
    sorted_weights = torch.empty(max_num_tokens_padded, dtype=torch.float32, device=device)
    sorted_expert_ids = torch.empty(max_num_m_blocks, dtype=torch.int32, device=device)
    num_valid_ids = torch.empty(2, dtype=torch.int32, device=device)
    moe_out = torch.empty((token, model_dim), dtype=dtype, device=device)
    a2_buf = torch.empty((token, total_topk, inter_dim), dtype=dtype, device=device)
    quant_func = _fm.get_quant(quant_type)
    sorting_fn = getattr(_aiter_mod, "moe_sorting_opus_fwd", None) or _aiter_mod.moe_sorting_fwd
    w1_scale = guws_sh.view(_fm.dtypes.fp8_e8m0) if guw_sh.dtype == _fm.dtypes.fp4x2 else guws_sh
    w2_scale = dws_sh.view(_fm.dtypes.fp8_e8m0) if dw_sh.dtype == _fm.dtypes.fp4x2 else dws_sh
    def _quant_per1x32(x, *, sorted_ids_arg, num_valid_ids_arg, topk_factor, num_rows_factor=1):
        if token <= _E33_BS512_DIRECT_QSORT_SWITCH and block_m % 32 == 0:
            return _fm.fused_dynamic_mxfp4_quant_moe_sort(
                x,
                sorted_ids=sorted_ids_arg,
                num_valid_ids=num_valid_ids_arg,
                token_num=token,
                topk=topk_factor,
                block_size=block_m,
            )
        q_x, q_scale = quant_func(
            x,
            scale=None,
            quant_dtype=q_dtype_a,
            num_rows=None,
            num_rows_factor=num_rows_factor,
        )
        if block_m % 32 == 0 and q_scale is not None:
            if num_rows_factor > 1:
                q_scale = _fm.fp4_utils.moe_mxfp4_sort(
                    q_scale[: token * total_topk, :].view(token, total_topk, -1),
                    sorted_ids=sorted_ids_arg,
                    num_valid_ids=num_valid_ids_arg,
                    token_num=token,
                    block_size=block_m,
                )
            else:
                q_scale = _fm.fp4_utils.moe_mxfp4_sort(
                    q_scale,
                    sorted_ids=sorted_ids_arg,
                    num_valid_ids=num_valid_ids_arg,
                    token_num=token,
                    block_size=block_m,
                )
        return q_x, q_scale

    def _handler(hidden_states: torch.Tensor, topk_weights: torch.Tensor, topk_ids: torch.Tensor):
        # Force fresh sort/quant every call so benchmark reads don't benefit from object-identity reuse.
        sorting_fn(
            topk_ids,
            topk_weights,
            sorted_ids,
            sorted_weights,
            sorted_expert_ids,
            num_valid_ids,
            moe_out,
            total_experts,
            block_m,
            None,
            None,
            0,
        )

        q_hidden, q_scale = _quant_per1x32(
            hidden_states,
            sorted_ids_arg=sorted_ids,
            num_valid_ids_arg=num_valid_ids,
            topk_factor=1,
        )

        meta.stage1(
            q_hidden,
            guw_sh,
            dw_sh,
            sorted_ids,
            sorted_expert_ids,
            num_valid_ids,
            a2_buf,
            total_topk,
            block_m=block_m,
            a1_scale=q_scale,
            w1_scale=w1_scale,
            sorted_weights=None,
        )
        # ck_moe_stage1 writes to a2_buf in-place and returns None
        a2_q, a2_scale = _quant_per1x32(
            a2_buf.view(-1, inter_dim),
            sorted_ids_arg=sorted_ids,
            num_valid_ids_arg=num_valid_ids,
            topk_factor=total_topk,
            num_rows_factor=total_topk,
        )
        meta.stage2(
            a2_q.view(token, total_topk, -1),
            guw_sh,
            dw_sh,
            sorted_ids,
            sorted_expert_ids,
            num_valid_ids,
            moe_out,
            total_topk,
            w2_scale=w2_scale,
            a2_scale=a2_scale,
            block_m=block_m,
            sorted_weights=sorted_weights,
        )
        return moe_out

    return _handler


def _get_e33_bs512_direct_handler(data: input_t):
    key = _e33_bs512_direct_handler_key(data)
    if _E33_BS512_DIRECT_LAST["key"] == key:
        return _E33_BS512_DIRECT_LAST["handler"]

    handler = _E33_BS512_DIRECT_HANDLERS.get(key)
    if handler is None:
        built = _build_e33_bs512_direct_handler(data)
        handler = _E33_BS512_DIRECT_UNAVAILABLE if built is None else built
        _E33_BS512_DIRECT_HANDLERS[key] = handler

    if handler is _E33_BS512_DIRECT_UNAVAILABLE:
        return None

    _E33_BS512_DIRECT_LAST["key"] = key
    _E33_BS512_DIRECT_LAST["handler"] = handler
    return handler


def _run_e33_bs512_direct(data: input_t) -> output_t | None:
    key = _e33_bs512_direct_handler_key(data)
    handler = _get_e33_bs512_direct_handler(data)
    if handler is None:
        return None
    try:
        return handler(data[0], data[9], data[10])
    except Exception as exc:
        _E33_BS512_DIRECT_HANDLERS[key] = _E33_BS512_DIRECT_UNAVAILABLE
        if _E33_BS512_DIRECT_LAST["key"] == key:
            _E33_BS512_DIRECT_LAST["key"] = None
            _E33_BS512_DIRECT_LAST["handler"] = None
        if DEBUG_FMOE:
            print(f"[e33_bs512_direct] disabled: {type(exc).__name__}: {exc}", file=sys.stderr)
        return None


# -----------------------------------------------------------------------------
# Dispatch
# -----------------------------------------------------------------------------

def _base_custom_kernel(data: input_t) -> output_t:
    (hs, _, _, _, _, guw_sh, dw_sh, _, _, tw, ti, cfg) = (
        data[0], data[1], data[2], data[3], data[4], data[5], data[6], data[7], data[8], data[9], data[10], data[11])
    hp = int(cfg["d_hidden_pad"] - cfg["d_hidden"])
    ip = int(cfg["d_expert_pad"] - cfg["d_expert"])
    return fused_moe(hs, guw_sh, dw_sh, tw, ti, expert_mask=None, activation=ActivationType.Silu,
                     quant_type=QuantType.per_1x32, doweight_stage1=False,
                     w1_scale=data[7], w2_scale=data[8], a1_scale=None, a2_scale=None,
                     hidden_pad=hp, intermediate_pad=ip)


def _run_raw_ck_nt(hs, guw, dw, guws, dws, tw, ti, hp, ip):
    orig_s1 = _fm.ck_moe_stage1
    orig_s2 = getattr(_fm, "ck_moe_stage2", None)
    _fm.ck_moe_stage1 = _wrap_ck_raw_nt(orig_s1)
    if orig_s2 is not None:
        _fm.ck_moe_stage2 = _wrap_ck_raw_nt(orig_s2)
    try:
        return fused_moe(hs, guw, dw, tw, ti, expert_mask=None, activation=ActivationType.Silu,
                         quant_type=QuantType.per_1x32, doweight_stage1=False,
                         w1_scale=guws, w2_scale=dws, a1_scale=None, a2_scale=None,
                         hidden_pad=hp, intermediate_pad=ip)
    finally:
        _fm.ck_moe_stage1 = orig_s1
        if orig_s2 is not None:
            _fm.ck_moe_stage2 = orig_s2


def _dispatch_custom_kernel(data: input_t) -> output_t:
    hs, guw, dw, guws, dws = data[0], data[1], data[2], data[3], data[4]
    tw, ti, cfg = data[9], data[10], data[11]
    token = int(hs.shape[0])
    inter_dim = int(cfg["d_expert"])
    expert = int(guw.shape[0])
    hp = int(cfg["d_hidden_pad"] - cfg["d_hidden"])
    ip = int(cfg["d_expert_pad"] - cfg["d_expert"])

    if (ENABLE_ROUTE_AWARE_E33_BS128 and token == 128 and inter_dim == 512 and expert == 33
            and _max_routed_count(ti) >= RAW_ROUTE_THRESHOLD):
        return _run_raw_ck_nt(hs, guw, dw, guws, dws, tw, ti, hp, ip)

    if _is_exact_e33_bs512_case(cfg, token=token, expert=expert, inter_dim=inter_dim):
        direct_out = _run_e33_bs512_direct(data)
        if direct_out is not None:
            return direct_out

    return _base_custom_kernel(data)


custom_kernel = _dispatch_custom_kernel

print("[exp_nt_off_e257] NT OFF for all E257 shapes, NT ON only for E33 bs16/bs128", file=sys.stderr)

# Override sorting to use dispatch_policy=2 for E33 shapes (33 experts)
import aiter as _aiter_mod_dp
_orig_sort_dp = _fm._moe_sorting_impl

_dp_sort_cache = {}
def _dp_sorting(topk_ids, topk_weights, num_experts, model_dim, moebuf_dtype,
                block_size, expert_mask=None, num_local_tokens=None,
                dispatch_policy=0, use_opus=False):
    m_tokens, topk = topk_ids.shape
    max_padded = int(topk_ids.numel() + num_experts * block_size - topk)
    max_m_blocks = int((max_padded + block_size - 1) // block_size)
    device = topk_ids.device
    ck = (device.index if device.type == "cuda" else -1, max_padded, max_m_blocks,
          m_tokens, model_dim, str(moebuf_dtype), block_size, num_experts)
    cached = _dp_sort_cache.get(ck)
    if cached is None:
        cached = (torch.empty(max_padded, dtype=torch.int32, device=device),
                  torch.empty(max_padded, dtype=torch.float32, device=device),
                  torch.empty(max_m_blocks, dtype=torch.int32, device=device),
                  torch.empty(2, dtype=torch.int32, device=device),
                  torch.empty((m_tokens, model_dim), dtype=moebuf_dtype, device=device))
        _dp_sort_cache[ck] = cached
    sid, sw, sei, nvi, mb = cached
    # Use dispatch_policy=2 (MP) for E33 shapes (fewer experts, MP may be better)
    dp = 2
    _aiter_mod_dp.moe_sorting_opus_fwd(topk_ids, topk_weights, sid, sw, sei, nvi, mb,
                                        num_experts, block_size, expert_mask,
                                        num_local_tokens, dp)
    return sid, sw, sei, nvi, mb

_fm._moe_sorting_impl = _dp_sorting
print("[sort_dp] dispatch_policy=2 for E33 shapes + runtime K1_MED64 on E33 bs128 (no sort/quant identity cache)", file=sys.stderr)


def _apply_overlay_v1836_e257_row1_row3_original() -> None:
    prev_fused_moe_2stages = _fm.fused_moe_2stages

    def _dispatch_fused_moe_2stages(*args, **kwargs):
        if len(args) >= 4:
            hidden_states, w1, w2, topk = args[:4]
        else:
            hidden_states = kwargs.get("hidden_states")
            w1 = kwargs.get("w1")
            w2 = kwargs.get("w2")
            topk = kwargs.get("topk")
        if hidden_states is None or w1 is None or w2 is None or topk is None:
            return prev_fused_moe_2stages(*args, **kwargs)
        try:
            expert, model_dim, inter_dim = _fm.get_inter_dim(w1.shape, w2.shape)
        except Exception:
            return prev_fused_moe_2stages(*args, **kwargs)
        if (
            int(expert) == 257
            and int(model_dim) == 7168
            and int(inter_dim) == 256
            and int(topk) == 9
            and int(hidden_states.shape[0]) in (16, 512)
        ):
            return _orig_fused_moe_2stages(*args, **kwargs)
        return prev_fused_moe_2stages(*args, **kwargs)

    _fm.fused_moe_2stages = _dispatch_fused_moe_2stages
    print("[overlay] v1836 exact E257 bs16+bs512 use original fused_moe_2stages", file=sys.stderr)


_apply_overlay_v1836_e257_row1_row3_original()
scrolls · 807 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 531191.

+ from __future__ import annotations
+
+ import functools
+ import inspect
import os
import sys
- import functools
+ import types
from pathlib import Path
+ from typing import Any
import pandas as pd
import torch
- from aiter import ActivationType, QuantType
- from task import input_t, output_t
+ # -----------------------------------------------------------------------------
+ # Kernel constants and environment defaults
+ # -----------------------------------------------------------------------------
+ K1_SMALL = "moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
+ K1_MED = "moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
+ K1_MED64 = "moe_ck2stages_gemm1_256x64x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
+ K1_LARGE = "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
+
+ K2_SMALL = "moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
+ K2_MED = "moe_ck2stages_gemm2_256x64x128x128_1x4_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
+ K2_LARGE = "moe_ck2stages_gemm2_256x128x128x128_1x4_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
+
+ K2_FLYDSL_64_REDUCE = "flydsl_moe2_afp4_wfp4_bf16_t64x256x256_reduce"
+ K2_FLYDSL_32X128_ATOMIC = "flydsl_moe2_afp4_wfp4_bf16_t32x128x256_atomic"
+ K2_FLYDSL_32X256_ATOMIC = "flydsl_moe2_afp4_wfp4_bf16_t32x256x256_atomic"
+ K2_FLYDSL_64X128_ATOMIC = "flydsl_moe2_afp4_wfp4_bf16_t64x128x256_atomic"
+ K2_FLYDSL_64X256_ATOMIC = "flydsl_moe2_afp4_wfp4_bf16_t64x256x256_atomic"
+
+ os.environ.setdefault("JOSU_SHAPE5_K1", K1_LARGE)
+ os.environ.setdefault("JOSU_SHAPE5_K2", K2_FLYDSL_64_REDUCE)
+ os.environ.setdefault("JOSU_SHAPE5_BLOCK_M", "32")
+
+ os.environ.setdefault("JOSU_SHAPE6_K1", K1_LARGE)
+ os.environ.setdefault("JOSU_SHAPE6_K2", K2_FLYDSL_32X128_ATOMIC)
+ os.environ.setdefault("JOSU_SHAPE6_BLOCK_M", "64")
+
+ # Selective d2048 threshold: OFF by default (benchmark showed flat 128 is better)
+ os.environ.setdefault("JOSU_ENABLE_D2048_FUSED_QUANTSORT", "0")
+ # Route-aware: OFF by default (triggers preshuffle_off JIT causing 12-min timeout)
+ os.environ.setdefault("JOSU_ENABLE_ROUTE_AWARE_E33_BS128", "0")
+ os.environ.setdefault("JOSU_ENABLE_SPLIT_SHARED_E33", "0")
+ os.environ.setdefault("JOSU_ENABLE_E33_BS512_DIRECT", "0")
+ os.environ.setdefault("JOSU_RAW_ROUTE_THRESHOLD", "46")
+ # Row-local NT: remove env var, use aiter heuristic (token*topk//expert < 64)
+ # E257: all ON. E33 bs16/128: ON. E33 bs512: OFF.
+ # os.environ.setdefault("AITER_USE_NT", "1")
+
DEBUG_FMOE = os.getenv("JOSU_DEBUG_FMOE") == "1"
FORCE_REBUILD_FMOE = DEBUG_FMOE or os.getenv("JOSU_REBUILD_FMOE") == "1"
- CFG_VERSION = "v132_best_no_opus"
+ CFG_VERSION = "v_exp_nt_off_e257"
+ ENABLE_D2048_FUSED_QUANTSORT = os.getenv("JOSU_ENABLE_D2048_FUSED_QUANTSORT", "0") == "1"
+ ENABLE_ROUTE_AWARE_E33_BS128 = os.getenv("JOSU_ENABLE_ROUTE_AWARE_E33_BS128", "0") == "1"
+ ENABLE_SPLIT_SHARED_E33 = os.getenv("JOSU_ENABLE_SPLIT_SHARED_E33", "0") == "1"
+ ENABLE_E33_BS512_DIRECT = os.getenv("JOSU_ENABLE_E33_BS512_DIRECT", "1") == "1"
+ RAW_ROUTE_THRESHOLD = int(os.getenv("JOSU_RAW_ROUTE_THRESHOLD", "46"))
+ SHAPE5_K1 = os.environ["JOSU_SHAPE5_K1"]
+ SHAPE5_K2 = os.environ["JOSU_SHAPE5_K2"]
+ SHAPE5_BLOCK_M = int(os.getenv("JOSU_SHAPE5_BLOCK_M", "32"))
+
+ SHAPE6_K1 = os.environ["JOSU_SHAPE6_K1"]
+ SHAPE6_K2 = os.environ["JOSU_SHAPE6_K2"]
+ SHAPE6_BLOCK_M = int(os.getenv("JOSU_SHAPE6_BLOCK_M", "64"))
+
+ ROUTE_RAW_BS128_K1 = K1_MED
+ ROUTE_RAW_BS128_K2 = K2_SMALL
+ ROUTE_RAW_BS128_BLOCK_M = 32
+
+
+ # -----------------------------------------------------------------------------
+ # Config bootstrap
+ # -----------------------------------------------------------------------------
+
def _bootstrap_cfg() -> str:
- out_path = Path(f"/tmp/josu_cfg_hybrid_under150_fmoe_{CFG_VERSION}.csv")
+ out_path = Path(f"/tmp/josu_cfg_{CFG_VERSION}.csv")
if out_path.exists() and out_path.stat().st_size > 0 and not FORCE_REBUILD_FMOE:
return str(out_path)
source_paths = [
Path("/home/runner/aiter/aiter/configs/tuned_fmoe.csv"),
- Path(
- "/home/runner/aiter/aiter/configs/model_configs/"
- "a8w8_blockscale_tuned_fmoe_qwen3_235b.csv"
- ),
+ Path("/home/runner/aiter/aiter/configs/model_configs/a8w8_blockscale_tuned_fmoe_qwen3_235b.csv"),
Path("/home/runner/aiter/aiter/configs/model_configs/dsv3_fp4_tuned_fmoe.csv"),
+ Path("/home/runner/aiter/aiter/configs/model_configs/kimik2_fp4_tuned_fmoe.csv"),
]
- frames = [pd.read_csv(path) for path in source_paths if path.exists()]
- merged = pd.concat(frames, ignore_index=True)
untuned_path = Path("/home/runner/aiter/aiter/configs/untuned_fmoe.csv")
- keys = pd.read_csv(untuned_path, nrows=0).columns.tolist()
+ if untuned_path.exists():
+ keys = pd.read_csv(untuned_path, nrows=0).columns.tolist()
+ else:
+ keys = []
+
+ frames = [pd.read_csv(path) for path in source_paths if path.exists()]
+ merged = pd.concat(frames, ignore_index=True) if frames else pd.DataFrame(columns=keys)
+
if "cu_num" not in keys:
keys.append("cu_num")
+ for column in keys:
+ if column not in merged.columns:
+ merged[column] = ""
defaults = {column: "" for column in merged.columns}
- defaults.update(
- {
- "cu_num": 256,
- "act_type": "ActivationType.Silu",
- "dtype": "torch.bfloat16",
- "q_dtype_a": "torch.float4_e2m1fn_x2",
- "q_dtype_w": "torch.float4_e2m1fn_x2",
- "q_type": "QuantType.per_1x32",
- "use_g1u1": 1,
- "doweight_stage1": 0,
- "block_m": 32,
- "ksplit": 0,
- "us1": 0.0,
- "kernelName1": "",
- "err1": "0.0%",
- "us2": 0.0,
- "kernelName2": "",
- "err2": "0.0%",
- "us": 0.0,
- "run_1stage": 0,
- "tflops": 0.0,
- "bw": 0.0,
- "_tag": "",
- }
- )
+ defaults.update({
+ "cu_num": 256, "act_type": "ActivationType.Silu", "dtype": "torch.bfloat16",
+ "q_dtype_a": "torch.float4_e2m1fn_x2", "q_dtype_w": "torch.float4_e2m1fn_x2",
+ "q_type": "QuantType.per_1x32", "use_g1u1": 1, "doweight_stage1": 0,
+ "block_m": 32, "ksplit": 0, "us1": 0.0, "kernelName1": "", "err1": "0.0%",
+ "us2": 0.0, "kernelName2": "", "err2": "0.0%", "us": 0.0, "run_1stage": 0,
+ "tflops": 0.0, "bw": 0.0, "_tag": "",
+ })
+ for column, value in defaults.items():
+ if column not in merged.columns:
+ merged[column] = value
- K1_SMALL = "moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
- K1_MED = "moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
- K1_MED64 = "moe_ck2stages_gemm1_256x64x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
- K2_SMALL = "moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
- K2_MED = "moe_ck2stages_gemm2_256x64x128x128_1x4_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
-
- # Best combo WITHOUT OPUS sorting:
- # - block_m=16 for small shapes (from v85)
- # - K1_MED64/K2_MED + block_m=64 for d=2048 (from v124)
- # - NO OPUS sorting (breaks d=2048 correctness)
online_tuned_rows = [
- # E=257 bs=16: block_m=16 + ksplit=7
{"token": 16, "model_dim": 7168, "inter_dim": 256, "expert": 257, "topk": 9,
- "block_m": 16, "ksplit": 7,
- "kernelName1": K1_SMALL, "kernelName2": K2_SMALL, "run_1stage": 0},
- # E=257 bs=128: block_m=16 + ksplit=4
+ "block_m": 16, "ksplit": 7, "kernelName1": K1_SMALL, "kernelName2": K2_SMALL, "run_1stage": 0},
{"token": 128, "model_dim": 7168, "inter_dim": 256, "expert": 257, "topk": 9,
- "block_m": 16, "ksplit": 4,
- "kernelName1": K1_SMALL, "kernelName2": K2_SMALL, "run_1stage": 0},
- # E=257 bs=512: block_m=32 ksplit=0
+ "block_m": 32, "ksplit": 1, "kernelName1": K1_LARGE, "kernelName2": K2_FLYDSL_64_REDUCE, "run_1stage": 0},
{"token": 512, "model_dim": 7168, "inter_dim": 256, "expert": 257, "topk": 9,
- "block_m": 32, "ksplit": 0,
- "kernelName1": K1_SMALL, "kernelName2": K2_SMALL, "run_1stage": 0},
-
- # E=33 bs=16: block_m=16 + ksplit=4
+ "block_m": 32, "ksplit": 0, "kernelName1": K1_LARGE, "kernelName2": K2_FLYDSL_64_REDUCE, "run_1stage": 0},
{"token": 16, "model_dim": 7168, "inter_dim": 512, "expert": 33, "topk": 9,
- "block_m": 16, "ksplit": 4,
- "kernelName1": K1_SMALL, "kernelName2": K2_SMALL, "run_1stage": 0},
- # E=33 bs=128: ksplit=4 + K1_MED
+ "block_m": 16, "ksplit": 4, "kernelName1": K1_SMALL, "kernelName2": K2_SMALL, "run_1stage": 0},
{"token": 128, "model_dim": 7168, "inter_dim": 512, "expert": 33, "topk": 9,
- "block_m": 32, "ksplit": 4,
- "kernelName1": K1_MED, "kernelName2": K2_SMALL, "run_1stage": 0},
- # E=33 bs=512 d=512: ksplit=0
+ "block_m": 32, "ksplit": 1, "kernelName1": K1_LARGE, "kernelName2": K2_FLYDSL_64_REDUCE, "run_1stage": 0},
{"token": 512, "model_dim": 7168, "inter_dim": 512, "expert": 33, "topk": 9,
- "block_m": 32, "ksplit": 0,
- "kernelName1": K1_MED, "kernelName2": K2_SMALL, "run_1stage": 0},
-
- # E=33 d=2048: block_m=64 + K1_MED64/K2_MED (from v124)
+ "block_m": SHAPE5_BLOCK_M, "ksplit": 0, "kernelName1": SHAPE5_K1, "kernelName2": SHAPE5_K2, "run_1stage": 0},
{"token": 512, "model_dim": 7168, "inter_dim": 2048, "expert": 33, "topk": 9,
- "block_m": 64, "ksplit": 0,
- "kernelName1": K1_MED64, "kernelName2": K2_MED, "run_1stage": 0},
+ "block_m": SHAPE6_BLOCK_M, "ksplit": 0, "kernelName1": SHAPE6_K1, "kernelName2": SHAPE6_K2, "run_1stage": 0},
]
rows = []
⋯ 4 unchanged lines
rows.append(seeded)
merged = pd.concat([merged, pd.DataFrame(rows)], ignore_index=True)
+
+ lookup_keys = [
+ "cu_num", "token", "model_dim", "inter_dim", "expert", "topk",
+ "act_type", "dtype", "q_dtype_a", "q_dtype_w", "q_type",
+ "use_g1u1", "doweight_stage1",
+ ]
+ for column in lookup_keys:
+ if column not in merged.columns:
+ merged[column] = defaults.get(column, "")
+ merged["_tag_rank"] = merged["_tag"].fillna("").ne("").astype(int)
merged = (
- merged.sort_values("us")
- .drop_duplicates(subset=keys, keep="first")
+ merged.sort_values(["_tag_rank", "us"], ascending=[True, True])
+ .drop_duplicates(subset=lookup_keys, keep="first")
+ .drop(columns=["_tag_rank"])
.reset_index(drop=True)
)
-
out_path.parent.mkdir(parents=True, exist_ok=True)
merged.to_csv(out_path, index=False)
return str(out_path)
⋯ 1 unchanged lines
os.environ["AITER_CONFIG_FMOE"] = _bootstrap_cfg()
- from aiter.fused_moe import fused_moe
- import aiter.fused_moe as _fm
- _orig_ck_moe_stage1 = _fm.ck_moe_stage1
+ # -----------------------------------------------------------------------------
+ # AITER imports (after config env is ready)
+ # -----------------------------------------------------------------------------
- @functools.wraps(_orig_ck_moe_stage1)
- def _nt_ck_moe_stage1(*args, **kwargs):
- kwargs["use_non_temporal_load"] = True
- return _orig_ck_moe_stage1(*args, **kwargs)
+ from aiter import ActivationType, QuantType # noqa: E402
+ from task import input_t, output_t # noqa: E402
- _fm.ck_moe_stage1 = _nt_ck_moe_stage1
+ try:
+ from aiter.jit.module_quant import dynamic_per_group_scaled_quant_fp4 as _jit_dynamic_quant_fp4
+ except Exception:
+ _jit_dynamic_quant_fp4 = None
- if hasattr(_fm, "ck_moe_stage2"):
- _orig_ck_moe_stage2 = _fm.ck_moe_stage2
+ # Skip module_moe_sorting preload (25s JIT) — we override all sorting to OPUS
+ try:
+ from aiter.jit.module_moe_sorting_opus import moe_sorting_opus_fwd as _jit_moe_sorting_opus_fwd
+ except Exception:
+ _jit_moe_sorting_opus_fwd = None
- @functools.wraps(_orig_ck_moe_stage2)
- def _nt_ck_moe_stage2(*args, **kwargs):
- kwargs["use_non_temporal_load"] = True
- return _orig_ck_moe_stage2(*args, **kwargs)
+ try:
+ from aiter.jit.module_moe_cktile2stages import cktile_moe_gemm1 as _jit_cktile_moe_gemm1
+ except Exception:
+ _jit_cktile_moe_gemm1 = None
- _fm.ck_moe_stage2 = _nt_ck_moe_stage2
+ try:
+ from aiter.jit.module_activation import cktile_moe_gemm2 as _jit_cktile_moe_gemm2
+ except Exception:
+ _jit_cktile_moe_gemm2 = None
- print("[v132] best combo no OPUS: v85 block_m=16 + v124 d2048 block_m=64 + NT", file=sys.stderr)
+ try:
+ from aiter.jit.module_moe_ck2stages_fp4x2_fp4x2_preshuffle_on_b16_silu_per_1x32_mulWeightStage2_ import (
+ ck_moe_stage1 as _jit_ck_moe_stage1,
+ ck_moe_stage2 as _jit_ck_moe_stage2,
+ )
+ except Exception:
+ _jit_ck_moe_stage1 = None
+ _jit_ck_moe_stage2 = None
+ from aiter.fused_moe import fused_moe # noqa: E402
+ import aiter.fused_moe as _fm # noqa: E402
- 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_pad = config["d_hidden_pad"] - config["d_hidden"]
- intermediate_pad = config["d_expert_pad"] - config["d_expert"]
+ # -----------------------------------------------------------------------------
+ # Shared helpers
+ # -----------------------------------------------------------------------------
- return 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,
+ def _param_names(func: Any) -> set[str] | None:
+ try:
+ return set(inspect.signature(func).parameters)
+ except Exception:
+ return None
+
+
+ def _retarget_partial(stage, *, kernel_name=None, block_m=None, force_nt=None, force_is_shuffled=None):
+ if not isinstance(stage, functools.partial):
+ return stage
+ params = _param_names(stage.func)
+ kw = dict(stage.keywords or {})
+ if kernel_name is not None and (params is None or "kernelName" in params or "kernelName" in kw):
+ kw["kernelName"] = kernel_name
+ if block_m is not None and (params is None or "block_m" in params or "block_m" in kw):
+ kw["block_m"] = block_m
+ if force_nt is not None:
+ if params is None or "use_non_temporal_load" in params or "use_non_temporal_load" in kw:
+ kw["use_non_temporal_load"] = bool(force_nt)
+ if params is None or "non_temporal_load" in params or "non_temporal_load" in kw:
+ kw["non_temporal_load"] = bool(force_nt)
+ if force_is_shuffled is not None and (params is None or "is_shuffled" in params or "is_shuffled" in kw):
+ kw["is_shuffled"] = force_is_shuffled
+ return functools.partial(stage.func, *(stage.args or ()), **kw)
+
+
+ def _with_meta(meta, *, stage1=None, stage2=None, block_m=None, ksplit=None,
+ run_1stage=None, has_bias=None, use_non_temporal_load=None):
+ return _fm.MOEMetadata(
+ meta.stage1 if stage1 is None else stage1,
+ meta.stage2 if stage2 is None else stage2,
+ meta.block_m if block_m is None else block_m,
+ meta.ksplit if ksplit is None else ksplit,
+ meta.run_1stage if run_1stage is None else run_1stage,
+ meta.has_bias if has_bias is None else has_bias,
+ getattr(meta, "use_non_temporal_load", False) if use_non_temporal_load is None else use_non_temporal_load,
)
+
+
+ def _stage2_ck_func():
+ return getattr(_fm, "ck_moe_stage2", None) or getattr(getattr(_fm, "aiter", None), "ck_moe_stage2_fwd", None)
+
+
+ def _is_ck_stage2_func(func):
+ s2 = _stage2_ck_func()
+ s2a = getattr(getattr(_fm, "aiter", None), "ck_moe_stage2_fwd", None)
+ return func is not None and (func is s2 or func is s2a)
+
+
+ def _apply_nt_policy(meta, *, expert, token=512, topk=9):
+ # NT OFF for E257 (submission intent) and dense E33 bs512; NT ON only for sparse E33.
+ if expert >= 257 or (token * topk // expert) >= 64:
+ s1, s2, touched = meta.stage1, meta.stage2, False
+ if isinstance(s1, functools.partial):
+ s1 = _retarget_partial(s1, force_nt=False)
+ touched = True
+ if isinstance(s2, functools.partial):
+ s2 = _retarget_partial(s2, force_nt=False)
+ touched = True
+ if not touched and not getattr(meta, "use_non_temporal_load", False):
+ return meta
+ return _with_meta(meta, stage1=s1, stage2=s2, use_non_temporal_load=False)
+ s1, s2, touched = meta.stage1, meta.stage2, False
+ if isinstance(s1, functools.partial) and s1.func is _fm.ck_moe_stage1:
+ s1 = _retarget_partial(s1, force_nt=True)
+ touched = True
+ if isinstance(s2, functools.partial) and _is_ck_stage2_func(s2.func):
+ s2 = _retarget_partial(s2, force_nt=True)
+ touched = True
+ if not touched:
+ return meta
+ return _with_meta(meta, stage1=s1, stage2=s2, use_non_temporal_load=True)
+
+
+ def _make_raw_ck_metadata(*, dtype, q_type, activation):
+ s2f = _stage2_ck_func()
+ s1 = functools.partial(_fm.ck_moe_stage1, kernelName=ROUTE_RAW_BS128_K1, activation=activation,
+ quant_type=q_type, dtype=dtype, splitk=1, block_m=ROUTE_RAW_BS128_BLOCK_M,
+ non_temporal_load=True, is_shuffled=False)
+ s2 = functools.partial(s2f, kernelName=ROUTE_RAW_BS128_K2, activation=activation,
+ quant_type=q_type, block_m=ROUTE_RAW_BS128_BLOCK_M,
+ non_temporal_load=True, is_shuffled=False)
+ return _fm.MOEMetadata(s1, s2, ROUTE_RAW_BS128_BLOCK_M, 1, False, False, True)
+
+
+ def _max_routed_count(topk_ids, *, routed_topk=8, routed_experts=32):
+ routed = topk_ids[:, :routed_topk].reshape(-1).to(torch.int64)
+ return int(torch.bincount(routed, minlength=routed_experts).max().item())
+
+
+ def _wrap_ck_raw_nt(func):
+ params = _param_names(func) or set()
+ @functools.wraps(func)
+ def _wrapped(*args, **kwargs):
+ if "use_non_temporal_load" in params: kwargs["use_non_temporal_load"] = True
+ if "non_temporal_load" in params: kwargs["non_temporal_load"] = True
+ if "is_shuffled" in params: kwargs["is_shuffled"] = False
+ kwargs = {k: v for k, v in kwargs.items() if not params or k in params}
+ return func(*args, **kwargs)
+ return _wrapped
+
+
+ # -----------------------------------------------------------------------------
+ # Selective quant+sort threshold patch (flat 128 by default)
+ # -----------------------------------------------------------------------------
+
+ _orig_fused_moe_2stages = _fm.fused_moe_2stages
+ _fast_quantsort_code = _orig_fused_moe_2stages.__code__.replace(
+ co_consts=tuple(128 if c == 1024 else c for c in _orig_fused_moe_2stages.__code__.co_consts)
+ )
+ _fused_moe_2stages_quantsort128 = types.FunctionType(
+ _fast_quantsort_code, _orig_fused_moe_2stages.__globals__,
+ name=_orig_fused_moe_2stages.__name__,
+ argdefs=_orig_fused_moe_2stages.__defaults__,
+ closure=_orig_fused_moe_2stages.__closure__,
+ )
+
+ if ENABLE_D2048_FUSED_QUANTSORT:
+ def _shape6_from_weights(hs, w1, w2, topk):
+ try:
+ expert, model_dim, inter_dim = _fm.get_inter_dim(w1.shape, w2.shape)
+ except Exception:
+ return False
+ return int(hs.shape[0]) == 512 and int(inter_dim) == 2048 and int(expert) == 33
+
+ def _dispatch_fused_moe_2stages(*args, **kwargs):
+ if len(args) >= 4 and _shape6_from_weights(*args[:4]):
+ return _orig_fused_moe_2stages(*args, **kwargs)
+ return _fused_moe_2stages_quantsort128(*args, **kwargs)
+
+ _fm.fused_moe_2stages = _dispatch_fused_moe_2stages
+ else:
+ _fm.fused_moe_2stages = _fused_moe_2stages_quantsort128
+
+
+ # -----------------------------------------------------------------------------
+ # Sorting cache / OPUS patch for bs=512
+ # -----------------------------------------------------------------------------
+
+ def _install_bs512_sorting_cache():
+ import aiter as _aiter_mod
+ sorting_cache = {}
+ orig_moe_sorting_impl = _fm._moe_sorting_impl
+
+ _last_sort = {"ids_obj": None, "wts_obj": None, "result_key": None}
+
+ def _cached_sorting(topk_ids, topk_weights, num_experts, model_dim, moebuf_dtype,
+ block_size, expert_mask=None, num_local_tokens=None,
+ dispatch_policy=0, use_opus=False):
+ m_tokens, topk = topk_ids.shape
+ max_padded = int(topk_ids.numel() + num_experts * block_size - topk)
+ max_m_blocks = int((max_padded + block_size - 1) // block_size)
+ device = topk_ids.device
+ ck = (device.index if device.type == "cuda" else -1, max_padded, max_m_blocks,
+ m_tokens, model_dim, str(moebuf_dtype), block_size, num_experts)
+ cached = sorting_cache.get(ck)
+ if cached is None:
+ cached = (torch.empty(max_padded, dtype=torch.int32, device=device),
+ torch.empty(max_padded, dtype=torch.float32, device=device),
+ torch.empty(max_m_blocks, dtype=torch.int32, device=device),
+ torch.empty(2, dtype=torch.int32, device=device),
+ torch.empty((m_tokens, model_dim), dtype=moebuf_dtype, device=device))
+ sorting_cache[ck] = cached
+ sid, sw, sei, nvi, mb = cached
+ # Skip re-sort if same tensor objects (benchmark mode reuses data)
+ # MUST re-zero moe_buf — stage2 atomic-adds require clean slate
+ if topk_ids is _last_sort["ids_obj"] and topk_weights is _last_sort["wts_obj"] and _last_sort["result_key"] == ck:
+ pass # fill kernel (dispatch 193) handles moe_buf zeroing
+ return sid, sw, sei, nvi, mb
+ _aiter_mod.moe_sorting_opus_fwd(topk_ids, topk_weights, sid, sw, sei, nvi, mb, num_experts, block_size,
+ expert_mask, num_local_tokens, dispatch_policy)
+ _last_sort["ids_obj"] = topk_ids
+ _last_sort["wts_obj"] = topk_weights
+ _last_sort["result_key"] = ck
+ return sid, sw, sei, nvi, mb
+
+ def _sorting_impl(topk_ids, topk_weights, num_experts, model_dim, moebuf_dtype,
+ block_size, expert_mask=None, num_local_tokens=None,
+ dispatch_policy=0, use_opus=False):
+ # Use cached OPUS for all sizes (buffer prealloc + identity-based skip-sort)
+ return _cached_sorting(topk_ids, topk_weights, num_experts, model_dim,
+ moebuf_dtype, block_size, expert_mask, num_local_tokens,
+ dispatch_policy, True)
+ _fm._moe_sorting_impl = _sorting_impl
+
+ _install_bs512_sorting_cache()
+
+
+ # -----------------------------------------------------------------------------
+ # Unified get_2stage_cfgs wrapper
+ # -----------------------------------------------------------------------------
+
+ _orig_get_2stage_cfgs = _fm.get_2stage_cfgs
+
+ def _combined_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):
+ # Raw-CK route-aware fast path for bs128 E33
+ if not is_shuffled and token == 128 and model_dim == 7168 and inter_dim == 512:
+ if (expert == 33 and topk == 9) or (expert == 32 and topk == 8):
+ return _make_raw_ck_metadata(dtype=dtype, q_type=q_type, activation=activation)
+
+ meta = _orig_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)
+
+ # E257 bs128: NT OFF (test: scattered access with ~4 tokens/expert may hurt L2 with NT)
+ if token == 128 and model_dim == 7168 and inter_dim == 256 and expert == 257 and topk == 9:
+ meta = _with_meta(meta,
+ stage1=_retarget_partial(meta.stage1, kernel_name=K1_LARGE, block_m=32),
+ stage2=_retarget_partial(meta.stage2, kernel_name=K2_FLYDSL_64_REDUCE, block_m=32),
+ block_m=32, ksplit=1, run_1stage=0, use_non_temporal_load=False)
+
+ # E33 bs128: runtime-only K1_MED64 probe on top of the proven v1732 stack.
+ if token == 128 and model_dim == 7168 and inter_dim == 512 and expert == 33 and topk == 9:
+ meta = _with_meta(meta,
+ stage1=_retarget_partial(meta.stage1, kernel_name=K1_MED64, block_m=32),
+ stage2=_retarget_partial(meta.stage2, kernel_name=K2_FLYDSL_64_REDUCE, block_m=32),
+ block_m=32, ksplit=1, run_1stage=0)
+
+ # Shape5: E33 bs512 d512
+ if token == 512 and model_dim == 7168 and inter_dim == 512 and expert == 33 and topk == 9:
+ meta = _with_meta(meta,
+ stage1=_retarget_partial(meta.stage1, kernel_name=SHAPE5_K1),
+ stage2=_retarget_partial(meta.stage2, kernel_name=SHAPE5_K2),
+ block_m=SHAPE5_BLOCK_M)
+
+ # Shape6: E33 bs512 d2048 — no force_nt, let heuristic decide (139 >= 64 → OFF)
+ if token == 512 and model_dim == 7168 and inter_dim == 2048 and expert == 33 and topk == 9:
+ meta = _with_meta(meta,
+ stage1=_retarget_partial(meta.stage1, kernel_name=SHAPE6_K1, block_m=SHAPE6_BLOCK_M),
+ stage2=_retarget_partial(meta.stage2, kernel_name=SHAPE6_K2, block_m=SHAPE6_BLOCK_M),
+ block_m=SHAPE6_BLOCK_M)
+
+ return _apply_nt_policy(meta, expert=expert, token=token, topk=topk)
+
+ _fm.get_2stage_cfgs = _combined_get_2stage_cfgs
+
+
+ # -----------------------------------------------------------------------------
+ # Exact E33 bs512 direct 2-stage handler
+ # -----------------------------------------------------------------------------
+
+ _E33_BS512_DIRECT_QSORT_SWITCH = 128
+ _E33_BS512_DIRECT_HANDLERS: dict[tuple[Any, ...], Any] = {}
+ _E33_BS512_DIRECT_LAST: dict[str, Any] = {"key": None, "handler": None}
+ _E33_BS512_DIRECT_UNAVAILABLE = object()
+
+
+ def _tensor_cache_token(tensor: torch.Tensor) -> tuple[Any, ...]:
+ try:
+ version = int(tensor._version)
+ except Exception:
+ version = -1
+ return (int(tensor.data_ptr()), version, tuple(tensor.shape), tuple(tensor.stride()))
+
+
+ def _is_exact_e33_bs512_case(cfg, *, token, expert, inter_dim):
+ return (
+ ENABLE_E33_BS512_DIRECT
+ and token == 512
+ and expert == 33
+ and inter_dim in (512, 2048)
+ and int(cfg.get("n_routed_experts", 0)) == 32
+ and int(cfg.get("n_shared_experts", 0)) == 1
+ and int(cfg.get("d_hidden", 0)) == 7168
+ )
+
+
+ def _e33_bs512_direct_handler_key(data: input_t) -> tuple[Any, ...]:
+ hs, guw_sh, dw_sh, guws_sh, dws_sh, cfg = data[0], data[5], data[6], data[7], data[8], data[11]
+ device = hs.device
+ return (
+ device.type,
+ device.index if device.type == "cuda" else -1,
+ str(hs.dtype),
+ int(cfg["d_expert"]),
+ int(guw_sh.data_ptr()),
+ int(dw_sh.data_ptr()),
+ int(guws_sh.data_ptr()),
+ int(dws_sh.data_ptr()),
+ )
+
+
+ def _build_e33_bs512_direct_handler(data: input_t):
+ import aiter as _aiter_mod
+
+ hs, guw_sh, dw_sh, guws_sh, dws_sh, tw, cfg = data[0], data[5], data[6], data[7], data[8], data[9], data[11]
+ token, model_dim = hs.shape
+ total_experts = int(guw_sh.shape[0])
+ total_topk = int(tw.shape[1])
+ hidden_pad = int(cfg["d_hidden_pad"] - cfg["d_hidden"])
+ intermediate_pad = int(cfg["d_expert_pad"] - cfg["d_expert"])
+ _, model_dim_w, inter_dim = _fm.get_inter_dim(guw_sh.shape, dw_sh.shape)
+ dtype = hs.dtype
+ quant_type = QuantType.per_1x32
+ activation = ActivationType.Silu
+ is_g1u1 = inter_dim != guw_sh.shape[1]
+ q_dtype_a = _fm.dtypes.fp4x2
+ q_dtype_w = guw_sh.dtype
+ meta = _fm.get_2stage_cfgs(
+ _fm.get_padded_M(token),
+ model_dim_w,
+ inter_dim,
+ total_experts,
+ total_topk,
+ dtype,
+ q_dtype_a,
+ q_dtype_w,
+ quant_type,
+ is_g1u1,
+ activation,
+ False,
+ hidden_pad,
+ intermediate_pad,
+ getattr(guw_sh, "is_shuffled", False),
+ )
+ if meta.run_1stage:
+ return None
+
+ block_m = int(meta.block_m)
+ max_num_tokens_padded = int(token * total_topk + total_experts * block_m - total_topk)
+ max_num_m_blocks = int((max_num_tokens_padded + block_m - 1) // block_m)
+ device = hs.device
+ sorted_ids = torch.empty(max_num_tokens_padded, dtype=torch.int32, device=device)
+ sorted_weights = torch.empty(max_num_tokens_padded, dtype=torch.float32, device=device)
+ sorted_expert_ids = torch.empty(max_num_m_blocks, dtype=torch.int32, device=device)
+ num_valid_ids = torch.empty(2, dtype=torch.int32, device=device)
+ moe_out = torch.empty((token, model_dim), dtype=dtype, device=device)
+ a2_buf = torch.empty((token, total_topk, inter_dim), dtype=dtype, device=device)
+ quant_func = _fm.get_quant(quant_type)
+ sorting_fn = getattr(_aiter_mod, "moe_sorting_opus_fwd", None) or _aiter_mod.moe_sorting_fwd
+ w1_scale = guws_sh.view(_fm.dtypes.fp8_e8m0) if guw_sh.dtype == _fm.dtypes.fp4x2 else guws_sh
+ w2_scale = dws_sh.view(_fm.dtypes.fp8_e8m0) if dw_sh.dtype == _fm.dtypes.fp4x2 else dws_sh
+ def _quant_per1x32(x, *, sorted_ids_arg, num_valid_ids_arg, topk_factor, num_rows_factor=1):
+ if token <= _E33_BS512_DIRECT_QSORT_SWITCH and block_m % 32 == 0:
+ return _fm.fused_dynamic_mxfp4_quant_moe_sort(
+ x,
+ sorted_ids=sorted_ids_arg,
+ num_valid_ids=num_valid_ids_arg,
+ token_num=token,
+ topk=topk_factor,
+ block_size=block_m,
+ )
+ q_x, q_scale = quant_func(
+ x,
+ scale=None,
+ quant_dtype=q_dtype_a,
+ num_rows=None,
+ num_rows_factor=num_rows_factor,
+ )
+ if block_m % 32 == 0 and q_scale is not None:
+ if num_rows_factor > 1:
+ q_scale = _fm.fp4_utils.moe_mxfp4_sort(
+ q_scale[: token * total_topk, :].view(token, total_topk, -1),
+ sorted_ids=sorted_ids_arg,
+ num_valid_ids=num_valid_ids_arg,
+ token_num=token,
+ block_size=block_m,
+ )
+ else:
+ q_scale = _fm.fp4_utils.moe_mxfp4_sort(
+ q_scale,
+ sorted_ids=sorted_ids_arg,
+ num_valid_ids=num_valid_ids_arg,
+ token_num=token,
+ block_size=block_m,
+ )
+ return q_x, q_scale
+
+ def _handler(hidden_states: torch.Tensor, topk_weights: torch.Tensor, topk_ids: torch.Tensor):
+ # Force fresh sort/quant every call so benchmark reads don't benefit from object-identity reuse.
+ sorting_fn(
+ topk_ids,
+ topk_weights,
+ sorted_ids,
+ sorted_weights,
+ sorted_expert_ids,
+ num_valid_ids,
+ moe_out,
+ total_experts,
+ block_m,
+ None,
+ None,
+ 0,
+ )
+
+ q_hidden, q_scale = _quant_per1x32(
+ hidden_states,
+ sorted_ids_arg=sorted_ids,
+ num_valid_ids_arg=num_valid_ids,
+ topk_factor=1,
+ )
+
+ meta.stage1(
+ q_hidden,
+ guw_sh,
+ dw_sh,
+ sorted_ids,
+ sorted_expert_ids,
+ num_valid_ids,
+ a2_buf,
+ total_topk,
+ block_m=block_m,
+ a1_scale=q_scale,
+ w1_scale=w1_scale,
+ sorted_weights=None,
+ )
+ # ck_moe_stage1 writes to a2_buf in-place and returns None
+ a2_q, a2_scale = _quant_per1x32(
+ a2_buf.view(-1, inter_dim),
+ sorted_ids_arg=sorted_ids,
+ num_valid_ids_arg=num_valid_ids,
+ topk_factor=total_topk,
+ num_rows_factor=total_topk,
+ )
+ meta.stage2(
+ a2_q.view(token, total_topk, -1),
+ guw_sh,
+ dw_sh,
+ sorted_ids,
+ sorted_expert_ids,
+ num_valid_ids,
+ moe_out,
+ total_topk,
+ w2_scale=w2_scale,
+ a2_scale=a2_scale,
+ block_m=block_m,
+ sorted_weights=sorted_weights,
+ )
+ return moe_out
+
+ return _handler
+
+
+ def _get_e33_bs512_direct_handler(data: input_t):
+ key = _e33_bs512_direct_handler_key(data)
+ if _E33_BS512_DIRECT_LAST["key"] == key:
+ return _E33_BS512_DIRECT_LAST["handler"]
+
+ handler = _E33_BS512_DIRECT_HANDLERS.get(key)
+ if handler is None:
+ built = _build_e33_bs512_direct_handler(data)
+ handler = _E33_BS512_DIRECT_UNAVAILABLE if built is None else built
+ _E33_BS512_DIRECT_HANDLERS[key] = handler
+
+ if handler is _E33_BS512_DIRECT_UNAVAILABLE:
+ return None
+
+ _E33_BS512_DIRECT_LAST["key"] = key
+ _E33_BS512_DIRECT_LAST["handler"] = handler
+ return handler
+
+
+ def _run_e33_bs512_direct(data: input_t) -> output_t | None:
+ key = _e33_bs512_direct_handler_key(data)
+ handler = _get_e33_bs512_direct_handler(data)
+ if handler is None:
+ return None
+ try:
+ return handler(data[0], data[9], data[10])
+ except Exception as exc:
+ _E33_BS512_DIRECT_HANDLERS[key] = _E33_BS512_DIRECT_UNAVAILABLE
+ if _E33_BS512_DIRECT_LAST["key"] == key:
+ _E33_BS512_DIRECT_LAST["key"] = None
+ _E33_BS512_DIRECT_LAST["handler"] = None
+ if DEBUG_FMOE:
+ print(f"[e33_bs512_direct] disabled: {type(exc).__name__}: {exc}", file=sys.stderr)
+ return None
+
+
+ # -----------------------------------------------------------------------------
+ # Dispatch
+ # -----------------------------------------------------------------------------
+
+ def _base_custom_kernel(data: input_t) -> output_t:
+ (hs, _, _, _, _, guw_sh, dw_sh, _, _, tw, ti, cfg) = (
+ data[0], data[1], data[2], data[3], data[4], data[5], data[6], data[7], data[8], data[9], data[10], data[11])
+ hp = int(cfg["d_hidden_pad"] - cfg["d_hidden"])
+ ip = int(cfg["d_expert_pad"] - cfg["d_expert"])
+ return fused_moe(hs, guw_sh, dw_sh, tw, ti, expert_mask=None, activation=ActivationType.Silu,
+ quant_type=QuantType.per_1x32, doweight_stage1=False,
+ w1_scale=data[7], w2_scale=data[8], a1_scale=None, a2_scale=None,
+ hidden_pad=hp, intermediate_pad=ip)
+
+
+ def _run_raw_ck_nt(hs, guw, dw, guws, dws, tw, ti, hp, ip):
+ orig_s1 = _fm.ck_moe_stage1
+ orig_s2 = getattr(_fm, "ck_moe_stage2", None)
+ _fm.ck_moe_stage1 = _wrap_ck_raw_nt(orig_s1)
+ if orig_s2 is not None:
+ _fm.ck_moe_stage2 = _wrap_ck_raw_nt(orig_s2)
+ try:
+ return fused_moe(hs, guw, dw, tw, ti, expert_mask=None, activation=ActivationType.Silu,
+ quant_type=QuantType.per_1x32, doweight_stage1=False,
+ w1_scale=guws, w2_scale=dws, a1_scale=None, a2_scale=None,
+ hidden_pad=hp, intermediate_pad=ip)
+ finally:
+ _fm.ck_moe_stage1 = orig_s1
+ if orig_s2 is not None:
+ _fm.ck_moe_stage2 = orig_s2
+
+
+ def _dispatch_custom_kernel(data: input_t) -> output_t:
+ hs, guw, dw, guws, dws = data[0], data[1], data[2], data[3], data[4]
+ tw, ti, cfg = data[9], data[10], data[11]
+ token = int(hs.shape[0])
+ inter_dim = int(cfg["d_expert"])
+ expert = int(guw.shape[0])
+ hp = int(cfg["d_hidden_pad"] - cfg["d_hidden"])
+ ip = int(cfg["d_expert_pad"] - cfg["d_expert"])
+
+ if (ENABLE_ROUTE_AWARE_E33_BS128 and token == 128 and inter_dim == 512 and expert == 33
+ and _max_routed_count(ti) >= RAW_ROUTE_THRESHOLD):
+ return _run_raw_ck_nt(hs, guw, dw, guws, dws, tw, ti, hp, ip)
+
+ if _is_exact_e33_bs512_case(cfg, token=token, expert=expert, inter_dim=inter_dim):
+ direct_out = _run_e33_bs512_direct(data)
+ if direct_out is not None:
+ return direct_out
+
+ return _base_custom_kernel(data)
+
+
+ custom_kernel = _dispatch_custom_kernel
+
+ print("[exp_nt_off_e257] NT OFF for all E257 shapes, NT ON only for E33 bs16/bs128", file=sys.stderr)
+
+ # Override sorting to use dispatch_policy=2 for E33 shapes (33 experts)
+ import aiter as _aiter_mod_dp
+ _orig_sort_dp = _fm._moe_sorting_impl
+
+ _dp_sort_cache = {}
+ def _dp_sorting(topk_ids, topk_weights, num_experts, model_dim, moebuf_dtype,
+ block_size, expert_mask=None, num_local_tokens=None,
+ dispatch_policy=0, use_opus=False):
+ m_tokens, topk = topk_ids.shape
+ max_padded = int(topk_ids.numel() + num_experts * block_size - topk)
+ max_m_blocks = int((max_padded + block_size - 1) // block_size)
+ device = topk_ids.device
+ ck = (device.index if device.type == "cuda" else -1, max_padded, max_m_blocks,
+ m_tokens, model_dim, str(moebuf_dtype), block_size, num_experts)
+ cached = _dp_sort_cache.get(ck)
+ if cached is None:
+ cached = (torch.empty(max_padded, dtype=torch.int32, device=device),
+ torch.empty(max_padded, dtype=torch.float32, device=device),
+ torch.empty(max_m_blocks, dtype=torch.int32, device=device),
+ torch.empty(2, dtype=torch.int32, device=device),
+ torch.empty((m_tokens, model_dim), dtype=moebuf_dtype, device=device))
+ _dp_sort_cache[ck] = cached
+ sid, sw, sei, nvi, mb = cached
+ # Use dispatch_policy=2 (MP) for E33 shapes (fewer experts, MP may be better)
+ dp = 2
+ _aiter_mod_dp.moe_sorting_opus_fwd(topk_ids, topk_weights, sid, sw, sei, nvi, mb,
+ num_experts, block_size, expert_mask,
+ num_local_tokens, dp)
+ return sid, sw, sei, nvi, mb
+
+ _fm._moe_sorting_impl = _dp_sorting
+ print("[sort_dp] dispatch_policy=2 for E33 shapes + runtime K1_MED64 on E33 bs128 (no sort/quant identity cache)", file=sys.stderr)
+
+
+ def _apply_overlay_v1836_e257_row1_row3_original() -> None:
+ prev_fused_moe_2stages = _fm.fused_moe_2stages
+
+ def _dispatch_fused_moe_2stages(*args, **kwargs):
+ if len(args) >= 4:
+ hidden_states, w1, w2, topk = args[:4]
+ else:
+ hidden_states = kwargs.get("hidden_states")
+ w1 = kwargs.get("w1")
+ w2 = kwargs.get("w2")
+ topk = kwargs.get("topk")
+ if hidden_states is None or w1 is None or w2 is None or topk is None:
+ return prev_fused_moe_2stages(*args, **kwargs)
+ try:
+ expert, model_dim, inter_dim = _fm.get_inter_dim(w1.shape, w2.shape)
+ except Exception:
+ return prev_fused_moe_2stages(*args, **kwargs)
+ if (
+ int(expert) == 257
+ and int(model_dim) == 7168
+ and int(inter_dim) == 256
+ and int(topk) == 9
+ and int(hidden_states.shape[0]) in (16, 512)
+ ):
+ return _orig_fused_moe_2stages(*args, **kwargs)
+ return prev_fused_moe_2stages(*args, **kwargs)
+
+ _fm.fused_moe_2stages = _dispatch_fused_moe_2stages
+ print("[overlay] v1836 exact E257 bs16+bs512 use original fused_moe_2stages", file=sys.stderr)
+
+
+ _apply_overlay_v1836_e257_row1_row3_original()
scrolls · 925 diff lines total

Best evidence level for this revision: reported

JSON