[AMD] Fix the dsv4 performance of MoE issue. (#27858)

Co-authored-by: wunhuang <wunhuang@amd.com>
This commit is contained in:
kk
2026-06-11 01:06:54 -07:00
committed by GitHub
co-authored by wunhuang
parent 2a51479a91
commit b8376aebd0
+19 -11
View File
@@ -127,9 +127,8 @@ def _require_fp4_dtype():
if _use_aiter or _use_hip_int4:
from aiter.ops.shuffle import (
shuffle_scale_a16w4,
shuffle_scale,
shuffle_weight,
shuffle_weight_a16w4,
)
if _use_aiter:
@@ -1148,6 +1147,7 @@ class Fp8MoEMethod(FusedMoEMethodBase):
def process_weights_after_loading_block_quant(self, layer: Module) -> None:
# AMD FP4 experts: use aiter's native MXFP4 MoE path
if _use_aiter and self.is_fp4_expert:
gu_intv = envs.SGLANG_USE_AITER_MOE_GU_ITLV.get()
fp4_weight_dtype = _require_fp4_dtype()
# CK FP4 MoE kernel requires K_packed divisible by 128
@@ -1160,6 +1160,11 @@ class Fp8MoEMethod(FusedMoEMethodBase):
padded_inter = (
(inter_per_part + fp4_k_align - 1) // fp4_k_align * fp4_k_align
)
# Record the padding so fused_moe is told the real intermediate size
# (aiter fused_moe needs intermediate_pad = padded - real; ATOM passes
# 128, SGLang previously defaulted to 0 -> computed the padded region).
layer.intermediate_pad = padded_inter - inter_per_part
layer.hidden_pad = 0
if padded_inter != inter_per_part:
pad_amount = padded_inter - inter_per_part
fp4_block_k = 32
@@ -1226,23 +1231,24 @@ class Fp8MoEMethod(FusedMoEMethodBase):
for scale_name in ("w13_weight_scale_inv", "w2_weight_scale_inv"):
scale = getattr(layer, scale_name)
num_experts, num_rows, _ = scale.shape
# a8w4: aiter flydsl scale layout
is_w13_scale = scale_name == "w13_weight_scale_inv"
scale.data = shuffle_scale_a16w4(
scale.view(num_experts * num_rows, -1), num_experts, is_w13_scale
)
scale_2d = scale.reshape(-1, scale.shape[-1])
scale.data = shuffle_scale(scale_2d, num_experts, gu_intv, is_w13_scale)
layer.w13_weight.data = layer.w13_weight.data.view(fp4_weight_dtype)
layer.w2_weight.data = layer.w2_weight.data.view(fp4_weight_dtype)
is_shuffled = _is_shuffle_moe_mxfp4
if is_shuffled:
# a8w4: aiter flydsl weight layout
layer.w13_weight.data = shuffle_weight_a16w4(
layer.w13_weight.contiguous(), 16, True
layer.w13_weight.data = shuffle_weight(
layer.w13_weight,
is_guinterleave=gu_intv,
gate_up=True,
)
layer.w2_weight.data = shuffle_weight_a16w4(
layer.w2_weight.contiguous(), 16, False
layer.w2_weight.data = shuffle_weight(
layer.w2_weight,
is_guinterleave=gu_intv,
gate_up=False,
)
layer.w13_weight.is_shuffled = is_shuffled
layer.w2_weight.is_shuffled = is_shuffled
@@ -2129,6 +2135,8 @@ class Fp8MoEMethod(FusedMoEMethodBase):
w2_scale=w2_scale,
expert_mask=layer.dispatcher.expert_mask_gpu if _use_aiter else None,
swiglu_limit=self.moe_runner_config.swiglu_limit or 0.0,
hidden_pad=getattr(layer, "hidden_pad", 0),
intermediate_pad=getattr(layer, "intermediate_pad", 0),
)