DeepSpeed
8cae9d28 - Implement the documented per-param-group lists in OneCycle (#8201)

Commit
6 days ago
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>
Author
Parents
Loading