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
6 changes: 6 additions & 0 deletions src/megatron/bridge/_compat/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,6 @@
"""Backports of newer megatron-core APIs.

Used as fallback imports when the runtime's Megatron-LM predates symbols
that newer megatron-bridge versions import (e.g. the radixark/miles:dev image's
bundled Megatron-LM is missing megatron.core.models.mimo.config.role).
"""
167 changes: 167 additions & 0 deletions src/megatron/bridge/_compat/mimo_role.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,167 @@
# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved.

"""Data classes for MIMO rank role management in multi-module pipeline parallelism."""

import logging
from dataclasses import dataclass, field
from enum import Enum
from typing import Dict, List, Optional

import torch.distributed as dist

from megatron.core.hyper_comm_grid import HyperCommGrid

logger = logging.getLogger(__name__)

# Fixed key for the language module in module_to_grid_map and RankRole.
# MIMO always has exactly one language model, so this is not configurable.
MIMO_LANGUAGE_MODULE_KEY = "language"


class ModuleLayout(Enum):
"""Pipeline mode for MIMO multi-module parallelism.

Determines how modules are distributed across ranks and which
forward path is used.

COLOCATED: All modules share the same ranks. Covers both legacy
(no grid map, global parallel_state) and heterogeneous TP/DP
(grid map with overlapping ranks). Uses _forward_all_modules.

NON_COLOCATED: module_to_grid_map is set with non-overlapping rank
ranges. Each rank runs EITHER encoder(s) OR the language model.
Uses role-based dispatch with separate forward paths.
"""

COLOCATED = "colocated"
NON_COLOCATED = "non_colocated"


@dataclass
class ModuleStageInfo:
"""Information about a rank's stage position within a module's pipeline.

Args:
is_first_stage: True if this rank is the first PP stage for this module.
is_last_stage: True if this rank is the last PP stage for this module.
"""

is_first_stage: bool
is_last_stage: bool


@dataclass
class RankRole:
"""Describes what modules this rank participates in for multi-module PP.

This class captures the role of a specific rank in a multi-module pipeline
parallel setup, tracking which modules the rank participates in and their
stage positions. The language module is always identified by MIMO_LANGUAGE_MODULE_KEY.

Args:
modules: Dict mapping module names to their stage info for modules
this rank participates in.
mode: Pipeline mode determining forward path dispatch.
"""

modules: Dict[str, ModuleStageInfo] = field(default_factory=dict)
mode: ModuleLayout = ModuleLayout.COLOCATED

@classmethod
def build(
cls,
modality_module_names: List[str],
module_to_grid_map: Optional[Dict[str, 'HyperCommGrid']] = None,
) -> 'RankRole':
"""Build a RankRole, dispatching by whether grids share ranks.

No grid map or all grids span the same ranks → COLOCATED.
Grids differ → NON_COLOCATED with PP-stage info per module.
"""
if module_to_grid_map is None or cls._all_grids_colocated(module_to_grid_map):
return cls._colocated(modality_module_names)
return cls._from_grid_map(module_to_grid_map)

@staticmethod
def _all_grids_colocated(module_to_grid_map: Dict[str, 'HyperCommGrid']) -> bool:
grids = list(module_to_grid_map.values())
first = grids[0]
return all(g.rank_offset == first.rank_offset and g.size == first.size for g in grids[1:])

@classmethod
def _colocated(cls, modality_module_names: List[str]) -> 'RankRole':
"""Colocated layout: every module on every rank, PP=1."""
all_module_names = list(modality_module_names) + [MIMO_LANGUAGE_MODULE_KEY]
return cls(
modules={
name: ModuleStageInfo(is_first_stage=True, is_last_stage=True)
for name in all_module_names
},
mode=ModuleLayout.COLOCATED,
)

@classmethod
def _from_grid_map(cls, module_to_grid_map: Dict[str, HyperCommGrid]) -> 'RankRole':
"""Non-colocated role for this rank from a module-to-grid mapping.

Grid map keys are validated by ``MimoModelConfig.__post_init__``.

Raises:
RuntimeError: If current rank is not in any module grid.
"""
current_rank = dist.get_rank()
modules = {}

for module_name, grid in module_to_grid_map.items():
if not (grid.rank_offset <= current_rank < grid.rank_offset + grid.size):
continue

