Implement the documented per-param-group lists in OneCycle (#8201)
## The bug
`OneCycle` documents four of its arguments as accepting a
per-param-group list:
```
cycle_min_lr (float or list): Initial learning rate which is the
lower boundary in the cycle for each parameter group.
cycle_max_lr (float or list): Upper learning rate boundaries in the cycle
for each parameter group.
cycle_min_mom (float or list): Initial momentum which is the
lower boundary in the cycle for each parameter group.
cycle_max_mom (float or list): Upper momentum boundaries in the cycle
for each parameter group.
```
`_initialize_lr` and `_initialize_momentum` only ever broadcast a
scalar:
```python
self.min_lrs = [cycle_min_lr] * len(optimizer.param_groups)
...
self.min_moms = [(cycle_min_mom, 0.99)] * len(optimizer.param_groups)
```
so the documented list is written whole into every param group, and the
optimizer is left holding a list where it expects a number:
```
param_group lrs after construction: [[0.001, 0.002], [0.001, 0.002]]
param_group betas after construction: [([0.8, 0.85], 0.99), ([0.8, 0.85], 0.99)]
scheduler.step() -> TypeError: unsupported operand type(s) for -: 'list' and 'list'
optimizer.step() -> TypeError: unsupported operand type(s) for -: 'int' and 'list'
```
The second line matters: the optimizer is corrupt from construction, so
even a plain `optimizer.step()` fails before the scheduler is stepped at
all.
This is reachable from a plain JSON config, not just the Python API.
`engine.py:1550` does `scheduler(optimizer, **scheduler_params)`, so
`"cycle_min_lr": [0.001, 0.002]` in `ds_config` deserializes to a Python
list and lands directly in `OneCycle.__init__`.
A wrong-length list is also accepted silently, where the siblings raise:
```
OneCycle: accepted 3 values for 2 param groups, no error
LRRangeTest: ValueError expected 2 lr_range_test_min_lr, got 3
WarmupLR: ValueError expected 2 value for min_lr, got [0.0, 0.1, 0.2]
```
## Why implement it rather than delete the docstring lines
Deleting the four "or list" claims would be a smaller diff, but the rest
of `OneCycle` is already per-group end to end: `_get_cycle_lr` zips
`min_lrs` with `max_lrs`, `_get_cycle_mom` zips `min_moms` with
`max_moms`, and `update_lr` walks the param groups. Only the two
initializers collapse the input. Both sibling schedulers in this file
implement the same documented contract, and the two most recent
multi-group fixes here (#7969 for `WarmupCosineLR`, #8171 for
`WarmupLR`) went in the same direction. This reads as an unfinished port
rather than a design decision.
## The fix
Reuse `_format_param`, which is how the siblings already honour this
contract. It was defined twice, identically: as a method on `WarmupLR`,
and again on `WarmupCosineLR` where nothing calls it (`_format_param`
appears in only two files repo-wide, and in the test file only inside a
comment). I promoted the single copy to module level next to `update_lr`
and `get_torch_optimizer`, dropped the dead one, and pointed `WarmupLR`
and `OneCycle` at it. Net result is 19 added, 22 removed, and one
implementation of this logic instead of two.
I chose promoting over leaving one-line delegate methods behind because
`_format_param` is private and has no callers outside this file, so a
delegate would be indirection with no consumer; happy to switch to
delegates if you would rather not remove the methods.
Three details worth calling out rather than leaving for review:
**The momentum call has to wrap the scalar, not the tuple.**
`_format_param` accepts tuples, and the default `cycle_min_mom` pairs
with `0.99` into a length-2 tuple, so wrapping the existing
`(cycle_min_mom, 0.99)` expression would raise at construction for 1 and
3 param groups, and for exactly 2 groups would silently write
`group['betas'] = 0.8` as a float and blow up later in `_get_cycle_mom`.
The correct form, which is what this PR uses, formats the scalar first:
```python
self.min_moms = [(mom, 0.99) for mom in _format_param(optimizer, cycle_min_mom, 'cycle_min_mom')]
```
**Both bounds are now validated before the optimizer is touched.**
`_initialize_lr` used to compute `min_lrs`, write `group['lr']`, and
only then look at `cycle_max_lr`, so a bad-length `cycle_max_lr` left
the param groups half updated. Moving the second `_format_param` call
above the mutation loop makes the constructor all-or-nothing:
```
before: lrs after a failed ctor = [[0.001, 0.002], [0.001, 0.002]]
after: ValueError, lrs after a failed ctor = [0.1, 0.2] (untouched)
```
**One token in `_format_param`'s error message.** Both copies
interpolate `FileNotFoundError(param_value)` where the wording promises
a count, so `WarmupLR` currently reports `expected 2 value for min_lr,
got [0.0, 0.1, 0.2]`. Since the two copies are collapsing into one
shared helper, I corrected it to `len(param_value)` rather than carry
the typo into the surviving copy. It is the only change to `WarmupLR`'s
behaviour and nothing asserts on that message (no `pytest.raises(...,
match=...)` anywhere in the file); say the word and I will drop it back
to verbatim.
**Not claiming this is strictly safer for momentum.** Because
`_format_param` accepts tuples, a betas-shaped `cycle_min_mom=(0.8,
0.999)` on a two-group optimizer goes from a loud `TypeError` to
silently training with per-group momenta. That hazard already exists
identically in `WarmupLR`, so I kept the behaviour symmetric rather than
diverging, but it is a real trade rather than a pure win.
## Tests
Added to `tests/unit/runtime/test_lr_schedulers.py` as module-level
functions, matching the existing plain tests there:
- `test_one_cycle_accepts_per_group_lr_and_momentum_lists`: two param
groups, per-group lists for all four arguments, asserting the
constructor sets each group's own lr and `betas[0]`, that the cycle peak
reaches each group's own `cycle_max_lr` with momentum at its own
`cycle_min_mom`, and that the bottom of the cycle returns each group to
its own `cycle_max_mom`.
- `test_one_cycle_rejects_wrong_length_per_group_lists`, parametrized
over all four arguments.
It uses `Adam` rather than `SGD` on purpose: `_initialize_momentum`
returns early when `'betas' not in optimizer.defaults`, so the momentum
half of the test would silently never run under SGD.
`pytest` cannot start on my machine (no GPU, and the `tests/unit`
conftest pulls in the distributed harness), so I ran the module-level
tests in this file directly against the real `lr_schedules.py`, with the
`DistributedTest` classes stripped and only `deepspeed.utils.logger`
stubbed. Three runs:
```
control upstream lr_schedules.py + upstream tests 21 passed, 0 failed
before upstream lr_schedules.py + these tests 21 passed, 5 failed
after this branch 26 passed, 0 failed
```
All 5 failures before are the new tests, and the 21 pre-existing ones
are unchanged by this diff. The `DistributedTest` OneCycle coverage
(`TestOneCycle.test_lr`, `test_mom`) and the other scalar-momentum users
(`test_fp16.py`, `test_bf16.py`, `test_pipeline.py`,
`test_other_optimizer.py`) all pass scalars, which take the unchanged
broadcast path; I am relying on CI for those since they need a GPU.
Lint: `yapf` 0.40.0 with the repo's `.style.yapf` reports no diff on
both files, and `flake8` with the repo's `.flake8` is clean on both
(also confirmed clean on the unmodified files, so that is a real result
rather than a config that checks nothing).
## Prior art
No open or closed PR implements list support here. `--search` over
`lr_schedules`, `_format_param`, `OneCycle`, `cycle_min_lr` and `lr
scheduler list param groups` turns up #8151, #8166, #8171, #7969, #8179,
#1455 and #4563, all merged and none touching these two initializers. No
open issue covers it either; the only open `OneCycle` issue is #3492, a
request for `CosineAnnealingLR` support.
This follows #8179 in the same class, so to be upfront about it: that
one was about the cycle shape (`_initialize_cycle` and
`_get_scale_factor`), this one is about the two value initializers, and
I did not see it while in there. If you would rather batch further
`lr_schedules.py` work, tell me and I will hold the rest.
Signed-off-by: Vineeth Sai <vineethsai4444@gmail.com>
Co-authored-by: Zhipeng Wang <zhipeng.rainbowserie@gmail.com>
Co-authored-by: Masahiro Tanaka <81312776+tohtana@users.noreply.github.com>