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,
)
| 739 | |
| 740 | |
| 741 | def 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 | |
| 796 | def adamw_decoupled_optimizer( |