if "pp" not in grid.dim_names:
modules[module_name] = ModuleStageInfo(is_first_stage=True, is_last_stage=True)
continue

pp_group = grid.get_pg("pp")
pp_rank = pp_group.rank()
pp_size = pp_group.size()
is_first = pp_rank == 0
is_last = pp_rank == pp_size - 1
logger.info(
f"[RankRole._from_grid_map] Rank {current_rank}: module={module_name}, "
f"pp_rank={pp_rank}/{pp_size}, is_first_stage={is_first}, is_last_stage={is_last}"
)
modules[module_name] = ModuleStageInfo(is_first_stage=is_first, is_last_stage=is_last)

if not modules:
raise RuntimeError(
f"Rank {current_rank} is not in any module grid. "
f"Check module_to_grid_map configuration."
)

return cls(modules=modules, mode=ModuleLayout.NON_COLOCATED)

@property
def has_modality_modules(self) -> bool:
"""Return True if this rank participates in any modality (non-language) module."""
return any(name != MIMO_LANGUAGE_MODULE_KEY for name in self.modules)

@property
def has_language_module(self) -> bool:
"""Return True if this rank participates in the language module."""
return MIMO_LANGUAGE_MODULE_KEY in self.modules

@property
def modality_module_names(self) -> List[str]:
"""Return names of modality modules (non-language) this rank participates in."""
return [name for name in self.modules if name != MIMO_LANGUAGE_MODULE_KEY]

def is_first_stage(self, module_name: str) -> bool:
"""Check if this rank is the first stage for a given module."""
if module_name not in self.modules:
return False
return self.modules[module_name].is_first_stage

def is_last_stage(self, module_name: str) -> bool:
"""Check if this rank is the last stage for a given module."""
if module_name not in self.modules:
return False
return self.modules[module_name].is_last_stage
7 changes: 6 additions & 1 deletion src/megatron/bridge/data/megatron_mimo/dp_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,12 @@

import torch
import torch.distributed as dist
from megatron.core.models.mimo.config.role import MIMO_LANGUAGE_MODULE_KEY
try:
from megatron.core.models.mimo.config.role import MIMO_LANGUAGE_MODULE_KEY
except ImportError:
# Backport for older Megatron-LM (e.g. radixark/miles:dev) that lacks
# megatron.core.models.mimo.config.role. Vendored copy lives in bridge.
from megatron.bridge._compat.mimo_role import MIMO_LANGUAGE_MODULE_KEY


if TYPE_CHECKING:
Expand Down
12 changes: 11 additions & 1 deletion src/megatron/bridge/models/mamba/mamba_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,17 @@
from megatron.core.pipeline_parallel.utils import is_pp_first_stage, is_pp_last_stage
from megatron.core.post_training.modelopt.mamba.model_specs import get_mamba_stack_modelopt_spec
from megatron.core.process_groups_config import ProcessGroupCollection
from megatron.core.ssm.mamba_hybrid_layer_allocation import Symbols, parse_hybrid_pattern
from megatron.core.ssm.mamba_hybrid_layer_allocation import Symbols

try:
from megatron.core.ssm.mamba_hybrid_layer_allocation import parse_hybrid_pattern
except ImportError:
# radixark/miles `miles-main` Megatron-LM fork lacks `parse_hybrid_pattern`
# (upstream NVIDIA added it in MTP commit 300d1b655, which miles-main
# branched off before). It's only used inside method paths Mamba/hybrid
# models take; non-Mamba consumers (gpt-oss, Qwen, Kimi) never call it.
# Degrade to None so the module loads everywhere.
parse_hybrid_pattern = None
from megatron.core.transformer import ModuleSpec
from megatron.core.transformer.enums import AttnBackend

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,12 @@
from dataclasses import dataclass, field
from typing import Optional

from megatron.core.models.mimo.config.role import MIMO_LANGUAGE_MODULE_KEY
try:
from megatron.core.models.mimo.config.role import MIMO_LANGUAGE_MODULE_KEY
except ImportError:
# Backport for older Megatron-LM (e.g. radixark/miles:dev) that lacks
# megatron.core.models.mimo.config.role. Vendored copy lives in bridge.
from megatron.bridge._compat.mimo_role import MIMO_LANGUAGE_MODULE_KEY


@dataclass
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,12 @@

from typing import TYPE_CHECKING, Dict, Optional

from megatron.core.models.mimo.config.role import MIMO_LANGUAGE_MODULE_KEY
try:
from megatron.core.models.mimo.config.role import MIMO_LANGUAGE_MODULE_KEY
except ImportError:
# Backport for older Megatron-LM (e.g. radixark/miles:dev) that lacks
# megatron.core.models.mimo.config.role. Vendored copy lives in bridge.
from megatron.bridge._compat.mimo_role import MIMO_LANGUAGE_MODULE_KEY


if TYPE_CHECKING:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,12 @@
from megatron.core.distributed import DistributedDataParallelConfig
from megatron.core.models.mimo import MimoModel
from megatron.core.models.mimo.config.base_configs import MimoModelConfig
from megatron.core.models.mimo.config.role import MIMO_LANGUAGE_MODULE_KEY
try:
from megatron.core.models.mimo.config.role import MIMO_LANGUAGE_MODULE_KEY
except ImportError:
# Backport for older Megatron-LM (e.g. radixark/miles:dev) that lacks
# megatron.core.models.mimo.config.role. Vendored copy lives in bridge.
from megatron.bridge._compat.mimo_role import MIMO_LANGUAGE_MODULE_KEY
from megatron.core.process_groups_config import ProcessGroupCollection
from megatron.core.transformer.module import MegatronModule
from megatron.core.transformer.spec_utils import ModuleSpec
Expand Down
71 changes: 62 additions & 9 deletions src/megatron/bridge/training/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,15 +34,68 @@
from megatron.core.transformer.module import MegatronModule
from megatron.core.transformer.transformer_config import MLATransformerConfig as MCoreMLATransformerConfig
from megatron.core.transformer.transformer_config import TransformerConfig as MCoreTransformerConfig
from megatron.training.config import CheckpointConfig as MTrainCheckpointConfig
from megatron.training.config import DistributedInitConfig as MTrainDistributedInitConfig
from megatron.training.config import LoggerConfig as MTrainLoggerConfig
from megatron.training.config import ProfilingConfig as MTrainProfilingConfig
from megatron.training.config import RerunStateMachineConfig as MTrainRerunStateMachineConfig
from megatron.training.config import RNGConfig, ValidationConfig
from megatron.training.config import SchedulerConfig as MTrainSchedulerConfig
from megatron.training.config import StragglerDetectionConfig as MTrainStragglerDetectionConfig
from megatron.training.config import TrainingConfig as MTrainTrainingConfig
try:
# Newer Megatron-LM (≥ commit 8b00c3ce, 2026-03-30): config dataclasses are
# consolidated into the megatron.training.config package.
from megatron.training.config import CheckpointConfig as MTrainCheckpointConfig
from megatron.training.config import DistributedInitConfig as MTrainDistributedInitConfig
from megatron.training.config import LoggerConfig as MTrainLoggerConfig
from megatron.training.config import ProfilingConfig as MTrainProfilingConfig
from megatron.training.config import RerunStateMachineConfig as MTrainRerunStateMachineConfig
from megatron.training.config import RNGConfig, ValidationConfig
from megatron.training.config import SchedulerConfig as MTrainSchedulerConfig
from megatron.training.config import StragglerDetectionConfig as MTrainStragglerDetectionConfig
from megatron.training.config import TrainingConfig as MTrainTrainingConfig
except ImportError:
# Older Megatron-LM (e.g. radixark/miles:dev image): dataclasses are scattered
# across individual modules. DistributedInitConfig was added in commit 8b00c3ce
# and has no counterpart, so we backport an inline stub mirroring the upstream
# definition (kept in sync with newer Megatron-LM).
from megatron.training.training_config import (
CheckpointConfig as MTrainCheckpointConfig,
LoggerConfig as MTrainLoggerConfig,
SchedulerConfig as MTrainSchedulerConfig,
TrainingConfig as MTrainTrainingConfig,
ValidationConfig,
)
from megatron.training.common_config import (
ProfilingConfig as MTrainProfilingConfig,
RNGConfig,
)
from megatron.training.resilience_config import (
RerunStateMachineConfig as MTrainRerunStateMachineConfig,
StragglerDetectionConfig as MTrainStragglerDetectionConfig,
)

