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

Function adam_optimizer

axlearn/common/optimizers.py:859–878  ·  view source on GitHub ↗

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

Source from the content-addressed store, hash-verified

857
858
859def 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
881class EmaState(NamedTuple):

Callers 2

test_adam_optimizerMethod · 0.90

Calls 5

l2_regularizerFunction · 0.85
adam_partitionFunction · 0.85
scale_by_scheduleFunction · 0.85
scale_from_learning_rateFunction · 0.85
chainFunction · 0.70

Tested by 2

test_adam_optimizerMethod · 0.72