diff --git a/litgpt/args.py b/litgpt/args.py index ef41a0f006..56729fe97e 100644 --- a/litgpt/args.py +++ b/litgpt/args.py @@ -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 @@ -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`." diff --git a/tests/test_args.py b/tests/test_args.py index e4749b818f..c1f6aca6fc 100644 --- a/tests/test_args.py +++ b/tests/test_args.py @@ -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)