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
40 changes: 32 additions & 8 deletions deepspeed/moe/sharded_moe.py
Original file line number Diff line number Diff line change
Expand Up @@ -293,7 +293,8 @@ def top2gating(logits: Tensor,
min_capacity: int,
drop_tokens: bool = True,
ep_group: Union[torch.distributed.ProcessGroup, None] = None,
top2_2nd_expert_sampling: bool = True) -> Tuple[Tensor, Tensor, Tensor, Tensor]:
top2_2nd_expert_sampling: bool = True,
use_tutel: bool = False) -> Tuple[Tensor, Tensor, Tensor, Tensor]:
"""Implements Top2Gating on logits."""
# everything is in fp32 in this function
gates = F.softmax(logits, dim=1)
Expand Down Expand Up @@ -359,6 +360,12 @@ def top2gating(logits: Tensor,
gates1_s /= denom_s
gates2_s /= denom_s

if use_tutel:
indices1_s = torch.where(mask1.sum(dim=1).bool(), indices1_s, torch.full_like(indices1_s, -1))
indices2_s = torch.where(mask2.sum(dim=1).bool(), indices2_s, torch.full_like(indices2_s, -1))
return l_aux, capacity, num_experts, [indices1_s, indices2_s], [locations1_s,
locations2_s], [gates1_s, gates2_s], exp_counts

# Calculate combine_weights and dispatch_mask
gates1 = einsum("s,se->se", gates1_s, mask1_float)
gates2 = einsum("s,se->se", gates2_s, mask2_float)
Expand All @@ -380,6 +387,7 @@ def topkgating(
drop_tokens: bool = True,
ep_group: Union[torch.distributed.ProcessGroup, None] = None,
drop_policy: str = "probs",
use_tutel: bool = False,
) -> Tuple[Tensor, Tensor, Tensor, Tensor]:
"""Implements TopKGating on logits."""

Expand Down Expand Up @@ -439,6 +447,20 @@ def topkgating(

if locations is None:
raise ValueError(f"Locations is not set: {locations}")

if use_tutel:
indices_ = []
locations_ = []
gates_ = []
for route in range(k):
indices_s = top_idx[:, route]
route_mask = F.one_hot(indices_s, num_classes=num_experts).bool() & mask
indices_s = torch.where(route_mask.any(dim=1), indices_s, torch.full_like(indices_s, -1))
indices_.append(indices_s)
locations_.append(torch.sum(locations * route_mask, dim=1))
gates_.append(torch.sum(gates_masked * route_mask, dim=1))
return l_aux, capacity, num_experts, indices_, locations_, gates_, exp_counts

# dispatch_mask
locations_sc = _one_hot_to_float((locations * mask), capacity)

Expand Down Expand Up @@ -520,11 +542,16 @@ def forward(self,

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)
self.min_capacity, self.drop_tokens, self.ep_group, self.top2_2nd_expert_sampling,
use_tutel)
else:
gate_output = topkgating(logits, self.k,
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)
self.min_capacity,
self.drop_tokens,
self.ep_group,
use_tutel=use_tutel)

if self.wall_clock_breakdown:
self.timers(TOPK_GATE_TIMER).stop()
Expand Down Expand Up @@ -571,16 +598,13 @@ def __init__(self,
self.timers = SynchronizedWallClockTimer()
self.wall_clock_breakdown = False

self.use_tutel = use_tutel and TUTEL_INSTALLED and gate.k == 1
self.use_tutel = use_tutel and TUTEL_INSTALLED

if self.use_tutel:
logger.info('Using Tutel optimizations.')
elif use_tutel and not TUTEL_INSTALLED:
logger.warning("Tutel optimization requested but not installed. "
"Proceeding without Tutel.")
elif use_tutel and TUTEL_INSTALLED and gate.k != 1:
logger.warning("To enable Tutel optimization, use top-1 instead of top-2 gate. "
"Proceeding without Tutel.")

def _set_ep_group(self, ep_group):
self.ep_group = ep_group
Expand Down
29 changes: 29 additions & 0 deletions tests/unit/moe/test_moe.py
Original file line number Diff line number Diff line change
Expand Up @@ -300,6 +300,35 @@ def fail_collective(*args, **kwargs):

class TestTopkGate(DistributedTest):

@staticmethod
def _assert_tutel_routes_match(combine_weights, tutel_output):
_, capacity, num_experts, indices_, locations_, gates_, _ = tutel_output
reconstructed = torch.zeros_like(combine_weights)
for indices_s, locations_s, gates_s in zip(indices_, locations_, gates_):
route_mask = indices_s.ge(0)
safe_indices = indices_s.clamp_min(0)
expert_mask = torch.nn.functional.one_hot(safe_indices, num_classes=num_experts)
expert_mask *= route_mask.unsqueeze(1)
location_mask = torch.nn.functional.one_hot(locations_s, num_classes=int(capacity))
reconstructed += torch.einsum("s,se,sc->sec", gates_s, expert_mask, location_mask)
torch.testing.assert_close(reconstructed, combine_weights)

def test_tutel_top2_routing(self):
logits = torch.tensor([[0.1, 0.9, 0.2], [0.8, 0.1, 0.3], [0.4, 0.6, 0.5], [0.7, 0.2, 0.1]])
native_output = top2gating(logits, 0.5, 0, top2_2nd_expert_sampling=False)
tutel_output = top2gating(logits, 0.5, 0, top2_2nd_expert_sampling=False, use_tutel=True)

assert len(tutel_output[3]) == 2
self._assert_tutel_routes_match(native_output[1], tutel_output)

def test_tutel_topk_routing(self):
logits = torch.tensor([[0.1, 0.9, 0.2, 0.3], [0.8, 0.1, 0.3, 0.2], [0.4, 0.6, 0.5, 0.7], [0.7, 0.2, 0.1, 0.8]])
native_output = topkgating(logits, 3, 0.5, 0)
tutel_output = topkgating(logits, 3, 0.5, 0, use_tutel=True)

assert len(tutel_output[3]) == 3
self._assert_tutel_routes_match(native_output[1], tutel_output)

def test(self):

def check_equal(logits, cap, sparse_truth, res):
Expand Down
Loading