SGD optimizer implementation. Args: learning_rate: the learning rate schedule. decouple_weight_decay: Decouples weight decay so that it is not part of the gradient and thus do not affect the gradient accumulators. A brief guidance: - If you ar
(
learning_rate: schedule.Schedule,
*,
decouple_weight_decay: bool,
momentum: float = 0,
weight_decay: float = 0,
weight_decay_per_param_scale: Optional[Callable[[NestedOptParam], Any]] = None,
)
| 683 | |
| 684 | |
| 685 | def sgd_optimizer( |
| 686 | learning_rate: schedule.Schedule, |
| 687 | *, |
| 688 | decouple_weight_decay: bool, |
| 689 | momentum: float = 0, |
| 690 | weight_decay: float = 0, |
| 691 | weight_decay_per_param_scale: Optional[Callable[[NestedOptParam], Any]] = None, |
| 692 | ) -> PartitionedGradientTransformation: |
| 693 | """SGD optimizer implementation. |
| 694 | |
| 695 | Args: |
| 696 | learning_rate: the learning rate schedule. |
| 697 | decouple_weight_decay: Decouples weight decay so that it is not |
| 698 | part of the gradient and thus do not affect the gradient |
| 699 | accumulators. A brief guidance: |
| 700 | - If you are trying to reproduce an existing model trained |
| 701 | with SGD, you probably want to set it to `False` to get the |
| 702 | same behavior as in Torch or TF; |
| 703 | - If you are tuning a new model and plan to tune the weight |
| 704 | decay rate, you may want to set it to `True` to get a |
| 705 | simpler behavior. |
| 706 | Reference: |
| 707 | https://arxiv.org/abs/1711.05101 |
| 708 | https://www.fast.ai/2018/07/02/adam-weight-decay/#adamw |
| 709 | momentum: the momentum for SGD update. |
| 710 | weight_decay: the weight decay rate. |
| 711 | weight_decay_per_param_scale: the per-param decay scale. The scale |
| 712 | will be applied on top of the global decay rate. |
| 713 | |
| 714 | Returns: |
| 715 | A corresponding `PartitionedGradientTransformation`. |
| 716 | """ |
| 717 | if decouple_weight_decay: |
| 718 | return chain( |
| 719 | trace_partition(optax.trace(decay=momentum)), |
| 720 | add_decayed_weights( |
| 721 | weight_decay=weight_decay, |
| 722 | # Weight decay updates will already be scaled by the learning rate below. |
| 723 | learning_rate_exponent=None, |
| 724 | per_param_scale=weight_decay_per_param_scale, |
| 725 | ), |
| 726 | scale_by_schedule(scale_from_learning_rate(learning_rate)), |
| 727 | ) |
| 728 | else: |
| 729 | return chain( |
| 730 | add_decayed_weights( |
| 731 | weight_decay=weight_decay, |
| 732 | # Weight decay updates will already be scaled by the learning rate below. |
| 733 | learning_rate_exponent=None, |
| 734 | per_param_scale=weight_decay_per_param_scale, |
| 735 | ), |
| 736 | trace_partition(optax.trace(decay=momentum)), |
| 737 | scale_by_schedule(scale_from_learning_rate(learning_rate)), |
| 738 | ) |
| 739 | |
| 740 | |
| 741 | def adamw_optimizer( |