From e5a5b6a5de1ab9999cd5b2e370584c4d77be9638 Mon Sep 17 00:00:00 2001 From: Xin Yao Date: Tue, 15 Jul 2025 22:28:54 -0700 Subject: [PATCH 1/6] 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 7fa4692ef2f..5bf8191ce2d 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 ( @@ -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..a07d8623a2a 100644 --- a/megatron/core/transformer/moe/token_dispatcher.py +++ b/megatron/core/transformer/moe/token_dispatcher.py @@ -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, @@ -1227,6 +1228,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]: @@ -1386,6 +1388,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 2a10c063ce57be1fb4f7f7082c3dd95a490d456b Mon Sep 17 00:00:00 2001 From: Xin Yao Date: Mon, 28 Jul 2025 02:02:53 +0000 Subject: [PATCH 2/6] 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 cd17d4e5b8e9eb268d867a129d9af2e28e96b8a8 Mon Sep 17 00:00:00 2001 From: Xin Yao Date: Mon, 20 Oct 2025 08:15:35 +0000 Subject: [PATCH 3/6] 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 a07d8623a2a..9efa3c19c74 100644 --- a/megatron/core/transformer/moe/token_dispatcher.py +++ b/megatron/core/transformer/moe/token_dispatcher.py @@ -1228,7 +1228,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 41d42263173c6902b28895e45746d3073e02007a Mon Sep 17 00:00:00 2001 From: Xin Yao Date: Fri, 24 Oct 2025 09:38:21 +0000 Subject: [PATCH 4/6] 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 a974d066bd76b0cf885dce9ec7eee5c211c894d6 Mon Sep 17 00:00:00 2001 From: Xin Yao Date: Mon, 27 Oct 2025 03:39:28 +0000 Subject: [PATCH 5/6] 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 5bf8191ce2d..3283f632e1e 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 9efa3c19c74..161c7216a9e 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 12ff69bd7b0482e3b0b0a74578b1367e54a30868 Mon Sep 17 00:00:00 2001 From: Xin Yao Date: Mon, 27 Oct 2025 03:45:16 +0000 Subject: [PATCH 6/6] 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 3283f632e1e..16fc9d9af8f 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 161c7216a9e..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-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. import logging from abc import ABC, abstractmethod