已合并
[bugfix][pynative] 将 MoE/MTP/index loss 追踪重构为模型实例方法并修复 recompute 重复累加问题 #8465
[bugfix][pynative] 将 MoE/MTP/index loss 追踪重构为模型实例方法并修复 recompute 重复累加问题 #8465
已合并
niujunhao创建于 6月25日
共 8 个文件变更+582-192
@@ -71,3 +71,23 @@ class PyNativeDeepseekV3ForCausalLM(TrainModelMixin, DeepseekV3PreTrainedModel):
71 loss_mask=loss_mask,71 loss_mask=loss_mask,
72 actual_seq_len=actual_seq_len72 actual_seq_len=actual_seq_len
73 )73 )
74+ 
75+ def _update_expert_bias(self, metric_group, metric_group_size):
76+ return self.model._update_expert_bias(metric_group, metric_group_size)
77+ 
78+ def get_load_balancing_loss(
79+ self, metric_group, metric_group_size, pp_metric_group, pp_metric_group_size, **kwargs
80+ ):
81+ return self.model.get_load_balancing_loss(
82+ metric_group, metric_group_size,
83+ pp_metric_group, pp_metric_group_size, **kwargs,
84+ )
85+ 
86+ def reset_model_temporary_tensors(self):
87+ return self.model.reset_model_temporary_tensors()
88+ 
89+ def get_mtp_loss(self, metric_group, metric_group_size):
90+ return self.model.get_mtp_loss(metric_group, metric_group_size)
91+ 
92+ def get_index_loss(self):
93+ return self.model.get_index_loss()
@@ -28,6 +28,7 @@ from hyper_parallel.core.dtensor.placement_types import Replicate
28 28 
29from mindspore import Tensor, dtype, nn, mint, ops29from mindspore import Tensor, dtype, nn, mint, ops
30from mindspore.mint.distributed import all_reduce, get_world_size30from mindspore.mint.distributed import all_reduce, get_world_size
31+from mindspore.graph.api import _no_grad
31 32 
32from mindformers.tools.logger import logger33from mindformers.tools.logger import logger
33from mindformers.pynative.loss.loss import CrossEntropyLoss, ChunkCrossEntropyLoss34from mindformers.pynative.loss.loss import CrossEntropyLoss, ChunkCrossEntropyLoss
@@ -40,11 +41,19 @@ from mindformers.pynative.base_models.common.embeddings.rotary_pos_embedding imp
40from mindformers.pynative.base_models.common.embeddings.yarn_rotary_pos_embedding import YarnRotaryEmbedding41from mindformers.pynative.base_models.common.embeddings.yarn_rotary_pos_embedding import YarnRotaryEmbedding
41from mindformers.pynative.base_models.gpt.gpt_layer_specs import get_gpt_mtp_block_spec42from mindformers.pynative.base_models.gpt.gpt_layer_specs import get_gpt_mtp_block_spec
42from mindformers.pynative.transformers.transformer_block import TransformerBlock, TransformerBlockSubmodules43from mindformers.pynative.transformers.transformer_block import TransformerBlock, TransformerBlockSubmodules
43-from mindformers.pynative.transformers.multi_token_prediction import MultiTokenPredictionBlock, MTPLossAutoScaler44+from mindformers.pynative.transformers.multi_token_prediction import (
45+ MultiTokenPredictionBlock,
46+ MTPLossAutoScaler,
47+ process_mtp_loss,
48+ track_mtp_metrics,
49+)
44from mindformers.pynative.layers.linear import Linear50from mindformers.pynative.layers.linear import Linear
45from mindformers.pynative.optimizer.muon_utils import make_muon_fns51from mindformers.pynative.optimizer.muon_utils import make_muon_fns
46-from mindformers.pynative.transformers.multi_token_prediction import process_mtp_loss
47from mindformers.pynative.dtensor_compat import inplace_copy52from mindformers.pynative.dtensor_compat import inplace_copy
53+from mindformers.pynative.transformers.moe.moe_utils import track_moe_metrics
54+from mindformers.pynative.transformers.experimental_attention_variant.utils import (
55+ track_indexer_metrics,
56+)
48 57 
49 58 
50class GPTModel(nn.Cell):59class GPTModel(nn.Cell):
@@ -637,6 +646,147 @@ class GPTModel(nn.Cell):
637 attention_mask = self.casual_mask(input_ids)646 attention_mask = self.casual_mask(input_ids)
638 return labels, attention_mask, loss_mask647 return labels, attention_mask, loss_mask
639 648 
649+ def _get_expert_bias_modules(self):
650+ """Return (and lazily cache) the MoE layers in this model that own an
651+ auxiliary-loss-free ``expert_bias`` buffer.
652+ 
653+ The cache is an instance attribute, so each model object (e.g. each
654+ virtual pipeline chunk under ``pipeline_parallel_interleave_num > 1``)
655+ keeps its own module list. This makes the per-instance caching correct
656+ under interleaving: one model instance per chunk.
657+ """
658+ cache = getattr(self, "_expert_bias_modules_cache", None)
659+ if cache is None:
660+ cache = [
661+ module for _, module in self.cells_and_names()
662+ if getattr(module, "enable_expert_bias", False)
663+ ]
664+ self._expert_bias_modules_cache = cache
665+ return cache
666+ 
667+ def _get_global_aux_loss_modules(self):
668+ """Return (and lazily cache) the modules in this model that own a
669+ global-aux-loss tracker to reset each step.
670+ 
671+ Cached as an instance attribute for the same per-chunk reason as
672+ :meth:`_get_expert_bias_modules`.
673+ """
674+ cache = getattr(self, "_global_aux_loss_modules_cache", None)
675+ if cache is None:
676+ cache = [
677+ module for _, module in self.cells_and_names()
678+ if hasattr(module, "reset_global_aux_loss_tracker")
679+ ]
680+ self._global_aux_loss_modules_cache = cache
681+ return cache
682+ 
683+ def _update_expert_bias(self, metric_group, metric_group_size):
684+ """Update auxiliary-loss-free ``expert_bias`` and reset per-step token counters.
685+ 
686+ Must be called on every step (regardless of log interval or loss
687+ availability) and on every PP stage / VPP chunk, so the routers stay
688+ in sync across the pipeline.
689+ 
690+ Args:
691+ metric_group: Process group spanning the dp x cp domain
692+ (``loss_mesh``). Each rank in this group sees a different sample
693+ shard of the global batch, so its local ``tokens_per_expert``
694+ histogram is only a partial count. We all-reduce SUM over this
695+ group to recover the global-batch histogram, otherwise the bias
696+ delta diverges from the single-card baseline. ``None`` (or a
697+ single-rank group) skips the all-reduce. TP is intentionally
698+ excluded: the MoE plan replicates router inputs across TP ranks,
699+ so each TP rank already sees the same tokens; reducing over TP
700+ would multiply the count by ``tp_size``.
701+ metric_group_size: Size of ``metric_group`` (skip all-reduce when
702+ ``<= 1``).
703+ """
704+ cache = self._get_expert_bias_modules()
705+ if not cache:
706+ return
707+ 
708+ # NOTE: Currently this sync is blocking (thus exposed) and happens on the
709+ # default compute stream. Need to assess if this is OK performance-wise.
710+ tokens_per_expert_list = [module.tokens_per_expert for module in cache]
711+ tokens_per_expert_by_layer = ops.vstack(tokens_per_expert_list)
712+ 
713+ # All-reduce the per-rank token counts over the dp x cp domain so the
714+ # bias delta is computed from global-batch statistics, matching the
715+ # single-card path (which has no group to reduce over).
716+ if metric_group is not None and (
717+ metric_group_size is None or metric_group_size > 1):
718+ if DTensor is not None and isinstance(tokens_per_expert_by_layer, DTensor):
719+ tokens_per_expert_by_layer = tokens_per_expert_by_layer.to_local()
720+ all_reduce(tokens_per_expert_by_layer, group=metric_group)
721+ 
722+ with _no_grad():
723+ for moe_layer_idx, module in enumerate(cache):
724+ tokens_per_expert = tokens_per_expert_by_layer[moe_layer_idx].float()
725+ 
726+ # update the expert bias
727+ # this is not exactly the same as https://arxiv.org/pdf/2408.15664 proposed
728+ expert_bias_delta = module.load_balance_coeff * mint.sign(
729+ tokens_per_expert.mean() - tokens_per_expert
730+ )
731+ 
732+ # NOTE: Megatron-LM does not apply zero-mean correction to the expert bias delta;
733+ # Torchtitan applies zero-mean correction here
734+ module.expert_bias.add_(expert_bias_delta)
735+ module.tokens_per_expert = mint.zeros_like(module.tokens_per_expert)
736+ 
737+ def get_load_balancing_loss(self, metric_group, metric_group_size,
738+ pp_metric_group, pp_metric_group_size, **kwargs):
739+ """Combine and reduce the MoE load-balancing (aux) loss for logging.
740+ 
741+ ``track_moe_metrics`` combines the per-layer tracker across PP stages
742+ (a PP-group collective) before reducing it to a scalar, so it must run
743+ on every stage every step to stay deadlock-free. Non-last stages get
744+ ``None`` back (they only contribute to the combine).
745+ 
746+ Returns the raw tracked aux loss (NOT divided by the number of
747+ gradient-accumulation steps); the caller applies that scaling.
748+ """
749+ has_last = kwargs.get("has_last", True)
750+ return track_moe_metrics(
751+ loss_scale=self.config.moe_aux_loss_coeff,
752+ num_layers=self.config.num_layers,
753+ moe_layer_freq=self.config.moe_layer_freq,
754+ mtp_num_layers=self.config.mtp_num_layers,
755+ group=metric_group,
756+ group_size=metric_group_size,
757+ pp_group=pp_metric_group,
758+ pp_group_size=pp_metric_group_size,
759+ has_last=has_last,
760+ )
761+ 
762+ def reset_model_temporary_tensors(self):
763+ """Reset per-step temporary tensors (global-aux-loss trackers).
764+ 
765+ No-op unless ``moe_router_load_balancing_type == "global_aux_loss"``.
766+ """
767+ if self.config.moe_router_load_balancing_type != "global_aux_loss":
768+ return
769+ for module in self._get_global_aux_loss_modules():
770+ module.reset_global_aux_loss_tracker()
771+ 
772+ def get_mtp_loss(self, metric_group, metric_group_size):
773+ """Return the reduced multi-token-prediction (MTP) loss for logging.
774+ 
775+ Returns the raw per-layer tracked MTP loss tensor (NOT divided by the
776+ number of gradient-accumulation steps); the caller applies that
777+ scaling and formatting. ``None`` when there is no MTP loss to report.
778+ """
779+ return track_mtp_metrics(group=metric_group, group_size=metric_group_size)
780+ 
781+ def get_index_loss(self):
782+ """Return the reduced sparse-attention indexer loss for logging.
783+ 
784+ Returns the raw tracked indexer loss (NOT divided by the number of
785+ gradient-accumulation steps); the caller applies that scaling.
786+ ``None`` when there is no indexer loss to report.
787+ """
788+ return track_indexer_metrics()
789+ 
640 def get_gpt_transformer_config(self):790 def get_gpt_transformer_config(self):
641 """Get the transformer config for GPT model.791 """Get the transformer config for GPT model.
642 792 
@@ -80,6 +80,7 @@ from mindformers.pynative.base_models.gpt.gpt_model import GPTModel
80from mindformers.pynative.transformers.experimental_attention_variant.deepseek_v4_hybrid_attention import (80from mindformers.pynative.transformers.experimental_attention_variant.deepseek_v4_hybrid_attention import (
81 DSv4HybridSelfAttention,81 DSv4HybridSelfAttention,
82)82)
83+from mindformers.pynative.transformers.moe.moe_utils import set_moe_aux_loss_group_info
83from mindformers.tools.logger import logger84from mindformers.tools.logger import logger
84 85 
85__all__ = ["parallelize_gptmodel"]86__all__ = ["parallelize_gptmodel"]
@@ -115,37 +116,49 @@ def _unwrap_gptmodel(model: nn.Cell) -> nn.Cell:
115 return gpt_models[0]116 return gpt_models[0]
116 117 
117 118 
118-def _setup_mtp_embedding_grad_sync(model, parallel_dims):119+def _setup_mtp_embedding_grad_sync(model_parts, parallel_dims):
119 """Tag the MTP-shared input embedding for cross-PP-stage gradient sync.120 """Tag the MTP-shared input embedding for cross-PP-stage gradient sync.
120 121 
121- Under pipeline parallelism with MTP enabled, the input embedding is122+ Under PP + MTP, the input embedding is replicated on stage 0 (main forward)
122- replicated on two PP stages: stage 0 (the main forward) and the last stage123+ and the last stage (MTP block reuses it for shifted tokens). Unlike the
123- (which hosts the MTP block and reuses the embedding to embed the shifted MTP124+ single-card baseline where one tensor receives ``main_grad + mtp_grad``,
124- tokens). The single-card baseline keeps ONE embedding tensor that receives125+ each PP copy only sees its own partial gradient, so without sync the two
125- ``main_grad + mtp_grad``; under PP the stage-0 copy only sees ``main_grad``126+ copies drift apart and the MTP loss diverges.
126- and the MTP-stage copy only sees ``mtp_grad``. With no gradient sync the two
127- copies drift apart and the MTP loss diverges from single-card (a slow,
128- one-sided, compounding error).
129 127 
130 This builds a process group over the embedding-owning PP ranks (per128 This builds a process group over the embedding-owning PP ranks (per
131 dp/tp/cp/ep coordinate) and tags the local ``word_embeddings.weight`` with129 dp/tp/cp/ep coordinate) and tags the local ``word_embeddings.weight`` with
132 ``_embedding_grad_sync_group`` / ``_pp_replica_count``. The grad-norm helper130 ``_embedding_grad_sync_group`` / ``_pp_replica_count``. The grad-norm helper
133 (`_calculate_global_grad_norm` -> `_get_grad_factor`) then all-reduces the131 (`_calculate_global_grad_norm` -> `_get_grad_factor`) then all-reduces the
134- tagged gradient before computing the norm/step and counts the replicated132+ tagged gradient and counts the replicated embedding exactly once, with no
135- embedding exactly once, with no changes required in the trainer.133+ trainer changes.
134+ 
135+ ``model_parts`` is the full list of virtual-pipeline chunks owned by this
136+ rank. Under interleaving the main and MTP embeddings live in different
137+ chunks on different PP stages, so ownership must be aggregated across ALL
138+ chunks: per-chunk detection would make each ``all_gather`` see at most one
139+ owner, ``embed_ranks`` would never reach length 2, and the sync group would
140+ never form -- silently reintroducing the MTP drift described above.
136 """141 """
137 if parallel_dims is None or not parallel_dims.pp_enabled:142 if parallel_dims is None or not parallel_dims.pp_enabled:
138 return143 return
139 144 
140- gpt_model = _unwrap_gptmodel(model)145+ if not isinstance(model_parts, (list, tuple)):
141- cfg = gpt_model.get_gpt_transformer_config()146+ model_parts = [model_parts]
147+ if not model_parts:
148+ return
149+ 
150+ cfg = _unwrap_gptmodel(model_parts[0]).get_gpt_transformer_config()
142 if not getattr(cfg, "mtp_num_layers", 0):151 if not getattr(cfg, "mtp_num_layers", 0):
143 return152 return
144 153 
145- # Locate this rank's input-embedding weight (present on stage 0 and the MTP154+ # Locate this rank's input-embedding weight(s) across ALL chunks it owns
146- # stage). At most one on a non-interleaved run.155+ # (present on stage 0 and the MTP stage). With interleaving the main and MTP
156+ # embeddings sit in separate chunks, so we must scan every part -- not just
157+ # one -- to detect that this rank owns an embedding at all.
147 embed_weights = [158 embed_weights = [
148- param for name, param in model.parameters_and_names()159+ param
160+ for part in model_parts
161+ for name, param in part.parameters_and_names()
149 if name.endswith("embedding.word_embeddings.weight")162 if name.endswith("embedding.word_embeddings.weight")
150 ]163 ]
151 164 
@@ -179,8 +192,8 @@ def _setup_mtp_embedding_grad_sync(model, parallel_dims):
179 weight._embedding_grad_sync_group = group192 weight._embedding_grad_sync_group = group
180 weight._embedding_grad_sync_size = len(embed_ranks)193 weight._embedding_grad_sync_size = len(embed_ranks)
181 logger.info(194 logger.info(
182- "[MTP-EmbedSync] rank %d tagged embedding for grad-sync group %s",195+ "[MTP-EmbedSync] rank %d tagged %d embedding weight(s) for grad-sync group %s",
183- local_rank, embed_ranks,196+ local_rank, len(embed_weights), embed_ranks,
184 )197 )
185 198 
186 199 
@@ -1540,6 +1553,71 @@ def apply_context_parallel_attention(
1540 for layer_idx, mtp_layer in enumerate(mtp.layers):1553 for layer_idx, mtp_layer in enumerate(mtp.layers):
1541 _apply_cp_to_block(mtp_layer, f"mtp.layers.{layer_idx}")1554 _apply_cp_to_block(mtp_layer, f"mtp.layers.{layer_idx}")
1542 1555 
1556+def _setup_moe_aux_loss_group(model, parallel_dims):
1557+ """Attach the cp reduction group to every MoE router for aux-loss computation.
1558+ 
1559+ This function equips each MoE router with the process group used to
1560+ all-reduce ``tokens_per_expert`` before the aux loss is computed.
1561+ 
1562+ The group scope is **cp only**. The router runs BEFORE the MoE dispatcher,
1563+ so its token view is NOT sharded by TP (router inputs are replicated across
1564+ TP ranks even under sequence parallelism). DP ranks hold independent
1565+ samples — their averaging is handled by the DP all-reduce in
1566+ ``track_moe_metrics`` (Step 3), not here.
1567+ 
1568+ CP *does* shard the sequence across ranks, so when CP > 1 we collect the
1569+ cp mesh group so ``get_tokens_per_expert_and_token_count`` can sum the
1570+ partial per-expert histograms across CP shards.
1571+ 
1572+ When cp is not enabled, ``aux_groups`` stays empty (size = 1) — no
1573+ all-reduce is needed, and each rank computes aux_loss from its local
1574+ token view.
1575+ """
1576+ if parallel_dims is None:
1577+ return
1578+ 
1579+ # Only CP shards the router's token view. TP replicates router inputs
1580+ # (verified: seq_length is always the global seqlen regardless of SP),
1581+ # and DP holds independent data (averaged elsewhere).
1582+ aux_groups: list = []
1583+ aux_group_size = 1
1584+ cp_mesh = parallel_dims.get_optional_mesh("cp")
1585+ if cp_mesh is not None:
1586+ aux_groups.append(cp_mesh.get_group())
1587+ aux_group_size = cp_mesh.size()
1588+ 
1589+ gpt_model = _unwrap_gptmodel(model)
1590+ moe_routers = []
1591+ for _, cell in gpt_model.cells_and_names():
1592+ router = getattr(cell, "router", None)
1593+ # Identify TopKRouters by their ``aux_loss_group`` attribute, without
1594+ # importing the moe package (keeps this helper decoupled).
1595+ if router is not None and hasattr(router, "moe_aux_loss_coeff"):
1596+ moe_routers.append(router)
1597+ 
1598+ mtp = getattr(gpt_model, "mtp", None)
1599+ if mtp is not None:
1600+ for _, cell in mtp.cells_and_names():
1601+ router = getattr(cell, "router", None)
1602+ if router is not None and hasattr(router, "moe_aux_loss_coeff"):
1603+ moe_routers.append(router)
1604+ 
1605+ if not moe_routers:
1606+ return
1607+ 
1608+ # Set global communication domain variables in moe_utils so both the
1609+ # router forward pass (``get_tokens_per_expert_and_token_count``) and
1610+ # the step-end aggregation (``track_moe_metrics``) share the same cp
1611+ # reduction group without per-router attribute plumbing.
1612+ set_moe_aux_loss_group_info(aux_groups, aux_group_size)
1613+ 
1614+ logger.info(
1615+ "[MoE-AuxLossGroup] set global cp aux_loss group "
1616+ "(size=%d, ) for %d MoE router(s).",
1617+ aux_group_size, len(moe_routers),
1618+ )
1619+ 
1620+ 
1543def _apply_spmd_parallelism(1621def _apply_spmd_parallelism(
1544 model: nn.Cell,1622 model: nn.Cell,
1545 parallel_dims: Any,1623 parallel_dims: Any,
@@ -1650,10 +1728,10 @@ def _apply_spmd_parallelism(
1650 "[QK-Clip] could not resolve loss_mesh reduce group (%s); "1728 "[QK-Clip] could not resolve loss_mesh reduce group (%s); "
1651 "falling back to world all-reduce.", exc)1729 "falling back to world all-reduce.", exc)
1652 1730 
1653- # Tag the MTP-shared input embedding so its gradient is summed across the1731+ # Tag every MoE router with the dp x cp group so the seq_aux_loss /
1654- # embedding-owning PP stages before the optimizer step (handled in the1732+ # global_aux_loss can all-reduce tokens_per_expert and aggregated probs
1655- # grad-norm helper), keeping the embedding identical to the single-card run.1733+ # before applying the Switch formula.
1656- _setup_mtp_embedding_grad_sync(model, parallel_dims)1734+ _setup_moe_aux_loss_group(model, parallel_dims)
1657 1735 
1658 logger.info("GPTModel parallelization completed.")1736 logger.info("GPTModel parallelization completed.")
1659 return model1737 return model
@@ -1822,6 +1900,15 @@ def apply_pp(
1822 overlap=overlap,1900 overlap=overlap,
1823 )1901 )
1824 1902 
1903+ # Tag the MTP-shared input embedding so its gradient is summed across the
1904+ # embedding-owning PP stages before the optimizer step (handled in the
1905+ # grad-norm helper), keeping the embedding identical to the single-card run.
1906+ # Run ONCE over all chunks this rank owns (not per-chunk inside
1907+ # _apply_spmd_parallelism): under interleaving the main and MTP embeddings
1908+ # live in different chunks/stages, so ownership must be aggregated across
1909+ # the whole rank for the cross-PP-stage sync group to form.
1910+ _setup_mtp_embedding_grad_sync(model_parts, parallel_dims)
1911+ 
1825 # Adjust the last-stage loss backward scaling so PP gradients match the single-card case (loss / grad_accum).1912 # Adjust the last-stage loss backward scaling so PP gradients match the single-card case (loss / grad_accum).
1826 # Note: DP/TP/CP are already included in get_loss_sense as 1/(dp*tp*cp)/grad_accum.1913 # Note: DP/TP/CP are already included in get_loss_sense as 1/(dp*tp*cp)/grad_accum.
1827 # However, PipelineStage.get_last_stage_sens also divides the loss DTensor by repeat_num.1914 # However, PipelineStage.get_last_stage_sens also divides the loss DTensor by repeat_num.
@@ -17,11 +17,9 @@
17import time17import time
18from typing import Dict, Any18from typing import Dict, Any
19from copy import deepcopy19from copy import deepcopy
20-import os
21 20 
22-from mindspore import ops
23from mindspore.nn.learning_rate_schedule import LearningRateSchedule21from mindspore.nn.learning_rate_schedule import LearningRateSchedule
24-from mindspore.mint.distributed import get_world_size, all_reduce22+from mindspore.mint.distributed import get_world_size
25 23 
26from mindformers.pynative.callback.callback import TrainerCallback24from mindformers.pynative.callback.callback import TrainerCallback
27from mindformers.tools.logger import logger25from mindformers.tools.logger import logger
@@ -29,9 +27,6 @@ from mindformers.models.utils import (
29 convert_transformer_config_to_args_for_tflops,27 convert_transformer_config_to_args_for_tflops,
30 num_floating_point_operations,28 num_floating_point_operations,
31)29)
32-from mindformers.pynative.transformers.moe.moe_utils import track_moe_metrics
33-from mindformers.pynative.transformers.multi_token_prediction import track_mtp_metrics
34-from mindformers.pynative.transformers.experimental_attention_variant.utils import track_indexer_metrics
35 30 
36 31 
37class LossCallback(TrainerCallback):32class LossCallback(TrainerCallback):
@@ -53,6 +48,13 @@ class LossCallback(TrainerCallback):
53 self.log_interval = log_interval48 self.log_interval = log_interval
54 self.step_time = time.time()49 self.step_time = time.time()
55 self.epoch_time = time.time()50 self.epoch_time = time.time()
51+ self.logger_record = {
52+ "_update_expert_bias": False,
53+ "reset_model_temporary_tensors": False,
54+ "get_load_balancing_loss": False,
55+ "get_mtp_loss": False,
56+ "get_index_loss": False,
57+ }
56 58 
57 def on_train_begin(self, args, state, **kwargs):59 def on_train_begin(self, args, state, **kwargs):
58 """60 """
@@ -129,55 +131,64 @@ class LossCallback(TrainerCallback):
129 step_time_cost = int((cur_time - self.step_time) * 1000)131 step_time_cost = int((cur_time - self.step_time) * 1000)
atomgit-botatomgit-bot
atomgit-botatomgit-bot6月25日

