diff --git a/megatron/core/fusions/fused_pad_routing_map.py b/megatron/core/fusions/fused_pad_routing_map.py index e7c3a7e48c9..8e4d1763270 100644 --- a/megatron/core/fusions/fused_pad_routing_map.py +++ b/megatron/core/fusions/fused_pad_routing_map.py @@ -1,9 +1,11 @@ -# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved. +# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + from unittest.mock import MagicMock import torch from packaging import version +from megatron.core.jit import jit_fuser from megatron.core.utils import experimental_fn, null_decorator try: @@ -69,6 +71,7 @@ def _pad_routing_map_kernel( @experimental_fn(introduced_with_version="0.13.0") +@jit_fuser def fused_pad_routing_map(routing_map: torch.Tensor, pad_multiple: int) -> torch.Tensor: """Fused version of pad_routing_map. Args: diff --git a/megatron/core/transformer/moe/moe_utils.py b/megatron/core/transformer/moe/moe_utils.py index dc857129834..17942fa5a3e 100644 --- a/megatron/core/transformer/moe/moe_utils.py +++ b/megatron/core/transformer/moe/moe_utils.py @@ -1,4 +1,4 @@ -# Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved. +# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. import math from typing import List, Optional, Union @@ -7,6 +7,7 @@ from megatron.core import parallel_state from megatron.core.process_groups_config import ProcessGroupCollection +from megatron.core.transformer.cuda_graphs import is_graph_capturing try: import transformer_engine as te # pylint: disable=unused-import @@ -905,12 +906,16 @@ class RandomSTE(torch.autograd.Function): """ generator = None + random_logits = None @staticmethod def forward(ctx, logits): """ Forward pass returns random logits with rank-specific seed. """ + if is_graph_capturing() and RandomSTE.random_logits is not None: + return RandomSTE.random_logits + if RandomSTE.generator is None: global_rank = torch.distributed.get_rank() base_seed = 42 @@ -918,8 +923,8 @@ def forward(ctx, logits): RandomSTE.generator = torch.Generator(device=logits.device) RandomSTE.generator.manual_seed(seed) - random_logits = logits.clone().normal_(generator=RandomSTE.generator) - return random_logits + RandomSTE.random_logits = logits.clone().normal_(generator=RandomSTE.generator) + return RandomSTE.random_logits @staticmethod def backward(ctx, grad_output): diff --git a/megatron/core/transformer/moe/router.py b/megatron/core/transformer/moe/router.py index 7fa4692ef2f..16fc9d9af8f 100644 --- a/megatron/core/transformer/moe/router.py +++ b/megatron/core/transformer/moe/router.py @@ -1,10 +1,11 @@ -# Copyright (c) 2023, NVIDIA CORPORATION. All rights reserved. +# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. from abc import ABC, abstractmethod from typing import Optional import torch +from megatron.core.jit import jit_fuser from megatron.core.tensor_parallel import reduce_from_tensor_model_parallel_region from megatron.core.transformer.module import MegatronModule from megatron.core.transformer.moe.moe_utils import ( @@ -468,6 +469,16 @@ def apply_input_jitter(self, input: torch.Tensor): else: return input + @jit_fuser + def _apply_expert_bias(self, routing_map: torch.Tensor): + """ + Update expert bias and tokens_per_expert + Prevent extra local tokens accumulation on evaluation or activation recomputation + """ + if self.enable_expert_bias and torch.is_grad_enabled(): + with torch.no_grad(): + self.local_tokens_per_expert += routing_map.sum(dim=0) + def routing(self, logits: torch.Tensor): """Top-k routing function @@ -526,11 +537,8 @@ def routing(self, logits: torch.Tensor): probs, scores_for_aux_loss, routing_map_for_aux_loss ) - # Update expert bias and tokens_per_expert - # Prevent extra local tokens accumulation on evaluation or activation recomputation - if self.enable_expert_bias and torch.is_grad_enabled(): - with torch.no_grad(): - self.local_tokens_per_expert += routing_map.sum(dim=0) + # Optionally apply expert bias + self._apply_expert_bias(routing_map) return probs, routing_map diff --git a/megatron/core/transformer/moe/token_dispatcher.py b/megatron/core/transformer/moe/token_dispatcher.py index 46f94ebe79a..bb034292715 100644 --- a/megatron/core/transformer/moe/token_dispatcher.py +++ b/megatron/core/transformer/moe/token_dispatcher.py @@ -1,4 +1,4 @@ -# Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved. +# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. import logging from abc import ABC, abstractmethod @@ -12,6 +12,7 @@ from megatron.core.fp8_utils import get_fp8_align_size from megatron.core.fusions.fused_indices_converter import fused_indices_to_multihot from megatron.core.fusions.fused_pad_routing_map import fused_pad_routing_map +from megatron.core.jit import jit_fuser from megatron.core.tensor_parallel import ( all_to_all, gather_from_sequence_parallel_region, @@ -1386,6 +1387,7 @@ def _initialize_metadata(self, routing_map: torch.Tensor, probs: torch.Tensor) - ).contiguous() return routing_map, probs + @jit_fuser def dispatch_preprocess( self, hidden_states: torch.Tensor, routing_map: torch.Tensor, probs: torch.Tensor ):