submission 611420
Aniket Sadashiva · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 460 lines, June 9 Researcher Reciprocity License v1.0.
submission_v503.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-611420?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:414d9a95897a0a36af49273b691afb5f941a6e8b0a5529706aa5dda12605afa0
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_v503.py460 lines
"""v503"""
import functools
import os
import sys
import torch
from dataclasses import replace
_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"
modified_lines.append(",".join(fields) + "\n")
elif expert_val == 257 and token_val == 512:
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")
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"[v490] dsv3 error: {e}", file=sys.stderr)
_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_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_t64x256x256_reduce"
_k2_2048_flydsl = "flydsl_moe2_afp4_wfp4_bf16_t32x128x256_atomic"
_csv_rows = [
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",
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
os.environ["AITER_USE_OPUS_MOE_SORTING"] = "1"
os.environ["AITER_USE_NT"] = "1"
os.environ.pop("FLIR_CK_LDS128", None)
os.environ["FLIR_MOE_STAGE1_SCHED"] = "1"
os.environ["FLIR_MOE_STAGE2_SCHED"] = "1"
os.environ["FLIR_MOE_STAGE2_PERSIST_M"] = "1"
from task import input_t, output_t
from aiter import ActivationType, QuantType, dtypes
import aiter
import aiter.fused_moe as fused_moe_mod
from aiter.ops.triton.quant.fused_mxfp4_quant import fused_dynamic_mxfp4_quant_moe_sort
from aiter.ops.flydsl.moe_kernels import flydsl_moe_stage2
_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
_ORIGINAL_GET_2STAGE_CFGS = fused_moe_mod.get_2stage_cfgs
_CKTILE_BUFS = {}
def _cached_cktile_moe_stage1(
hidden_states, w1, w2,
sorted_token_ids, sorted_expert_ids, num_valid_ids,
out, topk, block_m,
a1_scale, w1_scale, sorted_weights=None,
n_pad_zeros=0, k_pad_zeros=0, bias1=None,
activation=ActivationType.Silu, split_k=1, dtype=torch.bfloat16,
):
token_num = hidden_states.shape[0]
_, n1, k1 = w1.shape
_, k2, n2 = w2.shape
D = n2 if k2 == k1 else n2 * 2
if w1.dtype is torch.uint32:
D = D * 8
buf_key = (token_num, topk, D, w1.shape[1], split_k, hidden_states.device)
if buf_key not in _CKTILE_BUFS:
_CKTILE_BUFS[buf_key] = (
torch.empty((token_num, topk, D), dtype=dtype, device=hidden_states.device),
torch.zeros(
(token_num, topk, w1.shape[1]), dtype=hidden_states.dtype,
device=hidden_states.device,
) if split_k > 1 else None,
)
out_buf, tmp_buf = _CKTILE_BUFS[buf_key]
if split_k > 1:
tmp_buf.zero_()
aiter.moe_cktile2stages_gemm1(
hidden_states, w1, tmp_buf,
sorted_token_ids, sorted_expert_ids, num_valid_ids,
topk, n_pad_zeros, k_pad_zeros,
sorted_weights, a1_scale, w1_scale, bias1,
activation, block_m, split_k,
)
aiter.silu_and_mul(out_buf, tmp_buf)
else:
aiter.moe_cktile2stages_gemm1(
hidden_states, w1, out_buf,
sorted_token_ids, sorted_expert_ids, num_valid_ids,
topk, n_pad_zeros, k_pad_zeros,
sorted_weights, a1_scale, w1_scale, bias1,
activation, block_m, split_k,
)
return out_buf
def _make_cktile_metadata(hidden_pad, intermediate_pad, use_g1u1, activation, split_k=2):
return fused_moe_mod.MOEMetadata(
functools.partial(
_cached_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,
)
if expert == 257 and inter_dim == 256 and token == 16:
return _make_cktile_metadata(
hidden_pad, intermediate_pad, use_g1u1, activation, split_k=2,
)
if expert == 257 and inter_dim == 256 and token == 128:
return _make_cktile_metadata(
hidden_pad, intermediate_pad, use_g1u1, activation, split_k=4,
)
if expert == 33 and inter_dim == 512 and token <= 128:
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,
)
fused_moe_mod.get_2stage_cfgs = _patched_get_2stage_cfgs
_DIRECT_BUFS = {}
def _get_split_k(E, M, inter_dim):
if E == 257 and inter_dim == 256:
if M == 16:
return 2
if M == 128:
return 4
if E == 33 and inter_dim == 512:
if M == 16:
return 2
if M == 128:
return 1
return 0
def _direct_cktile_pipeline(
hidden_states, w1, w2, w1_scale, w2_scale,
topk_ids, topk_weights, E, M, topk, model_dim,
inter_dim, hidden_pad, intermediate_pad, split_k,
):
block_m = 16
sid, sw, sei, nvi, mb = _cached_moe_sorting_impl(
topk_ids, topk_weights, E, model_dim, torch.bfloat16,
block_m, None, None, 0, True,
)
_, n1, k1 = w1.shape
_, k2, n2 = w2.shape
D = n2 if k2 == k1 else n2 * 2
if w1.dtype is torch.uint32:
D = D * 8
n_pad1 = intermediate_pad // 64 * 64 * 2
k_pad1 = hidden_pad // 128 * 128
n_pad2 = hidden_pad // 64 * 64
k_pad2 = intermediate_pad // 128 * 128
buf_key = (M, topk, D, w1.shape[1], split_k)
if buf_key not in _DIRECT_BUFS:
dev = hidden_states.device
out = torch.empty((M, topk, D), dtype=torch.bfloat16, device=dev)
tmp = (
torch.zeros((M, topk, w1.shape[1]), dtype=torch.bfloat16, device=dev)
if split_k > 1 else None
)
_DIRECT_BUFS[buf_key] = (out, tmp)
out_buf, tmp_buf = _DIRECT_BUFS[buf_key]
w1_scale_e8m0 = w1_scale.view(dtypes.fp8_e8m0)
w2_scale_e8m0 = w2_scale.view(dtypes.fp8_e8m0)
if split_k > 1:
tmp_buf.zero_()
aiter.moe_cktile2stages_gemm1(
hidden_states, w1, tmp_buf,
sid, sei, nvi,
topk, n_pad1, k_pad1,
None, None, w1_scale_e8m0, None,
ActivationType.Silu, block_m, split_k,
)
aiter.silu_and_mul(out_buf, tmp_buf)
else:
aiter.moe_cktile2stages_gemm1(
hidden_states, w1, out_buf,
sid, sei, nvi,
topk, n_pad1, k_pad1,
None, None, w1_scale_e8m0, None,
ActivationType.Silu, block_m, 1,
)
aiter.moe_cktile2stages_gemm2(
out_buf, w2, mb,
sid, sei, nvi,
topk, n_pad2, k_pad2,
sw, None, w2_scale_e8m0, None,
ActivationType.Silu, block_m,
)
return mb
_A2_BUFS = {}
_S3_CK_K1 = (
"moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_"
"Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
_DIRECT_CK_CONFIGS = {
(257, 256): (_S3_CK_K1, 32, 16, 256, 256, "reduce"),
(33, 512): (_k1_512, 32, 16, 256, 256, "reduce"),
(33, 2048): (_k1_2048, 64, 16, 256, 256, "atomic"),
}
def _direct_ck_flydsl_pipeline(
hidden_states, w1, w2, w1_scale, w2_scale,
topk_ids, topk_weights, E, M, topk, model_dim, inter_dim,
):
ck_k1, block_m, fly_tm, fly_tn, fly_tk, fly_mode = _DIRECT_CK_CONFIGS[(E, inter_dim)]
sid, sw, sei, nvi, mb = _cached_moe_sorting_impl(
topk_ids, topk_weights, E, model_dim, torch.bfloat16,
block_m, None, None, 0, True,
)
a1, a1_scale = fused_dynamic_mxfp4_quant_moe_sort(
hidden_states, sorted_ids=sid, num_valid_ids=nvi,
token_num=M, topk=1, block_size=block_m,
)
buf_key = (M, topk, inter_dim)
if buf_key not in _A2_BUFS:
_A2_BUFS[buf_key] = torch.empty(
(M, topk, inter_dim), dtype=torch.bfloat16, device=hidden_states.device,
)
a2 = _A2_BUFS[buf_key]
w1_scale_e8m0 = w1_scale.view(dtypes.fp8_e8m0)
aiter.ck_moe_stage1_fwd(
a1, w1, w2, sid, sei, nvi, a2, topk,
ck_k1, w1_scale_e8m0, a1_scale, block_m,
None, QuantType.per_1x32, ActivationType.Silu, 0, True,
torch.bfloat16,
)
a2_flat = a2.view(-1, inter_dim)
a2_quant, a2_scale = fused_dynamic_mxfp4_quant_moe_sort(
a2_flat, sorted_ids=sid, num_valid_ids=nvi,
token_num=M, topk=topk, block_size=block_m,
)
a2_quant = a2_quant.view(M, topk, -1)
w2_scale_e8m0 = w2_scale.view(dtypes.fp8_e8m0)
flydsl_moe_stage2(
inter_states=a2_quant, w2=w2,
sorted_token_ids=sid, sorted_expert_ids=sei,
num_valid_ids=nvi, out=mb, topk=topk,
tile_m=fly_tm, tile_n=fly_tn, tile_k=fly_tk,
a_dtype="fp4", b_dtype="fp4", out_dtype="bf16",
mode=fly_mode,
w2_scale=w2_scale_e8m0, a2_scale=a2_scale,
sorted_weights=sw,
)
return mb
def custom_kernel(data: input_t) -> output_t:
(
hidden_states, gate_up_weight, down_weight,
gate_up_weight_scale, down_weight_scale,
gate_up_weight_shuffled, down_weight_shuffled,
gate_up_weight_scale_shuffled, down_weight_scale_shuffled,
topk_weights, topk_ids, config,
) = data
hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
intermediate_pad = config["d_expert_pad"] - config["d_expert"]
M = hidden_states.shape[0]
E = gate_up_weight.shape[0]
inter_dim = config["d_expert"]
model_dim = hidden_states.shape[1]
topk = topk_ids.shape[1]
split_k = _get_split_k(E, M, inter_dim)
if split_k > 0 and model_dim == 7168 and topk == 9:
return _direct_cktile_pipeline(
hidden_states,
gate_up_weight_shuffled, down_weight_shuffled,
gate_up_weight_scale_shuffled, down_weight_scale_shuffled,
topk_ids, topk_weights, E, M, topk, model_dim,
inter_dim, hidden_pad, intermediate_pad, split_k,
)
if M == 512 and model_dim == 7168 and topk == 9 and (E, inter_dim) in _DIRECT_CK_CONFIGS:
return _direct_ck_flydsl_pipeline(
hidden_states,
gate_up_weight_shuffled, down_weight_shuffled,
gate_up_weight_scale_shuffled, down_weight_scale_shuffled,
topk_ids, topk_weights, E, M, topk, model_dim, inter_dim,
)
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 · 460 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 601016.
- """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- """+ """v503"""import functoolsimport osimport sys⋯ 2 unchanged linesfrom 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:⋯ 13 unchanged linesmodified_lines.append(line)continueif expert_val == 257 and token_val <= 128:- fields[14] = "2" # ksplit=2+ fields[14] = "2"modified_lines.append(",".join(fields) + "\n")elif expert_val == 257 and token_val == 512:- # FlyDSL reduce stage2 for S3, with CK fallbackflydsl_fields = list(fields)flydsl_fields[19] = _flydsl_s3_stage2flydsl_fields[20] = "0.1%"if len(flydsl_fields) > 25:flydsl_fields[25] = ""modified_lines.append(",".join(flydsl_fields) + "\n")- # CK fallback rowfallback_fields = list(fields)if len(fallback_fields) > 25:fallback_fields[25] = "flydsl_fallback"⋯ 5 unchanged lineswith open(_dsv3_path, "w") as f:f.writelines(modified_lines)except Exception as e:- print(f"[v329] dsv3 error: {e}", file=sys.stderr)+ print(f"[v490] 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,"⋯ 3 unchanged lines"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"⋯ 10 unchanged lines"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"+ _k2_512_flydsl = "flydsl_moe2_afp4_wfp4_bf16_t64x256x256_reduce"+ _k2_2048_flydsl = "flydsl_moe2_afp4_wfp4_bf16_t32x128x256_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 fallbackf"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 fallbackf"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",]⋯ 6 unchanged linesexcept 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"+ os.environ.pop("FLIR_CK_LDS128", None)+ os.environ["FLIR_MOE_STAGE1_SCHED"] = "1"+ os.environ["FLIR_MOE_STAGE2_SCHED"] = "1"+ os.environ["FLIR_MOE_STAGE2_PERSIST_M"] = "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.flydsl- from aiter.ops.flydsl.utils import is_flydsl_available as _is_flydsl_available+ from aiter.ops.triton.quant.fused_mxfp4_quant import fused_dynamic_mxfp4_quant_moe_sort+ from aiter.ops.flydsl.moe_kernels import flydsl_moe_stage2- 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 = {}⋯ 31 unchanged linesfused_moe_mod._moe_sorting_impl = _cached_moe_sorting_impl- # ═══════════════════════════════════════════════════════════════════════- # Direct CKTile path for official S1/S2 only- # ═══════════════════════════════════════════════════════════════════════- _DIRECT_SMALL_DISABLED = False- _CKTILE_DIRECT_BUFS = {}+ _ORIGINAL_GET_2STAGE_CFGS = fused_moe_mod.get_2stage_cfgs+ _CKTILE_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):+ def _cached_cktile_moe_stage1(+ hidden_states, w1, w2,+ sorted_token_ids, sorted_expert_ids, num_valid_ids,+ out, topk, block_m,+ a1_scale, w1_scale, sorted_weights=None,+ n_pad_zeros=0, k_pad_zeros=0, bias1=None,+ activation=ActivationType.Silu, split_k=1, dtype=torch.bfloat16,+ ):+ token_num = hidden_states.shape[0]_, n1, k1 = w1.shape_, k2, n2 = w2.shape- model_dim = hidden_states.shape[1]- inter_dim = n2 if k2 == k1 else n2 * 2+ D = n2 if k2 == k1 else n2 * 2if w1.dtype is torch.uint32:- inter_dim *= 8+ D = D * 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+ buf_key = (token_num, topk, D, w1.shape[1], split_k, hidden_states.device)+ if buf_key not in _CKTILE_BUFS:+ _CKTILE_BUFS[buf_key] = (+ torch.empty((token_num, topk, D), dtype=dtype, device=hidden_states.device),+ torch.zeros(+ (token_num, topk, w1.shape[1]), dtype=hidden_states.dtype,+ device=hidden_states.device,+ ) if split_k > 1 else None,+ )+ out_buf, tmp_buf = _CKTILE_BUFS[buf_key]- 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,- )+ if split_k > 1:+ tmp_buf.zero_()+ aiter.moe_cktile2stages_gemm1(+ hidden_states, w1, tmp_buf,+ sorted_token_ids, sorted_expert_ids, num_valid_ids,+ topk, n_pad_zeros, k_pad_zeros,+ sorted_weights, a1_scale, w1_scale, bias1,+ activation, block_m, split_k,+ )+ aiter.silu_and_mul(out_buf, tmp_buf)+ else:+ aiter.moe_cktile2stages_gemm1(+ hidden_states, w1, out_buf,+ sorted_token_ids, sorted_expert_ids, num_valid_ids,+ topk, n_pad_zeros, k_pad_zeros,+ sorted_weights, a1_scale, w1_scale, bias1,+ activation, block_m, split_k,+ )+ return out_buf- 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,+ _cached_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,⋯ 29 unchanged linesdoweight_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:+ if expert == 257 and inter_dim == 256 and token == 16:return _make_cktile_metadata(hidden_pad, intermediate_pad, use_g1u1, activation, split_k=2,)+ if expert == 257 and inter_dim == 256 and token == 128:+ return _make_cktile_metadata(+ hidden_pad, intermediate_pad, use_g1u1, activation, split_k=4,+ )- # 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(+ 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,)- 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",- ),++ fused_moe_mod.get_2stage_cfgs = _patched_get_2stage_cfgs+++ _DIRECT_BUFS = {}+++ def _get_split_k(E, M, inter_dim):+ if E == 257 and inter_dim == 256:+ if M == 16:+ return 2+ if M == 128:+ return 4+ if E == 33 and inter_dim == 512:+ if M == 16:+ return 2+ if M == 128:+ return 1+ return 0+++ def _direct_cktile_pipeline(+ hidden_states, w1, w2, w1_scale, w2_scale,+ topk_ids, topk_weights, E, M, topk, model_dim,+ inter_dim, hidden_pad, intermediate_pad, split_k,+ ):+ block_m = 16++ sid, sw, sei, nvi, mb = _cached_moe_sorting_impl(+ topk_ids, topk_weights, E, model_dim, torch.bfloat16,+ block_m, None, None, 0, True,+ )++ _, n1, k1 = w1.shape+ _, k2, n2 = w2.shape+ D = n2 if k2 == k1 else n2 * 2+ if w1.dtype is torch.uint32:+ D = D * 8++ n_pad1 = intermediate_pad // 64 * 64 * 2+ k_pad1 = hidden_pad // 128 * 128+ n_pad2 = hidden_pad // 64 * 64+ k_pad2 = intermediate_pad // 128 * 128++ buf_key = (M, topk, D, w1.shape[1], split_k)+ if buf_key not in _DIRECT_BUFS:+ dev = hidden_states.device+ out = torch.empty((M, topk, D), dtype=torch.bfloat16, device=dev)+ tmp = (+ torch.zeros((M, topk, w1.shape[1]), dtype=torch.bfloat16, device=dev)+ if split_k > 1 else None)+ _DIRECT_BUFS[buf_key] = (out, tmp)- # 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)+ out_buf, tmp_buf = _DIRECT_BUFS[buf_key]- return default+ w1_scale_e8m0 = w1_scale.view(dtypes.fp8_e8m0)+ w2_scale_e8m0 = w2_scale.view(dtypes.fp8_e8m0)+ if split_k > 1:+ tmp_buf.zero_()+ aiter.moe_cktile2stages_gemm1(+ hidden_states, w1, tmp_buf,+ sid, sei, nvi,+ topk, n_pad1, k_pad1,+ None, None, w1_scale_e8m0, None,+ ActivationType.Silu, block_m, split_k,+ )+ aiter.silu_and_mul(out_buf, tmp_buf)+ else:+ aiter.moe_cktile2stages_gemm1(+ hidden_states, w1, out_buf,+ sid, sei, nvi,+ topk, n_pad1, k_pad1,+ None, None, w1_scale_e8m0, None,+ ActivationType.Silu, block_m, 1,+ )- fused_moe_mod.get_2stage_cfgs = _patched_get_2stage_cfgs+ aiter.moe_cktile2stages_gemm2(+ out_buf, w2, mb,+ sid, sei, nvi,+ topk, n_pad2, k_pad2,+ sw, None, w2_scale_e8m0, None,+ ActivationType.Silu, block_m,+ )+ return mb- # ═══════════════════════════════════════════════════════════════════════- # Dispatch — simple, no state switching- # ═══════════════════════════════════════════════════════════════════════- def custom_kernel(data: input_t) -> output_t:- global _DIRECT_SMALL_DISABLED+ _A2_BUFS = {}++ _S3_CK_K1 = (+ "moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_"+ "Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"+ )++ _DIRECT_CK_CONFIGS = {+ (257, 256): (_S3_CK_K1, 32, 16, 256, 256, "reduce"),+ (33, 512): (_k1_512, 32, 16, 256, 256, "reduce"),+ (33, 2048): (_k1_2048, 64, 16, 256, 256, "atomic"),+ }+++ def _direct_ck_flydsl_pipeline(+ hidden_states, w1, w2, w1_scale, w2_scale,+ topk_ids, topk_weights, E, M, topk, model_dim, inter_dim,+ ):+ ck_k1, block_m, fly_tm, fly_tn, fly_tk, fly_mode = _DIRECT_CK_CONFIGS[(E, inter_dim)]++ sid, sw, sei, nvi, mb = _cached_moe_sorting_impl(+ topk_ids, topk_weights, E, model_dim, torch.bfloat16,+ block_m, None, None, 0, True,+ )++ a1, a1_scale = fused_dynamic_mxfp4_quant_moe_sort(+ hidden_states, sorted_ids=sid, num_valid_ids=nvi,+ token_num=M, topk=1, block_size=block_m,+ )++ buf_key = (M, topk, inter_dim)+ if buf_key not in _A2_BUFS:+ _A2_BUFS[buf_key] = torch.empty(+ (M, topk, inter_dim), dtype=torch.bfloat16, device=hidden_states.device,+ )+ a2 = _A2_BUFS[buf_key]++ w1_scale_e8m0 = w1_scale.view(dtypes.fp8_e8m0)+ aiter.ck_moe_stage1_fwd(+ a1, w1, w2, sid, sei, nvi, a2, topk,+ ck_k1, w1_scale_e8m0, a1_scale, block_m,+ None, QuantType.per_1x32, ActivationType.Silu, 0, True,+ torch.bfloat16,+ )++ a2_flat = a2.view(-1, inter_dim)+ a2_quant, a2_scale = fused_dynamic_mxfp4_quant_moe_sort(+ a2_flat, sorted_ids=sid, num_valid_ids=nvi,+ token_num=M, topk=topk, block_size=block_m,+ )+ a2_quant = a2_quant.view(M, topk, -1)++ w2_scale_e8m0 = w2_scale.view(dtypes.fp8_e8m0)+ flydsl_moe_stage2(+ inter_states=a2_quant, w2=w2,+ sorted_token_ids=sid, sorted_expert_ids=sei,+ num_valid_ids=nvi, out=mb, topk=topk,+ tile_m=fly_tm, tile_n=fly_tn, tile_k=fly_tk,+ a_dtype="fp4", b_dtype="fp4", out_dtype="bf16",+ mode=fly_mode,+ w2_scale=w2_scale_e8m0, a2_scale=a2_scale,+ sorted_weights=sw,+ )++ return mb+++ def custom_kernel(data: input_t) -> output_t:(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)+ M = hidden_states.shape[0]+ E = gate_up_weight.shape[0]+ inter_dim = config["d_expert"]+ model_dim = hidden_states.shape[1]+ topk = topk_ids.shape[1]+ split_k = _get_split_k(E, M, inter_dim)++ if split_k > 0 and model_dim == 7168 and topk == 9:+ return _direct_cktile_pipeline(+ hidden_states,+ gate_up_weight_shuffled, down_weight_shuffled,+ gate_up_weight_scale_shuffled, down_weight_scale_shuffled,+ topk_ids, topk_weights, E, M, topk, model_dim,+ inter_dim, hidden_pad, intermediate_pad, split_k,+ )++ if M == 512 and model_dim == 7168 and topk == 9 and (E, inter_dim) in _DIRECT_CK_CONFIGS:+ return _direct_ck_flydsl_pipeline(+ hidden_states,+ gate_up_weight_shuffled, down_weight_shuffled,+ gate_up_weight_scale_shuffled, down_weight_scale_shuffled,+ topk_ids, topk_weights, E, M, topk, model_dim, inter_dim,+ )+return fused_moe_mod.fused_moe(hidden_states, gate_up_weight_shuffled, down_weight_shuffled,topk_weights, topk_ids,
scrolls · 686 diff lines total
Best evidence level for this revision: reported
JSON