РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
427 of 824 menu

Класс 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,
    который позволяет создавать колбэки на лету
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить