Skip to content
Closed
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
9 changes: 5 additions & 4 deletions megatron/core/optimizer/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -686,11 +686,12 @@ def init_state_fn(opt, config=None):
setattr(optimizer, 'grad_stats_parallel_group', model_parallel_group)

if pg_collection is None or not hasattr(pg_collection, 'tp'):
tp_group = parallel_state.get_tensor_model_parallel_group()
else:
tp_group = pg_collection.tp
# TODO(M4): plumb tp_group through optimizer constructors so this setattr disappears.
pg_collection = ProcessGroupCollection.use_mpu_process_groups()
tp_group = pg_collection.tp
expert_tp_group = getattr(pg_collection, 'expt_tp', tp_group)
# TODO(M4): plumb TP groups through optimizer constructors so these setattrs disappear.
setattr(optimizer, 'tp_group', tp_group)
setattr(optimizer, 'expert_tp_group', expert_tp_group)

return optimizer

Expand Down
5 changes: 4 additions & 1 deletion megatron/core/optimizer/clip_grads.py
Original file line number Diff line number Diff line change
Expand Up @@ -201,6 +201,7 @@ def count_zeros_fp32(
grad_stats_parallel_group: torch.distributed.ProcessGroup,
use_decoupled_grad: bool = False,
tp_group: Optional[torch.distributed.ProcessGroup] = None,
expert_tp_group: Optional[torch.distributed.ProcessGroup] = None,
) -> float:
"""Counts the number of zero values in the gradients of the given parameters.

Expand Down Expand Up @@ -245,7 +246,9 @@ def count_zeros_fp32(
total_num_zeros += num_zeros
continue
is_not_shared = param_is_not_shared(param)
is_not_tp_duplicate = param_is_not_tensor_parallel_duplicate(param, tp_group=tp_group)
is_not_tp_duplicate = param_is_not_tensor_parallel_duplicate(
param, tp_group=tp_group, expert_tp_group=expert_tp_group
)
if grad_not_none and is_not_shared and is_not_tp_duplicate:
grad_obj = getattr(param, grad_attr)
data_parallel_group = get_data_parallel_group_if_dtensor(grad_obj, data_parallel_group)
Expand Down
9 changes: 8 additions & 1 deletion megatron/core/optimizer/optimizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -124,6 +124,8 @@ def _is_separate_grad_norm_group(grad_norm_group: Optional[str]) -> bool:

def copy_optimizer_param_metadata(destination: torch.Tensor, source: torch.Tensor) -> None:
"""Copy optimizer-relevant metadata when creating param views/copies."""
if hasattr(source, 'allreduce'):
destination.allreduce = source.allreduce
if hasattr(source, 'shared'):
destination.shared = source.shared
if hasattr(source, GRAD_NORM_GROUP_ATTR):
Expand Down Expand Up @@ -213,7 +215,9 @@ def _filter_grads_for_norm(
grad_not_none = grad is not None
is_not_shared = param_is_not_shared(param)
is_not_tp_duplicate = tensor_parallel.param_is_not_tensor_parallel_duplicate(
param, getattr(self, 'tp_group', None)
param,
tp_group=getattr(self, 'tp_group', None),
expert_tp_group=getattr(self, 'expert_tp_group', None),
)
if grad_not_none and is_not_shared and is_not_tp_duplicate:
grads_for_norm.append(grad)
Expand Down Expand Up @@ -380,6 +384,7 @@ def count_zeros(self) -> float:
and getattr(params[0], "__fsdp_param__", False)
),
tp_group=getattr(self, 'tp_group', None),
expert_tp_group=getattr(self, 'expert_tp_group', None),
)

@abstractmethod
Expand Down Expand Up @@ -1695,6 +1700,8 @@ def count_zeros(self):
self.config.use_precision_aware_optimizer
and getattr(params[0], "__fsdp_param__", False)
),
tp_group=getattr(self.chained_optimizers[0], 'tp_group', None),
expert_tp_group=getattr(self.chained_optimizers[0], 'expert_tp_group', None),
)
else:
num_zeros_in_grad = 0
Expand Down
78 changes: 78 additions & 0 deletions tests/unit_tests/test_optimizer.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.

import math
import os
from unittest.mock import patch

Expand All @@ -12,6 +13,7 @@
# FP8 recipe will be used to test precision-aware-optimizer.
from transformer_engine.pytorch.fp8 import fp8_autocast

from megatron.core import parallel_state
from megatron.core.distributed import DistributedDataParallel, DistributedDataParallelConfig
from megatron.core.optimizer import (
ChainedOptimizer,
Expand All @@ -23,6 +25,7 @@
get_megatron_optimizer,
get_standard_config_overrides,
)
from megatron.core.optimizer.optimizer import copy_optimizer_param_metadata
from megatron.core.optimizer_param_scheduler import ParamGroupOverride
from megatron.core.process_groups_config import ProcessGroupCollection
from megatron.core.transformer import TransformerConfig
Expand Down Expand Up @@ -69,6 +72,81 @@ def forward(self, x):
return x


def test_copy_optimizer_param_metadata_preserves_allreduce():
source = torch.empty(1)
destination = torch.empty_like(source)
source.allreduce = False

copy_optimizer_param_metadata(destination, source)

assert destination.allreduce is False


@pytest.mark.skipif(
int(os.getenv('WORLD_SIZE', '1')) < 2, reason="test requires at least two distributed ranks"
)
@pytest.mark.parametrize("use_distributed_optimizer", (False, True), ids=("optimizer", "distopt"))
def test_expert_grad_stats_use_expert_tp_group(use_distributed_optimizer: bool):
"""Expert grad stats must deduplicate over ETP, not dense TP."""
world_size = int(os.environ['WORLD_SIZE'])
rank = int(os.environ['RANK'])
_init_distributed(world_size, rank)

class DenseAndExpertParameters(nn.Module):
def __init__(self):
super().__init__()
self.dense = nn.Parameter(torch.ones(1, dtype=torch.bfloat16, device='cuda'))
self.expert = nn.Parameter(torch.ones(1, dtype=torch.bfloat16, device='cuda'))
self.expert.allreduce = False

try:
Utils.initialize_model_parallel(
tensor_model_parallel_size=world_size,
expert_model_parallel_size=world_size,
expert_tensor_parallel_size=1,
)
model = DenseAndExpertParameters()
transformer_config = TransformerConfig(num_attention_heads=1, num_layers=1)
model = DistributedDataParallel(
transformer_config,
DistributedDataParallelConfig(use_distributed_optimizer=use_distributed_optimizer),
model,
)
optimizer = get_megatron_optimizer(
OptimizerConfig(
optimizer='adam',
lr=0.0,
bf16=True,
clip_grad=0.0,
log_num_zeros_in_grad=True,
use_distributed_optimizer=use_distributed_optimizer,
),
[model],
use_gloo_process_groups=False,
)

dense_optimizer, expert_optimizer = optimizer.chained_optimizers
assert dense_optimizer.tp_group is parallel_state.get_tensor_model_parallel_group()
assert expert_optimizer.expert_tp_group is parallel_state.get_expert_tensor_parallel_group()

for param in model.parameters():
param.main_grad.fill_(1.0)
assert optimizer.prepare_grads() is False

expected_grad_norm = math.sqrt(1 + world_size)
actual_grad_norm = optimizer.get_grad_norm()
if isinstance(actual_grad_norm, torch.Tensor):
actual_grad_norm = actual_grad_norm.item()
assert actual_grad_norm == pytest.approx(expected_grad_norm)

for param in model.parameters():
param.main_grad.zero_()
assert optimizer.prepare_grads() is False
assert optimizer.count_zeros() == 1 + world_size
finally:
Utils.destroy_model_parallel()


@patch('torch.distributed.get_world_size', return_value=1)
@patch(
'torch.distributed.all_gather_object', lambda output_list, obj: output_list.__setitem__(0, obj)
Expand Down
Loading