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
53 changes: 42 additions & 11 deletions megatron/core/optimizer/layer_wise_optimizer.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,12 @@
# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.

from typing import List, Optional
from typing import Callable, List, Optional

import torch

from megatron.core.dist_checkpointing import ShardedTensor
from megatron.core.dist_checkpointing.dict_utils import nested_values
from megatron.core.dist_checkpointing.mapping import ShardedStateDict
from megatron.core.process_groups_config import ProcessGroupCollection
from megatron.core.utils import get_pg_rank, get_pg_size

Expand Down Expand Up @@ -36,17 +39,36 @@ def __init__(
optimizers: List[MegatronOptimizer],
config: OptimizerConfig,
pg_collection: Optional[ProcessGroupCollection] = None,
init_state_fn_list: Optional[List[Callable]] = None,
) -> None:
"""
Initialize LayerWiseDistributedOptimizer.

Args:
optimizers: List of MegatronOptimizers.
config: OptimizerConfig.
pg_collection: ProcessGroupCollection.
init_state_fn_list: List of init state functions.
"""

self.pg_collection = pg_collection
self.shard_params(optimizers)
# wrap optimizer after sharding to avoid unnecessary master weight creation
# TODO(deyuf): check if underlying optimizer.config need to fixed and if so can use
# that instead of passing
if init_state_fn_list is None:
init_state_fn_list = [None] * len(optimizers)
else:
assert len(init_state_fn_list) == len(optimizers), (
"init_state_fn_list must be the " "same length as optimizers if provided"
)

if config.bf16:
if isinstance(optimizers[0], Float16OptimizerWithFloat16Params):
raise TypeError('LayerWiseDistributedOptimizer received Float16 optimizer already.')
optimizers = [
Float16OptimizerWithFloat16Params(optim, config, None, None) for optim in optimizers
Float16OptimizerWithFloat16Params(optim, config, None, init_state_fn_list[idx])
for idx, optim in enumerate(optimizers)
]
super().__init__(optimizers)

Expand Down Expand Up @@ -152,14 +174,23 @@ def step(self): # type: ignore[no-untyped-def]

return update_successful, grad_norm, num_zeros_in_grad

