File size: 6,360 Bytes
f4a39ee | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 | # Copyright 2024 Google LLC
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Configurable optimizers from JAX."""
import collections
import re
from typing import Sequence
import gin
import optax
gin.external_configurable(optax.adabelief, module='optax')
gin.external_configurable(optax.adam, module='optax')
gin.external_configurable(optax.adamw, module='optax')
gin.external_configurable(optax.constant_schedule, module='optax')
gin.external_configurable(optax.join_schedules, module='optax')
gin.external_configurable(optax.piecewise_constant_schedule, module='optax')
gin.external_configurable(optax.exponential_decay, module='optax')
gin.external_configurable(
optax.warmup_exponential_decay_schedule, module='optax'
)
class OptimizerError(Exception):
"""Raised if a custom Whirl optimizer encounters an error."""
@gin.configurable
def optimizer(value):
return value
OptState = collections.namedtuple('OptState', ['state', 'params'])
@gin.register
def piecewise_constant_schedule_specified_by_rates(
rates: Sequence[float],
boundaries: Sequence[int],
) -> optax.Schedule:
"""Schedule that is piecewise constant and specified by rates (not scales).
This is similar to optax.piecewise_constant_schedule, which requires users
to specify "scales" (ratio of old LR to new LR).
Args:
rates: Length K sequence of learning rates. `rates[i]` is used for steps
`0 <= step < boundaries[1]`, for i=0, and
`boundaries[i-1] <= step < boundaries[i]`, for 0 < i < len(boundaries)
`boundaries[i-1] <= step < ∞`, for i = len(boundaries)
boundaries: Length K-1 sequence of boundaries.
Returns:
Schedule to pass to optax optimizers.
"""
return optax.join_schedules(
schedules=[optax.constant_schedule(r) for r in rates],
boundaries=boundaries,
)
@gin.register
def delayed_constant_schedule(
turn_on_step: int,
rate: float,
) -> optax.Schedule:
"""Schedule that is zero until `turn_on_step` then `rate` thereafter."""
return piecewise_constant_schedule_specified_by_rates(
rates=[0., rate],
boundaries=[turn_on_step],
)
@gin.register
def top_level_multi_adam(
top_level_keys: Sequence[str] = (),
learning_rates: Sequence[optax.ScalarOrSchedule] = (),
default_learning_rate: optax.ScalarOrSchedule = 1e-4,
b1: float = 0.9,
b2: float = 0.95,
eps: float = 1e-6,
raise_if_keys_not_found: bool = True,
) -> optax.GradientTransformation:
"""Uses an Adam optimizer with different learning rates for different params.
Args:
top_level_keys: Keys to use non-default learning rates for. A key starting
with 'REGEX_', such as 'REGEX_cats' will use re.search to find keys, e.g.
re.search('cats', key).
learning_rates: Learning rates to use leafs under the `top_level_keys`.
default_learning_rate: Learning rate to use for keys not in `learning_rates`
b1: Exponential decay to track the first moment of past gradients.
b2: Exponential decay to track the second moment of past gradients.
eps: A small constant applied to denominator outside of the square root to
avoid dividing by zero when rescaling.
raise_if_keys_not_found: Whether to raise if some `top_level_keys` are not
found in params.
Returns:
optax optimizer with learning rate based on top level key in params dict.
"""
if len(top_level_keys) != len(learning_rates):
raise ValueError(
f'{top_level_keys=} had different length than {learning_rates=}'
)
if '' in top_level_keys:
raise ValueError('An empty string "" was found in `top_level_keys`.')
default_label = 'DEFAULT_LABEL'
if default_label in top_level_keys:
raise ValueError(f'{default_label=} should not be in `top_level_keys`')
def find_matching_top_level_key(param_name: str) -> str:
"""Searches for param_name in top_level_keys, returns the matching key."""
prefix = 'REGEX_'
matches = []
for k in top_level_keys:
if k.startswith(prefix) and re.search(k.lstrip(prefix), param_name):
matches.append(k)
elif k == param_name:
matches.append(k)
if not matches:
return default_label
elif len(matches) == 1:
return matches[0]
else:
raise ValueError(
f'{param_name=} had more than 1 ({len(matches)}) match '
f'({matches}). Only one `top_level_keys` should match, or else we '
'cannot choose a unique learning rate for these parameters.'
)
def get_prefix_labels(params):
"""Makes prefix labels to help optax match params with learning rates."""
# E.g. if top_level_keys = ['module_A', 'REGEX_special'],
# and params.keys() = ['module_A', 'special_A', 'special_B', 'module_C'],
# labels = {
# 'module_A': 'module_A',
# 'special_A': 'REGEX_special', 'special_B': 'REGEX_special',
# 'module_C': 'DEFAULT_LABEL', 'module_D': 'DEFAULT_LABEL',...
# }
# E.g. labels tells optax to use the learning rate 'REGEX_special' for
# parameters under the prefix 'module_C'.
labels = {
param_name: find_matching_top_level_key(param_name)
for param_name in params
}
top_level_keys_that_matched = [
k for k in labels.values() if k != default_label
]
missing_keys = set(top_level_keys).difference(top_level_keys_that_matched)
if raise_if_keys_not_found and missing_keys:
raise OptimizerError(
f'{missing_keys=} not found in params: {sorted(params)}'
)
return labels
def make_adam(lr):
return optax.adam(lr, b1=b1, b2=b2, eps=eps)
return optax.multi_transform(
transforms={ # pyrefly: ignore[bad-argument-type]
k: make_adam(lr) for k, lr in zip(top_level_keys, learning_rates)
}
| {default_label: make_adam(default_learning_rate)},
param_labels=get_prefix_labels,
)
|