From 79df454fae13131eab657db7fc0a23f71d88e57a Mon Sep 17 00:00:00 2001 From: ARAVINDHAN T Date: Tue, 6 Oct 2026 14:54:54 -0700 Subject: [PATCH] Add an opt-in route_norm flag for the legacy MoE gate top1gating leaves the chosen expert's softmax probability in the combine weight. top2gating and topkgating rescale the kept weights so they sum to 1. Making those agree changes the MoE output for existing checkpoints, so route_norm defaults to None and each gate keeps its historical rule. True or False overrides all three gates. Auxiliary loss is unchanged. Signed-off-by: ARAVINDHAN T --- deepspeed/moe/layer.py | 11 +- deepspeed/moe/sharded_moe.py | 102 ++++++++++---- docs/_tutorials/mixture-of-experts.md | 2 + tests/unit/v1/moe/test_moe.py | 186 +++++++++++++++++++++++++- 4 files changed, 274 insertions(+), 27 deletions(-) diff --git a/deepspeed/moe/layer.py b/deepspeed/moe/layer.py index 6777788ab885..393db5daecff 100644 --- a/deepspeed/moe/layer.py +++ b/deepspeed/moe/layer.py @@ -33,6 +33,12 @@ class MoE(nn.Module): use_tutel (bool, optional): default=False, whether to use Tutel optimizations (if installed). enable_expert_tensor_parallelism (bool, optional): default=False, whether to use tensor parallelism for experts top2_2nd_expert_sampling (bool, optional): default=True, whether to perform sampling for 2nd expert + route_norm (bool, optional): default=None. None keeps the historical combine weights: k=1 leaves the chosen + expert's softmax probability, and k>=2 rescales the kept weights so they sum to 1. True rescales for every k. + False leaves the raw softmax probabilities for every k. True or False changes the MoE output relative to a + checkpoint trained with the other setting. The auxiliary load-balancing loss is computed before this step + and does not change. This is the legacy gate. It is not expert_parallel.route_norm, which configures the + AutoEP router. """ def __init__(self, @@ -50,7 +56,8 @@ def __init__(self, use_rts: bool = True, use_tutel: bool = False, enable_expert_tensor_parallelism: bool = False, - top2_2nd_expert_sampling: bool = True) -> None: + top2_2nd_expert_sampling: bool = True, + route_norm: Optional[bool] = None) -> None: super(MoE, self).__init__() @@ -72,7 +79,7 @@ def __init__(self, experts = Experts(expert, self.num_local_experts, self.expert_group_name) self.deepspeed_moe = MOELayer(TopKGate(hidden_size, num_experts, k, capacity_factor, eval_capacity_factor, min_capacity, noisy_gate_policy, drop_tokens, use_rts, None, - top2_2nd_expert_sampling), + top2_2nd_expert_sampling, route_norm), experts, self.expert_group_name, self.ep_size, diff --git a/deepspeed/moe/sharded_moe.py b/deepspeed/moe/sharded_moe.py index 1744f5edb3c1..a49f4b5579ce 100644 --- a/deepspeed/moe/sharded_moe.py +++ b/deepspeed/moe/sharded_moe.py @@ -232,6 +232,13 @@ def _sparse_decode(expert_output: Tensor, slots: Tensor, gates: Tensor, num_toke return combined.to(expert_output.dtype) +def _route_norm_enabled(route_norm: Optional[bool], legacy_default: bool) -> bool: + """None keeps that gate's historical combine-weight rule. True or False overrides every gate.""" + if route_norm is None: + return legacy_default + return route_norm + + def top1gating(logits: Tensor, capacity_factor: float, min_capacity: int, @@ -241,8 +248,13 @@ def top1gating(logits: Tensor, use_rts: bool = True, ep_group: Union[torch.distributed.ProcessGroup, None] = None, sparse_routes: bool = False, - use_tutel: bool = False) -> Tuple[Tensor, Tensor, Tensor, Tensor]: - """Implements Top1Gating on logits.""" + use_tutel: bool = False, + route_norm: Optional[bool] = None) -> Tuple[Tensor, Tensor, Tensor, Tensor]: + """Implements Top1Gating on logits. + + route_norm=None keeps the historical weights: the chosen expert's softmax probability, with no + second normalization. The auxiliary loss is computed from those probabilities either way. + """ if noisy_gate_policy == 'RSample': logits_w_noise = logits + gumbel_rsample(logits.shape, device=logits.device) # everything is in fp32 in this function @@ -319,6 +331,11 @@ def top1gating(logits: Tensor, if sparse_routes: gates1_s = (gates * mask1).sum(dim=1) + # Historical top-1 does not renormalize. An explicit True makes a kept token's weight 1. + # A dropped token sums to 0, and 0 / eps stays 0. + if _route_norm_enabled(route_norm, False): + denom_s = torch.clamp(gates1_s, min=torch.finfo(gates1_s.dtype).eps) + gates1_s = gates1_s / denom_s locations1_s = torch.sum(locations1 * mask1, dim=1) return (l_aux, capacity, num_experts, indices1_s.to(torch.int32).unsqueeze(0), locations1_s.to(torch.int32).unsqueeze(0), gates1_s.unsqueeze(0), exp_counts) @@ -326,9 +343,13 @@ def top1gating(logits: Tensor, # Store the capacity location for each token locations1_s = torch.sum(locations1 * mask1, dim=1) - # Normalize gate probabilities + # Keep the assigned expert's softmax probability. route_norm rescales a kept row to 1. mask1_float = mask1.float() gates = gates * mask1_float + if _route_norm_enabled(route_norm, False): + denom_s = torch.sum(gates, dim=1, keepdim=True) + denom_s = torch.clamp(denom_s, min=torch.finfo(denom_s.dtype).eps) + gates = gates / denom_s locations1_sc = _one_hot_to_float(locations1_s, capacity) combine_weights = einsum("se,sc->sec", gates, locations1_sc) @@ -344,8 +365,13 @@ def top2gating(logits: Tensor, drop_tokens: bool = True, ep_group: Union[torch.distributed.ProcessGroup, None] = None, top2_2nd_expert_sampling: bool = True, - sparse_routes: bool = False) -> Tuple[Tensor, Tensor, Tensor, Tensor]: - """Implements Top2Gating on logits.""" + sparse_routes: bool = False, + route_norm: Optional[bool] = None) -> Tuple[Tensor, Tensor, Tensor, Tensor]: + """Implements Top2Gating on logits. + + route_norm=None keeps the historical weights: the two kept probabilities are rescaled to sum to 1. + The auxiliary loss is computed before that rescaling. + """ # everything is in fp32 in this function gates = F.softmax(logits, dim=1) @@ -399,16 +425,17 @@ def top2gating(logits: Tensor, locations1_s = torch.sum(locations1 * mask1, dim=1) locations2_s = torch.sum(locations2 * mask2, dim=1) - # Normalize gate probabilities + # Normalize gate probabilities. route_norm=None keeps this rescaling; False leaves the raw softmax values. mask1_float = mask1.float() mask2_float = mask2.float() gates1_s = einsum("se,se->s", gates, mask1_float) gates2_s = einsum("se,se->s", gates, mask2_float) - denom_s = gates1_s + gates2_s - # Avoid divide-by-zero - denom_s = torch.clamp(denom_s, min=torch.finfo(denom_s.dtype).eps) - gates1_s /= denom_s - gates2_s /= denom_s + if _route_norm_enabled(route_norm, True): + denom_s = gates1_s + gates2_s + # Avoid divide-by-zero. A dropped token is 0 / eps and stays 0. + denom_s = torch.clamp(denom_s, min=torch.finfo(denom_s.dtype).eps) + gates1_s /= denom_s + gates2_s /= denom_s if sparse_routes: # Routes evicted by the capacity limit are flagged with a negative expert index. @@ -440,8 +467,13 @@ def topkgating( ep_group: Union[torch.distributed.ProcessGroup, None] = None, drop_policy: str = "probs", sparse_routes: bool = False, + route_norm: Optional[bool] = None, ) -> Tuple[Tensor, Tensor, Tensor, Tensor]: - """Implements TopKGating on logits.""" + """Implements TopKGating on logits. + + route_norm=None keeps the historical weights: the kept probabilities are rescaled to sum to 1. + The auxiliary loss is computed before that rescaling. + """ # everything is in fp32 in this function # gating decisions @@ -494,11 +526,12 @@ def topkgating( capacity = new_capacity locations = torch.cumsum(mask, dim=0) - 1 - # normalize gates + # normalize gates. route_norm=None keeps this rescaling; False leaves the raw softmax values. gates_masked = gates * mask - gates_s = torch.sum(gates_masked, dim=-1, keepdim=True) - denom_s = torch.clamp(gates_s, min=torch.finfo(gates_masked.dtype).eps) - gates_masked = gates_masked / denom_s + if _route_norm_enabled(route_norm, True): + gates_s = torch.sum(gates_masked, dim=-1, keepdim=True) + denom_s = torch.clamp(gates_s, min=torch.finfo(gates_masked.dtype).eps) + gates_masked = gates_masked / denom_s if locations is None: raise ValueError(f"Locations is not set: {locations}") @@ -539,6 +572,10 @@ class TopKGate(Module): size of model embedding dimension num_experts (int): number of experts in model + route_norm (bool, optional): + None keeps the historical combine weights (no renorm for k=1, renorm for k>=2). + True renormalizes the kept weights for every k. False leaves the raw softmax values. + The auxiliary loss does not use this flag. """ wg: torch.nn.Linear @@ -554,7 +591,8 @@ def __init__(self, drop_tokens: bool = True, use_rts: bool = True, ep_group: Union[torch.distributed.ProcessGroup, None] = None, - top2_2nd_expert_sampling: bool = True) -> None: + top2_2nd_expert_sampling: bool = True, + route_norm: Optional[bool] = None) -> None: super().__init__() self.wg = torch.nn.Linear(model_dim, num_experts, bias=False) @@ -570,6 +608,8 @@ def __init__(self, self.drop_tokens = drop_tokens self.use_rts = use_rts self.top2_2nd_expert_sampling = top2_2nd_expert_sampling + # None is forwarded so each gating function keeps its own historical default. + self.route_norm = route_norm def _set_ep_group(self, ep_group): assert self.ep_group is None, 'Attempting to override an existing ep_group' @@ -591,14 +631,27 @@ def forward(self, logits = torch.nn.functional.linear(input_fp32, weight=self.wg.weight.float(), bias=None) if self.k == 1: - gate_output = top1gating(logits, self.capacity_factor if self.training else self.eval_capacity_factor, - self.min_capacity, used_token, self.noisy_gate_policy if self.training else None, - self.drop_tokens, self.use_rts, self.ep_group, sparse_routes, use_tutel) + gate_output = top1gating(logits, + self.capacity_factor if self.training else self.eval_capacity_factor, + self.min_capacity, + used_token, + self.noisy_gate_policy if self.training else None, + self.drop_tokens, + self.use_rts, + self.ep_group, + sparse_routes, + use_tutel, + route_norm=self.route_norm) elif self.k == 2: - gate_output = top2gating(logits, self.capacity_factor if self.training else self.eval_capacity_factor, - self.min_capacity, self.drop_tokens, self.ep_group, self.top2_2nd_expert_sampling, - sparse_routes) + gate_output = top2gating(logits, + self.capacity_factor if self.training else self.eval_capacity_factor, + self.min_capacity, + self.drop_tokens, + self.ep_group, + self.top2_2nd_expert_sampling, + sparse_routes, + route_norm=self.route_norm) else: gate_output = topkgating(logits, self.k, @@ -606,7 +659,8 @@ def forward(self, self.min_capacity, self.drop_tokens, self.ep_group, - sparse_routes=sparse_routes) + sparse_routes=sparse_routes, + route_norm=self.route_norm) if self.wall_clock_breakdown: self.timers(TOPK_GATE_TIMER).stop() diff --git a/docs/_tutorials/mixture-of-experts.md b/docs/_tutorials/mixture-of-experts.md index d7e1365ad36f..a47d6c6f1734 100644 --- a/docs/_tutorials/mixture-of-experts.md +++ b/docs/_tutorials/mixture-of-experts.md @@ -63,6 +63,8 @@ Updated with MoE Layers self.fc4 = nn.Linear(84, 10) ``` +`MoE(..., route_norm=...)` is optional and defaults to `None`. That default keeps the historical combine weights: `k=1` uses the chosen expert's softmax probability, and `k>=2` rescales the kept weights so they sum to 1. `True` rescales for every `k`, and `False` leaves the raw softmax probabilities for every `k`. Either explicit value changes the layer output relative to a checkpoint trained with the other setting. The auxiliary load-balancing loss is unchanged. This constructor argument is the legacy gate. It is not the AutoEP config key `expert_parallel.route_norm`. + ### Pyramid-Residual MoE Recently, we proposed a novel [Pyramid-Residual MoE](https://arxiv.org/abs/2201.05596) (PR-MoE) model architecture. To create such an MoE model, the users need to do two additional things: diff --git a/tests/unit/v1/moe/test_moe.py b/tests/unit/v1/moe/test_moe.py index 7651b90e0ce6..f30583011d12 100644 --- a/tests/unit/v1/moe/test_moe.py +++ b/tests/unit/v1/moe/test_moe.py @@ -14,7 +14,7 @@ import deepspeed.moe.sharded_moe as sharded_moe from deepspeed import get_accelerator from deepspeed.moe.layer import MoE -from deepspeed.moe.sharded_moe import (top1gating, top2gating, topkgating, _route_slots, _sparse_encode, +from deepspeed.moe.sharded_moe import (top1gating, top2gating, topkgating, TopKGate, _route_slots, _sparse_encode, _sparse_decode) from deepspeed.moe.utils import split_params_into_different_moe_groups_for_optimizer, is_moe_param from deepspeed.runtime.fp16.fused_optimizer import FP16_Optimizer @@ -526,6 +526,190 @@ def test_top1gating_preserves_tensor_parallel_capacity(): assert dispatch_mask.shape[-1] == 6 +def _routes_equal(left, right): + assert len(left) == len(right) + for lhs, rhs in zip(left, right): + if torch.is_tensor(lhs) or torch.is_tensor(rhs): + assert torch.is_tensor(lhs) and torch.is_tensor(rhs) + assert lhs.shape == rhs.shape + assert torch.equal(lhs, rhs) + else: + assert lhs == rhs + + +def _fixed_logits(): + return torch.tensor([[2.0, 0.0, -1.0, 0.5], [0.2, 3.0, 0.1, -0.4], [1.0, 1.1, 0.2, 0.0], [-0.5, 0.3, 2.5, 0.1]]) + + +_OMIT_ROUTE_NORM = object() + + +def _gate_call(logits, k, sparse_routes, route_norm=_OMIT_ROUTE_NORM): + # Omitting the flag must hit the real default argument, not an explicit None. + kwargs = {} if route_norm is _OMIT_ROUTE_NORM else {"route_norm": route_norm} + if k == 1: + return top1gating(logits, 1.0, 0, drop_tokens=False, use_rts=False, sparse_routes=sparse_routes, **kwargs) + if k == 2: + return top2gating(logits, + 1.0, + 0, + drop_tokens=False, + top2_2nd_expert_sampling=False, + sparse_routes=sparse_routes, + **kwargs) + return topkgating(logits, k, 1.0, 0, drop_tokens=False, sparse_routes=sparse_routes, **kwargs) + + +def _legacy_default(k): + # top-1 historically leaves the softmax probability. top-2 and top-k rescale the kept weights to 1. + return False if k == 1 else True + + +def test_route_norm_default_matches_historical_gates(): + logits = _fixed_logits() + for k in (1, 2, 3): + for sparse_routes in (False, True): + omitted = _gate_call(logits, k, sparse_routes) + unset = _gate_call(logits, k, sparse_routes, route_norm=None) + explicit = _gate_call(logits, k, sparse_routes, route_norm=_legacy_default(k)) + _routes_equal(omitted, unset) + _routes_equal(omitted, explicit) + + +def test_route_norm_changes_only_combine_weights(): + logits = _fixed_logits() + probs = torch.softmax(logits, dim=1) + for k in (1, 2, 3): + for sparse_routes in (False, True): + baseline = _gate_call(logits, k, sparse_routes, route_norm=None) + enabled = _gate_call(logits, k, sparse_routes, route_norm=True) + disabled = _gate_call(logits, k, sparse_routes, route_norm=False) + # Selection, capacity, and the load-balancing loss stay on the pre-renorm gate. + assert torch.equal(baseline[0], enabled[0]) + assert torch.equal(baseline[0], disabled[0]) + weight = 5 if sparse_routes else 1 + flipped = enabled if _legacy_default(k) is False else disabled + assert not torch.equal(baseline[weight], flipped[weight]) + if sparse_routes: + _routes_equal(baseline[1:5], enabled[1:5]) + _routes_equal(baseline[1:5], disabled[1:5]) + _routes_equal(baseline[6:], enabled[6:]) + _routes_equal(baseline[6:], disabled[6:]) + else: + assert torch.equal(baseline[2], enabled[2]) + assert torch.equal(baseline[2], disabled[2]) + assert torch.equal(baseline[3], enabled[3]) + assert torch.equal(baseline[3], disabled[3]) + + sparse_top1 = _gate_call(logits, 1, True, route_norm=None) + chosen = probs.gather(1, sparse_top1[3].reshape(-1).long().unsqueeze(1)).reshape(-1) + assert torch.equal(sparse_top1[5].reshape(-1), chosen) + assert not torch.allclose(chosen, torch.ones_like(chosen)) + + renorm_top1 = _gate_call(logits, 1, True, route_norm=True) + assert torch.equal(renorm_top1[5].reshape(-1), torch.ones_like(chosen)) + dense_top1 = _gate_call(logits, 1, False, route_norm=True) + assert torch.equal(dense_top1[1].sum(dim=(1, 2)), torch.ones(logits.shape[0])) + + raw_top2 = _gate_call(logits, 2, True, route_norm=False) + raw_mass = raw_top2[5].sum(dim=0) + assert not torch.allclose(raw_mass, torch.ones_like(raw_mass)) + renorm_top2 = _gate_call(logits, 2, True, route_norm=None) + assert torch.allclose(renorm_top2[5].sum(dim=0), torch.ones(logits.shape[0])) + + +def test_route_norm_dropped_token_stays_zero(): + # Four tokens, four experts, capacity 1, every token prefers expert 0. Three tokens are dropped. + logits = torch.tensor([[5.0, 0.0, 0.0, 0.0], [4.0, 0.1, 0.0, 0.0], [3.0, 0.0, 0.1, 0.0], [2.0, 0.0, 0.0, 0.1]]) + sparse = top1gating(logits, 1.0, 0, drop_tokens=True, use_rts=False, sparse_routes=True, route_norm=True) + gates = sparse[5].reshape(-1) + assert int((gates > 0).sum()) == 1 + assert torch.equal(gates, (gates > 0).to(gates.dtype)) + dense = top1gating(logits, 1.0, 0, drop_tokens=True, use_rts=False, sparse_routes=False, route_norm=True) + row_mass = dense[1].sum(dim=(1, 2)) + assert torch.equal(row_mass, (row_mass > 0).to(row_mass.dtype)) + assert int(row_mass.sum()) == 1 + + +def test_route_norm_does_not_change_aux_loss_between_top1_and_topk(): + logits = _fixed_logits() + top1_loss = top1gating(logits, 1.0, 0, drop_tokens=False, use_rts=False)[0] + topk_loss = topkgating(logits, 1, 1.0, 0, drop_tokens=False)[0] + # The k=1 load-balancing formulas agree, so the loss matches even though top-k renormalizes and top-1 does not. + assert torch.equal(top1_loss, topk_loss) + assert torch.equal(top1_loss, top1gating(logits, 1.0, 0, drop_tokens=False, use_rts=False, route_norm=True)[0]) + assert torch.equal(topk_loss, topkgating(logits, 1, 1.0, 0, drop_tokens=False, route_norm=False)[0]) + top2_loss = top2gating(logits, 1.0, 0, drop_tokens=False, top2_2nd_expert_sampling=False)[0] + topk2_loss = topkgating(logits, 2, 1.0, 0, drop_tokens=False)[0] + # The two k=2 formulas do not agree. route_norm must not be used to paper over that. + assert not torch.equal(top2_loss, topk2_loss) + assert torch.equal( + top2_loss, + top2gating(logits, 1.0, 0, drop_tokens=False, top2_2nd_expert_sampling=False, route_norm=False)[0]) + assert torch.equal(topk2_loss, topkgating(logits, 2, 1.0, 0, drop_tokens=False, route_norm=False)[0]) + + +def _moe_layer(k, route_norm, state): + expert = torch.nn.Linear(4, 4, bias=False) + layer = MoE(hidden_size=4, + expert=expert, + num_experts=4, + ep_size=1, + k=k, + min_capacity=0, + drop_tokens=False, + use_rts=False, + top2_2nd_expert_sampling=False, + route_norm=route_norm) + layer.load_state_dict(state) + return layer + + +def test_moe_route_norm_default_preserves_output_and_opt_in_changes_it(): + torch.manual_seed(7) + seed = MoE(hidden_size=4, + expert=torch.nn.Linear(4, 4, bias=False), + num_experts=4, + ep_size=1, + k=1, + min_capacity=0, + drop_tokens=False, + use_rts=False, + top2_2nd_expert_sampling=False) + hidden = torch.randn(2, 3, 4) + baseline = seed(hidden)[0] + same = _moe_layer(1, None, seed.state_dict()) + assert torch.equal(same(hidden)[0], baseline) + assert torch.equal(same(hidden)[1], seed(hidden)[1]) + + flipped = _moe_layer(1, True, seed.state_dict()) + assert not torch.equal(flipped(hidden)[0], baseline) + assert torch.equal(flipped(hidden)[1], seed(hidden)[1]) + + torch.manual_seed(8) + seed_k2 = MoE(hidden_size=4, + expert=torch.nn.Linear(4, 4, bias=False), + num_experts=4, + ep_size=1, + k=2, + min_capacity=0, + drop_tokens=False, + use_rts=False, + top2_2nd_expert_sampling=False) + hidden_k2 = torch.randn(2, 3, 4) + baseline_k2 = seed_k2(hidden_k2)[0] + assert torch.equal(_moe_layer(2, True, seed_k2.state_dict())(hidden_k2)[0], baseline_k2) + assert not torch.equal(_moe_layer(2, False, seed_k2.state_dict())(hidden_k2)[0], baseline_k2) + assert torch.equal(_moe_layer(2, False, seed_k2.state_dict())(hidden_k2)[1], seed_k2(hidden_k2)[1]) + + +def test_topk_gate_positional_constructor_leaves_route_norm_unset(): + # deepspeed/ops/transformer/inference/moe_inference.py builds TopKGate with these positional args. + gate = TopKGate(8, 4, 1, 1.0, 1.0, 4, None, True, True, None) + assert gate.route_norm is None + assert gate.top2_2nd_expert_sampling is True + + class TestExpertWeightGradWithZero(DistributedTest): world_size = 2