| @@ -71,3 +71,23 @@ class PyNativeDeepseekV3ForCausalLM(TrainModelMixin, DeepseekV3PreTrainedModel): | |||
| 71 | loss_mask=loss_mask, | 71 | loss_mask=loss_mask, |
| 72 | actual_seq_len=actual_seq_len | 72 | 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 | ||
| 29 | from mindspore import Tensor, dtype, nn, mint, ops | 29 | from mindspore import Tensor, dtype, nn, mint, ops |
| 30 | from mindspore.mint.distributed import all_reduce, get_world_size | 30 | from mindspore.mint.distributed import all_reduce, get_world_size |
| 31 | +from mindspore.graph.api import _no_grad | ||
| 31 | 32 | ||
| 32 | from mindformers.tools.logger import logger | 33 | from mindformers.tools.logger import logger |
| 33 | from mindformers.pynative.loss.loss import CrossEntropyLoss, ChunkCrossEntropyLoss | 34 | from mindformers.pynative.loss.loss import CrossEntropyLoss, ChunkCrossEntropyLoss |
| @@ -40,11 +41,19 @@ from mindformers.pynative.base_models.common.embeddings.rotary_pos_embedding imp | |||
| 40 | from mindformers.pynative.base_models.common.embeddings.yarn_rotary_pos_embedding import YarnRotaryEmbedding | 41 | from mindformers.pynative.base_models.common.embeddings.yarn_rotary_pos_embedding import YarnRotaryEmbedding |
| 41 | from mindformers.pynative.base_models.gpt.gpt_layer_specs import get_gpt_mtp_block_spec | 42 | from mindformers.pynative.base_models.gpt.gpt_layer_specs import get_gpt_mtp_block_spec |
| 42 | from mindformers.pynative.transformers.transformer_block import TransformerBlock, TransformerBlockSubmodules | 43 | from mindformers.pynative.transformers.transformer_block import TransformerBlock, TransformerBlockSubmodules |
| 43 | -from mindformers.pynative.transformers.multi_token_prediction import MultiTokenPredictionBlock, MTPLossAutoScaler | 44 | +from mindformers.pynative.transformers.multi_token_prediction import ( |
| 45 | + MultiTokenPredictionBlock, | ||
| 46 | + MTPLossAutoScaler, | ||
| 47 | + process_mtp_loss, | ||
| 48 | + track_mtp_metrics, | ||
| 49 | +) | ||
| 44 | from mindformers.pynative.layers.linear import Linear | 50 | from mindformers.pynative.layers.linear import Linear |
| 45 | from mindformers.pynative.optimizer.muon_utils import make_muon_fns | 51 | from mindformers.pynative.optimizer.muon_utils import make_muon_fns |
| 46 | -from mindformers.pynative.transformers.multi_token_prediction import process_mtp_loss | ||
| 47 | from mindformers.pynative.dtensor_compat import inplace_copy | 52 | from 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 | ||
| 50 | class GPTModel(nn.Cell): | 59 | class 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_mask | 647 | 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 | |||
| 80 | from mindformers.pynative.transformers.experimental_attention_variant.deepseek_v4_hybrid_attention import ( | 80 | from 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 | ||
| 83 | from mindformers.tools.logger import logger | 84 | from 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 is | 122 | + 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 stage | 123 | + 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 MTP | 124 | + single-card baseline where one tensor receives ``main_grad + mtp_grad``, |
| 124 | - tokens). The single-card baseline keeps ONE embedding tensor that receives | 125 | + 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 (per | 128 | 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`` with | 129 | dp/tp/cp/ep coordinate) and tags the local ``word_embeddings.weight`` with |
| 132 | ``_embedding_grad_sync_group`` / ``_pp_replica_count``. The grad-norm helper | 130 | ``_embedding_grad_sync_group`` / ``_pp_replica_count``. The grad-norm helper |
| 133 | (`_calculate_global_grad_norm` -> `_get_grad_factor`) then all-reduces the | 131 | (`_calculate_global_grad_norm` -> `_get_grad_factor`) then all-reduces the |
| 134 | - tagged gradient before computing the norm/step and counts the replicated | 132 | + 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 | return | 143 | 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 | return | 152 | return |
| 144 | 153 | ||
| 145 | - # Locate this rank's input-embedding weight (present on stage 0 and the MTP | 154 | + # 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 = group | 192 | 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 | + | ||
| 1543 | def _apply_spmd_parallelism( | 1621 | def _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 the | 1731 | + # 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 the | 1732 | + # 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 model | 1737 | 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 @@ | |||
| 17 | import time | 17 | import time |
| 18 | from typing import Dict, Any | 18 | from typing import Dict, Any |
| 19 | from copy import deepcopy | 19 | from copy import deepcopy |
| 20 | -import os | ||
| 21 | 20 | ||
| 22 | -from mindspore import ops | ||
| 23 | from mindspore.nn.learning_rate_schedule import LearningRateSchedule | 21 | from mindspore.nn.learning_rate_schedule import LearningRateSchedule |
| 24 | -from mindspore.mint.distributed import get_world_size, all_reduce | 22 | +from mindspore.mint.distributed import get_world_size |
| 25 | 23 | ||
| 26 | from mindformers.pynative.callback.callback import TrainerCallback | 24 | from mindformers.pynative.callback.callback import TrainerCallback |
| 27 | from mindformers.tools.logger import logger | 25 | from 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 | ||
| 37 | class LossCallback(TrainerCallback): | 32 | class LossCallback(TrainerCallback): |
| @@ -53,6 +48,13 @@ class LossCallback(TrainerCallback): | |||
| 53 | self.log_interval = log_interval | 48 | 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) |
🔵 Low Priority 变更在 证据链:
建议:将 ![]() ![]() 不准确? | |||
| 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 of | 140 | # Update auxiliary-loss-free expert_bias on every step, regardless of |
| 137 | # log interval or loss availability. Non-last PP stages return loss=None | 141 | # log interval or loss availability. Non-last PP stages return loss=None |
| 138 | # but still hold MoE layers whose ``tokens_per_expert`` accumulators must | 142 | # 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 sync | 143 | # 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_steps | 173 | 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 | return | 176 | 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 loss | 185 | # 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 on | 188 | + mtp_loss = model[0].get_mtp_loss(metric_group, metric_group_size) |
| 178 | - # the config value also avoids the element-wise truthiness of the Tensor | 189 | + 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 = None | 199 | mtp_loss = None |
| 189 | 200 | ||
| 190 | # process indexer loss | 201 | # 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_steps | 209 | 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 | |||
| 20 | from mindformers.parallel_core.transformer_config import TransformerConfig | 20 | from mindformers.parallel_core.transformer_config import TransformerConfig |
| 21 | from mindformers.pynative.layers.linear import Linear | 21 | from mindformers.pynative.layers.linear import Linear |
| 22 | from mindformers.pynative.transformers.mlp import MLPSubmodules | 22 | from mindformers.pynative.transformers.mlp import MLPSubmodules |
| 23 | +from mindformers.pynative.distributed.activation_checkpoint import is_in_recompute | ||
| 23 | from .router import TopKRouter | 24 | from .router import TopKRouter |
| 24 | from .experts import GroupedMLP | 25 | from .experts import GroupedMLP |
| 25 | from .shared_experts import SharedExpertMLP | 26 | from .shared_experts import SharedExpertMLP |
| @@ -125,7 +126,8 @@ class MoELayer(nn.Cell): | |||
| 125 | hidden_states, self.expert_bias, input_ids | 126 | 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_expert | 133 | 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 logging | 31 | # 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 | ||
| 35 | def switch_load_balancing_loss_func( | 40 | def switch_load_balancing_loss_func( |
| 36 | probs: Tensor, | 41 | probs: Tensor, |
| @@ -104,6 +109,173 @@ def switch_load_balancing_loss_func( | |||
| 104 | return aux_loss | 109 | 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 | + | ||
| 107 | class _MoEAuxLossAutoScaler(_Function): | 279 | class _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] + loss | 415 | + cur_val = tracker["values"][layer_number] + loss |
| 244 | - | 416 | + tracker["values"][layer_number] = cur_val |
| 245 | 417 | ||
| 246 | def clear_aux_losses_tracker() -> None: | 418 | def 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_layers | 475 | num_pp_layers = num_layers |
| 296 | if mtp_num_layers: | 476 | if mtp_num_layers: |
| 297 | num_pp_layers += mtp_num_layers | 477 | 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 None | 482 | return None |
| 303 | 483 | ||
| 304 | - # No MoE layer contributed aux losses (e.g. an all-dense model where every layer | 484 | + # 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 None | 486 | return None |
| 309 | 487 | ||
| @@ -323,6 +501,7 @@ def track_moe_metrics( | |||
| 323 | num_moe_layers += mtp_num_layers | 501 | num_moe_layers += mtp_num_layers |
| 324 | 502 | ||
| 325 | aux_losses = tracker["values"].sum() / num_moe_layers | 503 | 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 | + | ||
| 25 | class TopKRouter(nn.Cell): | 29 | class TopKRouter(nn.Cell): |
| 26 | """This class implements token-choice routing. In token-choice top-K routing, each token is | 30 | """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-card | 538 | + 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 | - / bsz | 550 | + 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_coeff | 566 | 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, bsz | 591 | _ = 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_expert | 593 | + get_tokens_per_expert_and_token_count( |
| 594 | + routing_map=routing_map, | ||
| 595 | + reduce_group=get_moe_aux_loss_group(), | ||
W 在 是否需要在此两行外加上 ![]() ![]() | |||
| 596 | + topk=self.top_k, | ||
| 597 | + ) | ||
| 598 | + ) | ||
| 599 | + self.global_tokens_per_expert += global_tokens_per_expert | ||
| 580 | self.ga_steps += 1 | 600 | self.ga_steps += 1 |
| 581 | averated_tokens_per_expert = self.global_tokens_per_expert / self.ga_steps | 601 | 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_layers | 638 | 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 by | 654 | # 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_scores | 665 | return top_scores |
| 643 | 666 | ||
| @@ -44,6 +44,7 @@ from mindformers.pynative.transformers.transformer_block import ( | |||
| 44 | from mindformers.parallel_core.utils.spec_utils import ModuleSpec, build_module | 44 | from mindformers.parallel_core.utils.spec_utils import ModuleSpec, build_module |
| 45 | from mindformers.parallel_core.transformer_config import TransformerConfig | 45 | from mindformers.parallel_core.transformer_config import TransformerConfig |
| 46 | from mindformers.pynative.layers.linear import Linear | 46 | from mindformers.pynative.layers.linear import Linear |
| 47 | +from mindformers.pynative.distributed.activation_checkpoint import is_in_recompute | ||
| 47 | 48 | ||
| 48 | # MTP logging | 49 | # 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 | return | 121 | 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) |


🔵 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 之前紧邻的位置。