From cb962589e91560733b97f4559920f74d729890f7 Mon Sep 17 00:00:00 2001 From: the magician <82004885@qq.com> Date: Wed, 15 Jul 2026 04:21:03 +0800 Subject: [PATCH] revert moe router topk softmax experiment --- qwen3_6_scripts/qwen3_5.py | 13 ++++++------- 1 file changed, 6 insertions(+), 7 deletions(-) diff --git a/qwen3_6_scripts/qwen3_5.py b/qwen3_6_scripts/qwen3_5.py index 8d6fa3b..7eb4e17 100644 --- a/qwen3_6_scripts/qwen3_5.py +++ b/qwen3_6_scripts/qwen3_5.py @@ -849,13 +849,12 @@ class Qwen3_5MoeSparseBlock(nn.Module): Output is partial (pre-all-reduce), same contract as FusedMoE with reduce_results=False. """ - # Routing: top-k logits -> softmax over selected experts. - # This is mathematically equivalent to softmax(all experts) -> top-k -> - # renormalise, but avoids a full (T, num_experts) softmax in decode. - topk_logits, topk_ids = torch.topk( - router_logits.float(), self.top_k, dim=-1) # (T, top_k) - topk_weights = torch.softmax( - topk_logits, dim=-1).to(hidden_states.dtype) # (T, top_k) + # Routing: softmax -> topk -> renormalise + routing_weights = torch.softmax(router_logits.float(), dim=-1) + topk_weights, topk_ids = torch.topk( + routing_weights, self.top_k, dim=-1) # (T, top_k) + topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True) + topk_weights = topk_weights.to(hidden_states.dtype) w13 = self.experts.w13_weight # (E, 2*I, H) w2 = self.experts.w2_weight # (E, H, I)