Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 18 additions & 0 deletions aiter/configs/model_configs/kimik3_a4w4_tuned_fmoe.csv
Original file line number Diff line number Diff line change
@@ -0,0 +1,18 @@
gfx,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,xbf16,flat,tflops,bw,_tag
gfx950,256,1,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,32,0,10.4812,flydsl_moe1_afp4_wfp4_bf16_t32x32x256_w4_kw2_fp4,0.1%,9.5512,flydsl_moe2_afp4_wfp4_bf16_t32x128x128_reduce,0.0%,20.0324,0,0,0,6.6,184670.18,
gfx950,256,2,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,32,0,12.2582,flydsl_moe1_afp4_wfp4_bf16_t32x32x256_w2_kw2_fp4,0.1%,11.1954,flydsl_moe2_afp4_wfp4_bf16_t32x256x256_reduce,0.0%,23.4536,0,0,0,11.27,157732.61,
gfx950,256,3,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,32,0,15.7919,flydsl_moe1_afp4_wfp4_bf16_t32x32x256_w4_kw2_fp4,0.1%,13.3351,flydsl_moe2_afp4_wfp4_bf16_t32x128x256_reduce,0.0%,29.127,0,0,0,13.61,127009.59,
gfx950,256,4,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,32,0,17.9122,flydsl_moe1_afp4_wfp4_bf16_t32x32x256_w3_kw2_fp4,0.1%,14.9478,flydsl_moe2_afp4_wfp4_bf16_t32x128x128_reduce_bnt2,0.0%,32.86,0,0,0,16.08,112581.23,
gfx950,256,8,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,32,0,30.6297,flydsl_moe1_afp4_wfp4_bf16_t32x64x256_w3_kw2_fp4,0.1%,22.749,flydsl_moe2_afp4_wfp4_bf16_t32x128x128_reduce_bnt2,0.0%,53.3787,0,0,0,19.8,69305.96,
gfx950,256,16,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,32,0,55.6312,flydsl_moe1_afp4_wfp4_bf16_t32x128x256_w3_fp4,0.1%,35.569,flydsl_moe2_afp4_wfp4_bf16_t32x128x128_atomic_bnt2,0.2%,91.2002,0,0,0,23.18,40565.13,
gfx950,256,32,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,32,0,91.5468,flydsl_moe1_afp4_wfp4_bf16_t32x64x256_w4_kw2_fp4,0.1%,55.7707,flydsl_moe2_afp4_wfp4_bf16_t32x256x128_reduce_bnt2_persist,0.0%,147.3175,0,0,0,28.7,25113.92,
gfx950,256,64,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,32,0,139.5307,flydsl_moe1_afp4_wfp4_bf16_t32x32x256_w4_kw2_fp4,0.1%,86.3251,flydsl_moe2_afp4_wfp4_bf16_t32x128x128_atomic_bnt2_persist,0.3%,225.8558,0,0,0,37.44,16382.42,
gfx950,256,128,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,32,0,180.2132,flydsl_moe1_afp4_wfp4_bf16_t32x64x256_w4_kw2_fp4,0.2%,109.1208,flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic_sbm32,0.3%,289.334,0,0,0,58.45,12790.59,
gfx950,256,256,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,32,0,199.9643,flydsl_moe1_afp4_wfp4_bf16_t32x64x256_w3_kw2_fp4,0.1%,125.3707,flydsl_moe2_afp4_wfp4_bf16_t32x256x256_atomic_bnt2_persist,0.3%,325.335,0,0,0,103.96,11379.44,
gfx950,256,512,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,32,0,199.4349,flydsl_moe1_afp4_wfp4_bf16_t32x64x256_w4_kw2_fp4,0.2%,129.6391,flydsl_moe2_afp4_wfp4_bf16_t32x128x256_atomic_bnt2,0.3%,329.074,0,0,0,205.56,11258.5,
gfx950,256,1024,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,32,0,205.3733,flydsl_moe1_afp4_wfp4_bf16_t32x128x256_w4_fp4,0.1%,144.5268,flydsl_moe2_afp4_wfp4_bf16_t32x256x128_atomic_persist,0.3%,349.9001,0,0,0,386.66,10604.13,
gfx950,256,2048,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,64,0,215.5931,flydsl_moe1_afp4_wfp4_bf16_t64x128x256_w3_fp4,0.1%,212.6502,flydsl_moe2_afp4_wfp4_bf16_t64x256x256_reduce,0.0%,428.2433,0,0,0,631.84,8689.91,
gfx950,256,4096,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,128,0,314.3113,flydsl_moe1_afp4_wfp4_bf16_t128x128x256_w4_fp4,0.1%,338.8907,flydsl_moe2_afp4_wfp4_bf16_t64x256x128_reduce_xcd4_persist_sbm128,0.0%,653.202,0,0,0,828.48,5730.87,
gfx950,256,8192,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,64,0,422.9114,flydsl_moe1_afp4_wfp4_bf16_t64x128x256_w4_bnt0_xcd4_fp4,0.1%,597.3431,flydsl_moe2_afp4_wfp4_bf16_t64x256x256_reduce_bnt2_xcd4,0.0%,1020.2545,0,0,0,1060.84,3712.27,
gfx950,256,16384,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,64,0,642.4294,flydsl_moe1_afp4_wfp4_bf16_t64x128x256_w3_bnt0_xcd4_fp4,0.1%,1074.0772,flydsl_moe2_afp4_wfp4_bf16_t64x256x256_reduce_bnt2_xcd4_persist,0.0%,1716.5066,0,0,0,1261.09,2257.8,
gfx950,256,32768,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,64,0,1111.7422,flydsl_moe1_afp4_wfp4_bf16_t64x128x256_w2_bnt0_xcd4_fp4,0.1%,2185.1161,flydsl_moe2_afp4_wfp4_bf16_t64x256x256_reduce_xcd4,0.0%,3296.8583,0,0,0,1313.17,1228.96,
18 changes: 18 additions & 0 deletions aiter/configs/model_configs/kimik3_a4w4_untuned_fmoe.csv
Original file line number Diff line number Diff line change
@@ -0,0 +1,18 @@
token,model_dim,inter_dim,expert,topk,act_type,dtype,q_dtype_a,q_dtype_w,q_type,use_g1u1,doweight_stage1
1,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0
2,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0
3,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0
4,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0
8,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0
16,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0
32,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0
64,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0
128,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0
256,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0
512,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0
1024,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0
2048,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0
4096,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0
8192,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0
16384,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0
Comment thread
lalala-sh marked this conversation as resolved.
32768,3584,384,896,16,ActivationType.Situv2,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0
28 changes: 14 additions & 14 deletions aiter/fused_moe.py
Original file line number Diff line number Diff line change
Expand Up @@ -714,26 +714,26 @@ def _fused_moe_impl(
# mxfp8: both activation and weight are fp8 (per-1x32 e8m0 microscale).
q_dtype_a = dtypes.fp8
elif quant_type == QuantType.per_1x32:
if activation == ActivationType.Swiglu and gate_mode == GateMode.SEPARATED:
if activation == ActivationType.Situv2:
Comment thread
lalala-sh marked this conversation as resolved.
# SiTUv2 defaults to a16w4 (bf16 activation x mxfp4 weight) on the
# mixed_moe kernels. AITER_SITUV2_A8W4 / AITER_SITUV2_A4W4 select the
# fp8 / fp4 activation instead; each has its own tuned config
# (kimik3_{a8w4,a4w4}_tuned_fmoe.csv). Tested before the INTERLEAVE
# branch below, which would otherwise claim SiTUv2 and pick the
# activation dtype itself.
Comment thread
XiaobingSuper marked this conversation as resolved.
if os.environ.get("AITER_SITUV2_A8W4", "0") == "1":
q_dtype_a = dtypes.fp8
elif os.environ.get("AITER_SITUV2_A4W4", "0") == "1":
q_dtype_a = dtypes.fp4x2
else:
q_dtype_a = dtypes.bf16
Comment thread
XiaobingSuper marked this conversation as resolved.
elif activation == ActivationType.Swiglu and gate_mode == GateMode.SEPARATED:
q_dtype_a = dtypes.bf16 if M < _SWIGLU_MXFP4_BF16_BOUND else dtypes.fp4x2
elif activation == ActivationType.Swiglu or gate_mode == GateMode.INTERLEAVE:
if get_gfx() != "gfx950" or M < bf16_fp8_bound:
q_dtype_a = dtypes.bf16
else:
q_dtype_a = dtypes.fp8
elif activation == ActivationType.Situv2:
# SiTUv2 + separated == a16w4 (bf16 activation x mxfp4 weight); keep
# the activation in bf16 (no fp4 quant). a4w4 SiTUv2 full 2-stage is
# unsupported (no CK situv2 stage2), so separated-mode SiTUv2 always
# maps to the mixed_moe a16w4 kernel. AITER_SITUV2_A8W4=1 overrides
# to fp8 activation (a8w4) via the tuned flydsl afp8_wfp4 config.
# NB: on gfx1250 this is overridden below (fp4x2 / a8w4), so K3 is
# unaffected by this branch.
q_dtype_a = (
dtypes.fp8
if os.environ.get("AITER_SITUV2_A8W4", "0") == "1"
else dtypes.bf16
)
else:
q_dtype_a = dtypes.fp4x2

Expand Down
60 changes: 57 additions & 3 deletions csrc/ck_gemm_moe_2stages_codegen/gemm_moe_tune.py
Original file line number Diff line number Diff line change
Expand Up @@ -100,6 +100,13 @@ def is_flydsl_available():

COS_DIFF_THRESHOLD = 1e-1

# SiTUv2 params (Kimi-K3 config.json), passed to BOTH the torch reference and
# the kernel launch -- they used to fall back to their own differing defaults
# ((2.0, 1.5) vs (1.0, 1.0)), which scored every SiTUv2 stage1 candidate at
# ~35% err and failed them all against --errRatio.
_TUNER_SITU_BETA = 4.0
_TUNER_SITU_LINEAR_BETA = 25.0
Comment thread
XiaobingSuper marked this conversation as resolved.


def _manifest_flat_by_kernel(df: pd.DataFrame) -> dict:
"""Map ``knl_name`` -> 0/1 when the manifest has a ``flat`` column.
Expand Down Expand Up @@ -141,11 +148,31 @@ def torch_dynamic_mxfp8_quant(x: torch.Tensor):
)


_E2M1_MAG = (0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0)
_E2M1_LUT = _E2M1_MAG + tuple(-v for v in _E2M1_MAG)


def _to_f64_flat(t):
"""Flatten to float64, unpacking fp4x2 by hand.

``fp4x2.double()`` faults the queue (HIP unspecified launch failure), and
viewing it as uint8 would compare two packed e2m1 codes as one byte value,
so decode the nibbles into their represented values instead.
"""
if t.dtype == dtypes.fp4x2:
b = t.reshape(-1).view(torch.uint8)
lut = torch.tensor(_E2M1_LUT, device=b.device, dtype=torch.float64)
return torch.stack(
[lut[(b & 0xF).long()], lut[((b >> 4) & 0xF).long()]], dim=-1
).flatten()
return t.double().flatten()


def cosine_diff_compare(ref, res, msg="", printLog=True):
from aiter import logger

x = ref.double().flatten()
y = res.double().flatten()
x = _to_f64_flat(ref)
y = _to_f64_flat(res)
cos_diff = 1 - 2 * (x * y).sum().item() / max((x * x + y * y).sum().item(), 1e-12)
if printLog:
if cos_diff < COS_DIFF_THRESHOLD:
Expand Down Expand Up @@ -643,6 +670,8 @@ def run_flydsl_stage1_out(
b_dtype=kparams["b_dtype"],
out_dtype=_out_dtype,
act=act,
situ_beta=_TUNER_SITU_BETA,
situ_linear_beta=_TUNER_SITU_LINEAR_BETA,
w1_scale=w1_scale_aiter,
a1_scale=a1_scale,
sorted_weights=sorted_weights,
Expand Down Expand Up @@ -1702,6 +1731,8 @@ def run_torch_moe_stage1(
w1_scale=w1_scale,
w1_bias=w1_bias,
doweight=doweight_stage1,
situ_beta=_TUNER_SITU_BETA,
situ_linear_beta=_TUNER_SITU_LINEAR_BETA,
)
token_num = a1_qt.shape[0]
if fuse_fp4:
Expand Down Expand Up @@ -1980,6 +2011,8 @@ def torch_moe_2stages(
a1_scale=a1_scale,
w1_scale=w1_scale,
doweight=doweight_stage1,
situ_beta=_TUNER_SITU_BETA,
situ_linear_beta=_TUNER_SITU_LINEAR_BETA,
)
AQDType = hidden_states.dtype

Expand Down Expand Up @@ -2997,6 +3030,24 @@ def gen_flydsl_2stages_task(self, info, blockMs):
for kname, kparams in flydsl_s1_kernels.items():
is_splitk = kparams.get("k_batch", 1) > 1

# Drop k_batch/k_wave splits that do not divide the K axis.
# compile_mixed_moe_gemm1 rejects them, but only after the
# candidate has been dispatched, and on some shapes the launch
# faults the queue (HSA_STATUS_ERROR_EXCEPTION) instead of
# raising, taking the worker pool down with it.
_kb = kparams.get("k_batch", 1)
_kw = kparams.get("k_wave", 1)
_tk = kparams["tile_k"]
if model_dim % _kb != 0:
continue
_k_per_batch = model_dim // _kb
if _k_per_batch % _tk != 0:
continue
if _kw > 1 and (
_k_per_batch % _kw != 0 or (_k_per_batch // _kw) % _tk != 0
):
continue

# (kernel_name, kparams, is_fp4, is_fp8)
# out_dtype encodes fused quant type: "fp4" or "fp8"
# a8w4 (a_dtype_str="fp8"): stage2 expects fp8 activations -> out_dtype="fp8"
Expand Down Expand Up @@ -3038,7 +3089,10 @@ def gen_flydsl_2stages_task(self, info, blockMs):

for s1_name, s1_params, is_fp4, is_fp8 in s1_variants:
s1_compare_fn = None
if is_fp8 or a_dtype_str == "fp8":
if is_fp8 or is_fp4 or a_dtype_str in ("fp8", "fp4"):
# Fused stage1 emits packed mx values; the default
# comparison reads fp4x2 as raw bytes (two e2m1 codes
# each), which rejects >99% code-identical candidates.
s1_compare_fn = cosine_diff_compare
ref_args_extra = (
[
Expand Down
58 changes: 35 additions & 23 deletions op_tests/test_moe_2stage.py
Original file line number Diff line number Diff line change
Expand Up @@ -230,9 +230,7 @@ def weight_per_128x128_quant(weight, quant_dtype):
# a16w4 by the caller but dispatched as a8w4 on gfx950.
reference_aq_dtype = AQDType
if actType == aiter.ActivationType.Situv2:
runtime_aq_dtype = _runtime_situv2_mxfp4_q_dtype_a(
token, gateMode, qType, WQDType
)
runtime_aq_dtype = _runtime_situv2_mxfp4_q_dtype_a(qType, WQDType)
if runtime_aq_dtype is not None:
reference_aq_dtype = runtime_aq_dtype

Expand Down Expand Up @@ -885,25 +883,37 @@ def _iter_csv_cases():
)
continue
# The reference path below uses the CSV q_dtype_a directly, while
# fused_moe selects q_dtype_a from the current Swiglu MXFP4 runtime mode.
# Skip CSV rows that are tuned for a different mode to avoid comparing
# e.g. an fp4x2 reference against a bf16/fp8 runtime dispatch.
expected_aq_dtype = _runtime_swiglu_mxfp4_q_dtype_a(
kwargs["token"],
kwargs["actType"],
kwargs["gateMode"],
kwargs["qType"],
kwargs["AQDType"],
kwargs["WQDType"],
)
# fused_moe selects q_dtype_a from the current runtime mode. Skip CSV
# rows tuned for a different mode (e.g. a4w4/a8w4 without the opt-in env).
if kwargs["actType"] == aiter.ActivationType.Situv2:
expected_aq_dtype = _runtime_situv2_mxfp4_q_dtype_a(
kwargs["qType"], kwargs["WQDType"]
)
runtime_mode = "SiTUv2 MXFP4"
# SiTUv2 a16w4 never ran before this ordering fix and every row
# fails: _effective_gate_mode asks for INTERLEAVE while stage1 binds
# gate_mode="separated" for non-fp8 activations.
if kwargs["AQDType"] == dtypes.bf16 and kwargs["WQDType"] == dtypes.fp4x2:
continue
else:
expected_aq_dtype = _runtime_swiglu_mxfp4_q_dtype_a(
kwargs["token"],
kwargs["actType"],
kwargs["gateMode"],
kwargs["qType"],
kwargs["AQDType"],
kwargs["WQDType"],
)
runtime_mode = "Swiglu MXFP4"
if expected_aq_dtype is not None and kwargs["AQDType"] != expected_aq_dtype:
aiter.logger.info(
"skip row token=%s dim=(%s,%s): q_dtype_a=%s does not match "
"current Swiglu MXFP4 runtime mode (expected %s)",
"current %s runtime mode (expected %s)",
row.get("token"),
row.get("model_dim"),
row.get("inter_dim"),
kwargs["AQDType"],
runtime_mode,
expected_aq_dtype,
)
continue
Expand Down Expand Up @@ -955,7 +965,7 @@ def _effective_swiglu_limit(quant_type, aq_dtype, wq_dtype, swiglu_limit):
return None


def _runtime_situv2_mxfp4_q_dtype_a(token, gate_mode, q_type, wq_dtype):
def _runtime_situv2_mxfp4_q_dtype_a(q_type, wq_dtype):
"""Mirror fused_moe's SiTUv2 MXFP4 activation-dtype routing."""
if q_type != aiter.QuantType.per_1x32 or wq_dtype != dtypes.fp4x2:
return None
Expand All @@ -967,13 +977,15 @@ def _runtime_situv2_mxfp4_q_dtype_a(token, gate_mode, q_type, wq_dtype):
else dtypes.fp4x2
)

if GateMode(gate_mode) == GateMode.INTERLEAVE:
bound = int(os.environ.get("AITER_BF16_FP8_MOE_BOUND", "256"))
return dtypes.bf16 if get_gfx() != "gfx950" or token < bound else dtypes.fp8

return (
dtypes.fp8 if os.environ.get("AITER_SITUV2_A8W4", "0") == "1" else dtypes.bf16
)
# fused_moe tests SiTUv2 ahead of the Swiglu/INTERLEAVE branch, so gate mode
# and token count do not enter into it -- mirror that order here, otherwise
# a4w4/a8w4 rows are skipped as "mode mismatch" under gate_mode=INTERLEAVE
# and the opt-in paths go untested.
if os.environ.get("AITER_SITUV2_A8W4", "0") == "1":
return dtypes.fp8
if os.environ.get("AITER_SITUV2_A4W4", "0") == "1":
return dtypes.fp4x2
return dtypes.bf16


def _runtime_swiglu_mxfp4_q_dtype_a(
Expand Down
Loading