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,
)
| 1707 | |
| 1708 | |
| 1709 | def 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 | |
| 1760 | def adastar_optimizer( |