Skip to content
Open
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
14 changes: 13 additions & 1 deletion megatron/core/optimizer/fully_sharded_optimizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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."""
4 changes: 2 additions & 2 deletions tests/unit_tests/distributed/mfsdp_v2/test_mcore_adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 = []
Expand Down Expand Up @@ -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)
Loading