Set learning rate(s) for optimizer.
(self, new_lrs: float | list)
| 376 | self.reset() |
| 377 | |
| 378 | def _set_learning_rate(self, new_lrs: float | list) -> None: |
| 379 | """Set learning rate(s) for optimizer.""" |
| 380 | if not isinstance(new_lrs, list): |
| 381 | new_lrs = [new_lrs] * len(self.optimizer.param_groups) |
| 382 | if len(new_lrs) != len(self.optimizer.param_groups): |
| 383 | raise ValueError( |
| 384 | "Length of `new_lrs` is not equal to the number of parameter groups " + "in the given optimizer" |
| 385 | ) |
| 386 | |
| 387 | for param_group, new_lr in zip(self.optimizer.param_groups, new_lrs): |
| 388 | param_group["lr"] = new_lr |
| 389 | |
| 390 | def _check_for_scheduler(self): |
| 391 | """Check optimizer doesn't already have scheduler.""" |