🔵 Low Priority

变更前,cur_time = time.time() 位于 _update_expert_bias + track_moe_metrics 之后(旧代码约 L169),step_time_cost 包含了这些每步必执行的集合通信(DP×CP all-reduce 更新 expert_bias、PP-group all-reduce 组合 aux loss)的开销,吞吐量指标反映完整步耗时。

变更后(L129-131),cur_time 被移到 on_step_end 开头,位于 _update_expert_bias 循环(L146-151)和 get_load_balancing_loss 调用(L160-167)之前。这两个操作每步都执行且包含跨卡集合通信(token 直方图 all-reduce + PP group all-reduce),在大规模训练中耗时不可忽略(通常数毫秒到数十毫秒),导致上报的 step_time 偏小、throughput 偏高,指标不再反映真实步耗时。

建议:将 cur_time 和 step_time_cost 的计算移到 get_load_balancing_loss 调用之后(约 L173),使其包含完整的每步操作开销,与旧行为一致。或者至少移到 L175 的 early return 之前紧邻的位置。

likedislike
不准确?
atomgit-botatomgit-bot6月25日

🔵 Low Priority

变更在 on_step_end 中将 cur_time = time.time() 从旧位置(_update_expert_bias + track_moe_metrics 集合通信之后)提前到了函数开头(L129-131)。这导致 step_time_cost 不再包含 expert bias 更新(含 dp×cp all-reduce)和 MoE aux loss PP-group all-reduce 的耗时,上报的 step_time 偏小、throughput 偏高。

