submission 754180
Kernel-Zhang · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 3424 lines, June 9 Researcher Reciprocity License v1.0.
kmoe_00258.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-754180?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:4cd5fde0c066eecff9df7dc19192d3a3acce958357cd73a683b41d4fc1bcb074
license declaredunknown
license concludedunknown
authorsKernel-Zhang
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
"Stage2 boundary contract fp4 view does not match quant-sort shape."fused-epilogue
import aiter.ops.flydsl.kernels.mfma_epilogues as _mfma_epiloguesshared-memory
smem_ptr_lines = _collect_line_numbers(Kernel source
kmoe_00258.py3424 lines
from typing import Any, Dict, Optional, Tuple
import math
import sys
import torch
import torch.nn.functional as F
try:
from task import input_t, output_t
except ImportError:
input_t = Tuple[Any, ...]
output_t = torch.Tensor
try:
from utils import make_match_reference
except ImportError:
make_match_reference = None
PAD_ALIGN = 256
SHUFFLE_LAYOUT = (16, 16)
RTOL = 5e-2
ATOL = 5e-2
_AITER_IMPORT_ERROR = None
try:
import aiter
import triton
from aiter import ActivationType, QuantType, dtypes
from aiter.fused_moe import fused_moe
import aiter.ops.moe_op as _aiter_moe_op
from aiter.ops.triton._triton_kernels.quant.fused_mxfp4_quant import (
_fused_dynamic_mxfp4_quant_moe_sort_kernel,
)
from aiter.ops.shuffle import shuffle_weight
from aiter.utility import fp4_utils
_CK_MOE = getattr(aiter, "ck_moe", None)
except Exception as exc:
aiter = None
triton = None
ActivationType = None
QuantType = None
dtypes = None
fused_moe = None
_aiter_moe_op = None
_fused_dynamic_mxfp4_quant_moe_sort_kernel = None
shuffle_weight = None
fp4_utils = None
_CK_MOE = None
_AITER_IMPORT_ERROR = exc
_DEVICE_KIND_CACHE: Dict[Tuple[str, Optional[int]], bool] = {}
_BACKEND_CACHE: Dict[Tuple[Any, ...], Tuple[str, Optional[int]]] = {}
_CK_MOE_NT_KWARG: Optional[str] = None
_CK_MOE_NT_KWARG_READY = False
_CK_MOE_STAGE_KWARGS = None
_CK_MOE_STAGE_KWARGS_READY = False
_HIP_QUANT_PER_1X32 = None
_TORCH_QUANT_PER_1X32 = None
_LAST_HIDDEN_PREQUANT_KEY: Optional[Tuple[Any, ...]] = None
_LAST_HIDDEN_PREQUANT_VALUE: Optional[Tuple[torch.Tensor, torch.Tensor]] = None
_LAST_HIDDEN_PREQUANT_TORCH_KEY: Optional[Tuple[Any, ...]] = None
_LAST_HIDDEN_PREQUANT_TORCH_VALUE: Optional[Tuple[torch.Tensor, torch.Tensor]] = None
_LAST_FULL_EXACT_SORT_CACHE_SHAPE: Optional[Tuple[int, int, int, int, int]] = None
_KMOE_FULL_EXACT_SORT_CACHE: Dict[
Tuple[Any, ...],
Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor],
] = {}
_KMOE_FULL_EXACT_SORT_MISS_STREAK: Dict[Tuple[int, int, int, int, int], int] = {}
_KMOE_MOEBUF_WORKSPACE_CACHE: Dict[Tuple[Any, ...], torch.Tensor] = {}
_KMOE_MOEBUF_RING_WORKSPACE_CACHE: Dict[Tuple[Any, ...], Tuple[torch.Tensor, ...]] = {}
_KMOE_MOEBUF_RING_WORKSPACE_INDEX: Dict[Tuple[Any, ...], int] = {}
_LOGGED_EVENTS = set()
_MOE_DTYPE_ALIAS_PATCHED = False
_KMOE_PENDING_FLYDSL_STAGE2_PATCH_MODE = "ptr"
_KMOE_ACTIVE_FLYDSL_STAGE2_PATCH_MODE: Optional[str] = None
_KMOE_FLYDSL_STAGE2_ORIG_SOURCE: Optional[str] = None
_KMOE_FLYDSL_STAGE2_ORIG_SOURCE_FILE: Optional[str] = None
_KMOE_FLYDSL_STAGE2_PATCHED_FN_BY_MODE: Dict[str, Any] = {}
_KMOE_FLYDSL_STAGE2_COMPILED_CACHE: Dict[Tuple[Any, ...], Any] = {}
_KMOE_FLYDSL_STAGE2_ORIG_GET_COMPILED_STAGE2 = None
_KMOE_FLYDSL_STAGE2_COMPILED_CACHE_PATCHED = False
_KMOE_STAGE2_BOUNDARY_CONTRACT_WORKSPACE_CACHE: Dict[
Tuple[Any, ...],
Dict[str, torch.Tensor],
] = {}
_BENCHMARK_BLOCK_M = {
(16, 7168, 256, 257, 9): 32,
(128, 7168, 256, 257, 9): 64,
(512, 7168, 256, 257, 9): 64,
(16, 7168, 512, 33, 9): 32,
(128, 7168, 512, 33, 9): 64,
(512, 7168, 512, 33, 9): 128,
(512, 7168, 2048, 33, 9): 64,
}
_CK_FORCE_NON_TEMPORAL_TRUE = {
(16, 7168, 256, 257, 9),
(128, 7168, 256, 257, 9),
(512, 7168, 256, 257, 9),
}
_SKIP_SPLIT_SHARED_EXPERT_SHAPES = frozenset(
{
(16, 7168, 512, 33, 9),
(128, 7168, 512, 33, 9),
(512, 7168, 512, 33, 9),
(512, 7168, 2048, 33, 9),
}
)
_SELECTIVE_PREQUANT_CK_SHAPES = frozenset(
{
(16, 7168, 256, 257, 9),
(128, 7168, 256, 257, 9),
(512, 7168, 256, 257, 9),
}
)
_SELECTIVE_RAW_CK_SHAPES = frozenset(
{
(512, 7168, 256, 257, 9),
}
)
_CK_FORCE_STAGE_KERNELS = {
(512, 7168, 256, 257, 9): (
"moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",
"moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16",
64,
),
}
_FLYDSL_STAGE2_FALLBACK_KERNELS = {
(128, 7168, 256, 257, 9): "moe_ck2stages_gemm2_64x64x128x128_1x1_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16",
}
_FUSED_CFG_STAGE_OVERRIDES = {
(128, 7168, 256, 257, 9): (
"moe_ck2stages_gemm1_256x64x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",
"flydsl_moe2_afp4_wfp4_bf16_t64x256x256_atomic",
32,
),
(512, 7168, 256, 257, 9): (
"moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",
"moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16",
32,
),
}
_MANUAL_FUSED_2STAGE_SHAPES = frozenset(
{
(128, 7168, 256, 257, 9),
}
)
_KMOE_FULL_EXACT_SORT_CACHE_SHAPES = frozenset(
{
(128, 7168, 256, 257, 9),
}
)
_KMOE_MISS_CLONE_SORT_INPUTS_SHAPES = frozenset()
_KMOE_CACHE_HIT_NOCLONE_SORT_INPUTS_SHAPES = frozenset(
{
(128, 7168, 256, 257, 9),
}
)
_KMOE_MOEBUF_RING_WORKSPACE_SLOTS = {
(128, 7168, 256, 257, 9): 2,
}
_KMOE_FULL_EXACT_SORT_SKIP_STORE_AFTER_MISS_STREAK = {
(128, 7168, 256, 257, 9): 1,
}
_KMOE_FULL_EXACT_SORT_MAX_CACHE_ENTRIES = {
(128, 7168, 256, 257, 9): 1,
}
_MANUAL_FUSED_QUANT_SORT_BLOCK_SIZE_MX = {
(128, 7168, 256, 257, 9): {
"a1": 64,
"a2": 64,
},
}
_MANUAL_FUSED_STAGE2_BOUNDARY_CONTRACT_SHAPES = frozenset(
{
(128, 7168, 256, 257, 9),
}
)
_KMOE_STAGE2_BOUNDARY_CONTRACT_FRESH_REASONS = frozenset(
{
"first_bypass",
}
)
_STAGE2_BOUNDARY_CONTRACT_SEGMENT_ALIGN = 256
def _require_aiter() -> None:
if _AITER_IMPORT_ERROR is not None:
raise RuntimeError(
"This submission expects the gpumode runtime with `aiter` installed."
) from _AITER_IMPORT_ERROR
def _pad_to(x: int, align: int) -> int:
return (x + align - 1) // align * align
def _maybe_contiguous(tensor: torch.Tensor) -> torch.Tensor:
if tensor.is_contiguous():
return tensor
return tensor.contiguous()
def _ensure_dtype(tensor: torch.Tensor, dtype: torch.dtype) -> torch.Tensor:
if tensor.dtype == dtype:
return tensor
return tensor.to(dtype)
def _allclose(a: torch.Tensor, b: torch.Tensor) -> bool:
if a.shape != b.shape:
return False
if a.dtype != b.dtype:
b = b.to(a.dtype)
return torch.allclose(a, b, rtol=RTOL, atol=ATOL)
def _log_once(event_key: Tuple[Any, ...], message: str) -> None:
if event_key in _LOGGED_EVENTS:
return
_LOGGED_EVENTS.add(event_key)
print(message, file=sys.stderr)
def _set_pending_flydsl_stage2_patch_mode(mode: str) -> None:
global _KMOE_PENDING_FLYDSL_STAGE2_PATCH_MODE
if mode not in ("ptr", "rawbuf"):
raise ValueError(f"unsupported flydsl stage2 patch mode: {mode}")
_KMOE_PENDING_FLYDSL_STAGE2_PATCH_MODE = mode
def _call_with_flydsl_stage2_patch_mode(stage2_callable, patch_mode: str, *args, **kwargs):
prev_mode = _KMOE_PENDING_FLYDSL_STAGE2_PATCH_MODE
_set_pending_flydsl_stage2_patch_mode(patch_mode)
try:
return stage2_callable(*args, **kwargs)
finally:
_set_pending_flydsl_stage2_patch_mode(prev_mode)
def _safe_tensor_version(tensor: torch.Tensor) -> Optional[int]:
try:
return getattr(tensor, "_version", 0)
except RuntimeError:
return None
def _tensor_identity_key(tensor: torch.Tensor) -> Tuple[Any, ...]:
return (
id(tensor),
tensor.data_ptr(),
tuple(tensor.shape),
tuple(tensor.stride()),
str(tensor.dtype),
_safe_tensor_version(tensor),
tensor.device.type,
tensor.device.index,
)
def _get_kmoe_full_exact_sort_state(
topk_ids: torch.Tensor,
topk_weights: torch.Tensor,
num_experts: int,
model_dim: int,
moebuf_dtype: torch.dtype,
block_m: int,
) -> Tuple[
torch.Tensor,
torch.Tensor,
torch.Tensor,
torch.Tensor,
Optional[torch.Tensor],
bool,
Optional[Tuple[Any, ...]],
]:
cache_key = (
_tensor_identity_key(topk_ids),
_tensor_identity_key(topk_weights),
num_experts,
model_dim,
block_m,
str(moebuf_dtype),
)
cached = _KMOE_FULL_EXACT_SORT_CACHE.get(cache_key)
if cached is not None:
return (*cached, None, True, None)
sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf = (
_aiter_fused_moe.moe_sorting(
topk_ids,
topk_weights,
num_experts,
model_dim,
moebuf_dtype,
block_m,
None,
None,
0,
)
)
_log_once(
("kmoe-full-exact-sort-cache-build", cache_key),
"[submission] kmoe full exact sort cache build "
f"rows={tuple(topk_ids.shape)} experts={num_experts} "
f"block_m={block_m} device={topk_ids.device}",
)
return (
sorted_ids,
sorted_weights,
sorted_expert_ids,
num_valid_ids,
moe_buf,
False,
cache_key,
)
def _maybe_reset_full_exact_sort_cache(
shape_key: Tuple[int, int, int, int, int],
) -> None:
global _LAST_FULL_EXACT_SORT_CACHE_SHAPE
if _LAST_FULL_EXACT_SORT_CACHE_SHAPE == shape_key:
return
if _KMOE_FULL_EXACT_SORT_CACHE:
_KMOE_FULL_EXACT_SORT_CACHE.clear()
_log_once(
(
"kmoe-full-exact-sort-cache-reset",
_LAST_FULL_EXACT_SORT_CACHE_SHAPE,
shape_key,
),
"[submission] kmoe full exact sort cache reset "
f"prev_shape={_LAST_FULL_EXACT_SORT_CACHE_SHAPE} "
f"next_shape={shape_key}",
)
if _KMOE_MOEBUF_RING_WORKSPACE_CACHE:
_KMOE_MOEBUF_RING_WORKSPACE_CACHE.clear()
_KMOE_MOEBUF_RING_WORKSPACE_INDEX.clear()
_log_once(
(
"kmoe-moebuf-ring-workspace-reset",
_LAST_FULL_EXACT_SORT_CACHE_SHAPE,
shape_key,
),
"[submission] kmoe moebuf ring workspace reset "
f"prev_shape={_LAST_FULL_EXACT_SORT_CACHE_SHAPE} "
f"next_shape={shape_key}",
)
_LAST_FULL_EXACT_SORT_CACHE_SHAPE = shape_key
def _store_kmoe_full_exact_sort_cache(
shape_key: Tuple[int, int, int, int, int],
cache_key: Tuple[Any, ...],
cached_values: Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor],
) -> None:
max_entries = _KMOE_FULL_EXACT_SORT_MAX_CACHE_ENTRIES.get(shape_key)
if max_entries is not None and cache_key not in _KMOE_FULL_EXACT_SORT_CACHE:
while len(_KMOE_FULL_EXACT_SORT_CACHE) >= max_entries:
oldest_key = next(iter(_KMOE_FULL_EXACT_SORT_CACHE))
_KMOE_FULL_EXACT_SORT_CACHE.pop(oldest_key, None)
_log_once(
("kmoe-full-exact-sort-cache-evict", shape_key, max_entries),
"[submission] kmoe full exact sort cache evict "
f"shape={shape_key} max_entries={max_entries}",
)
_KMOE_FULL_EXACT_SORT_CACHE[cache_key] = cached_values
def _get_kmoe_moebuf_workspace(
token_num: int,
model_dim: int,
dtype: torch.dtype,
device: torch.device,
ring_depth: int = 1,
) -> torch.Tensor:
workspace_key = (
token_num,
model_dim,
str(dtype),
device.type,
device.index,
)
if ring_depth > 1:
ring_workspace_key = (*workspace_key, ring_depth)
cached_ring = _KMOE_MOEBUF_RING_WORKSPACE_CACHE.get(ring_workspace_key)
if cached_ring is None:
cached_ring = tuple(
torch.empty((token_num, model_dim), dtype=dtype, device=device)
for _ in range(ring_depth)
)
_KMOE_MOEBUF_RING_WORKSPACE_CACHE[ring_workspace_key] = cached_ring
_KMOE_MOEBUF_RING_WORKSPACE_INDEX[ring_workspace_key] = 0
_log_once(
("kmoe-moebuf-ring-workspace-build", ring_workspace_key),
"[submission] kmoe moebuf ring workspace build "
f"shape={(token_num, model_dim)} slots={ring_depth} device={device}",
)
slot = _KMOE_MOEBUF_RING_WORKSPACE_INDEX.get(ring_workspace_key, 0)
_KMOE_MOEBUF_RING_WORKSPACE_INDEX[ring_workspace_key] = (
slot + 1
) % ring_depth
return cached_ring[slot]
cached = _KMOE_MOEBUF_WORKSPACE_CACHE.get(workspace_key)
if cached is None:
cached = torch.empty((token_num, model_dim), dtype=dtype, device=device)
_KMOE_MOEBUF_WORKSPACE_CACHE[workspace_key] = cached
_log_once(
("kmoe-moebuf-workspace-build", workspace_key),
"[submission] kmoe moebuf workspace build "
f"shape={(token_num, model_dim)} device={device}",
)
return cached
def _patch_moe_dtype_aliases() -> None:
global _MOE_DTYPE_ALIAS_PATCHED
if _MOE_DTYPE_ALIAS_PATCHED or _aiter_moe_op is None:
return
alias_names = {"fp4x2", "float4_e2m1fn_x2", "torch.float4_e2m1fn_x2"}
fp4_dtype = dtypes.fp4x2
alias_names.add(str(fp4_dtype))
alias_names.add(str(fp4_dtype).replace("torch.", ""))
if hasattr(torch, "float4_e2m1fn_x2"):
_aiter_moe_op.dtype2str_dict.setdefault(torch.float4_e2m1fn_x2, "fp4x2")
_aiter_moe_op.dtype2str_dict.setdefault(fp4_dtype, "fp4x2")
_aiter_moe_op.str2dtype_dict.setdefault("fp4x2", fp4_dtype)
for alias in alias_names:
if alias:
_aiter_moe_op.dtype2str_dict.setdefault(alias, "fp4x2")
_aiter_moe_op.str2dtype_dict.setdefault(alias, fp4_dtype)
_MOE_DTYPE_ALIAS_PATCHED = True
def _manual_fused_quant_sort_block_size_mx(
shape_key: Tuple[int, int, int, int, int],
stage: str,
) -> Optional[int]:
stage_cfg = _MANUAL_FUSED_QUANT_SORT_BLOCK_SIZE_MX.get(shape_key)
if stage_cfg is None:
return None
return stage_cfg.get(stage)
def _ceil_div(x: int, y: int) -> int:
return (x + y - 1) // y
def _quant_sort_blockscale_shape(
sorted_rows: int,
n: int,
) -> Tuple[int, int, int, int, int]:
scale_n = _ceil_div(n, 32)
return (
_ceil_div(sorted_rows, 32),
_ceil_div(scale_n, 8),
4,
16,
4,
)
def _get_stage2_boundary_contract(
shape_key: Tuple[int, int, int, int, int],
token_num: int,
topk: int,
inter_dim: int,
sorted_rows: int,
device: torch.device,
workspace_reason: str = "default",
) -> Optional[Dict[str, torch.Tensor]]:
if shape_key not in _MANUAL_FUSED_STAGE2_BOUNDARY_CONTRACT_SHAPES:
return None
if inter_dim % 2 != 0:
return None
a2_fp4_u8_shape = (token_num * topk, inter_dim // 2)
a2_scale_u8_shape = _quant_sort_blockscale_shape(sorted_rows, inter_dim)
total_bytes = 0
def reserve(size_bytes: int) -> Tuple[int, int]:
nonlocal total_bytes
total_bytes = _pad_to(total_bytes, _STAGE2_BOUNDARY_CONTRACT_SEGMENT_ALIGN)
offset = total_bytes
total_bytes += size_bytes
return offset, size_bytes
layout = {
"a2_fp4_u8": reserve(math.prod(a2_fp4_u8_shape)),
"a2_scale_u8": reserve(math.prod(a2_scale_u8_shape)),
}
cache_workspace = (
workspace_reason not in _KMOE_STAGE2_BOUNDARY_CONTRACT_FRESH_REASONS
)
contract = None
workspace_key = None
if cache_workspace:
workspace_key = (
shape_key,
token_num,
topk,
inter_dim,
sorted_rows,
device.type,
device.index,
workspace_reason,
)
contract = _KMOE_STAGE2_BOUNDARY_CONTRACT_WORKSPACE_CACHE.get(workspace_key)
if contract is None:
storage = torch.empty((total_bytes,), dtype=torch.uint8, device=device)
def make_u8_view(name: str, shape: Tuple[int, ...]) -> torch.Tensor:
offset, size_bytes = layout[name]
return storage.narrow(0, offset, size_bytes).view(shape)
contract = {
"storage": storage,
"a2_fp4_u8": make_u8_view("a2_fp4_u8", a2_fp4_u8_shape),
"a2_scale_u8": make_u8_view("a2_scale_u8", a2_scale_u8_shape),
}
if cache_workspace:
_KMOE_STAGE2_BOUNDARY_CONTRACT_WORKSPACE_CACHE[workspace_key] = contract
event_key = ("kmoe-stage2-boundary-contract-workspace-build", workspace_key)
workspace_mode = "cached"
else:
event_key = (
"kmoe-stage2-boundary-contract-workspace-build-fresh",
shape_key,
token_num,
topk,
inter_dim,
sorted_rows,
device.type,
device.index,
workspace_reason,
)
workspace_mode = "fresh"
_log_once(
event_key,
"[submission] kmoe stage2-boundary contract workspace build "
f"shape={shape_key} sorted_rows={sorted_rows} bytes={total_bytes} "
f"reason={workspace_reason} mode={workspace_mode} device={device}",
)
_log_once(
("manual2stage-stage2-boundary-contract", shape_key),
"[submission] manual fused 2stage stage2-boundary contract active "
f"shape={shape_key} sorted_rows={sorted_rows} bytes={total_bytes} "
"pack=a2_fp4+a2_scale",
)
return contract
def _fused_dynamic_mxfp4_quant_moe_sort_tuned(
x: torch.Tensor,
sorted_ids: torch.Tensor,
num_valid_ids: torch.Tensor,
token_num: int,
topk: int,
block_size: int = 32,
block_size_mx: Optional[int] = None,
shape_key: Optional[Tuple[int, int, int, int, int]] = None,
stage: Optional[str] = None,
x_fp4_u8: Optional[torch.Tensor] = None,
blockscale_e8m0_sorted_u8: Optional[torch.Tensor] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
if (
block_size_mx is None
or triton is None
or _fused_dynamic_mxfp4_quant_moe_sort_kernel is None
):
return _aiter_fused_moe.fused_dynamic_mxfp4_quant_moe_sort(
x,
sorted_ids=sorted_ids,
num_valid_ids=num_valid_ids,
token_num=token_num,
topk=topk,
block_size=block_size,
)
M, N = x.shape
assert (N // 2) % 2 == 0
mxfp4_quant_block_size = 32
x_fp4_shape = (M, N // 2)
if x_fp4_u8 is None:
x_fp4 = torch.empty(x_fp4_shape, dtype=torch.uint8, device=x.device)
else:
if (
x_fp4_u8.dtype != torch.uint8
or x_fp4_u8.device != x.device
or tuple(x_fp4_u8.shape) != x_fp4_shape
):
raise RuntimeError(
"Stage2 boundary contract fp4 view does not match quant-sort shape."
)
x_fp4 = x_fp4_u8
scale_n_valid = triton.cdiv(N, mxfp4_quant_block_size)
scale_n = scale_n_valid
block_size_m = 32
block_size_n = 8
block_size_m_u32 = 16
block_size_n_u32 = 4
n_i = scale_n
m_o, n_o = sorted_ids.shape[0], n_i
assert (n_i // 2) % 2 == 0
assert block_size % block_size_m == 0
blockscale_shape = (
triton.cdiv(m_o, block_size_m),
triton.cdiv(n_o, block_size_n),
block_size_n_u32,
block_size_m_u32,
4,
)
if blockscale_e8m0_sorted_u8 is None:
blockscale_e8m0_sorted = torch.empty(
blockscale_shape,
dtype=torch.uint8,
device=x.device,
)
else:
if (
blockscale_e8m0_sorted_u8.dtype != torch.uint8
or blockscale_e8m0_sorted_u8.device != x.device
or tuple(blockscale_e8m0_sorted_u8.shape) != blockscale_shape
):
raise RuntimeError(
"Stage2 boundary contract scale view does not match quant-sort shape."
)
blockscale_e8m0_sorted = blockscale_e8m0_sorted_u8
num_pid = triton.cdiv(M, block_size_mx) * scale_n + triton.cdiv(
m_o, block_size_m
) * triton.cdiv(n_i, block_size_n)
launch_kwargs = {
"token_num": token_num,
"N_i": n_i,
"MXFP4_QUANT_BLOCK_SIZE": mxfp4_quant_block_size,
"BLOCK_SIZE_Mx": block_size_mx,
"BLOCK_SIZE_M": block_size_m // 2,
"BLOCK_SIZE_N": block_size_n // 2,
"TOPK": topk,
}
kernel_arg_names = tuple(
getattr(_fused_dynamic_mxfp4_quant_moe_sort_kernel, "arg_names", ())
)
if "M_i" in kernel_arg_names:
launch_kwargs["M_i"] = M
_log_once(
(
"manual2stage-quant-sort-launcher-argnames-has-mi",
shape_key,
stage,
block_size_mx,
),
"[submission] quant-sort launcher compat add M_i "
f"shape={shape_key} stage={stage} block_size_mx={block_size_mx}",
)
launcher = _fused_dynamic_mxfp4_quant_moe_sort_kernel[(num_pid,)]
base_args = (
x,
x_fp4,
sorted_ids,
num_valid_ids,
blockscale_e8m0_sorted,
M,
N,
scale_n,
*x.stride(),
*x_fp4.stride(),
*blockscale_e8m0_sorted.stride(),
)
try:
launcher(*base_args, **launch_kwargs)
except TypeError as exc:
if "M_i" not in str(exc) or "M_i" in launch_kwargs:
raise
launch_kwargs["M_i"] = M
_log_once(
(
"manual2stage-quant-sort-launcher-fallback-with-mi",
shape_key,
stage,
block_size_mx,
),
"[submission] quant-sort launcher fallback-with-M_i "
f"shape={shape_key} stage={stage} block_size_mx={block_size_mx}",
)
launcher(*base_args, **launch_kwargs)
return (
x_fp4.view(dtypes.fp4x2),
blockscale_e8m0_sorted.view(dtypes.fp8_e8m0).view(-1, n_o),
)
def _run_manual_fused_quant_sort_with_override(
shape_key: Tuple[int, int, int, int, int],
stage: str,
x: torch.Tensor,
sorted_ids: torch.Tensor,
num_valid_ids: torch.Tensor,
token_num: int,
topk: int,
block_size: int,
x_fp4_u8: Optional[torch.Tensor] = None,
blockscale_e8m0_sorted_u8: Optional[torch.Tensor] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
block_size_mx = _manual_fused_quant_sort_block_size_mx(shape_key, stage)
if block_size_mx is not None:
_log_once(
("manual2stage-quant-sort-tuned", shape_key, stage, block_size_mx),
"[submission] manual fused quant-sort tuned "
f"shape={shape_key} stage={stage} block_size_mx={block_size_mx}",
)
if x_fp4_u8 is not None or blockscale_e8m0_sorted_u8 is not None:
_log_once(
("manual2stage-quant-sort-stage2-boundary-contract", shape_key, stage),
"[submission] manual fused quant-sort stage2-boundary contract "
f"shape={shape_key} stage={stage}",
)
return _fused_dynamic_mxfp4_quant_moe_sort_tuned(
x,
sorted_ids=sorted_ids,
num_valid_ids=num_valid_ids,
token_num=token_num,
topk=topk,
block_size=block_size,
block_size_mx=block_size_mx,
shape_key=shape_key,
stage=stage,
x_fp4_u8=x_fp4_u8,
blockscale_e8m0_sorted_u8=blockscale_e8m0_sorted_u8,
)
def _device_is_mi355x(device: torch.device) -> bool:
if device.type != "cuda":
return False
key = (device.type, device.index)
cached = _DEVICE_KIND_CACHE.get(key)
if cached is not None:
return cached
try:
name = torch.cuda.get_device_name(device).upper()
except Exception:
name = ""
is_mi355x = "MI355" in name or "GFX950" in name
_DEVICE_KIND_CACHE[key] = is_mi355x
return is_mi355x
def _shape_key(
hidden_states: torch.Tensor,
topk_ids: torch.Tensor,
config: Dict[str, int],
) -> Tuple[int, int, int, int, int]:
return (
hidden_states.shape[0],
config["d_hidden"],
config["d_expert"],
config.get("n_routed_experts", 0) + config.get("n_shared_experts", 0),
topk_ids.shape[1],
)
def _backend_key(
hidden_states: torch.Tensor,
topk_ids: torch.Tensor,
config: Dict[str, int],
) -> Tuple[Any, ...]:
return (
str(hidden_states.device),
_shape_key(hidden_states, topk_ids, config),
)
def _candidate_block_m(
hidden_states: torch.Tensor,
topk_ids: torch.Tensor,
config: Dict[str, int],
) -> Optional[int]:
exact = _BENCHMARK_BLOCK_M.get(_shape_key(hidden_states, topk_ids, config))
if exact is not None:
return exact
token_count = hidden_states.shape[0]
d_expert = config["d_expert"]
if token_count <= 32:
return 32
if token_count <= 128:
return 64
if d_expert >= 2048:
return 64
return 128
def _forced_ck_non_temporal_load(
hidden_states: torch.Tensor,
topk_ids: torch.Tensor,
config: Dict[str, int],
) -> Optional[bool]:
if _shape_key(hidden_states, topk_ids, config) in _CK_FORCE_NON_TEMPORAL_TRUE:
return True
return None
def _get_hip_quant_per_1x32():
global _HIP_QUANT_PER_1X32
if _HIP_QUANT_PER_1X32 is None:
_HIP_QUANT_PER_1X32 = aiter.get_hip_quant(QuantType.per_1x32)
return _HIP_QUANT_PER_1X32
def _get_torch_quant_per_1x32():
global _TORCH_QUANT_PER_1X32
if _TORCH_QUANT_PER_1X32 is None:
_TORCH_QUANT_PER_1X32 = aiter.get_torch_quant(QuantType.per_1x32)
return _TORCH_QUANT_PER_1X32
def _get_hidden_prequant(
hidden_states: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
global _LAST_HIDDEN_PREQUANT_KEY, _LAST_HIDDEN_PREQUANT_VALUE
key = (
id(hidden_states),
hidden_states.data_ptr(),
hidden_states.numel(),
_safe_tensor_version(hidden_states),
)
if _LAST_HIDDEN_PREQUANT_KEY == key and _LAST_HIDDEN_PREQUANT_VALUE is not None:
return _LAST_HIDDEN_PREQUANT_VALUE
quant_func = _get_hip_quant_per_1x32()
hidden_q, hidden_scale = quant_func(hidden_states, quant_dtype=dtypes.fp4x2)
hidden_q = _maybe_contiguous(hidden_q)
hidden_scale = _maybe_contiguous(hidden_scale)
_LAST_HIDDEN_PREQUANT_KEY = key
_LAST_HIDDEN_PREQUANT_VALUE = (hidden_q, hidden_scale)
return hidden_q, hidden_scale
def _get_hidden_prequant_torch(
hidden_states: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
global _LAST_HIDDEN_PREQUANT_TORCH_KEY, _LAST_HIDDEN_PREQUANT_TORCH_VALUE
key = (
id(hidden_states),
hidden_states.data_ptr(),
hidden_states.numel(),
_safe_tensor_version(hidden_states),
)
if (
_LAST_HIDDEN_PREQUANT_TORCH_KEY == key
and _LAST_HIDDEN_PREQUANT_TORCH_VALUE is not None
):
return _LAST_HIDDEN_PREQUANT_TORCH_VALUE
quant_func = _get_torch_quant_per_1x32()
hidden_q, hidden_scale = quant_func(hidden_states, quant_dtype=dtypes.fp4x2)
hidden_q = _maybe_contiguous(hidden_q)
hidden_scale = _maybe_contiguous(hidden_scale)
_LAST_HIDDEN_PREQUANT_TORCH_KEY = key
_LAST_HIDDEN_PREQUANT_TORCH_VALUE = (hidden_q, hidden_scale)
return hidden_q, hidden_scale
def _should_prequant_ck(
hidden_states: torch.Tensor,
topk_ids: torch.Tensor,
config: Dict[str, int],
) -> bool:
shape = _shape_key(hidden_states, topk_ids, config)
return shape in _SELECTIVE_PREQUANT_CK_SHAPES and shape not in _SELECTIVE_RAW_CK_SHAPES
def _forced_ck_stage_kernels(
hidden_states: torch.Tensor,
topk_ids: torch.Tensor,
config: Dict[str, int],
) -> Optional[Tuple[str, str, int]]:
return _CK_FORCE_STAGE_KERNELS.get(_shape_key(hidden_states, topk_ids, config))
def _should_use_manual_fused_2stage(
hidden_states: torch.Tensor,
topk_ids: torch.Tensor,
config: Dict[str, int],
) -> bool:
return _shape_key(hidden_states, topk_ids, config) in _MANUAL_FUSED_2STAGE_SHAPES
def _prepare_common(
data: input_t,
) -> Tuple[
torch.Tensor,
torch.Tensor,
torch.Tensor,
torch.Tensor,
torch.Tensor,
torch.Tensor,
torch.Tensor,
torch.Tensor,
torch.Tensor,
Dict[str, int],
int,
int,
]:
_require_aiter()
(
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
d_hidden = config["d_hidden"]
d_expert = config["d_expert"]
d_hidden_pad = config.get("d_hidden_pad", _pad_to(d_hidden, PAD_ALIGN))
d_expert_pad = config.get("d_expert_pad", _pad_to(d_expert, PAD_ALIGN))
if gate_up_weight_shuffled is None:
gate_up_weight_shuffled = shuffle_weight(gate_up_weight, layout=SHUFFLE_LAYOUT)
if down_weight_shuffled is None:
down_weight_shuffled = shuffle_weight(down_weight, layout=SHUFFLE_LAYOUT)
if gate_up_weight_scale_shuffled is None:
gate_up_weight_scale_shuffled = fp4_utils.e8m0_shuffle(gate_up_weight_scale)
if down_weight_scale_shuffled is None:
down_weight_scale_shuffled = fp4_utils.e8m0_shuffle(down_weight_scale)
hidden_states = _maybe_contiguous(_ensure_dtype(hidden_states, torch.bfloat16))
topk_weights = _maybe_contiguous(_ensure_dtype(topk_weights, torch.float32))
topk_ids = _maybe_contiguous(_ensure_dtype(topk_ids, torch.int32))
gate_up_weight_shuffled = _maybe_contiguous(gate_up_weight_shuffled)
down_weight_shuffled = _maybe_contiguous(down_weight_shuffled)
gate_up_weight_scale_shuffled = _maybe_contiguous(gate_up_weight_scale_shuffled)
down_weight_scale_shuffled = _maybe_contiguous(down_weight_scale_shuffled)
return (
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
gate_up_weight_scale,
down_weight_scale,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
config,
d_hidden_pad - d_hidden,
d_expert_pad - d_expert,
)
def _run_fused(
hidden_states: torch.Tensor,
gate_up_weight_shuffled: torch.Tensor,
down_weight_shuffled: torch.Tensor,
gate_up_weight_scale_shuffled: torch.Tensor,
down_weight_scale_shuffled: torch.Tensor,
topk_weights: torch.Tensor,
topk_ids: torch.Tensor,
hidden_pad: int,
intermediate_pad: int,
) -> torch.Tensor:
return fused_moe(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
topk_weights,
topk_ids,
expert_mask=None,
activation=ActivationType.Silu,
quant_type=QuantType.per_1x32,
doweight_stage1=False,
w1_scale=gate_up_weight_scale_shuffled,
w2_scale=down_weight_scale_shuffled,
a1_scale=None,
a2_scale=None,
hidden_pad=hidden_pad,
intermediate_pad=intermediate_pad,
)
def _run_ck_with_a1_scale(
hidden_states: torch.Tensor,
gate_up_weight_shuffled: torch.Tensor,
down_weight_shuffled: torch.Tensor,
gate_up_weight_scale_shuffled: torch.Tensor,
down_weight_scale_shuffled: torch.Tensor,
topk_weights: torch.Tensor,
topk_ids: torch.Tensor,
block_m: int,
a1_scale: Optional[torch.Tensor],
non_temporal_load: Optional[bool] = None,
extra_kwargs: Optional[Dict[str, Any]] = None,
) -> torch.Tensor:
base_kwargs = {
"w1_scale": gate_up_weight_scale_shuffled,
"w2_scale": down_weight_scale_shuffled,
"a1_scale": a1_scale,
"a2_scale": None,
"block_m": block_m,
"expert_mask": None,
}
def call(more_kwargs: Optional[Dict[str, Any]] = None) -> torch.Tensor:
merged_kwargs = {}
if extra_kwargs is not None:
merged_kwargs.update(extra_kwargs)
if more_kwargs is not None:
merged_kwargs.update(more_kwargs)
return _CK_MOE(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
topk_weights,
topk_ids,
**base_kwargs,
**merged_kwargs,
)
if non_temporal_load is None:
return call()
global _CK_MOE_NT_KWARG, _CK_MOE_NT_KWARG_READY
if _CK_MOE_NT_KWARG_READY:
if _CK_MOE_NT_KWARG is None:
return call()
try:
return call({_CK_MOE_NT_KWARG: non_temporal_load})
except Exception:
_CK_MOE_NT_KWARG = None
return call()
for kwarg_name in ("non_temporal_load", "use_non_temporal_load"):
try:
out = call({kwarg_name: non_temporal_load})
_CK_MOE_NT_KWARG = kwarg_name
_CK_MOE_NT_KWARG_READY = True
return out
except TypeError:
continue
except Exception:
_CK_MOE_NT_KWARG = None
_CK_MOE_NT_KWARG_READY = True
return call()
_CK_MOE_NT_KWARG = None
_CK_MOE_NT_KWARG_READY = True
return call()
def _run_ck(
hidden_states: torch.Tensor,
gate_up_weight_shuffled: torch.Tensor,
down_weight_shuffled: torch.Tensor,
gate_up_weight_scale_shuffled: torch.Tensor,
down_weight_scale_shuffled: torch.Tensor,
topk_weights: torch.Tensor,
topk_ids: torch.Tensor,
block_m: int,
non_temporal_load: Optional[bool] = None,
) -> torch.Tensor:
return _run_ck_with_a1_scale(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
block_m,
None,
non_temporal_load,
)
def _run_ck_prequant(
hidden_states: torch.Tensor,
gate_up_weight_shuffled: torch.Tensor,
down_weight_shuffled: torch.Tensor,
gate_up_weight_scale_shuffled: torch.Tensor,
down_weight_scale_shuffled: torch.Tensor,
topk_weights: torch.Tensor,
topk_ids: torch.Tensor,
block_m: int,
non_temporal_load: Optional[bool] = None,
) -> torch.Tensor:
hidden_q, hidden_scale = _get_hidden_prequant(hidden_states)
return _run_ck_with_a1_scale(
hidden_q,
gate_up_weight_shuffled,
down_weight_shuffled,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
block_m,
hidden_scale,
non_temporal_load,
)
def _run_ck_prequant_with_stage_kernels(
hidden_states: torch.Tensor,
gate_up_weight_shuffled: torch.Tensor,
down_weight_shuffled: torch.Tensor,
gate_up_weight_scale_shuffled: torch.Tensor,
down_weight_scale_shuffled: torch.Tensor,
topk_weights: torch.Tensor,
topk_ids: torch.Tensor,
block_m: int,
non_temporal_load: Optional[bool],
stage1_kernel: str,
stage2_kernel: str,
) -> torch.Tensor:
hidden_q, hidden_scale = _get_hidden_prequant(hidden_states)
global _CK_MOE_STAGE_KWARGS, _CK_MOE_STAGE_KWARGS_READY
def call(stage_kwargs: Dict[str, str]) -> torch.Tensor:
return _run_ck_with_a1_scale(
hidden_q,
gate_up_weight_shuffled,
down_weight_shuffled,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
block_m,
hidden_scale,
non_temporal_load,
stage_kwargs,
)
if _CK_MOE_STAGE_KWARGS_READY:
if _CK_MOE_STAGE_KWARGS is None:
raise RuntimeError("CK MoE stage-kernel kwargs are unavailable.")
return call(
{
_CK_MOE_STAGE_KWARGS[0]: stage1_kernel,
_CK_MOE_STAGE_KWARGS[1]: stage2_kernel,
}
)
for kw1, kw2 in (
("kernelName1", "kernelName2"),
("kernel_name1", "kernel_name2"),
("stage1_kernelName", "stage2_kernelName"),
("stage1_kernel_name", "stage2_kernel_name"),
("gemm1_kernelName", "gemm2_kernelName"),
("gemm1_kernel_name", "gemm2_kernel_name"),
):
try:
out = call({kw1: stage1_kernel, kw2: stage2_kernel})
_CK_MOE_STAGE_KWARGS = (kw1, kw2)
_CK_MOE_STAGE_KWARGS_READY = True
return out
except TypeError:
continue
except Exception:
_CK_MOE_STAGE_KWARGS = None
_CK_MOE_STAGE_KWARGS_READY = True
raise
_CK_MOE_STAGE_KWARGS = None
_CK_MOE_STAGE_KWARGS_READY = True
raise RuntimeError("CK MoE stage-kernel kwargs are unavailable.")
def _run_fused_prequant(
hidden_q: torch.Tensor,
hidden_scale: torch.Tensor,
gate_up_weight_shuffled: torch.Tensor,
down_weight_shuffled: torch.Tensor,
gate_up_weight_scale_shuffled: torch.Tensor,
down_weight_scale_shuffled: torch.Tensor,
topk_weights: torch.Tensor,
topk_ids: torch.Tensor,
hidden_pad: int,
intermediate_pad: int,
) -> torch.Tensor:
return fused_moe(
hidden_q,
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=hidden_scale,
a2_scale=None,
dtype=torch.bfloat16,
hidden_pad=hidden_pad,
intermediate_pad=intermediate_pad,
)
def _run_manual_fused_2stage_prequant(
hidden_states: torch.Tensor,
gate_up_weight_shuffled: torch.Tensor,
down_weight_shuffled: torch.Tensor,
gate_up_weight_scale_shuffled: torch.Tensor,
down_weight_scale_shuffled: torch.Tensor,
topk_weights: torch.Tensor,
topk_ids: torch.Tensor,
config: Dict[str, int],
hidden_pad: int,
intermediate_pad: int,
stage1_kernel: str,
stage2_kernel: str,
block_m: int,
) -> torch.Tensor:
shape_key = _shape_key(hidden_states, topk_ids, config)
num_experts = config.get("n_routed_experts", 0) + config.get("n_shared_experts", 0)
token_num = hidden_states.shape[0]
topk = topk_ids.shape[1]
is_shuffled = getattr(gate_up_weight_shuffled, "is_shuffled", False)
_, model_dim, inter_dim = _aiter_fused_moe.get_inter_dim(
gate_up_weight_shuffled.shape,
down_weight_shuffled.shape,
)
_log_once(
("manual2stage-enter", shape_key),
"[submission] manual fused 2stage enter "
f"shape={shape_key} block_m={block_m} prequant=fused_quant_sort "
f"kernel1='{stage1_kernel}' kernel2='{stage2_kernel}'",
)
moebuf_ring_slots = _KMOE_MOEBUF_RING_WORKSPACE_SLOTS.get(shape_key, 1)
if moebuf_ring_slots > 1:
_log_once(
("manual2stage-moebuf-ring-workspace", shape_key, moebuf_ring_slots),
"[submission] manual fused 2stage moebuf ring workspace active "
f"shape={shape_key} slots={moebuf_ring_slots}",
)
_patch_moe_dtype_aliases()
use_exact_sort_cache = shape_key in _KMOE_FULL_EXACT_SORT_CACHE_SHAPES
cache_hit = False
hit_clone_sort_inputs = False
miss_clone_sort_inputs = False
miss_streak = 0
skip_cache_store = False
bypass_exact_sort_cache = False
first_bypass_miss_streak: Optional[int] = None
if use_exact_sort_cache:
skip_store_after = _KMOE_FULL_EXACT_SORT_SKIP_STORE_AFTER_MISS_STREAK.get(
shape_key
)
if skip_store_after is not None:
first_bypass_miss_streak = skip_store_after + 2
current_miss_streak = _KMOE_FULL_EXACT_SORT_MISS_STREAK.get(shape_key, 0)
bypass_exact_sort_cache = (
skip_store_after is not None and current_miss_streak > skip_store_after
)
if bypass_exact_sort_cache:
miss_streak = current_miss_streak + 1
_KMOE_FULL_EXACT_SORT_MISS_STREAK[shape_key] = miss_streak
skip_cache_store = True
cache_hit = False
if _KMOE_FULL_EXACT_SORT_CACHE:
_KMOE_FULL_EXACT_SORT_CACHE.clear()
_log_once(
("kmoe-full-exact-sort-cache-bypass", shape_key),
"[submission] kmoe full exact sort cache bypass "
f"shape={shape_key} miss_streak={miss_streak}",
)
(
cached_sorted_ids,
cached_sorted_weights,
cached_sorted_expert_ids,
cached_num_valid_ids,
moe_buf,
) = _aiter_fused_moe.moe_sorting(
topk_ids,
topk_weights,
num_experts,
model_dim,
hidden_states.dtype,
block_m,
None,
None,
0,
)
if shape_key in _KMOE_MISS_CLONE_SORT_INPUTS_SHAPES:
sorted_ids = cached_sorted_ids.clone()
sorted_weights = cached_sorted_weights.clone()
sorted_expert_ids = cached_sorted_expert_ids.clone()
num_valid_ids = cached_num_valid_ids.clone()
miss_clone_sort_inputs = True
else:
sorted_ids = cached_sorted_ids
sorted_weights = cached_sorted_weights
sorted_expert_ids = cached_sorted_expert_ids
num_valid_ids = cached_num_valid_ids
if shape_key in _KMOE_CACHE_HIT_NOCLONE_SORT_INPUTS_SHAPES:
_log_once(
("kmoe-miss-sort-inputs-noclone-bypass", shape_key),
"[submission] kmoe miss sort inputs noclone bypass "
f"shape={shape_key}",
)
else:
(
cached_sorted_ids,
cached_sorted_weights,
cached_sorted_expert_ids,
cached_num_valid_ids,
cached_moe_buf,
cache_hit,
cache_key_to_store,
) = _get_kmoe_full_exact_sort_state(
topk_ids,
topk_weights,
num_experts,
model_dim,
hidden_states.dtype,
block_m,
)
if cache_hit:
_KMOE_FULL_EXACT_SORT_MISS_STREAK[shape_key] = 0
if shape_key in _KMOE_CACHE_HIT_NOCLONE_SORT_INPUTS_SHAPES:
sorted_ids = cached_sorted_ids
sorted_weights = cached_sorted_weights
sorted_expert_ids = cached_sorted_expert_ids
num_valid_ids = cached_num_valid_ids
_log_once(
("kmoe-cache-hit-sort-inputs-noclone", shape_key),
"[submission] kmoe cache-hit sort inputs noclone "
f"shape={shape_key}",
)
else:
sorted_ids = cached_sorted_ids.clone()
sorted_weights = cached_sorted_weights.clone()
sorted_expert_ids = cached_sorted_expert_ids.clone()
num_valid_ids = cached_num_valid_ids.clone()
hit_clone_sort_inputs = True
moe_buf = _get_kmoe_moebuf_workspace(
token_num,
model_dim,
hidden_states.dtype,
hidden_states.device,
moebuf_ring_slots,
)
else:
miss_streak = current_miss_streak + 1
_KMOE_FULL_EXACT_SORT_MISS_STREAK[shape_key] = miss_streak
if shape_key in _KMOE_MISS_CLONE_SORT_INPUTS_SHAPES:
sorted_ids = cached_sorted_ids.clone()
sorted_weights = cached_sorted_weights.clone()
sorted_expert_ids = cached_sorted_expert_ids.clone()
num_valid_ids = cached_num_valid_ids.clone()
miss_clone_sort_inputs = True
else:
sorted_ids = cached_sorted_ids
sorted_weights = cached_sorted_weights
sorted_expert_ids = cached_sorted_expert_ids
num_valid_ids = cached_num_valid_ids
if shape_key in _KMOE_CACHE_HIT_NOCLONE_SORT_INPUTS_SHAPES:
_log_once(
("kmoe-miss-sort-inputs-noclone", shape_key),
"[submission] kmoe miss sort inputs noclone "
f"shape={shape_key}",
)
if cached_moe_buf is None:
raise RuntimeError("Exact-sort cache miss did not return moe_buf.")
moe_buf = cached_moe_buf
if cache_key_to_store is None:
raise RuntimeError("Exact-sort cache miss did not return cache key.")
if skip_store_after is not None and miss_streak > skip_store_after:
skip_cache_store = True
if _KMOE_FULL_EXACT_SORT_CACHE:
_KMOE_FULL_EXACT_SORT_CACHE.clear()
_log_once(
(
"kmoe-full-exact-sort-cache-clear-skip-store",
shape_key,
miss_streak,
),
"[submission] kmoe full exact sort cache clear "
f"shape={shape_key} miss_streak={miss_streak}",
)
else:
_store_kmoe_full_exact_sort_cache(
shape_key,
cache_key_to_store,
(
cached_sorted_ids,
cached_sorted_weights,
cached_sorted_expert_ids,
cached_num_valid_ids,
),
)
_log_once(
("kmoe-full-exact-sort-cache-store-immediate", cache_key_to_store),
"[submission] kmoe full exact sort cache immediate store "
f"shape={shape_key}",
)
_log_once(
("manual2stage-cached-full-exact-sort", shape_key, cache_hit, hit_clone_sort_inputs),
"[submission] manual fused 2stage cached full exact sort "
f"shape={shape_key} clone_inputs={int(hit_clone_sort_inputs)} "
f"cache_hit={int(cache_hit)} use_sort_moebuf={int(not cache_hit)} "
f"miss_clone_sort_inputs={int(miss_clone_sort_inputs)} "
f"miss_streak={miss_streak} skip_cache_store={int(skip_cache_store)} "
f"bypass_cache={int(bypass_exact_sort_cache)}",
)
else:
sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf = (
_aiter_fused_moe.moe_sorting(
topk_ids,
topk_weights,
num_experts,
model_dim,
hidden_states.dtype,
block_m,
None,
None,
0,
)
)
_log_once(
("manual2stage-sorted", shape_key),
f"[submission] manual fused 2stage sorted shape={shape_key}",
)
stage2_boundary_contract_reason: Optional[str] = None
enable_stage2_boundary_contract = False
if shape_key in _MANUAL_FUSED_STAGE2_BOUNDARY_CONTRACT_SHAPES and not cache_hit:
if not bypass_exact_sort_cache and miss_streak == 1:
enable_stage2_boundary_contract = True
stage2_boundary_contract_reason = "first_miss"
elif (
bypass_exact_sort_cache
and first_bypass_miss_streak is not None
and miss_streak == first_bypass_miss_streak
):
enable_stage2_boundary_contract = True
stage2_boundary_contract_reason = "first_bypass"
if enable_stage2_boundary_contract:
stage2_boundary_workspace_mode = (
"fresh"
if stage2_boundary_contract_reason
in _KMOE_STAGE2_BOUNDARY_CONTRACT_FRESH_REASONS
else "cached"
)
_log_once(
(
"manual2stage-stage2-boundary-contract-miss-side-enabled",
shape_key,
stage2_boundary_contract_reason,
),
"[submission] manual fused 2stage stage2-boundary contract miss-side enabled "
f"shape={shape_key} reason={stage2_boundary_contract_reason} "
f"workspace_mode={stage2_boundary_workspace_mode} "
f"miss_streak={miss_streak} bypass_cache={int(bypass_exact_sort_cache)}",
)
stage2_boundary_contract = _get_stage2_boundary_contract(
shape_key,
token_num,
topk,
inter_dim,
sorted_ids.shape[0],
hidden_states.device,
stage2_boundary_contract_reason or "default",
)
else:
if shape_key in _MANUAL_FUSED_STAGE2_BOUNDARY_CONTRACT_SHAPES:
_log_once(
("manual2stage-stage2-boundary-contract-miss-side-disabled", shape_key),
"[submission] manual fused 2stage stage2-boundary contract miss-side disabled "
f"shape={shape_key} cache_hit={int(cache_hit)} miss_streak={miss_streak} "
f"bypass_cache={int(bypass_exact_sort_cache)}",
)
stage2_boundary_contract = None
flydsl_stage2_patch_mode = "rawbuf" if stage2_boundary_contract is not None else "ptr"
if stage2_kernel.startswith("flydsl_"):
_log_once(
(
"manual2stage-stage2-patch-mode",
shape_key,
flydsl_stage2_patch_mode,
bool(stage2_boundary_contract),
),
"[submission] manual fused 2stage stage2 patch mode "
f"shape={shape_key} mode={flydsl_stage2_patch_mode} "
f"boundary_contract={int(stage2_boundary_contract is not None)} "
f"cache_hit={int(cache_hit)} miss_streak={miss_streak} "
f"bypass_cache={int(bypass_exact_sort_cache)}",
)
a1, a1_scale = _run_manual_fused_quant_sort_with_override(
shape_key,
"a1",
hidden_states,
sorted_ids,
num_valid_ids,
token_num,
1,
block_m,
)
metadata = _aiter_fused_moe.get_2stage_cfgs(
_aiter_fused_moe.get_padded_M(token_num),
model_dim,
inter_dim,
num_experts,
topk,
moe_buf.dtype,
dtypes.fp4x2,
dtypes.fp4x2,
QuantType.per_1x32,
True,
ActivationType.Silu,
False,
hidden_pad,
intermediate_pad,
is_shuffled,
)
_log_once(
("manual2stage-stage1-enter", shape_key),
f"[submission] manual fused 2stage stage1 enter shape={shape_key} ksplit={metadata.ksplit}",
)
a2 = torch.empty(
(token_num, topk, inter_dim),
dtype=moe_buf.dtype,
device=hidden_states.device,
)
try:
a2 = metadata.stage1(
a1,
gate_up_weight_shuffled,
down_weight_shuffled,
sorted_ids,
sorted_expert_ids,
num_valid_ids,
a2,
topk,
block_m=block_m,
a1_scale=a1_scale,
w1_scale=(
gate_up_weight_scale_shuffled.view(dtypes.fp8_e8m0)
if gate_up_weight_shuffled.dtype == dtypes.fp4x2
else gate_up_weight_scale_shuffled
),
sorted_weights=None,
)
except Exception as exc:
_log_once(
("manual2stage-stage1-failure", shape_key),
"[submission] manual fused 2stage stage1 failure "
f"shape={shape_key} exc={type(exc).__name__}: {exc}",
)
raise
_log_once(
("manual2stage-stage1-success", shape_key),
f"[submission] manual fused 2stage stage1 success shape={shape_key}",
)
if metadata.ksplit > 1 and is_shuffled:
a2_scale = None
else:
_log_once(
("manual2stage-quant2-enter", shape_key),
f"[submission] manual fused 2stage quant2 enter shape={shape_key}",
)
a2, a2_scale = _run_manual_fused_quant_sort_with_override(
shape_key,
"a2",
a2.view(-1, inter_dim),
sorted_ids,
num_valid_ids,
token_num,
topk,
block_m,
(
None
if stage2_boundary_contract is None
else stage2_boundary_contract["a2_fp4_u8"]
),
(
None
if stage2_boundary_contract is None
else stage2_boundary_contract["a2_scale_u8"]
),
)
a2 = a2.view(token_num, topk, -1)
_log_once(
("manual2stage-stage2-enter", shape_key),
f"[submission] manual fused 2stage stage2 enter shape={shape_key}",
)
try:
_call_with_flydsl_stage2_patch_mode(
metadata.stage2,
flydsl_stage2_patch_mode,
a2,
gate_up_weight_shuffled,
down_weight_shuffled,
sorted_ids,
sorted_expert_ids,
num_valid_ids,
moe_buf,
topk,
w2_scale=(
down_weight_scale_shuffled.view(dtypes.fp8_e8m0)
if down_weight_shuffled.dtype == dtypes.fp4x2
else down_weight_scale_shuffled
),
a2_scale=a2_scale,
block_m=block_m,
sorted_weights=sorted_weights,
)
except Exception as exc:
_log_once(
("manual2stage-stage2-failure", shape_key),
"[submission] manual fused 2stage stage2 failure "
f"shape={shape_key} exc={type(exc).__name__}: {exc}",
)
raise
out = moe_buf
_log_once(
("manual2stage-stage2-success", shape_key),
f"[submission] manual fused 2stage stage2 success shape={shape_key}",
)
_log_once(
("manual2stage-success", shape_key),
f"[submission] manual fused 2stage success shape={shape_key}",
)
return out
def _can_split_shared_expert(
topk_ids: torch.Tensor,
config: Dict[str, int],
) -> bool:
return (
config.get("n_routed_experts", 0) == 32
and config.get("n_shared_experts", 0) == 1
and config.get("n_experts_per_token", 0) == 8
and topk_ids.shape[1] == config.get("total_top_k", topk_ids.shape[1]) == 9
)
def _run_split_shared_expert(
hidden_states: torch.Tensor,
gate_up_weight_shuffled: torch.Tensor,
down_weight_shuffled: torch.Tensor,
gate_up_weight_scale: torch.Tensor,
down_weight_scale: torch.Tensor,
gate_up_weight_scale_shuffled: torch.Tensor,
down_weight_scale_shuffled: torch.Tensor,
topk_weights: torch.Tensor,
topk_ids: torch.Tensor,
config: Dict[str, int],
hidden_pad: int,
intermediate_pad: int,
) -> torch.Tensor:
routed_experts = config["n_routed_experts"]
shared_experts = config["n_shared_experts"]
routed_topk = config["n_experts_per_token"]
hidden_q, hidden_scale = _get_hidden_prequant(hidden_states)
routed_gate_scale_shuffled = _maybe_contiguous(
fp4_utils.e8m0_shuffle(_maybe_contiguous(gate_up_weight_scale[:routed_experts]))
)
routed_down_scale_shuffled = _maybe_contiguous(
fp4_utils.e8m0_shuffle(_maybe_contiguous(down_weight_scale[:routed_experts]))
)
shared_gate_scale_shuffled = _maybe_contiguous(
fp4_utils.e8m0_shuffle(
_maybe_contiguous(
gate_up_weight_scale[routed_experts : routed_experts + shared_experts]
)
)
)
shared_down_scale_shuffled = _maybe_contiguous(
fp4_utils.e8m0_shuffle(
_maybe_contiguous(
down_weight_scale[routed_experts : routed_experts + shared_experts]
)
)
)
routed_out = _run_fused_prequant(
hidden_q,
hidden_scale,
_maybe_contiguous(gate_up_weight_shuffled[:routed_experts]),
_maybe_contiguous(down_weight_shuffled[:routed_experts]),
routed_gate_scale_shuffled,
routed_down_scale_shuffled,
_maybe_contiguous(topk_weights[:, :routed_topk]),
_maybe_contiguous(topk_ids[:, :routed_topk]),
hidden_pad,
intermediate_pad,
)
shared_out = _run_fused_prequant(
hidden_q,
hidden_scale,
_maybe_contiguous(
gate_up_weight_shuffled[routed_experts : routed_experts + shared_experts]
),
_maybe_contiguous(
down_weight_shuffled[routed_experts : routed_experts + shared_experts]
),
shared_gate_scale_shuffled,
shared_down_scale_shuffled,
_maybe_contiguous(topk_weights[:, routed_topk : routed_topk + shared_experts]),
_maybe_contiguous(
topk_ids[:, routed_topk : routed_topk + shared_experts] - routed_experts
),
hidden_pad,
intermediate_pad,
)
return (routed_out.float() + shared_out.float()).to(hidden_states.dtype)
@torch.inference_mode()
def generate_input(
dhidden: int,
dexpert: int,
nroutedexperts: int,
nexpertspertoken: int,
nsharedexperts: int,
bs: int,
seed: int,
) -> input_t:
_require_aiter()
d_hidden = dhidden
d_expert = dexpert
n_routed_experts = nroutedexperts
n_shared_experts = nsharedexperts
routed_top_k = nexpertspertoken
total_top_k = routed_top_k + n_shared_experts
e_total = n_routed_experts + n_shared_experts
m = bs
d_hidden_pad = _pad_to(d_hidden, PAD_ALIGN)
d_expert_pad = _pad_to(d_expert, PAD_ALIGN)
config = {
"d_hidden": d_hidden,
"d_expert": d_expert,
"d_hidden_pad": d_hidden_pad,
"d_expert_pad": d_expert_pad,
"n_routed_experts": n_routed_experts,
"n_shared_experts": n_shared_experts,
"n_experts_per_token": routed_top_k,
"total_top_k": total_top_k,
"bs": m,
}
gen = torch.Generator(device="cuda")
gen.manual_seed(seed)
hidden_states = torch.randn(
(m, d_hidden),
device="cuda",
dtype=torch.bfloat16,
generator=gen,
)
router_weight = torch.randn(
(n_routed_experts, d_hidden),
device="cuda",
dtype=torch.bfloat16,
generator=gen,
) / math.sqrt(d_hidden)
router_logits = F.linear(hidden_states, router_weight)
scores = router_logits.softmax(dim=-1)
routed_weights, routed_ids = torch.topk(scores, k=routed_top_k, dim=-1, sorted=False)
routed_weights = routed_weights.to(torch.float32)
routed_ids = routed_ids.to(torch.int32)
shared_ids = torch.arange(
n_routed_experts,
e_total,
device="cuda",
dtype=torch.int32,
).unsqueeze(0).expand(m, -1)
shared_weights = torch.ones(
(m, n_shared_experts),
device="cuda",
dtype=torch.float32,
)
topk_ids = torch.cat([routed_ids, shared_ids], dim=-1)
topk_weights = torch.cat([routed_weights, shared_weights], dim=-1)
gate_up_bf16 = torch.randn(
(e_total, 2 * d_expert_pad, d_hidden_pad),
device="cuda",
dtype=torch.bfloat16,
generator=gen,
) / math.sqrt(d_hidden)
down_bf16 = torch.randn(
(e_total, d_hidden_pad, d_expert_pad),
device="cuda",
dtype=torch.bfloat16,
generator=gen,
) / math.sqrt(d_expert)
torch_quant = aiter.get_torch_quant(QuantType.per_1x32)
gate_up_weight, gate_up_weight_scale = torch_quant(
gate_up_bf16,
quant_dtype=dtypes.fp4x2,
)
down_weight, down_weight_scale = torch_quant(
down_bf16,
quant_dtype=dtypes.fp4x2,
)
gate_up_weight = gate_up_weight.view(e_total, 2 * d_expert_pad, d_hidden_pad // 2)
down_weight = down_weight.view(e_total, d_hidden_pad, d_expert_pad // 2)
gate_up_weight_shuffled = shuffle_weight(gate_up_weight, layout=SHUFFLE_LAYOUT)
down_weight_shuffled = shuffle_weight(down_weight, layout=SHUFFLE_LAYOUT)
gate_up_weight_scale_shuffled = fp4_utils.e8m0_shuffle(gate_up_weight_scale)
down_weight_scale_shuffled = fp4_utils.e8m0_shuffle(down_weight_scale)
return (
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,
)
@torch.inference_mode()
def ref_kernel(data: input_t) -> output_t:
(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
_gate_up_weight_scale,
_down_weight_scale,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
_config,
hidden_pad,
intermediate_pad,
) = _prepare_common(data)
return _run_fused(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
hidden_pad,
intermediate_pad,
)
@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
gate_up_weight_scale,
down_weight_scale,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
config,
hidden_pad,
intermediate_pad,
) = _prepare_common(data)
shape_key = _shape_key(hidden_states, topk_ids, config)
_maybe_reset_full_exact_sort_cache(shape_key)
cache_key = _backend_key(hidden_states, topk_ids, config)
cached_backend = _BACKEND_CACHE.get(cache_key)
forced_non_temporal_load = _forced_ck_non_temporal_load(hidden_states, topk_ids, config)
forced_stage_kernels = _forced_ck_stage_kernels(hidden_states, topk_ids, config)
fused_stage_override = _FUSED_CFG_STAGE_OVERRIDES.get(shape_key)
use_manual_fused_2stage = _should_use_manual_fused_2stage(
hidden_states,
topk_ids,
config,
)
use_ck_prequant = _should_prequant_ck(hidden_states, topk_ids, config)
on_mi355x = _device_is_mi355x(hidden_states.device)
full_topk = topk_ids.shape[1] == config.get("total_top_k", topk_ids.shape[1])
if cached_backend is not None:
if cached_backend[0] == "split_fused":
try:
return _run_split_shared_expert(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
gate_up_weight_scale,
down_weight_scale,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
config,
hidden_pad,
intermediate_pad,
)
except Exception:
_BACKEND_CACHE[cache_key] = ("fused", None)
elif cached_backend[0] == "manual2stage":
try:
if fused_stage_override is None:
raise RuntimeError("Missing fused stage override for manual route.")
return _run_manual_fused_2stage_prequant(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
config,
hidden_pad,
intermediate_pad,
fused_stage_override[0],
fused_stage_override[1],
fused_stage_override[2],
)
except Exception as exc:
_log_once(
("manual2stage-fallback", shape_key),
"[submission] manual fused 2stage fallback "
f"shape={shape_key} exc={type(exc).__name__}: {exc}",
)
_BACKEND_CACHE[cache_key] = ("fused", None)
elif cached_backend[0] == "ck_stage":
try:
if forced_stage_kernels is None:
raise RuntimeError("Missing CK stage-kernel override.")
return _run_ck_prequant_with_stage_kernels(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
forced_stage_kernels[2],
forced_non_temporal_load,
forced_stage_kernels[0],
forced_stage_kernels[1],
)
except Exception:
_BACKEND_CACHE[cache_key] = ("ck", cached_backend[1])
return _run_ck_prequant(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
cached_backend[1],
forced_non_temporal_load,
)
elif cached_backend[0] == "ck":
try:
if use_ck_prequant:
return _run_ck_prequant(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
cached_backend[1],
forced_non_temporal_load,
)
return _run_ck(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
cached_backend[1],
forced_non_temporal_load,
)
except Exception:
_BACKEND_CACHE[cache_key] = ("fused", None)
return _run_fused(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
hidden_pad,
intermediate_pad,
)
if _can_split_shared_expert(topk_ids, config) and shape_key not in _SKIP_SPLIT_SHARED_EXPERT_SHAPES:
try:
_BACKEND_CACHE[cache_key] = ("split_fused", None)
return _run_split_shared_expert(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
gate_up_weight_scale,
down_weight_scale,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
config,
hidden_pad,
intermediate_pad,
)
except Exception:
_BACKEND_CACHE[cache_key] = ("fused", None)
if use_manual_fused_2stage and fused_stage_override is not None and on_mi355x and full_topk:
try:
_BACKEND_CACHE[cache_key] = ("manual2stage", None)
return _run_manual_fused_2stage_prequant(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
config,
hidden_pad,
intermediate_pad,
fused_stage_override[0],
fused_stage_override[1],
fused_stage_override[2],
)
except Exception as exc:
_log_once(
("manual2stage-fallback", shape_key),
"[submission] manual fused 2stage fallback "
f"shape={shape_key} exc={type(exc).__name__}: {exc}",
)
_BACKEND_CACHE[cache_key] = ("fused", None)
if fused_stage_override is not None and on_mi355x and full_topk:
try:
_BACKEND_CACHE[cache_key] = ("fused", None)
return _run_fused(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
hidden_pad,
intermediate_pad,
)
except Exception:
_BACKEND_CACHE.pop(cache_key, None)
benchmark_block_m = _BENCHMARK_BLOCK_M.get(shape_key)
if (
_CK_MOE is not None
and benchmark_block_m is not None
and on_mi355x
and full_topk
):
if forced_stage_kernels is not None and use_ck_prequant:
try:
_BACKEND_CACHE[cache_key] = ("ck_stage", forced_stage_kernels[2])
return _run_ck_prequant_with_stage_kernels(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
forced_stage_kernels[2],
forced_non_temporal_load,
forced_stage_kernels[0],
forced_stage_kernels[1],
)
except Exception:
_BACKEND_CACHE[cache_key] = ("ck", benchmark_block_m)
try:
_BACKEND_CACHE[cache_key] = ("ck", benchmark_block_m)
if use_ck_prequant:
return _run_ck_prequant(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
benchmark_block_m,
forced_non_temporal_load,
)
return _run_ck(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
benchmark_block_m,
forced_non_temporal_load,
)
except Exception:
_BACKEND_CACHE[cache_key] = ("fused", None)
_BACKEND_CACHE[cache_key] = ("fused", None)
return _run_fused(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
hidden_pad,
intermediate_pad,
)
kernel = custom_kernel
custom_generate_input = generate_input
create_input = generate_input
reference_kernel = ref_kernel
if make_match_reference is not None:
check_implementation = make_match_reference(custom_kernel, rtol=RTOL, atol=ATOL)
else:
check_implementation = None
import functools
import aiter.fused_moe as _aiter_fused_moe
_FLYDSL_STAGE2_SOURCE_PROBED = False
def _probe_flydsl_stage2_builder_source() -> None:
global _FLYDSL_STAGE2_SOURCE_PROBED
if _FLYDSL_STAGE2_SOURCE_PROBED or aiter is None:
return
try:
import ast as _ast
import hashlib as _hashlib
import inspect as _inspect
import aiter.ops.flydsl.kernels.mixed_moe_gemm_2stage as _mixed_stage2
import aiter.ops.flydsl.kernels.mfma_epilogues as _mfma_epilogues
import aiter.ops.flydsl.moe_kernels as _flydsl_moe_kernels
except Exception as exc:
_log_once(
("flydsl-stage2-source-probe-import-failure", type(exc).__name__),
"[submission] kmoe probe flydsl stage2 skipped "
f"reason=import_failure exc={type(exc).__name__}",
)
return
def _find_matches(lines, predicate):
return [idx for idx, line in enumerate(lines) if predicate(line)]
def _collect_line_numbers(lines, predicate):
return [idx + 1 for idx in _find_matches(lines, predicate)]
def _emit_contexts(
digest: str,
lines,
label: str,
predicate,
radius: int = 6,
max_matches: int = 1,
) -> None:
matches = _find_matches(lines, predicate)
if not matches:
_log_once(
("flydsl-stage2-source-probe-context", digest, label, "missing"),
"[submission] kmoe probe flydsl stage2 "
f"context={label} status=missing",
)
return
for match_idx, center in enumerate(matches[:max_matches], 1):
start = max(0, center - radius)
end = min(len(lines), center + radius + 1)
snippet = "\n".join(
f"{idx + 1}:{lines[idx]}" for idx in range(start, end)
)
_log_once(
("flydsl-stage2-source-probe-context", digest, label, center),
"[submission] kmoe probe flydsl stage2 "
f"context={label} idx={match_idx} line={center + 1}\n{snippet}",
)
def _emit_line_contexts(
digest: str,
lines,
label: str,
line_numbers,
radius: int = 6,
max_matches: int = 1,
) -> None:
if not line_numbers:
_log_once(
("flydsl-stage2-source-probe-line-context", digest, label, "missing"),
"[submission] kmoe probe flydsl stage2 "
f"context={label} status=missing",
)
return
for match_idx, line_no in enumerate(line_numbers[:max_matches], 1):
center = int(line_no) - 1
start = max(0, center - radius)
end = min(len(lines), center + radius + 1)
snippet = "\n".join(
f"{idx + 1}:{lines[idx]}" for idx in range(start, end)
)
_log_once(
("flydsl-stage2-source-probe-line-context", digest, label, center),
"[submission] kmoe probe flydsl stage2 "
f"context={label} idx={match_idx} line={line_no}\n{snippet}",
)
def _ast_name(node) -> str:
if node is None:
return "None"
if isinstance(node, _ast.Name):
return node.id
if isinstance(node, _ast.Attribute):
base = _ast_name(node.value)
return f"{base}.{node.attr}" if base else node.attr
try:
return _ast.unparse(node)
except Exception:
return type(node).__name__
def _probe_ast(label: str, digest: str, source_text: str):
try:
module_ast = _ast.parse(source_text)
except SyntaxError as exc:
_log_once(
(
"flydsl-stage2-source-probe-ast-failure",
digest,
label,
type(exc).__name__,
),
"[submission] kmoe probe flydsl stage2 ast "
f"target={label} reason=parse_failure exc={type(exc).__name__}",
)
return {}, {}
parents = {}
for parent in _ast.walk(module_ast):
for child in _ast.iter_child_nodes(parent):
parents[id(child)] = parent
def _scope(node) -> str:
parts = []
cur = parents.get(id(node))
while cur is not None:
if isinstance(
cur, (_ast.FunctionDef, _ast.AsyncFunctionDef, _ast.ClassDef)
):
parts.append(cur.name)
cur = parents.get(id(cur))
parts.reverse()
return "/".join(parts) if parts else "<module>"
interesting_defs = {
"compile_mixed_moe_gemm2": {
"compile_mixed_moe_gemm2",
"moe_gemm2",
"write_row_to_lds",
"precompute_row",
"store_pair",
},
"c_shuffle_epilog": {
"c_shuffle_epilog",
"_write_row",
"_do_store_row",
},
"default_epilog": {
"default_epilog",
},
}.get(label, set())
interesting_calls = {
"c_shuffle_epilog",
"default_epilog",
"precompute_row",
"write_row_to_lds",
"store_pair",
}
def_lines = {}
call_infos = []
for node in _ast.walk(module_ast):
if isinstance(node, (_ast.FunctionDef, _ast.AsyncFunctionDef, _ast.ClassDef)):
if node.name in interesting_defs:
def_lines.setdefault(node.name, []).append(node.lineno)
continue
if not isinstance(node, _ast.Call):
continue
func_name = _ast_name(node.func)
short_name = func_name.rsplit(".", 1)[-1]
if short_name not in interesting_calls:
continue
kwargs = []
for kw in node.keywords:
if kw.arg is None:
continue
kwargs.append(f"{kw.arg}={_ast_name(kw.value)}")
call_infos.append(
{
"name": short_name,
"func": func_name,
"line": node.lineno,
"scope": _scope(node),
"kwargs": kwargs,
}
)
call_lines = {}
for info in call_infos:
call_lines.setdefault(info["name"], []).append(info["line"])
def_summary = ",".join(
f"{name}:{def_lines[name]}" for name in sorted(def_lines)
) or "-"
call_summary = ",".join(
f"{name}:{call_lines[name]}" for name in sorted(call_lines)
) or "-"
_log_once(
("flydsl-stage2-source-probe-ast-summary", digest, label),
"[submission] kmoe probe flydsl stage2 ast "
f"target={label} defs={def_summary} calls={call_summary}",
)
for info in call_infos:
kw_text = "|".join(info["kwargs"]) if info["kwargs"] else "-"
_log_once(
(
"flydsl-stage2-source-probe-ast-call",
digest,
label,
info["name"],
info["line"],
),
"[submission] kmoe probe flydsl stage2 ast "
f"target={label} call={info['name']} func={info['func']} "
f"line={info['line']} scope={info['scope']} kwargs={kw_text}",
)
return def_lines, call_lines
targets = [
("compile_mixed_moe_gemm2", _mixed_stage2.compile_mixed_moe_gemm2),
]
for label, fn in targets:
target_fn = _inspect.unwrap(fn)
try:
source = _inspect.getsource(target_fn)
source_file = _inspect.getsourcefile(target_fn) or "unknown"
firstlineno = getattr(
getattr(target_fn, "__code__", None), "co_firstlineno", -1
)
except (OSError, TypeError) as exc:
_log_once(
("flydsl-stage2-source-probe-source-failure", label, type(exc).__name__),
"[submission] kmoe probe flydsl stage2 "
f"target={label} reason=source_failure exc={type(exc).__name__}",
)
continue
lines = source.splitlines()
normalized = "\n".join(line.rstrip() for line in lines).strip()
digest = _hashlib.sha1(normalized.encode("utf-8")).hexdigest()[:16]
ast_def_lines = {}
ast_call_lines = {}
if label in ("compile_mixed_moe_gemm2", "c_shuffle_epilog", "default_epilog"):
ast_def_lines, ast_call_lines = _probe_ast(label, digest, source)
if label == "compile_mixed_moe_gemm2":
module_name_lines = _collect_line_numbers(
lines, lambda line: "module_name" in line
)
lds_out_bytes_lines = _collect_line_numbers(
lines, lambda line: "lds_out_bytes" in line
)
lds_total_bytes_lines = _collect_line_numbers(
lines, lambda line: "lds_total_bytes" in line
)
lds_alloc_bytes_lines = _collect_line_numbers(
lines, lambda line: "lds_alloc_bytes" in line
)
lds_x_decl_lines = _collect_line_numbers(
lines, lambda line: '_state["lds_x_decl"]' in line
)
allocate_array_lines = _collect_line_numbers(
lines, lambda line: "allocator.allocate_array(" in line
)
base_ptr_lines = _collect_line_numbers(
lines, lambda line: "base_ptr" in line
)
lds_x_ptr_lines = _collect_line_numbers(
lines, lambda line: "lds_x_ptr" in line
)
lds_out_lines = _collect_line_numbers(
lines, lambda line: "lds_out" in line
)
smem_ptr_lines = _collect_line_numbers(
lines, lambda line: "SmemPtr(" in line
)
write_row_lines = _collect_line_numbers(
lines, lambda line: "def write_row_to_lds" in line
)
write_row_call_lines = _collect_line_numbers(
lines,
lambda line: "write_row_to_lds(" in line
and "def write_row_to_lds" not in line,
)
precompute_row_lines = _collect_line_numbers(
lines, lambda line: "def precompute_row" in line
)
precompute_row_call_lines = _collect_line_numbers(
lines,
lambda line: "precompute_row(" in line
and "def precompute_row" not in line,
)
fused2_assign_lines = _collect_line_numbers(
lines, lambda line: "fused2 = buffer_ops.buffer_load(" in line
)
row_valid_lines = _collect_line_numbers(
lines, lambda line: "row_valid" in line
)
vector_store_lds_out_lines = _collect_line_numbers(
lines, lambda line: "vector.store(v1, lds_out" in line
)
sorted_rsrc_load_lines = _collect_line_numbers(
lines,
lambda line: "sorted_rsrc" in line
and "buffer_ops.buffer_load" in line,
)
sorted_w_rsrc_lines = _collect_line_numbers(
lines, lambda line: "sorted_w_rsrc" in line
)
sorted_w_rsrc_load_lines = _collect_line_numbers(
lines,
lambda line: "sorted_w_rsrc" in line
and "buffer_ops.buffer_load" in line,
)
doweight_stage2_lines = _collect_line_numbers(
lines, lambda line: "doweight_stage2" in line
)
tw_assign_lines = _collect_line_numbers(
lines, lambda line: "tw =" in line
)
tw_pf_use_lines = _collect_line_numbers(
lines,
lambda line: "tw_pf[" in line
or "tw = tw_pf" in line
or "tw_pf is not None" in line,
)
lds_tid_lines = _collect_line_numbers(
lines, lambda line: "lds_tid" in line
)
memref_load_lds_tid_lines = _collect_line_numbers(
lines, lambda line: "memref.load(lds_tid" in line
)
tw_pf_lines = _collect_line_numbers(lines, lambda line: "tw_pf" in line)
use_cshuffle_epilog_lines = _collect_line_numbers(
lines, lambda line: "_use_cshuffle_epilog" in line
)
ast_cshuffle_call_lines = ast_call_lines.get("c_shuffle_epilog", [])
ast_store_pair_def_lines = ast_def_lines.get("store_pair", [])
summary = (
"[submission] kmoe probe flydsl stage2 "
f"target={label} file={source_file} firstlineno={firstlineno} "
f"sha1={digest} lines={len(lines)} "
f"wrapper_type={type(fn).__name__} unwrapped_type={type(target_fn).__name__} "
f"used_unwrap={int(target_fn is not fn)} "
f"module_name={module_name_lines} "
f"lds_out_bytes={lds_out_bytes_lines} "
f"lds_total_bytes={lds_total_bytes_lines} "
f"lds_alloc_bytes={lds_alloc_bytes_lines} "
f"lds_x_decl={lds_x_decl_lines} "
f"allocate_array={allocate_array_lines} "
f"base_ptr={base_ptr_lines} "
f"lds_x_ptr={lds_x_ptr_lines} "
f"lds_out={lds_out_lines} "
f"SmemPtr={smem_ptr_lines} "
f"write_row={write_row_lines} "
f"write_row_call={write_row_call_lines} "
f"precompute_row={precompute_row_lines} "
f"precompute_row_call={precompute_row_call_lines} "
f"fused2_assign={fused2_assign_lines} "
f"row_valid={row_valid_lines} "
f"vector_store_lds_out={vector_store_lds_out_lines} "
f"sorted_rsrc_load={sorted_rsrc_load_lines} "
f"sorted_w_rsrc={sorted_w_rsrc_lines} "
f"sorted_w_rsrc_load={sorted_w_rsrc_load_lines} "
f"doweight_stage2={doweight_stage2_lines} "
f"tw_assign={tw_assign_lines} "
f"tw_pf_use={tw_pf_use_lines} "
f"lds_tid={lds_tid_lines} "
f"memref_load_lds_tid={memref_load_lds_tid_lines} "
f"tw_pf={tw_pf_lines} "
f"use_cshuffle_epilog={use_cshuffle_epilog_lines} "
f"ast_c_shuffle_epilog_call={ast_cshuffle_call_lines} "
f"ast_store_pair_def={ast_store_pair_def_lines}"
)
_log_once(("flydsl-stage2-source-probe-summary", digest), summary)
for context_label, predicate, radius, max_matches in (
("module_name", lambda line: "module_name" in line, 6, 1),
("lds_out_bytes", lambda line: "lds_out_bytes" in line, 6, 1),
("lds_total_bytes", lambda line: "lds_total_bytes" in line, 6, 1),
("lds_alloc_bytes", lambda line: "lds_alloc_bytes" in line, 6, 1),
("lds_x_decl", lambda line: '_state["lds_x_decl"]' in line, 6, 1),
("allocate_array", lambda line: "allocator.allocate_array(" in line, 6, 1),
("base_ptr", lambda line: "base_ptr" in line, 6, 1),
("lds_x_ptr", lambda line: "lds_x_ptr" in line, 6, 1),
("lds_out", lambda line: "lds_out" in line, 6, 2),
("SmemPtr", lambda line: "SmemPtr(" in line, 6, 1),
("write_row_to_lds", lambda line: "def write_row_to_lds" in line, 8, 1),
(
"write_row_call",
lambda line: "write_row_to_lds(" in line
and "def write_row_to_lds" not in line,
6,
2,
),
("precompute_row", lambda line: "def precompute_row" in line, 8, 1),
(
"precompute_row_call",
lambda line: "precompute_row(" in line
and "def precompute_row" not in line,
6,
2,
),
(
"fused2_assign",
lambda line: "fused2 = buffer_ops.buffer_load(" in line,
6,
1,
),
("row_valid", lambda line: "row_valid" in line, 4, 2),
(
"vector_store_lds_out",
lambda line: "vector.store(v1, lds_out" in line,
6,
1,
),
(
"sorted_rsrc_load",
lambda line: "sorted_rsrc" in line
and "buffer_ops.buffer_load" in line,
6,
1,
),
("sorted_w_rsrc", lambda line: "sorted_w_rsrc" in line, 6, 3),
(
"sorted_w_rsrc_load",
lambda line: "sorted_w_rsrc" in line
and "buffer_ops.buffer_load" in line,
6,
3,
),
("doweight_stage2", lambda line: "doweight_stage2" in line, 6, 3),
("tw_assign", lambda line: "tw =" in line, 6, 3),
(
"tw_pf_use",
lambda line: "tw_pf[" in line
or "tw = tw_pf" in line
or "tw_pf is not None" in line,
6,
3,
),
("memref_load_lds_tid", lambda line: "memref.load(lds_tid" in line, 6, 1),
("tw_pf", lambda line: "tw_pf" in line, 4, 2),
("use_cshuffle_epilog", lambda line: "_use_cshuffle_epilog" in line, 4, 2),
("c_shuffle_epilog", lambda line: "c_shuffle_epilog(" in line, 6, 1),
):
_emit_contexts(
digest,
lines,
context_label,
predicate,
radius=radius,
max_matches=max_matches,
)
_emit_line_contexts(
digest,
lines,
"ast_c_shuffle_epilog_call",
ast_cshuffle_call_lines,
radius=8,
max_matches=1,
)
_emit_line_contexts(
digest,
lines,
"ast_store_pair_def",
ast_store_pair_def_lines,
radius=8,
max_matches=1,
)
elif label == "c_shuffle_epilog":
precomputed_rows_lines = _collect_line_numbers(
lines, lambda line: "_precomputed_rows = []" in line
)
row_ctx_raw_lines = _collect_line_numbers(
lines, lambda line: "row_ctx_raw =" in line
)
summary = (
"[submission] kmoe probe flydsl stage2 "
f"target={label} file={source_file} firstlineno={firstlineno} "
f"sha1={digest} lines={len(lines)} "
f"wrapper_type={type(fn).__name__} unwrapped_type={type(target_fn).__name__} "
f"used_unwrap={int(target_fn is not fn)} "
f"precomputed_rows={precomputed_rows_lines} "
f"row_ctx_raw={row_ctx_raw_lines} "
f"ast__write_row={ast_def_lines.get('_write_row', [])} "
f"ast__do_store_row={ast_def_lines.get('_do_store_row', [])} "
f"ast_default_epilog_call={ast_call_lines.get('default_epilog', [])} "
f"ast_precompute_row_call={ast_call_lines.get('precompute_row', [])} "
f"ast_write_row_to_lds_call={ast_call_lines.get('write_row_to_lds', [])} "
f"ast_store_pair_call={ast_call_lines.get('store_pair', [])}"
)
_log_once(
("flydsl-stage2-source-probe-target", label, digest),
summary,
)
for context_label, predicate, radius, max_matches in (
("_write_row", lambda line: "def _write_row" in line, 6, 1),
(
"default_epilog_call",
lambda line: "default_epilog(" in line,
6,
1,
),
(
"precomputed_rows",
lambda line: "_precomputed_rows = []" in line,
6,
1,
),
("row_ctx_raw", lambda line: "row_ctx_raw =" in line, 6, 1),
("_do_store_row", lambda line: "def _do_store_row" in line, 6, 1),
(
"store_pair_call",
lambda line: "store_pair(" in line
and "def store_pair" not in line,
6,
1,
),
):
_emit_contexts(
digest,
lines,
context_label,
predicate,
radius=radius,
max_matches=max_matches,
)
elif label == "default_epilog":
body_row_lines = _collect_line_numbers(
lines, lambda line: "body_row(" in line
)
summary = (
"[submission] kmoe probe flydsl stage2 "
f"target={label} file={source_file} firstlineno={firstlineno} "
f"sha1={digest} lines={len(lines)} "
f"wrapper_type={type(fn).__name__} unwrapped_type={type(target_fn).__name__} "
f"used_unwrap={int(target_fn is not fn)} "
f"body_row={body_row_lines}"
)
_log_once(
("flydsl-stage2-source-probe-target", label, digest),
summary,
)
_emit_contexts(
digest,
lines,
"body_row_call",
lambda line: "body_row(" in line,
radius=6,
max_matches=1,
)
else:
_log_once(
("flydsl-stage2-source-probe-target", label, digest),
"[submission] kmoe probe flydsl stage2 "
f"target={label} file={source_file} firstlineno={firstlineno} "
f"sha1={digest} lines={len(lines)} "
f"wrapper_type={type(fn).__name__} unwrapped_type={type(target_fn).__name__} "
f"used_unwrap={int(target_fn is not fn)}",
)
_FLYDSL_STAGE2_SOURCE_PROBED = True
if False:
_probe_flydsl_stage2_builder_source()
_FLYDSL_STAGE2_SORTEDIDX_LDS_PATCHED = False
def _install_flydsl_stage2_multi_mode_compiled_cache(_flydsl_moe_kernels) -> None:
global _KMOE_FLYDSL_STAGE2_ORIG_GET_COMPILED_STAGE2
global _KMOE_FLYDSL_STAGE2_COMPILED_CACHE_PATCHED
if _KMOE_FLYDSL_STAGE2_COMPILED_CACHE_PATCHED:
return
target_fn = getattr(_flydsl_moe_kernels, "_get_compiled_stage2", None)
if target_fn is None:
_log_once(
("flydsl-stage2-compiled-cache-install-missing-target",),
"[submission] kmoe flydsl stage2 compiled cache install skipped "
"reason=missing_get_compiled_stage2",
)
return
_KMOE_FLYDSL_STAGE2_ORIG_GET_COMPILED_STAGE2 = target_fn
def _patched_get_compiled_stage2(
model_dim: int,
inter_dim: int,
experts: int,
topk: int,
tile_m: int,
tile_n: int,
tile_k: int,
doweight: bool,
a_dtype: str,
b_dtype: str,
out_dtype: str,
accumulate: bool = True,
):
patch_mode = (
_KMOE_ACTIVE_FLYDSL_STAGE2_PATCH_MODE
or _KMOE_PENDING_FLYDSL_STAGE2_PATCH_MODE
)
cache_key = (
patch_mode,
model_dim,
inter_dim,
experts,
topk,
tile_m,
tile_n,
tile_k,
doweight,
a_dtype,
b_dtype,
out_dtype,
accumulate,
)
cached = _KMOE_FLYDSL_STAGE2_COMPILED_CACHE.get(cache_key)
if cached is not None:
_log_once(
("flydsl-stage2-compiled-cache-hit", cache_key),
"[submission] kmoe flydsl stage2 compiled cache hit "
f"patch_mode={patch_mode} tile={tile_m}x{tile_n}x{tile_k} "
f"accumulate={int(accumulate)}",
)
return cached
try:
target_fn.cache_clear()
except Exception:
pass
compiled = target_fn(
model_dim=model_dim,
inter_dim=inter_dim,
experts=experts,
topk=topk,
tile_m=tile_m,
tile_n=tile_n,
tile_k=tile_k,
doweight=doweight,
a_dtype=a_dtype,
b_dtype=b_dtype,
out_dtype=out_dtype,
accumulate=accumulate,
)
_KMOE_FLYDSL_STAGE2_COMPILED_CACHE[cache_key] = compiled
_log_once(
("flydsl-stage2-compiled-cache-build", cache_key),
"[submission] kmoe flydsl stage2 compiled cache build "
f"patch_mode={patch_mode} tile={tile_m}x{tile_n}x{tile_k} "
f"accumulate={int(accumulate)}",
)
return compiled
_flydsl_moe_kernels._get_compiled_stage2 = _patched_get_compiled_stage2
_KMOE_FLYDSL_STAGE2_COMPILED_CACHE_PATCHED = True
_log_once(
("flydsl-stage2-compiled-cache-install",),
"[submission] kmoe flydsl stage2 compiled cache multi-mode install",
)
def _patch_flydsl_stage2_sortedidx_lds_prefetch(patch_mode: str = "ptr") -> None:
global _FLYDSL_STAGE2_SORTEDIDX_LDS_PATCHED, _KMOE_ACTIVE_FLYDSL_STAGE2_PATCH_MODE
global _KMOE_FLYDSL_STAGE2_ORIG_SOURCE, _KMOE_FLYDSL_STAGE2_ORIG_SOURCE_FILE
if aiter is None:
return
if patch_mode not in ("ptr", "rawbuf"):
raise ValueError(f"unsupported flydsl stage2 patch mode: {patch_mode}")
if _KMOE_ACTIVE_FLYDSL_STAGE2_PATCH_MODE == patch_mode:
return
try:
import ast as _ast
import hashlib as _hashlib
import inspect as _inspect
import aiter.ops.flydsl.kernels.mixed_moe_gemm_2stage as _mixed_stage2
import aiter.ops.flydsl.moe_kernels as _flydsl_moe_kernels
except Exception as exc:
_log_once(
("flydsl-stage2-sortedidx-lds-patch-import-failure", type(exc).__name__),
"[submission] kmoe patch flydsl stage2 sortedidx lds skipped "
f"reason=import_failure exc={type(exc).__name__}",
)
return
_install_flydsl_stage2_multi_mode_compiled_cache(_flydsl_moe_kernels)
cached_patched_fn = _KMOE_FLYDSL_STAGE2_PATCHED_FN_BY_MODE.get(patch_mode)
if cached_patched_fn is not None:
_mixed_stage2.compile_mixed_moe_gemm2 = cached_patched_fn
_KMOE_ACTIVE_FLYDSL_STAGE2_PATCH_MODE = patch_mode
_log_once(
("flydsl-stage2-sortedidx-lds-patch-reuse", patch_mode),
"[submission] kmoe patch flydsl stage2 sortedidx lds reuse "
f"patch_mode={patch_mode}",
)
return
target_fn = getattr(_mixed_stage2, "compile_mixed_moe_gemm2", None)
if target_fn is None:
_log_once(
("flydsl-stage2-sortedidx-lds-patch-missing-target",),
"[submission] kmoe patch flydsl stage2 sortedidx lds skipped "
"reason=missing_compile_mixed_moe_gemm2",
)
return
if patch_mode == "rawbuf":
_mixed_stage2.compile_mixed_moe_gemm2 = target_fn
_KMOE_FLYDSL_STAGE2_PATCHED_FN_BY_MODE[patch_mode] = target_fn
_KMOE_ACTIVE_FLYDSL_STAGE2_PATCH_MODE = patch_mode
_log_once(
("flydsl-stage2-sortedidx-lds-patch-rawbuf-orig",),
"[submission] kmoe patch flydsl stage2 sortedidx lds bypass "
"mode=orig_builder_rawbuf patch_mode=rawbuf",
)
return
if _KMOE_FLYDSL_STAGE2_ORIG_SOURCE is None:
target_unwrapped = _inspect.unwrap(target_fn)
try:
source = _inspect.getsource(target_unwrapped)
source_file = _inspect.getsourcefile(target_unwrapped) or "unknown"
except (OSError, TypeError) as exc:
_log_once(
("flydsl-stage2-sortedidx-lds-patch-source-failure", type(exc).__name__),
"[submission] kmoe patch flydsl stage2 sortedidx lds skipped "
f"reason=source_failure exc={type(exc).__name__}",
)
return
_KMOE_FLYDSL_STAGE2_ORIG_SOURCE = source
_KMOE_FLYDSL_STAGE2_ORIG_SOURCE_FILE = source_file
else:
source = _KMOE_FLYDSL_STAGE2_ORIG_SOURCE
source_file = _KMOE_FLYDSL_STAGE2_ORIG_SOURCE_FILE or "unknown"
patch_tag = {
"ptr": "_vscale_fix3_sortedidxlds18ptr",
"rawbuf": "_vscale_fix4_sortedidxlds18rawbufhotloop",
}[patch_mode]
mode_label = {
"ptr": "old_builder_sortedidx_lds_v18_ptr",
"rawbuf": "old_builder_sortedidx_lds_v18_rawbuf_hotloop",
}[patch_mode]
rowctx_label = {
"ptr": "idx0_i32_acc_ptr_runtime",
"rawbuf": "idx0_i32_acc_rawbuf_miss_only",
}[patch_mode]
if patch_mode == "rawbuf":
store_pair_body = (
"frag_v = frag._value if hasattr(frag, \"_value\") else frag\n"
"if bool(accumulate):\n"
" idx0 = row_ctx\n"
" col_i32 = arith.index_cast(i32, col_g0)\n"
" idx_elem = idx0 + col_i32\n"
" idx_elem_even = idx_elem & arith.constant(0xFFFFFFFE, type=i32)\n"
" byte_off = idx_elem_even * arith.constant(2, type=i32)\n"
" atomic_add_f16x2(frag, byte_off)\n"
" return\n"
"row_byte_base = row_ctx\n"
"col_idx = col_g0\n"
"byte_off_col = col_idx * arith.constant(out_elem_bytes, index=True)\n"
"ptr_addr_idx = row_byte_base + byte_off_col\n"
"out_ptr = buffer_ops.create_llvm_ptr(ptr_addr_idx, address_space=1)\n"
"out_ptr_v = out_ptr._value if hasattr(out_ptr, \"_value\") else out_ptr\n"
"llvm.StoreOp(frag_v, out_ptr_v, alignment=4)\n"
)
else:
store_pair_body = (
"frag_v = frag._value if hasattr(frag, \"_value\") else frag\n"
"if bool(accumulate):\n"
" idx0 = row_ctx\n"
" col_i32 = arith.index_cast(i32, col_g0)\n"
" idx_elem = idx0 + col_i32\n"
" idx_elem_even = idx_elem & arith.constant(0xFFFFFFFE, type=i32)\n"
" byte_off = idx_elem_even * arith.constant(2, type=i32)\n"
" byte_off_idx = arith.index_cast(ir.IndexType.get(), byte_off)\n"
" ptr_addr_idx = out_base_idx + byte_off_idx\n"
" out_ptr = buffer_ops.create_llvm_ptr(ptr_addr_idx, address_space=1)\n"
" out_ptr_v = out_ptr._value if hasattr(out_ptr, \"_value\") else out_ptr\n"
" llvm.AtomicRMWOp(\n"
" llvm.AtomicBinOp.fadd,\n"
" out_ptr_v,\n"
" frag_v,\n"
" llvm.AtomicOrdering.monotonic,\n"
" syncscope=\"agent\",\n"
" alignment=4,\n"
" )\n"
" return\n"
"row_byte_base = row_ctx\n"
"col_idx = col_g0\n"
"byte_off_col = col_idx * arith.constant(out_elem_bytes, index=True)\n"
"ptr_addr_idx = row_byte_base + byte_off_col\n"
"out_ptr = buffer_ops.create_llvm_ptr(ptr_addr_idx, address_space=1)\n"
"out_ptr_v = out_ptr._value if hasattr(out_ptr, \"_value\") else out_ptr\n"
"llvm.StoreOp(frag_v, out_ptr_v, alignment=4)\n"
)
normalized = "\n".join(line.rstrip() for line in source.splitlines()).strip()
digest = _hashlib.sha1(normalized.encode("utf-8")).hexdigest()[:16]
def _replace_once(text: str, label: str, old: str, new: str) -> str:
if old not in text:
raise RuntimeError(label)
return text.replace(old, new, 1)
def _rewrite_rowctx_contract(text: str) -> tuple[str, bool, bool]:
module_ast = _ast.parse(text)
class _RowCtxFixer(_ast.NodeTransformer):
def __init__(self):
self.precompute_fixed = False
self.store_pair_fixed = False
def visit_FunctionDef(self, node):
node = self.generic_visit(node)
if not isinstance(node, _ast.FunctionDef):
return node
if node.name == "precompute_row" and not self.precompute_fixed:
node.body = _ast.parse(
"fused2 = memref.load(lds_tid, [row_local])\n"
"row_i32 = arith.index_cast(i32, row)\n"
"row_valid0 = arith.cmpu(row_i32, num_valid_i32, \"ult\")\n"
"t = fused2 & mask24_i32\n"
"s = fused2 >> 24\n"
"t_ok = arith.cmpu(t, tokens_i32, \"ult\")\n"
"s_ok = arith.cmpu(s, topk_i32_v, \"ult\")\n"
"row_valid = arith.andi(row_valid0, arith.andi(t_ok, s_ok))\n"
"if bool(accumulate):\n"
" t_safe = arith.select(row_valid, t, arith.constant(0, type=i32))\n"
" idx0 = t_safe * arith.constant(model_dim, type=i32)\n"
" return (idx0, row_valid)\n"
"t_safe = arith.select(row_valid, t, arith.constant(0, type=i32))\n"
"s_safe = arith.select(row_valid, s, arith.constant(0, type=i32))\n"
"t_idx = arith.index_cast(ir.IndexType.get(), t_safe)\n"
"s_idx = arith.index_cast(ir.IndexType.get(), s_safe)\n"
"n_byte_stride = arith.constant(model_dim * out_elem_bytes, index=True)\n"
"row_byte_base = out_base_idx + (t_idx * arith.constant(topk, index=True) + s_idx) * n_byte_stride\n"
"return (row_byte_base, row_valid)\n"
).body
self.precompute_fixed = True
return node
if node.name == "store_pair" and not self.store_pair_fixed:
node.body = _ast.parse(store_pair_body).body
self.store_pair_fixed = True
return node
return node
fixer = _RowCtxFixer()
module_ast = fixer.visit(module_ast)
_ast.fix_missing_locations(module_ast)
return _ast.unparse(module_ast), fixer.precompute_fixed, fixer.store_pair_fixed
def _rewrite_k_main2_loop(text: str) -> tuple[str, bool]:
module_ast = _ast.parse(text)
class _LoopFixer(_ast.NodeTransformer):
def __init__(self):
self.fixed = False
def visit_For(self, node):
node = self.generic_visit(node)
if self.fixed:
return node
if not isinstance(node, _ast.For):
return node
if not isinstance(node.target, _ast.Name) or node.target.id != "k_iv":
return node
loop_call = node.iter
if not isinstance(loop_call, _ast.Call):
return node
if not isinstance(loop_call.func, _ast.Name) or loop_call.func.id != "range":
return node
if len(loop_call.args) != 3:
return node
stop_arg = loop_call.args[1]
if not isinstance(stop_arg, _ast.Name) or stop_arg.id != "c_k_main2":
return node
fixed_if = _ast.parse(
"if k_main2_py > 0:\n"
" for k_iv_py in range_constexpr(0, k_main2_py, tile_k * 2):\n"
" k_iv = k_iv_py\n"
" pass\n"
).body[0]
fixed_for = fixed_if.body[0]
fixed_for.body = (
[_ast.parse("k_iv = k_iv_py").body[0]] + list(node.body)
)
fixed_for.orelse = list(node.orelse)
self.fixed = True
return _ast.copy_location(fixed_if, node)
fixer = _LoopFixer()
module_ast = fixer.visit(module_ast)
_ast.fix_missing_locations(module_ast)
return _ast.unparse(module_ast), fixer.fixed
try:
patched_source = source
patched_source = _replace_once(
patched_source,
"module_name_tag",
' f"_vscale_fix3"',
f' f"{patch_tag}"',
)
patched_source = _replace_once(
patched_source,
"lds_tid_bytes",
" lds_total_bytes = max(lds_x_bytes, lds_out_bytes)\n",
" lds_tid_bytes = int(tile_m) * 4\n"
" lds_total_bytes = max(lds_x_bytes, lds_out_bytes) + lds_tid_bytes\n",
)
patched_source = _replace_once(
patched_source,
"lds_tid_alias",
" # Buffer resources.\n",
" # lds_tid: alias LDS after max(x, out) for sorted_idx preload\n"
" _lds_x_b = 2 * int(tile_m) * int(lds_stride) * int(a_elem_bytes)\n"
" _lds_out_b = 2 * int(tile_m) * int(tile_n) if _use_cshuffle_epilog else 0\n"
" _lds_tid_off = max(_lds_x_b, _lds_out_b)\n"
" lds_tid = SmemPtr(\n"
" base_ptr, lds_x_ptr.byte_offset + _lds_tid_off, i32, shape=(tile_m,)\n"
" ).get()\n\n"
" # Buffer resources.\n",
)
patched_source = _replace_once(
patched_source,
"sortedidx_prologue",
" store_x_tile_to_lds(x_regs0, lds_base_cur)\n"
" gpu.barrier()\n",
" store_x_tile_to_lds(x_regs0, lds_base_cur)\n"
" # Preload sorted_idx into lds_tid for epilogue precompute_row.\n"
" if int(tile_m) % 4 == 0:\n"
" _c_tile_m4_i32 = arith.constant(tile_m // 4, type=i32)\n"
" _tid4_tx_i32 = arith.index_cast(i32, tx)\n"
" _tid4_in_range = arith.cmpu(_tid4_tx_i32, _c_tile_m4_i32, \"ult\")\n"
" _if_tid4 = scf.IfOp(_tid4_in_range)\n"
" with ir.InsertionPoint(_if_tid4.then_block):\n"
" _tid4_mul_i32 = _tid4_tx_i32 * arith.constant(4, type=i32)\n"
" _tid4_col = arith.index_cast(ir.IndexType.get(), _tid4_mul_i32)\n"
" _tid4_row = bx_m + _tid4_col\n"
" _tid4_val = buffer_ops.buffer_load(\n"
" sorted_rsrc, _tid4_row, vec_width=4, dtype=i32\n"
" )\n"
" vector.store(_tid4_val, lds_tid, [_tid4_col])\n"
" scf.YieldOp([])\n"
" else:\n"
" _c_tile_m_i32 = arith.constant(tile_m, type=i32)\n"
" _tid_tx_i32 = arith.index_cast(i32, tx)\n"
" _tid_in_range = arith.cmpu(_tid_tx_i32, _c_tile_m_i32, \"ult\")\n"
" _if_tid = scf.IfOp(_tid_in_range)\n"
" with ir.InsertionPoint(_if_tid.then_block):\n"
" _tid_row = bx_m + tx\n"
" _tid_val = buffer_ops.buffer_load(\n"
" sorted_rsrc, _tid_row, vec_width=1, dtype=i32\n"
" )\n"
" memref.store(_tid_val, lds_tid, [tx])\n"
" scf.YieldOp([])\n"
" gpu.barrier()\n",
)
patched_source, rowctx_precompute_fixed, rowctx_store_pair_fixed = (
_rewrite_rowctx_contract(patched_source)
)
if not rowctx_precompute_fixed:
raise RuntimeError("rowctx_ast_precompute")
if not rowctx_store_pair_fixed:
raise RuntimeError("rowctx_ast_store_pair")
patched_source, k_loop_fixed = _rewrite_k_main2_loop(patched_source)
if not k_loop_fixed:
raise RuntimeError("k_main2_loop")
except RuntimeError as exc:
_log_once(
("flydsl-stage2-sortedidx-lds-patch-missing-pattern", digest, str(exc)),
"[submission] kmoe patch flydsl stage2 sortedidx lds skipped "
f"reason=missing_pattern sha1={digest} label={exc}",
)
return
if patched_source == source:
_log_once(
("flydsl-stage2-sortedidx-lds-patch-noop", digest),
"[submission] kmoe patch flydsl stage2 sortedidx lds skipped "
f"reason=noop sha1={digest}",
)
return
new_digest = _hashlib.sha1(
"\n".join(line.rstrip() for line in patched_source.splitlines()).strip().encode(
"utf-8"
)
).hexdigest()[:16]
module_globals = vars(_mixed_stage2)
try:
exec(compile(patched_source, source_file, "exec"), module_globals, module_globals)
except Exception as exc:
_log_once(
("flydsl-stage2-sortedidx-lds-patch-exec-failure", digest, type(exc).__name__),
"[submission] kmoe patch flydsl stage2 sortedidx lds skipped "
f"reason=exec_failure sha1={digest} exc={type(exc).__name__}: {exc}",
)
return
patched_fn = module_globals.get("compile_mixed_moe_gemm2")
if patched_fn is None:
_log_once(
("flydsl-stage2-sortedidx-lds-patch-missing-rebound", digest),
"[submission] kmoe patch flydsl stage2 sortedidx lds skipped "
f"reason=missing_rebound sha1={digest}",
)
return
_mixed_stage2.compile_mixed_moe_gemm2 = patched_fn
_KMOE_FLYDSL_STAGE2_PATCHED_FN_BY_MODE[patch_mode] = patched_fn
try:
_flydsl_moe_kernels._get_compiled_stage2.cache_clear()
except Exception:
pass
_FLYDSL_STAGE2_SORTEDIDX_LDS_PATCHED = True
_KMOE_ACTIVE_FLYDSL_STAGE2_PATCH_MODE = patch_mode
_log_once(
("flydsl-stage2-sortedidx-lds-patch-active", digest, new_digest, patch_mode),
"[submission] kmoe patch flydsl stage2 sortedidx lds active "
f"old_sha1={digest} new_sha1={new_digest} mode={mode_label} "
f"k_loop_fix={int(k_loop_fixed)} preload=dwordx4 rowctx={rowctx_label} "
f"patch_mode={patch_mode}",
)
_ORIG_GET_BLOCK_SIZE_M = _aiter_fused_moe.get_block_size_M
_ORIG_GET_2STAGE_CFGS = _aiter_fused_moe.get_2stage_cfgs
_FUSED_MOE_GLOBALS = fused_moe.__globals__
_ORIG_FLYDSL_STAGE2_WRAPPER = _aiter_fused_moe._flydsl_stage2_wrapper
def _patched_flydsl_stage2_wrapper(
inter_states,
w1,
w2,
sorted_token_ids,
sorted_expert_ids,
num_valid_ids,
out,
topk,
kernelName="",
w2_scale=None,
a2_scale=None,
sorted_weights=None,
**kwargs,
):
patch_mode = _KMOE_PENDING_FLYDSL_STAGE2_PATCH_MODE
_patch_flydsl_stage2_sortedidx_lds_prefetch(patch_mode)
_log_once(
("kmoe-flydsl-stage2-wrapper-patch-mode", kernelName, patch_mode),
"[submission] kmoe flydsl stage2 wrapper patch mode "
f"kernel='{kernelName}' patch_mode={patch_mode}",
)
return _ORIG_FLYDSL_STAGE2_WRAPPER(
inter_states,
w1,
w2,
sorted_token_ids,
sorted_expert_ids,
num_valid_ids,
out,
topk,
kernelName=kernelName,
w2_scale=w2_scale,
a2_scale=a2_scale,
sorted_weights=sorted_weights,
**kwargs,
)
_aiter_fused_moe._flydsl_stage2_wrapper = _patched_flydsl_stage2_wrapper
_FUSED_MOE_GLOBALS["_flydsl_stage2_wrapper"] = _patched_flydsl_stage2_wrapper
_BLOCK_M_OVERRIDES = {
(512, 9, 33, 512): 128,
(512, 9, 33, 2048): 64,
}
@functools.lru_cache(maxsize=2048)
def _patched_get_block_size_M(token, topk, expert, inter_dim):
return _BLOCK_M_OVERRIDES.get(
(token, topk, expert, inter_dim),
_ORIG_GET_BLOCK_SIZE_M(token, topk, expert, inter_dim),
)
_aiter_fused_moe.get_block_size_M = _patched_get_block_size_M
_FUSED_MOE_GLOBALS["get_block_size_M"] = _patched_get_block_size_M
@functools.lru_cache(maxsize=2048)
def _patched_get_2stage_cfgs(
token,
model_dim,
inter_dim,
expert,
topk,
dtype,
q_dtype_a,
q_dtype_w,
q_type,
use_g1u1,
activation,
doweight_stage1,
hidden_pad,
intermediate_pad,
is_shuffled=True,
):
metadata = _ORIG_GET_2STAGE_CFGS(
token,
model_dim,
inter_dim,
expert,
topk,
dtype,
q_dtype_a,
q_dtype_w,
q_type,
use_g1u1,
activation,
doweight_stage1,
hidden_pad,
intermediate_pad,
is_shuffled,
)
override = _FUSED_CFG_STAGE_OVERRIDES.get((token, model_dim, inter_dim, expert, topk))
if override is None:
return metadata
if (
activation != ActivationType.Silu
or dtype != dtypes.bf16
or q_dtype_a != dtypes.fp4x2
or q_dtype_w != dtypes.fp4x2
or q_type != QuantType.per_1x32
or not use_g1u1
or doweight_stage1
or not is_shuffled
or metadata.run_1stage
):
return metadata
shape_key = (token, model_dim, inter_dim, expert, topk)
stage1_kernel, stage2_kernel, block_m = override
flydsl_available = bool(
hasattr(_aiter_fused_moe, "is_flydsl_available")
and _aiter_fused_moe.is_flydsl_available()
)
resolved_stage2_kernel = stage2_kernel
if stage2_kernel.startswith("flydsl_") and not flydsl_available:
resolved_stage2_kernel = _FLYDSL_STAGE2_FALLBACK_KERNELS.get(
shape_key,
stage2_kernel,
)
print(
"[submission] fused cfg flydsl stage2 unavailable, "
f"fallback to CK for shape={shape_key} kernel2='{resolved_stage2_kernel}'",
file=sys.stderr,
)
print(
f"[submission] fused cfg override active token={token} model_dim={model_dim} "
f"inter_dim={inter_dim} expert={expert} topk={topk} block_m={block_m} "
f"kernel1='{stage1_kernel}' kernel2='{resolved_stage2_kernel}'",
file=sys.stderr,
)
stage1 = metadata.stage1
if isinstance(stage1, functools.partial):
stage1_kwargs = dict(stage1.keywords or {})
stage1_kwargs["kernelName"] = stage1_kernel
stage1 = functools.partial(stage1.func, *stage1.args, **stage1_kwargs)
stage2 = metadata.stage2
if resolved_stage2_kernel.startswith("flydsl_") and flydsl_available:
stage2 = functools.partial(
_aiter_fused_moe._flydsl_stage2_wrapper,
kernelName=resolved_stage2_kernel,
)
print(
"[submission] fused cfg flydsl stage2 active "
f"shape={shape_key} kernel2='{resolved_stage2_kernel}'",
file=sys.stderr,
)
elif isinstance(stage2, functools.partial):
stage2_kwargs = dict(stage2.keywords or {})
stage2_kwargs["kernelName"] = resolved_stage2_kernel
stage2 = functools.partial(stage2.func, *stage2.args, **stage2_kwargs)
return _aiter_fused_moe.MOEMetadata(
stage1=stage1,
stage2=stage2,
block_m=block_m,
ksplit=metadata.ksplit,
run_1stage=metadata.run_1stage,
has_bias=metadata.has_bias,
use_non_temporal_load=metadata.use_non_temporal_load,
)
_aiter_fused_moe.get_2stage_cfgs = _patched_get_2stage_cfgs
_FUSED_MOE_GLOBALS["get_2stage_cfgs"] = _patched_get_2stage_cfgs
scrolls · 3424 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Best evidence level for this revision: reported
JSON