submission 697388
.jonnss · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 325 lines, June 9 Researcher Reciprocity License v1.0.
Submission_v290.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-697388?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:1d004d6dc15c1571b164adceaa235d5797276236261e5c7029c8201390e112b6
license declaredunknown
license concludedunknown
authors.jonnss
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
Kernel source
Submission_v290.py325 lines
import dataclasses
import functools
import gc
import inspect
import importlib
import os
import torch
from task import input_t, output_t
os.environ["VLLM_MOE_CHUNK_SIZE"] = "512"
os.environ["HIP_FORCE_DEV_KERNARG"] = "1"
os.environ["AITER_BYPASS_TUNE_CONFIG"] = "1"
os.environ["AITER_KSPLIT"] = "2"
from aiter import ActivationType, QuantType
from aiter.fused_moe import fused_moe
gc.disable()
_PATCH_DONE = False
_CURRENT_SORT_BLOCK_SIZE = 32
_CURRENT_SORT_EXPERT = -1
_ADAPTIVE_QUANT_THRESHOLD = 1_000_000
_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,
}
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
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 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,
)
except Exception:
pass
return meta
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"
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)
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)
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,
mode="atomic",
MPerBlock=16,
)
kernel_params[_FLYDSL_STAGE2_KEY] = base
def _precompile_flydsl_defs(moe_kernels) -> None:
get_compiled_stage2 = getattr(moe_kernels, "_get_compiled_stage2", None)
if get_compiled_stage2 is None:
return
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
def _choose_block_size_m(hidden_states: torch.Tensor, gate_up_weight_shuffled: torch.Tensor) -> int | None:
m = int(hidden_states.shape[0])
experts = int(gate_up_weight_shuffled.shape[0])
if m <= 16:
return 16
if m <= 128 and experts > 64:
return 32
return None
def _desired_ksplit(token: int, expert: int) -> int:
if token <= 16:
return 2
if token <= 128 and expert > 64:
return 2
return 0
def _install_patch() -> None:
global _PATCH_DONE
if _PATCH_DONE:
return
_PATCH_DONE = True
try:
aiter_mod = importlib.import_module("aiter")
fused_moe_module = importlib.import_module("aiter.fused_moe")
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)
original_cfgs_sig = inspect.signature(original_cfgs) if original_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
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)
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):
meta = original_cfgs(*args, **kwargs)
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)
ksplit = _desired_ksplit(token, expert)
block_m = getattr(meta, "block_m", None)
stage1 = _retune_partial(getattr(meta, "stage1", None), splitk=ksplit)
stage2 = getattr(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)
return _replace_meta(meta, block_m=block_m, ksplit=ksplit, stage1=stage1, stage2=stage2)
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),
):
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
except Exception:
pass
@torch.inference_mode()
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)
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,
)
scrolls · 325 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 692250.
+ import dataclassesimport functools+ import gc+ import inspectimport importlibimport os⋯ 2 unchanged linesfrom task import input_t, output_tos.environ["VLLM_MOE_CHUNK_SIZE"] = "512"+ os.environ["HIP_FORCE_DEV_KERNARG"] = "1"+ os.environ["AITER_BYPASS_TUNE_CONFIG"] = "1"+ os.environ["AITER_KSPLIT"] = "2"from aiter import ActivationType, QuantTypefrom aiter.fused_moe import fused_moe+ gc.disable()++_PATCH_DONE = False+ _CURRENT_SORT_BLOCK_SIZE = 32+ _CURRENT_SORT_EXPERT = -1+ _ADAPTIVE_QUANT_THRESHOLD = 1_000_000+ _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,+ }+ 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++ 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 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,+ )+ except Exception:+ pass++ return meta+++ 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"+++ 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)+++ 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)+++ 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,+ mode="atomic",+ MPerBlock=16,+ )+ kernel_params[_FLYDSL_STAGE2_KEY] = base+++ def _precompile_flydsl_defs(moe_kernels) -> None:+ get_compiled_stage2 = getattr(moe_kernels, "_get_compiled_stage2", None)+ if get_compiled_stage2 is None:+ return++ 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+++ def _choose_block_size_m(hidden_states: torch.Tensor, gate_up_weight_shuffled: torch.Tensor) -> int | None:+ m = int(hidden_states.shape[0])+ experts = int(gate_up_weight_shuffled.shape[0])+ if m <= 16:+ return 16+ if m <= 128 and experts > 64:+ return 32+ return None+++ def _desired_ksplit(token: int, expert: int) -> int:+ if token <= 16:+ return 2+ if token <= 128 and expert > 64:+ return 2+ return 0++def _install_patch() -> None:global _PATCH_DONEif _PATCH_DONE:⋯ 1 unchanged lines_PATCH_DONE = Truetry:+ aiter_mod = importlib.import_module("aiter")fused_moe_module = importlib.import_module("aiter.fused_moe")- original = getattr(fused_moe_module, "get_2stage_cfgs", None)- if original is None:- return+ fp4_utils = importlib.import_module("aiter.utility.fp4_utils")+ moe_kernels = importlib.import_module("aiter.ops.flydsl.moe_kernels")- def patched_get_2stage_cfgs(*args, **kwargs):- meta = original(*args, **kwargs)+ _register_flydsl_t16(moe_kernels)+ _precompile_flydsl_defs(moe_kernels)- token = kwargs.get("token", args[0] if len(args) > 0 else None)- model_dim = kwargs.get("model_dim", args[1] if len(args) > 1 else None)- inter_dim = kwargs.get("inter_dim", args[2] if len(args) > 2 else None)- expert = kwargs.get("expert", args[3] if len(args) > 3 else None)- topk = kwargs.get("topk", args[4] if len(args) > 4 else None)+ original_cfgs = getattr(fused_moe_module, "get_2stage_cfgs", None)+ original_cfgs_sig = inspect.signature(original_cfgs) if original_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+ 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)+ if original_cfgs is None or flydsl_stage2_wrapper is None:+ return- if token == 512 and model_dim == 7168 and inter_dim == 512 and expert == 33 and topk == 9:+ def _extract_int(+ sig: inspect.Signature | None,+ call_args,+ call_kwargs,+ name: str,+ fallback_index: int,+ ) -> int:+ if sig is not None:try:- meta.block_m = 128+ bound_args = sig.bind_partial(*call_args, **call_kwargs).arguments+ if name in bound_args:+ return int(bound_args[name])except Exception:pass- for attr in ("stage1", "stage2"):- fn = getattr(meta, attr, None)- if isinstance(fn, functools.partial):- keywords = dict(fn.keywords or {})- keywords["block_m"] = 128- setattr(meta, attr, functools.partial(fn.func, *(fn.args or ()), **keywords))- return meta+ 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):+ meta = original_cfgs(*args, **kwargs)+ 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)++ ksplit = _desired_ksplit(token, expert)+ block_m = getattr(meta, "block_m", None)+ stage1 = _retune_partial(getattr(meta, "stage1", None), splitk=ksplit)+ stage2 = getattr(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)++ return _replace_meta(meta, block_m=block_m, ksplit=ksplit, stage1=stage1, stage2=stage2)++ 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- fused_moe_globals = getattr(fused_moe, "__globals__", None)- if isinstance(fused_moe_globals, dict) and fused_moe_globals.get("get_2stage_cfgs") is original:- fused_moe_globals["get_2stage_cfgs"] = patched_get_2stage_cfgs+ for maybe_globals in (+ getattr(fused_moe, "__globals__", None),+ getattr(fused_moe_internal, "__globals__", None),+ getattr(fused_moe_2stages, "__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_ksplitexcept Exception:pass⋯ 19 unchanged lineshidden_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)return fused_moe(hidden_states,⋯ 9 unchanged linesw2_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,)
scrolls · 319 diff lines total
Best evidence level for this revision: reported
JSON