Adafactor optimizer. References: https://arxiv.org/abs/1804.04235 https://github.com/deepmind/optax/blob/c4a4790b85ad69cda00a425cc3dcf9c9f9465120/optax/_src/alias.py#L77 WARNING: unlike adamw_optimizer, decay bias correction will *not* be applied if b1 or b2 is set to a constan
(
learning_rate: schedule.Schedule,
*,
b1: Optional[schedule.Schedule],
b2: schedule.Schedule,
multiply_by_parameter_scale: bool,
clipping_threshold: Optional[float],
dtype_momentum: Any = jnp.float32,
weight_decay: Optional[float] = None,
weight_decay_scale_by_learning_rate_exponent: Optional[float] = None,
weight_decay_per_param_scale: Optional[Callable[[NestedOptParam], Any]] = None,
eps: float = 1e-30,
factored: bool = True,
apply_scale_by_trust_ratio: bool = False,
)
| 1056 | |
| 1057 | |
| 1058 | def adafactor_optimizer( |
| 1059 | learning_rate: schedule.Schedule, |
| 1060 | *, |
| 1061 | b1: Optional[schedule.Schedule], |
| 1062 | b2: schedule.Schedule, |
| 1063 | multiply_by_parameter_scale: bool, |
| 1064 | clipping_threshold: Optional[float], |
| 1065 | dtype_momentum: Any = jnp.float32, |
| 1066 | weight_decay: Optional[float] = None, |
| 1067 | weight_decay_scale_by_learning_rate_exponent: Optional[float] = None, |
| 1068 | weight_decay_per_param_scale: Optional[Callable[[NestedOptParam], Any]] = None, |
| 1069 | eps: float = 1e-30, |
| 1070 | factored: bool = True, |
| 1071 | apply_scale_by_trust_ratio: bool = False, |
| 1072 | ) -> PartitionedGradientTransformation: |
| 1073 | """Adafactor optimizer. |
| 1074 | |
| 1075 | References: |
| 1076 | https://arxiv.org/abs/1804.04235 |
| 1077 | https://github.com/deepmind/optax/blob/c4a4790b85ad69cda00a425cc3dcf9c9f9465120/optax/_src/alias.py#L77 |
| 1078 | |
| 1079 | WARNING: unlike adamw_optimizer, decay bias correction will *not* be applied if b1 or b2 |
| 1080 | is set to a constant. However, users can enable bias correction by setting b1/b2 to |
| 1081 | config_for_function(schedule.decay_bias_correction).set(decay=<constant decay>). |
| 1082 | |
| 1083 | Args: |
| 1084 | learning_rate: (Schedule) The learning rate schedule. |
| 1085 | b1: (Schedule) first-moment exponential decay (beta1) schedule. If not None, enables |
| 1086 | momentum and uses extra memory. |
| 1087 | b2: (Schedule) second-moment exponential decay (beta2) schedule. |
| 1088 | multiply_by_parameter_scale: (bool): if True, then scale learning_rate by |
| 1089 | parameter norm. if False, provided learning_rate is absolute step size. |
| 1090 | Usually this should be set to False. |
| 1091 | clipping_threshold: (float>=1) optional value; if None, clipping disabled. |
| 1092 | dtype_momentum: (dtype) dtype of momentum buffers. |
| 1093 | weight_decay: (float) optional rate at which to decay weights. |
| 1094 | weight_decay_scale_by_learning_rate_exponent: (float) optional scale weight decay rate by |
| 1095 | (learning_rate ** exponent). Must not be None if weight_decay is not None. |
| 1096 | If set to 1, replicates the behavior of Adafactor weight decay in Lingvo and Praxis, |
| 1097 | https://github.com/google/praxis/blob/8fa3eb2e9ade0fd9a89a2ca56187882b12871605/praxis/optimizers.py#L1991-L1995. |
| 1098 | Set to 0 to disable scaling. |
| 1099 | weight_decay_per_param_scale: (optional) a Callable that returns a tree with same structure |
| 1100 | as the params PyTree, where each leaf is a float representing the per-param decay scale. |
| 1101 | The scale will be applied on top of the global decay rate: |
| 1102 | effective_decay_rate = global_decay_rate * per_param_scale. |
| 1103 | If None, all leaves will have a scale of 1. |
| 1104 | eps: (float) regularization constant for root mean squared gradient. |
| 1105 | factored: (bool) whether to use factored second-moment estimates. |
| 1106 | apply_scale_by_trust_ratio: (bool) whether to use variable-wise adaptive moments (LAMB): |
| 1107 | https://arxiv.org/abs/1904.00962. |
| 1108 | |
| 1109 | Returns: |
| 1110 | A PartitionedGradientTransformation representing an Adafactor optimizer. |
| 1111 | |
| 1112 | Raises: |
| 1113 | ValueError: If weight_decay_scale_by_learning_rate_exponent is not specified with |
| 1114 | weight_decay. |
| 1115 | """ |