Unified API to get any scheduler from its name. Args: name (`str` or `SchedulerType`): The name of the scheduler to use. optimizer (`torch.optim.Optimizer`): The optimizer that will be used during training. num_warmup_steps (`int`, *optional*
(
name: Union[str, SchedulerType],
optimizer: Optimizer,
num_warmup_steps: Optional[int] = None,
num_training_steps: Optional[int] = None,
scheduler_specific_kwargs: Optional[dict] = None,
)
| 469 | |
| 470 | |
| 471 | def get_scheduler( |
| 472 | name: Union[str, SchedulerType], |
| 473 | optimizer: Optimizer, |
| 474 | num_warmup_steps: Optional[int] = None, |
| 475 | num_training_steps: Optional[int] = None, |
| 476 | scheduler_specific_kwargs: Optional[dict] = None, |
| 477 | ): |
| 478 | """ |
| 479 | Unified API to get any scheduler from its name. |
| 480 | |
| 481 | Args: |
| 482 | name (`str` or `SchedulerType`): |
| 483 | The name of the scheduler to use. |
| 484 | optimizer (`torch.optim.Optimizer`): |
| 485 | The optimizer that will be used during training. |
| 486 | num_warmup_steps (`int`, *optional*): |
| 487 | The number of warmup steps to do. This is not required by all schedulers (hence the argument being |
| 488 | optional), the function will raise an error if it's unset and the scheduler type requires it. |
| 489 | num_training_steps (`int``, *optional*): |
| 490 | The number of training steps to do. This is not required by all schedulers (hence the argument being |
| 491 | optional), the function will raise an error if it's unset and the scheduler type requires it. |
| 492 | scheduler_specific_kwargs (`dict`, *optional*): |
| 493 | Extra parameters for schedulers such as cosine with restarts. Mismatched scheduler types and scheduler |
| 494 | parameters will cause the scheduler function to raise a TypeError. |
| 495 | """ |
| 496 | name = SchedulerType(name) |
| 497 | schedule_func = TYPE_TO_SCHEDULER_FUNCTION[name] |
| 498 | |
| 499 | # If a `LayerWiseDummyOptimizer` is passed we extract the optimizer dict and |
| 500 | # recursively call `get_scheduler` to get the proper schedulers on each parameter |
| 501 | if optimizer is not None and isinstance(optimizer, LayerWiseDummyOptimizer): |
| 502 | optimizer_dict = optimizer.optimizer_dict |
| 503 | scheduler_dict = {} |
| 504 | |
| 505 | for param in optimizer_dict.keys(): |
| 506 | scheduler_dict[param] = get_scheduler( |
| 507 | name, |
| 508 | optimizer=optimizer_dict[param], |
| 509 | num_warmup_steps=num_warmup_steps, |
| 510 | num_training_steps=num_training_steps, |
| 511 | ) |
| 512 | |
| 513 | def scheduler_hook(param): |
| 514 | # Since the optimizer hook has been already attached we only need to |
| 515 | # attach the scheduler hook, the gradients have been zeroed here |
| 516 | scheduler_dict[param].step() |
| 517 | |
| 518 | for param in optimizer_dict.keys(): |
| 519 | if param.requires_grad: |
| 520 | param.register_post_accumulate_grad_hook(scheduler_hook) |
| 521 | |
| 522 | return LayerWiseDummyScheduler(optimizer_dict=optimizer_dict, lr=optimizer.defaults["lr"]) |
| 523 | |
| 524 | if name == SchedulerType.CONSTANT: |
| 525 | return schedule_func(optimizer) |
| 526 | |
| 527 | if scheduler_specific_kwargs is None: |
| 528 | scheduler_specific_kwargs = {} |