diff --git a/src/megatron/bridge/training/config.py b/src/megatron/bridge/training/config.py index f4e31a27d2..25091106bd 100644 --- a/src/megatron/bridge/training/config.py +++ b/src/megatron/bridge/training/config.py @@ -1400,6 +1400,17 @@ def validate(self) -> None: if hasattr(self.model, "finalize"): self.model.finalize() + from megatron.bridge.training.gtp import is_gtp_remat_active + + if is_gtp_remat_active(self.model): + if self.dist.use_decentralized_pg: + raise ValueError( + "GTP is not supported with dist.use_decentralized_pg=True. " + "Set dist.use_decentralized_pg=False to use the standard MCore process-group runtime." + ) + if self.ddp.average_in_collective: + raise ValueError("GTP requires ddp.average_in_collective=False.") + self.logger.finalize() self.train.finalize() self.scheduler.finalize() diff --git a/src/megatron/bridge/training/eval.py b/src/megatron/bridge/training/eval.py index a128c382e2..468ebfb254 100644 --- a/src/megatron/bridge/training/eval.py +++ b/src/megatron/bridge/training/eval.py @@ -34,6 +34,7 @@ from megatron.bridge.training.callbacks import CallbackContext, CallbackManager, should_fire from megatron.bridge.training.config import ConfigContainer from megatron.bridge.training.forward_step_func_types import ForwardStepCallable +from megatron.bridge.training.gtp import get_data_distribution_group from megatron.bridge.training.state import GlobalState from megatron.bridge.training.utils.mlflow_utils import _sanitize_mlflow_metrics from megatron.bridge.training.utils.pg_utils import get_pg_collection @@ -141,7 +142,11 @@ def evaluate( eval_micro_batch_size = state.cfg.validation.eval_micro_batch_size # MegatronMIMO has heterogeneous per-module DP groups and intentionally owns # global-batch accounting through the container-level DP size. - eval_data_parallel_size = state.cfg.data_parallel_size if is_multimodule else pg_collection.dp.size() + eval_data_parallel_size = ( + state.cfg.data_parallel_size + if is_multimodule + else get_data_distribution_group(pg_collection, state.cfg.model).size() + ) eval_num_microbatches = eval_batch_size // (eval_micro_batch_size * eval_data_parallel_size) if is_multimodule and not isinstance(p2p_communicator, MultiModulePipelineCommunicator): @@ -290,7 +295,9 @@ def evaluate( if is_multimodule: dp_cp_group = pg_collection.get_language_model_collection().dp_cp else: - dp_cp_group = pg_collection.dp_cp + dp_cp_group = get_data_distribution_group( + pg_collection, state.cfg.model, with_context_parallel=True + ) for key in loss_dicts[0].keys(): if key not in total_loss_dict: diff --git a/src/megatron/bridge/training/gtp.py b/src/megatron/bridge/training/gtp.py new file mode 100644 index 0000000000..85e31471e3 --- /dev/null +++ b/src/megatron/bridge/training/gtp.py @@ -0,0 +1,86 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Generalized Tensor Parallelism helpers for the standard Bridge runtime.""" + +from typing import Any + +import torch +from megatron.core import parallel_state +from megatron.core.process_groups_config import ProcessGroupCollection + + +def get_transformer_config(model_config: Any) -> Any: + """Return the MCore transformer config nested in a Bridge model config.""" + model_fields = getattr(type(model_config), "__dataclass_fields__", {}) + if "transformer" in model_fields: + return model_config.transformer + return model_config + + +def is_gtp_remat_active(model_config: Any) -> bool: + """Return whether dense or expert GTP weight rematerialization is enabled.""" + transformer_config = get_transformer_config(model_config) + dense_size = getattr(transformer_config, "gtp_weight_remat_size", 1) + expert_size = getattr(transformer_config, "expert_gtp_weight_remat_size", 1) + return any(isinstance(size, int) and size > 1 for size in (dense_size, expert_size)) + + +def configure_gtp_remat(model_config: Any) -> None: + """Configure process-global GTP state before constructing model modules.""" + if not is_gtp_remat_active(model_config): + return + + transformer_config = get_transformer_config(model_config) + from megatron.core.tensor_parallel import gtp_api + + if not gtp_api.HAVE_GTP: + raise RuntimeError("GTP requires TransformerEngine >= 2.19.") + + gtp_api.configure_gtp_remat_from_recipe( + fp4=transformer_config.fp4 is not None, + fp8_recipe=transformer_config.fp8_recipe, + fp8=transformer_config.fp8 is not None, + calculate_per_token_loss=transformer_config.calculate_per_token_loss, + ) + + +def classify_gtp_remat_chains(model: list[torch.nn.Module], model_config: Any) -> None: + """Classify all model chunks after distributed wrapping and before first forward.""" + if not is_gtp_remat_active(model_config): + return + + transformer_config = get_transformer_config(model_config) + from megatron.core.tensor_parallel import gtp_api + + gtp_api.classify_gtp_remat_chains( + model, + cuda_graph_modules=transformer_config.cuda_graph_modules, + moe_shared_expert_overlap=transformer_config.moe_shared_expert_overlap, + cuda_graph_impl=transformer_config.cuda_graph_impl, + ) + + +def get_data_distribution_group( + pg_collection: ProcessGroupCollection, + model_config: Any, + *, + with_context_parallel: bool = False, +) -> torch.distributed.ProcessGroup: + """Return the group spanning every rank that consumes distinct input data.""" + if not is_gtp_remat_active(model_config): + return pg_collection.dp_cp if with_context_parallel else pg_collection.dp + if with_context_parallel: + return pg_collection.dp_cp_gtp_remat + return parallel_state.get_data_parallel_group(with_gtp_remat=True) diff --git a/src/megatron/bridge/training/initialize.py b/src/megatron/bridge/training/initialize.py index da842f1fbb..f4f7dfe993 100644 --- a/src/megatron/bridge/training/initialize.py +++ b/src/megatron/bridge/training/initialize.py @@ -50,6 +50,7 @@ from megatron.bridge.models.hybrid.hybrid_builder import HybridModelConfig from megatron.bridge.models.transformer_config import TransformerConfig, _set_moe_expert_tensor_parallel_default from megatron.bridge.training.config import ConfigContainer, DistributedInitConfig, RerunStateMachineConfig, RNGConfig +from megatron.bridge.training.gtp import is_gtp_remat_active from megatron.bridge.training.utils.pg_utils import DistTrainProcessGroupCollection from megatron.bridge.utils.common_utils import ( get_local_rank_preinit, @@ -759,6 +760,18 @@ def _initialize_distributed( if dist_config.use_decentralized_pg or dist_config.distributed_backend == "nccl": raise RuntimeError("Cannot initialize parallel groups with no CUDA devices available (device_count=0)") + if dist_config.use_decentralized_pg and is_gtp_remat_active(model_config): + raise NotImplementedError( + "GTP is not supported with dist.use_decentralized_pg=True. " + "Use the standard MCore process-group runtime by setting dist.use_decentralized_pg=False." + ) + + if is_gtp_remat_active(model_config): + from megatron.core.tensor_parallel.gtp_api import HAVE_GTP + + if not HAVE_GTP: + raise RuntimeError("GTP requires TransformerEngine >= 2.19.") + if dist_config.use_decentralized_pg: # Use HyperCommGrid to create local parallel groups passed through functions # instead of relying on mcore's global parallel state (mpu) variables. @@ -812,6 +825,8 @@ def _initialize_distributed( expert_model_parallel_size=model_config.expert_model_parallel_size, num_distributed_optimizer_instances=num_distributed_optimizer_instances, expert_tensor_parallel_size=model_config.expert_tensor_parallel_size, + gtp_remat_size=model_config.gtp_weight_remat_size, + expert_gtp_remat_size=model_config.expert_gtp_weight_remat_size, distributed_timeout_minutes=dist_config.distributed_timeout_minutes, nccl_communicator_config_path=dist_config.nccl_communicator_config_path, order="tp-cp-ep-dp-pp" if not dist_config.use_tp_pp_dp_mapping else "tp-cp-ep-pp-dp", diff --git a/src/megatron/bridge/training/setup.py b/src/megatron/bridge/training/setup.py index 59c458c094..bdc01f218e 100644 --- a/src/megatron/bridge/training/setup.py +++ b/src/megatron/bridge/training/setup.py @@ -55,6 +55,11 @@ ) from megatron.bridge.training.config import ConfigContainer from megatron.bridge.training.fsdp_compat import MEGATRON_FSDP_TYPES +from megatron.bridge.training.gtp import ( + classify_gtp_remat_chains, + configure_gtp_remat, + get_data_distribution_group, +) from megatron.bridge.training.initialize import initialize_megatron, set_jit_fusion_options from megatron.bridge.training.optim import ( memory_efficient_fp32_optimizer_state_loading, @@ -500,7 +505,7 @@ def modelopt_pre_wrap_hook(model): train_state=state.train_state, model_length=len(model), train_valid_test_datasets_provider=train_valid_test_datasets_provider, - dp_group=pg_collection.dp, + dp_group=get_data_distribution_group(pg_collection, cfg.model), eval_dp_group=state._eval_pgs.dp if state._eval_pgs is not None else None, ) timers("train/valid/test-data-iterators-setup").stop() @@ -584,10 +589,13 @@ def _register_setup_pre_wrap_hook( def _build_distributed_model(cfg: ConfigContainer, pg_collection: ProcessGroupCollection) -> list[MegatronModule]: """Build distributed model from either ModelConfig or ModelProviderMixin.""" model_config = cfg.model + if not isinstance(model_config, ModelConfig): + model_config.finalize() + configure_gtp_remat(model_config) if isinstance(model_config, ModelConfig): builder_cls = model_config.get_builder_cls() builder = builder_cls(model_config) - return builder.build_distributed_models( + model = builder.build_distributed_models( pg_collection=pg_collection, ddp_config=cfg.ddp, overlap_param_gather_with_optimizer_step=cfg.optimizer.overlap_param_gather_with_optimizer_step, @@ -596,8 +604,7 @@ def _build_distributed_model(cfg: ConfigContainer, pg_collection: ProcessGroupCo data_parallel_random_init=cfg.rng.data_parallel_random_init, ) else: - model_config.finalize() - return model_config.provide_distributed_model( + model = model_config.provide_distributed_model( ddp_config=cfg.ddp, use_megatron_fsdp=cfg.dist.use_megatron_fsdp, use_torch_fsdp2=cfg.dist.use_torch_fsdp2, @@ -605,6 +612,8 @@ def _build_distributed_model(cfg: ConfigContainer, pg_collection: ProcessGroupCo data_parallel_random_init=cfg.rng.data_parallel_random_init, pg_collection=pg_collection, ) + classify_gtp_remat_chains(model, model_config) + return model def _update_model_config_funcs( diff --git a/src/megatron/bridge/training/train.py b/src/megatron/bridge/training/train.py index 973b7fa8c2..a1e5e00d23 100644 --- a/src/megatron/bridge/training/train.py +++ b/src/megatron/bridge/training/train.py @@ -74,6 +74,7 @@ from megatron.bridge.training.eval import evaluate_and_print_results from megatron.bridge.training.forward_step_func_types import ForwardStepCallable from megatron.bridge.training.fsdp_compat import MEGATRON_FSDP_TYPES +from megatron.bridge.training.gtp import get_data_distribution_group from megatron.bridge.training.initialize import destroy_global_state from megatron.bridge.training.nvrx_straggler import ( check_nvrx_straggler_detection, @@ -336,7 +337,8 @@ def train( start_iteration = global_state.train_state.step print_rank_0(f"Starting training loop at iteration {start_iteration}") p2p_communicator = P2PCommunicator(pp_group=pg_collection.pp, config=model_config) - dp_size = pg_collection.dp.size() + data_distribution_group = get_data_distribution_group(pg_collection, config.model) + dp_size = data_distribution_group.size() # Anchor for interval-average throughput logging: training_log reports the FLOPS # performed over each logging interval as the delta of # floating_point_operations_so_far. Seed it with the current cumulative (0 fresh, @@ -588,7 +590,7 @@ def train( global_state, data_parallel_size=dp_size, vp_size=config.model.virtual_pipeline_model_parallel_size, - dp_group=pg_collection.dp, + dp_group=data_distribution_group, include_vision_patch_stats=True, include_cross_attention_stats=hasattr( config.model, "_get_num_floating_point_operations_with_runtime_stats" @@ -993,7 +995,7 @@ def train_step( # there is one dict per microbatch. in new reporting, we average # over the total number of tokens across the global batch. val = torch.vstack(val).sum(dim=0) - dp_cp_group = pg_collection.dp_cp + dp_cp_group = get_data_distribution_group(pg_collection, cfg.model, with_context_parallel=True) torch.distributed.all_reduce(val, group=dp_cp_group) loss_reduced[key] = val[0] / val[1] elif val[0].numel() == 1: @@ -1576,7 +1578,7 @@ def _should_skip_and_handle_iteration( # Update step and sample counters global_state.train_state.step += 1 - dp_size = pg_collection.dp.size() + dp_size = get_data_distribution_group(pg_collection, cfg.model).size() batch_size = dp_size * cfg.train.micro_batch_size * get_num_microbatches() global_state.train_state.consumed_train_samples += batch_size global_state.train_state.skipped_train_samples += batch_size diff --git a/tests/functional_tests/test_groups/training/test_pretrain.py b/tests/functional_tests/test_groups/training/test_pretrain.py index 321a709850..ddbf0b8650 100644 --- a/tests/functional_tests/test_groups/training/test_pretrain.py +++ b/tests/functional_tests/test_groups/training/test_pretrain.py @@ -19,8 +19,13 @@ import pytest import torch import torch.nn.functional as F +from megatron.core import parallel_state +from megatron.core.tensor_parallel import gtp_api +from megatron.bridge.models.gpt.model_config import BridgeGPTModelConfig from megatron.bridge.models.gpt_provider import GPTModelProvider +from megatron.bridge.models.transformer_config import TransformerConfig +from megatron.bridge.training.callbacks import Callback, CallbackContext from megatron.bridge.training.config import ( CheckpointConfig, ConfigContainer, @@ -35,6 +40,7 @@ ValidationConfig, ) from megatron.bridge.training.gpt_step import forward_step +from megatron.bridge.training.gtp import get_transformer_config from megatron.bridge.training.pretrain import pretrain from tests.functional_tests.utils import ( broadcast_path, @@ -73,11 +79,128 @@ class Llama32TestModelProvider(GPTModelProvider): vocab_size: int | None = None +class GTPValidationCallback(Callback): + """Validate GTP groups, parameters, and finite per-step training results. + + This callback is a GTP runtime smoke test, not an MLM numerical-parity or + tensor-parity assertion. + """ + + def __init__(self) -> None: + self.num_steps = 0 + + def on_train_start(self, context: CallbackContext) -> None: + transformer_config = get_transformer_config(context.state.cfg.model) + gtp_size = transformer_config.gtp_weight_remat_size + assert gtp_size == 2 + assert parallel_state.get_gtp_weight_remat_world_size() == gtp_size + assert parallel_state.get_data_parallel_world_size(with_gtp_remat=False) == 2 // gtp_size + assert parallel_state.get_data_parallel_world_size(with_gtp_remat=True) == 2 + gtp_params = [param for chunk in context.model for param in chunk.parameters() if gtp_api.is_gtp_param(param)] + assert gtp_params + assert all(param.chain_id is not None for param in gtp_params) + + def on_train_step_end(self, context: CallbackContext) -> None: + assert context.skipped_iter == 0 + assert context.grad_norm is not None and torch.isfinite(torch.tensor(context.grad_norm)) + assert context.loss_dict + assert all(torch.isfinite(loss).all() for loss in context.loss_dict.values()) + self.num_steps += 1 + + def on_train_end(self, context: CallbackContext) -> None: + assert self.num_steps == 3 + assert context.state.train_state.consumed_train_samples == 6 + + class TestPretrain: """ Test end to end training with checkpoint functionality. """ + @pytest.mark.run_only_on("GPU") + @pytest.mark.skipif(not gtp_api.HAVE_GTP, reason="GTP requires TransformerEngine >= 2.19") + def test_pretrain_with_generalized_tensor_parallelism(self): + """Train three finite BF16 steps with dense weights sharded over a two-rank GTP axis.""" + initialize_distributed() + total_iters = 3 + seq_length = 64 + transformer_cfg = TransformerConfig( + num_layers=2, + hidden_size=128, + ffn_hidden_size=256, + num_attention_heads=4, + tensor_model_parallel_size=1, + tensor_parallel_num_weight_shards=2, + pipeline_model_parallel_size=1, + context_parallel_size=1, + sequence_parallel=False, + attention_dropout=0.0, + hidden_dropout=0.0, + add_bias_linear=False, + gradient_accumulation_fusion=False, + params_dtype=torch.bfloat16, + pipeline_dtype=torch.bfloat16, + bf16=True, + ) + model_cfg = BridgeGPTModelConfig( + transformer=transformer_cfg, + vocab_size=128, + seq_length=seq_length, + share_embeddings_and_output_weights=False, + ) + cfg = ConfigContainer( + model=model_cfg, + train=TrainingConfig( + train_iters=total_iters, + global_batch_size=2, + micro_batch_size=1, + ), + validation=ValidationConfig(eval_interval=100, eval_iters=0), + optimizer=OptimizerConfig( + optimizer="adam", + bf16=True, + fp16=False, + params_dtype=torch.bfloat16, + use_distributed_optimizer=False, + clip_grad=1.0, + lr=3e-4, + min_lr=3e-5, + ), + scheduler=SchedulerConfig( + start_weight_decay=0.01, + end_weight_decay=0.01, + weight_decay_incr_style="constant", + lr_decay_style="cosine", + lr_warmup_iters=1, + lr_decay_iters=total_iters, + override_opt_param_scheduler=True, + ), + ddp=DistributedDataParallelConfig( + check_for_nan_in_grad=True, + grad_reduce_in_fp32=True, + overlap_grad_reduce=False, + overlap_param_gather=False, + average_in_collective=False, + use_distributed_optimizer=False, + ), + dataset=MockGPTDatasetConfig( + random_seed=1234, + seq_length=seq_length, + reset_position_ids=False, + reset_attention_mask=False, + eod_mask_loss=False, + dataloader_type="single", + num_workers=0, + ), + logger=LoggerConfig(log_interval=1), + tokenizer=TokenizerConfig(tokenizer_type="NullTokenizer", vocab_size=128), + checkpoint=CheckpointConfig(save=None, load=None), + rng=RNGConfig(seed=1234), + ) + + callback = GTPValidationCallback() + pretrain(cfg, forward_step, callbacks=[callback]) + @pytest.mark.run_only_on("GPU") def test_pretrain_with_checkpoint(self, tmp_path): """ diff --git a/tests/unit_tests/training/test_decentralized_pg.py b/tests/unit_tests/training/test_decentralized_pg.py index 81c017dbda..537605ffd9 100644 --- a/tests/unit_tests/training/test_decentralized_pg.py +++ b/tests/unit_tests/training/test_decentralized_pg.py @@ -414,6 +414,8 @@ def test_initialize_distributed_raises_on_no_cuda_devices(self, mock_get_rank, m mock_model_config.tensor_model_parallel_size = 1 mock_model_config.pipeline_model_parallel_size = 1 mock_model_config.context_parallel_size = 1 + mock_model_config.gtp_weight_remat_size = 1 + mock_model_config.expert_gtp_weight_remat_size = 1 mock_dist_config = MagicMock() @@ -452,6 +454,8 @@ def test_uses_hyper_comm_grid_when_decentralized_pg_enabled( mock_model_config.tensor_model_parallel_size = 1 mock_model_config.pipeline_model_parallel_size = 1 mock_model_config.context_parallel_size = 1 + mock_model_config.gtp_weight_remat_size = 1 + mock_model_config.expert_gtp_weight_remat_size = 1 mock_dist_config = MagicMock() mock_dist_config.use_decentralized_pg = True @@ -481,6 +485,7 @@ def test_uses_hyper_comm_grid_when_decentralized_pg_enabled( @patch("megatron.bridge.training.initialize.ProcessGroupCollection") @patch("megatron.bridge.training.initialize._create_pg_collection") @patch("megatron.bridge.training.initialize.parallel_state") + @patch("megatron.core.tensor_parallel.gtp_api.HAVE_GTP", True) @patch("torch.cuda.device_count", return_value=1) @patch("torch.distributed.is_initialized", return_value=True) @patch("megatron.bridge.training.initialize.get_rank_safe", return_value=0) @@ -507,6 +512,8 @@ def test_uses_mpu_when_decentralized_pg_disabled( mock_model_config.hierarchical_context_parallel_sizes = None mock_model_config.expert_model_parallel_size = 4 mock_model_config.expert_tensor_parallel_size = expert_tensor_parallel_size + mock_model_config.gtp_weight_remat_size = 2 + mock_model_config.expert_gtp_weight_remat_size = 3 mock_dist_config = MagicMock() mock_dist_config.use_decentralized_pg = False @@ -538,6 +545,42 @@ def test_uses_mpu_when_decentralized_pg_disabled( mock_parallel_state.initialize_model_parallel.call_args.kwargs["expert_tensor_parallel_size"] == expected_expert_tensor_parallel_size ) + assert mock_parallel_state.initialize_model_parallel.call_args.kwargs["gtp_remat_size"] == 2 + assert mock_parallel_state.initialize_model_parallel.call_args.kwargs["expert_gtp_remat_size"] == 3 + + @patch("megatron.bridge.training.initialize._create_pg_collection") + @patch("megatron.bridge.training.initialize.parallel_state") + @patch("torch.cuda.device_count", return_value=1) + @patch("torch.distributed.is_initialized", return_value=True) + @patch("megatron.bridge.training.initialize.get_rank_safe", return_value=0) + def test_rejects_gtp_with_decentralized_process_groups( + self, + mock_get_rank, + mock_is_init, + mock_device_count, + mock_parallel_state, + mock_create_pg_collection, + ): + """GTP must not silently use process groups that lack its remat axes.""" + from megatron.bridge.training.initialize import _initialize_distributed + + mock_model_config = MagicMock() + mock_model_config.expert_tensor_parallel_size = 1 + mock_model_config.gtp_weight_remat_size = 2 + mock_model_config.expert_gtp_weight_remat_size = 1 + mock_dist_config = MagicMock(use_decentralized_pg=True) + + with pytest.raises(NotImplementedError, match="standard MCore process-group runtime"): + _initialize_distributed( + model_config=mock_model_config, + dist_config=mock_dist_config, + num_distributed_optimizer_instances=1, + get_embedding_ranks=None, + get_position_embedding_ranks=None, + ) + + mock_create_pg_collection.assert_not_called() + mock_parallel_state.initialize_model_parallel.assert_not_called() class TestSetupUsesDecentralizedPg: diff --git a/tests/unit_tests/training/test_gtp.py b/tests/unit_tests/training/test_gtp.py new file mode 100644 index 0000000000..ba1b0a2c5a --- /dev/null +++ b/tests/unit_tests/training/test_gtp.py @@ -0,0 +1,149 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Tests for Generalized Tensor Parallelism runtime wiring.""" + +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +import pytest + +from megatron.bridge.models.transformer_config import TransformerConfig +from megatron.bridge.training.gtp import ( + classify_gtp_remat_chains, + configure_gtp_remat, + get_data_distribution_group, +) + + +def _gtp_config(*, dense_size: int = 2, expert_size: int = 1) -> SimpleNamespace: + return SimpleNamespace( + gtp_weight_remat_size=dense_size, + expert_gtp_weight_remat_size=expert_size, + fp4=None, + fp8_recipe=None, + fp8=None, + calculate_per_token_loss=True, + cuda_graph_modules=["attn"], + moe_shared_expert_overlap=False, + cuda_graph_impl="none", + ) + + +def test_transformer_config_derives_gtp_sizes_from_weight_shards(): + config = TransformerConfig( + num_layers=2, + hidden_size=64, + num_attention_heads=4, + tensor_model_parallel_size=2, + tensor_parallel_num_weight_shards=8, + expert_tensor_parallel_size=2, + expert_tensor_parallel_num_weight_shards=6, + ) + + config.finalize() + + assert config.gtp_weight_remat_size == 4 + assert config.expert_gtp_weight_remat_size == 3 + + +@pytest.mark.parametrize("num_weight_shards", [1, 3]) +def test_transformer_config_rejects_invalid_gtp_weight_shards(num_weight_shards): + config = TransformerConfig( + num_layers=2, + hidden_size=64, + num_attention_heads=4, + tensor_model_parallel_size=2, + tensor_parallel_num_weight_shards=num_weight_shards, + ) + + with pytest.raises(ValueError, match="tensor_parallel_num_weight_shards"): + config.finalize() + + +@patch("megatron.core.tensor_parallel.gtp_api.configure_gtp_remat_from_recipe") +@patch("megatron.core.tensor_parallel.gtp_api.HAVE_GTP", True) +def test_configure_gtp_remat_forwards_transformer_recipe(mock_configure): + config = _gtp_config() + + configure_gtp_remat(config) + + mock_configure.assert_called_once_with( + fp4=False, + fp8_recipe=None, + fp8=False, + calculate_per_token_loss=True, + ) + + +@patch("megatron.core.tensor_parallel.gtp_api.classify_gtp_remat_chains") +def test_classify_gtp_remat_chains_receives_all_model_chunks(mock_classify): + config = _gtp_config() + model = [MagicMock(), MagicMock()] + + classify_gtp_remat_chains(model, config) + + mock_classify.assert_called_once_with( + model, + cuda_graph_modules=["attn"], + moe_shared_expert_overlap=False, + cuda_graph_impl="none", + ) + + +def test_gtp_off_preserves_existing_data_parallel_groups(): + config = _gtp_config(dense_size=1, expert_size=1) + pg_collection = SimpleNamespace(dp=object(), dp_cp=object()) + + assert get_data_distribution_group(pg_collection, config) is pg_collection.dp + assert get_data_distribution_group(pg_collection, config, with_context_parallel=True) is pg_collection.dp_cp + + +@patch("megatron.bridge.training.gtp.parallel_state.get_data_parallel_group") +def test_gtp_uses_full_data_distribution_groups(mock_get_data_parallel_group): + config = _gtp_config() + full_dp_group = object() + full_dp_cp_group = object() + pg_collection = SimpleNamespace(dp_cp_gtp_remat=full_dp_cp_group) + mock_get_data_parallel_group.return_value = full_dp_group + + assert get_data_distribution_group(pg_collection, config) is full_dp_group + assert get_data_distribution_group(pg_collection, config, with_context_parallel=True) is full_dp_cp_group + mock_get_data_parallel_group.assert_called_once_with(with_gtp_remat=True) + + +@patch("megatron.bridge.training.setup.classify_gtp_remat_chains") +@patch("megatron.bridge.training.setup.configure_gtp_remat") +def test_distributed_model_build_obeys_gtp_lifecycle(mock_configure, mock_classify): + from megatron.bridge.training.setup import _build_distributed_model + + events = [] + model = [MagicMock(), MagicMock()] + model_config = SimpleNamespace() + model_config.finalize = MagicMock(side_effect=lambda: events.append("finalize")) + model_config.provide_distributed_model = MagicMock(side_effect=lambda **_kwargs: events.append("build") or model) + mock_configure.side_effect = lambda _config: events.append("configure") + mock_classify.side_effect = lambda _model, _config: events.append("classify") + cfg = SimpleNamespace( + model=model_config, + ddp=object(), + optimizer=SimpleNamespace(overlap_param_gather_with_optimizer_step=False), + dist=SimpleNamespace(use_megatron_fsdp=False, use_torch_fsdp2=False), + rng=SimpleNamespace(data_parallel_random_init=False), + ) + + result = _build_distributed_model(cfg, MagicMock()) + + assert result is model + assert events == ["finalize", "configure", "build", "classify"] diff --git a/tests/unit_tests/training/test_pg_collection_wiring.py b/tests/unit_tests/training/test_pg_collection_wiring.py index 19e6f30d82..aa78c03571 100644 --- a/tests/unit_tests/training/test_pg_collection_wiring.py +++ b/tests/unit_tests/training/test_pg_collection_wiring.py @@ -35,28 +35,33 @@ def test_should_skip_iteration_uses_passed_pg_collection(monkeypatch): # Set up a minimal config needed by _should_skip_and_handle_iteration # iterations_to_skip uses 1-based iteration numbers (matching MLM convention). # {1} means "skip the 1st iteration", which fires when step=0 (step+1==1). + model_config = SimpleNamespace() state.cfg = SimpleNamespace( + model=model_config, train=SimpleNamespace( iterations_to_skip={1}, micro_batch_size=4, exit_signal_handler=False, exit_signal=signal.SIGTERM, - ) + ), ) - # Fake pg_collection with a DP size + # Fake full data-distribution group with a DP size. class _DP: def size(self): return 3 - class _PG: - def __init__(self): - self.dp = _DP() + fake_pg = SimpleNamespace() + data_distribution_group = _DP() + group_calls = [] - fake_pg = _PG() + def _get_data_distribution_group(pg_collection, received_model_config): + group_calls.append((pg_collection, received_model_config)) + return data_distribution_group # Ensure deterministic microbatch count without touching global calculators monkeypatch.setattr(train_module, "get_num_microbatches", lambda: 2) + monkeypatch.setattr(train_module, "get_data_distribution_group", _get_data_distribution_group) # Avoid any distributed or pipeline logic inside the dummy step monkeypatch.setattr(train_module, "_dummy_train_step", lambda *args, **kwargs: None) @@ -71,6 +76,7 @@ def __init__(self): # Assert assert did_skip is True + assert group_calls == [(fake_pg, model_config)] # One iteration skipped assert state.train_state.step == 1 # Batch size = dp.size * micro_batch_size * num_microbatches = 3 * 4 * 2 = 24