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

Function adafactor_optimizer

axlearn/common/optimizers.py:1058–1146  ·  view source on GitHub ↗

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

Source from the content-addressed store, hash-verified

1056
1057
1058def 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 """

Calls 10

scale_by_factored_rmsFunction · 0.90
clip_by_block_rmsFunction · 0.85
scale_by_scheduleFunction · 0.85
scale_from_learning_rateFunction · 0.85
scale_by_param_block_rmsFunction · 0.85
emaFunction · 0.85
add_decayed_weightsFunction · 0.85
scale_by_trust_ratioFunction · 0.85
scaleFunction · 0.85
chainFunction · 0.70

Tested by 5

_compare_layersMethod · 0.72