Класс TerminateOnNaN
Класс TerminateOnNaN представляет собой
колбэк, который прерывает процесс обучения,
если значение функции потерь становится
равным NaN или бесконечности.
Это помогает вовремя остановить обучение,
когда численные значения выходят за пределы
допустимого диапазона, например из-за слишком
большой скорости обучения. Колбэк не принимает
обязательных параметров и передается в метод
fit через аргумент callbacks.
Синтаксис
tf.keras.callbacks.TerminateOnNaN()
Пример
Давайте создадим простую модель и передадим
колбэк TerminateOnNaN в метод fit,
чтобы обучение остановилось при появлении NaN:
import tensorflow as tf
import numpy as np
tf.random.set_seed(0)
model = tf.keras.Sequential([
tf.keras.layers.Dense(1, input_shape=(1,))
])
model.compile(optimizer='sgd', loss='mse')
x = np.array([1, 2, 3, 4, 5], dtype=np.float32)
y = np.array([2, 4, 6, 8, 10], dtype=np.float32)
callback = tf.keras.callbacks.TerminateOnNaN()
history = model.fit(x, y, epochs=5, verbose=0, callbacks=[callback])
print(history.history['loss'])
Результат выполнения кода:
[13.74276065826416, 12.945170402526855, 12.204591751098633, 11.516819953918457, 10.878060340881348]
Пример
Давайте искусственно вызовем появление NaN,
установив слишком большую скорость обучения,
и убедимся, что колбэк остановит обучение:
<+python+>
import tensorflow as tf
import numpy as np
tf.random.set_seed(0)
model = tf.keras.Sequential([
tf.keras.layers.Dense(1, input_shape=(1,))
])
model.compile(optimizer=tf.keras.optimizers.SGD(learning_rate=1e10), loss='mse')
x = np.array([1, 2, 3, 4, 5], dtype=np.float32)
y = np.array([2, 4, 6, 8, 10], dtype=np.float32)
callback = tf.keras.callbacks.TerminateOnNaN()
history = model.fit(x, y, epochs=10, verbose=0, callbacks=[callback])
print(history.history['loss'])
<-python+>
Результат выполнения кода:
<+python+>
[nan]
<-python+>
Смотрите также
-
класс
EarlyStopping,
который останавливает обучение при отсутствии улучшений -
класс
ModelCheckpoint,
который сохраняет модель во время обучения -
класс
ReduceLROnPlateau,
который уменьшает скорость обучения при застое -
класс
LambdaCallback,
который позволяет создавать колбэки на лету