submission 699035
.jonnss · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 455 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-699035?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:3325dd9909d9164506db0326b8df63bf72e2ec9578d01306d68023791ebb0867
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.py455 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"
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
_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,
}
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 _use_bypass(token: int, expert: int) -> bool:
return token <= 16 or (token <= 128 and expert > 64)
def _desired_ksplit(token: int, expert: int) -> int:
return 2 if _use_bypass(token, expert) else 0
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:
try:
bound = sig.bind_partial(*call_args, **call_kwargs)
return (
use_bypass,
tuple((name, _cacheable(value)) for name, value in bound.arguments.items()),
)
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())),
)
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
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)
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)
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),
):
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),
):
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
@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 · 455 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 697388.
⋯ 10 unchanged linesos.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⋯ 6 unchanged lines_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"⋯ 143 unchanged linesreturn None+ def _use_bypass(token: int, expert: int) -> bool:+ return token <= 16 or (token <= 128 and expert > 64)++def _desired_ksplit(token: int, expert: int) -> int:- if token <= 16:- return 2- if token <= 128 and expert > 64:- return 2- return 0+ return 2 if _use_bypass(token, expert) else 0+ 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:+ try:+ bound = sig.bind_partial(*call_args, **call_kwargs)+ return (+ use_bypass,+ tuple((name, _cacheable(value)) for name, value in bound.arguments.items()),+ )+ 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())),+ )+++ 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++def _install_patch() -> None:global _PATCH_DONEif _PATCH_DONE:⋯ 10 unchanged lines_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 Noneoriginal_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)⋯ 21 unchanged linesreturn -1def patched_get_2stage_cfgs(*args, **kwargs):- meta = original_cfgs(*args, **kwargs)+ global _CURRENT_SORT_BLOCK_SIZE, _CURRENT_SORT_EXPERTtoken = _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)- 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- 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)- 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- return _replace_meta(meta, block_m=block_m, ksplit=ksplit, stage1=stage1, stage2=stage2)+ _CURRENT_SORT_BLOCK_SIZE = int(getattr(meta, "block_m", None) or 32)+ _CURRENT_SORT_EXPERT = expert+ return metadef patched_get_ksplit(*args, **kwargs):token = _extract_int(original_get_ksplit_sig, args, kwargs, "token", 0)⋯ 17 unchanged linesmaybe_globals["get_2stage_cfgs"] = patched_get_2stage_cfgsif 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),+ ):+ 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_sortexcept Exception:pass
scrolls · 234 diff lines total
Best evidence level for this revision: reported
JSON