diff --git a/megatron/core/transformer/cuda_graphs.py b/megatron/core/transformer/cuda_graphs.py index baa295b8ae8..0de90c9cde4 100644 --- a/megatron/core/transformer/cuda_graphs.py +++ b/megatron/core/transformer/cuda_graphs.py @@ -859,13 +859,12 @@ def create_fwd_graph(self, args, kwargs, outputs=None, clone_inputs=True): is_moe = isinstance(self.base_module, MoETransformerLayer) if is_moe: - from megatron.core.transformer.moe.moe_utils import get_moe_layer_wise_logging_tracker + from megatron.core.transformer.moe.moe_logging import get_moe_metrics_tracker - tracker = get_moe_layer_wise_logging_tracker() + moe_metrics_tracker = get_moe_metrics_tracker() cached_aux_losses = {} - for name in tracker: - if "values" in tracker[name]: - cached_aux_losses[name] = torch.clone(tracker[name]["values"]) + for name, entry in moe_metrics_tracker.metrics.items(): + cached_aux_losses[name] = entry.values.clone() self.fwd_graph = torch.cuda.CUDAGraph() @@ -1047,8 +1046,11 @@ def clone_ten(ten): buf.copy_(buf_copy) if is_moe: - for name in tracker: - tracker[name]["values"].copy_(cached_aux_losses[name]) + for name, cached_values in cached_aux_losses.items(): + assert ( + name in moe_metrics_tracker.metrics + ), "cached metrics must be found in the tracker." + moe_metrics_tracker.metrics[name].values.copy_(cached_values) def create_bwd_graph(self): """Create a bwd cudagraph for this runner. Should be called inside @@ -2334,13 +2336,13 @@ def _reset_after_capture(self): Reset the model and optimizer state after capturing CUDA Graphs. """ from megatron.core.distributed.finalize_model_grads import reset_model_temporary_tensors - from megatron.core.transformer.moe.moe_utils import clear_aux_losses_tracker + from megatron.core.transformer.moe.moe_logging import get_moe_metrics_tracker for model_chunk in self.model: model_chunk.zero_grad_buffer() for optimizer in self.optimizers: optimizer.zero_grad() - clear_aux_losses_tracker() + get_moe_metrics_tracker().clear() reset_model_temporary_tensors(self.config, self.model) def _finish_capturing(self, start_time): diff --git a/megatron/core/transformer/moe/moe_logging.py b/megatron/core/transformer/moe/moe_logging.py new file mode 100644 index 00000000000..b1f2b27000b --- /dev/null +++ b/megatron/core/transformer/moe/moe_logging.py @@ -0,0 +1,379 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""MoE metrics tracking and logging. + +Collects per-layer MoE metrics during forward passes, synchronizes them across +distributed ranks, and writes scalar summaries to TensorBoard / W&B. + +Usage: + tracker = get_moe_metrics_tracker() + + # In router forward pass: + tracker.record("load_balancing_loss", loss, layer_number=1, num_layers=32, + reduce_group=tp_cp_group) + + # At end of training step: + log_str = tracker.report( + loss_scale=1 / num_microbatches, + iteration=step, + writer=tb_writer, + num_layers=32, + ) +""" + +from dataclasses import dataclass +from typing import Dict, List, Optional, Union + +import torch + +from megatron.core import parallel_state +from megatron.core.process_groups_config import ProcessGroupCollection + + +@dataclass +class MetricEntry: + """Per-layer metric with distributed reduction configuration.""" + + values: torch.Tensor + reduce_group: Optional[torch.distributed.ProcessGroup] = None + avg_group: Optional[torch.distributed.ProcessGroup] = None + needs_dp_avg: bool = True + + +# --------------------------------------------------------------------------- +# Module-level global tracker (follows parallel_state / global_vars pattern) +# --------------------------------------------------------------------------- +_MOE_METRICS_TRACKER: Optional['MoEMetricsTracker'] = None + + +def get_moe_metrics_tracker() -> 'MoEMetricsTracker': + """Return the global MoE metrics tracker, creating it lazily if needed.""" + global _MOE_METRICS_TRACKER + if _MOE_METRICS_TRACKER is None: + _MOE_METRICS_TRACKER = MoEMetricsTracker() + return _MOE_METRICS_TRACKER + + +def set_moe_metrics_tracker(tracker: 'MoEMetricsTracker') -> None: + """Set the global MoE metrics tracker.""" + global _MOE_METRICS_TRACKER + _MOE_METRICS_TRACKER = tracker + + +def destroy_moe_metrics_tracker() -> None: + """Reset the global MoE metrics tracker to ``None``.""" + global _MOE_METRICS_TRACKER + _MOE_METRICS_TRACKER = None + + +class MoEMetricsTracker: + """Tracker for MoE layer-wise metrics. + + Lifecycle: ``record()`` per-layer values during forward → ``report()`` at + step end (sync, aggregate, log, clear) → repeat. + + Example: + tracker = get_moe_metrics_tracker() + tracker.record("load_balancing_loss", loss, layer_number=1, num_layers=32) + log_str = tracker.report(loss_scale=1/8, iteration=100, writer=tb_writer, + num_layers=32) + """ + + def __init__(self): + self._metrics: Dict[str, MetricEntry] = {} + + # ========================================================================= + # Public API + # ========================================================================= + + @property + def metrics(self) -> Dict[str, MetricEntry]: + """Read-only access to the underlying metric entries.""" + return self._metrics + + def record( + self, + name: str, + value: torch.Tensor, + layer_number: int, + num_layers: int, + reduce_group: Optional[torch.distributed.ProcessGroup] = None, + avg_group: Optional[torch.distributed.ProcessGroup] = None, + needs_dp_avg: bool = True, + ) -> None: + """Accumulate a metric value for a specific layer. + + Called during the router forward pass. Lazily creates the metric entry + on first call for each metric name. + + Args: + name: Metric name (e.g. ``"load_balancing_loss"``). + value: Scalar tensor to accumulate (will be detached). + layer_number: 1-based layer index. + num_layers: Total number of layers (determines tensor size). + reduce_group: Process group for sum-reduction (e.g. tp_cp_group). + avg_group: Process group for average-reduction. + needs_dp_avg: Whether to average across DP ranks after other reductions. + """ + if layer_number is None: + return + + if name not in self._metrics: + self._metrics[name] = MetricEntry(values=torch.zeros(num_layers, device=value.device)) + + entry = self._metrics[name] + entry.values[layer_number - 1] += value.detach() + entry.reduce_group = reduce_group + entry.avg_group = avg_group + entry.needs_dp_avg = needs_dp_avg + + def report( + self, + loss_scale: float, + iteration: int, + writer=None, + wandb_writer=None, + per_layer_logging: bool = False, + force_initialize: bool = False, + track_names: Optional[Union[str, List[str]]] = None, + num_layers: Optional[int] = None, + moe_layer_freq: Optional[Union[int, List[int]]] = None, + mtp_num_layers: Optional[int] = None, + total_loss_dict: Optional[dict[str, torch.Tensor]] = None, + percentiles: Optional[Dict[str, List[float]]] = None, + pg_collection: Optional[ProcessGroupCollection] = None, + ) -> str: + """Sync metrics across ranks, aggregate, log, and clear. + + This is the main entry point called once per training step. It pairs + with :meth:`record`: you *record* individual data points during forward, + then *report* the summary at step end. + + Args: + loss_scale: Scale factor for averaging across microbatches + (usually ``1 / num_microbatches``). + iteration: Current training iteration. + writer: TensorBoard ``SummaryWriter`` (optional). + wandb_writer: Weights & Biases run object (optional). + per_layer_logging: Whether to also write per-layer values. + force_initialize: If True, pre-create metric entries for *track_names* + that don't exist yet. Required for PP ranks without MoE layers + whose tensor sizes must match ranks that do have MoE layers. + track_names: Metric name(s) to report. ``None`` reports all. + num_layers: Total transformer layers (required when *force_initialize*). + moe_layer_freq: MoE layer frequency or binary pattern list. + mtp_num_layers: Extra layers from Multi-Token Prediction. + total_loss_dict: Megatron training-loop accumulator. Metrics + ending with ``"loss"`` are accumulated here and excluded from + the returned console log string. + percentiles: Per-metric percentiles to compute, e.g. + ``{"load_imbalance": [0.5, 0.95]}``. + pg_collection: Custom process-group collection for reduction. + + Returns: + Formatted log string for console output. + """ + metric_names = self._resolve_names(track_names) + + # Pre-create entries on PP ranks that lack MoE layers. + # Tensor size must be (num_layers + mtp_num_layers) to match ranks that + # recorded via record(), otherwise all_reduce across PP will hang. + if force_initialize: + if num_layers is None: + raise ValueError("num_layers must be provided when force_initialize=True.") + init_size = num_layers + (mtp_num_layers or 0) + for name in metric_names: + self.ensure_initialized(name, init_size) + + self._sync_metrics(metric_names, pg_collection) + + num_moe_layers = self._count_moe_layers(num_layers, moe_layer_freq, mtp_num_layers) + scalars = self._aggregate(loss_scale, num_moe_layers, metric_names, percentiles) + + # Megatron integration: accumulate loss metrics into total_loss_dict + console_scalars = dict(scalars) + if total_loss_dict is not None: + for k, v in scalars.items(): + if k.lower().endswith("loss"): + if k in total_loss_dict: + total_loss_dict[k] += v + else: + total_loss_dict[k] = v + console_scalars.pop(k) + + self._log_scalars(scalars, iteration, writer, wandb_writer) + if per_layer_logging: + self._log_per_layer( + loss_scale, metric_names, iteration, writer, wandb_writer, percentiles + ) + + log_string = self._format(console_scalars) + self.clear() + return log_string + + def clear(self) -> None: + """Zero out all metric values (entries are kept for reuse).""" + for entry in self._metrics.values(): + entry.values.zero_() + + def ensure_initialized( + self, name: str, num_layers: int, device: Optional[Union[str, torch.device, int]] = None + ) -> None: + """Pre-create a metric entry if it does not already exist. + + This is needed for PP ranks that have no MoE layers -- their tensor + size must match ranks that do, otherwise ``all_reduce`` across PP hangs. + + Args: + name: Metric name. + num_layers: Tensor size (should include MTP layers). + device: Device for the zero tensor. Defaults to current CUDA device. + """ + if name not in self._metrics: + if device is None: + device = torch.cuda.current_device() if torch.cuda.is_available() else "cpu" + self._metrics[name] = MetricEntry(values=torch.zeros(num_layers, device=device)) + + # ========================================================================= + # Private implementation + # ========================================================================= + + def _resolve_names(self, track_names: Optional[Union[str, List[str]]]) -> List[str]: + """Normalize *track_names* argument to a list of strings.""" + if track_names is None: + return list(self._metrics.keys()) + if isinstance(track_names, str): + return [track_names] + return track_names + + def _sync_metrics( + self, metric_names: List[str], pg_collection: Optional[ProcessGroupCollection] = None + ) -> None: + """All-reduce metrics across distributed ranks. + + Reduction order: PP collect → reduce_group sum → avg_group avg → DP avg. + """ + if pg_collection is None: + pp_group = parallel_state.get_pipeline_model_parallel_group() + dp_group = parallel_state.get_data_parallel_group( + with_context_parallel=False, partial_data_parallel=False + ) + else: + pp_group = pg_collection.pp + dp_group = pg_collection.dp + + for name in metric_names: + if name not in self._metrics: + continue + + entry = self._metrics[name] + v = entry.values + + torch.distributed.all_reduce(v, group=pp_group) + + if entry.reduce_group is not None: + torch.distributed.all_reduce(v, group=entry.reduce_group) + + if entry.avg_group is not None: + torch.distributed.all_reduce( + v, group=entry.avg_group, op=torch.distributed.ReduceOp.AVG + ) + + if entry.needs_dp_avg: + torch.distributed.all_reduce(v, group=dp_group, op=torch.distributed.ReduceOp.AVG) + + @staticmethod + def _count_moe_layers( + num_layers: Optional[int], + moe_layer_freq: Optional[Union[int, List[int]]], + mtp_num_layers: Optional[int], + ) -> int: + """Compute the effective number of MoE layers from configuration.""" + if moe_layer_freq is None: + n = num_layers + elif isinstance(moe_layer_freq, int): + assert isinstance(num_layers, int) + n = sum(1 for i in range(num_layers) if i % moe_layer_freq == 0) + elif isinstance(moe_layer_freq, list): + n = sum(moe_layer_freq) + else: + raise ValueError(f"Invalid moe_layer_freq: {moe_layer_freq}") + + if mtp_num_layers is not None: + n += mtp_num_layers + + return n + + def _aggregate( + self, + loss_scale: float, + num_moe_layers: int, + metric_names: List[str], + percentiles: Optional[Dict[str, List[float]]] = None, + ) -> Dict[str, Union[float, torch.Tensor]]: + """Aggregate per-layer values into scalar summaries. + + Always computes the mean across MoE layers. If *percentiles* specifies + quantiles for a metric, those are computed over non-zero layer values and + added as ``"{name}_p{pct}"`` keys. + """ + result: Dict[str, Union[float, torch.Tensor]] = {} + + for name in metric_names: + if name not in self._metrics: + continue + + values = self._metrics[name].values.float() * loss_scale + + if percentiles and name in percentiles: + nonzero = values[values > 0] + if nonzero.numel() > 0: + pcts = percentiles[name] + pct_vals = torch.quantile( + nonzero, torch.tensor(pcts, device=nonzero.device) + ).tolist() + for pct, pct_val in zip(pcts, pct_vals): + result[f"{name}_p{int(pct * 100)}"] = pct_val + + result[name] = values.sum() / num_moe_layers + + return result + + def _log_scalars( + self, scalars: Dict[str, Union[float, torch.Tensor]], iteration: int, writer, wandb_writer + ) -> None: + """Write scalar metrics to TensorBoard and/or W&B.""" + for name, value in scalars.items(): + if writer is not None: + writer.add_scalar(name, value, iteration) + if wandb_writer is not None: + wandb_writer.log({name: value}, iteration) + + def _log_per_layer( + self, + loss_scale: float, + metric_names: List[str], + iteration: int, + writer, + wandb_writer, + percentiles: Optional[Dict[str, List[float]]] = None, + ) -> None: + """Write per-layer metric values to TensorBoard and/or W&B.""" + for name in metric_names: + if name not in self._metrics: + continue + + values = self._metrics[name].values.float() * loss_scale + is_sparse = percentiles is not None and name in percentiles + for i, val in enumerate(values.tolist()): + if is_sparse and val == 0: + continue + if writer is not None: + writer.add_scalar(f"moe/{name}_layer_{i}", val, iteration) + if wandb_writer is not None: + wandb_writer.log({f"moe/{name}_layer_{i}": val}, iteration) + + @staticmethod + def _format(scalars: Dict[str, Union[float, torch.Tensor]]) -> str: + """Format aggregated metrics as a console log string.""" + return "".join(f" {k}: {v:.2f} |" for k, v in scalars.items()) diff --git a/megatron/core/transformer/moe/moe_utils.py b/megatron/core/transformer/moe/moe_utils.py index fbadcb7d3da..ca004bd10b1 100644 --- a/megatron/core/transformer/moe/moe_utils.py +++ b/megatron/core/transformer/moe/moe_utils.py @@ -20,9 +20,10 @@ from megatron.core.tensor_parallel.mappings import reduce_from_tensor_model_parallel_region from megatron.core.transformer.cuda_graphs import is_graph_capturing from megatron.core.transformer.enums import CudaGraphModule +from megatron.core.transformer.moe.moe_logging import get_moe_metrics_tracker from megatron.core.transformer.moe.router_replay import RouterReplay from megatron.core.transformer.transformer_config import TransformerConfig -from megatron.core.utils import internal_api, is_te_min_version +from megatron.core.utils import deprecated, internal_api, is_te_min_version if HAVE_TE: from megatron.core.extensions.transformer_engine import ( @@ -52,10 +53,6 @@ ) = (None, None, None, None, None, None, None, None, None, None) -# MOE logging -_MOE_LAYER_WISE_LOGGING_TRACKER: dict = {} - - def switch_load_balancing_loss_func( probs: torch.Tensor, tokens_per_expert: torch.Tensor, @@ -967,6 +964,9 @@ def apply_router_token_dropping( return final_probs, final_map +@deprecated( + version="0.16", removal_version="0.18", alternative="get_moe_metrics_tracker().record()" +) def save_to_aux_losses_tracker( name: str, loss: torch.Tensor, @@ -983,38 +983,36 @@ def save_to_aux_losses_tracker( layer_number (int): Layer index of the loss. num_layers (int): The number of total layers. reduce_group (torch.distributed.ProcessGroup, optional): The group for reducing the loss. - Defaults to None. + Defaults to None. avg_group (torch.distributed.ProcessGroup, optional): The group for averaging the loss. - Defaults to None. - reduce_group_has_dp (bool, optional): Whether the reduce group has data parallel ranks. - Set this to True if the reduce group has data parallel ranks. This flag is used to - ensure the correct reduction in aux loss tracking. Defaults to False. + Defaults to None. + reduce_group_has_dp (bool, optional): Whether the reduce group already includes DP ranks. + If True, DP averaging is skipped. Defaults to False. """ - # Skip aux loss logging if layer_number is None. - if layer_number is None: - return - - tracker = get_moe_layer_wise_logging_tracker() - if name not in tracker: - tracker[name] = {} - tracker[name]["values"] = torch.zeros(num_layers, device=loss.device) - tracker[name]["values"][layer_number - 1] += loss.detach() # Aggregate the loss for the layer. - tracker[name]["reduce_group"] = reduce_group - tracker[name]["avg_group"] = avg_group - tracker[name]["reduce_group_has_dp"] = reduce_group_has_dp + get_moe_metrics_tracker().record( + name=name, + value=loss, + layer_number=layer_number, + num_layers=num_layers, + reduce_group=reduce_group, + avg_group=avg_group, + needs_dp_avg=not reduce_group_has_dp, + ) +@deprecated(version="0.16", removal_version="0.18", alternative="get_moe_metrics_tracker().clear()") def clear_aux_losses_tracker() -> None: """Clear the auxiliary losses.""" - tracker = get_moe_layer_wise_logging_tracker() - for name in tracker: - tracker[name]["values"].zero_() + get_moe_metrics_tracker().clear() +@deprecated( + version="0.16", removal_version="0.18", alternative="get_moe_metrics_tracker()._sync_metrics()" +) def reduce_aux_losses_tracker_across_ranks( track_names: Optional[List[str]] = None, pg_collection: Optional[ProcessGroupCollection] = None ) -> None: - """Collect and reduce the auxiliary losses across ranks. + """Reduce the auxiliary losses across ranks. Args: track_names (Optional[List[str]], optional): @@ -1022,40 +1020,28 @@ def reduce_aux_losses_tracker_across_ranks( pg_collection (Optional[ProcessGroupCollection], optional): The process group collection. Defaults to None. """ - tracker = get_moe_layer_wise_logging_tracker() - if track_names is None: - track_names = tracker.keys() - - if pg_collection is None: - # Use parallel_state groups - pp_group = parallel_state.get_pipeline_model_parallel_group() - dp_group = parallel_state.get_data_parallel_group( - with_context_parallel=False, partial_data_parallel=False - ) - else: - pp_group = pg_collection.pp - dp_group = pg_collection.dp - - for name in track_names: - values = tracker[name]["values"] - # TODO(Hepteract): delete the usage of the global parallel_state. - # Collect aux losses across PP. - torch.distributed.all_reduce(values, group=pp_group) - # Reduce aux losses across ranks. - if tracker[name].get('reduce_group') is not None: - torch.distributed.all_reduce(values, group=tracker[name].get('reduce_group')) - # Need to conduct reduction across data parallel ranks. When the reduce_group - # does not have 'dp' attribute, do it manually. - if not tracker[name].get('reduce_group_has_dp', False): - torch.distributed.all_reduce( - values, group=dp_group, op=torch.distributed.ReduceOp.AVG - ) - if tracker[name].get('avg_group') is not None: - torch.distributed.all_reduce( - values, group=tracker[name]['avg_group'], op=torch.distributed.ReduceOp.AVG - ) - - + tracker = get_moe_metrics_tracker() + names_list = track_names if track_names is not None else list(tracker.metrics.keys()) + tracker._sync_metrics(names_list, pg_collection) + + +@deprecated(version="0.16", removal_version="0.18", alternative="get_moe_metrics_tracker().metrics") +def get_moe_layer_wise_logging_tracker(): + """Return the moe layer wise tracker in legacy dict format.""" + return { + name: { + "values": entry.values, + "reduce_group": entry.reduce_group, + "avg_group": entry.avg_group, + "needs_dp_avg": entry.needs_dp_avg, + } + for name, entry in get_moe_metrics_tracker().metrics.items() + } + + +@deprecated( + version="0.15", removal_version="0.17", alternative="get_moe_metrics_tracker().report()" +) def track_moe_metrics( loss_scale: float, iteration: int, @@ -1069,95 +1055,25 @@ def track_moe_metrics( moe_layer_freq: Optional[Union[int, List[int]]] = None, mtp_num_layers: Optional[int] = None, pg_collection: Optional[ProcessGroupCollection] = None, -) -> None: +) -> str: """Track the MoE metrics for logging. - Args: - loss_scale (float): The loss scale. - iteration (int): The iteration. - writer (SummaryWriter, optional): The tensorboard writer. Defaults to None. - wandb_writer (wandb.Run, optional): The wandb writer. Defaults to None. - total_loss_dict (dict[str, torch.Tensor], optional): The total loss dictionary. - Defaults to None. - per_layer_logging (bool, optional): Whether to log per layer. Defaults to False. - force_initialize (bool, optional): Whether to force initialize the tracker. - Defaults to False. - track_names (List[str], optional): The names of the losses to track. Defaults to None. - num_layers (int, optional): The number of layers. Defaults to None. - moe_layer_freq (Union[int, List[int]], optional): The frequency of the MoE layers. - Defaults to None. - mtp_num_layers (int, optional): The number of layers in the model parallel group. - Defaults to None. - pg_collection (ProcessGroupCollection, optional): The process group collection. - Defaults to None. + Deprecated: Use get_moe_metrics_tracker().report() directly. """ - # Aux loss logging - tracker = get_moe_layer_wise_logging_tracker() - # Initialize the tracker if force_initialize is True. - # The values tensor size must match what the router creates in save_to_aux_losses_tracker, - # which uses (num_layers + mtp_num_layers). This is important for PP ranks that have no - # MoE layers (so the tracker is empty and force_initialize creates the entry); their tensor - # size must match ranks that do have MoE layers, otherwise all_reduce across PP will hang. - tracker_num_layers = num_layers - if mtp_num_layers is not None: - tracker_num_layers += mtp_num_layers - if force_initialize: - if track_names is not None: - for key in track_names: - if key not in tracker: - tracker[key] = {} - tracker[key]["values"] = torch.zeros(tracker_num_layers, device="cuda") - tracker[key]["reduce_group"] = None - tracker[key]["avg_group"] = None - tracker[key]["reduce_group_has_dp"] = False - reduce_aux_losses_tracker_across_ranks(track_names, pg_collection=pg_collection) - - # Get number of MoE layers - if moe_layer_freq is None: - num_moe_layers = num_layers - elif isinstance(moe_layer_freq, int): - assert isinstance(num_layers, int) - moe_layer_pattern = [1 if (i % moe_layer_freq == 0) else 0 for i in range(num_layers)] - num_moe_layers = sum(moe_layer_pattern) - elif isinstance(moe_layer_freq, list): - num_moe_layers = sum(moe_layer_freq) - else: - raise ValueError(f"Invalid moe_layer_freq: {moe_layer_freq}") - - if mtp_num_layers is not None: - num_moe_layers += mtp_num_layers - - aux_losses = {k: v['values'].float() * loss_scale for k, v in tracker.items()} - for name, loss_list in aux_losses.items(): - if total_loss_dict is not None: - if name not in total_loss_dict: - total_loss_dict[name] = loss_list.sum() / num_moe_layers - else: - total_loss_dict[name] += loss_list.sum() / num_moe_layers - if writer is not None: - # currently when using add_scalars, - # torch.utils.add_scalars makes each timer its own run, which - # polutes the runs list, so we just add each as a scalar - writer.add_scalar(name, loss_list.sum() / num_moe_layers, iteration) - if per_layer_logging: - for i, loss in enumerate(loss_list.tolist()): - writer.add_scalar(f"moe/{name}_layer_{i}", loss, iteration) - - # W&B logging lacks support for logging multiple scalars simultaneously. - # As a workaround, we log each scalar individually first, then we can create - # a custom panel to manually group them to a single plot. - if wandb_writer: - wandb_writer.log({f"{name}": loss_list.sum() / num_moe_layers}, iteration) - if per_layer_logging: - wandb_writer.log( - { - f"moe/{name}_layer_{i}": loss - for i, loss in enumerate(loss_list.tolist()) - }, - iteration, - ) - - clear_aux_losses_tracker() + return get_moe_metrics_tracker().report( + loss_scale=loss_scale, + iteration=iteration, + writer=writer, + wandb_writer=wandb_writer, + per_layer_logging=per_layer_logging, + force_initialize=force_initialize, + track_names=track_names, + num_layers=num_layers, + moe_layer_freq=moe_layer_freq, + mtp_num_layers=mtp_num_layers, + pg_collection=pg_collection, + total_loss_dict=total_loss_dict, + ) def get_updated_expert_bias( @@ -1218,12 +1134,6 @@ def maybe_move_tensor_to_cpu( return tensor -def get_moe_layer_wise_logging_tracker() -> dict: - """Return the moe layer wise tracker.""" - global _MOE_LAYER_WISE_LOGGING_TRACKER - return _MOE_LAYER_WISE_LOGGING_TRACKER - - @internal_api class RandomSTE(torch.autograd.Function): """ diff --git a/megatron/core/transformer/moe/router.py b/megatron/core/transformer/moe/router.py index 0348739be80..235a616dfbf 100644 --- a/megatron/core/transformer/moe/router.py +++ b/megatron/core/transformer/moe/router.py @@ -8,6 +8,7 @@ from megatron.core.inference.utils import InferenceMode from megatron.core.jit import jit_fuser from megatron.core.transformer.module import MegatronModule +from megatron.core.transformer.moe.moe_logging import get_moe_metrics_tracker from megatron.core.transformer.moe.moe_utils import ( MoEAuxLossAutoScaler, ProcessGroupCollection, @@ -17,7 +18,6 @@ compute_routing_scores_for_aux_loss, get_tokens_per_expert_and_token_count, router_gating_linear, - save_to_aux_losses_tracker, sinkhorn, switch_load_balancing_loss_func, topk_routing_with_score_function, @@ -419,7 +419,7 @@ def _apply_global_aux_loss( global_aux_loss, "global_load_balancing_loss", self.tp_dp_cp_group, - reduce_group_has_dp=True, + needs_dp_avg=False, valid_token_count=local_num_tokens, ) return probs @@ -431,7 +431,7 @@ def attach_and_log_load_balancing_loss( aux_loss: torch.Tensor, aux_loss_name: str, reduce_group: torch.distributed.ProcessGroup, - reduce_group_has_dp: bool = False, + needs_dp_avg: bool = True, valid_token_count: Optional[Union[int, torch.Tensor]] = None, ): """Attach aux loss function to activation and add to logging. @@ -442,9 +442,7 @@ def attach_and_log_load_balancing_loss( aux_loss (torch.Tensor): Computed aux loss. aux_loss_name (str): Name of the aux loss for logging. reduce_group (torch.distributed.ProcessGroup): Process group for reduction. - reduce_group_has_dp (bool): Whether the reduce group has data parallel ranks. - Set this to True if the reduce group has data parallel ranks. This flag is used to - ensure the correct reduction in aux loss tracking. + needs_dp_avg (bool): Whether to average this metric across DP ranks after reduce_group. valid_token_count (int or torch.Tensor, optional): Number of valid tokens excluding padding tokens. Can be a Python int or a torch.Tensor (typically 0-d tensor). If None, uses activation.shape[0]. Defaults to None. @@ -472,13 +470,13 @@ def attach_and_log_load_balancing_loss( else: layer_number = self.layer_number - save_to_aux_losses_tracker( + get_moe_metrics_tracker().record( aux_loss_name, aux_loss / aux_loss_coeff, layer_number, num_layers, reduce_group=reduce_group, - reduce_group_has_dp=reduce_group_has_dp, + needs_dp_avg=needs_dp_avg, ) if self.calculate_per_token_loss: # Scale the aux_loss by the number of tokens. @@ -545,7 +543,7 @@ def apply_z_loss(self, logits, padding_mask: Optional[torch.Tensor] = None): else: layer_number = self.layer_number - save_to_aux_losses_tracker( + get_moe_metrics_tracker().record( "z_loss", z_loss / moe_z_loss_coeff, layer_number, num_layers ) return logits diff --git a/megatron/training/training.py b/megatron/training/training.py index f69e4f30f6a..42216e0510d 100644 --- a/megatron/training/training.py +++ b/megatron/training/training.py @@ -214,7 +214,7 @@ def set_startup_timestamps(program_start=None, main_entry=None): from megatron.core.resharding.refit import swap_model_weights from megatron.core.transformer.experimental_attention_variant.dsa import DSAIndexerLossLoggingHelper from megatron.core.transformer.moe import upcycling_utils -from megatron.core.transformer.moe.moe_utils import clear_aux_losses_tracker, track_moe_metrics +from megatron.core.transformer.moe.moe_logging import get_moe_metrics_tracker from megatron.core.transformer.multi_token_prediction import MTPLossLoggingHelper from megatron.core.utils import get_batch_on_this_cp_rank, get_batch_on_this_tp_rank, unwrap_model from megatron.training.config import FaultInjectorConfig @@ -2525,8 +2525,8 @@ def training_log( writer.add_scalar('max_attention_logit', max_attention_logit, iteration) if wandb_writer: wandb_writer.log({'max_attention_logit': max_attention_logit}, iteration) - # Log MoE metrics. + moe_log_string = "" if args.num_experts is not None: moe_loss_scale = 1 / get_num_microbatches() track_names = [] @@ -2550,12 +2550,11 @@ def training_log( else: layers = args.num_layers - track_moe_metrics( + moe_log_string = get_moe_metrics_tracker().report( loss_scale=moe_loss_scale, iteration=iteration, writer=writer, wandb_writer=wandb_writer, - total_loss_dict=total_loss_dict, per_layer_logging=args.moe_per_layer_logging, force_initialize=True, track_names=track_names, @@ -2563,6 +2562,7 @@ def training_log( moe_layer_freq=args.moe_layer_freq, mtp_num_layers=args.mtp_num_layers, pg_collection=pg_collection, + total_loss_dict=total_loss_dict, ) # Log MTP metrics. @@ -2655,6 +2655,8 @@ def training_log( log_string += ' {}: {:.6E} |'.format(key, avg) if should_reset: total_loss_dict[key] = torch.tensor([0.0], dtype=torch.float, device='cuda') + if args.num_experts is not None and moe_log_string: + log_string += moe_log_string log_string += f' loss scale: {loss_scale:.1f} |' if grad_norm is not None: log_string += f' grad norm: {grad_norm:.3f} |' @@ -3728,7 +3730,7 @@ def trace_handler(p): if args.log_energy: energy_monitor.resume() if args.num_experts is not None: - clear_aux_losses_tracker() + get_moe_metrics_tracker().clear() # Miscellaneous post-training-step functions (e.g., FT heartbeats, GC). # Some of these only happen at specific iterations. Capture updated FLOPs accumulator diff --git a/tests/unit_tests/models/test_hybrid_moe_model.py b/tests/unit_tests/models/test_hybrid_moe_model.py index 59a6da45a1e..5aef44eccde 100644 --- a/tests/unit_tests/models/test_hybrid_moe_model.py +++ b/tests/unit_tests/models/test_hybrid_moe_model.py @@ -16,6 +16,7 @@ from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer import TransformerConfig from megatron.core.transformer.enums import AttnBackend +from megatron.core.transformer.moe.moe_logging import destroy_moe_metrics_tracker from megatron.training.arguments import core_transformer_config_from_args, parse_args, validate_args from megatron.training.global_vars import ( destroy_global_vars, @@ -509,6 +510,7 @@ def create_test_args(self): def setup_method(self, method): os.environ['CUDA_DEVICE_MAX_CONNECTIONS'] = '1' + destroy_moe_metrics_tracker() args = self.create_test_args() set_args(args)