submission 589135
Aniket Sadashiva · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 425 lines, June 9 Researcher Reciprocity License v1.0.
submission_v311.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-589135?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:ec18a7493209dbdcc53de9c706978566719203dcfe7055398a08480ebab91490
license declaredunknown
license concludedunknown
authorsAniket Sadashiva
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
Kernel source
submission_v311.py425 lines
"""v311: v252 + metadata-backed native BF16/FP4 E=33 bs=512 path.
Key insight: per-shape OPUS switching (v249) may cause correctness issues from
cache_clear() interactions. Instead, set OPUS=1 globally before import.
OPUS sorting helps S6 (~181µs vs ~210µs) and is harmless for CKTile shapes.
- dsv3 CSV ksplit=2 for E=257 bs<=128
- CKTile split_k=2 for S1/S2/S4/S5 via get_2stage_cfgs patch
- E=33 CSV: S6 block_m=32, S7 block_m=128
- OPUS=1 globally (no switching)
- NT=1 globally (helps S3/S6/S7, harmless for CKTile shapes)
- block_m=64+NT for S7 via patch (overrides CSV)
- Cached sorting buffers
- Direct native BF16/FP4 2-stage metadata path for E=33 bs=512
"""
import functools
import os
import sys
import torch
from dataclasses import replace
# ── Step 1: Modify dsv3 CSV for E=257 ksplit=2 on bs<=128 ──
_dsv3_path = "/home/runner/aiter/aiter/configs/model_configs/dsv3_fp4_tuned_fmoe.csv"
try:
with open(_dsv3_path, "r") as f:
lines = f.readlines()
header = lines[0].strip()
modified_lines = [header + "\n"]
for line in lines[1:]:
stripped = line.strip()
if not stripped:
continue
fields = stripped.split(",")
try:
token_val = int(fields[1])
expert_val = int(fields[4])
except (ValueError, IndexError):
modified_lines.append(line)
continue
if expert_val == 257 and token_val <= 128:
fields[14] = "2" # ksplit=2
modified_lines.append(",".join(fields) + "\n")
else:
modified_lines.append(line)
with open(_dsv3_path, "w") as f:
f.writelines(modified_lines)
except Exception as e:
print(f"[v252] dsv3 error: {e}", file=sys.stderr)
# ── Step 2: E=33 CSV for S4/S5 ksplit=2, S6/S7 tuned configs ──
_csv_header = (
"cu_num,token,model_dim,inter_dim,expert,topk,act_type,dtype,q_dtype_a,"
"q_dtype_w,q_type,use_g1u1,doweight_stage1,block_m,ksplit,us1,kernelName1,"
"err1,us2,kernelName2,err2,us,run_1stage,tflops,bw,_tag"
)
_common = (
"ActivationType.Silu,torch.bfloat16,torch.float4_e2m1fn_x2,"
"torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0"
)
_k1_small = (
"moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_"
"Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
_k2_small = (
"moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_"
"Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
)
_k1_512 = (
"moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_"
"Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
_k2_512 = (
"moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_"
"Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
)
_k1_2048 = (
"moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_"
"Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
_k2_2048 = (
"moe_ck2stages_gemm2_256x128x128x128_1x4_MulABScaleExpertWeightShuffled_v3_"
"Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
)
_csv_rows = [
# S4 (bs=16, d=512): ksplit=2
f"256,16,7168,512,33,9,{_common},32,2,0,{_k1_small},0.0%,0,{_k2_small},0.0%,47,0,67,7691,",
# S5 (bs=128, d=512): ksplit=2
f"256,128,7168,512,33,9,{_common},32,2,0,{_k1_small},0.0%,0,{_k2_small},0.0%,58,0,434,6263,",
# S6 (bs=512, d=512): block_m=32, ksplit=0
f"256,512,7168,512,33,9,{_common},32,0,0,{_k1_512},0.0%,0,{_k2_512},0.0%,129.79,0,781.78,2884.18,",
# S7 (bs=512, d=2048): block_m=128 (default kernel), ksplit=0
f"256,512,7168,2048,33,9,{_common},128,0,0,{_k1_2048},0.0%,0,{_k2_2048},0.0%,275.08,0,1475.47,5323.27,",
]
try:
e33_path = "/home/runner/aiter/aiter/configs/model_configs/e33_fp4_tuned_fmoe.csv"
with open(e33_path, "w") as f:
f.write(_csv_header + "\n")
for row in _csv_rows:
f.write(row + "\n")
except Exception:
pass
# ── Step 3: Global OPUS=1, NT=1 (no per-shape switching!) ──
os.environ["AITER_USE_OPUS_MOE_SORTING"] = "1"
os.environ["AITER_USE_NT"] = "1"
from task import input_t, output_t
from aiter import ActivationType, QuantType, dtypes
import aiter
import aiter.fused_moe as fused_moe_mod
# ═══════════════════════════════════════════════════════════════════════
# Cached sorting buffers (zero alloc after warmup)
# ═══════════════════════════════════════════════════════════════════════
_SORT_BUFS = {}
_DIRECT_E33_CK_BUFS = {}
_DIRECT_E33_CK_DISABLED = False
def _cached_moe_sorting_impl(
topk_ids, topk_weights, num_experts, model_dim, moebuf_dtype,
block_size, expert_mask, num_local_tokens, dispatch_policy, use_opus,
):
device = topk_ids.device
M, topk = topk_ids.shape
key = (M, num_experts, block_size, model_dim)
if key not in _SORT_BUFS:
max_num_tokens_padded = int(M * topk + num_experts * block_size - topk)
max_num_m_blocks = int(
(max_num_tokens_padded + block_size - 1) // block_size
)
_SORT_BUFS[key] = (
torch.empty(max_num_tokens_padded, dtype=dtypes.i32, device=device),
torch.empty(max_num_tokens_padded, dtype=dtypes.fp32, device=device),
torch.empty(max_num_m_blocks, dtype=dtypes.i32, device=device),
torch.empty(2, dtype=dtypes.i32, device=device),
torch.empty((M, model_dim), dtype=moebuf_dtype, device=device),
)
sid, sw, sei, nvi, mb = _SORT_BUFS[key]
fwd = aiter.moe_sorting_opus_fwd if use_opus else aiter.moe_sorting_fwd
fwd(
topk_ids, topk_weights, sid, sw, sei, nvi, mb,
num_experts, int(block_size), expert_mask, num_local_tokens,
dispatch_policy,
)
return sid, sw, sei, nvi, mb
fused_moe_mod._moe_sorting_impl = _cached_moe_sorting_impl
def _get_direct_e33_ck_bufs(hidden_states, config, topk, block_m):
token_num = hidden_states.shape[0]
num_experts = config["n_routed_experts"] + config["n_shared_experts"]
model_dim = config["d_hidden"]
inter_dim = config["d_expert"]
key = (
hidden_states.device,
hidden_states.dtype,
token_num,
topk,
num_experts,
model_dim,
inter_dim,
block_m,
)
bufs = _DIRECT_E33_CK_BUFS.get(key)
if bufs is None:
bufs = {
"inter_buf": torch.empty(
(token_num, topk, inter_dim),
dtype=hidden_states.dtype,
device=hidden_states.device,
),
"out": torch.empty(
(token_num, model_dim),
dtype=hidden_states.dtype,
device=hidden_states.device,
),
}
_DIRECT_E33_CK_BUFS[key] = bufs
return bufs
def _direct_e33_bs512_ck(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
config,
):
token_num = hidden_states.shape[0]
topk = topk_ids.shape[1]
model_dim = config["d_hidden"]
inter_dim = config["d_expert"]
hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
intermediate_pad = config["d_expert_pad"] - config["d_expert"]
total_experts = config["n_routed_experts"] + config["n_shared_experts"]
policy_metadata = fused_moe_mod.get_2stage_cfgs(
token_num,
model_dim,
inter_dim,
total_experts,
topk,
hidden_states.dtype,
dtypes.fp4x2,
dtypes.fp4x2,
QuantType.per_1x32,
True,
ActivationType.Silu,
False,
hidden_pad,
intermediate_pad,
True,
)
native_metadata = fused_moe_mod.get_2stage_cfgs(
token_num,
model_dim,
inter_dim,
total_experts,
topk,
hidden_states.dtype,
hidden_states.dtype,
dtypes.fp4x2,
QuantType.per_1x32,
True,
ActivationType.Silu,
False,
hidden_pad,
intermediate_pad,
True,
)
if native_metadata.run_1stage:
raise RuntimeError("unexpected 1-stage metadata for direct BF16/FP4 path")
block_m = policy_metadata.block_m
bufs = _get_direct_e33_ck_bufs(hidden_states, config, topk, block_m)
sid, sw, sei, nvi, _ = _cached_moe_sorting_impl(
topk_ids,
topk_weights,
total_experts,
model_dim,
hidden_states.dtype,
block_m,
None,
None,
0,
True,
)
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
)
w2_scale = (
down_weight_scale_shuffled.view(dtypes.fp8_e8m0)
if down_weight_shuffled.dtype == dtypes.fp4x2
else down_weight_scale_shuffled
)
bufs["inter_buf"].zero_()
bufs["out"].zero_()
native_metadata.stage1(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
sid,
sei,
nvi,
bufs["inter_buf"],
topk,
block_m=block_m,
a1_scale=None,
w1_scale=w1_scale,
sorted_weights=None,
)
native_metadata.stage2(
bufs["inter_buf"],
gate_up_weight_shuffled,
down_weight_shuffled,
sid,
sei,
nvi,
bufs["out"],
topk,
block_m=block_m,
a2_scale=None,
w2_scale=w2_scale,
sorted_weights=sw,
)
return bufs["out"]
# ═══════════════════════════════════════════════════════════════════════
# Patch get_2stage_cfgs: CKTile for S1/S2/S4/S5, tuned S7
# ═══════════════════════════════════════════════════════════════════════
_ORIGINAL_GET_2STAGE_CFGS = fused_moe_mod.get_2stage_cfgs
def _make_cktile_metadata(hidden_pad, intermediate_pad, use_g1u1, activation, split_k=2):
return fused_moe_mod.MOEMetadata(
functools.partial(
fused_moe_mod.cktile_moe_stage1,
n_pad_zeros=intermediate_pad // 64 * 64 * (2 if use_g1u1 else 1),
k_pad_zeros=hidden_pad // 128 * 128,
activation=activation,
split_k=split_k,
),
functools.partial(
fused_moe_mod.cktile_moe_stage2,
n_pad_zeros=hidden_pad // 64 * 64,
k_pad_zeros=intermediate_pad // 128 * 128,
activation=activation,
),
16, split_k, False, False, True,
)
@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,
):
common = (
model_dim == 7168 and topk == 9
and dtype == dtypes.bf16 and q_dtype_a == dtypes.fp4x2
and q_dtype_w == dtypes.fp4x2 and q_type == QuantType.per_1x32
and use_g1u1 and not doweight_stage1 and is_shuffled
)
if not common:
return _ORIGINAL_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,
)
# E=257 bs<512 → CKTile split_k=2 (quant-skip)
if expert == 257 and inter_dim == 256 and token < 512:
return _make_cktile_metadata(
hidden_pad, intermediate_pad, use_g1u1, activation, split_k=2,
)
# E=33 d=512 bs<=128 → CKTile split_k=2 (quant-skip)
if expert == 33 and inter_dim == 512 and token <= 128:
return _make_cktile_metadata(
hidden_pad, intermediate_pad, use_g1u1, activation, split_k=2,
)
default = _ORIGINAL_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,
)
# S7 (E=33, d=2048, bs=512): block_m=64 + NT
if expert == 33 and inter_dim == 2048 and token == 512:
return replace(default, block_m=64, use_non_temporal_load=True)
return default
fused_moe_mod.get_2stage_cfgs = _patched_get_2stage_cfgs
# ═══════════════════════════════════════════════════════════════════════
# Dispatch — simple, no state switching
# ═══════════════════════════════════════════════════════════════════════
def custom_kernel(data: input_t) -> output_t:
global _DIRECT_E33_CK_DISABLED
(
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
total_experts = config["n_routed_experts"] + config["n_shared_experts"]
hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
intermediate_pad = config["d_expert_pad"] - config["d_expert"]
if total_experts == 33 and config["bs"] == 512 and not _DIRECT_E33_CK_DISABLED:
try:
return _direct_e33_bs512_ck(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
config,
)
except Exception as e:
print(f"[v311] direct CK E33 path failed: {e}", file=sys.stderr)
_DIRECT_E33_CK_DISABLED = True
return fused_moe_mod.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,
)
scrolls · 425 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 573469.
- """Stable default submission.+ """v311: v252 + metadata-backed native BF16/FP4 E=33 bs=512 path.- This keeps the best known shape-specialized routing from the local experiment log:- - E=257, bs<512: force CKTile split-k path- - E=33, d=512, bs<=128: force CKTile split-k path- - E=33, large shapes: keep the tuned CSV-backed CK 2-stage kernels+ Key insight: per-shape OPUS switching (v249) may cause correctness issues from+ cache_clear() interactions. Instead, set OPUS=1 globally before import.+ OPUS sorting helps S6 (~181µs vs ~210µs) and is harmless for CKTile shapes.++ - dsv3 CSV ksplit=2 for E=257 bs<=128+ - CKTile split_k=2 for S1/S2/S4/S5 via get_2stage_cfgs patch+ - E=33 CSV: S6 block_m=32, S7 block_m=128+ - OPUS=1 globally (no switching)+ - NT=1 globally (helps S3/S6/S7, harmless for CKTile shapes)+ - block_m=64+NT for S7 via patch (overrides CSV)+ - Cached sorting buffers+ - Direct native BF16/FP4 2-stage metadata path for E=33 bs=512"""import functoolsimport os+ import sys- from task import input_t, output_t+ import torch- _CSV_PATH = "/home/runner/aiter/aiter/configs/model_configs/e33_fp4_tuned_fmoe.csv"- _CSV_HEADER = (+ from dataclasses import replace++ # ── Step 1: Modify dsv3 CSV for E=257 ksplit=2 on bs<=128 ──+ _dsv3_path = "/home/runner/aiter/aiter/configs/model_configs/dsv3_fp4_tuned_fmoe.csv"+ try:+ with open(_dsv3_path, "r") as f:+ lines = f.readlines()+ header = lines[0].strip()+ modified_lines = [header + "\n"]+ for line in lines[1:]:+ stripped = line.strip()+ if not stripped:+ continue+ fields = stripped.split(",")+ try:+ token_val = int(fields[1])+ expert_val = int(fields[4])+ except (ValueError, IndexError):+ modified_lines.append(line)+ continue+ if expert_val == 257 and token_val <= 128:+ fields[14] = "2" # ksplit=2+ modified_lines.append(",".join(fields) + "\n")+ else:+ modified_lines.append(line)+ with open(_dsv3_path, "w") as f:+ f.writelines(modified_lines)+ except Exception as e:+ print(f"[v252] dsv3 error: {e}", file=sys.stderr)++ # ── Step 2: E=33 CSV for S4/S5 ksplit=2, S6/S7 tuned configs ──+ _csv_header = ("cu_num,token,model_dim,inter_dim,expert,topk,act_type,dtype,q_dtype_a,""q_dtype_w,q_type,use_g1u1,doweight_stage1,block_m,ksplit,us1,kernelName1,""err1,us2,kernelName2,err2,us,run_1stage,tflops,bw,_tag")- _COMMON = (+ _common = ("ActivationType.Silu,torch.bfloat16,torch.float4_e2m1fn_x2,""torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0")- _K1_512 = (+ _k1_small = (+ "moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_"+ "Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"+ )+ _k2_small = (+ "moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_"+ "Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"+ )+ _k1_512 = ("moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_""Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16")- _K2_512 = (+ _k2_512 = ("moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_""Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16")- _K1_2048 = (+ _k1_2048 = ("moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_""Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16")- _K2_2048 = (+ _k2_2048 = ("moe_ck2stages_gemm2_256x128x128x128_1x4_MulABScaleExpertWeightShuffled_v3_""Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16")- _CSV_ROWS = [- f"256,512,7168,512,33,9,{_COMMON},32,0,0,{_K1_512},0.0%,0,{_K2_512},0.0%,129.79,0,781.78,2884.18,",- f"256,512,7168,2048,33,9,{_COMMON},128,0,0,{_K1_2048},0.0%,0,{_K2_2048},0.0%,275.08,0,1475.47,5323.27,",+ _csv_rows = [+ # S4 (bs=16, d=512): ksplit=2+ f"256,16,7168,512,33,9,{_common},32,2,0,{_k1_small},0.0%,0,{_k2_small},0.0%,47,0,67,7691,",+ # S5 (bs=128, d=512): ksplit=2+ f"256,128,7168,512,33,9,{_common},32,2,0,{_k1_small},0.0%,0,{_k2_small},0.0%,58,0,434,6263,",+ # S6 (bs=512, d=512): block_m=32, ksplit=0+ f"256,512,7168,512,33,9,{_common},32,0,0,{_k1_512},0.0%,0,{_k2_512},0.0%,129.79,0,781.78,2884.18,",+ # S7 (bs=512, d=2048): block_m=128 (default kernel), ksplit=0+ f"256,512,7168,2048,33,9,{_common},128,0,0,{_k1_2048},0.0%,0,{_k2_2048},0.0%,275.08,0,1475.47,5323.27,",]-try:- with open(_CSV_PATH, "w") as f:- f.write(_CSV_HEADER + "\n")- for row in _CSV_ROWS:+ e33_path = "/home/runner/aiter/aiter/configs/model_configs/e33_fp4_tuned_fmoe.csv"+ with open(e33_path, "w") as f:+ f.write(_csv_header + "\n")+ for row in _csv_rows:f.write(row + "\n")except Exception:pass- os.environ["AITER_USE_OPUS_MOE_SORTING"] = "0"- os.environ["AITER_USE_NT"] = "0"+ # ── Step 3: Global OPUS=1, NT=1 (no per-shape switching!) ──+ os.environ["AITER_USE_OPUS_MOE_SORTING"] = "1"+ os.environ["AITER_USE_NT"] = "1"+ from task import input_t, output_tfrom aiter import ActivationType, QuantType, dtypes+ import aiterimport aiter.fused_moe as fused_moe_mod- _LAST_RUNTIME_STATE = None++ # ═══════════════════════════════════════════════════════════════════════+ # Cached sorting buffers (zero alloc after warmup)+ # ═══════════════════════════════════════════════════════════════════════+ _SORT_BUFS = {}+ _DIRECT_E33_CK_BUFS = {}+ _DIRECT_E33_CK_DISABLED = False+++ def _cached_moe_sorting_impl(+ topk_ids, topk_weights, num_experts, model_dim, moebuf_dtype,+ block_size, expert_mask, num_local_tokens, dispatch_policy, use_opus,+ ):+ device = topk_ids.device+ M, topk = topk_ids.shape+ key = (M, num_experts, block_size, model_dim)++ if key not in _SORT_BUFS:+ max_num_tokens_padded = int(M * topk + num_experts * block_size - topk)+ max_num_m_blocks = int(+ (max_num_tokens_padded + block_size - 1) // block_size+ )+ _SORT_BUFS[key] = (+ torch.empty(max_num_tokens_padded, dtype=dtypes.i32, device=device),+ torch.empty(max_num_tokens_padded, dtype=dtypes.fp32, device=device),+ torch.empty(max_num_m_blocks, dtype=dtypes.i32, device=device),+ torch.empty(2, dtype=dtypes.i32, device=device),+ torch.empty((M, model_dim), dtype=moebuf_dtype, device=device),+ )++ sid, sw, sei, nvi, mb = _SORT_BUFS[key]+ fwd = aiter.moe_sorting_opus_fwd if use_opus else aiter.moe_sorting_fwd+ fwd(+ topk_ids, topk_weights, sid, sw, sei, nvi, mb,+ num_experts, int(block_size), expert_mask, num_local_tokens,+ dispatch_policy,+ )+ return sid, sw, sei, nvi, mb+++ fused_moe_mod._moe_sorting_impl = _cached_moe_sorting_impl+++ def _get_direct_e33_ck_bufs(hidden_states, config, topk, block_m):+ token_num = hidden_states.shape[0]+ num_experts = config["n_routed_experts"] + config["n_shared_experts"]+ model_dim = config["d_hidden"]+ inter_dim = config["d_expert"]++ key = (+ hidden_states.device,+ hidden_states.dtype,+ token_num,+ topk,+ num_experts,+ model_dim,+ inter_dim,+ block_m,+ )+ bufs = _DIRECT_E33_CK_BUFS.get(key)+ if bufs is None:+ bufs = {+ "inter_buf": torch.empty(+ (token_num, topk, inter_dim),+ dtype=hidden_states.dtype,+ device=hidden_states.device,+ ),+ "out": torch.empty(+ (token_num, model_dim),+ dtype=hidden_states.dtype,+ device=hidden_states.device,+ ),+ }+ _DIRECT_E33_CK_BUFS[key] = bufs+ return bufs+++ def _direct_e33_bs512_ck(+ hidden_states,+ gate_up_weight_shuffled,+ down_weight_shuffled,+ gate_up_weight_scale_shuffled,+ down_weight_scale_shuffled,+ topk_weights,+ topk_ids,+ config,+ ):+ token_num = hidden_states.shape[0]+ topk = topk_ids.shape[1]+ model_dim = config["d_hidden"]+ inter_dim = config["d_expert"]+ hidden_pad = config["d_hidden_pad"] - config["d_hidden"]+ intermediate_pad = config["d_expert_pad"] - config["d_expert"]+ total_experts = config["n_routed_experts"] + config["n_shared_experts"]++ policy_metadata = fused_moe_mod.get_2stage_cfgs(+ token_num,+ model_dim,+ inter_dim,+ total_experts,+ topk,+ hidden_states.dtype,+ dtypes.fp4x2,+ dtypes.fp4x2,+ QuantType.per_1x32,+ True,+ ActivationType.Silu,+ False,+ hidden_pad,+ intermediate_pad,+ True,+ )+ native_metadata = fused_moe_mod.get_2stage_cfgs(+ token_num,+ model_dim,+ inter_dim,+ total_experts,+ topk,+ hidden_states.dtype,+ hidden_states.dtype,+ dtypes.fp4x2,+ QuantType.per_1x32,+ True,+ ActivationType.Silu,+ False,+ hidden_pad,+ intermediate_pad,+ True,+ )+ if native_metadata.run_1stage:+ raise RuntimeError("unexpected 1-stage metadata for direct BF16/FP4 path")+ block_m = policy_metadata.block_m+ bufs = _get_direct_e33_ck_bufs(hidden_states, config, topk, block_m)++ sid, sw, sei, nvi, _ = _cached_moe_sorting_impl(+ topk_ids,+ topk_weights,+ total_experts,+ model_dim,+ hidden_states.dtype,+ block_m,+ None,+ None,+ 0,+ True,+ )++ 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+ )+ w2_scale = (+ down_weight_scale_shuffled.view(dtypes.fp8_e8m0)+ if down_weight_shuffled.dtype == dtypes.fp4x2+ else down_weight_scale_shuffled+ )++ bufs["inter_buf"].zero_()+ bufs["out"].zero_()++ native_metadata.stage1(+ hidden_states,+ gate_up_weight_shuffled,+ down_weight_shuffled,+ sid,+ sei,+ nvi,+ bufs["inter_buf"],+ topk,+ block_m=block_m,+ a1_scale=None,+ w1_scale=w1_scale,+ sorted_weights=None,+ )++ native_metadata.stage2(+ bufs["inter_buf"],+ gate_up_weight_shuffled,+ down_weight_shuffled,+ sid,+ sei,+ nvi,+ bufs["out"],+ topk,+ block_m=block_m,+ a2_scale=None,+ w2_scale=w2_scale,+ sorted_weights=sw,+ )++ return bufs["out"]+++ # ═══════════════════════════════════════════════════════════════════════+ # Patch get_2stage_cfgs: CKTile for S1/S2/S4/S5, tuned S7+ # ═══════════════════════════════════════════════════════════════════════_ORIGINAL_GET_2STAGE_CFGS = fused_moe_mod.get_2stage_cfgs⋯ 12 unchanged linesk_pad_zeros=intermediate_pad // 128 * 128,activation=activation,),- 16,- split_k,- False,- False,- True,+ 16, split_k, False, False, True,)@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,+ 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,):- if (- expert == 257- and token < 512- and model_dim == 7168- and inter_dim == 256- and topk == 9- and dtype == dtypes.bf16- and q_dtype_a == dtypes.fp4x2- and q_dtype_w == dtypes.fp4x2- and q_type == QuantType.per_1x32- and use_g1u1- and not doweight_stage1- and is_shuffled- ):- return _make_cktile_metadata(hidden_pad, intermediate_pad, use_g1u1, activation, split_k=2)-- if (- expert == 33- and token <= 128- and model_dim == 7168- and inter_dim == 512- and topk == 9- and dtype == dtypes.bf16- and q_dtype_a == dtypes.fp4x2- and q_dtype_w == dtypes.fp4x2- and q_type == QuantType.per_1x32- and use_g1u1- and not doweight_stage1- and is_shuffled- ):- return _make_cktile_metadata(hidden_pad, intermediate_pad, use_g1u1, activation, split_k=2)-- return _ORIGINAL_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,+ common = (+ model_dim == 7168 and topk == 9+ and dtype == dtypes.bf16 and q_dtype_a == dtypes.fp4x2+ and q_dtype_w == dtypes.fp4x2 and q_type == QuantType.per_1x32+ and use_g1u1 and not doweight_stage1 and is_shuffled)+ if not common:+ return _ORIGINAL_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,+ )- fused_moe_mod.get_2stage_cfgs = _patched_get_2stage_cfgs+ # E=257 bs<512 → CKTile split_k=2 (quant-skip)+ if expert == 257 and inter_dim == 256 and token < 512:+ return _make_cktile_metadata(+ hidden_pad, intermediate_pad, use_g1u1, activation, split_k=2,+ )+ # E=33 d=512 bs<=128 → CKTile split_k=2 (quant-skip)+ if expert == 33 and inter_dim == 512 and token <= 128:+ return _make_cktile_metadata(+ hidden_pad, intermediate_pad, use_g1u1, activation, split_k=2,+ )- def _set_runtime_state(use_opus, use_nt, ksplit=None):- global _LAST_RUNTIME_STATE+ default = _ORIGINAL_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,+ )- state = (use_opus, use_nt, ksplit)- if state == _LAST_RUNTIME_STATE:- return+ # S7 (E=33, d=2048, bs=512): block_m=64 + NT+ if expert == 33 and inter_dim == 2048 and token == 512:+ return replace(default, block_m=64, use_non_temporal_load=True)- fused_moe_mod._USE_OPUS_MOE_SORTING = use_opus- os.environ["AITER_USE_NT"] = "1" if use_nt else "0"+ return default- if ksplit in (None, 0):- os.environ.pop("AITER_KSPLIT", None)- else:- os.environ["AITER_KSPLIT"] = str(ksplit)- fused_moe_mod.use_nt.cache_clear()- fused_moe_mod.get_ksplit.cache_clear()- fused_moe_mod.get_2stage_cfgs.cache_clear()- _ORIGINAL_GET_2STAGE_CFGS.cache_clear()- _LAST_RUNTIME_STATE = state+ fused_moe_mod.get_2stage_cfgs = _patched_get_2stage_cfgs- def _policy(config):- total_experts = config["n_routed_experts"] + config["n_shared_experts"]-- if total_experts == 257:- return False, False, None-- if config["d_expert"] == 512 and config["bs"] <= 128:- return False, True, None-- return True, True, None--+ # ═══════════════════════════════════════════════════════════════════════+ # Dispatch — simple, no state switching+ # ═══════════════════════════════════════════════════════════════════════def custom_kernel(data: input_t) -> output_t:+ global _DIRECT_E33_CK_DISABLED+(- 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,+ 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- use_opus, use_nt, ksplit = _policy(config)- _set_runtime_state(use_opus=use_opus, use_nt=use_nt, ksplit=ksplit)-+ total_experts = config["n_routed_experts"] + config["n_shared_experts"]hidden_pad = config["d_hidden_pad"] - config["d_hidden"]intermediate_pad = config["d_expert_pad"] - config["d_expert"]+ if total_experts == 33 and config["bs"] == 512 and not _DIRECT_E33_CK_DISABLED:+ try:+ return _direct_e33_bs512_ck(+ hidden_states,+ gate_up_weight_shuffled,+ down_weight_shuffled,+ gate_up_weight_scale_shuffled,+ down_weight_scale_shuffled,+ topk_weights,+ topk_ids,+ config,+ )+ except Exception as e:+ print(f"[v311] direct CK E33 path failed: {e}", file=sys.stderr)+ _DIRECT_E33_CK_DISABLED = True+return fused_moe_mod.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,+ 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,+ a1_scale=None, a2_scale=None,+ hidden_pad=hidden_pad, intermediate_pad=intermediate_pad,)
scrolls · 561 diff lines total
Best evidence level for this revision: reported
JSON