From bb85dbd6c13e6402ce8f974eb60f717b15bbbb38 Mon Sep 17 00:00:00 2001 From: Jingyue Wu Date: Fri, 7 Aug 2026 04:16:24 +0000 Subject: [PATCH] fix(mfsdp): refresh V2 compute weights after optimizer step Signed-off-by: Jingyue Wu --- megatron/core/optimizer/fully_sharded_optimizer.py | 14 +++++++++++++- .../distributed/mfsdp_v2/test_mcore_adapter.py | 4 ++-- 2 files changed, 15 insertions(+), 3 deletions(-) diff --git a/megatron/core/optimizer/fully_sharded_optimizer.py b/megatron/core/optimizer/fully_sharded_optimizer.py index 18c2354dcb2..b15deafa961 100644 --- a/megatron/core/optimizer/fully_sharded_optimizer.py +++ b/megatron/core/optimizer/fully_sharded_optimizer.py @@ -8,6 +8,9 @@ from ..config_logger import has_config_logger_enabled, log_config_to_disk from ..dist_checkpointing.mapping import ShardedStateDict +from ..distributed.fsdp.src.megatron_fsdp.experimental.parameter_group import ( + get_containing_parameter_group, +) from ..transformer.module import MegatronModule from .grad_scaler import MegatronGradScaler from .optimizer import MixedPrecisionOptimizer @@ -120,7 +123,16 @@ def _copy_model_grads_to_main_grads(self) -> None: """No-op: MFSDP v2 reduces directly into optimizer-visible sharded grads.""" def _copy_main_params_to_model_params(self) -> None: - """No-op: MFSDP v2 currently syncs compute weights in its forward pre-hook.""" + """Refresh MFSDP V2 compute weights after updating optimizer weights.""" + # TODO(wujingyue): Reuse experimental.checkpoint.resync_compute_weights after #6024 merges. + seen_parameter_groups = set() + for parameter in self.get_parameters(): + if (parameter_group := get_containing_parameter_group(parameter)) is None: + continue + if parameter_group in seen_parameter_groups: + continue + seen_parameter_groups.add(parameter_group) + parameter_group.sync_model_weight_from_main_weight() def _copy_model_params_to_main_params(self, state_dict=None) -> None: """No-op: model loads already write into MFSDP v2's main weights.""" diff --git a/tests/unit_tests/distributed/mfsdp_v2/test_mcore_adapter.py b/tests/unit_tests/distributed/mfsdp_v2/test_mcore_adapter.py index a132ff88139..18f236d6bb6 100644 --- a/tests/unit_tests/distributed/mfsdp_v2/test_mcore_adapter.py +++ b/tests/unit_tests/distributed/mfsdp_v2/test_mcore_adapter.py @@ -164,7 +164,7 @@ def test_build_train_and_step(self): torch.randn(8, 2, config.hidden_size, device="cuda", dtype=torch.bfloat16) for _ in range(2) ] - for _ in range(3) + for _ in range(10) ] reference_losses = [] @@ -198,4 +198,4 @@ def test_build_train_and_step(self): reference_losses = torch.stack(reference_losses) assert torch.isfinite(losses).all() assert torch.isfinite(reference_losses).all() - torch.testing.assert_close(losses, reference_losses, rtol=1e-2, atol=0) + torch.testing.assert_close(losses, reference_losses, rtol=1e-3, atol=0)