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

Function lion_optimizer

axlearn/common/optimizers.py:1709–1757  ·  view source on GitHub ↗

Lion optimizer with parameter scaling. https://arxiv.org/abs/2302.06675 Adapted from https://github.com/google/automl/blob/master/lion/lion_optax.py Args: learning_rate: The learning rate schedule. b1: The exponential decay rate for the 1st moment estimates. b2:

(
    learning_rate: schedule.Schedule,
    b1: float,
    b2: float,
    mu_dtype: Optional[jnp.dtype] = None,
    weight_decay: float = 0.0,
    multiply_by_parameter_scale: bool = False,
    weight_decay_per_param_scale: Optional[Callable[[NestedOptParam], Any]] = None,
)

Source from the content-addressed store, hash-verified

1707
1708
1709def lion_optimizer(
1710 learning_rate: schedule.Schedule,
1711 b1: float,
1712 b2: float,
1713 mu_dtype: Optional[jnp.dtype] = None,
1714 weight_decay: float = 0.0,
1715 multiply_by_parameter_scale: bool = False,
1716 weight_decay_per_param_scale: Optional[Callable[[NestedOptParam], Any]] = None,
1717) -> PartitionedGradientTransformation:
1718 """Lion optimizer with parameter scaling.
1719
1720 https://arxiv.org/abs/2302.06675
1721 Adapted from https://github.com/google/automl/blob/master/lion/lion_optax.py
1722
1723 Args:
1724 learning_rate: The learning rate schedule.
1725 b1: The exponential decay rate for the 1st moment estimates.
1726 b2: The exponential decay rate for the 2nd moment estimates.
1727 weight_decay: Optional rate at which to decay weights.
1728 weight_decay_per_param_scale: A Callable that returns a tree with same structure
1729 as the params PyTree, where each leaf is a float representing the per-param decay scale.
1730 The scale will be applied on top of the global decay rate:
1731 effective_decay_rate = global_decay_rate * per_param_scale.
1732 If None, all leaves will have a scale of 1.
1733 mu_dtype: Optional `dtype` to be used for the first order accumulator;
1734 if `None` then the dtype is inferred from params and updates.
1735 multiply_by_parameter_scale: If `True`, then scale learning_rate by
1736 parameter RMS. if `False`, provided learning_rate is absolute step size.
1737 Usually this should be left as False.
1738
1739 Returns:
1740 A PartitionedGradientTransformation representing an Lion optimizer with parameter scalin.
1741 """
1742 tx = [scale_by_lion(b1=b1, b2=b2, mu_dtype=mu_dtype)]
1743 if multiply_by_parameter_scale:
1744 tx.append(scale_by_param_block_rms())
1745 tx.extend(
1746 [
1747 add_decayed_weights(
1748 weight_decay=weight_decay,
1749 # Weight decay updates will already be scaled by the learning rate below.
1750 learning_rate_exponent=None,
1751 per_param_scale=weight_decay_per_param_scale,
1752 ),
1753 scale_by_schedule(scale_from_learning_rate(learning_rate)),
1754 ]
1755 )
1756
1757 return chain(*tx)
1758
1759
1760def adastar_optimizer(

Callers 2

test_lion_optimizerMethod · 0.90

Calls 6

scale_by_lionFunction · 0.85
scale_by_param_block_rmsFunction · 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_lion_optimizerMethod · 0.72