From cb6668d190f7b809922077df96de2496de0dffd3 Mon Sep 17 00:00:00 2001 From: Philip Petrakian Date: Mon, 10 Aug 2026 20:50:14 +0000 Subject: [PATCH] Remove PagedStashRunner root-config assumption Signed-off-by: Philip Petrakian --- megatron/core/transformer/moe/paged_stash.py | 29 ++- .../moe/test_paged_stash_runner.py | 198 ++++++++++++++++++ 2 files changed, 219 insertions(+), 8 deletions(-) create mode 100644 tests/unit_tests/transformer/moe/test_paged_stash_runner.py diff --git a/megatron/core/transformer/moe/paged_stash.py b/megatron/core/transformer/moe/paged_stash.py index 5bef3d321b5..3badf7b5e09 100644 --- a/megatron/core/transformer/moe/paged_stash.py +++ b/megatron/core/transformer/moe/paged_stash.py @@ -976,15 +976,15 @@ def __init__(self, config, copy_main_params, model, optimizer, forward_backward_ self.optimizer = optimizer self.forward_backward_func = forward_backward_func self.moe_layers = [] - # TransformerConfig objects that must stay in sync for moe_paged_stash: the training - # loop `config` (schedules / paged_stash_reset) plus each VP chunk's GPT root config - # (GPTModel.forward). MoE mlps use the same config reference as that root, so we do - # not track mlp.config separately. + # Config objects that must stay in sync for moe_paged_stash: the training loop config + # (schedules / paged_stash_reset), each model chunk's root config (model forward), and + # every MoE layer config (expert forward). Some models may use a distinct config for + # each layer. seen_cfg_ids = set() self._configs_to_sync_moe_paged_stash = [] def _track_cfg(c): - if c is None: + if c is None or not hasattr(c, 'moe_paged_stash'): return cid = id(c) if cid not in seen_cfg_ids: @@ -1002,6 +1002,16 @@ def _track_cfg(c): model_chunk, "decoder", allow_none=False, return_model_obj=True ) _track_cfg(model_with_decoder.config) + + # Track MoE configs independently from the existing structural discovery below. + # This keeps overflow and retry behavior unchanged for models whose modules share + # the root config while allowing distinct module configs to stay synchronized. + for module in model_with_decoder.modules(): + token_dispatcher = getattr(module, 'token_dispatcher', None) + if token_dispatcher is None or not hasattr(token_dispatcher, 'check_over_budget'): + continue + _track_cfg(getattr(module, 'config', None)) + for layer in model_with_decoder.decoder.layers: transformer_layer = ( layer.mtp_model_layer if isinstance(layer, MultiTokenPredictionLayer) else layer @@ -1029,7 +1039,7 @@ def _track_cfg(c): self.moe_layers.append(mlp) def _set_moe_paged_stash_all(self, value: bool) -> None: - """Set moe_paged_stash on every tracked config (train + per VP chunk root).""" + """Set moe_paged_stash on every tracked training, model, and MoE config.""" for c in self._configs_to_sync_moe_paged_stash: c.moe_paged_stash = value @@ -1179,7 +1189,9 @@ def __call__(self, *args, **kwargs): training = not kwargs['forward_only'] data_iterator = kwargs['data_iterator'] - saved_moe_paged_stash = self.config.moe_paged_stash + saved_moe_paged_stash_values = [ + (config, config.moe_paged_stash) for config in self._configs_to_sync_moe_paged_stash + ] num_tries = 0 while True: assert ( @@ -1214,7 +1226,8 @@ def __call__(self, *args, **kwargs): mlp.token_dispatcher._comm_manager.moe_expert_rank_capacity_factor = ( mlp.token_dispatcher.config.moe_expert_rank_capacity_factor ) - self._set_moe_paged_stash_all(saved_moe_paged_stash) + for config, value in saved_moe_paged_stash_values: + config.moe_paged_stash = value break # Overflow or over-budget: prepare_for_rerun clears capacity factor and paged stash. diff --git a/tests/unit_tests/transformer/moe/test_paged_stash_runner.py b/tests/unit_tests/transformer/moe/test_paged_stash_runner.py new file mode 100644 index 00000000000..fd80d283ee8 --- /dev/null +++ b/tests/unit_tests/transformer/moe/test_paged_stash_runner.py @@ -0,0 +1,198 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +from types import SimpleNamespace + +import torch + +from megatron.core.transformer.moe.paged_stash import PagedStashManager, PagedStashRunner + + +def _config(moe_paged_stash): + return SimpleNamespace(moe_paged_stash=moe_paged_stash, moe_expert_rank_capacity_factor=1.5) + + +class _FakeTokenDispatcher: + + def __init__(self, config): + self.config = config + self._comm_manager = SimpleNamespace(moe_expert_rank_capacity_factor=1.5) + self.reset_count = 0 + + def check_over_budget(self): + return None + + def reset_over_budget(self): + self.reset_count += 1 + + +class _FakeMoELayer(torch.nn.Module): + + def __init__(self, config): + super().__init__() + self.config = config + self.token_dispatcher = _FakeTokenDispatcher(config) + + +class _FakeTransformerLayer(torch.nn.Module): + + def __init__(self, mlp): + super().__init__() + self.mlp = mlp + + +class _FakeStack(torch.nn.Module): + + def __init__(self, layer): + super().__init__() + self.layers = torch.nn.ModuleList([layer]) + + +class _FakeMTPPredictionLayer(torch.nn.Module): + + def __init__(self, mtp_model_layer): + super().__init__() + self.mtp_model_layer = mtp_model_layer + + +class _FakeModelChunk(torch.nn.Module): + + def __init__(self, config, decoder_moe, mtp_moe, nested_mtp): + super().__init__() + self.config = config + self.decoder = _FakeStack(_FakeTransformerLayer(decoder_moe)) + mtp_model_layer = _FakeTransformerLayer(mtp_moe) + if nested_mtp: + mtp_model_layer = _FakeStack(mtp_model_layer) + self.mtp = _FakeStack(_FakeMTPPredictionLayer(mtp_model_layer)) + self.mtp_process = True + # Register the decoder MoE through a second path to verify identity deduplication. + self.duplicate_decoder_moe = decoder_moe + self.zero_grad_count = 0 + + def zero_grad_buffer(self): + self.zero_grad_count += 1 + + +def _run_retry( + monkeypatch, training_config, model_config, decoder_config, mtp_config, nested_mtp=False +): + monkeypatch.setattr( + "megatron.core.transformer.multi_token_prediction.MultiTokenPredictionLayer", + _FakeMTPPredictionLayer, + ) + decoder_moe = _FakeMoELayer(decoder_config) + mtp_moe = _FakeMoELayer(mtp_config) + model = _FakeModelChunk(model_config, decoder_moe, mtp_moe, nested_mtp=nested_mtp) + + values_seen_by_forward = [] + + def forward_backward_func(**_): + values_seen_by_forward.append( + ( + training_config.moe_paged_stash, + model_config.moe_paged_stash, + decoder_config.moe_paged_stash, + mtp_config.moe_paged_stash, + ) + ) + return len(values_seen_by_forward) + + release_stash_buffer_calls = [] + fake_stash_manager = SimpleNamespace( + overflow=None, + host_spill=None, + release_stash_buffers=lambda: release_stash_buffer_calls.append(None), + ) + monkeypatch.setattr(PagedStashManager, 'STASH_MGR', fake_stash_manager) + runner = PagedStashRunner( + config=training_config, + copy_main_params=False, + model=[model], + optimizer=None, + forward_backward_func=forward_backward_func, + ) + overflow_results = iter([(1, 0, 0), (0, 0, 0)]) + runner.check_moe_overflow = lambda: next(overflow_results) + + result = runner( + model=[model], data_iterator=None, num_microbatches=1, seq_length=1, forward_only=False + ) + + return SimpleNamespace( + runner=runner, + result=result, + model=model, + decoder_moe=decoder_moe, + mtp_moe=mtp_moe, + values_seen_by_forward=values_seen_by_forward, + release_stash_buffer_calls=release_stash_buffer_calls, + ) + + +def test_retry_preserves_shared_root_config_behavior(monkeypatch): + """Models whose MoE modules share the root config retain their existing behavior.""" + training_config = _config(True) + model_config = _config(True) + + run = _run_retry( + monkeypatch, + training_config=training_config, + model_config=model_config, + decoder_config=model_config, + mtp_config=model_config, + ) + + assert run.runner.moe_layers == [run.decoder_moe, run.mtp_moe] + assert [id(config) for config in run.runner._configs_to_sync_moe_paged_stash] == [ + id(training_config), + id(model_config), + ] + assert run.values_seen_by_forward == [(True, True, True, True), (False, False, False, False)] + assert run.result == 2 + assert run.model.zero_grad_count == 1 + assert len(run.release_stash_buffer_calls) == 1 + assert run.decoder_moe.token_dispatcher.reset_count == 1 + assert run.mtp_moe.token_dispatcher.reset_count == 1 + assert run.decoder_moe.token_dispatcher._comm_manager.moe_expert_rank_capacity_factor == 1.5 + assert run.mtp_moe.token_dispatcher._comm_manager.moe_expert_rank_capacity_factor == 1.5 + assert training_config.moe_paged_stash is True + assert model_config.moe_paged_stash is True + + +def test_retry_disables_and_restores_per_module_configs(monkeypatch): + """Retry must disable direct and nested-MTP MoE configs, then restore each value.""" + training_config = _config(True) + model_config = _config(True) + decoder_moe_config = _config(True) + mtp_moe_config = _config(False) + + run = _run_retry( + monkeypatch, + training_config=training_config, + model_config=model_config, + decoder_config=decoder_moe_config, + mtp_config=mtp_moe_config, + nested_mtp=True, + ) + + assert run.runner.moe_layers == [run.decoder_moe] + assert [id(config) for config in run.runner._configs_to_sync_moe_paged_stash] == [ + id(training_config), + id(model_config), + id(decoder_moe_config), + id(mtp_moe_config), + ] + assert run.values_seen_by_forward == [(True, True, True, False), (False, False, False, False)] + assert run.result == 2 + assert run.model.zero_grad_count == 1 + assert len(run.release_stash_buffer_calls) == 1 + assert run.decoder_moe.token_dispatcher.reset_count == 1 + assert run.mtp_moe.token_dispatcher.reset_count == 0 + assert run.decoder_moe.token_dispatcher._comm_manager.moe_expert_rank_capacity_factor == 1.5 + assert run.mtp_moe.token_dispatcher._comm_manager.moe_expert_rank_capacity_factor == 1.5 + assert ( + training_config.moe_paged_stash, + model_config.moe_paged_stash, + decoder_moe_config.moe_paged_stash, + mtp_moe_config.moe_paged_stash, + ) == (True, True, True, False)