[9/N] MoE Refactor: cleanup dispatcher interfaces (#11847)
This commit is contained in:
@@ -219,7 +219,7 @@ class Qwen2MoeSparseMoeBlock(nn.Module):
|
||||
# router_logits: (num_tokens, n_experts)
|
||||
router_logits, _ = self.gate(hidden_states)
|
||||
shared_output = self._forward_shared_experts(hidden_states)
|
||||
topk_weights, topk_idx, _ = self.topk(
|
||||
topk_output = self.topk(
|
||||
hidden_states,
|
||||
router_logits,
|
||||
num_token_non_padded=forward_batch.num_token_non_padded,
|
||||
@@ -228,14 +228,10 @@ class Qwen2MoeSparseMoeBlock(nn.Module):
|
||||
),
|
||||
)
|
||||
else:
|
||||
topk_weights, topk_idx, _ = self.topk.empty_topk_output(
|
||||
hidden_states.device
|
||||
)
|
||||
topk_output = self.topk.empty_topk_output(hidden_states.device)
|
||||
final_hidden_states = self.experts(
|
||||
hidden_states=hidden_states,
|
||||
topk_idx=topk_idx,
|
||||
topk_weights=topk_weights,
|
||||
forward_batch=forward_batch,
|
||||
topk_output=topk_output,
|
||||
)
|
||||
|
||||
if shared_output is not None:
|
||||
|
||||
Reference in New Issue
Block a user