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
11 changes: 9 additions & 2 deletions litgpt/args.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,8 +16,9 @@ class TrainArgs:
"""Number of samples between optimizer steps across data-parallel ranks"""
micro_batch_size: int = 4
"""Number of samples per data-parallel rank"""
lr_warmup_steps: int | None = 100
"""Number of iterations with learning rate warmup active"""
lr_warmup_steps: int | None = None
"""Number of iterations with learning rate warmup active. Defaults to 100 when
``lr_warmup_fraction`` is not set; ignored when ``lr_warmup_fraction`` is set."""
lr_warmup_fraction: float | None = None
"""The fraction of an epoch to use for learning rate warmup"""
epochs: int | None = None
Expand Down Expand Up @@ -46,6 +47,12 @@ def __post_init__(self) -> None:
if self.lr_warmup_fraction and not (0 <= self.lr_warmup_fraction <= 1):
raise ValueError("`--train.lr_warmup_fraction` must be between 0 and 1.")

# Apply the documented default for ``lr_warmup_steps`` only when the user did not
# provide ``lr_warmup_fraction``. This keeps a sensible default (100 steps) without
# conflicting with users who explicitly opt into fraction-based warmup.
if self.lr_warmup_steps is None:
self.lr_warmup_steps = 0 if self.lr_warmup_fraction else 100

if self.lr_warmup_steps and self.max_steps and (self.lr_warmup_steps >= self.max_steps):
warnings.warn(
"`--train.lr_warmup_steps` should be less than `--train.max_steps`."
Expand Down
18 changes: 18 additions & 0 deletions tests/test_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,3 +34,21 @@ def test_compute_warmup_iters():
assert train.warmup_iters(devices=1, num_nodes=1, max_iters=20, train_dataloader=range(100)) == 20
# lr_warmup_fraction rounds up
assert train.warmup_iters(devices=1, num_nodes=1, max_iters=1000, train_dataloader=range(5)) == 2


def test_lr_warmup_fraction_works_without_overriding_default_steps():
"""Setting only `lr_warmup_fraction` must succeed even though `lr_warmup_steps` defaults to 100.

Previously a user passing `--train.lr_warmup_fraction 0.1` would hit a ValueError because the
mutually-exclusive validation considered the default `lr_warmup_steps=100` as user-provided.
"""
# only lr_warmup_fraction is provided
train = TrainArgs(global_batch_size=1, micro_batch_size=1, lr_warmup_fraction=0.3)
# lr_warmup_fraction must be honored, not silently overridden by the default lr_warmup_steps
assert train.warmup_iters(devices=1, num_nodes=1, max_iters=1000, train_dataloader=range(100)) == 30


def test_lr_warmup_conflict_only_when_both_explicitly_set():
"""Passing both lr_warmup_steps and lr_warmup_fraction explicitly must still raise."""
with pytest.raises(ValueError, match="Can't provide both `--train.lr_warmup_fraction`"):
TrainArgs(lr_warmup_steps=50, lr_warmup_fraction=0.2)