submission 749286
Hamza · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 221 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-749286?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:e7f9dcd5fdf50ce4677ee5f738087b1f09da5c8dc6bb83ce4a7b5313db72d8c2
license declaredunknown
license concludedunknown
authorsHamza
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
tile_m=32, tile_n=256, tile_k=128, a_dtype="fp4", b_dtype="fp4", out_dtype="bf16",Kernel source
submission.py221 lines
#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X
import gc
import os
import sys
import subprocess
import shutil
import re
import functools
sys.setswitchinterval(1.0)
os.environ["HIP_FORCE_DEV_KERNARG"] = "1"
os.environ["AITER_USE_NT"] = "1"
os.environ["HSA_ENABLE_INTERRUPT"] = "0"
_AITER_DIR = '/home/runner/aiter'
_NEW_DIR = '/tmp/aiter_fresh'
_FLYDSL_DIR = '/tmp/flydsl_new'
# ========== Install flydsl 0.1.1 ==========
os.makedirs(_FLYDSL_DIR, exist_ok=True)
subprocess.run([sys.executable, '-m', 'pip', 'install', 'flydsl==0.1.1',
'--target', _FLYDSL_DIR, '--upgrade', '--break-system-packages'],
capture_output=True, text=True, timeout=120)
sys.path.insert(0, _FLYDSL_DIR)
for key in list(sys.modules.keys()):
if 'flydsl' in key:
del sys.modules[key]
# ========== Clone AITER + pin to old API ==========
subprocess.run(['git', 'checkout', '--', '.'], capture_output=True, cwd=_AITER_DIR, timeout=10)
if not os.path.exists(os.path.join(_NEW_DIR, 'aiter')):
subprocess.run(['git', 'clone', '--depth=50', 'https://github.com/ROCm/aiter.git', _NEW_DIR],
capture_output=True, text=True, timeout=120)
_mk_path = os.path.join(_NEW_DIR, 'aiter', 'ops', 'flydsl', 'moe_kernels.py')
if os.path.exists(_mk_path):
with open(_mk_path) as f:
if '_get_compiled_stage1' not in f.read():
r = subprocess.run(['git', 'log', '--format=%H', '-50'],
capture_output=True, text=True, cwd=_NEW_DIR, timeout=10)
for commit in r.stdout.strip().split('\n')[1:]:
if not commit.strip(): continue
subprocess.run(['git', 'checkout', commit.strip()], capture_output=True, cwd=_NEW_DIR, timeout=10)
with open(_mk_path) as f2:
if '_get_compiled_stage1' in f2.read():
print(f"[moe] Pinned to {commit.strip()[:12]}", file=sys.stderr)
break
# Copy NEWER fused_moe.py (fixes stage1 calling convention)
shutil.copy2(os.path.join(_NEW_DIR, 'aiter', 'fused_moe.py'),
os.path.join(_AITER_DIR, 'aiter', 'fused_moe.py'))
# Copy FlyDSL kernel files
flydsl_src = os.path.join(_NEW_DIR, 'aiter', 'ops', 'flydsl')
flydsl_dst = os.path.join(_AITER_DIR, 'aiter', 'ops', 'flydsl')
for fname in ['moe_kernels.py', 'utils.py']:
shutil.copy2(os.path.join(flydsl_src, fname), os.path.join(flydsl_dst, fname))
init_path = os.path.join(flydsl_dst, '__init__.py')
with open(init_path) as f:
init_src = f.read()
init_src = re.sub(r'raise ImportError\([^)]*\)', 'pass # version check bypassed', init_src)
with open(init_path, 'w') as f:
f.write(init_src)
kernels_src = os.path.join(flydsl_src, 'kernels')
kernels_dst = os.path.join(flydsl_dst, 'kernels')
if os.path.exists(kernels_src):
os.makedirs(kernels_dst, exist_ok=True)
for fname in os.listdir(kernels_src):
if fname.endswith('.py'):
shutil.copy2(os.path.join(kernels_src, fname), os.path.join(kernels_dst, fname))
# ========== Source-patch fused_moe.py ==========
_FMOE_PATH = os.path.join(_AITER_DIR, 'aiter', 'fused_moe.py')
with open(_FMOE_PATH, 'r') as f:
src = f.read()
# Append: FlyDSL stage1 for E≤100 ONLY, stage2 FlyDSL for all
tail = r'''
# FlyDSL stage1 E<=100 only, CK stage1 for E=257
def _fly_s1(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 _fly_s2_atomic(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=16, 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)
_fly_orig_g2sc = get_2stage_cfgs
try: _fly_orig_g2sc.cache_clear()
except: pass
_fly_ready = [False]
_fly_s1_ready = [False]
import functools as _ft
@_ft.lru_cache(maxsize=2048)
def _fly_g2sc(*a, **kw):
md = _fly_orig_g2sc(*a, **kw)
E_val = a[4] if len(a) >= 5 else 0
# Stage1: FlyDSL ONLY for E<=100, CK for E=257
if _fly_s1_ready[0] and E_val <= 100 and not md.run_1stage and md.stage1 is not None:
md.stage1 = _ft.partial(_fly_s1)
# Stage2: FlyDSL for ALL shapes
if _fly_ready[0] and not md.run_1stage and md.stage2 is not None:
md.stage2 = _ft.partial(_fly_s2_atomic)
return md
get_2stage_cfgs = _fly_g2sc
import sys as _s
print("[moe] FlyDSL stage1 E<=100, stage2 all — source-patched", file=_s.stderr)
'''
src += tail
with open(_FMOE_PATH, 'w') as f:
f.write(src)
print("[moe] Source patches applied", file=sys.stderr)
# ========== Fix Triton constexpr ==========
try:
import triton.language.core as _tlc
_ce = _tlc.constexpr
try:
_ce(1).__lt__(_ce(2), _semantic=None)
except TypeError:
for _mn in ['__lt__', '__le__', '__gt__', '__ge__', '__eq__', '__ne__',
'__add__', '__radd__', '__sub__', '__rsub__', '__mul__', '__rmul__',
'__truediv__', '__floordiv__', '__mod__', '__pow__',
'__lshift__', '__rshift__', '__and__', '__or__', '__xor__']:
_ofn = getattr(_ce, _mn, None)
if _ofn is not None and not getattr(_ofn, '_patched', False):
def _mkw(fn):
def _w(self, *a, _semantic=None, **kw):
return fn(self, *a, **kw)
_w._patched = True
return _w
setattr(_ce, _mn, _mkw(_ofn))
del _ce, _tlc
except Exception:
pass
# ========== Import AITER (with newer fused_moe.py, source-patched) ==========
import torch
torch.set_grad_enabled(False)
from task import input_t, output_t
from aiter import ActivationType, QuantType
from aiter.fused_moe import fused_moe
import aiter.fused_moe as _fmoe
from aiter.ops.quant import per_1x32_f4_quant_hip as _quant_hip
from aiter.utility import fp4_utils as _fp4_utils
# ========== Pre-compile FlyDSL ==========
try:
from aiter.ops.flydsl.moe_kernels import _get_compiled_stage1, _get_compiled_stage2
# Stage1 ONLY for E=33
for inter, exp in [(512, 33), (2048, 33)]:
try:
_get_compiled_stage1(model_dim=7168, inter_dim=inter, experts=exp, topk=9,
tile_m=32, tile_n=256, tile_k=128, doweight=True,
a_dtype="fp4", b_dtype="fp4", out_dtype="bf16", act="silu")
print(f"[moe] Stage1 compiled: d={inter} E={exp}", file=sys.stderr)
except Exception as e:
print(f"[moe] Stage1 FAILED: d={inter} E={exp}: {str(e)[:80]}", file=sys.stderr)
_fmoe._fly_s1_ready[0] = True
# Stage2 for all
for inter, exp in [(256, 257), (512, 33), (2048, 33)]:
try:
_get_compiled_stage2(model_dim=7168, inter_dim=inter, experts=exp, topk=9,
tile_m=16, tile_n=256, tile_k=128, doweight=True,
a_dtype="fp4", b_dtype="fp4", out_dtype="bf16", accumulate=True)
except: pass
_fmoe._fly_ready[0] = True
print("[moe] FlyDSL pre-compiled", file=sys.stderr)
except Exception as e:
print(f"[moe] Pre-compile failed: {e}", file=sys.stderr)
# Warmup
_warmed = set()
def _warmup(data):
(hs, guw, dw, guws, dws, guw_sh, dw_sh, guws_sh, dws_sh, tw, ti, cfg) = data
key = (cfg["bs"], cfg["n_routed_experts"], cfg["d_expert"])
if key in _warmed: return
_warmed.add(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=cfg["d_hidden_pad"]-cfg["d_hidden"],
intermediate_pad=cfg["d_expert_pad"]-cfg["d_expert"])
torch.cuda.synchronize()
except: pass
gc.disable()
print("[moe] Ready", file=sys.stderr)
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=cfg["d_hidden_pad"]-cfg["d_hidden"],
intermediate_pad=cfg["d_expert_pad"]-cfg["d_expert"])
scrolls · 221 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 704012.
#!POPCORN leaderboard amd-moe-mxfp4#!POPCORN gpu MI355X- # v90: Fix shape 5 bypass regression — restore per-E routing from v55/v67- # Shape 5 (M=128, E=32) must use 2-stage FlyDSL, NOT bypass (~15µs faster)-import gcimport osimport sys+ import subprocess+ import shutil+ import re+ import functools+ sys.setswitchinterval(1.0)+os.environ["HIP_FORCE_DEV_KERNARG"] = "1"os.environ["AITER_USE_NT"] = "1"os.environ["HSA_ENABLE_INTERRUPT"] = "0"- import torch- from task import input_t, output_t+ _AITER_DIR = '/home/runner/aiter'+ _NEW_DIR = '/tmp/aiter_fresh'+ _FLYDSL_DIR = '/tmp/flydsl_new'- from aiter import ActivationType, QuantType- from aiter.fused_moe import fused_moe, moe_sorting- import aiter.fused_moe as _fm+ # ========== Install flydsl 0.1.1 ==========+ os.makedirs(_FLYDSL_DIR, exist_ok=True)+ subprocess.run([sys.executable, '-m', 'pip', 'install', 'flydsl==0.1.1',+ '--target', _FLYDSL_DIR, '--upgrade', '--break-system-packages'],+ capture_output=True, text=True, timeout=120)+ sys.path.insert(0, _FLYDSL_DIR)+ for key in list(sys.modules.keys()):+ if 'flydsl' in key:+ del sys.modules[key]- from aiter.ops.quant import per_1x32_f4_quant_hip as _quant_hip- from aiter.utility import fp4_utils as _fp4_utils+ # ========== Clone AITER + pin to old API ==========+ subprocess.run(['git', 'checkout', '--', '.'], capture_output=True, cwd=_AITER_DIR, timeout=10)- # Adaptive quant: split HIP for large tensors, fused Triton for small- _original_fused_quant = _fm.fused_dynamic_mxfp4_quant_moe_sort+ if not os.path.exists(os.path.join(_NEW_DIR, 'aiter')):+ subprocess.run(['git', 'clone', '--depth=50', 'https://github.com/ROCm/aiter.git', _NEW_DIR],+ capture_output=True, text=True, timeout=120)- def _adaptive_quant_moe_sort(x, sorted_ids, num_valid_ids, token_num, topk, block_size=32, scaling_mode="even"):- if x.numel() > 1_000_000:- x_fp4, scale = _quant_hip(x)- sorted_scale = _fp4_utils.moe_mxfp4_sort(- scale, sorted_ids=sorted_ids, num_valid_ids=num_valid_ids,- token_num=token_num, block_size=block_size,- )- return x_fp4, sorted_scale- else:- return _original_fused_quant(x, sorted_ids, num_valid_ids, token_num, topk, block_size, scaling_mode)+ _mk_path = os.path.join(_NEW_DIR, 'aiter', 'ops', 'flydsl', 'moe_kernels.py')+ if os.path.exists(_mk_path):+ with open(_mk_path) as f:+ if '_get_compiled_stage1' not in f.read():+ r = subprocess.run(['git', 'log', '--format=%H', '-50'],+ capture_output=True, text=True, cwd=_NEW_DIR, timeout=10)+ for commit in r.stdout.strip().split('\n')[1:]:+ if not commit.strip(): continue+ subprocess.run(['git', 'checkout', commit.strip()], capture_output=True, cwd=_NEW_DIR, timeout=10)+ with open(_mk_path) as f2:+ if '_get_compiled_stage1' in f2.read():+ print(f"[moe] Pinned to {commit.strip()[:12]}", file=sys.stderr)+ break- _fm.fused_dynamic_mxfp4_quant_moe_sort = _adaptive_quant_moe_sort- print("[moe] Adaptive quant: split for numel>1M, fused for numel<=1M", file=sys.stderr)+ # Copy NEWER fused_moe.py (fixes stage1 calling convention)+ shutil.copy2(os.path.join(_NEW_DIR, 'aiter', 'fused_moe.py'),+ os.path.join(_AITER_DIR, 'aiter', 'fused_moe.py'))- # Register FlyDSL t16x256x128 kernel params- try:- from aiter.ops.flydsl.moe_kernels import _KERNEL_PARAMS- _t16_name = "flydsl_moe2_afp4_wfp4_bf16_t16x256x128_atomic"- if _t16_name not in _KERNEL_PARAMS:- _KERNEL_PARAMS[_t16_name] = {- "stage": 2, "a_dtype": "fp4", "b_dtype": "fp4", "out_dtype": "bf16",- "tile_m": 16, "tile_n": 256, "tile_k": 128,- "mode": "atomic", "MPerBlock": 16,- }- except Exception as e:- print(f"[moe] _KERNEL_PARAMS registration failed: {e}", file=sys.stderr)+ # Copy FlyDSL kernel files+ flydsl_src = os.path.join(_NEW_DIR, 'aiter', 'ops', 'flydsl')+ flydsl_dst = os.path.join(_AITER_DIR, 'aiter', 'ops', 'flydsl')+ for fname in ['moe_kernels.py', 'utils.py']:+ shutil.copy2(os.path.join(flydsl_src, fname), os.path.join(flydsl_dst, fname))+ init_path = os.path.join(flydsl_dst, '__init__.py')+ with open(init_path) as f:+ init_src = f.read()+ init_src = re.sub(r'raise ImportError\([^)]*\)', 'pass # version check bypassed', init_src)+ with open(init_path, 'w') as f:+ f.write(init_src)+ kernels_src = os.path.join(flydsl_src, 'kernels')+ kernels_dst = os.path.join(flydsl_dst, 'kernels')+ if os.path.exists(kernels_src):+ os.makedirs(kernels_dst, exist_ok=True)+ for fname in os.listdir(kernels_src):+ if fname.endswith('.py'):+ shutil.copy2(os.path.join(kernels_src, fname), os.path.join(kernels_dst, fname))- # CSV append for shapes 5,6,7 with tuned CK stage1 + FlyDSL t16 stage2- _tune_file = _fm.AITER_CONFIGS.AITER_CONFIG_FMOE_FILE- try:- with open(_tune_file, "a") as f:- # Shape 5: block_m=64, CK 256x64 stage1- f.write("256,128,7168,512,33,9,ActivationType.Silu,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,64,0,0,moe_ck2stages_gemm1_256x64x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16,0.0%,0,flydsl_moe2_afp4_wfp4_bf16_t16x256x128_atomic,0.0%,0,0,0,0\n")- # Shape 6: block_m=128, CK 256x128 stage1- f.write("256,512,7168,512,33,9,ActivationType.Silu,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,128,0,0,moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16,0.0%,0,flydsl_moe2_afp4_wfp4_bf16_t16x256x128_atomic,0.0%,0,0,0,0\n")- # Shape 7: block_m=64, CK 256x64 stage1- f.write("256,512,7168,2048,33,9,ActivationType.Silu,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,64,0,0,moe_ck2stages_gemm1_256x64x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16,0.0%,0,flydsl_moe2_afp4_wfp4_bf16_t16x256x128_atomic,0.0%,0,0,0,0\n")- except Exception as e:- print(f"[moe] CSV append failed: {e}", file=sys.stderr)+ # ========== Source-patch fused_moe.py ==========+ _FMOE_PATH = os.path.join(_AITER_DIR, 'aiter', 'fused_moe.py')+ with open(_FMOE_PATH, 'r') as f:+ src = f.read()- # Pre-compile FlyDSL MLIR for both inter_dim variants- try:- from aiter.ops.flydsl.moe_kernels import _get_compiled_stage2- _get_compiled_stage2(model_dim=7168, inter_dim=512, experts=33, topk=9, tile_m=16, tile_n=256, tile_k=128, doweight=False, a_dtype="fp4", b_dtype="fp4", out_dtype="bf16", accumulate=True)- _get_compiled_stage2(model_dim=7168, inter_dim=2048, experts=33, topk=9, tile_m=16, tile_n=256, tile_k=128, doweight=False, a_dtype="fp4", b_dtype="fp4", out_dtype="bf16", accumulate=True)- except Exception:- pass+ # Append: FlyDSL stage1 for E≤100 ONLY, stage2 FlyDSL for all+ tail = r'''- # GPU binary warmup: trigger @flyc.jit compilation with dummy tensors- try:+ # FlyDSL stage1 E<=100 only, CK stage1 for E=257+ def _fly_s1(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 _fly_s2_atomic(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- _dev = "cuda"- for _inter_dim in (512, 2048):- _ip = _inter_dim // 2- _di = torch.zeros((16, 9, _ip), dtype=torch.uint8, device=_dev)- _dw = torch.zeros((33, 7168, _ip), dtype=torch.uint8, device=_dev)- _dws = torch.ones((33, 7168, _inter_dim // 32), dtype=torch.uint8, device=_dev)- _das = torch.ones((144, _inter_dim // 32), dtype=torch.uint8, device=_dev)- _dsi = torch.arange(16, dtype=torch.int32, device=_dev)- _dei = torch.zeros(1, dtype=torch.int32, device=_dev)- _dnv = torch.tensor([16] + [0] * 32, dtype=torch.int32, device=_dev)- flydsl_moe_stage2(_di, _dw, _dsi, _dei, _dnv, topk=9, tile_m=16, tile_n=256, tile_k=128, a_dtype="fp4", b_dtype="fp4", out_dtype="bf16", mode="atomic", w2_scale=_dws, a2_scale=_das)- del _di, _dw, _dws, _das, _dsi, _dei, _dnv- torch.cuda.empty_cache()- except Exception:- pass+ 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=16, 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)- # HIP quant kernel warmup+ _fly_orig_g2sc = get_2stage_cfgs+ try: _fly_orig_g2sc.cache_clear()+ except: pass+ _fly_ready = [False]+ _fly_s1_ready = [False]++ import functools as _ft++ @_ft.lru_cache(maxsize=2048)+ def _fly_g2sc(*a, **kw):+ md = _fly_orig_g2sc(*a, **kw)+ E_val = a[4] if len(a) >= 5 else 0+ # Stage1: FlyDSL ONLY for E<=100, CK for E=257+ if _fly_s1_ready[0] and E_val <= 100 and not md.run_1stage and md.stage1 is not None:+ md.stage1 = _ft.partial(_fly_s1)+ # Stage2: FlyDSL for ALL shapes+ if _fly_ready[0] and not md.run_1stage and md.stage2 is not None:+ md.stage2 = _ft.partial(_fly_s2_atomic)+ return md+ get_2stage_cfgs = _fly_g2sc++ import sys as _s+ print("[moe] FlyDSL stage1 E<=100, stage2 all — source-patched", file=_s.stderr)+ '''+ src += tail+ with open(_FMOE_PATH, 'w') as f:+ f.write(src)++ print("[moe] Source patches applied", file=sys.stderr)++ # ========== Fix Triton constexpr ==========try:- _d = torch.randn(512, 7168, dtype=torch.bfloat16, device="cuda")- _quant_hip(_d)- del _d- torch.cuda.empty_cache()+ import triton.language.core as _tlc+ _ce = _tlc.constexpr+ try:+ _ce(1).__lt__(_ce(2), _semantic=None)+ except TypeError:+ for _mn in ['__lt__', '__le__', '__gt__', '__ge__', '__eq__', '__ne__',+ '__add__', '__radd__', '__sub__', '__rsub__', '__mul__', '__rmul__',+ '__truediv__', '__floordiv__', '__mod__', '__pow__',+ '__lshift__', '__rshift__', '__and__', '__or__', '__xor__']:+ _ofn = getattr(_ce, _mn, None)+ if _ofn is not None and not getattr(_ofn, '_patched', False):+ def _mkw(fn):+ def _w(self, *a, _semantic=None, **kw):+ return fn(self, *a, **kw)+ _w._patched = True+ return _w+ setattr(_ce, _mn, _mkw(_ofn))+ del _ce, _tlcexcept Exception:pass- # moe_sorting kernel warmup (trigger JIT compile for both E configs)- try:- for _E in (33, 257):- _ids = torch.zeros((16, 9), dtype=torch.int32, device="cuda")- _wts = torch.ones((16, 9), dtype=torch.float32, device="cuda")- moe_sorting(_ids, _wts, _E, 7168, torch.bfloat16, 32)- del _ids, _wts- torch.cuda.empty_cache()- print("[moe] moe_sorting warmed up", file=sys.stderr)- except Exception as e:- print(f"[moe] moe_sorting warmup failed: {e}", file=sys.stderr)+ # ========== Import AITER (with newer fused_moe.py, source-patched) ==========+ import torch+ torch.set_grad_enabled(False)+ from task import input_t, output_t- # Disable GC to prevent pauses during benchmark- gc.disable()- print("[moe] v90: fix shape5 routing + warmup + gc.disable", file=sys.stderr)+ from aiter import ActivationType, QuantType+ from aiter.fused_moe import fused_moe+ import aiter.fused_moe as _fmoe+ from aiter.ops.quant import per_1x32_f4_quant_hip as _quant_hip+ from aiter.utility import fp4_utils as _fp4_utils- _PAD_CACHE = {}+ # ========== Pre-compile FlyDSL ==========+ try:+ from aiter.ops.flydsl.moe_kernels import _get_compiled_stage1, _get_compiled_stage2+ # Stage1 ONLY for E=33+ for inter, exp in [(512, 33), (2048, 33)]:+ try:+ _get_compiled_stage1(model_dim=7168, inter_dim=inter, experts=exp, topk=9,+ tile_m=32, tile_n=256, tile_k=128, doweight=True,+ a_dtype="fp4", b_dtype="fp4", out_dtype="bf16", act="silu")+ print(f"[moe] Stage1 compiled: d={inter} E={exp}", file=sys.stderr)+ except Exception as e:+ print(f"[moe] Stage1 FAILED: d={inter} E={exp}: {str(e)[:80]}", file=sys.stderr)- 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+ _fmoe._fly_s1_ready[0] = True- M = topk_ids.shape[0]+ # Stage2 for all+ for inter, exp in [(256, 257), (512, 33), (2048, 33)]:+ try:+ _get_compiled_stage2(model_dim=7168, inter_dim=inter, experts=exp, topk=9,+ tile_m=16, tile_n=256, tile_k=128, doweight=True,+ a_dtype="fp4", b_dtype="fp4", out_dtype="bf16", accumulate=True)+ except: pass- cfg_id = id(config)- cached = _PAD_CACHE.get(cfg_id)- if cached is not None:- hidden_pad, intermediate_pad = cached- else:- hidden_pad = config["d_hidden_pad"] - config["d_hidden"]- intermediate_pad = config["d_expert_pad"] - config["d_expert"]- _PAD_CACHE[cfg_id] = (hidden_pad, intermediate_pad)+ _fmoe._fly_ready[0] = True+ print("[moe] FlyDSL pre-compiled", file=sys.stderr)+ except Exception as e:+ print(f"[moe] Pre-compile failed: {e}", file=sys.stderr)- E = gate_up_weight_shuffled.shape[0] # E+1 (includes shared expert)+ # Warmup+ _warmed = set()+ def _warmup(data):+ (hs, guw, dw, guws, dws, guw_sh, dw_sh, guws_sh, dws_sh, tw, ti, cfg) = data+ key = (cfg["bs"], cfg["n_routed_experts"], cfg["d_expert"])+ if key in _warmed: return+ _warmed.add(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=cfg["d_hidden_pad"]-cfg["d_hidden"],+ intermediate_pad=cfg["d_expert_pad"]-cfg["d_expert"])+ torch.cuda.synchronize()+ except: pass- # Per-shape routing (v55/v67 condition):- # - Shapes 1,2 (M<=128, E=257): BYPASS (cktile faster for large E)- # - Shape 4 (M=16, E=33): BYPASS (too small for 2-stage overhead)- # - Shape 5 (M=128, E=33): 2-STAGE (FlyDSL t16 is ~15µs faster than bypass)- # - Shapes 3,6,7 (M>128): 2-STAGE (always)- if M <= 16 or (M <= 128 and E > 64):- os.environ["AITER_KSPLIT"] = "2"- os.environ["AITER_BYPASS_TUNE_CONFIG"] = "1"- else:- os.environ["AITER_KSPLIT"] = "0"- os.environ["AITER_BYPASS_TUNE_CONFIG"] = "0"+ gc.disable()+ print("[moe] Ready", file=sys.stderr)- return fused_moe(- hidden_states, gate_up_weight_shuffled, down_weight_shuffled,- topk_weights, topk_ids,- activation=ActivationType.Silu, quant_type=QuantType.per_1x32,- w1_scale=gate_up_weight_scale_shuffled,- w2_scale=down_weight_scale_shuffled,- hidden_pad=hidden_pad, intermediate_pad=intermediate_pad,- )+ 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=cfg["d_hidden_pad"]-cfg["d_hidden"],+ intermediate_pad=cfg["d_expert_pad"]-cfg["d_expert"])+
scrolls · 349 diff lines total
Best evidence level for this revision: reported
JSON