def save_state_dict_to_file(self, filename: str) -> None:
"""Save the parameter state of the optimizer.

Args:
filename: The filename to save the parameter state.
def sharded_state_dict(
self, model_sharded_state_dict: ShardedStateDict, is_loading: bool = False, **kwargs
):
"""
Sharded state dict for torch_dist format checkpointing.
For fixed DP usage only, set replica_id to 0 for all ShardedTensor.
"""
torch.save(super().state_dict(), filename)
sharded_state_dict = super().sharded_state_dict(
model_sharded_state_dict, is_loading, **kwargs
)

# for fixed DP usage only
for sh_base in nested_values(sharded_state_dict):
if isinstance(sh_base, ShardedTensor):
assert (
len(sh_base.replica_id) == 3
), f'Expected replica_id format (PP, TP, DP), got: {sh_base}'
sh_base.replica_id = (*sh_base.replica_id[:2], 0)

def load_state_dict_from_file(self, filename: str) -> None:
"""Load the parameter state of the optimizer."""
super().load_state_dict(torch.load(filename))
# return sharded_state_dict
36 changes: 32 additions & 4 deletions megatron/core/optimizer/muon.py
Original file line number Diff line number Diff line change
Expand Up @@ -280,22 +280,45 @@ def get_megatron_muon_optimizer(
# TODO(deyuf): allow user to select optimizer mix and relax ChainedOptimizer design
config.optimizer = 'adam'

# Needed for torch_dist ckpt_format, unlike torch ckpt_format
# For other emerging optimizers, need to implement init_state_fn as well
# TODO(boxiangw): Improve usability after optimizer refactor
# TODO(boxiangw): support precision aware optimizer
def muon_init_state_fn(opt, config=None):
for group in opt.param_groups:
for p in group['params']:
if len(opt.state[p]) == 0:
opt.state[p]['momentum_buffer'] = torch.zeros_like(p.data)

def adam_init_state_fn(opt, config=None):
for group in opt.param_groups:
for p in group['params']:
if len(opt.state[p]) == 0:
if config is None or not config.use_precision_aware_optimizer:
opt.state[p]['exp_avg'] = torch.zeros_like(p.data)
opt.state[p]['exp_avg_sq'] = torch.zeros_like(p.data)
else:
opt.initialize_state(p)

# need to wrap into megatron mix precision optimizer. (only support bf16 w/o loss scale now)
if config.fp16:
raise Exception('muon with fp16 is not supported.')

reset_config_bf16 = False
if config.bf16:
if layer_wise_distributed_optimizer:
# creating master weight before layerwise sharding will lead to unnecessary master
# weight so here we delay master weight creation into layer_wise unset config.bf16
# weight so here we delay master weight creation into layer_wise unset config.bf16
# will also result in all optimizers below(adam) to also not be wrapped
config.bf16 = False
reset_config_bf16 = True
else:
# if not using layer_wise wrapper, just create master weight here is fine
optimizer = Float16OptimizerWithFloat16Params(optimizer, config, None, None)
optimizer = Float16OptimizerWithFloat16Params(
optimizer, config, None, muon_init_state_fn
)
else:
optimizer = FP32Optimizer(optimizer, config, None)
optimizer = FP32Optimizer(optimizer, config, muon_init_state_fn)

optimizers.append(optimizer)

Expand All @@ -321,5 +344,10 @@ def get_megatron_muon_optimizer(
log_single_rank(logger, logging.INFO, 'Using LayerWiseDistributedOptimizer for Muon')
if reset_config_bf16:
config.bf16 = True
return LayerWiseDistributedOptimizer(optimizers, config, pg_collection)
return LayerWiseDistributedOptimizer(
optimizers,
config,
pg_collection,
init_state_fn_list=[muon_init_state_fn, adam_init_state_fn],
)
return ChainedOptimizer(optimizers)
2 changes: 1 addition & 1 deletion megatron/training/arguments.py
Original file line number Diff line number Diff line change
Expand Up @@ -1181,7 +1181,7 @@ def validate_args(args, defaults={}):
assert not args.use_distributed_optimizer, "Muon optimizer does not support distributed optimizer for now."
assert not args.use_torch_fsdp2, "Muon optimizer does not support Torch-FSDP2 for now."
assert not args.use_megatron_fsdp, "Muon optimizer does not support Megatron-FSDP for now."
assert args.ckpt_format == "torch", "Muon optimizer only supports torch checkpoint format for now."
assert args.ckpt_format in ["torch", "torch_dist"], "Muon optimizer supports torch and torch_dist checkpoint format."

# Optimizer CPU offload check
if args.optimizer_cpu_offload:
Expand Down
18 changes: 5 additions & 13 deletions megatron/training/checkpointing.py
Original file line number Diff line number Diff line change
Expand Up @@ -479,14 +479,6 @@ def save_checkpoint(iteration, model, optimizer, opt_param_scheduler, num_floati
if not optimizer.is_stub_optimizer:
optimizer.save_parameter_state(optim_checkpoint_name)

# LayerWiseDistributedOptimizer save
if getattr(args, "optimizer", "adam").startswith("dist_"):
dp_rank = mpu.get_data_parallel_rank()
optim_checkpoint_name = os.path.join(os.path.dirname(checkpoint_name), f"layer_wise_optimizer_{dp_rank}.pt")
ensure_directory_exists(optim_checkpoint_name)
if not optimizer.is_stub_optimizer:
optimizer.save_state_dict_to_file(optim_checkpoint_name)

async_save_request = None
if args.async_save:
if ckpt_type == CheckpointType.LEGACY:
Expand Down Expand Up @@ -837,6 +829,10 @@ def generate_state_dict(
else:
optimizer_sd = optimizer.state_dict()

# check if optimizer_sd is not None
if optimizer_sd is None:
raise ValueError(f"optimizer_sd is None for rank {torch.distributed.get_rank()}")

state_dict['optimizer'] = optimizer_sd

if opt_param_scheduler is not None:
Expand Down Expand Up @@ -1661,11 +1657,7 @@ def load_model_state_dict(module, state_dict, strict: bool):
if not release and not args.finetune and not args.no_load_optim:
try:
# Load state dict.
if getattr(args, "optimizer", "adam").startswith("dist_"):
dp_rank = mpu.get_data_parallel_rank()
optim_checkpoint_name = os.path.join(os.path.dirname(checkpoint_name), f"layer_wise_optimizer_{dp_rank}.pt")
optimizer.load_state_dict_from_file(optim_checkpoint_name)
elif not skip_load_to_model_and_opt and optimizer is not None and not optimizer.is_stub_optimizer:
if not skip_load_to_model_and_opt and optimizer is not None and not optimizer.is_stub_optimizer:
optimizer.load_state_dict(state_dict['optimizer'])

# Load distributed optimizer's custom parameter state.
Expand Down
45 changes: 27 additions & 18 deletions tests/unit_tests/test_layer_wise_optimizer.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.
import os
import tempfile

import pytest
import torch
Expand Down Expand Up @@ -224,29 +224,38 @@ def test_state_dict(self):
# TODO(deyuf): fix this. not going through get() will cause missing keys like wd_mult
# optimizer.load_state_dict(state_dict)

def test_save_load_file(self):
"""Test LayerWiseDistributedOptimizer save and load state dict to/from file."""
def test_sharded_state_dict(self):
"""Test LayerWiseDistributedOptimizer sharded_state_dict method."""
model, optimizer, pg_collection = self.create_model_and_optimizer()

for param in model.parameters():
param.grad = torch.randn_like(param)
optimizer.step()

# Test save to file
with tempfile.NamedTemporaryFile(delete=False, suffix='.pt') as tmp_file:
temp_filename = tmp_file.name

try:
optimizer.save_state_dict_to_file(temp_filename)
assert os.path.exists(temp_filename), "State dict file should be created"

# Test load from file
# TODO(deyuf): fix this. not going through get() will cause missing keys like wd_mult
# optimizer.load_state_dict_from_file(temp_filename)
finally:
# Clean up temporary file
if os.path.exists(temp_filename):
os.remove(temp_filename)
# Get model sharded state dict
model_sharded_state_dict = model.sharded_state_dict()

# Test sharded_state_dict
sharded_state_dict = optimizer.sharded_state_dict(model_sharded_state_dict)

# Verify the sharded_state_dict is not None and has expected structure
assert sharded_state_dict is not None, "Sharded state dict should not be None"
assert (
'optimizer' in sharded_state_dict
), "Sharded state dict should contain 'optimizer' key"

# Verify that replica_id is set correctly (should be 0 for DP dimension)
from megatron.core.dist_checkpointing import ShardedTensor
from megatron.core.dist_checkpointing.dict_utils import nested_values

for sh_base in nested_values(sharded_state_dict):
if isinstance(sh_base, ShardedTensor):
assert (
len(sh_base.replica_id) == 3
), f'Expected replica_id format (PP, TP, DP), got: {sh_base.replica_id}'
assert (
sh_base.replica_id[2] == 0
), f'Expected DP replica_id to be 0 for layer-wise optimizer, got: {sh_base.replica_id[2]}'

def test_multiple_optimizers(self):
"""Test LayerWiseDistributedOptimizer with multiple chained optimizers.
Expand Down
Loading