证据链:

  • 变更行 L129-131:cur_time / step_time_cost 计算前移
  • 旧代码:cur_time 在 _update_expert_bias + track_moe_metrics 之后
  • 影响行为:每步必执行的集合通信(expert bias token 直方图 all-reduce、PP-group aux loss all-reduce)的开销不再计入步耗时
  • 失败模式:"step_time" 和 "throughput" 日志指标系统性偏低/偏高,掩盖真实性能,可能干扰性能回归判断

建议:将 cur_time = time.time() 和 step_time_cost 的计算移到 L173(load_balancing_loss 处理完成之后),或至少紧邻 L175 的 early return 之前。这确保 step_time 包含所有每步必执行的集合通信开销,与旧行为一致。

likedislike
不准确?
130 132 
131 model = model if isinstance(model, list) else [model]133 model = model if isinstance(model, list) else [model]
134+ model_cls_name = type(model[0]).__name__
132 model_config = deepcopy(model[0].get_gpt_transformer_config())135 model_config = deepcopy(model[0].get_gpt_transformer_config())
133 if model_config is None:136 if model_config is None:
134- raise ValueError("model_config is None, please check the model type.")137+ raise ValueError(f"{model_cls_name} model_config is None, please check the model type.")
135- for m in model:138+ 
139+ if getattr(model_config, "moe_router_enable_expert_bias", False):
136 # Update auxiliary-loss-free expert_bias on every step, regardless of140 # Update auxiliary-loss-free expert_bias on every step, regardless of
137 # log interval or loss availability. Non-last PP stages return loss=None141 # log interval or loss availability. Non-last PP stages return loss=None
138 # but still hold MoE layers whose ``tokens_per_expert`` accumulators must142 # but still hold MoE layers whose ``tokens_per_expert`` accumulators must
139 # be drained and converted into a bias delta to keep the router in sync143 # be drained and converted into a bias delta to keep the router in sync
140- # with the rest of the pipeline.144+ # with the rest of the pipeline. The model keeps a per-instance module
141- if getattr(model_config, "moe_router_enable_expert_bias", False):145+ # cache, so each virtual pipeline chunk is updated independently.
142- _update_expert_bias(m, metric_group, metric_group_size)146+ for m in model:
147+ if hasattr(m, "_update_expert_bias"):
148+ m._update_expert_bias(metric_group, metric_group_size)
149+ elif not self.logger_record["_update_expert_bias"]:
150+ logger.warning(f"{model_cls_name} does not have _update_expert_bias method.")
151+ self.logger_record["_update_expert_bias"] = True
152+ 
153+ # Process the MoE aux loss. ``get_load_balancing_loss`` combines the
154+ # per-layer tracker across PP stages (a PP-group collective) before
155+ # reducing it to a scalar, so it must run on every stage every step --
156+ # ahead of the log-interval / loss-None early returns below -- to stay
157+ # deadlock-free. Non-last stages get ``None`` back (they only contribute
158+ # to the combine).
159+ load_balancing_loss = None
160+ if hasattr(model[0], "get_load_balancing_loss"):
161+ load_balancing_loss = model[0].get_load_balancing_loss(
162+ metric_group,
163+ metric_group_size,
164+ pp_metric_group,
165+ pp_metric_group_size,
166+ has_last=has_last,
167+ )
168+ elif not self.logger_record["get_load_balancing_loss"]:
169+ logger.warning(f"{model_cls_name} does not have get_load_balancing_loss method.")
170+ self.logger_record["get_load_balancing_loss"] = True
143 171 
144- # Process the MoE aux loss. ``track_moe_metrics`` combines the per-layer
145- # tracker across PP stages (a PP-group collective) before reducing it to
146- # a scalar, so it must run on every stage every step -- ahead of the
147- # log-interval / loss-None early returns below -- to stay deadlock-free.
148- # Non-last stages get ``None`` back (they only contribute to the combine).
149- load_balancing_loss = track_moe_metrics(
150- loss_scale=model_config.moe_aux_loss_coeff,
151- num_layers=model_config.num_layers,
152- moe_layer_freq=model_config.moe_layer_freq,
153- mtp_num_layers=model_config.mtp_num_layers,
154- group=metric_group,
155- group_size=metric_group_size,
156- pp_group=pp_metric_group,
157- pp_group_size=pp_metric_group_size,
158- has_last=has_last,
159- )
160 if load_balancing_loss is not None:172 if load_balancing_loss is not None:
161 load_balancing_loss /= state.num_accumulation_steps173 load_balancing_loss /= state.num_accumulation_steps
162 174 
163 if loss is None or state.global_step % self.log_interval != 0:175 if loss is None or state.global_step % self.log_interval != 0:
164 return176 return
165 177 
166- grad_norm = kwargs.get("grad_norm")
167- 
168- # Calculate the time cost for the current step in milliseconds
169- cur_time = time.time()
170- step_time_cost = int((cur_time - self.step_time) * 1000)
171 for m in model:178 for m in model:
172- reset_model_temporary_tensors(model_config, m)179+ if hasattr(m, "reset_model_temporary_tensors"):
180+ m.reset_model_temporary_tensors()
181+ elif not self.logger_record["reset_model_temporary_tensors"]:
182+ logger.warning(f"{model_cls_name} does not have reset_model_temporary_tensors method.")
183+ self.logger_record["reset_model_temporary_tensors"] = True
173 184 
174 # process mtp loss185 # process mtp loss
175- mtp_loss = track_mtp_metrics(group=metric_group, group_size=metric_group_size,)186+ mtp_loss = None
176- # mtp_loss_scaling_factor == 0 means the MTP params are frozen (not updated),187+ if hasattr(model[0], "get_mtp_loss"):
177- # so the tracked loss carries no signal and is not worth logging. Guarding on188+ mtp_loss = model[0].get_mtp_loss(metric_group, metric_group_size)
178- # the config value also avoids the element-wise truthiness of the Tensor189+ elif not self.logger_record["get_mtp_loss"]:
179- # returned by track_mtp_metrics, which crashed the previous `if mtp_loss:`190+ logger.warning(f"{model_cls_name} does not have get_mtp_loss method.")
180- # check by leaking a raw Tensor into _print_log's join.191+ self.logger_record["get_mtp_loss"] = True
181 if mtp_loss is not None and model_config.mtp_loss_scaling_factor:192 if mtp_loss is not None and model_config.mtp_loss_scaling_factor:
182 mtp_loss_values = []193 mtp_loss_values = []
183 for ind, val in enumerate(mtp_loss):194 for ind, val in enumerate(mtp_loss):
@@ -188,7 +199,12 @@ class LossCallback(TrainerCallback):
188 mtp_loss = None199 mtp_loss = None
189 200 
190 # process indexer loss201 # process indexer loss
191- indexer_loss = track_indexer_metrics()202+ indexer_loss = None
203+ if hasattr(model[0], "get_index_loss"):
204+ indexer_loss = model[0].get_index_loss()
205+ elif not self.logger_record["get_index_loss"]:
206+ logger.warning(f"{model_cls_name} does not have get_index_loss method.")
207+ self.logger_record["get_index_loss"] = True
192 if indexer_loss:208 if indexer_loss:
193 indexer_loss /= state.num_accumulation_steps209 indexer_loss /= state.num_accumulation_steps
194 210 
@@ -309,97 +325,3 @@ class LossCallback(TrainerCallback):
309 if hasattr(data, "item"):325 if hasattr(data, "item"):
310 return data.item()326 return data.item()
311 return float(data)327 return float(data)
312- 
313- 
314-def reset_model_temporary_tensors(config, model):
315- """
316- Reset the temporary tensors of the model.
317- 
318- Uses cached module list to avoid full cell tree traversal on every step.
319- """
320- if config.moe_router_load_balancing_type != "global_aux_loss":
321- return
322- 
323- cache = getattr(reset_model_temporary_tensors, '_cache', None)
324- if cache is None:
325- cache = [
326- module for _, module in model.cells_and_names()
327- if hasattr(module, 'reset_global_aux_loss_tracker')
328- ]
329- reset_model_temporary_tensors._cache = cache
330- 
331- for module in cache:
332- module.reset_global_aux_loss_tracker()
333- 
334- 
335-def _update_expert_bias(model, metric_reduce_group=None, metric_reduce_group_size=None):
336- """
337- Update expert bias for load-balanced routing and reset per-step token counters.
338- 
339- Uses cached module list to avoid full cell tree traversal on every step.
340- Args:
341- model: Root model module to walk for MoE layers.
342- metric_reduce_group: Process group spanning the dp x cp domain
343- (``loss_mesh``). Each rank in this group sees a different sample
344- shard of the global batch, so its local ``tokens_per_expert``
345- histogram is only a partial count. We all-reduce SUM over this
346- group to recover the global-batch histogram, otherwise the bias
347- delta diverges from the single-card baseline. ``None`` (or
348- single-rank group) skips the all-reduce. TP is intentionally
349- excluded: the MoE plan replicates router inputs across TP ranks,
350- so each TP rank already sees the same tokens; reducing over TP
351- would multiply the count by ``tp_size``.
352- metric_reduce_group_size: Size of ``metric_reduce_group`` (skip
353- all-reduce when ``<= 1``).
354- """
355- from mindspore import mint
356- from mindspore.graph.api import _no_grad
357- 
358- try:
359- from hyper_parallel.core.dtensor.dtensor import DTensor
360- except ImportError:
361- DTensor = None
362- 
363- cache = getattr(_update_expert_bias, '_cache', None)
364- if cache is None:
365- cache = [
366- module for _, module in model.cells_and_names()
367- if getattr(module, 'enable_expert_bias', False)
368- ]
369- _update_expert_bias._cache = cache
370- 
371- if not cache:
372- return
373- 
374- # NOTE: Currently this sync is blocking (thus exposed) and happens on the
375- # default compute stream. Need to assess if this is OK performance-wise.
376- tokens_per_expert_list = [module.tokens_per_expert for module in cache]
377- tokens_per_expert_by_layer = ops.vstack(tokens_per_expert_list)
378- 
379- # All-reduce the per-rank token counts over the dp x cp domain so the
380- # bias delta is computed from global-batch statistics, matching the
381- # single-card path (which has no group to reduce over).
382- 
383- if metric_reduce_group is not None and (
384- metric_reduce_group_size is None or metric_reduce_group_size > 1):
385- if DTensor is not None and isinstance(tokens_per_expert_by_layer, DTensor):
386- tokens_per_expert_by_layer = tokens_per_expert_by_layer.to_local()
387- all_reduce(tokens_per_expert_by_layer, group=metric_reduce_group)
388- 
389- with _no_grad():
390- for moe_layer_idx, module in enumerate(cache):
391- tokens_per_expert = tokens_per_expert_by_layer[moe_layer_idx].float()
392- 
393- # update the expert bias
394- # this is not exactly the same as https://arxiv.org/pdf/2408.15664 proposed
395- # pyrefly: ignore [missing-attribute]
396- expert_bias_delta = module.load_balance_coeff * mint.sign(
397- tokens_per_expert.mean() - tokens_per_expert
398- )
399- # NOTE: Megatron-LM does not apply zero-mean correction to the expert bias delta;
400- # Torchtitan applies zero-mean correction here
401- 
402- # pyrefly: ignore [missing-attribute]
403- module.expert_bias.add_(expert_bias_delta)
404- # pyrefly: ignore [missing-attribute]
405- module.tokens_per_expert = mint.zeros_like(module.tokens_per_expert)
@@ -20,6 +20,7 @@ from mindspore.common.parameter import Parameter
20from mindformers.parallel_core.transformer_config import TransformerConfig20from mindformers.parallel_core.transformer_config import TransformerConfig
21from mindformers.pynative.layers.linear import Linear21from mindformers.pynative.layers.linear import Linear
22from mindformers.pynative.transformers.mlp import MLPSubmodules22from mindformers.pynative.transformers.mlp import MLPSubmodules
23+from mindformers.pynative.distributed.activation_checkpoint import is_in_recompute
23from .router import TopKRouter24from .router import TopKRouter
24from .experts import GroupedMLP25from .experts import GroupedMLP
25from .shared_experts import SharedExpertMLP26from .shared_experts import SharedExpertMLP
@@ -125,7 +126,8 @@ class MoELayer(nn.Cell):
125 hidden_states, self.expert_bias, input_ids126 hidden_states, self.expert_bias, input_ids
126 )127 )
127 128 
128- self.tokens_per_expert.add_(num_tokens_per_expert)129+ if not is_in_recompute():
130+ self.tokens_per_expert.add_(num_tokens_per_expert)
129 131 
130 routed_output = self.experts(132 routed_output = self.experts(
131 hidden_states, top_scores, selected_experts_indices, num_tokens_per_expert133 hidden_states, top_scores, selected_experts_indices, num_tokens_per_expert
@@ -31,6 +31,11 @@ from mindformers.pynative.distributed.activation_checkpoint import is_in_recompu
31# MOE logging31# MOE logging
32_MOE_LAYER_WISE_LOGGING_TRACKER: dict = {}32_MOE_LAYER_WISE_LOGGING_TRACKER: dict = {}
33 33 
34+# Global communication domain variables for MoE aux loss.
35+# Set via ``set_moe_aux_loss_group_info`` from the model parallelize step.
36+_AUX_LOSS_GROUP = None
37+_AUX_LOSS_GROUP_SIZE = 1
38+ 
34 39 
35def switch_load_balancing_loss_func(40def switch_load_balancing_loss_func(
36 probs: Tensor,41 probs: Tensor,
@@ -104,6 +109,173 @@ def switch_load_balancing_loss_func(
104 return aux_loss109 return aux_loss
105 110 
106 111 
112+def get_tokens_per_expert_and_token_count(
113+ routing_map: Tensor,
114+ reduce_group,
115+ topk: int = None,
116+ with_padding_mask: bool = False,
117+) -> Tuple[Tensor, int, int]:
118+ """Compute per-expert global token counts and local/total token counts for MoE aux loss.
119+ 
120+ This is the core statistics entry point for the Mixture-of-Experts load-balancing
121+ auxiliary loss. It performs two main tasks:
122+ 
123+ 1. Sums ``routing_map`` along the expert dimension to obtain the local per-expert
124+ token histogram (``local_tokens_per_expert``), then all-reduces (SUM) over
125+ ``reduce_group`` to aggregate the global per-expert token count
126+ (``global_tokens_per_expert``) across the parallel domain.
127+ 2. Derives the local token count (``local_num_tokens``) from the row count of
128+ ``routing_map``, and the global token total (``total_num_tokens``) by scaling
129+ it with the world size of ``reduce_group``.
130+ 
131+ ``reduce_group`` spans only the TP x CP domain: routing happens BEFORE the MoE
132+ dispatcher, so only TP/CP shard the router's token view, and summing over the
133+ group recovers the per-(batch, expert) histogram for the full batch. DP is
134+ intentionally excluded -- each DP rank contributes its own per-sequence aux loss,
135+ averaged by the caller via ``/bsz``; all-reducing over DP would double-count.
136+ 
137+ Args:
138+ routing_map (Tensor): Local token-expert assignment map with shape
139+ ``[slen, bsz*E]`` (after the per-sequence reshape) or ``[T, E]``,
140+ matching the format expected by the calling aux-loss variant.
141+ reduce_group: Either a single process group covering the tp×cp domain,
142+ OR a list of 1D process groups (one per axis) covering the same
143+ domain. The list form is used when the mesh is 2D and only per-axis
144+ groups are available. ``None`` or empty means "no reduction needed;
145+ the local histogram is already the global view".
146+ topk (int): Number of experts selected per token. Required when
147+ ``with_padding_mask=True``; the ``seq_aux_loss`` variant passes
148+ ``topk * bsz`` to account for the per-batch flatten.
149+ with_padding_mask (bool): Whether the routing_map is padded. Currently
150+ unsupported in this codebase.
151+ 
152+ Returns:
153+ Tuple of ``(global_tokens_per_expert, local_num_tokens, total_num_tokens)``:
154+ - ``global_tokens_per_expert``: per-expert token count after all-reduce over
155+ ``reduce_group``.
156+ - ``local_num_tokens``: row count of ``routing_map`` (local token count).
157+ - ``total_num_tokens``: ``local_num_tokens * reduce_group.size()`` (or the
158+ histogram-derived equivalent when padding is masked).
159+ """
160+ _ = topk
161+ if with_padding_mask:
162+ raise NotImplementedError(
163+ "Padding-mask aware token count is not implemented in this "
164+ "codebase; ``get_tokens_per_expert_and_token_count`` falls back "
165+ "to the row-count derivation used in the non-padded path."
166+ )
167+ local_tokens_per_expert = routing_map.sum(dim=0)
168+ global_tokens_per_expert = local_tokens_per_expert
169+ # Normalize ``reduce_group`` to a list of 1D groups so the nested-reduce
170+ # logic below is uniform regardless of whether the caller passed a single
171+ # group (legacy path) or a list of axis groups (tp×cp decomposition).
172+ if reduce_group is None:
173+ reduce_groups = []
174+ elif isinstance(reduce_group, (list, tuple)):
175+ reduce_groups = list(reduce_group)
176+ else:
177+ reduce_groups = [reduce_group]
178+ if reduce_groups:
179+ for sub_group in reduce_groups:
180+ all_reduce(global_tokens_per_expert, op=ops.ReduceOp.SUM, group=sub_group)
181+ local_num_tokens = routing_map.shape[0]
182+ total_num_tokens = local_num_tokens * (
183+ get_world_size_from_group(reduce_group)
184+ )
185+ return global_tokens_per_expert, local_num_tokens, total_num_tokens
186+ 
187+ 
188+def get_world_size_from_group(group) -> int:
189+ """Best-effort ``world_size`` lookup for a process group.
190+ 
191+ Falls back to ``1`` if the size cannot be resolved (e.g. ``group`` is
192+ ``None`` or the build lacks a world-size API). Accepts a list/tuple of
193+ groups (the tp×cp decomposition produced by
194+ ``set_moe_aux_loss_group_info``) and returns the product of the
195+ per-axis sizes, i.e. ``cp_size * tp_size``.
196+ """
197+ # pylint: disable=broad-exception-caught
198+ if group is None:
199+ return 1
200+ if isinstance(group, (list, tuple)):
201+ # Early-exit on a degenerate entry to avoid inflating the product
202+ # with a stray 0 from a group whose size could not be resolved.
203+ sub_sizes = [get_world_size_from_group(g) for g in group]
204+ if any(s <= 0 for s in sub_sizes):
205+ return 1
206+ size = 1
207+ for s in sub_sizes:
208+ size *= s
209+ return size
210+ try:
211+ return group.size()
212+ except Exception:
213+ pass
214+ try:
215+ return get_world_size(group)
216+ except Exception:
217+ return 1
218+ 
219+ 
220+def set_moe_aux_loss_group_info(groups, group_size):
221+ """Set global communication domain variables for MoE aux loss.
222+ 
223+ Called once during model parallelization to establish the tp×cp reduction
224+ group(s) used by both the router-side histogram reduction
225+ (``get_tokens_per_expert_and_token_count``) and the step-end logging
226+ aggregation (``track_moe_metrics``).
227+ 
228+ The hyper-parallel ``DeviceMesh`` does not expose a single process group
229+ covering an entire multi-axis mesh, and the per-axis ``get_group()`` API
230+ raises when the mesh is 2D. We therefore store the list of 1D axis groups
231+ (e.g. ``[cp_group, tp_group]``) and reduce over them sequentially -- SUM
232+ is associative, so the result is identical to a single tp×cp all-reduce.
233+ 
234+ Args:
235+ groups: List of 1D process groups whose Cartesian product spans the
236+ tp×cp domain. ``None`` or an empty list means "no reduction
237+ needed" (e.g. tp=cp=1). A single-element list is also accepted
238+ (e.g. only tp is enabled).
239+ group_size: World size of the full tp×cp domain
240+ (``cp_size * tp_size``). Used by the router to scale the aux
241+ loss gradient so it is independent of the mesh shape.
242+ """
243+ # pylint: disable=W0603
244+ global _AUX_LOSS_GROUP, _AUX_LOSS_GROUP_SIZE
245+ if groups is None:
246+ groups = []
247+ # Drop None entries (size-1 axes) and keep at most one entry per axis.
248+ _AUX_LOSS_GROUP = [g for g in groups if g is not None]
249+ _AUX_LOSS_GROUP_SIZE = max(group_size, 1)
250+ 
251+ 
252+def get_moe_aux_loss_group():
253+ """Return the list of tp×cp process groups for aux-loss histogram reduction.
254+ 
255+ Returns a list (possibly empty) of 1D groups. Callers that perform a
256+ single all-reduce should use ``reduce_over_aux_loss_groups`` instead,
257+ which iterates the list and applies each group sequentially.
258+ """
259+ return _AUX_LOSS_GROUP
260+ 
261+ 
262+def get_moe_aux_loss_group_size():
263+ """Return the world size of the tp×cp aux-loss group (1 if unset)."""
264+ return _AUX_LOSS_GROUP_SIZE
265+ 
266+ 
267+def reduce_over_aux_loss_groups(tensor, op=ops.ReduceOp.SUM):
268+ """Apply SUM/other reduction sequentially over every tp×cp axis group.
269+ 
270+ Equivalent to a single all-reduce over the full tp×cp mesh, but works
271+ with hyper-parallel's per-axis 1D process groups. SUM is associative so
272+ the order of axes does not affect the result. No-op when there are no
273+ enabled groups (e.g. tp=cp=1).
274+ """
275+ for group in _AUX_LOSS_GROUP:
276+ all_reduce(tensor, op=op, group=group)
277+ 
278+ 
107class _MoEAuxLossAutoScaler(_Function):279class _MoEAuxLossAutoScaler(_Function):
108 """An AutoScaler that triggers the backward pass and scales the grad for auxiliary loss."""280 """An AutoScaler that triggers the backward pass and scales the grad for auxiliary loss."""
109 281 
@@ -237,11 +409,11 @@ def save_to_aux_losses_tracker(
237 if isinstance(loss, DTensor):409 if isinstance(loss, DTensor):
238 loss = loss.to_local()410 loss = loss.to_local()
239 if hasattr(loss, "detach"):411 if hasattr(loss, "detach"):
240- tracker["values"][layer_number] = tracker["values"][412+ cur_val = tracker["values"][layer_number] + loss.detach()
241- layer_number] + loss.detach() # Aggregate the loss for the layer.413+ tracker["values"][layer_number] = cur_val # Aggregate the loss for the layer.
242 else:414 else:
243- tracker["values"][layer_number] = tracker["values"][layer_number] + loss415+ cur_val = tracker["values"][layer_number] + loss
244- 416+ tracker["values"][layer_number] = cur_val
245 417 
246def clear_aux_losses_tracker() -> None:418def clear_aux_losses_tracker() -> None:
247 """Clear the auxiliary losses."""419 """Clear the auxiliary losses."""
@@ -286,12 +458,20 @@ def track_moe_metrics(
286 458 
287 tracker = get_moe_layer_wise_logging_tracker()459 tracker = get_moe_layer_wise_logging_tracker()
288 460 
289- # Combine the per-layer tracker across PP stages (collective; see docstring).461+ # Step 1: CP all-reduce (SUM) — combine context-parallel shards.
462+ # Router inputs are replicated across TP and independent across DP;
463+ # only CP shards the token view, so only CP ranks (if any) are combined here.
464+ if _AUX_LOSS_GROUP and _AUX_LOSS_GROUP_SIZE > 1:
465+ if "values" not in tracker:
466+ num_total_layers = num_layers
467+ if mtp_num_layers:
468+ num_total_layers += mtp_num_layers
469+ tracker["values"] = mint.zeros(num_total_layers)
470+ for sub_group in _AUX_LOSS_GROUP:
471+ all_reduce(tracker["values"], op=ops.ReduceOp.SUM, group=sub_group)
472+ # Step 2: Combine the per-layer tracker across PP stages.
290 if pp_group_size and pp_group_size > 1:473 if pp_group_size and pp_group_size > 1:
291 if "values" not in tracker:474 if "values" not in tracker:
292- # This stage owns no MoE layers that contributed aux losses yet.
293- # Seed a zero vector sized to the global layer count so the
294- # all_reduce shape matches the stages that do hold values.
295 num_pp_layers = num_layers475 num_pp_layers = num_layers
296 if mtp_num_layers:476 if mtp_num_layers:
297 num_pp_layers += mtp_num_layers477 num_pp_layers += mtp_num_layers
@@ -301,9 +481,7 @@ def track_moe_metrics(
301 clear_aux_losses_tracker()481 clear_aux_losses_tracker()
302 return None482 return None
303 483 
304- # No MoE layer contributed aux losses (e.g. an all-dense model where every layer484+ # No MoE layer contributed aux losses.
305- # is below first_k_dense_replace). The tracker is empty, so there is nothing to
306- # report; returning None makes the loss callback simply omit load_balancing_loss.
307 if "values" not in tracker:485 if "values" not in tracker:
308 return None486 return None
309 487 
@@ -323,6 +501,7 @@ def track_moe_metrics(
323 num_moe_layers += mtp_num_layers501 num_moe_layers += mtp_num_layers
324 502 
325 aux_losses = tracker["values"].sum() / num_moe_layers503 aux_losses = tracker["values"].sum() / num_moe_layers
504+ 
326 if group_size is None:505 if group_size is None:
327 group_size = get_world_size()506 group_size = get_world_size()
328 if group_size > 1:507 if group_size > 1:
@@ -19,9 +19,13 @@ from .moe_utils import (
19 switch_load_balancing_loss_func,19 switch_load_balancing_loss_func,
20 save_to_aux_losses_tracker,20 save_to_aux_losses_tracker,
21 MoEAuxLossAutoScaler,21 MoEAuxLossAutoScaler,
22+ get_tokens_per_expert_and_token_count,
23+ get_moe_aux_loss_group,
24+ get_moe_aux_loss_group_size,
22)25)
23 26 
24 27 
28+ 
25class TopKRouter(nn.Cell):29class TopKRouter(nn.Cell):
26 """This class implements token-choice routing. In token-choice top-K routing, each token is30 """This class implements token-choice routing. In token-choice top-K routing, each token is
27 routed to top K experts based on the router scores.31 routed to top K experts based on the router scores.
@@ -531,23 +535,33 @@ class TopKRouter(nn.Cell):
531 Returns:535 Returns:
532 Tensor: top_scores with aux loss gradient injected.536 Tensor: top_scores with aux loss gradient injected.
533 """537 """
534- # NOTE: AllReduce tokens_per_expert over tp_cp_group for multi-card538+ scores_for_aux_loss = scores_for_aux_loss.reshape(seq_length, -1)
535- scores_for_aux_loss = scores_for_aux_loss.reshape((seq_length, -1),)539+ routing_map = routing_map.reshape(seq_length, -1)
536- tokens_per_expert = routing_map.reshape((seq_length, -1),).sum(dim=0)
537- total_num_tokens = seq_length
538 540 
539- aux_loss = (541+ # ``topk * bsz`` -- the routing_map has shape ``(slen, bsz*E)`` after
540- switch_load_balancing_loss_func(542+ # the reshape (one row per sequence, ``bsz`` experts' worth of columns),
541- probs=scores_for_aux_loss,543+ # so each token contributes ``topk * bsz`` to the histogram sum. Used
542- tokens_per_expert=tokens_per_expert,544+ # by ``get_tokens_per_expert_and_token_count`` to derive the token
543- total_num_tokens=total_num_tokens,545+ # count from the histogram when padding is masked; kept for parity
544- topk=self.top_k,546+ # with the Megatron helper signature even though we currently only
545- num_experts=self.config.num_moe_experts,547+ # support the non-padded path.
546- moe_aux_loss_coeff=self.moe_aux_loss_coeff,548+ global_tokens_per_expert, _, total_num_tokens = (
547- )549+ get_tokens_per_expert_and_token_count(
548- / bsz550+ routing_map=routing_map,
551+ reduce_group=get_moe_aux_loss_group(),
552+ topk=self.top_k * bsz,
553+ )
549 )554 )
550 555 
556+ aux_loss = switch_load_balancing_loss_func(
557+ probs=scores_for_aux_loss,
558+ tokens_per_expert=global_tokens_per_expert,
559+ total_num_tokens=total_num_tokens,
560+ topk=self.top_k,
561+ num_experts=self.config.num_moe_experts,
562+ moe_aux_loss_coeff=self.moe_aux_loss_coeff,
563+ ) / bsz
564+ 
551 top_scores = self._attach_and_log_aux_loss(565 top_scores = self._attach_and_log_aux_loss(
552 top_scores, aux_loss, self.moe_aux_loss_coeff566 top_scores, aux_loss, self.moe_aux_loss_coeff
553 )567 )
@@ -575,15 +589,17 @@ class TopKRouter(nn.Cell):
575 Tensor: top_scores with aux loss gradient injected.589 Tensor: top_scores with aux loss gradient injected.
576 """590 """
577 _ = seq_length, bsz591 _ = seq_length, bsz
578- tokens_per_expert = routing_map.sum(dim=0)592+ global_tokens_per_expert, _, total_num_tokens = (
579- self.global_tokens_per_expert += tokens_per_expert593+ get_tokens_per_expert_and_token_count(
594+ routing_map=routing_map,
595+ reduce_group=get_moe_aux_loss_group(),
W
Wwei_zhuoyi6月30日

在 _apply_global_aux_loss 中,self.global_tokens_per_expert 和 self.ga_steps 的累加操作缺少 is_in_recompute() 守卫。recompute 期间 Router 的 forward 被重新执行,这两个累加会被重复计数。虽然 reset_global_aux_loss_tracker() 每步会重置,但在 grad accumulation 中,一个 micro-batch 内触发 recompute 会导致本 micro-batch 的 EMA (averated_tokens_per_expert = global_tokens_per_expert / ga_steps) 偏移,进而影响 global_aux_loss 的计算。

是否需要在此两行外加上 if not is_in_recompute(): 守卫?

likedislike
niujunhao
6月30日 评论:
596+ topk=self.top_k,
597+ )
598+ )
599+ self.global_tokens_per_expert += global_tokens_per_expert
580 self.ga_steps += 1600 self.ga_steps += 1
581 averated_tokens_per_expert = self.global_tokens_per_expert / self.ga_steps601 averated_tokens_per_expert = self.global_tokens_per_expert / self.ga_steps
582 602 
583- num_tokens = scores_for_aux_loss.shape[0]
584- # total_num_tokens = num_tokens * self.tp_dp_cp_group.size()
585- total_num_tokens = num_tokens
586- 
587 global_aux_loss = switch_load_balancing_loss_func(603 global_aux_loss = switch_load_balancing_loss_func(
588 probs=scores_for_aux_loss,604 probs=scores_for_aux_loss,
589 tokens_per_expert=averated_tokens_per_expert,605 tokens_per_expert=averated_tokens_per_expert,
@@ -621,6 +637,9 @@ class TopKRouter(nn.Cell):
621 if self.config.mtp_num_layers is not None:637 if self.config.mtp_num_layers is not None:
622 num_layers += self.config.mtp_num_layers638 num_layers += self.config.mtp_num_layers
623 639 
640+ # Compensate for the tp×cp-dependent scaling in the Switch formula.
641+ tp_cp_size = get_moe_aux_loss_group_size()
642+ 
624 save_to_aux_losses_tracker(643 save_to_aux_losses_tracker(
625 aux_loss / aux_loss_coeff,644 aux_loss / aux_loss_coeff,
626 self.layer_number,645 self.layer_number,
@@ -635,9 +654,13 @@ class TopKRouter(nn.Cell):
635 # which scales both the main_loss gradient and aux_loss gradient by654 # which scales both the main_loss gradient and aux_loss gradient by
636 # 1/(num_local_tokens * dp_size * num_micro_batches) in finalize_model_grads function.655 # 1/(num_local_tokens * dp_size * num_micro_batches) in finalize_model_grads function.
637 # To correct this scaling, we need to scale the aux_loss by num_local_tokens here.656 # To correct this scaling, we need to scale the aux_loss by num_local_tokens here.
638- top_scores = self.moe_aux_loss_auto_scaler(top_scores, aux_loss * top_scores.shape[0])657+ top_scores = self.moe_aux_loss_auto_scaler(
658+ top_scores, aux_loss * top_scores.shape[0] * tp_cp_size
659+ )
639 else:660 else:
640- top_scores = self.moe_aux_loss_auto_scaler(top_scores, aux_loss)661+ top_scores = self.moe_aux_loss_auto_scaler(
662+ top_scores, aux_loss * tp_cp_size
663+ )
641 664 
642 return top_scores665 return top_scores
643 666 
@@ -44,6 +44,7 @@ from mindformers.pynative.transformers.transformer_block import (
44from mindformers.parallel_core.utils.spec_utils import ModuleSpec, build_module44from mindformers.parallel_core.utils.spec_utils import ModuleSpec, build_module
45from mindformers.parallel_core.transformer_config import TransformerConfig45from mindformers.parallel_core.transformer_config import TransformerConfig
46from mindformers.pynative.layers.linear import Linear46from mindformers.pynative.layers.linear import Linear
47+from mindformers.pynative.distributed.activation_checkpoint import is_in_recompute
47 48 
48# MTP logging49# MTP logging
49_MTP_LAYER_WISE_LOGGING_TRACKER: dict = {}50_MTP_LAYER_WISE_LOGGING_TRACKER: dict = {}
@@ -119,6 +120,12 @@ def save_to_mtp_losses_tracker(
119 if layer_number is None:120 if layer_number is None:
120 return121 return
121 122 
123+ # Skip during activation recompute: the recomputed forward would otherwise add
124+ # this layer's MTP loss to the tracker a second time, doubling mtp_*_loss for
125+ # recomputed MTP layers. Mirrors the guard in ``save_to_aux_losses_tracker``.
126+ if is_in_recompute():
127+ return
128+ 
122 tracker = get_mtp_layer_wise_logging_tracker()129 tracker = get_mtp_layer_wise_logging_tracker()
123 if not tracker:130 if not tracker:
124 tracker["values"] = mint.zeros(num_layers)131 tracker["values"] = mint.zeros(num_layers)