Класс LearningRateScheduler
Класс LearningRateScheduler применяется для
гибкого управления скоростью обучения во время
обучения модели. Он вызывается в конце каждой
эпохи и получает на вход текущий номер эпохи.
В качестве параметра класс принимает функцию
schedule, которая возвращает новое
значение скорости обучения в зависимости от
номера эпохи. Вторым необязательным параметром
является verbose, который управляет выводом
информации об изменении скорости обучения.
Синтаксис
tf.keras.callbacks.LearningRateScheduler(schedule, verbose=0)
Пример
Давайте создадим простую модель и применим
колбэк LearningRateScheduler с функцией,
которая уменьшает скорость обучения вдвое
каждые 5 эпох:
import tensorflow as tf
tf.random.set_seed(0)
def scheduler(epoch, lr):
if epoch % 5 == 0 and epoch > 0:
return lr * 0.5
return lr
model = tf.keras.Sequential([
tf.keras.layers.Dense(10, activation='relu', input_shape=(5,)),
tf.keras.layers.Dense(1)
])
model.compile(optimizer='sgd', loss='mse')
callback = tf.keras.callbacks.LearningRateScheduler(scheduler)
x = tf.constant([[1.0, 2.0, 3.0, 4.0, 5.0]])
y = tf.constant([[1.0]])
model.fit(x, y, epochs=10, verbose=0, callbacks=[callback])
res = model.optimizer.learning_rate.numpy()
print(res)
Результат выполнения кода:
0.03125
Пример
Давайте создадим колбэк с функцией экспоненциального
затухания и выведем значение скорости обучения
на каждой эпохе с помощью verbose=1:
import tensorflow as tf
tf.random.set_seed(0)
def scheduler(epoch, lr):
return 0.1 * (0.9 ** epoch)
model = tf.keras.Sequential([
tf.keras.layers.Dense(1, input_shape=(1,))
])
model.compile(optimizer='sgd', loss='mse')
callback = tf.keras.callbacks.LearningRateScheduler(scheduler, verbose=1)
x = tf.constant([[1.0], [2.0], [3.0]])
y = tf.constant([[2.0], [4.0], [6.0]])
model.fit(x, y, epochs=3, verbose=0, callbacks=[callback])
Результат выполнения кода:
"Epoch 00001: LearningRateScheduler setting learning rate to 0.1."
"Epoch 00002: LearningRateScheduler setting learning rate to 0.09."
"Epoch 00003: LearningRateScheduler setting learning rate to 0.081."
Смотрите также
-
класс
ReduceLROnPlateau,
который уменьшает скорость обучения при отсутствии улучшений -
класс
ExponentialDecay,
который реализует экспоненциальное затухание скорости обучения -
класс
CosineDecay,
который реализует косинусное затухание скорости обучения -
класс
EarlyStopping,
который останавливает обучение при отсутствии улучшений