Adam optimizer with l2 regularization.
(
learning_rate: schedule.Schedule,
*,
b1: float,
b2: float,
eps: float,
l2_regularizer_weight: float = 0,
l2_regularizer_per_param_scale: Optional[Callable[[NestedOptParam], Any]] = None,
mu_dtype: Optional[jnp.dtype] = None,
)
| 857 | |
| 858 | |
| 859 | def adam_optimizer( |
| 860 | learning_rate: schedule.Schedule, |
| 861 | *, |
| 862 | b1: float, |
| 863 | b2: float, |
| 864 | eps: float, |
| 865 | l2_regularizer_weight: float = 0, |
| 866 | l2_regularizer_per_param_scale: Optional[Callable[[NestedOptParam], Any]] = None, |
| 867 | mu_dtype: Optional[jnp.dtype] = None, |
| 868 | ) -> PartitionedGradientTransformation: |
| 869 | """Adam optimizer with l2 regularization.""" |
| 870 | tx = [ |
| 871 | l2_regularizer( |
| 872 | regularizer_weight=l2_regularizer_weight, |
| 873 | per_param_scale=l2_regularizer_per_param_scale, |
| 874 | ), |
| 875 | adam_partition(optax.scale_by_adam(b1=b1, b2=b2, eps=eps, mu_dtype=mu_dtype)), |
| 876 | scale_by_schedule(scale_from_learning_rate(learning_rate)), |
| 877 | ] |
| 878 | return chain(*tx) |
| 879 | |
| 880 | |
| 881 | class EmaState(NamedTuple): |