From e0fdf45f42031622207cbe88d041c6bf2d38d4d4 Mon Sep 17 00:00:00 2001 From: Deyu Fu Date: Mon, 27 Oct 2025 19:57:39 +0800 Subject: [PATCH 01/10] speed and memory improvements for layerwise distopt: - switch to faster all gather - sort param size and assign to rank in a ping-pong loop for balanced memory usage Signed-off-by: Deyu Fu --- .../core/optimizer/layer_wise_optimizer.py | 137 ++++++++++-------- 1 file changed, 80 insertions(+), 57 deletions(-) diff --git a/megatron/core/optimizer/layer_wise_optimizer.py b/megatron/core/optimizer/layer_wise_optimizer.py index 2b311dfe659..9e95cbebfa5 100644 --- a/megatron/core/optimizer/layer_wise_optimizer.py +++ b/megatron/core/optimizer/layer_wise_optimizer.py @@ -1,8 +1,9 @@ -# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved. +# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. from typing import Callable, List, Optional import torch +from torch._utils import _flatten_dense_tensors, _unflatten_dense_tensors from megatron.core.dist_checkpointing import ShardedTensor from megatron.core.dist_checkpointing.dict_utils import nested_values @@ -18,20 +19,17 @@ class LayerWiseDistributedOptimizer(ChainedOptimizer): """Layer-wise distributed optimizer for Megatron-core models. - This is a experimental distributed optimizer wrapper that distributes weight to DP ranks - by full layer. Implemented as ChainedOptimizer to support different weights use different - optimizers (e.g. muon+adam). When using, keep all megatron distributed optimizer related - options OFF. + Experimental distributed optimizer wrapper that distributes weight to DP ranks by layer. + Implemented as ChainedOptimizer to support multiple optimizers (e.g. muon + adamW) + When using, keep all megatron distributed-optimizer related options OFF. How LayerWiseDistributedOptimizer work: 1. weights are splited into lists and each rank only keep its shard in its optimizer - 2. Megatron DDP handle allreduce grad for all params, note that each rank have full model - and grad. + 2. Megatron DDP handle allreduce grad, note that each rank have full model and grad 3. optimizer is already modified so only param belong to this DP rank is updated - 3. grad_norm and zero counting will reduce metrics globally in step function - 4. Do regular update with chained optimizers, optimizer is already modified so partial update - happens. - 5. allgather updated params to every rank(currently through broadcast loop) + 4. grad_norm and zero counting will reduce metrics globally in step function + 5. Do regular update with chained optimizers, modified optimizer only update shard + 6. allgather updated params to every rank(currently through broadcast loop) """ def __init__( @@ -73,20 +71,17 @@ def __init__( super().__init__(optimizers) # TODO(kunlun, deyuf): potential future perf optimization - # since allreduce is unchanged and handled by megatron DDP, they're already in contiguous - # gbuf, so instead of shard param by layer randomly, we can still shard by buf range but - # keep some "extras" to keep boundary weight not sharded. This way each rank do some - # duplicated work but we can call single allgather later and all current distopt - # optimization can be applied. + # since allreduce is unchanged and handled by megatron DDP, they're already in + # contiguous gbuf. So instead of shard param by layer randomly, we can shard by + # buf range but keep some "extras" to keep boundary weight not sharded. + # This way each rank do some duplicated work but allgather_v is no longer needed + # All current distopt optimization can also be potentially applied def shard_params(self, optimizers): """Shard all params into lists by rank.""" - # We'll optimize sharding later if there is perf issue. should be ok since linear are - # grouped already. - # Key is to create separate sharding for dp/expt parallel, saved in dp_cp_params_list, - # expt_dp_params_list. - # Example of 4 dp rank and 10 non-expert parameters p0-p9, then dp_cp_params_list will - # look like: [[p0, p4, p8], [p1, p5, p9], [p2, p6], [p3, p7]] + # list of parameter are sorted by numel and assigned to ranks in ping-pong style + # example of 4 ranks and 10 parameters p0-p9 after sorting, then dp_cp_params_list will be + # [[p0, p7, p8], [p1, p6, p9], [p2, p5], [p3, p4]] # simplify when dp_cp group size is 1 if get_pg_size(self.pg_collection.dp_cp) == 1: @@ -97,52 +92,80 @@ def shard_params(self, optimizers): dp_cp_idx, expt_dp_idx = 0, 0 dp_cp_size = get_pg_size(self.pg_collection.dp_cp) expt_dp_size = get_pg_size(self.pg_collection.expt_dp) + # create ping-pong style loop so memory is more balanced + dp_cp_loop = list(range(dp_cp_size)) + list(range(dp_cp_size))[::-1] + expt_dp_loop = list(range(expt_dp_size)) + list(range(expt_dp_size))[::-1] self.dp_cp_params_list = [[] for _ in range(dp_cp_size)] self.expt_dp_params_list = [[] for _ in range(expt_dp_size)] - # get all param groups, this is called before init so cannot rely on - # Chained optimizer method + # get all param groups param_groups = [] for optimizer in optimizers: param_groups += optimizer.param_groups - for group in param_groups: - params_this_rank = [] - if group.get("is_expert_parallel", False): - for p in group["params"]: - if expt_dp_idx == get_pg_rank(self.pg_collection.expt_dp): - params_this_rank.append(p) - self.expt_dp_params_list[expt_dp_idx].append(p) - expt_dp_idx = (expt_dp_idx + 1) % expt_dp_size + + # sort param in all groups by param numel and assign to each rank evenly + param_list = [] + for group_index, group in enumerate(param_groups): + for p in group["params"]: + param_list.append((p, group_index)) + param_list.sort(key=lambda x: x[0].numel()) + param_groups_this_rank = [[] for g in param_groups] + + # assign params to rank in ping-pong style loop + for p, group_index in param_list: + if param_groups[group_index].get("is_expert_parallel", False): + if expt_dp_loop[expt_dp_idx] == get_pg_rank(self.pg_collection.expt_dp): + param_groups_this_rank[group_index].append(p) + self.expt_dp_params_list[expt_dp_loop[expt_dp_idx]].append(p) + expt_dp_idx = (expt_dp_idx + 1) % len(expt_dp_loop) else: - for p in group["params"]: - if dp_cp_idx == get_pg_rank(self.pg_collection.dp_cp): - params_this_rank.append(p) - self.dp_cp_params_list[dp_cp_idx].append(p) - dp_cp_idx = (dp_cp_idx + 1) % dp_cp_size - # now we modify the group to only handle local params - group["params"] = params_this_rank + if dp_cp_loop[dp_cp_idx] == get_pg_rank(self.pg_collection.dp_cp): + param_groups_this_rank[group_index].append(p) + self.dp_cp_params_list[dp_cp_loop[dp_cp_idx]].append(p) + dp_cp_idx = (dp_cp_idx + 1) % len(dp_cp_loop) + + # now we modify the group to only handle local params + for groups, params in zip(param_groups, param_groups_this_rank): + groups["params"] = params # simplify when expt_dp group size is 1 or expert parallel is off if expt_dp_size == 1 or len(self.expt_dp_params_list[0]) == 0: self.expt_dp_params_list = None @torch.no_grad() - def broadcast_params(self): - """All rank broadcast updated local params(allgatherv).""" - # Broadcast linear layer weights to all other ranks. - # This may not be slower than PyTorch allgatherv which calls broadcast internally. - # TODO(skyw): Profile and implement more efficient version. - if self.dp_cp_params_list is None: - return - for i, params in enumerate(self.dp_cp_params_list): - src_global_rank = torch.distributed.get_global_rank(self.pg_collection.dp_cp, i) - for p in params: - torch.distributed.broadcast(p, src_global_rank, self.pg_collection.dp_cp) - if self.expt_dp_params_list is None: - return - for i, params in enumerate(self.expt_dp_params_list): - src_global_rank = torch.distributed.get_global_rank(self.pg_collection.expt_dp, i) - for p in params: - torch.distributed.broadcast(p, src_global_rank, self.pg_collection.expt_dp) + def allgather_params(self) -> None: + """All-gather updated params from all ranks.""" + + # helper function to flatten local params, allgather, unflatten and copy to model params + def _allgather_helper(params_list, group): + # flatten this rank's params and create empty tensor output list + device = params_list[0][0].device + dtype = params_list[0][0].dtype + rank = get_pg_rank(group) + # for rank without params create empty tensor and participate in allgather + src = ( + _flatten_dense_tensors(params_list[rank]) + if len(params_list[rank]) > 0 + else torch.empty(0, device=device, dtype=dtype) + ) + output_list = [ + torch.empty(sum([p.numel() for p in params]), device=device, dtype=dtype) + for params in params_list + ] + # single all_gather_v to collect all updated params + torch.distributed.all_gather(output_list, src, group=group) + # unflatten and copy gathered params for each rank i + for idx, (flat_params, params) in enumerate(zip(output_list, params_list)): + # skip local params and empty tensors + if len(params) == 0 or idx == rank: + continue + updated_params = _unflatten_dense_tensors(flat_params, params) + for updated_p, model_p in zip(updated_params, params): + model_p.data.copy_(updated_p) + + if self.dp_cp_params_list: + _allgather_helper(self.dp_cp_params_list, self.pg_collection.dp_cp) + if self.expt_dp_params_list: + _allgather_helper(self.expt_dp_params_list, self.pg_collection.expt_dp) @torch.no_grad() def get_grad_norm(self): @@ -170,7 +193,7 @@ def step(self): # type: ignore[no-untyped-def] update_successful, grad_norm, num_zeros_in_grad = super().step() # All gather updated params. - self.broadcast_params() + self.allgather_params() return update_successful, grad_norm, num_zeros_in_grad From 9b7f095aa0d91eabd2c54ef586e0eadda2865093 Mon Sep 17 00:00:00 2001 From: Deyu Fu Date: Mon, 3 Nov 2025 20:40:27 +0800 Subject: [PATCH 02/10] add back broadcast as reference --- .../core/optimizer/layer_wise_optimizer.py | 17 ++++++++ tests/unit_tests/test_layer_wise_optimizer.py | 39 ++++++++++++++++++- 2 files changed, 55 insertions(+), 1 deletion(-) diff --git a/megatron/core/optimizer/layer_wise_optimizer.py b/megatron/core/optimizer/layer_wise_optimizer.py index 9e95cbebfa5..f3ecabbb8b7 100644 --- a/megatron/core/optimizer/layer_wise_optimizer.py +++ b/megatron/core/optimizer/layer_wise_optimizer.py @@ -167,6 +167,23 @@ def _allgather_helper(params_list, group): if self.expt_dp_params_list: _allgather_helper(self.expt_dp_params_list, self.pg_collection.expt_dp) + @torch.no_grad() + def broadcast_params(self): + """All rank broadcast updated local params.""" + # Broadcast linear layer weights to all other ranks. Kept as reference test. + if self.dp_cp_params_list is None: + return + for i, params in enumerate(self.dp_cp_params_list): + src_global_rank = torch.distributed.get_global_rank(self.pg_collection.dp_cp, i) + for p in params: + torch.distributed.broadcast(p, src_global_rank, self.pg_collection.dp_cp) + if self.expt_dp_params_list is None: + return + for i, params in enumerate(self.expt_dp_params_list): + src_global_rank = torch.distributed.get_global_rank(self.pg_collection.expt_dp, i) + for p in params: + torch.distributed.broadcast(p, src_global_rank, self.pg_collection.expt_dp) + @torch.no_grad() def get_grad_norm(self): # similar to dist opt, always aggregate globally diff --git a/tests/unit_tests/test_layer_wise_optimizer.py b/tests/unit_tests/test_layer_wise_optimizer.py index c9dd542cf25..301af44cd03 100644 --- a/tests/unit_tests/test_layer_wise_optimizer.py +++ b/tests/unit_tests/test_layer_wise_optimizer.py @@ -1,4 +1,5 @@ -# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved. +# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + import os import pytest @@ -401,3 +402,39 @@ def test_parameter_updates_insufficient_parameters(self): This will be insufficient when world size > 2. """ self._run_parameter_update_test(model_class=TinyModel) + + def test_broadcast_vs_allgather(self): + """Test LayerWiseDistributedOptimizer allgather code agains broadcast code.""" + model, optimizer, pg_collection = self.create_model_and_optimizer(model_class=SimpleModel) + + # Create reference model and optimizer using the same function + reference_model, reference_optimizer, _ = self.create_model_and_optimizer( + model_class=SimpleModel, copy_from=model + ) + + # Set same gradients on both models + for param, ref_param in zip(model.parameters(), reference_model.parameters()): + assert torch.equal(param.data, ref_param.data) + torch.testing.assert_close(param.data, ref_param.data, rtol=1e-5, atol=1e-5) + grad_value = torch.randn_like(param) + torch.distributed.broadcast(grad_value, src=0, group=pg_collection.dp_cp) + param.main_grad = grad_value.clone().detach() + ref_param.main_grad = grad_value.clone().detach() + + optimizer.step() + + # Verify at least some parameters were updated + params_updated = 0 + for param, ref_param in zip(model.parameters(), reference_model.parameters()): + if not torch.equal(param.data, ref_param.data): + params_updated += 1 + + assert params_updated > 0, "At least some parameters should be updated" + + # step() internal call allgather_params. replace reference object with bcast + reference_optimizer.allgather_params = reference_optimizer.broadcast_params + reference_optimizer.step() + + # Verify updated values match reference optimizer + for param, ref_param in zip(model.parameters(), reference_model.parameters()): + torch.testing.assert_close(param.data, ref_param.data, rtol=1e-5, atol=1e-5) From 70333ffd0b7387cf8c137be5f97da6bc9252f18b Mon Sep 17 00:00:00 2001 From: Deyu Fu Date: Tue, 4 Nov 2025 05:35:04 -0800 Subject: [PATCH 03/10] fix unnecessary FP32 megatron optimizer wrap and undetected ShardedTensorFactory replica_id change --- .../core/optimizer/layer_wise_optimizer.py | 31 +++++++++++-------- 1 file changed, 18 insertions(+), 13 deletions(-) diff --git a/megatron/core/optimizer/layer_wise_optimizer.py b/megatron/core/optimizer/layer_wise_optimizer.py index f3ecabbb8b7..785dfd78305 100644 --- a/megatron/core/optimizer/layer_wise_optimizer.py +++ b/megatron/core/optimizer/layer_wise_optimizer.py @@ -5,14 +5,13 @@ import torch from torch._utils import _flatten_dense_tensors, _unflatten_dense_tensors -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 from .clip_grads import count_zeros_fp32, get_grad_norm_fp32 -from .optimizer import ChainedOptimizer, Float16OptimizerWithFloat16Params, MegatronOptimizer +from .optimizer import ChainedOptimizer, FP32Optimizer, Float16OptimizerWithFloat16Params, MegatronOptimizer from .optimizer_config import OptimizerConfig @@ -51,9 +50,6 @@ def __init__( 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: @@ -61,13 +57,20 @@ def __init__( "init_state_fn_list must be the " "same length as optimizers if provided" ) + # wrap optimizer after sharding to avoid unnecessary master weight creation + # for higher precision, optimizers are wrapped with megatron already if config.bf16: if isinstance(optimizers[0], Float16OptimizerWithFloat16Params): raise TypeError('LayerWiseDistributedOptimizer received Float16 optimizer already.') - optimizers = [ - Float16OptimizerWithFloat16Params(optim, config, None, init_state_fn_list[idx]) - for idx, optim in enumerate(optimizers) - ] + # unwrap FP32 optimizer, possibly from reusing get_megatron_optimizer for adam + for i, opt in enumerate(optimizers): + if isinstance(opt, Float16OptimizerWithFloat16Params): + raise TypeError('LayerWiseDistributedOptimizer received Float16 optimizer already.') + # unwrap FP32 optimizer from reusing get_megatron_optimizer for adam + if isinstance(opt, FP32Optimizer): + opt = opt.optimizer + optimizers[i] = Float16OptimizerWithFloat16Params(opt, config, None, init_state_fn_list[i]) + super().__init__(optimizers) # TODO(kunlun, deyuf): potential future perf optimization @@ -227,10 +230,12 @@ def sharded_state_dict( # for fixed DP usage only for sh_base in nested_values(sharded_state_dict): - if isinstance(sh_base, ShardedTensor): + if hasattr(sh_base, 'replica_id'): 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) + isinstance(sh_base.replica_id, int) or len(sh_base.replica_id) == 3 + ), f'Expected replica_id as int or (PP, TP, DP), got: {sh_base}' + sh_base.replica_id = ( + 0 if isinstance(sh_base.replica_id, int) else (*sh_base.replica_id[:2], 0) + ) return sharded_state_dict From 8cb4004b77a3cdd557c01327e29c3d47f781ce24 Mon Sep 17 00:00:00 2001 From: Deyu Fu Date: Tue, 4 Nov 2025 05:43:19 -0800 Subject: [PATCH 04/10] lint --- megatron/core/optimizer/layer_wise_optimizer.py | 15 ++++++++++++--- 1 file changed, 12 insertions(+), 3 deletions(-) diff --git a/megatron/core/optimizer/layer_wise_optimizer.py b/megatron/core/optimizer/layer_wise_optimizer.py index 785dfd78305..16ff8f525c7 100644 --- a/megatron/core/optimizer/layer_wise_optimizer.py +++ b/megatron/core/optimizer/layer_wise_optimizer.py @@ -11,7 +11,12 @@ from megatron.core.utils import get_pg_rank, get_pg_size from .clip_grads import count_zeros_fp32, get_grad_norm_fp32 -from .optimizer import ChainedOptimizer, FP32Optimizer, Float16OptimizerWithFloat16Params, MegatronOptimizer +from .optimizer import ( + ChainedOptimizer, + FP32Optimizer, + Float16OptimizerWithFloat16Params, + MegatronOptimizer, +) from .optimizer_config import OptimizerConfig @@ -65,11 +70,15 @@ def __init__( # unwrap FP32 optimizer, possibly from reusing get_megatron_optimizer for adam for i, opt in enumerate(optimizers): if isinstance(opt, Float16OptimizerWithFloat16Params): - raise TypeError('LayerWiseDistributedOptimizer received Float16 optimizer already.') + raise TypeError( + 'LayerWiseDistributedOptimizer received Float16 optimizer already.' + ) # unwrap FP32 optimizer from reusing get_megatron_optimizer for adam if isinstance(opt, FP32Optimizer): opt = opt.optimizer - optimizers[i] = Float16OptimizerWithFloat16Params(opt, config, None, init_state_fn_list[i]) + optimizers[i] = Float16OptimizerWithFloat16Params( + opt, config, None, init_state_fn_list[i] + ) super().__init__(optimizers) From 9fbb48e4f7bb1a93f6c1d2e538ccbe4f787039ed Mon Sep 17 00:00:00 2001 From: Deyu Fu Date: Tue, 4 Nov 2025 06:13:43 -0800 Subject: [PATCH 05/10] lint --- megatron/core/optimizer/layer_wise_optimizer.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/megatron/core/optimizer/layer_wise_optimizer.py b/megatron/core/optimizer/layer_wise_optimizer.py index 16ff8f525c7..06fc4b40994 100644 --- a/megatron/core/optimizer/layer_wise_optimizer.py +++ b/megatron/core/optimizer/layer_wise_optimizer.py @@ -13,8 +13,8 @@ from .clip_grads import count_zeros_fp32, get_grad_norm_fp32 from .optimizer import ( ChainedOptimizer, - FP32Optimizer, Float16OptimizerWithFloat16Params, + FP32Optimizer, MegatronOptimizer, ) from .optimizer_config import OptimizerConfig @@ -55,9 +55,7 @@ def __init__( self.pg_collection = pg_collection self.shard_params(optimizers) - if init_state_fn_list is None: - init_state_fn_list = [None] * len(optimizers) - else: + if init_state_fn_list: assert len(init_state_fn_list) == len(optimizers), ( "init_state_fn_list must be the " "same length as optimizers if provided" ) @@ -77,7 +75,7 @@ def __init__( if isinstance(opt, FP32Optimizer): opt = opt.optimizer optimizers[i] = Float16OptimizerWithFloat16Params( - opt, config, None, init_state_fn_list[i] + opt, config, None, init_state_fn_list[i] if init_state_fn_list else None ) super().__init__(optimizers) @@ -174,6 +172,8 @@ def _allgather_helper(params_list, group): for updated_p, model_p in zip(updated_params, params): model_p.data.copy_(updated_p) + if self.pg_collection is None: + return if self.dp_cp_params_list: _allgather_helper(self.dp_cp_params_list, self.pg_collection.dp_cp) if self.expt_dp_params_list: From 15fc11e1a1ab00ce41f6a10db5e8d1c0b66a5640 Mon Sep 17 00:00:00 2001 From: Deyu Fu Date: Tue, 4 Nov 2025 22:45:27 -0800 Subject: [PATCH 06/10] fix forgot to remove check --- megatron/core/optimizer/layer_wise_optimizer.py | 2 -- 1 file changed, 2 deletions(-) diff --git a/megatron/core/optimizer/layer_wise_optimizer.py b/megatron/core/optimizer/layer_wise_optimizer.py index 06fc4b40994..2e1e1824847 100644 --- a/megatron/core/optimizer/layer_wise_optimizer.py +++ b/megatron/core/optimizer/layer_wise_optimizer.py @@ -63,8 +63,6 @@ def __init__( # wrap optimizer after sharding to avoid unnecessary master weight creation # for higher precision, optimizers are wrapped with megatron already if config.bf16: - if isinstance(optimizers[0], Float16OptimizerWithFloat16Params): - raise TypeError('LayerWiseDistributedOptimizer received Float16 optimizer already.') # unwrap FP32 optimizer, possibly from reusing get_megatron_optimizer for adam for i, opt in enumerate(optimizers): if isinstance(opt, Float16OptimizerWithFloat16Params): From 9f43c2f8ba359e9ad815e01d4a843546d9d451fb Mon Sep 17 00:00:00 2001 From: Deyu Fu Date: Wed, 5 Nov 2025 00:09:06 -0800 Subject: [PATCH 07/10] fix empty group breaks common_dict after sharding --- megatron/core/optimizer/layer_wise_optimizer.py | 16 ++++++++++++++++ 1 file changed, 16 insertions(+) diff --git a/megatron/core/optimizer/layer_wise_optimizer.py b/megatron/core/optimizer/layer_wise_optimizer.py index 2e1e1824847..7b9aaa76fd3 100644 --- a/megatron/core/optimizer/layer_wise_optimizer.py +++ b/megatron/core/optimizer/layer_wise_optimizer.py @@ -245,4 +245,20 @@ def sharded_state_dict( 0 if isinstance(sh_base.replica_id, int) else (*sh_base.replica_id[:2], 0) ) + # Adjust dict due to possible empty rank 0 which output common_dict + for sd in sharded_state_dict.values(): + # Drop empty group state to avoid save in common dict (non-empty rank still save) + sd['fp32_from_fp16_params'][:] = [ + group for group in sd['fp32_from_fp16_params'] if group + ] + # TODO(deyuf): 'common_step' code path is broken and 'step' is saved in 'param_groups' + # Find next 'step' if present. note this still break if rank0 adam is fully empty + step = next( + (group['step'] for group in sd['optimizer']['param_groups'] if 'step' in group), + None, + ) + if step is not None: + for group in sd['optimizer']['param_groups']: + group['step'] = step + return sharded_state_dict From 325366cdaad663c56b39b4d60f2d399c15214c82 Mon Sep 17 00:00:00 2001 From: Deyu Fu Date: Wed, 5 Nov 2025 20:53:05 +0800 Subject: [PATCH 08/10] address comments --- megatron/core/optimizer/layer_wise_optimizer.py | 11 ++++++----- tests/unit_tests/test_layer_wise_optimizer.py | 4 ++-- 2 files changed, 8 insertions(+), 7 deletions(-) diff --git a/megatron/core/optimizer/layer_wise_optimizer.py b/megatron/core/optimizer/layer_wise_optimizer.py index 7b9aaa76fd3..3965ce091c7 100644 --- a/megatron/core/optimizer/layer_wise_optimizer.py +++ b/megatron/core/optimizer/layer_wise_optimizer.py @@ -33,7 +33,7 @@ class LayerWiseDistributedOptimizer(ChainedOptimizer): 3. optimizer is already modified so only param belong to this DP rank is updated 4. grad_norm and zero counting will reduce metrics globally in step function 5. Do regular update with chained optimizers, modified optimizer only update shard - 6. allgather updated params to every rank(currently through broadcast loop) + 6. allgather updated params to every rank """ def __init__( @@ -56,15 +56,16 @@ def __init__( self.pg_collection = pg_collection self.shard_params(optimizers) if init_state_fn_list: - assert len(init_state_fn_list) == len(optimizers), ( - "init_state_fn_list must be the " "same length as optimizers if provided" - ) + assert len(init_state_fn_list) == len( + optimizers + ), "init_state_fn_list must be the same length as optimizers if provided" # wrap optimizer after sharding to avoid unnecessary master weight creation # for higher precision, optimizers are wrapped with megatron already if config.bf16: # unwrap FP32 optimizer, possibly from reusing get_megatron_optimizer for adam - for i, opt in enumerate(optimizers): + for i in range(len(optimizers)): + opt = optimizers[i] if isinstance(opt, Float16OptimizerWithFloat16Params): raise TypeError( 'LayerWiseDistributedOptimizer received Float16 optimizer already.' diff --git a/tests/unit_tests/test_layer_wise_optimizer.py b/tests/unit_tests/test_layer_wise_optimizer.py index 301af44cd03..05ce26bcfa0 100644 --- a/tests/unit_tests/test_layer_wise_optimizer.py +++ b/tests/unit_tests/test_layer_wise_optimizer.py @@ -415,7 +415,7 @@ def test_broadcast_vs_allgather(self): # Set same gradients on both models for param, ref_param in zip(model.parameters(), reference_model.parameters()): assert torch.equal(param.data, ref_param.data) - torch.testing.assert_close(param.data, ref_param.data, rtol=1e-5, atol=1e-5) + torch.testing.assert_close(param.data, ref_param.data, rtol=0, atol=0) grad_value = torch.randn_like(param) torch.distributed.broadcast(grad_value, src=0, group=pg_collection.dp_cp) param.main_grad = grad_value.clone().detach() @@ -437,4 +437,4 @@ def test_broadcast_vs_allgather(self): # Verify updated values match reference optimizer for param, ref_param in zip(model.parameters(), reference_model.parameters()): - torch.testing.assert_close(param.data, ref_param.data, rtol=1e-5, atol=1e-5) + torch.testing.assert_close(param.data, ref_param.data, rtol=0, atol=0) From 2c01a2d47dcccaf8901f531a6339c9bdd9b507bc Mon Sep 17 00:00:00 2001 From: Deyu Fu Date: Wed, 5 Nov 2025 21:22:12 +0800 Subject: [PATCH 09/10] fix --- megatron/core/optimizer/layer_wise_optimizer.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/megatron/core/optimizer/layer_wise_optimizer.py b/megatron/core/optimizer/layer_wise_optimizer.py index 3965ce091c7..1d75e488714 100644 --- a/megatron/core/optimizer/layer_wise_optimizer.py +++ b/megatron/core/optimizer/layer_wise_optimizer.py @@ -249,9 +249,10 @@ def sharded_state_dict( # Adjust dict due to possible empty rank 0 which output common_dict for sd in sharded_state_dict.values(): # Drop empty group state to avoid save in common dict (non-empty rank still save) - sd['fp32_from_fp16_params'][:] = [ - group for group in sd['fp32_from_fp16_params'] if group - ] + if 'fp32_from_fp16_params' in sd: + sd['fp32_from_fp16_params'][:] = [ + group for group in sd['fp32_from_fp16_params'] if group + ] # TODO(deyuf): 'common_step' code path is broken and 'step' is saved in 'param_groups' # Find next 'step' if present. note this still break if rank0 adam is fully empty step = next( From ac643811e487af482304406f6260b2cea07750b3 Mon Sep 17 00:00:00 2001 From: Deyu Fu Date: Wed, 5 Nov 2025 08:36:49 -0800 Subject: [PATCH 10/10] fix when there is only 1 optimizer --- megatron/core/optimizer/layer_wise_optimizer.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/megatron/core/optimizer/layer_wise_optimizer.py b/megatron/core/optimizer/layer_wise_optimizer.py index 1d75e488714..64eac03d626 100644 --- a/megatron/core/optimizer/layer_wise_optimizer.py +++ b/megatron/core/optimizer/layer_wise_optimizer.py @@ -246,8 +246,12 @@ def sharded_state_dict( 0 if isinstance(sh_base.replica_id, int) else (*sh_base.replica_id[:2], 0) ) + if len(self.chained_optimizers) == 1: + wrapped_sharded_state_dict = {1: sharded_state_dict} + else: + wrapped_sharded_state_dict = sharded_state_dict # Adjust dict due to possible empty rank 0 which output common_dict - for sd in sharded_state_dict.values(): + for sd in wrapped_sharded_state_dict.values(): # Drop empty group state to avoid save in common dict (non-empty rank still save) if 'fp32_from_fp16_params' in sd: sd['fp32_from_fp16_params'][:] = [