diff --git a/megatron/core/optimizer_param_scheduler.py b/megatron/core/optimizer_param_scheduler.py index 7ff6fee35a7..e01a708ce79 100644 --- a/megatron/core/optimizer_param_scheduler.py +++ b/megatron/core/optimizer_param_scheduler.py @@ -34,6 +34,26 @@ class ParamGroupOverride(TypedDict): wd_mult: float +def get_canonical_lr_for_logging(param_groups: list[dict]) -> float | None: + """Return the lr of the first ``default_config=True`` param group. + + All ``default_config`` groups share the same LR schedule, so the first one + is representative. This includes empty rank-alignment stub groups, which + the scheduler still writes a valid lr onto. + + Args: + param_groups (list[dict]): parameter groups from the optimizer. + + Returns: + float | None: the canonical learning rate, or None if no + ``default_config=True`` group is found. + """ + for param_group in param_groups: + if param_group.get('default_config', False): + return param_group.get('lr') + return None + + def param_group_override_to_tuple( param_group_override: ParamGroupOverride | None, ) -> tuple[tuple[str, Any], ...] | None: @@ -265,6 +285,10 @@ def step(self, increment: int) -> None: increment (int): number of steps to increment """ self.num_steps += increment + # Do not skip empty param groups: get_canonical_lr_for_logging reads lr + # from default_config groups regardless of whether they hold parameters. + # This is important for logging under model parallelism that may leave + # some ranks with empty default_config parameter groups. for param_group in self.optimizer.param_groups: param_group['lr'] = self.get_lr(param_group) param_group['weight_decay'] = self.get_wd(param_group) * param_group.get('wd_mult', 1.0) diff --git a/megatron/training/training.py b/megatron/training/training.py index 378dcfb3593..2c68c70735d 100644 --- a/megatron/training/training.py +++ b/megatron/training/training.py @@ -49,6 +49,7 @@ def set_startup_timestamps(program_start=None, main_entry=None): import torch.distributed from megatron.core.optimizer.distrib_optimizer import DistributedOptimizer +from megatron.core.optimizer_param_scheduler import get_canonical_lr_for_logging from .log_handler import CustomHandler # Make default logging level INFO, but filter out all log messages not from MCore. @@ -1838,7 +1839,7 @@ def train_step(forward_step_func, data_iterator, model, optimizer, opt_param_sch def training_log( loss_dict, total_loss_dict, - learning_rate, + learning_rate: float | None, iteration, loss_scale, report_memory_flag, @@ -1939,15 +1940,16 @@ def training_log( total_iterations = total_loss_dict[advanced_iters_key] + total_loss_dict[skipped_iters_key] # learning rate will be None on ranks without trainable params, so we must gather across mp ranks - learning_rate = reduce_max_stat_across_model_parallel_group(learning_rate) + learning_rate: float | None = reduce_max_stat_across_model_parallel_group(learning_rate) # Tensorboard values. if writer and (iteration % args.tensorboard_log_interval == 0): if wandb_writer: wandb_writer.log({'samples vs steps': args.consumed_train_samples}, iteration) - writer.add_scalar('learning-rate', learning_rate, iteration) - writer.add_scalar('learning-rate vs samples', learning_rate, args.consumed_train_samples) - if wandb_writer: - wandb_writer.log({'learning-rate': learning_rate}, iteration) + if learning_rate is not None: + writer.add_scalar('learning-rate', learning_rate, iteration) + writer.add_scalar('learning-rate vs samples', learning_rate, args.consumed_train_samples) + if wandb_writer: + wandb_writer.log({'learning-rate': learning_rate}, iteration) if args.skipped_train_samples > 0: writer.add_scalar('skipped-train-samples', args.skipped_train_samples, iteration) if wandb_writer: @@ -2973,12 +2975,7 @@ def trace_handler(p): if args.log_params_norm: params_norm = calc_params_l2_norm(model) - learning_rate = None - for param_group in optimizer.param_groups: - if len(param_group['params']) == 0: - continue - if param_group['default_config']: - learning_rate = param_group['lr'] + learning_rate = get_canonical_lr_for_logging(optimizer.param_groups) report_memory_flag = training_log( loss_dict, total_loss_dict, diff --git a/megatron/training/utils.py b/megatron/training/utils.py index 8af1a44f1fb..edd50dc831f 100644 --- a/megatron/training/utils.py +++ b/megatron/training/utils.py @@ -240,7 +240,7 @@ def average_losses_across_data_parallel_group(losses): return averaged_losses -def reduce_max_stat_across_model_parallel_group(stat: float) -> float: +def reduce_max_stat_across_model_parallel_group(stat: float) -> float | None: """ Ranks without an optimizer will have no grad_norm or num_zeros_in_grad stats. We need to ensure the logging and writer rank has those values. @@ -255,6 +255,7 @@ def reduce_max_stat_across_model_parallel_group(stat: float) -> float: stat, op=torch.distributed.ReduceOp.MAX, group=mpu.get_model_parallel_group() ) if stat.item() == -1.0: + # No rank has a valid stat, so return None to indicate that it is None across all ranks. return None else: return stat.item() diff --git a/tests/unit_tests/test_optimizer_param_scheduler.py b/tests/unit_tests/test_optimizer_param_scheduler.py index 9b781694546..670ca92c992 100644 --- a/tests/unit_tests/test_optimizer_param_scheduler.py +++ b/tests/unit_tests/test_optimizer_param_scheduler.py @@ -1,10 +1,13 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + import math from unittest.mock import MagicMock import pytest -from megatron.core.optimizer_param_scheduler import ( # Adjust import according to your module path +from megatron.core.optimizer_param_scheduler import ( OptimizerParamScheduler, + get_canonical_lr_for_logging, ) @@ -182,6 +185,51 @@ def test_step_function(mock_optimizer): assert math.isclose(param_group['weight_decay'], 0.01, rel_tol=1e-5) +def test_step_updates_empty_param_groups(): + """Empty param groups (rank-alignment stubs) must still receive lr updates. + + get_canonical_lr_for_logging reads lr from default_config groups regardless + of whether they hold parameters, so step() must not skip them. + """ + optimizer = MagicMock() + # lr and weight_decay are set by the scheduler's step() method + optimizer.param_groups = [ + # Non-default group with its own max_lr override (lr will differ from the canonical schedule) + {'params': [1, 2], "min_lr": 0.001, "max_lr": 0.2, "default_config": False}, + # Model parallelism may leave default_config groups empty on some ranks + {'params': [], "wd_mult": 0.0, 'default_config': True}, + ] + scheduler = OptimizerParamScheduler( + optimizer=optimizer, + init_lr=0.01, + max_lr=0.1, + min_lr=0.001, + lr_warmup_steps=100, + lr_decay_steps=1000, + lr_decay_style='linear', + start_wd=0.0, + end_wd=0.1, + wd_incr_steps=1000, + wd_incr_style='linear', + ) + + scheduler.step(100) + non_empty, empty = optimizer.param_groups + + # Verify learning rates: at step 100 warmup is complete so lr == max_lr + assert "lr" in non_empty, "non-empty param group must have an lr" + assert "lr" in empty, "empty param group must have an lr" + assert non_empty['lr'] == pytest.approx(0.2) # warmup complete → this group's max_lr override + assert empty['lr'] == pytest.approx(0.1) # warmup complete → scheduler's default max_lr + assert get_canonical_lr_for_logging(optimizer.param_groups) == pytest.approx(0.1) + + # Verify weight decay: linear from 0.0 to 0.1 over 1000 steps → base wd is 0.01 at step 100 + assert "weight_decay" in non_empty, "non-empty param group must have a weight decay" + assert "weight_decay" in empty, "empty param group must have a weight decay" + assert non_empty['weight_decay'] == pytest.approx(0.01) # base wd, no wd_mult override + assert empty['weight_decay'] == pytest.approx(0.0) # base wd * wd_mult=0.0 + + def test_state_dict(mock_optimizer): scheduler = OptimizerParamScheduler( optimizer=mock_optimizer, @@ -249,3 +297,54 @@ def test_load_state_dict(mock_optimizer): assert scheduler.end_wd == 0.2 assert scheduler.wd_incr_steps == 500 assert scheduler.wd_incr_style == 'cosine' + + +# ── get_canonical_lr_for_logging tests ────────────────────────────────────── +# +# Returns the lr of the first default_config=True param group. In practice +# the scheduler always sets a valid lr on every group (including empty +# rank-alignment stubs), so a default_config=True group with a float lr is +# always present. + + +class TestGetCanonicalLrForLogging: + """Tests for get_canonical_lr_for_logging.""" + + def test_single_default_config_group(self): + """Typical case: one default_config group with a valid lr.""" + param_groups = [{'lr': 0.05, 'default_config': True}] + assert get_canonical_lr_for_logging(param_groups) == 0.05 + + def test_default_config_with_non_default_groups(self): + """default_config group is returned even when non-default groups are present.""" + param_groups = [{'lr': 0.001, 'default_config': True}, {'lr': 0.999}] + assert get_canonical_lr_for_logging(param_groups) == 0.001 + + def test_default_config_after_non_default(self): + """default_config group is found even when it is not first in the list.""" + param_groups = [{'lr': 0.50}, {'lr': 0.01, 'default_config': True}] + assert get_canonical_lr_for_logging(param_groups) == 0.01 + + def test_no_default_config_groups(self): + """Returns None when no group has default_config=True.""" + param_groups = [{'lr': 0.50}, {'lr': 0.01}] + assert get_canonical_lr_for_logging(param_groups) is None + + def test_missing_lr_key(self): + """Returns None (not KeyError) when the default_config group lacks an 'lr' key.""" + param_groups = [{'default_config': True}] + assert get_canonical_lr_for_logging(param_groups) is None + + def test_empty_param_groups(self): + """Returns None when there are no param groups at all.""" + assert get_canonical_lr_for_logging([]) is None + + def test_no_default_config_no_lr(self): + """Returns None when groups exist but none are default_config.""" + param_groups = [{'params': []}] + assert get_canonical_lr_for_logging(param_groups) is None + + def test_lr_zero_is_valid(self): + """lr=0.0 is a legitimate value, not to be confused with None.""" + param_groups = [{'lr': 0.0, 'default_config': True}] + assert get_canonical_lr_for_logging(param_groups) == 0.0