submission 601016
Aniket Sadashiva · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 556 lines, June 9 Researcher Reciprocity License v1.0.
submission_v367.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-601016?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:81349b5b28ce77df78b72e963580a07438a87c92cef8d00a5f5109d5d3f3b863
license declaredunknown
license concludedunknown
authorsAniket Sadashiva
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
a_dtype="fp4",fused-epilogue
" prefetch_epilogue: bool = False,\n"split-k
- CKTile split_k=2 for S1/S2/S4/S5 via get_2stage_cfgs patchKernel source
submission_v367.py556 lines
"""v367: v363 + enable dormant FlyDSL stage1 scheduler calls for S3.
- dsv3 CSV: ksplit=2 for E=257 bs<=128, FlyDSL t64x256x256_reduce for S3
- CKTile split_k=2 for S1/S2/S4/S5 via get_2stage_cfgs patch
- E=33 CSV: S6 FlyDSL reduce, S7 FlyDSL t64x128x256_atomic
- OPUS=1 globally
- NT=1 globally
- Cached sorting buffers
- Direct `moe_cktile2stages_gemm1/2` path for S1/S2 only
"""
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, FlyDSL stage2 on bs=512 ──
_dsv3_path = "/home/runner/aiter/aiter/configs/model_configs/dsv3_fp4_tuned_fmoe.csv"
_flydsl_s3_stage2 = "flydsl_moe2_afp4_wfp4_bf16_t64x256x256_reduce"
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")
elif expert_val == 257 and token_val == 512:
# FlyDSL reduce stage2 for S3, with CK fallback
flydsl_fields = list(fields)
flydsl_fields[19] = _flydsl_s3_stage2
flydsl_fields[20] = "0.1%"
if len(flydsl_fields) > 25:
flydsl_fields[25] = ""
modified_lines.append(",".join(flydsl_fields) + "\n")
# CK fallback row
fallback_fields = list(fields)
if len(fallback_fields) > 25:
fallback_fields[25] = "flydsl_fallback"
else:
fallback_fields.append("flydsl_fallback")
modified_lines.append(",".join(fallback_fields) + "\n")
else:
modified_lines.append(line)
with open(_dsv3_path, "w") as f:
f.writelines(modified_lines)
except Exception as e:
print(f"[v329] 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"
)
_k2_512_flydsl = "flydsl_moe2_afp4_wfp4_bf16_t64x128x256_reduce"
_k2_2048_flydsl = "flydsl_moe2_afp4_wfp4_bf16_t64x128x256_atomic"
_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): probe FlyDSL reduce stage2 behind a tagged CK fallback
f"256,512,7168,512,33,9,{_common},32,0,0,{_k1_512},0.0%,90.0,{_k2_512_flydsl},0.1%,219.79,0,781.78,2884.18,",
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,flydsl_fallback",
# S7 (bs=512, d=2048): probe FlyDSL atomic stage2 behind a tagged CK fallback
f"256,512,7168,2048,33,9,{_common},64,0,0,{_k1_2048},0.0%,180.0,{_k2_2048_flydsl},0.1%,455.08,0,1475.47,5323.27,",
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,flydsl_fallback",
]
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"
# ── Step 4: Repair the known FlyDSL stage1 FP4 source bug on the runner
# and enable its dormant hot-loop scheduler calls. ──
_flydsl_stage1_bug_path = (
"/home/runner/aiter/aiter/ops/flydsl/kernels/mixed_moe_gemm_2stage.py"
)
try:
with open(_flydsl_stage1_bug_path, "r") as f:
_flydsl_src = f.read()
_old_stage1_sig = (
"def compute_f8f6f4_tile(\n"
" acc_gate_in,\n"
" acc_up_in,\n"
" b_tile_in_gate,\n"
" b_tile_in_up,\n"
" lds_base,\n"
" *,\n"
" a0_prefetch=None,\n"
" a_scale=None,\n"
" b_scale_gate=None,\n"
" b_scale_up=None,\n"
" prefetch_epilogue: bool = False,\n"
" ):"
)
_new_stage1_sig = (
"def compute_f8f6f4_tile(\n"
" acc_gate_in,\n"
" acc_up_in,\n"
" b_tile_in,\n"
" lds_base,\n"
" *,\n"
" a0_prefetch=None,\n"
" a_scale=None,\n"
" b_scale=None,\n"
" prefetch_epilogue: bool = False,\n"
" ):"
)
if _old_stage1_sig in _flydsl_src:
_flydsl_src = _flydsl_src.replace(_old_stage1_sig, _new_stage1_sig)
_stage1_sched_disabled = (
" # hot_loop_scheduler()\n"
" gpu.barrier()"
)
_stage1_sched_enabled = (
" hot_loop_scheduler()\n"
" gpu.barrier()"
)
if _stage1_sched_disabled in _flydsl_src:
_flydsl_src = _flydsl_src.replace(
_stage1_sched_disabled,
_stage1_sched_enabled,
3,
)
with open(_flydsl_stage1_bug_path, "w") as f:
f.write(_flydsl_src)
except Exception as e:
print(f"[v367] flydsl stage1 source patch skipped: {e}", file=sys.stderr)
from task import input_t, output_t
from aiter import ActivationType, QuantType, dtypes
import aiter
import aiter.fused_moe as fused_moe_mod
import aiter.ops.flydsl
from aiter.ops.flydsl.utils import is_flydsl_available as _is_flydsl_available
print(f"[v329] flydsl available: {_is_flydsl_available()}", file=sys.stderr)
# ═══════════════════════════════════════════════════════════════════════
# FlyDSL stage1 wrapper (matches CK stage1 signature)
# ═══════════════════════════════════════════════════════════════════════
def _flydsl_stage1_wrapper(
hidden_states,
w1,
w2, # unused by stage1
sorted_token_ids,
sorted_expert_ids,
num_valid_ids,
out,
topk,
block_m=32,
a1_scale=None,
w1_scale=None,
kernelName="",
sorted_weights=None,
tile_m=32,
tile_n=256,
tile_k=256,
a_dtype="fp4",
b_dtype="fp4",
out_dtype="bf16",
**_kwargs,
):
return aiter.ops.flydsl.flydsl_moe_stage1(
a=hidden_states,
w1=w1,
sorted_token_ids=sorted_token_ids,
sorted_expert_ids=sorted_expert_ids,
num_valid_ids=num_valid_ids,
out=out,
topk=topk,
tile_m=tile_m,
tile_n=tile_n,
tile_k=tile_k,
a_dtype=a_dtype,
b_dtype=b_dtype,
out_dtype=out_dtype,
w1_scale=w1_scale,
a1_scale=a1_scale,
sorted_weights=sorted_weights,
)
# ═══════════════════════════════════════════════════════════════════════
# Cached sorting buffers (zero alloc after warmup)
# ═══════════════════════════════════════════════════════════════════════
_SORT_BUFS = {}
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
# ═══════════════════════════════════════════════════════════════════════
# Direct CKTile path for official S1/S2 only
# ═══════════════════════════════════════════════════════════════════════
_DIRECT_SMALL_DISABLED = False
_CKTILE_DIRECT_BUFS = {}
def _is_direct_small_shape(config, topk):
total_experts = config["n_routed_experts"] + config["n_shared_experts"]
return (
config["d_hidden"] == 7168
and topk == 9
and total_experts == 257
and config["d_expert"] == 256
and config["bs"] in (16, 128)
)
def _get_direct_cktile_workspace(hidden_states, w1, w2, topk, total_experts, block_m):
_, n1, k1 = w1.shape
_, k2, n2 = w2.shape
model_dim = hidden_states.shape[1]
inter_dim = n2 if k2 == k1 else n2 * 2
if w1.dtype is torch.uint32:
inter_dim *= 8
token_num = hidden_states.shape[0]
max_num_tokens_padded = int(token_num * topk + total_experts * block_m - topk)
max_num_m_blocks = int((max_num_tokens_padded + block_m - 1) // block_m)
key = (
hidden_states.device,
hidden_states.dtype,
token_num,
topk,
total_experts,
model_dim,
n1,
inter_dim,
block_m,
)
workspace = _CKTILE_DIRECT_BUFS.get(key)
if workspace is None:
workspace = {
"sorted_ids": torch.empty(max_num_tokens_padded, dtype=dtypes.i32, device=hidden_states.device),
"sorted_weights": torch.empty(max_num_tokens_padded, dtype=dtypes.fp32, device=hidden_states.device),
"sorted_expert_ids": torch.empty(max_num_m_blocks, dtype=dtypes.i32, device=hidden_states.device),
"num_valid_ids": torch.empty(2, dtype=dtypes.i32, device=hidden_states.device),
"moe_buf": torch.empty((token_num, model_dim), dtype=hidden_states.dtype, device=hidden_states.device),
"tmp_out": torch.empty((token_num, topk, n1), dtype=hidden_states.dtype, device=hidden_states.device),
"stage1_out": torch.empty((token_num, topk, inter_dim), dtype=hidden_states.dtype, device=hidden_states.device),
}
_CKTILE_DIRECT_BUFS[key] = workspace
return workspace
def _direct_small_cktile(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
config,
):
block_m = 16
split_k = 2
topk = topk_ids.shape[1]
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"]
workspace = _get_direct_cktile_workspace(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
topk,
total_experts,
block_m,
)
sorted_ids = workspace["sorted_ids"]
sorted_weights = workspace["sorted_weights"]
sorted_expert_ids = workspace["sorted_expert_ids"]
num_valid_ids = workspace["num_valid_ids"]
moe_buf = workspace["moe_buf"]
tmp_out = workspace["tmp_out"]
stage1_out = workspace["stage1_out"]
aiter.moe_sorting_fwd(
topk_ids,
topk_weights,
sorted_ids,
sorted_weights,
sorted_expert_ids,
num_valid_ids,
moe_buf,
total_experts,
block_m,
None,
None,
0,
)
tmp_out.zero_()
aiter.moe_cktile2stages_gemm1(
hidden_states,
gate_up_weight_shuffled,
tmp_out,
sorted_ids,
sorted_expert_ids,
num_valid_ids,
topk,
intermediate_pad // 64 * 64 * 2,
hidden_pad // 128 * 128,
None,
None,
gate_up_weight_scale_shuffled.view(dtypes.fp8_e8m0),
None,
ActivationType.Silu,
block_m,
split_k,
)
aiter.silu_and_mul(stage1_out, tmp_out)
aiter.moe_cktile2stages_gemm2(
stage1_out,
down_weight_shuffled,
moe_buf,
sorted_ids,
sorted_expert_ids,
num_valid_ids,
topk,
hidden_pad // 64 * 64,
intermediate_pad // 128 * 128,
sorted_weights,
None,
down_weight_scale_shuffled.view(dtypes.fp8_e8m0),
None,
ActivationType.Silu,
block_m,
1,
)
return moe_buf
# ═══════════════════════════════════════════════════════════════════════
# Patch get_2stage_cfgs: CKTile for S1/S2/S4/S5, FlyDSL stage1 for bs=512
# ═══════════════════════════════════════════════════════════════════════
_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,
)
if (
token == 512
and expert == 257
and inter_dim == 256
and _is_flydsl_available()
):
default = replace(
default,
stage1=functools.partial(
_flydsl_stage1_wrapper,
tile_m=128,
tile_n=256,
tile_k=256,
a_dtype="fp4",
b_dtype="fp4",
out_dtype="bf16",
),
)
# 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_SMALL_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
hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
intermediate_pad = config["d_expert_pad"] - config["d_expert"]
if not _DIRECT_SMALL_DISABLED and _is_direct_small_shape(config, topk_ids.shape[1]):
try:
return _direct_small_cktile(
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:
_DIRECT_SMALL_DISABLED = True
print(f"[v350] direct S1/S2 path disabled: {e}", file=sys.stderr)
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 · 556 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 600014.
- """v328: v327 + v326 best tiles combined.+ """v367: v363 + enable dormant FlyDSL stage1 scheduler calls for S3.- dsv3 CSV: ksplit=2 for E=257 bs<=128, FlyDSL t64x256x256_reduce for S3- CKTile split_k=2 for S1/S2/S4/S5 via get_2stage_cfgs patch⋯ 1 unchanged lines- OPUS=1 globally- NT=1 globally- Cached sorting buffers+ - Direct `moe_cktile2stages_gemm1/2` path for S1/S2 only"""import functoolsimport os⋯ 108 unchanged linesos.environ["AITER_USE_OPUS_MOE_SORTING"] = "1"os.environ["AITER_USE_NT"] = "1"+ # ── Step 4: Repair the known FlyDSL stage1 FP4 source bug on the runner+ # and enable its dormant hot-loop scheduler calls. ──+ _flydsl_stage1_bug_path = (+ "/home/runner/aiter/aiter/ops/flydsl/kernels/mixed_moe_gemm_2stage.py"+ )+ try:+ with open(_flydsl_stage1_bug_path, "r") as f:+ _flydsl_src = f.read()+ _old_stage1_sig = (+ "def compute_f8f6f4_tile(\n"+ " acc_gate_in,\n"+ " acc_up_in,\n"+ " b_tile_in_gate,\n"+ " b_tile_in_up,\n"+ " lds_base,\n"+ " *,\n"+ " a0_prefetch=None,\n"+ " a_scale=None,\n"+ " b_scale_gate=None,\n"+ " b_scale_up=None,\n"+ " prefetch_epilogue: bool = False,\n"+ " ):"+ )+ _new_stage1_sig = (+ "def compute_f8f6f4_tile(\n"+ " acc_gate_in,\n"+ " acc_up_in,\n"+ " b_tile_in,\n"+ " lds_base,\n"+ " *,\n"+ " a0_prefetch=None,\n"+ " a_scale=None,\n"+ " b_scale=None,\n"+ " prefetch_epilogue: bool = False,\n"+ " ):"+ )+ if _old_stage1_sig in _flydsl_src:+ _flydsl_src = _flydsl_src.replace(_old_stage1_sig, _new_stage1_sig)+ _stage1_sched_disabled = (+ " # hot_loop_scheduler()\n"+ " gpu.barrier()"+ )+ _stage1_sched_enabled = (+ " hot_loop_scheduler()\n"+ " gpu.barrier()"+ )+ if _stage1_sched_disabled in _flydsl_src:+ _flydsl_src = _flydsl_src.replace(+ _stage1_sched_disabled,+ _stage1_sched_enabled,+ 3,+ )+ with open(_flydsl_stage1_bug_path, "w") as f:+ f.write(_flydsl_src)+ except Exception as e:+ print(f"[v367] flydsl stage1 source patch skipped: {e}", file=sys.stderr)+from task import input_t, output_tfrom aiter import ActivationType, QuantType, dtypesimport aiterimport aiter.fused_moe as fused_moe_mod+ import aiter.ops.flydslfrom aiter.ops.flydsl.utils import is_flydsl_available as _is_flydsl_availableprint(f"[v329] flydsl available: {_is_flydsl_available()}", file=sys.stderr)# ═══════════════════════════════════════════════════════════════════════+ # FlyDSL stage1 wrapper (matches CK stage1 signature)+ # ═══════════════════════════════════════════════════════════════════════+ def _flydsl_stage1_wrapper(+ hidden_states,+ w1,+ w2, # unused by stage1+ sorted_token_ids,+ sorted_expert_ids,+ num_valid_ids,+ out,+ topk,+ block_m=32,+ a1_scale=None,+ w1_scale=None,+ kernelName="",+ sorted_weights=None,+ tile_m=32,+ tile_n=256,+ tile_k=256,+ a_dtype="fp4",+ b_dtype="fp4",+ out_dtype="bf16",+ **_kwargs,+ ):+ return aiter.ops.flydsl.flydsl_moe_stage1(+ a=hidden_states,+ w1=w1,+ sorted_token_ids=sorted_token_ids,+ sorted_expert_ids=sorted_expert_ids,+ num_valid_ids=num_valid_ids,+ out=out,+ topk=topk,+ tile_m=tile_m,+ tile_n=tile_n,+ tile_k=tile_k,+ a_dtype=a_dtype,+ b_dtype=b_dtype,+ out_dtype=out_dtype,+ w1_scale=w1_scale,+ a1_scale=a1_scale,+ sorted_weights=sorted_weights,+ )+++ # ═══════════════════════════════════════════════════════════════════════# Cached sorting buffers (zero alloc after warmup)# ═══════════════════════════════════════════════════════════════════════_SORT_BUFS = {}⋯ 34 unchanged lines# ═══════════════════════════════════════════════════════════════════════- # Patch get_2stage_cfgs: CKTile for S1/S2/S4/S5, tuned S7+ # Direct CKTile path for official S1/S2 only# ═══════════════════════════════════════════════════════════════════════+ _DIRECT_SMALL_DISABLED = False+ _CKTILE_DIRECT_BUFS = {}+++ def _is_direct_small_shape(config, topk):+ total_experts = config["n_routed_experts"] + config["n_shared_experts"]+ return (+ config["d_hidden"] == 7168+ and topk == 9+ and total_experts == 257+ and config["d_expert"] == 256+ and config["bs"] in (16, 128)+ )+++ def _get_direct_cktile_workspace(hidden_states, w1, w2, topk, total_experts, block_m):+ _, n1, k1 = w1.shape+ _, k2, n2 = w2.shape+ model_dim = hidden_states.shape[1]+ inter_dim = n2 if k2 == k1 else n2 * 2+ if w1.dtype is torch.uint32:+ inter_dim *= 8++ token_num = hidden_states.shape[0]+ max_num_tokens_padded = int(token_num * topk + total_experts * block_m - topk)+ max_num_m_blocks = int((max_num_tokens_padded + block_m - 1) // block_m)+ key = (+ hidden_states.device,+ hidden_states.dtype,+ token_num,+ topk,+ total_experts,+ model_dim,+ n1,+ inter_dim,+ block_m,+ )+ workspace = _CKTILE_DIRECT_BUFS.get(key)+ if workspace is None:+ workspace = {+ "sorted_ids": torch.empty(max_num_tokens_padded, dtype=dtypes.i32, device=hidden_states.device),+ "sorted_weights": torch.empty(max_num_tokens_padded, dtype=dtypes.fp32, device=hidden_states.device),+ "sorted_expert_ids": torch.empty(max_num_m_blocks, dtype=dtypes.i32, device=hidden_states.device),+ "num_valid_ids": torch.empty(2, dtype=dtypes.i32, device=hidden_states.device),+ "moe_buf": torch.empty((token_num, model_dim), dtype=hidden_states.dtype, device=hidden_states.device),+ "tmp_out": torch.empty((token_num, topk, n1), dtype=hidden_states.dtype, device=hidden_states.device),+ "stage1_out": torch.empty((token_num, topk, inter_dim), dtype=hidden_states.dtype, device=hidden_states.device),+ }+ _CKTILE_DIRECT_BUFS[key] = workspace+ return workspace+++ def _direct_small_cktile(+ hidden_states,+ gate_up_weight_shuffled,+ down_weight_shuffled,+ gate_up_weight_scale_shuffled,+ down_weight_scale_shuffled,+ topk_weights,+ topk_ids,+ config,+ ):+ block_m = 16+ split_k = 2+ topk = topk_ids.shape[1]+ 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"]+ workspace = _get_direct_cktile_workspace(+ hidden_states,+ gate_up_weight_shuffled,+ down_weight_shuffled,+ topk,+ total_experts,+ block_m,+ )++ sorted_ids = workspace["sorted_ids"]+ sorted_weights = workspace["sorted_weights"]+ sorted_expert_ids = workspace["sorted_expert_ids"]+ num_valid_ids = workspace["num_valid_ids"]+ moe_buf = workspace["moe_buf"]+ tmp_out = workspace["tmp_out"]+ stage1_out = workspace["stage1_out"]++ aiter.moe_sorting_fwd(+ topk_ids,+ topk_weights,+ sorted_ids,+ sorted_weights,+ sorted_expert_ids,+ num_valid_ids,+ moe_buf,+ total_experts,+ block_m,+ None,+ None,+ 0,+ )++ tmp_out.zero_()+ aiter.moe_cktile2stages_gemm1(+ hidden_states,+ gate_up_weight_shuffled,+ tmp_out,+ sorted_ids,+ sorted_expert_ids,+ num_valid_ids,+ topk,+ intermediate_pad // 64 * 64 * 2,+ hidden_pad // 128 * 128,+ None,+ None,+ gate_up_weight_scale_shuffled.view(dtypes.fp8_e8m0),+ None,+ ActivationType.Silu,+ block_m,+ split_k,+ )+ aiter.silu_and_mul(stage1_out, tmp_out)++ aiter.moe_cktile2stages_gemm2(+ stage1_out,+ down_weight_shuffled,+ moe_buf,+ sorted_ids,+ sorted_expert_ids,+ num_valid_ids,+ topk,+ hidden_pad // 64 * 64,+ intermediate_pad // 128 * 128,+ sorted_weights,+ None,+ down_weight_scale_shuffled.view(dtypes.fp8_e8m0),+ None,+ ActivationType.Silu,+ block_m,+ 1,+ )+ return moe_buf+++ # ═══════════════════════════════════════════════════════════════════════+ # Patch get_2stage_cfgs: CKTile for S1/S2/S4/S5, FlyDSL stage1 for bs=512+ # ═══════════════════════════════════════════════════════════════════════_ORIGINAL_GET_2STAGE_CFGS = fused_moe_mod.get_2stage_cfgs⋯ 54 unchanged linesdoweight_stage1, hidden_pad, intermediate_pad, is_shuffled,)+ if (+ token == 512+ and expert == 257+ and inter_dim == 256+ and _is_flydsl_available()+ ):+ default = replace(+ default,+ stage1=functools.partial(+ _flydsl_stage1_wrapper,+ tile_m=128,+ tile_n=256,+ tile_k=256,+ a_dtype="fp4",+ b_dtype="fp4",+ out_dtype="bf16",+ ),+ )+# S7 (E=33, d=2048, bs=512): block_m=64 + NTif expert == 33 and inter_dim == 2048 and token == 512:return replace(default, block_m=64, use_non_temporal_load=True)⋯ 8 unchanged lines# Dispatch — simple, no state switching# ═══════════════════════════════════════════════════════════════════════def custom_kernel(data: input_t) -> output_t:+ global _DIRECT_SMALL_DISABLED+(hidden_states, gate_up_weight, down_weight,gate_up_weight_scale, down_weight_scale,⋯ 5 unchanged lineshidden_pad = config["d_hidden_pad"] - config["d_hidden"]intermediate_pad = config["d_expert_pad"] - config["d_expert"]+ if not _DIRECT_SMALL_DISABLED and _is_direct_small_shape(config, topk_ids.shape[1]):+ try:+ return _direct_small_cktile(+ 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:+ _DIRECT_SMALL_DISABLED = True+ print(f"[v350] direct S1/S2 path disabled: {e}", file=sys.stderr)+return fused_moe_mod.fused_moe(hidden_states, gate_up_weight_shuffled, down_weight_shuffled,topk_weights, topk_ids,
scrolls · 346 diff lines total
Best evidence level for this revision: reported
JSON