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

Function adamw_optimizer

axlearn/common/optimizers.py:741–793  ·  view source on GitHub ↗

AdamW optimizer with parameter scaling. N.B. The default weight-decay implementation is consistent with those in e.g. PyTorch & Optax, but inconsistent with the "decoupled" adamw weight decay formulation in Algorithm 2. To faithfully re

(
    learning_rate: schedule.Schedule,
    *,
    b1: float,
    b2: float,
    eps: float,
    weight_decay: float = 0,
    weight_decay_per_param_scale: Optional[Callable[[NestedOptParam], Any]] = None,
    mu_dtype: Optional[jnp.dtype] = None,
    adam_update_transformation: Optional[ConfigOr[PartitionedGradientTransformation]] = None,
)

Source from the content-addressed store, hash-verified

739
740
741def adamw_optimizer(
742 learning_rate: schedule.Schedule,
743 *,
744 b1: float,
745 b2: float,
746 eps: float,
747 weight_decay: float = 0,
748 weight_decay_per_param_scale: Optional[Callable[[NestedOptParam], Any]] = None,
749 mu_dtype: Optional[jnp.dtype] = None,
750 adam_update_transformation: Optional[ConfigOr[PartitionedGradientTransformation]] = None,
751) -> PartitionedGradientTransformation:
752 """AdamW optimizer with parameter scaling.
753
754 N.B. The default weight-decay implementation is consistent with
755 those in e.g. PyTorch & Optax, but inconsistent with the "decoupled"
756 adamw weight decay formulation in <https://arxiv.org/abs/1711.05101> Algorithm 2.
757 To faithfully replicate Algorithm 2, use `adamw_decoupled_optimizer`.
758
759 Args:
760 learning_rate: the learning rate schedule.
761 b1: the exponential decay rate for the 1st moment estimates.
762 b2: the exponential decay rate for the 2nd moment estimates.
763 eps: a small constant for numerical stability.
764 weight_decay: optional rate at which to decay weights.
765 weight_decay_per_param_scale: a Callable that returns a tree with same structure
766 as the params PyTree, where each leaf is a float representing the per-param decay scale.
767 The scale will be applied on top of the global decay rate:
768 effective_decay_rate = global_decay_rate * per_param_scale.
769 If None, all leaves will have a scale of 1.
770 mu_dtype: optional `dtype` to be used for the first order accumulator;
771 if `None` then the dtype is inferred from params and updates.
772 adam_update_transformation: A transformation applied directly on the adam updates
773 (but before weight decay). If None, no transformation is applied.
774
775 Returns:
776 A PartitionedGradientTransformation representing an AdamW optimizer with parameter scaling.
777 """
778 tx = [adam_partition(optax.scale_by_adam(b1=b1, b2=b2, eps=eps, mu_dtype=mu_dtype))]
779 if adam_update_transformation is not None:
780 tx.append(maybe_instantiate(adam_update_transformation))
781 tx.extend(
782 [
783 add_decayed_weights(
784 weight_decay=weight_decay,
785 # Weight decay updates will already be scaled by the learning rate below.
786 learning_rate_exponent=None,
787 per_param_scale=weight_decay_per_param_scale,
788 ),
789 scale_by_schedule(scale_from_learning_rate(learning_rate)),
790 ]
791 )
792
793 return chain(*tx)
794
795
796def adamw_decoupled_optimizer(

Calls 6

maybe_instantiateFunction · 0.90
adam_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_adamw_optimizerMethod · 0.72