From 2c64b4c1cc18bc25ff98f1835993d7545a8b394c Mon Sep 17 00:00:00 2001 From: Hongxia Yang <62075498+hongxiayang@users.noreply.github.com> Date: Wed, 8 Jul 2026 03:26:17 -0400 Subject: [PATCH] [ROCm] fixed aiter master flag and expert parallelism compatibility on minimax-m3-mxfp8 (#47158) Signed-off-by: Hongxia Yang --- tests/kernels/test_minimax_m3_amd_ops.py | 94 +++++++++++++++++++ .../fused_moe/experts/aiter_mxfp8_moe.py | 25 ++--- 2 files changed, 107 insertions(+), 12 deletions(-) diff --git a/tests/kernels/test_minimax_m3_amd_ops.py b/tests/kernels/test_minimax_m3_amd_ops.py index bbd3e923419..60fbf53d4ad 100644 --- a/tests/kernels/test_minimax_m3_amd_ops.py +++ b/tests/kernels/test_minimax_m3_amd_ops.py @@ -422,3 +422,97 @@ def test_mxfp8_linear_emulation_bf16_at_load( out = kernel.apply_weights(layer, x) assert out.dtype == act_dtype # dtype-match preserved (no tl.dot/F.linear crash) assert _relerr(out.float(), out_ref.float()) < 2e-2 + + +# ── EP expert_mask handling for the FlyDSL (AITER_MXFP8) MoE ──────────────── +# Regression for the EP + aiter-master-switch interaction: under expert +# parallelism ``RoutedExperts.expert_map`` hands the experts either the 0/1 +# ``expert_mask`` (aiter master ON, ``rocm_aiter_fmoe_enabled``) or vLLM's -1 +# index map (master OFF). ``AiterMxfp8Experts.apply`` must forward the right 0/1 +# mask to aiter in BOTH cases. The old code always rebuilt the mask via +# ``(expert_map >= 0)``; on the already-0/1 mask that collapses to all-ones (no +# experts masked out) and EP output becomes garbage (no accuracy). +def _capture_expert_mask(expert_map, *, rocm_aiter_fmoe_enabled, global_num_experts): + """Drive the real ``AiterMxfp8Experts.apply`` mask branch and capture the + ``expert_mask`` it forwards to ``rocm_aiter_ops.fused_moe``.""" + from types import SimpleNamespace + from unittest import mock + + from vllm._aiter_ops import rocm_aiter_ops + from vllm.model_executor.layers.fused_moe.experts.aiter_mxfp8_moe import ( + AiterMxfp8Experts, + ) + + experts = object.__new__(AiterMxfp8Experts) # bypass heavy __init__ + experts.moe_config = SimpleNamespace( + rocm_aiter_fmoe_enabled=rocm_aiter_fmoe_enabled + ) + experts.quant_config = SimpleNamespace(gemm1_clamp_limit=None) + experts.w1_scale_val = None + experts.w2_scale_val = None + + captured = {} + + def _fake_fused_moe(hidden_states, w1, w2, tw, ti, *, expert_mask, **kw): + captured["expert_mask"] = expert_mask + return torch.zeros_like(hidden_states) + + w1 = torch.zeros(1, device=DEVICE) + w2 = torch.zeros(1, device=DEVICE) + out = torch.zeros(4, 8, device=DEVICE, dtype=torch.bfloat16) + hidden = torch.zeros(4, 8, device=DEVICE, dtype=torch.bfloat16) + tw = torch.ones(4, 2, device=DEVICE) + ti = torch.zeros(4, 2, dtype=torch.int32, device=DEVICE) + + with mock.patch.object(rocm_aiter_ops, "fused_moe", side_effect=_fake_fused_moe): + experts.apply( + output=out, + hidden_states=hidden, + w1=w1, + w2=w2, + topk_weights=tw, + topk_ids=ti, + activation=None, + global_num_experts=global_num_experts, + expert_map=expert_map, + a1q_scale=None, + a2_scale=None, + workspace13=None, + workspace2=None, + expert_tokens_meta=None, + apply_router_weight_on_input=False, + ) + return captured["expert_mask"] + + +@pytest.mark.skipif(not current_platform.is_rocm(), reason="ROCm only") +def test_aiter_mxfp8_ep_expert_mask_both_master_modes(): + """Both aiter-master forms must yield the SAME correct 0/1 aiter mask; + guards the EP+master regression (mask must not collapse to all-ones).""" + from vllm.model_executor.layers.fused_moe.expert_map_manager import ( + determine_expert_map, + ) + + E, ep_size, ep_rank = 8, 2, 0 # rank owns global experts 0..3 + # master OFF: vLLM's -1 index map + _, idx_map, _ = determine_expert_map(ep_size, ep_rank, E, return_expert_mask=False) + # master ON: 0/1 mask (+ trailing sentinel) that RoutedExperts forwards + _, _, ep_mask = determine_expert_map(ep_size, ep_rank, E, return_expert_mask=True) + + idx_map = idx_map.to(DEVICE) + ep_mask = ep_mask.to(DEVICE) + + # Expected aiter expert_mask: 0/1 over global ids + trailing sentinel slot. + expected = torch.tensor([1, 1, 1, 1, 0, 0, 0, 0, 0], dtype=torch.int32) + + got_off = _capture_expert_mask( + idx_map, rocm_aiter_fmoe_enabled=False, global_num_experts=E + ) + got_on = _capture_expert_mask( + ep_mask, rocm_aiter_fmoe_enabled=True, global_num_experts=E + ) + + assert torch.equal(got_off.cpu().to(torch.int32), expected) + # master ON forwards the prebuilt mask unchanged (NOT collapsed to all-ones) + assert torch.equal(got_on.cpu().to(torch.int32), ep_mask.cpu().to(torch.int32)) + assert got_on.sum().item() == 4 # exactly the 4 local experts, not all 9 diff --git a/vllm/model_executor/layers/fused_moe/experts/aiter_mxfp8_moe.py b/vllm/model_executor/layers/fused_moe/experts/aiter_mxfp8_moe.py index 3cbab0a0d54..c5330f3b438 100644 --- a/vllm/model_executor/layers/fused_moe/experts/aiter_mxfp8_moe.py +++ b/vllm/model_executor/layers/fused_moe/experts/aiter_mxfp8_moe.py @@ -78,9 +78,6 @@ class AiterMxfp8Experts(Mxfp8TritonExpertsBase): @staticmethod def _supports_parallel_config(moe_parallel_config) -> bool: - # Both TP (expert_map=None) and EP are supported: apply() forwards the - # expert_map as aiter's ``expert_mask`` (the per-rank local-expert - # selection), mirroring the native rocm_aiter_moe path. return True @staticmethod @@ -129,17 +126,21 @@ class AiterMxfp8Experts(Mxfp8TritonExpertsBase): limit = self.quant_config.gemm1_clamp_limit swiglu_limit = 0.0 if limit is None else float(limit) - # Under EP, aiter expects ``expert_mask`` as a 0/1 *local-expert* mask - # over global ids with a trailing fake-expert sentinel slot - # (shape ``[global_num_experts + 1]``), NOT vLLM's expert_map (a - # global->local index map with -1 for non-local). Convert it; aiter - # derives the global->local compaction from the mask itself. ``None`` - # under pure TP. - if expert_map is not None: + # Under EP, aiter expects ``expert_mask``: a 0/1 *local-expert* mask over + # global ids with a trailing fake-expert sentinel slot (shape + # ``[global_num_experts + 1]``), from which it derives the global->local + # compaction. What ``RoutedExperts.expert_map`` hands us depends on the + # aiter master switch (``rocm_aiter_fmoe_enabled``). + # Branching on the (static) master flag — not the tensor contents — + # keeps this HIP-graph/torch.compile safe (no data-dependent sync). + # ``None`` under pure TP. + if expert_map is None: + expert_mask = None + elif self.moe_config.rocm_aiter_fmoe_enabled: + expert_mask = expert_map + else: local_mask = (expert_map >= 0).to(torch.int32) expert_mask = torch.cat([local_mask, local_mask.new_zeros(1)]) - else: - expert_mask = None # Route through the graph-safe ``rocm_aiter_fused_moe`` custom op so the # call is captured under HIP graphs / torch.compile (a direct