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
88 changes: 44 additions & 44 deletions deepspeed/module_inject/auto_ep_layer.py
Original file line number Diff line number Diff line change
Expand Up @@ -641,8 +641,12 @@ def _deepep_route(self, tokens: torch.Tensor, ro: "RouterOutput") -> torch.Tenso
assert_dtype_supported(tokens.dtype)

# The configured worst-case capacity is identical across ranks, so
# buffer construction needs no rank-local decision or synchronization.
# buffer construction needs no rank-local resize decision.
if self._deepep_exchange is None:
# An externally initialized process group may still have a lazy
# NCCL communicator. DeepEP needs it before constructing its team;
# the removed split-count collective used to initialize it for us.
dist.barrier(group=self.ep_group, device_ids=[tokens.device.index])

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Badge Add the mandatory sign-off trailer

This is a one-parent, non-merge commit, but its commit message contains no Signed-off-by: trailer. The repository requires every non-merge commit to carry one, so the commit does not satisfy the project’s contribution requirements and must be recreated with --signoff.

AGENTS.md reference: AGENTS.md:L6-L8

Useful? React with 👍 / 👎.

self._deepep_exchange = DeepEPExchange(
ep_group=self.ep_group,
num_experts=self.num_experts,
Expand Down Expand Up @@ -696,6 +700,34 @@ def _deepep_route(self, tokens: torch.Tensor, ro: "RouterOutput") -> torch.Tenso

return deepep_combine(exchange, expert_output, handle)

def _finalize_output(self, output: torch.Tensor, x: torch.Tensor, hidden_states: torch.Tensor,
hdim: int) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
"""Apply the model-specific output tail shared by every communication backend."""
if self.moe_output_shape == "flat":
output = output.reshape(-1, hdim)
shared_expert_input = x
elif self.shared_experts_gate is not None:
shared_expert_input = x
else:
shared_expert_input = hidden_states

if self.shared_experts is not None:
shared_expert_output = self.shared_experts(shared_expert_input)
if self.shared_experts_gate is not None:
shared_expert_gate = torch.sigmoid(self.shared_experts_gate(shared_expert_input))
shared_expert_output = shared_expert_gate * shared_expert_output
if shared_expert_output.shape != output.shape:
shared_expert_output = shared_expert_output.reshape_as(output)
output = output + shared_expert_output

if self.return_router_logits:
logits = self._cached_router_logits
self._cached_router_logits = None
return output, logits

self._cached_router_logits = None
return output

def forward(
self,
hidden_states: torch.Tensor,
Expand Down Expand Up @@ -724,15 +756,16 @@ def forward(
with torch.no_grad():
self.tokens_per_expert.add_(ro.num_tokens_per_expert)

if self.ep_size > 1 and self.comm_backend == DEEPEP_BACKEND:
output = self._deepep_route(x, ro).reshape(bsz, seqlen, hdim)
return self._finalize_output(output, x, hidden_states, hdim)

# Reorder tokens into expert-contiguous order.
token_indices_sorted = torch.argsort(ro.selected_experts.view(-1), stable=True)
top_scores_sorted = ro.top_scores.view(-1)[token_indices_sorted]
expert_indices_sorted = ro.selected_experts.reshape(-1).index_select(0, token_indices_sorted)

folded_tp = self.folding_group_handles is not None and self.folding_group_handles.spec.tp_size > 1
# Set only where DeepEP's combine actually produced the output, since
# that decides whether the reduction below has already happened.
deepep_combined = False
restore_ctx = None
if folded_tp:
from deepspeed.moe.ep_tp_dispatch import (
Expand Down Expand Up @@ -810,31 +843,21 @@ def forward(
num_tokens_per_expert=ro.num_tokens_per_expert,
)

if self.comm_backend == DEEPEP_BACKEND:
expert_output = self._deepep_route(x, ro)
deepep_combined = True
else:
routed_input = _AllToAllV.apply(self.ep_group, routed_input, plan.input_splits, plan.output_splits)
routed_input = _AllToAllV.apply(self.ep_group, routed_input, plan.input_splits, plan.output_splits)

routed_input, perm_indices, aligned_counts, n_tokens = permute_by_local_expert(
routed_input, plan.local_counts_by_source)
expert_output = self.experts(routed_input, aligned_counts)
expert_output = unpermute_by_local_expert(expert_output, perm_indices, n_tokens)
routed_input, perm_indices, aligned_counts, n_tokens = permute_by_local_expert(
routed_input, plan.local_counts_by_source)
expert_output = self.experts(routed_input, aligned_counts)
expert_output = unpermute_by_local_expert(expert_output, perm_indices, n_tokens)

expert_output = _AllToAllV.apply(self.ep_group, expert_output, plan.output_splits, plan.input_splits)
expert_output = _AllToAllV.apply(self.ep_group, expert_output, plan.output_splits, plan.input_splits)

if folded_tp:
output = restore_combined(expert_output,
restore_ctx,
tp_group=self.tp_group,
validate_coverage=self.validate_folding_routing).reshape(bsz, seqlen, hdim)
self._last_folding_dispatch_counters = dispatch_counters(restore_ctx)
elif deepep_combined:
# DeepEP's combine already reduced over top-k and restored token
# order. This is keyed on the route having run rather than on the
# backend being selected: with ep_size == 1 the local path runs
# instead and still has one row per assignment to reduce.
output = expert_output.reshape(bsz, seqlen, hdim)
elif self.combine_impl == "fused_weighted_sum":
output = fused_token_ops.fused_weighted_restore(
expert_output,
Expand All @@ -854,30 +877,7 @@ def forward(
shape=(bsz, seqlen, hdim),
)

if self.moe_output_shape == "flat":
output = output.reshape(-1, hdim)
shared_expert_input = x
elif self.shared_experts_gate is not None:
shared_expert_input = x
else:
shared_expert_input = hidden_states

if self.shared_experts is not None:
shared_expert_output = self.shared_experts(shared_expert_input)
if self.shared_experts_gate is not None:
shared_expert_gate = torch.sigmoid(self.shared_experts_gate(shared_expert_input))
shared_expert_output = shared_expert_gate * shared_expert_output
if shared_expert_output.shape != output.shape:
shared_expert_output = shared_expert_output.reshape_as(output)
output = output + shared_expert_output

if self.return_router_logits:
logits = self._cached_router_logits
self._cached_router_logits = None
return output, logits

self._cached_router_logits = None
return output
return self._finalize_output(output, x, hidden_states, hdim)


class ReplacementSourceMap:
Expand Down
11 changes: 10 additions & 1 deletion docs/code-docs/source/autoep.rst
Original file line number Diff line number Diff line change
Expand Up @@ -98,7 +98,8 @@ that set nothing keep the existing path unchanged.
"autoep_size": 8,
"comm_backend": "deepep",
"comm_num_sm": 12,
"comm_qp_margin": 4
"comm_qp_margin": 4,
"comm_max_tokens_per_rank": 4096
}
}

Expand All @@ -112,6 +113,14 @@ that set nothing keep the existing path unchanged.
DeepEP buffer is sized statically and must use the same capacity on every
rank. A batch that exceeds it is an error.

For ``autoep_size > 1``, DeepEP receives the router output directly, bypassing
the collective backend's sorting, token expansion, and split-count exchange.
Shared experts and router-logit outputs retain the same behavior. The EP
communicator is initialized once before each layer's first DeepEP buffer is
constructed, including when the caller supplied a lazily initialized process
group. This initialization does not run on subsequent forwards. The standard
``comm`` and ``autoep_size=1`` paths are unchanged.

On 16 H100s across two nodes, replaying routing captured from real training,
DeepEP reduced payload AllToAll time from roughly 100 ms to 48 ms per step. A
full SFT step on Qwen3.5-MoE went from roughly 325 ms to 266 ms, a 1.2x speedup
Expand Down
114 changes: 107 additions & 7 deletions tests/unit/module_inject/test_auto_ep_comm.py
Original file line number Diff line number Diff line change
Expand Up @@ -308,6 +308,93 @@ def test_row_count_comes_from_the_handle_not_the_device(self):
self.assertNotIn("psum_num_recv_tokens_per_expert[-1]", source)


class TestDeepEPEarlyRoute(unittest.TestCase):
"""DeepEP must branch before collective-only token preparation."""

class Router(torch.nn.Module):

def __init__(self, output):
super().__init__()
self.output = output

def forward(self, *_args):
return self.output

@staticmethod
def layer(ep_size=2, comm_backend=DEEPEP_BACKEND, *, return_router_logits=False):
layer = object.__new__(auto_ep_layer.AutoEPMoELayer)
torch.nn.Module.__init__(layer)
scores = torch.tensor([[0.75, 0.25], [0.6, 0.4]])
experts = torch.tensor([[1, 0], [1, 0]], dtype=torch.long)
counts = torch.tensor([2, 2], dtype=torch.int32)
layer.router = TestDeepEPEarlyRoute.Router((scores, experts, counts))
layer.expert_bias = None
layer.register_buffer("tokens_per_expert", torch.zeros_like(counts, dtype=torch.float32))
layer.combine_impl = "weighted_sum"
layer._fused_combine_checked = False
layer.ep_size = ep_size
layer.comm_backend = comm_backend
# Keep the fixture valid when composed with opt-in async split planning.
layer.async_split_plan = False
layer._async_split_plan_pending = None
layer.top_k = 2
layer.num_experts = 2
layer.num_local_experts = 1
layer.folding_group_handles = None
layer.score_apply = "post"
layer.shared_experts = None
layer.shared_experts_gate = None
layer.moe_output_shape = "batched"
layer.return_router_logits = return_router_logits
layer._cached_router_logits = torch.randn(2, 2) if return_router_logits else None
layer._deepep_route = mock.Mock(return_value=torch.ones((2, 4)))
return layer

def test_deepep_bypasses_collective_preparation(self):
layer = self.layer()
hidden = torch.randn(1, 2, 4)

with mock.patch.object(auto_ep_layer.torch, "argsort", side_effect=AssertionError("argsort ran")), \
mock.patch.object(auto_ep_layer, "compute_split_plan",
side_effect=AssertionError("split plan ran")), \
mock.patch.object(auto_ep_layer, "compute_split_plan_from_expert_indices",
side_effect=AssertionError("folded split plan ran")), \
mock.patch.object(auto_ep_layer, "apply_scores_before_experts_if_enabled",
side_effect=AssertionError("score application ran")):
output = auto_ep_layer.AutoEPMoELayer.forward(layer, hidden)

self.assertEqual(tuple(output.shape), (1, 2, 4))
tokens, router_output = layer._deepep_route.call_args.args
self.assertEqual(tuple(tokens.shape), (2, 4))
self.assertTrue(torch.equal(tokens, hidden.reshape(2, 4)))
self.assertTrue(torch.equal(router_output.selected_experts, torch.tensor([[1, 0], [1, 0]])))
self.assertTrue(torch.equal(layer.tokens_per_expert, torch.tensor([2.0, 2.0])))

def test_standard_comm_and_ep1_keep_the_existing_path(self):
for ep_size, backend in ((2, COMM_BACKEND), (1, DEEPEP_BACKEND)):
with self.subTest(ep_size=ep_size, backend=backend):
layer = self.layer(ep_size=ep_size, comm_backend=backend)
with mock.patch.object(auto_ep_layer.torch, "argsort", side_effect=RuntimeError("sort reached")):
with self.assertRaisesRegex(RuntimeError, "sort reached"):
auto_ep_layer.AutoEPMoELayer.forward(layer, torch.randn(1, 2, 4))
layer._deepep_route.assert_not_called()

def test_shared_tail_and_router_logits_are_preserved(self):
layer = self.layer(return_router_logits=True)
expected_logits = layer._cached_router_logits
layer.moe_output_shape = "flat"
layer.shared_experts = mock.Mock(side_effect=lambda x: x * 2)
hidden = torch.arange(8, dtype=torch.float32).reshape(1, 2, 4)

output, logits = auto_ep_layer.AutoEPMoELayer.forward(layer, hidden)

self.assertEqual(tuple(output.shape), (2, 4))
self.assertTrue(torch.equal(output, torch.ones_like(output) + hidden.reshape(2, 4) * 2))
self.assertIs(logits, expected_logits)
self.assertIsNone(layer._cached_router_logits)
layer.shared_experts.assert_called_once()


class TestBufferLifecycle(unittest.TestCase):
"""The statically sized buffer never makes a rank-local resize decision."""

Expand Down Expand Up @@ -336,22 +423,35 @@ def test_outgrowing_the_buffer_names_the_remedy(self):
self.assertIn("600", message)

def test_the_configured_capacity_sizes_the_buffer(self):
layer = mock.Mock(_deepep_exchange=None, comm_max_tokens_per_rank=4096, comm_num_sm=12, comm_qp_margin=4)

built = mock.Mock(return_value=mock.Mock(num_max_tokens_per_rank=4096))
with mock.patch.object(auto_ep_layer, "DeepEPExchange", built), \
layer = TestDeepEPEarlyRoute.layer()
layer._deepep_exchange = None
layer.ep_group = object()
layer.hidden_size = 8
layer.comm_max_tokens_per_rank = 4096
layer.comm_num_sm = 12
layer.comm_qp_margin = 4
tokens = torch.ones((8, 8), dtype=torch.bfloat16)

def build_exchange(**kwargs):
barrier.assert_called_once_with(group=layer.ep_group, device_ids=[tokens.device.index])
return mock.Mock(num_max_tokens_per_rank=4096)

with mock.patch.object(auto_ep_layer.dist, "barrier") as barrier, \
mock.patch.object(auto_ep_layer, "DeepEPExchange", side_effect=build_exchange) as built, \
mock.patch.object(auto_ep_layer, "deepep_dispatch", side_effect=RuntimeError("stop here")), \
mock.patch.object(auto_ep_layer.dist, "all_reduce", lambda *a, **k: None):
router_output = auto_ep_layer.RouterOutput(
top_scores=torch.ones((8, 1)),
selected_experts=torch.zeros((8, 1), dtype=torch.long),
num_tokens_per_expert=torch.zeros(4, dtype=torch.long),
)
with self.assertRaises(RuntimeError):
auto_ep_layer.AutoEPMoELayer._deepep_route(layer, torch.ones((8, 8), dtype=torch.bfloat16),
router_output)
for _ in range(2):
with self.assertRaisesRegex(RuntimeError, "stop here"):
auto_ep_layer.AutoEPMoELayer._deepep_route(layer, tokens, router_output)

self.assertEqual(built.call_args.kwargs["num_max_tokens_per_rank"], 4096)
built.assert_called_once()
barrier.assert_called_once_with(group=layer.ep_group, device_ids=[tokens.device.index])


class TestAutogradSignatures(unittest.TestCase):
Expand Down
Loading
Loading