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
24 changes: 24 additions & 0 deletions megatron/core/optimizer_param_scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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)
Expand Down
21 changes: 9 additions & 12 deletions megatron/training/training.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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,
Expand Down
3 changes: 2 additions & 1 deletion megatron/training/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand All @@ -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()
Expand Down
101 changes: 100 additions & 1 deletion tests/unit_tests/test_optimizer_param_scheduler.py
Original file line number Diff line number Diff line change
@@ -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,
)


Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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
Loading