MCPcopy Create free account
hub / github.com/apple/axlearn / sgd_optimizer

Function sgd_optimizer

axlearn/common/optimizers.py:685–738  ·  view source on GitHub ↗

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

Source from the content-addressed store, hash-verified

683
684
685def 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
741def adamw_optimizer(

Callers 2

test_sgd_optimizerMethod · 0.90

Calls 5

trace_partitionFunction · 0.85
add_decayed_weightsFunction · 0.85
scale_by_scheduleFunction · 0.85
scale_from_learning_rateFunction · 0.85
chainFunction · 0.70

Tested by 2

test_sgd_optimizerMethod · 0.72