Класс 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 в потере