MCPcopy Create free account
hub / github.com/ModalityDance/Omni-R1 / get_scheduler

Function get_scheduler

src/transformers/src/transformers/optimization.py:471–555  ·  view source on GitHub ↗

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,
)

Source from the content-addressed store, hash-verified

469
470
471def 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 = {}

Callers 15

test_get_schedulerMethod · 0.90
finetuneFunction · 0.90
mainFunction · 0.90
mainFunction · 0.90
mainFunction · 0.90
mainFunction · 0.90
mainFunction · 0.90
mainFunction · 0.90
mainFunction · 0.90
mainFunction · 0.90

Calls 3

SchedulerTypeClass · 0.85
keysMethod · 0.45

Tested by 1

test_get_schedulerMethod · 0.72