From 6cd34b177ee124686ffa0c16ffa9ec4ba2d5f482 Mon Sep 17 00:00:00 2001 From: nie-zhentao Date: Thu, 29 Jan 2026 21:47:09 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BF=AE=E5=A4=8Daux=5Floss=E6=B2=A1=E4=BC=A0?= =?UTF-8?q?=E9=80=92=E7=9A=84bug?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../pynative/base_models/gpt/gpt_model.py | 10 +++++++--- .../pynative/transformers/moe/moe_layer.py | 9 ++++++++- .../pynative/transformers/transformer_block.py | 10 +++++++--- .../pynative/transformers/transformer_layer.py | 16 ++++++++++++---- 4 files changed, 34 insertions(+), 11 deletions(-) diff --git a/mindformers/pynative/base_models/gpt/gpt_model.py b/mindformers/pynative/base_models/gpt/gpt_model.py index 896c187e9..1b1ac69cd 100644 --- a/mindformers/pynative/base_models/gpt/gpt_model.py +++ b/mindformers/pynative/base_models/gpt/gpt_model.py @@ -262,7 +262,7 @@ class GPTModel(nn.Cell): labels, attention_mask, loss_mask = self._preprocess_input_labels_and_masks( input_ids, labels, attention_mask, loss_mask) - hidden_states, _ = self.language_model( + hidden_states, _, aux_loss = self.language_model( input_ids, position_ids, attention_mask, @@ -293,7 +293,11 @@ class GPTModel(nn.Cell): if self.calculate_per_token_loss: numerator0, denominator0 = loss + if aux_loss is not None: + numerator0 = numerator0 + aux_loss return numerator0, denominator0 + if aux_loss is not None: + loss = loss + aux_loss return loss, logits, hidden_states def language_model( @@ -349,7 +353,7 @@ class GPTModel(nn.Cell): attn_mask = self.concat_prefix((prefix_mask, attn_mask)) # Run decoder. - hidden_states = self.decoder( + hidden_states, aux_loss = self.decoder( decoder_input, attn_mask, rotary_pos_emb, @@ -358,7 +362,7 @@ class GPTModel(nn.Cell): input_ids, ) - return hidden_states, rotary_pos_emb + return hidden_states, rotary_pos_emb, aux_loss def shared_embedding_or_output_weight(self): """Gets the embedding weight or output logit weights when share embedding and output weights set to True. diff --git a/mindformers/pynative/transformers/moe/moe_layer.py b/mindformers/pynative/transformers/moe/moe_layer.py index f45c03723..d1732752a 100644 --- a/mindformers/pynative/transformers/moe/moe_layer.py +++ b/mindformers/pynative/transformers/moe/moe_layer.py @@ -124,7 +124,14 @@ class MoELayer(nn.Cell): else: final_out = out_experts - return final_out, aux_loss + # NOTE: + # MoELayer is used as the MLP module inside TransformerLayer. To allow the + # transformer stack to propagate aux_loss (load balancing loss) upwards + # without changing the MLP bias semantics, we return a three-tuple here: + # (final_out, mlp_output_bias, aux_loss) + # where mlp_output_bias is kept as None, and aux_loss is a scalar tensor + # or None when load balancing loss is disabled. + return final_out, None, aux_loss class HashRoutedMoELayer(MoELayer): """ diff --git a/mindformers/pynative/transformers/transformer_block.py b/mindformers/pynative/transformers/transformer_block.py index e333374af..9de09968d 100644 --- a/mindformers/pynative/transformers/transformer_block.py +++ b/mindformers/pynative/transformers/transformer_block.py @@ -152,13 +152,15 @@ class TransformerBlock(nn.Cell): input_ids (optional): Input index, only required when using hash router. Default: None. Returns: - Tuple[Tensor, Tensor]: A tuple containing: + Tuple[Tensor, Optional[Tensor]]: A tuple containing: - hidden_states (Tensor): Output tensor of shape (S, B, H). + - aux_loss (Tensor | None): Optional auxiliary loss accumulated over layers. """ + aux_loss = None for index in range(self.num_layers): layer = self._get_layer(index) prefix_kv = prefix_keys_values[index] if prefix_keys_values is not None else None - hidden_states, _, = layer( + hidden_states, _, layer_aux_loss = layer( hidden_states, attention_mask, rotary_pos_emb=rotary_pos_emb, @@ -166,9 +168,11 @@ class TransformerBlock(nn.Cell): actual_seq_len=actual_seq_len, input_ids=input_ids, ) + if layer_aux_loss is not None: + aux_loss = layer_aux_loss if aux_loss is None else aux_loss + layer_aux_loss # final layernorm. if self.post_layer_norm: hidden_states = self.final_layernorm(hidden_states) - return hidden_states + return hidden_states, aux_loss diff --git a/mindformers/pynative/transformers/transformer_layer.py b/mindformers/pynative/transformers/transformer_layer.py index 01eda37a1..bfb3d709d 100644 --- a/mindformers/pynative/transformers/transformer_layer.py +++ b/mindformers/pynative/transformers/transformer_layer.py @@ -178,10 +178,12 @@ class TransformerLayer(nn.Cell, BaseTransformerLayer): input_ids (optional): Input index, only required when using hash router. Default: None. Returns: - Tuple[Tensor, Tensor, float]: A tuple containing: + Tuple[Tensor, Tensor, Optional[Tensor]]: A tuple containing: - output (Tensor): Transformed hidden states of shape [s, b, h]. - context (Tensor): Updated context tensor if cross-attention is used, otherwise the same as input context. + - aux_loss (Tensor | None): Optional auxiliary loss (e.g. MoE load balancing + loss) produced by the MLP module. None when not applicable. """ # Note: context parameter is currently unused but kept for API compatibility. # It may be used in future cross-attention implementations. @@ -249,9 +251,15 @@ class TransformerLayer(nn.Cell, BaseTransformerLayer): residual = norm_input if isinstance(self.mlp, HashRoutedMoELayer): - mlp_output, mlp_output_bias = self.mlp(pre_mlp_layernorm_output, input_ids) + mlp_outputs = self.mlp(pre_mlp_layernorm_output, input_ids) else: - mlp_output, mlp_output_bias = self.mlp(pre_mlp_layernorm_output) + mlp_outputs = self.mlp(pre_mlp_layernorm_output) + + if len(mlp_outputs) == 3: + mlp_output, mlp_output_bias, aux_loss = mlp_outputs + else: + mlp_output, mlp_output_bias = mlp_outputs + aux_loss = None if mlp_output_bias is not None: mlp_output = self.add(mlp_output, mlp_output_bias) @@ -267,4 +275,4 @@ class TransformerLayer(nn.Cell, BaseTransformerLayer): output = self.add(residual, dropout_output) # Note: context parameter is returned for API compatibility but currently unused. # It may be deprecated in future versions. - return output, context + return output, context, aux_loss -- Gitee