submission 754731
.jonnss · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 297 lines, June 9 Researcher Reciprocity License v1.0.
v84_v2_sub.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-754731?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
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:7816cbaf712d98a692ff897c527a5d6471d231f7ce2a96c6ff7fd37dc589b794
license declaredunknown
license concludedunknown
authors.jonnss
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
MoE MXFP4 two-stage kernel — MI355XKernel source
v84_v2_sub.py297 lines
"""
MoE MXFP4 two-stage kernel — MI355X
Patches aiter in-place so FlyDSL handles both GEMM stages for E=33 shapes.
E=257 shapes continue to use the CK fallback path.
"""
import gc
import os
import re
import shutil
import subprocess
import sys
# --- runtime knobs ---
os.environ.setdefault("HIP_FORCE_DEV_KERNARG", "1")
os.environ.setdefault("HSA_ENABLE_INTERRUPT", "0")
# ------------------------------------------------------------------ #
# Step 1 – bring in a recent flydsl wheel so the MLIR compiler works #
# ------------------------------------------------------------------ #
_WHEEL_DIR = "/tmp/_flydsl_wheel"
os.makedirs(_WHEEL_DIR, exist_ok=True)
subprocess.run(
[sys.executable, "-m", "pip", "install", "flydsl==0.1.1",
"--target", _WHEEL_DIR, "--upgrade", "--break-system-packages"],
capture_output=True, timeout=120,
)
sys.path.insert(0, _WHEEL_DIR)
# flush any stale cached module objects
for _k in [k for k in sys.modules if "flydsl" in k]:
del sys.modules[_k]
# ------------------------------------------------------------------ #
# Step 2 – pull a recent aiter checkout and copy the FlyDSL files #
# ------------------------------------------------------------------ #
_AITER_ROOT = "/home/runner/aiter"
_CLONE_DIR = "/tmp/_aiter_clone"
# reset any local edits from previous runs
subprocess.run(["git", "checkout", "--", "."],
capture_output=True, cwd=_AITER_ROOT, timeout=15)
if not os.path.isdir(os.path.join(_CLONE_DIR, "aiter")):
subprocess.run(
["git", "clone", "--depth=50",
"https://github.com/ROCm/aiter.git", _CLONE_DIR],
capture_output=True, timeout=120,
)
# make sure the clone has _get_compiled_stage1 – pin to a commit that does
_kernels_py = os.path.join(_CLONE_DIR, "aiter", "ops", "flydsl", "moe_kernels.py")
if os.path.exists(_kernels_py):
with open(_kernels_py) as _fh:
if "_get_compiled_stage1" not in _fh.read():
_log = subprocess.run(
["git", "log", "--format=%H", "-50"],
capture_output=True, text=True, cwd=_CLONE_DIR, timeout=15,
)
for _sha in _log.stdout.strip().split("\n")[1:]:
_sha = _sha.strip()
if not _sha:
continue
subprocess.run(["git", "checkout", _sha],
capture_output=True, cwd=_CLONE_DIR, timeout=15)
with open(_kernels_py) as _fh2:
if "_get_compiled_stage1" in _fh2.read():
break
# copy the updated fused_moe + FlyDSL kernel files into the live aiter tree
shutil.copy2(
os.path.join(_CLONE_DIR, "aiter", "fused_moe.py"),
os.path.join(_AITER_ROOT, "aiter", "fused_moe.py"),
)
_src_flydsl = os.path.join(_CLONE_DIR, "aiter", "ops", "flydsl")
_dst_flydsl = os.path.join(_AITER_ROOT, "aiter", "ops", "flydsl")
for _fname in ("moe_kernels.py", "utils.py"):
shutil.copy2(os.path.join(_src_flydsl, _fname),
os.path.join(_dst_flydsl, _fname))
# patch out the version guard in __init__.py
_init = os.path.join(_dst_flydsl, "__init__.py")
with open(_init) as _fh:
_txt = _fh.read()
_txt = re.sub(r"raise ImportError\([^)]*\)", "pass # version check skipped", _txt)
with open(_init, "w") as _fh:
_fh.write(_txt)
# copy kernel MLIR sources if present
_src_k = os.path.join(_src_flydsl, "kernels")
_dst_k = os.path.join(_dst_flydsl, "kernels")
if os.path.isdir(_src_k):
os.makedirs(_dst_k, exist_ok=True)
for _fname in os.listdir(_src_k):
if _fname.endswith(".py"):
shutil.copy2(os.path.join(_src_k, _fname),
os.path.join(_dst_k, _fname))
# ------------------------------------------------------------------ #
# Step 3 – source-patch fused_moe.py #
# Inject FlyDSL dispatch functions and override get_2stage_cfgs. #
# ------------------------------------------------------------------ #
_fmoe_path = os.path.join(_AITER_ROOT, "aiter", "fused_moe.py")
with open(_fmoe_path) as _fh:
_fmoe_src = _fh.read()
_patch = r'''
# ---- FlyDSL stage-1: gate+up+SiLU (E<=100 shapes only) ----
def _moe_stage1_fly(hidden_states, w1, w2, sorted_token_ids,
sorted_expert_ids, num_valid_ids, out, topk,
w1_scale=None, a1_scale=None, sorted_weights=None, **_kw):
from aiter.ops.flydsl.moe_kernels import flydsl_moe_stage1
flydsl_moe_stage1(
a=hidden_states, w1=w1,
sorted_token_ids=sorted_token_ids,
sorted_expert_ids=sorted_expert_ids,
num_valid_ids=num_valid_ids, out=out, topk=topk,
tile_m=32, tile_n=256, tile_k=128,
a_dtype="fp4", b_dtype="fp4", out_dtype="bf16",
act="silu",
w1_scale=w1_scale, a1_scale=a1_scale,
sorted_weights=sorted_weights,
)
return out
# ---- FlyDSL stage-2: down-projection ----
def _moe_stage2_fly(inter_states, w1, w2, sorted_token_ids,
sorted_expert_ids, num_valid_ids, out, topk,
w2_scale=None, a2_scale=None, sorted_weights=None, **_kw):
from aiter.ops.flydsl.moe_kernels import flydsl_moe_stage2
flydsl_moe_stage2(
inter_states=inter_states, w2=w2,
sorted_token_ids=sorted_token_ids,
sorted_expert_ids=sorted_expert_ids,
num_valid_ids=num_valid_ids, out=out, topk=topk,
tile_m=32, tile_n=256, tile_k=128,
a_dtype="fp4", b_dtype="fp4", out_dtype="bf16",
mode="atomic",
w2_scale=w2_scale, a2_scale=a2_scale,
sorted_weights=sorted_weights,
)
_orig_g2sc = get_2stage_cfgs
try:
_orig_g2sc.cache_clear()
except AttributeError:
pass
_stage1_ready = [False]
_stage2_ready = [False]
import functools as _ft
@_ft.lru_cache(maxsize=2048)
def _patched_g2sc(*args, **kwargs):
meta = _orig_g2sc(*args, **kwargs)
n_experts = args[4] if len(args) >= 5 else 0
if (_stage1_ready[0] and n_experts <= 100
and not meta.run_1stage and meta.stage1 is not None):
meta.stage1 = _ft.partial(_moe_stage1_fly)
if (_stage2_ready[0] and n_experts <= 100 and not meta.run_1stage
and meta.stage2 is not None):
meta.stage2 = _ft.partial(_moe_stage2_fly)
return meta
get_2stage_cfgs = _patched_g2sc
'''
_fmoe_src += _patch
with open(_fmoe_path, "w") as _fh:
_fh.write(_fmoe_src)
# ------------------------------------------------------------------ #
# Step 4 – fix Triton constexpr operator signatures if needed #
# ------------------------------------------------------------------ #
try:
import triton.language.core as _tlc
_ce = _tlc.constexpr
try:
_ce(1).__lt__(_ce(2), _semantic=None)
except TypeError:
_ops = [
"__lt__", "__le__", "__gt__", "__ge__", "__eq__", "__ne__",
"__add__", "__radd__", "__sub__", "__rsub__", "__mul__", "__rmul__",
"__truediv__", "__floordiv__", "__mod__", "__pow__",
"__lshift__", "__rshift__", "__and__", "__or__", "__xor__",
]
for _op in _ops:
_orig_op = getattr(_ce, _op, None)
if _orig_op and not getattr(_orig_op, "_patched", False):
def _wrap(fn):
def _inner(self, *a, _semantic=None, **kw):
return fn(self, *a, **kw)
_inner._patched = True
return _inner
setattr(_ce, _op, _wrap(_orig_op))
del _ce, _tlc
except Exception:
pass
# ------------------------------------------------------------------ #
# Step 5 – normal imports (see patched aiter on disk) #
# ------------------------------------------------------------------ #
import torch
from task import input_t, output_t
from aiter import ActivationType, QuantType
from aiter.fused_moe import fused_moe
import aiter.fused_moe as _fm
# ------------------------------------------------------------------ #
# Step 6 – pre-compile FlyDSL MLIR + set ready flags #
# ------------------------------------------------------------------ #
try:
from aiter.ops.flydsl.moe_kernels import (
_get_compiled_stage1,
_get_compiled_stage2,
)
# compile stage-1 for both E=33 inter-dim variants
for _idim in (512, 2048):
try:
_get_compiled_stage1(
model_dim=7168, inter_dim=_idim, experts=33, topk=9,
tile_m=32, tile_n=256, tile_k=128,
doweight=True, a_dtype="fp4", b_dtype="fp4",
out_dtype="bf16", act="silu",
)
except Exception as _exc:
print(f"[kernel] stage1 compile failed (inter={_idim}): {_exc}",
file=sys.stderr)
_fm._stage1_ready[0] = True
# compile stage-2 for all relevant shapes
for _idim, _nexp in [(256, 257), (512, 33), (2048, 33)]:
try:
_get_compiled_stage2(
model_dim=7168, inter_dim=_idim, experts=_nexp, topk=9,
tile_m=32, tile_n=256, tile_k=128,
doweight=True, a_dtype="fp4", b_dtype="fp4",
out_dtype="bf16", accumulate=True,
)
except Exception:
pass
_fm._stage2_ready[0] = True
print("[kernel] FlyDSL compilation done", file=sys.stderr)
except Exception as _exc:
print(f"[kernel] FlyDSL setup error: {_exc}", file=sys.stderr)
# ------------------------------------------------------------------ #
# Step 7 – per-shape warmup call #
# ------------------------------------------------------------------ #
_seen_shapes: set = set()
def _warmup(data):
(hs, _guw, _dw, _guws, _dws,
guw_sh, dw_sh, guws_sh, dws_sh, tw, ti, cfg) = data
shape_key = (cfg["bs"], cfg["n_routed_experts"], cfg["d_expert"])
if shape_key in _seen_shapes:
return
_seen_shapes.add(shape_key)
try:
fused_moe(
hs[:2], guw_sh, dw_sh, tw[:2], ti[:2],
activation=ActivationType.Silu,
quant_type=QuantType.per_1x32,
w1_scale=guws_sh, w2_scale=dws_sh,
hidden_pad=int(cfg["d_hidden_pad"]) - int(cfg["d_hidden"]),
intermediate_pad=int(cfg["d_expert_pad"]) - int(cfg["d_expert"]),
)
torch.cuda.synchronize()
except Exception:
pass
# ------------------------------------------------------------------ #
# Kernel entry point #
# ------------------------------------------------------------------ #
def custom_kernel(data: input_t) -> output_t:
_warmup(data)
(hs, _guw, _dw, _guws, _dws,
guw_sh, dw_sh, guws_sh, dws_sh, tw, ti, cfg) = data
return fused_moe(
hs, guw_sh, dw_sh, tw, ti,
activation=ActivationType.Silu,
quant_type=QuantType.per_1x32,
w1_scale=guws_sh,
w2_scale=dws_sh,
hidden_pad=int(cfg["d_hidden_pad"]) - int(cfg["d_hidden"]),
intermediate_pad=int(cfg["d_expert_pad"]) - int(cfg["d_expert"]),
)
scrolls · 297 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 715459.
- import dataclasses- import functools+ """+ MoE MXFP4 two-stage kernel — MI355X+ Patches aiter in-place so FlyDSL handles both GEMM stages for E=33 shapes.+ E=257 shapes continue to use the CK fallback path.+ """+import gc- import inspect- import importlibimport os+ import re+ import shutil+ import subprocess+ import sys- import torch+ # --- runtime knobs ---+ os.environ.setdefault("HIP_FORCE_DEV_KERNARG", "1")+ os.environ.setdefault("HSA_ENABLE_INTERRUPT", "0")- from task import input_t, output_t-- os.environ["VLLM_MOE_CHUNK_SIZE"] = "512"- os.environ["HIP_FORCE_DEV_KERNARG"] = "1"-- from aiter import ActivationType, QuantType- from aiter.fused_moe import fused_moe--- gc.disable()--- _PATCH_DONE = False- _DIRECT_FUSED_MOE = None- _FUSED_MOE_MODULE = None- _CURRENT_SORT_BLOCK_SIZE = 32- _CURRENT_SORT_EXPERT = -1- _ADAPTIVE_QUANT_THRESHOLD = 1_000_000- _CFG_CACHE = {}- _SORT_BLOCK_CACHE = {}- _FLYDSL_STAGE2_KEY = "flydsl_moe2_afp4_wfp4_bf16_t16x256x128_atomic"- _FLYDSL_BASE_KEY = "flydsl_moe2_afp4_wfp4_bf16_t32x256x256_atomic"- _E33_STAGE2_SHAPES = {- (128, 7168, 512, 33, 9): 64,- (512, 7168, 512, 33, 9): 128,- (512, 7168, 2048, 33, 9): 64,- }- _DIRECT_PATH_SHAPES = frozenset(- {- (512, 7168, 256, 257, 9),- (128, 7168, 512, 33, 9),- (512, 7168, 512, 33, 9),- (512, 7168, 2048, 33, 9),- }+ # ------------------------------------------------------------------ #+ # Step 1 – bring in a recent flydsl wheel so the MLIR compiler works #+ # ------------------------------------------------------------------ #+ _WHEEL_DIR = "/tmp/_flydsl_wheel"+ os.makedirs(_WHEEL_DIR, exist_ok=True)+ subprocess.run(+ [sys.executable, "-m", "pip", "install", "flydsl==0.1.1",+ "--target", _WHEEL_DIR, "--upgrade", "--break-system-packages"],+ capture_output=True, timeout=120,)+ sys.path.insert(0, _WHEEL_DIR)+ # flush any stale cached module objects+ for _k in [k for k in sys.modules if "flydsl" in k]:+ del sys.modules[_k]+ # ------------------------------------------------------------------ #+ # Step 2 – pull a recent aiter checkout and copy the FlyDSL files #+ # ------------------------------------------------------------------ #+ _AITER_ROOT = "/home/runner/aiter"+ _CLONE_DIR = "/tmp/_aiter_clone"- def _replace_meta(meta, *, block_m=None, ksplit=None, stage1=None, stage2=None):- try:- if dataclasses.is_dataclass(meta):- updates = {}- if block_m is not None:- updates["block_m"] = block_m- if ksplit is not None:- updates["ksplit"] = ksplit- if stage1 is not None:- updates["stage1"] = stage1- if stage2 is not None:- updates["stage2"] = stage2- if updates:- return dataclasses.replace(meta, **updates)- except Exception:- pass+ # reset any local edits from previous runs+ subprocess.run(["git", "checkout", "--", "."],+ capture_output=True, cwd=_AITER_ROOT, timeout=15)- try:- if block_m is not None:- meta.block_m = block_m- if ksplit is not None:- meta.ksplit = ksplit- if stage1 is not None:- meta.stage1 = stage1- if stage2 is not None:- meta.stage2 = stage2- return meta- except Exception:- pass+ if not os.path.isdir(os.path.join(_CLONE_DIR, "aiter")):+ subprocess.run(+ ["git", "clone", "--depth=50",+ "https://github.com/ROCm/aiter.git", _CLONE_DIR],+ capture_output=True, timeout=120,+ )- if all(hasattr(meta, name) for name in ("stage1", "stage2", "block_m", "ksplit", "run_1stage")):- has_bias = getattr(meta, "has_bias", False)- use_non_temporal_load = getattr(meta, "use_non_temporal_load", True)- try:- return type(meta)(- 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,- has_bias,- use_non_temporal_load,+ # make sure the clone has _get_compiled_stage1 – pin to a commit that does+ _kernels_py = os.path.join(_CLONE_DIR, "aiter", "ops", "flydsl", "moe_kernels.py")+ if os.path.exists(_kernels_py):+ with open(_kernels_py) as _fh:+ if "_get_compiled_stage1" not in _fh.read():+ _log = subprocess.run(+ ["git", "log", "--format=%H", "-50"],+ capture_output=True, text=True, cwd=_CLONE_DIR, timeout=15,)- except Exception:- pass+ for _sha in _log.stdout.strip().split("\n")[1:]:+ _sha = _sha.strip()+ if not _sha:+ continue+ subprocess.run(["git", "checkout", _sha],+ capture_output=True, cwd=_CLONE_DIR, timeout=15)+ with open(_kernels_py) as _fh2:+ if "_get_compiled_stage1" in _fh2.read():+ break- return meta+ # copy the updated fused_moe + FlyDSL kernel files into the live aiter tree+ shutil.copy2(+ os.path.join(_CLONE_DIR, "aiter", "fused_moe.py"),+ os.path.join(_AITER_ROOT, "aiter", "fused_moe.py"),+ )+ _src_flydsl = os.path.join(_CLONE_DIR, "aiter", "ops", "flydsl")+ _dst_flydsl = os.path.join(_AITER_ROOT, "aiter", "ops", "flydsl")+ for _fname in ("moe_kernels.py", "utils.py"):+ shutil.copy2(os.path.join(_src_flydsl, _fname),+ os.path.join(_dst_flydsl, _fname))- def _split_kw_name(fn) -> str:- if isinstance(fn, functools.partial):- keywords = dict(fn.keywords or {})- if "splitk" in keywords:- return "splitk"- if "split_k" in keywords:- return "split_k"- target = fn.func if isinstance(fn, functools.partial) else fn- name = getattr(target, "__name__", "")- if name in {"cktile_moe_gemm1", "cktile_moe_stage1"}:- return "split_k"- return "splitk"+ # patch out the version guard in __init__.py+ _init = os.path.join(_dst_flydsl, "__init__.py")+ with open(_init) as _fh:+ _txt = _fh.read()+ _txt = re.sub(r"raise ImportError\([^)]*\)", "pass # version check skipped", _txt)+ with open(_init, "w") as _fh:+ _fh.write(_txt)+ # copy kernel MLIR sources if present+ _src_k = os.path.join(_src_flydsl, "kernels")+ _dst_k = os.path.join(_dst_flydsl, "kernels")+ if os.path.isdir(_src_k):+ os.makedirs(_dst_k, exist_ok=True)+ for _fname in os.listdir(_src_k):+ if _fname.endswith(".py"):+ shutil.copy2(os.path.join(_src_k, _fname),+ os.path.join(_dst_k, _fname))- def _retune_partial(fn, *, block_m=None, splitk=None):- if not isinstance(fn, functools.partial):- return fn- keywords = dict(fn.keywords or {})- if block_m is not None:- keywords["block_m"] = block_m- if splitk is not None:- keywords.pop("splitk", None)- keywords.pop("split_k", None)- keywords[_split_kw_name(fn)] = splitk- return functools.partial(fn.func, *(fn.args or ()), **keywords)+ # ------------------------------------------------------------------ #+ # Step 3 – source-patch fused_moe.py #+ # Inject FlyDSL dispatch functions and override get_2stage_cfgs. #+ # ------------------------------------------------------------------ #+ _fmoe_path = os.path.join(_AITER_ROOT, "aiter", "fused_moe.py")+ with open(_fmoe_path) as _fh:+ _fmoe_src = _fh.read()+ _patch = r'''- def _wrap_flydsl_stage2(stage2_wrapper, original_stage2, kernel_name: str):- keywords = {}- if isinstance(original_stage2, functools.partial):- keywords.update(original_stage2.keywords or {})- keywords["kernelName"] = kernel_name- return functools.partial(stage2_wrapper, **keywords)+ # ---- FlyDSL stage-1: gate+up+SiLU (E<=100 shapes only) ----+ def _moe_stage1_fly(hidden_states, w1, w2, sorted_token_ids,+ sorted_expert_ids, num_valid_ids, out, topk,+ w1_scale=None, a1_scale=None, sorted_weights=None, **_kw):+ from aiter.ops.flydsl.moe_kernels import flydsl_moe_stage1+ flydsl_moe_stage1(+ a=hidden_states, w1=w1,+ sorted_token_ids=sorted_token_ids,+ sorted_expert_ids=sorted_expert_ids,+ num_valid_ids=num_valid_ids, out=out, topk=topk,+ tile_m=32, tile_n=256, tile_k=128,+ a_dtype="fp4", b_dtype="fp4", out_dtype="bf16",+ act="silu",+ w1_scale=w1_scale, a1_scale=a1_scale,+ sorted_weights=sorted_weights,+ )+ return out-- def _register_flydsl_t16(moe_kernels) -> None:- kernel_params = getattr(moe_kernels, "_KERNEL_PARAMS", None)- if not isinstance(kernel_params, dict):- return- if _FLYDSL_STAGE2_KEY in kernel_params:- return-- base = dict(kernel_params.get(_FLYDSL_BASE_KEY, {}))- if not base:- return-- base.update(- tile_m=16,- tile_n=256,- tile_k=128,+ # ---- FlyDSL stage-2: down-projection ----+ def _moe_stage2_fly(inter_states, w1, w2, sorted_token_ids,+ sorted_expert_ids, num_valid_ids, out, topk,+ w2_scale=None, a2_scale=None, sorted_weights=None, **_kw):+ from aiter.ops.flydsl.moe_kernels import flydsl_moe_stage2+ flydsl_moe_stage2(+ inter_states=inter_states, w2=w2,+ sorted_token_ids=sorted_token_ids,+ sorted_expert_ids=sorted_expert_ids,+ num_valid_ids=num_valid_ids, out=out, topk=topk,+ tile_m=32, tile_n=256, tile_k=128,+ a_dtype="fp4", b_dtype="fp4", out_dtype="bf16",mode="atomic",- MPerBlock=16,+ w2_scale=w2_scale, a2_scale=a2_scale,+ sorted_weights=sorted_weights,)- kernel_params[_FLYDSL_STAGE2_KEY] = base+ _orig_g2sc = get_2stage_cfgs+ try:+ _orig_g2sc.cache_clear()+ except AttributeError:+ pass- def _precompile_flydsl_defs(moe_kernels) -> None:- get_compiled_stage2 = getattr(moe_kernels, "_get_compiled_stage2", None)- if get_compiled_stage2 is None:- return+ _stage1_ready = [False]+ _stage2_ready = [False]- for inter_dim in (512, 2048):- try:- get_compiled_stage2(- model_dim=7168,- inter_dim=inter_dim,- experts=33,- topk=9,- tile_m=16,- tile_n=256,- tile_k=128,- doweight=True,- a_dtype="fp4",- b_dtype="fp4",- out_dtype="bf16",- )- except Exception:- pass+ import functools as _ft+ @_ft.lru_cache(maxsize=2048)+ def _patched_g2sc(*args, **kwargs):+ meta = _orig_g2sc(*args, **kwargs)+ n_experts = args[4] if len(args) >= 5 else 0+ if (_stage1_ready[0] and n_experts <= 100+ and not meta.run_1stage and meta.stage1 is not None):+ meta.stage1 = _ft.partial(_moe_stage1_fly)+ if (_stage2_ready[0] and n_experts <= 100 and not meta.run_1stage+ and meta.stage2 is not None):+ meta.stage2 = _ft.partial(_moe_stage2_fly)+ return meta- def _choose_block_size_m(hidden_states: torch.Tensor, gate_up_weight_shuffled: torch.Tensor) -> int | None:- m = int(hidden_states.shape[0])- if m <= 16:- return 16- return None+ get_2stage_cfgs = _patched_g2sc+ '''+ _fmoe_src += _patch+ with open(_fmoe_path, "w") as _fh:+ _fh.write(_fmoe_src)- def _should_use_direct_path(- hidden_states: torch.Tensor,- gate_up_weight_shuffled: torch.Tensor,- topk_ids: torch.Tensor,- config,- ) -> bool:- shape = (- int(hidden_states.shape[0]),- int(config["d_hidden"]),- int(config["d_expert"]),- int(gate_up_weight_shuffled.shape[0]),- int(topk_ids.shape[1]),- )- return shape in _DIRECT_PATH_SHAPES+ # ------------------------------------------------------------------ #+ # Step 4 – fix Triton constexpr operator signatures if needed #+ # ------------------------------------------------------------------ #+ try:+ import triton.language.core as _tlc+ _ce = _tlc.constexpr+ try:+ _ce(1).__lt__(_ce(2), _semantic=None)+ except TypeError:+ _ops = [+ "__lt__", "__le__", "__gt__", "__ge__", "__eq__", "__ne__",+ "__add__", "__radd__", "__sub__", "__rsub__", "__mul__", "__rmul__",+ "__truediv__", "__floordiv__", "__mod__", "__pow__",+ "__lshift__", "__rshift__", "__and__", "__or__", "__xor__",+ ]+ for _op in _ops:+ _orig_op = getattr(_ce, _op, None)+ if _orig_op and not getattr(_orig_op, "_patched", False):+ def _wrap(fn):+ def _inner(self, *a, _semantic=None, **kw):+ return fn(self, *a, **kw)+ _inner._patched = True+ return _inner+ setattr(_ce, _op, _wrap(_orig_op))+ del _ce, _tlc+ except Exception:+ pass+ # ------------------------------------------------------------------ #+ # Step 5 – normal imports (see patched aiter on disk) #+ # ------------------------------------------------------------------ #+ import torch- def _use_bypass(token: int, expert: int) -> bool:- return token <= 16 or (token <= 128 and expert > 64)+ from task import input_t, output_t+ from aiter import ActivationType, QuantType+ from aiter.fused_moe import fused_moe+ import aiter.fused_moe as _fm+ # ------------------------------------------------------------------ #+ # Step 6 – pre-compile FlyDSL MLIR + set ready flags #+ # ------------------------------------------------------------------ #+ try:+ from aiter.ops.flydsl.moe_kernels import (+ _get_compiled_stage1,+ _get_compiled_stage2,+ )- def _desired_ksplit(token: int, expert: int) -> int:- return 2 if _use_bypass(token, expert) else 0+ # compile stage-1 for both E=33 inter-dim variants+ for _idim in (512, 2048):+ try:+ _get_compiled_stage1(+ model_dim=7168, inter_dim=_idim, experts=33, topk=9,+ tile_m=32, tile_n=256, tile_k=128,+ doweight=True, a_dtype="fp4", b_dtype="fp4",+ out_dtype="bf16", act="silu",+ )+ except Exception as _exc:+ print(f"[kernel] stage1 compile failed (inter={_idim}): {_exc}",+ file=sys.stderr)+ _fm._stage1_ready[0] = True- def _cacheable(value):- if isinstance(value, (str, int, float, bool, type(None))):- return value- if isinstance(value, tuple):- return tuple(_cacheable(item) for item in value)- return repr(value)--- def _make_cfg_cache_key(sig: inspect.Signature | None, use_bypass: bool, call_args, call_kwargs):- if sig is not None:+ # compile stage-2 for all relevant shapes+ for _idim, _nexp in [(256, 257), (512, 33), (2048, 33)]:try:- bound = sig.bind_partial(*call_args, **call_kwargs)- return (- use_bypass,- tuple((name, _cacheable(value)) for name, value in bound.arguments.items()),+ _get_compiled_stage2(+ model_dim=7168, inter_dim=_idim, experts=_nexp, topk=9,+ tile_m=32, tile_n=256, tile_k=128,+ doweight=True, a_dtype="fp4", b_dtype="fp4",+ out_dtype="bf16", accumulate=True,)except Exception:pass- return (- use_bypass,- tuple(_cacheable(value) for value in call_args),- tuple(sorted((key, _cacheable(value)) for key, value in call_kwargs.items())),- )+ _fm._stage2_ready[0] = True+ print("[kernel] FlyDSL compilation done", file=sys.stderr)- def _call_with_bypass_env(use_bypass: bool, fn, *args, **kwargs):- previous = os.environ.get("AITER_BYPASS_TUNE_CONFIG")- os.environ["AITER_BYPASS_TUNE_CONFIG"] = "1" if use_bypass else "0"- try:- return fn(*args, **kwargs)- finally:- if previous is None:- os.environ.pop("AITER_BYPASS_TUNE_CONFIG", None)- else:- os.environ["AITER_BYPASS_TUNE_CONFIG"] = previous+ except Exception as _exc:+ print(f"[kernel] FlyDSL setup error: {_exc}", file=sys.stderr)+ # ------------------------------------------------------------------ #+ # Step 7 – per-shape warmup call #+ # ------------------------------------------------------------------ #+ _seen_shapes: set = set()- def _install_patch() -> None:- global _PATCH_DONE, _DIRECT_FUSED_MOE, _FUSED_MOE_MODULE- if _PATCH_DONE:+ def _warmup(data):+ (hs, _guw, _dw, _guws, _dws,+ guw_sh, dw_sh, guws_sh, dws_sh, tw, ti, cfg) = data+ shape_key = (cfg["bs"], cfg["n_routed_experts"], cfg["d_expert"])+ if shape_key in _seen_shapes:return- _PATCH_DONE = True-+ _seen_shapes.add(shape_key)try:- aiter_mod = importlib.import_module("aiter")- fused_moe_module = importlib.import_module("aiter.fused_moe")- _FUSED_MOE_MODULE = fused_moe_module- fp4_utils = importlib.import_module("aiter.utility.fp4_utils")- moe_kernels = importlib.import_module("aiter.ops.flydsl.moe_kernels")-- _register_flydsl_t16(moe_kernels)- _precompile_flydsl_defs(moe_kernels)-- original_cfgs = getattr(fused_moe_module, "get_2stage_cfgs", None)- raw_cfgs = getattr(original_cfgs, "__wrapped__", original_cfgs)- original_cfgs_sig = inspect.signature(original_cfgs) if original_cfgs is not None else None- raw_cfgs_sig = inspect.signature(raw_cfgs) if raw_cfgs is not None else None- original_get_ksplit = getattr(fused_moe_module, "get_ksplit", None)- original_get_ksplit_sig = inspect.signature(original_get_ksplit) if original_get_ksplit is not None else None- original_quant_sort = getattr(fused_moe_module, "fused_dynamic_mxfp4_quant_moe_sort", None)- quant_hip = getattr(aiter_mod, "per_1x32_f4_quant_hip", None)- moe_mxfp4_sort = getattr(fp4_utils, "moe_mxfp4_sort", None)- flydsl_stage2_wrapper = getattr(fused_moe_module, "_flydsl_stage2_wrapper", None)- fused_moe_internal = getattr(fused_moe_module, "fused_moe_", None)- fused_moe_2stages = getattr(fused_moe_module, "fused_moe_2stages", None)- direct_fused_moe = None- if fused_moe_internal is not None:- try:- direct_fused_moe = inspect.unwrap(fused_moe_internal)- except Exception:- direct_fused_moe = getattr(fused_moe_internal, "__wrapped__", None)- if direct_fused_moe is None:- direct_fused_moe = fused_moe_internal- _DIRECT_FUSED_MOE = direct_fused_moe- if original_cfgs is None or flydsl_stage2_wrapper is None:- return-- def _extract_int(- sig: inspect.Signature | None,- call_args,- call_kwargs,- name: str,- fallback_index: int,- ) -> int:- if sig is not None:- try:- bound_args = sig.bind_partial(*call_args, **call_kwargs).arguments- if name in bound_args:- return int(bound_args[name])- except Exception:- pass- value = call_kwargs.get(name, call_args[fallback_index] if len(call_args) > fallback_index else -1)- try:- return int(value)- except Exception:- return -1-- def patched_get_2stage_cfgs(*args, **kwargs):- global _CURRENT_SORT_BLOCK_SIZE, _CURRENT_SORT_EXPERT- token = _extract_int(original_cfgs_sig, args, kwargs, "token", 0)- model_dim = _extract_int(original_cfgs_sig, args, kwargs, "model_dim", 1)- inter_dim = _extract_int(original_cfgs_sig, args, kwargs, "inter_dim", 2)- expert = _extract_int(original_cfgs_sig, args, kwargs, "expert", 3)- topk = _extract_int(original_cfgs_sig, args, kwargs, "topk", 4)- shape = (token, model_dim, inter_dim, expert, topk)- use_bypass = _use_bypass(token, expert)- cache_key = _make_cfg_cache_key(raw_cfgs_sig, use_bypass, args, kwargs)- meta = _CFG_CACHE.get(cache_key)- if meta is None:- raw_meta = _call_with_bypass_env(use_bypass, raw_cfgs, *args, **kwargs)- ksplit = _desired_ksplit(token, expert)- block_m = getattr(raw_meta, "block_m", None)- stage1 = _retune_partial(getattr(raw_meta, "stage1", None), splitk=ksplit)- stage2 = getattr(raw_meta, "stage2", None)-- override_block_m = _E33_STAGE2_SHAPES.get(shape)- if override_block_m is not None:- block_m = override_block_m- stage1 = _retune_partial(stage1, block_m=override_block_m, splitk=0)- stage2 = _retune_partial(stage2, block_m=override_block_m)- stage2 = _wrap_flydsl_stage2(flydsl_stage2_wrapper, stage2, _FLYDSL_STAGE2_KEY)- ksplit = 0-- if shape == (512, 7168, 256, 257, 9):- ksplit = 0- stage1 = _retune_partial(stage1, splitk=0)-- meta = _replace_meta(raw_meta, block_m=block_m, ksplit=ksplit, stage1=stage1, stage2=stage2)- _CFG_CACHE[cache_key] = meta-- _CURRENT_SORT_BLOCK_SIZE = int(getattr(meta, "block_m", None) or 32)- _CURRENT_SORT_EXPERT = expert- return meta-- def patched_get_ksplit(*args, **kwargs):- token = _extract_int(original_get_ksplit_sig, args, kwargs, "token", 0)- expert = _extract_int(original_get_ksplit_sig, args, kwargs, "expert", 3)- if token >= 0 and expert >= 0:- return _desired_ksplit(token, expert)- if original_get_ksplit is None:- return 0- return original_get_ksplit(*args, **kwargs)-- fused_moe_module.get_2stage_cfgs = patched_get_2stage_cfgs- if original_get_ksplit is not None:- fused_moe_module.get_ksplit = patched_get_ksplit-- for maybe_globals in (- getattr(fused_moe, "__globals__", None),- getattr(fused_moe_internal, "__globals__", None),- getattr(fused_moe_2stages, "__globals__", None),- getattr(direct_fused_moe, "__globals__", None),- ):- if isinstance(maybe_globals, dict) and maybe_globals.get("get_2stage_cfgs") is original_cfgs:- maybe_globals["get_2stage_cfgs"] = patched_get_2stage_cfgs- if isinstance(maybe_globals, dict) and maybe_globals.get("get_ksplit") is original_get_ksplit:- maybe_globals["get_ksplit"] = patched_get_ksplit-- if original_quant_sort is None or quant_hip is None or moe_mxfp4_sort is None:- return-- def split_quant_sort(- x: torch.Tensor,- sorted_ids: torch.Tensor,- num_valid_ids: torch.Tensor,- token_num: int,- topk: int,- block_size: int = 32,- scaling_mode: str = "even",- ):- if scaling_mode != "even":- return original_quant_sort(- x,- sorted_ids,- num_valid_ids,- token_num,- topk,- block_size=block_size,- scaling_mode=scaling_mode,- )-- quant_out = quant_hip(x, scale=None, shuffle=False)- x_fp4, blockscale_e8m0 = quant_out[:2]- requested_block_size = int(_CURRENT_SORT_BLOCK_SIZE or block_size)- cache_key = (requested_block_size, int(_CURRENT_SORT_EXPERT))- candidates = []- for candidate in (- _SORT_BLOCK_CACHE.get(cache_key),- requested_block_size,- int(block_size),- max(requested_block_size, 128),- 64,- 128,- 256,- 512,- ):- if candidate and candidate not in candidates:- candidates.append(candidate)-- last_error = None- for candidate in candidates:- try:- blockscale_sorted = moe_mxfp4_sort(- blockscale_e8m0,- sorted_ids,- num_valid_ids,- token_num,- candidate,- )- _SORT_BLOCK_CACHE[cache_key] = candidate- return x_fp4, blockscale_sorted- except AssertionError as exc:- last_error = exc-- if last_error is not None:- raise last_error- return x_fp4, blockscale_e8m0-- def adaptive_quant_sort(*args, **kwargs):- x = kwargs.get("x", args[0] if len(args) > 0 else None)- scaling_mode = kwargs.get("scaling_mode", args[6] if len(args) > 6 else "even")- if x is None or scaling_mode != "even":- return original_quant_sort(*args, **kwargs)- if int(x.numel()) > _ADAPTIVE_QUANT_THRESHOLD:- return split_quant_sort(*args, **kwargs)- return original_quant_sort(*args, **kwargs)-- fused_moe_module.fused_dynamic_mxfp4_quant_moe_sort = adaptive_quant_sort- for maybe_globals in (- getattr(fused_moe_internal, "__globals__", None),- getattr(fused_moe_2stages, "__globals__", None),- getattr(direct_fused_moe, "__globals__", None),- ):- if (- isinstance(maybe_globals, dict)- and maybe_globals.get("fused_dynamic_mxfp4_quant_moe_sort") is original_quant_sort- ):- maybe_globals["fused_dynamic_mxfp4_quant_moe_sort"] = adaptive_quant_sort- except Exception:- pass--- def _call_moe(- hidden_states: torch.Tensor,- gate_up_weight_shuffled: torch.Tensor,- down_weight_shuffled: torch.Tensor,- topk_weights: torch.Tensor,- topk_ids: torch.Tensor,- gate_up_weight_scale_shuffled: torch.Tensor,- down_weight_scale_shuffled: torch.Tensor,- hidden_pad: int,- intermediate_pad: int,- block_size_m: int | None,- ):- direct_fused_moe = _DIRECT_FUSED_MOE- if direct_fused_moe is None:- return fused_moe(- hidden_states,- gate_up_weight_shuffled,- down_weight_shuffled,- topk_weights,- topk_ids,- expert_mask=None,+ fused_moe(+ hs[:2], guw_sh, dw_sh, tw[:2], ti[:2],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,- block_size_M=block_size_m,- hidden_pad=hidden_pad,- intermediate_pad=intermediate_pad,+ w1_scale=guws_sh, w2_scale=dws_sh,+ hidden_pad=int(cfg["d_hidden_pad"]) - int(cfg["d_hidden"]),+ intermediate_pad=int(cfg["d_expert_pad"]) - int(cfg["d_expert"]),)-- block_size_arg = -1 if not block_size_m else int(block_size_m)- try:- return direct_fused_moe(- hidden_states=hidden_states,- w1=gate_up_weight_shuffled,- w2=down_weight_shuffled,- topk_weight=topk_weights,- topk_ids=topk_ids,- expert_mask=None,- activation=ActivationType.Silu.value,- quant_type=QuantType.per_1x32.value,- doweight_stage1=False,- w1_scale=gate_up_weight_scale_shuffled,- w2_scale=down_weight_scale_shuffled,- a1_scale=None,- a2_scale=None,- block_size_M=block_size_arg,- num_local_tokens=None,- moe_sorting_dispatch_policy=0,- dtype=None,- hidden_pad=hidden_pad,- intermediate_pad=intermediate_pad,- bias1=None,- bias2=None,- )+ torch.cuda.synchronize()except Exception:- 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,- block_size_M=block_size_m,- hidden_pad=hidden_pad,- intermediate_pad=intermediate_pad,- )+ pass- def _call_moe_inlined(- hidden_states: torch.Tensor,- gate_up_weight_shuffled: torch.Tensor,- down_weight_shuffled: torch.Tensor,- topk_weights: torch.Tensor,- topk_ids: torch.Tensor,- gate_up_weight_scale_shuffled: torch.Tensor,- down_weight_scale_shuffled: torch.Tensor,- hidden_pad: int,- intermediate_pad: int,- block_size_m: int | None,- ):- fused_moe_module = _FUSED_MOE_MODULE- if fused_moe_module is None:- return _call_moe(- hidden_states,- gate_up_weight_shuffled,- down_weight_shuffled,- topk_weights,- topk_ids,- gate_up_weight_scale_shuffled,- down_weight_scale_shuffled,- hidden_pad,- intermediate_pad,- block_size_m,- )- try:- get_inter_dim = getattr(fused_moe_module, "get_inter_dim")- dtypes = getattr(fused_moe_module, "dtypes")- quant_remap = getattr(fused_moe_module, "quant_remap")- get_gfx = getattr(fused_moe_module, "get_gfx")- get_padded_M = getattr(fused_moe_module, "get_padded_M")- get_2stage_cfgs = getattr(fused_moe_module, "get_2stage_cfgs")- moe_sorting = getattr(fused_moe_module, "moe_sorting")- fused_moe_2stages = getattr(fused_moe_module, "fused_moe_2stages")-- activation = ActivationType.Silu- quant_type = QuantType.per_1x32- M, topk = topk_ids.shape- E, model_dim, inter_dim = get_inter_dim(- gate_up_weight_shuffled.shape,- down_weight_shuffled.shape,- )-- assert gate_up_weight_shuffled.shape[1] in [inter_dim, inter_dim * 2]- is_g1u1 = inter_dim != gate_up_weight_shuffled.shape[1]- is_shuffled = getattr(gate_up_weight_shuffled, "is_shuffled", False)- dtype = hidden_states.dtype- global_e = E- q_dtype_w = gate_up_weight_shuffled.dtype- q_dtype_a = gate_up_weight_shuffled.dtype if gate_up_weight_shuffled.dtype != torch.uint32 else dtypes.fp8- quant_type = quant_remap.get(quant_type, quant_type)- if quant_type == QuantType.per_1x32:- if activation == ActivationType.Swiglu:- if get_gfx() != "gfx950" or M < 512:- q_dtype_a = dtypes.bf16- else:- q_dtype_a = dtypes.fp8- else:- q_dtype_a = dtypes.fp4x2-- metadata = get_2stage_cfgs(- get_padded_M(M),- model_dim,- inter_dim,- E,- topk,- dtype,- q_dtype_a,- q_dtype_w,- quant_type,- is_g1u1,- activation,- False,- hidden_pad,- intermediate_pad,- is_shuffled,- )-- block_size_eff = getattr(metadata, "block_m", None) if block_size_m is None else block_size_m- if block_size_eff is not None:- block_size_eff = int(block_size_eff)-- sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf = moe_sorting(- topk_ids,- topk_weights,- global_e,- model_dim,- dtype,- block_size_eff,- None,- None,- 0,- )-- if getattr(metadata, "run_1stage", False):- return metadata.stage1(- hidden_states,- gate_up_weight_shuffled,- down_weight_shuffled,- topk,- sorted_ids,- sorted_weights,- sorted_expert_ids,- num_valid_ids,- moe_buf,- is_g1u1,- block_size_eff,- q_dtype_a=q_dtype_a,- q_dtype_w=q_dtype_w,- w1_scale=gate_up_weight_scale_shuffled,- w2_scale=down_weight_scale_shuffled,- a1_scale=None,- a2_scale=None,- num_local_tokens=None,- M=M,- device=topk_ids.device,- doweight_stage1=False,- )-- return fused_moe_2stages(- hidden_states,- gate_up_weight_shuffled,- down_weight_shuffled,- topk,- sorted_ids,- sorted_weights,- sorted_expert_ids,- num_valid_ids,- moe_buf,- is_g1u1,- block_size_eff,- activation=activation,- quant_type=quant_type,- doweight_stage1=False,- q_dtype_a=q_dtype_a,- q_dtype_w=q_dtype_w,- w1_scale=gate_up_weight_scale_shuffled,- w2_scale=down_weight_scale_shuffled,- a1_scale=None,- a2_scale=None,- num_local_tokens=None,- hidden_pad=hidden_pad,- intermediate_pad=intermediate_pad,- bias1=None,- bias2=None,- )- except Exception:- return _call_moe(- hidden_states,- gate_up_weight_shuffled,- down_weight_shuffled,- topk_weights,- topk_ids,- gate_up_weight_scale_shuffled,- down_weight_scale_shuffled,- hidden_pad,- intermediate_pad,- block_size_m,- )--- @torch.inference_mode()+ # ------------------------------------------------------------------ #+ # Kernel entry point #+ # ------------------------------------------------------------------ #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-- _install_patch()-- hidden_pad = int(config["d_hidden_pad"]) - int(config["d_hidden"])- intermediate_pad = int(config["d_expert_pad"]) - int(config["d_expert"])- block_size_m = _choose_block_size_m(hidden_states, gate_up_weight_shuffled)-- if not _should_use_direct_path(hidden_states, gate_up_weight_shuffled, topk_ids, config):- 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,- block_size_M=block_size_m,- hidden_pad=hidden_pad,- intermediate_pad=intermediate_pad,- )-- return _call_moe_inlined(- hidden_states,- gate_up_weight_shuffled,- down_weight_shuffled,- topk_weights,- topk_ids,- gate_up_weight_scale_shuffled,- down_weight_scale_shuffled,- hidden_pad,- intermediate_pad,- block_size_m,+ _warmup(data)+ (hs, _guw, _dw, _guws, _dws,+ guw_sh, dw_sh, guws_sh, dws_sh, tw, ti, cfg) = data+ return fused_moe(+ hs, guw_sh, dw_sh, tw, ti,+ activation=ActivationType.Silu,+ quant_type=QuantType.per_1x32,+ w1_scale=guws_sh,+ w2_scale=dws_sh,+ hidden_pad=int(cfg["d_hidden_pad"]) - int(cfg["d_hidden"]),+ intermediate_pad=int(cfg["d_expert_pad"]) - int(cfg["d_expert"]),)
scrolls · 989 diff lines total
Best evidence level for this revision: reported
JSON