[NVIDIA][comm] Merge EP+MoE-TP post-experts all-reduces into one _TP reduction (#32963)
This commit is contained in:
@@ -1,9 +1,18 @@
|
||||
import contextlib
|
||||
import types
|
||||
import unittest
|
||||
from unittest.mock import patch
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.layers import communicator as comm
|
||||
from sglang.srt.layers.communicator import LayerCommunicator, ScatterMode
|
||||
from sglang.srt.layers.moe import (
|
||||
can_merge_post_experts_all_reduce,
|
||||
deferred_post_experts_all_reduce,
|
||||
post_experts_all_reduce,
|
||||
)
|
||||
from sglang.srt.layers.moe import utils as moe_utils
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
@@ -20,19 +29,218 @@ def _fake_communicator(mlp_mode=ScatterMode.TP_ATTN_FULL):
|
||||
)
|
||||
|
||||
|
||||
class TestFuseMlpAllReduceGate(CustomTestCase):
|
||||
"""Hybrid EP+TP must not fuse the post-experts all-reduce away.
|
||||
@contextlib.contextmanager
|
||||
def _recorded_all_reduces(called, *, moe_ep_size, moe_tp_size, moe_dp_size):
|
||||
"""Log which group each all-reduce helper reduces over, under a fixed topo."""
|
||||
|
||||
The fused residual+LN reduces over a single group, but with moe_ep_size > 1
|
||||
and moe_tp_size > 1 the post-experts reduction spans two disjoint groups
|
||||
(_MOE_EP then _MOE_TP) and should_skip_post_experts_all_reduce() drops both
|
||||
once fusion is published. The result is activations reduced over only half
|
||||
the peers -- wrong output, no crash. Observed as garbage completions on
|
||||
Qwen3-30B-A3B with --tp-size 4 --ep-size 2.
|
||||
def record(name):
|
||||
return lambda x: called.append(name) or x
|
||||
|
||||
with (
|
||||
patch(
|
||||
"sglang.srt.distributed.communication_op.tensor_model_parallel_all_reduce",
|
||||
side_effect=record("tp"),
|
||||
),
|
||||
patch(
|
||||
"sglang.srt.distributed.communication_op.moe_expert_parallel_all_reduce",
|
||||
side_effect=record("ep"),
|
||||
),
|
||||
patch(
|
||||
"sglang.srt.distributed.communication_op.moe_tensor_model_parallel_all_reduce",
|
||||
side_effect=record("moe_tp"),
|
||||
),
|
||||
get_parallel().override(
|
||||
moe_ep_size=moe_ep_size,
|
||||
moe_tp_size=moe_tp_size,
|
||||
moe_dp_size=moe_dp_size,
|
||||
tp_size=moe_ep_size * moe_tp_size * moe_dp_size,
|
||||
),
|
||||
):
|
||||
yield
|
||||
|
||||
|
||||
class TestPostExpertsAllReduceMerge(CustomTestCase):
|
||||
"""The two post-experts reductions collapse into one _TP reduction.
|
||||
|
||||
_MOE_EP and _MOE_TP are orthogonal subgroups of _TP, so with
|
||||
moe_dp_size == 1 reducing over each in turn equals one _TP reduction --
|
||||
one collective instead of two. With moe_dp_size > 1 they cover only part of
|
||||
_TP and merging would sum across DP replicas, which hold different tokens.
|
||||
"""
|
||||
|
||||
def _calls(self, *, moe_ep_size, moe_tp_size, moe_dp_size=1, skip=False):
|
||||
called = []
|
||||
with (
|
||||
patch.object(
|
||||
moe_utils, "should_skip_post_experts_all_reduce", return_value=skip
|
||||
),
|
||||
_recorded_all_reduces(
|
||||
called,
|
||||
moe_ep_size=moe_ep_size,
|
||||
moe_tp_size=moe_tp_size,
|
||||
moe_dp_size=moe_dp_size,
|
||||
),
|
||||
):
|
||||
post_experts_all_reduce(torch.zeros(2, 2))
|
||||
return called
|
||||
|
||||
def test_hybrid_issues_one_tp_reduction(self):
|
||||
self.assertEqual(self._calls(moe_ep_size=2, moe_tp_size=2), ["tp"])
|
||||
|
||||
def test_moe_dp_keeps_the_two_step_form(self):
|
||||
# Server args reject moe_ep_size > 1 together with moe_tp_size > 1 and
|
||||
# moe_dp_size > 1 (they force ep_size * moe_dp_size == tp_size), so this
|
||||
# pins the guard rather than a topology that can be launched today.
|
||||
self.assertEqual(
|
||||
self._calls(moe_ep_size=2, moe_tp_size=2, moe_dp_size=2), ["ep", "moe_tp"]
|
||||
)
|
||||
|
||||
def test_single_dimension_issues_one_reduction(self):
|
||||
self.assertEqual(self._calls(moe_ep_size=1, moe_tp_size=4), ["moe_tp"])
|
||||
self.assertEqual(self._calls(moe_ep_size=4, moe_tp_size=1), ["ep"])
|
||||
|
||||
def test_skipped_when_deferred_to_fusion(self):
|
||||
self.assertEqual(self._calls(moe_ep_size=2, moe_tp_size=2, skip=True), [])
|
||||
|
||||
|
||||
class TestDeferredPostExpertsAllReduce(CustomTestCase):
|
||||
"""The inline fallback must reduce over the same peers the fused kernel would.
|
||||
|
||||
_MOE_TP holds a single rank under pure EP, so reducing over it there is a
|
||||
no-op that drops the deferred reduction instead of performing it.
|
||||
"""
|
||||
|
||||
def _calls(self, *, moe_ep_size, moe_tp_size, moe_dp_size=1):
|
||||
called = []
|
||||
with _recorded_all_reduces(
|
||||
called,
|
||||
moe_ep_size=moe_ep_size,
|
||||
moe_tp_size=moe_tp_size,
|
||||
moe_dp_size=moe_dp_size,
|
||||
):
|
||||
deferred_post_experts_all_reduce(torch.zeros(2, 2))
|
||||
return called
|
||||
|
||||
def test_hybrid_reduces_over_tp(self):
|
||||
self.assertEqual(self._calls(moe_ep_size=2, moe_tp_size=2), ["tp"])
|
||||
|
||||
def test_pure_ep_reduces_over_ep(self):
|
||||
self.assertEqual(self._calls(moe_ep_size=4, moe_tp_size=1), ["ep"])
|
||||
|
||||
def test_pure_tp_reduces_over_moe_tp(self):
|
||||
self.assertEqual(self._calls(moe_ep_size=1, moe_tp_size=4), ["moe_tp"])
|
||||
|
||||
def test_moe_dp_reduces_over_moe_tp(self):
|
||||
self.assertEqual(
|
||||
self._calls(moe_ep_size=1, moe_tp_size=2, moe_dp_size=2), ["moe_tp"]
|
||||
)
|
||||
|
||||
|
||||
class TestCanMergePostExpertsAllReduce(CustomTestCase):
|
||||
def _can_merge(self, *, moe_ep_size, moe_tp_size, moe_dp_size=1):
|
||||
with get_parallel().override(
|
||||
moe_ep_size=moe_ep_size,
|
||||
moe_tp_size=moe_tp_size,
|
||||
moe_dp_size=moe_dp_size,
|
||||
tp_size=moe_ep_size * moe_tp_size * moe_dp_size,
|
||||
):
|
||||
return can_merge_post_experts_all_reduce()
|
||||
|
||||
def test_hybrid_ep_tp_merges(self):
|
||||
self.assertTrue(self._can_merge(moe_ep_size=2, moe_tp_size=2))
|
||||
|
||||
def test_single_dimension_does_not_merge(self):
|
||||
self.assertFalse(self._can_merge(moe_ep_size=1, moe_tp_size=4))
|
||||
self.assertFalse(self._can_merge(moe_ep_size=4, moe_tp_size=1))
|
||||
|
||||
|
||||
class TestResolveFusionGroup(CustomTestCase):
|
||||
"""EP2/MoE-TP2/DP1 (e.g. DeepSeek-V4-Flash with --tp-size 4 --ep-size 2) must
|
||||
resolve to the _TP group with world_size=4 and the TP rank."""
|
||||
|
||||
def _resolve(self, *, moe_ep_size, moe_tp_size, moe_dp_size=1, tp_rank=0):
|
||||
from sglang.srt.layers.flashinfer_comm_fusion import (
|
||||
resolve_fusion_group,
|
||||
resolve_fusion_world_size,
|
||||
)
|
||||
|
||||
fake_tp_group = MagicMock(name="tp_group")
|
||||
fake_ep_group = MagicMock(name="ep_group")
|
||||
fake_moe_tp_group = MagicMock(name="moe_tp_group")
|
||||
tp_size = moe_ep_size * moe_tp_size * moe_dp_size
|
||||
|
||||
with (
|
||||
get_parallel().override(
|
||||
moe_ep_size=moe_ep_size,
|
||||
moe_tp_size=moe_tp_size,
|
||||
moe_dp_size=moe_dp_size,
|
||||
tp_size=tp_size,
|
||||
tp_rank=tp_rank,
|
||||
moe_ep_rank=tp_rank % moe_ep_size,
|
||||
moe_tp_rank=tp_rank % moe_tp_size,
|
||||
),
|
||||
patch(
|
||||
"sglang.srt.layers.flashinfer_comm_fusion.get_tp_group",
|
||||
return_value=fake_tp_group,
|
||||
),
|
||||
patch(
|
||||
"sglang.srt.layers.flashinfer_comm_fusion.get_moe_ep_group",
|
||||
return_value=fake_ep_group,
|
||||
),
|
||||
patch(
|
||||
"sglang.srt.layers.flashinfer_comm_fusion.get_moe_tp_group",
|
||||
return_value=fake_moe_tp_group,
|
||||
),
|
||||
):
|
||||
ws = resolve_fusion_world_size(use_attn_tp_group=False)
|
||||
group_tuple = resolve_fusion_group(use_attn_tp_group=False)
|
||||
return ws, group_tuple, (fake_tp_group, fake_ep_group, fake_moe_tp_group)
|
||||
|
||||
def test_hybrid_ep2_tp2_dp1_resolves_to_tp_ws4(self):
|
||||
# EP2/MoE-TP2/DP1 (DeepSeek-V4-Flash on 4 GPUs): workspace must sit on
|
||||
# _TP (ws=4) so the fused kernel reduces over all 4 peers.
|
||||
ws, (size, rank, group), (tp_grp, ep_grp, moe_tp_grp) = self._resolve(
|
||||
moe_ep_size=2, moe_tp_size=2, moe_dp_size=1, tp_rank=3
|
||||
)
|
||||
self.assertEqual(ws, 4)
|
||||
self.assertEqual(size, 4)
|
||||
self.assertEqual(rank, 3)
|
||||
self.assertIs(group, tp_grp)
|
||||
|
||||
def test_pure_ep_resolves_to_ep_group(self):
|
||||
ws, (size, rank, group), (tp_grp, ep_grp, moe_tp_grp) = self._resolve(
|
||||
moe_ep_size=4, moe_tp_size=1, moe_dp_size=1, tp_rank=2
|
||||
)
|
||||
self.assertEqual(ws, 4)
|
||||
self.assertEqual(size, 4)
|
||||
self.assertIs(group, ep_grp)
|
||||
|
||||
def test_pure_tp_resolves_to_moe_tp_group(self):
|
||||
ws, (size, rank, group), (tp_grp, ep_grp, moe_tp_grp) = self._resolve(
|
||||
moe_ep_size=1, moe_tp_size=4, moe_dp_size=1, tp_rank=1
|
||||
)
|
||||
self.assertEqual(ws, 4)
|
||||
self.assertEqual(size, 4)
|
||||
self.assertIs(group, moe_tp_grp)
|
||||
|
||||
|
||||
class TestFuseMlpAllReduceGate(CustomTestCase):
|
||||
"""Fusion is allowed only when one group covers the whole reduction.
|
||||
|
||||
The fused residual+LN reduces over a single group. Hybrid EP+TP produces two
|
||||
reductions over disjoint groups; merging collapses them to one _TP reduction
|
||||
that the fused kernel can absorb. When merging does not apply
|
||||
(moe_dp_size > 1) there is no such group and fusion must stay off --
|
||||
otherwise the fused reduce covers half the peers and silently under-reduces.
|
||||
"""
|
||||
|
||||
def _should_fuse(
|
||||
self, *, moe_ep_size, moe_tp_size, mlp_mode=ScatterMode.TP_ATTN_FULL
|
||||
self,
|
||||
*,
|
||||
moe_ep_size,
|
||||
moe_tp_size,
|
||||
moe_dp_size=1,
|
||||
mlp_mode=ScatterMode.TP_ATTN_FULL,
|
||||
):
|
||||
forward_batch = types.SimpleNamespace(
|
||||
input_ids=types.SimpleNamespace(shape=(8,))
|
||||
@@ -46,15 +254,24 @@ class TestFuseMlpAllReduceGate(CustomTestCase):
|
||||
return_value=types.SimpleNamespace(input_scattered=False),
|
||||
),
|
||||
get_parallel().override(
|
||||
moe_ep_size=moe_ep_size, moe_tp_size=moe_tp_size, tp_size=4
|
||||
moe_ep_size=moe_ep_size,
|
||||
moe_tp_size=moe_tp_size,
|
||||
moe_dp_size=moe_dp_size,
|
||||
tp_size=moe_ep_size * moe_tp_size * moe_dp_size,
|
||||
),
|
||||
):
|
||||
return LayerCommunicator.should_fuse_mlp_allreduce_with_next_layer(
|
||||
_fake_communicator(mlp_mode), forward_batch
|
||||
)
|
||||
|
||||
def test_hybrid_ep_tp_does_not_fuse(self):
|
||||
self.assertFalse(self._should_fuse(moe_ep_size=2, moe_tp_size=2))
|
||||
def test_hybrid_ep_tp_fuses_when_mergeable(self):
|
||||
self.assertTrue(self._should_fuse(moe_ep_size=2, moe_tp_size=2))
|
||||
|
||||
def test_hybrid_ep_tp_does_not_fuse_when_moe_dp_blocks_the_merge(self):
|
||||
# Same caveat as test_moe_dp_keeps_the_two_step_form: unreachable today,
|
||||
# kept so a future relaxation cannot silently re-enable fusion over a
|
||||
# reduction that no single group covers.
|
||||
self.assertFalse(self._should_fuse(moe_ep_size=2, moe_tp_size=2, moe_dp_size=2))
|
||||
|
||||
def test_pure_tp_still_fuses(self):
|
||||
self.assertTrue(self._should_fuse(moe_ep_size=1, moe_tp_size=4))
|
||||
|
||||
Reference in New Issue
Block a user