From aa5350c50f7fa901434c9afc39e8c51b3b139231 Mon Sep 17 00:00:00 2001 From: Hao Zhang Date: Fri, 27 Mar 2026 10:16:35 +0800 Subject: [PATCH] refactor: remove gate normalization in DecoderUnit Simplify the gate computation by removing the normalization step. The gate values are now used directly from softmax scatter without the additional sum normalization. --- qmp/networks/transformers.py | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/qmp/networks/transformers.py b/qmp/networks/transformers.py index 3ddba315..8ed2aa9d 100644 --- a/qmp/networks/transformers.py +++ b/qmp/networks/transformers.py @@ -143,9 +143,8 @@ def forward( similarity = torch.nn.functional.softmax(x @ self.centroid.t(), dim=-1) # top_k_indices: batch * site * selected _, top_k_indices = torch.topk(similarity + self.bias, self.selected_num, dim=-1) - # gate_prime, gate: batch * site * routed - gate_prime = torch.zeros_like(similarity).scatter_(-1, top_k_indices, similarity.gather(-1, top_k_indices)) - gate = gate_prime / gate_prime.sum(dim=-1).unsqueeze(-1) + # gate, gate: batch * site * routed + gate = torch.zeros_like(similarity).scatter_(-1, top_k_indices, similarity.gather(-1, top_k_indices)) for i, expert in enumerate(self.feed_forward_routed): y = y + expert(x) * gate[:, :, i].unsqueeze(-1) x = self.norm2(y)