From 48edcb2878536c027d482bf12d9696bc3ffa4e18 Mon Sep 17 00:00:00 2001 From: Xin Yao Date: Tue, 15 Jul 2025 22:28:54 -0700 Subject: [PATCH 1/7] use torch.compile Signed-off-by: Xin Yao --- megatron/core/transformer/moe/router.py | 18 +++++++++++++----- .../core/transformer/moe/token_dispatcher.py | 3 +++ 2 files changed, 16 insertions(+), 5 deletions(-) diff --git a/megatron/core/transformer/moe/router.py b/megatron/core/transformer/moe/router.py index 068d680c798..5d1f906969d 100644 --- a/megatron/core/transformer/moe/router.py +++ b/megatron/core/transformer/moe/router.py @@ -5,6 +5,7 @@ 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 ( @@ -466,6 +467,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 @@ -524,11 +535,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 82fb7b00583..fbadd3a7632 100644 --- a/megatron/core/transformer/moe/token_dispatcher.py +++ b/megatron/core/transformer/moe/token_dispatcher.py @@ -11,6 +11,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, @@ -1065,6 +1066,7 @@ def combine( self.handle = None return hidden_states + @jit_fuser def _pad_routing_map( self, routing_map: torch.Tensor, tokens_per_expert: torch.Tensor ) -> Tuple[torch.Tensor, torch.Tensor]: @@ -1208,6 +1210,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 ): From d94eb19bb1e3bf082233c6830e062d1f06649aca Mon Sep 17 00:00:00 2001 From: Xin Yao Date: Mon, 28 Jul 2025 02:02:53 +0000 Subject: [PATCH 2/7] patch RandomSTE Signed-off-by: Xin Yao --- megatron/core/transformer/moe/moe_utils.py | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/megatron/core/transformer/moe/moe_utils.py b/megatron/core/transformer/moe/moe_utils.py index dc857129834..c20881398f2 100644 --- a/megatron/core/transformer/moe/moe_utils.py +++ b/megatron/core/transformer/moe/moe_utils.py @@ -905,12 +905,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 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 +922,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): From 08dae185c8542364bcf0e86052848ef7738bd079 Mon Sep 17 00:00:00 2001 From: Xin Yao Date: Mon, 20 Oct 2025 08:15:35 +0000 Subject: [PATCH 3/7] fix Signed-off-by: Xin Yao --- megatron/core/fusions/fused_pad_routing_map.py | 2 ++ megatron/core/transformer/moe/token_dispatcher.py | 1 - 2 files changed, 2 insertions(+), 1 deletion(-) diff --git a/megatron/core/fusions/fused_pad_routing_map.py b/megatron/core/fusions/fused_pad_routing_map.py index e7c3a7e48c9..eb41e5b657f 100644 --- a/megatron/core/fusions/fused_pad_routing_map.py +++ b/megatron/core/fusions/fused_pad_routing_map.py @@ -4,6 +4,7 @@ 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 +70,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/token_dispatcher.py b/megatron/core/transformer/moe/token_dispatcher.py index fbadd3a7632..f5c7fa825b4 100644 --- a/megatron/core/transformer/moe/token_dispatcher.py +++ b/megatron/core/transformer/moe/token_dispatcher.py @@ -1066,7 +1066,6 @@ def combine( self.handle = None return hidden_states - @jit_fuser def _pad_routing_map( self, routing_map: torch.Tensor, tokens_per_expert: torch.Tensor ) -> Tuple[torch.Tensor, torch.Tensor]: From bb20758d0738847aaab9f008b238156dcff962f3 Mon Sep 17 00:00:00 2001 From: Xin Yao Date: Fri, 24 Oct 2025 09:38:21 +0000 Subject: [PATCH 4/7] add is_graph_capturing Signed-off-by: Xin Yao --- megatron/core/transformer/moe/moe_utils.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/megatron/core/transformer/moe/moe_utils.py b/megatron/core/transformer/moe/moe_utils.py index c20881398f2..ed482705ae9 100644 --- a/megatron/core/transformer/moe/moe_utils.py +++ b/megatron/core/transformer/moe/moe_utils.py @@ -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 @@ -912,7 +913,7 @@ def forward(ctx, logits): """ Forward pass returns random logits with rank-specific seed. """ - if RandomSTE.random_logits is not None: + if is_graph_capturing() and RandomSTE.random_logits is not None: return RandomSTE.random_logits if RandomSTE.generator is None: From b802a613142634d953073c4beb02a25cf571c18a Mon Sep 17 00:00:00 2001 From: Xin Yao Date: Mon, 27 Oct 2025 03:39:28 +0000 Subject: [PATCH 5/7] update copyright Signed-off-by: Xin Yao --- megatron/core/fusions/fused_pad_routing_map.py | 2 +- megatron/core/transformer/moe/moe_utils.py | 2 +- megatron/core/transformer/moe/router.py | 2 +- megatron/core/transformer/moe/token_dispatcher.py | 2 +- 4 files changed, 4 insertions(+), 4 deletions(-) diff --git a/megatron/core/fusions/fused_pad_routing_map.py b/megatron/core/fusions/fused_pad_routing_map.py index eb41e5b657f..82856e1b9b8 100644 --- a/megatron/core/fusions/fused_pad_routing_map.py +++ b/megatron/core/fusions/fused_pad_routing_map.py @@ -1,4 +1,4 @@ -# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved. +# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. from unittest.mock import MagicMock import torch diff --git a/megatron/core/transformer/moe/moe_utils.py b/megatron/core/transformer/moe/moe_utils.py index ed482705ae9..4eb818ea5bc 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) 2024-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. import math from typing import List, Optional, Union diff --git a/megatron/core/transformer/moe/router.py b/megatron/core/transformer/moe/router.py index 5d1f906969d..99e79f998c9 100644 --- a/megatron/core/transformer/moe/router.py +++ b/megatron/core/transformer/moe/router.py @@ -1,4 +1,4 @@ -# Copyright (c) 2023, NVIDIA CORPORATION. All rights reserved. +# Copyright (c) 2023-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. from abc import ABC, abstractmethod from typing import Optional diff --git a/megatron/core/transformer/moe/token_dispatcher.py b/megatron/core/transformer/moe/token_dispatcher.py index f5c7fa825b4..ea8a13e3788 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) 2024-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. import logging from abc import ABC, abstractmethod From 04b3623b4e6dac540c7bff34489023c910526ae6 Mon Sep 17 00:00:00 2001 From: Xin Yao Date: Mon, 27 Oct 2025 03:45:16 +0000 Subject: [PATCH 6/7] update copyright Signed-off-by: Xin Yao --- megatron/core/fusions/fused_pad_routing_map.py | 1 + megatron/core/transformer/moe/moe_utils.py | 2 +- megatron/core/transformer/moe/router.py | 2 +- megatron/core/transformer/moe/token_dispatcher.py | 2 +- 4 files changed, 4 insertions(+), 3 deletions(-) diff --git a/megatron/core/fusions/fused_pad_routing_map.py b/megatron/core/fusions/fused_pad_routing_map.py index 82856e1b9b8..8e4d1763270 100644 --- a/megatron/core/fusions/fused_pad_routing_map.py +++ b/megatron/core/fusions/fused_pad_routing_map.py @@ -1,4 +1,5 @@ # Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + from unittest.mock import MagicMock import torch diff --git a/megatron/core/transformer/moe/moe_utils.py b/megatron/core/transformer/moe/moe_utils.py index 4eb818ea5bc..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-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. import math from typing import List, Optional, Union diff --git a/megatron/core/transformer/moe/router.py b/megatron/core/transformer/moe/router.py index 99e79f998c9..65533b98d60 100644 --- a/megatron/core/transformer/moe/router.py +++ b/megatron/core/transformer/moe/router.py @@ -1,4 +1,4 @@ -# Copyright (c) 2023-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. from abc import ABC, abstractmethod from typing import Optional diff --git a/megatron/core/transformer/moe/token_dispatcher.py b/megatron/core/transformer/moe/token_dispatcher.py index ea8a13e3788..b80a3d6f44a 100644 --- a/megatron/core/transformer/moe/token_dispatcher.py +++ b/megatron/core/transformer/moe/token_dispatcher.py @@ -1,4 +1,4 @@ -# Copyright (c) 2024-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. import logging from abc import ABC, abstractmethod From 71dcec0ea2461e921a2384078d0dbe65d100387e Mon Sep 17 00:00:00 2001 From: Xin Yao Date: Tue, 18 Nov 2025 03:25:42 +0000 Subject: [PATCH 7/7] replace print with log Signed-off-by: Xin Yao --- megatron/core/transformer/moe/token_dispatcher.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/megatron/core/transformer/moe/token_dispatcher.py b/megatron/core/transformer/moe/token_dispatcher.py index b80a3d6f44a..7b95393a2c4 100644 --- a/megatron/core/transformer/moe/token_dispatcher.py +++ b/megatron/core/transformer/moe/token_dispatcher.py @@ -44,6 +44,8 @@ num_global_tokens: num_local_tokens*TP*EP """ +logger = logging.getLogger(__name__) + class MoETokenDispatcher: """ @@ -990,7 +992,9 @@ def dispatch( # DeepEP only supports float32 probs if self.token_probs.dtype != torch.float32: if self.token_probs.dtype in [torch.bfloat16, torch.float16]: - print("DeepEP only supports float32 probs, please set --moe-router-dtype=fp32") + logger.warning( + "DeepEP only supports float32 probs, please set --moe-router-dtype=fp32" + ) self.token_probs = self.token_probs.float() # downcast or upcast hidden_states, dispatched_indices, dispatched_probs, num_tokens_per_expert, handle = ( fused_dispatch( @@ -1082,7 +1086,6 @@ def _pad_routing_map( # Check if there are enough tokens to pad enough_tokens_to_pad = torch.all(target_tokens_per_expert <= num_input_tokens) if not enough_tokens_to_pad: - logger = logging.getLogger(__name__) logger.warning( "Not enough tokens to pad. The total number of tokens received in this rank " "is smaller than the target number of tokens for each expert. "