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

Класс EarlyStopping

Класс EarlyStopping представляет собой колбэк, который останавливает процесс обучения модели, когда выбранная метрика перестает улучшаться на протяжении заданного числа эпох. Первым параметром передается имя метрики для наблюдения, например val_loss. Параметр patience задает количество эпох ожидания улучшения, а restore_best_weights определяет, нужно ли вернуть веса лучшей эпохи. Класс также принимает monitor, min_delta, mode, baseline и verbose.

Синтаксис

tf.keras.callbacks.EarlyStopping( monitor='val_loss', min_delta=0, patience=0, verbose=0, mode='auto', baseline=None, restore_best_weights=False )

Пример

Давайте создадим простую модель и обучим ее с колбэком EarlyStopping, который останавливает обучение после двух эпох без улучшения:

import tensorflow as tf tf.random.set_seed(0) model = tf.keras.Sequential([ tf.keras.layers.Dense(4, activation='relu', input_shape=(3,)), tf.keras.layers.Dense(1) ]) model.compile(optimizer='adam', loss='mse') x = tf.constant([[1, 2, 3], [4, 5, 6], [7, 8, 9], [1, 1, 1]]) y = tf.constant([[1], [2], [3], [1]]) callback = tf.keras.callbacks.EarlyStopping( monitor='loss', patience=2, restore_best_weights=True ) history = model.fit(x, y, epochs=50, verbose=0, callbacks=[callback]) print(len(history.history['loss']))

Результат выполнения кода:

4

Пример

Давайте обучим модель с отслеживанием val_loss и выведем информацию об остановке:

import tensorflow as tf tf.random.set_seed(0) model = tf.keras.Sequential([ tf.keras.layers.Dense(4, activation='relu', input_shape=(3,)), tf.keras.layers.Dense(1) ]) model.compile(optimizer='adam', loss='mse') x = tf.constant([[1, 2, 3], [4, 5, 6], [7, 8, 9], [1, 1, 1]]) y = tf.constant([[1], [2], [3], [1]]) callback = tf.keras.callbacks.EarlyStopping( monitor='val_loss', patience=3, verbose=1, restore_best_weights=True ) history = model.fit( x, y, validation_split=0.25, epochs=100, verbose=0, callbacks=[callback] ) print(history.history.keys())

Результат выполнения кода:

dict_keys(['loss', 'val_loss'])

Пример

Давайте используем параметр min_delta, чтобы учитывать только значимые улучшения:

import tensorflow as tf tf.random.set_seed(0) model = tf.keras.Sequential([ tf.keras.layers.Dense(4, activation='relu', input_shape=(3,)), tf.keras.layers.Dense(1) ]) model.compile(optimizer='adam', loss='mse') x = tf.constant([[1, 2, 3], [4, 5, 6], [7, 8, 9], [1, 1, 1]]) y = tf.constant([[1], [2], [3], [1]]) callback = tf.keras.callbacks.EarlyStopping( monitor='loss', min_delta=0.01, patience=2, restore_best_weights=True ) history = model.fit(x, y, epochs=50, verbose=0, callbacks=[callback]) print(callback.best_epoch)

Результат выполнения кода:

3

Смотрите также

  • класс ModelCheckpoint,
    который сохраняет модель во время обучения
  • класс ReduceLROnPlateau,
    который снижает скорость обучения при остановке улучшений
  • класс TensorBoard,
    который записывает логи обучения для визуализации
  • класс TerminateOnNaN,
    который останавливает обучение при появлении NaN в потере
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить