Skip to content
Open
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
11 changes: 9 additions & 2 deletions deepspeed/moe/layer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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:

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 required Signed-off-by trailer

Because this SHA has one parent, it is a non-merge commit, but git log --format='%b' b11e0a1924c934183ecfbe01e225c81a04a123af shows that its message has no Signed-off-by trailer. Recreate the commit with --signoff using the configured identity so that it satisfies the repository's commit/CI policy.

AGENTS.md reference: AGENTS.md:L6-L9

Useful? React with 👍 / 👎.


super(MoE, self).__init__()

Expand All @@ -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,
Expand Down
102 changes: 78 additions & 24 deletions deepspeed/moe/sharded_moe.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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
Expand Down Expand Up @@ -319,16 +331,25 @@ 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)

# 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)
Expand All @@ -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)

Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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}")
Expand Down Expand Up @@ -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
Expand All @@ -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)
Expand All @@ -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'
Expand All @@ -591,22 +631,36 @@ 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,
self.capacity_factor if self.training else self.eval_capacity_factor,
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()
Expand Down
2 changes: 2 additions & 0 deletions docs/_tutorials/mixture-of-experts.md
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
Loading
Loading