@dataclass(kw_only=True)
class MTrainDistributedInitConfig:
"""Backport of upstream megatron.training.config.DistributedInitConfig
for older Megatron-LM that doesn't ship a megatron.training.config package."""

distributed_backend: Literal["nccl", "gloo"] = "nccl"
distributed_timeout_minutes: int = 10
align_grad_reduce: bool = True
local_rank: int = field(default_factory=lambda: int(os.getenv("LOCAL_RANK", "0")))
lazy_mpu_init: bool = False
use_megatron_fsdp: bool = False
use_torch_fsdp2: bool = False
nccl_communicator_config_path: Optional[str] = None
use_tp_pp_dp_mapping: bool = False
use_gloo_process_groups: bool = field(
default=True,
metadata={"argparse_meta": {"arg_names": ["--disable-gloo-process-groups"]}},
)
use_sharp: bool = False
sharp_enabled_group: Optional[Literal["dp", "dp_replica"]] = None
high_priority_stream_groups: Optional[list[str]] = field(default_factory=list)
distributed_timeout_seconds_after_init: Optional[int] = None
flight_recorder_dump_path: Optional[str] = None
flight_recorder_trace_buffer_size: int = 2000
flight_recorder_dump_on_timeout: bool = True
flight_recorder_include_stack_trace: bool = False
flight_recorder_include_only_active: bool = True
flight_recorder_extra_dump_on_exec: bool = True
disable_jit_fuser: bool = False

from megatron.bridge.data.datasets.packed_sequence import PackedSequenceSpecs
from megatron.bridge.models import GPTModelProvider, T5ModelProvider
Expand Down
7 changes: 6 additions & 1 deletion src/megatron/bridge/training/megatron_mimo_parallel_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,12 @@
import torch.distributed as dist
from megatron.core.distributed.finalize_model_grads import finalize_model_grads as _finalize_model_grads
from megatron.core.models.mimo import MimoModel
from megatron.core.models.mimo.config.role import MIMO_LANGUAGE_MODULE_KEY
try:
from megatron.core.models.mimo.config.role import MIMO_LANGUAGE_MODULE_KEY
except ImportError:
# Backport for older Megatron-LM (e.g. radixark/miles:dev) that lacks
# megatron.core.models.mimo.config.role. Vendored copy lives in bridge.
from megatron.bridge._compat.mimo_role import MIMO_LANGUAGE_MODULE_KEY

from megatron.bridge.models.megatron_mimo.megatron_mimo_provider import MegatronMIMOInfra

Expand Down
7 changes: 6 additions & 1 deletion src/megatron/bridge/training/megatron_mimo_step.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,12 @@

import torch
from megatron.core.models.mimo import MimoModel
from megatron.core.models.mimo.config.role import MIMO_LANGUAGE_MODULE_KEY
try:
from megatron.core.models.mimo.config.role import MIMO_LANGUAGE_MODULE_KEY
except ImportError:
# Backport for older Megatron-LM (e.g. radixark/miles:dev) that lacks
# megatron.core.models.mimo.config.role. Vendored copy lives in bridge.
from megatron.bridge._compat.mimo_role import MIMO_LANGUAGE_MODULE_KEY

from megatron.bridge.data.megatron_mimo.dp_utils import slice_batch_for_megatron_mimo
from megatron.bridge.training.megatron_mimo_parallel_utils import unwrap_megatron_mimo_model
Expand Down
7 changes: 6 additions & 1 deletion src/megatron/bridge/training/train_megatron_mimo.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,12 @@

import torch
import torch.distributed as dist
from megatron.core.models.mimo.config.role import MIMO_LANGUAGE_MODULE_KEY
try:
from megatron.core.models.mimo.config.role import MIMO_LANGUAGE_MODULE_KEY
except ImportError:
# Backport for older Megatron-LM (e.g. radixark/miles:dev) that lacks
# megatron.core.models.mimo.config.role. Vendored copy lives in bridge.
from megatron.bridge._compat.mimo_role import MIMO_LANGUAGE_MODULE_KEY
from megatron.core.num_microbatches_calculator import get_num_microbatches
from megatron.core.pipeline_parallel.schedules import forward_backward_pipelining_without_interleaving

Expand Down
Loading