Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 4 additions & 1 deletion megatron/core/fusions/fused_pad_routing_map.py
Original file line number Diff line number Diff line change
@@ -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:
Expand Down Expand Up @@ -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:
Expand Down
11 changes: 8 additions & 3 deletions megatron/core/transformer/moe/moe_utils.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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
Expand Down Expand Up @@ -905,21 +906,25 @@ 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
seed = base_seed + global_rank
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):
Expand Down
20 changes: 14 additions & 6 deletions megatron/core/transformer/moe/router.py
Original file line number Diff line number Diff line change
@@ -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 (
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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

Expand Down
4 changes: 3 additions & 1 deletion megatron/core/transformer/moe/token_dispatcher.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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,
Expand Down Expand Up @@ -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
):
Expand Down
Loading