Следите за новинками
в нашем Telegram канале. Жми, чтобы подписаться:)
429 of 824 menu
◀ ▶

Класс LambdaCallback

Класс LambdaCallback позволяет создавать простые, не сохраняемые колбэки для использования во время обучения модели. Конструктор принимает набор необязательных именованных параметров, каждый из которых является функцией, вызываемой в определенный момент обучения. Это может быть начало или конец эпохи, начало или конец пакета данных, а также начало или конец обучения.

Основные параметры: on_epoch_begin, on_epoch_end, on_batch_begin, on_batch_end, on_train_begin, on_train_end. Все они принимают функции обратного вызова, которые получают аргументы epoch, logs или batch в зависимости от типа колбэка.

Синтаксис

tf.keras.callbacks.LambdaCallback( on_epoch_begin=None, on_epoch_end=None, on_batch_begin=None, on_batch_end=None, on_train_begin=None, on_train_end=None, )

Пример

Давайте создадим простой колбэк, который выводит сообщение в начале и в конце каждой эпохи обучения:

import tensorflow as tf tf.random.set_seed(0) # Define callback functions def on_epoch_begin(epoch, logs): print(f"Start of epoch {epoch}") def on_epoch_end(epoch, logs): print(f"End of epoch {epoch}, loss: {logs['loss']:.4f}") # Create LambdaCallback callback = tf.keras.callbacks.LambdaCallback( on_epoch_begin=on_epoch_begin, on_epoch_end=on_epoch_end ) # Create and train a simple model model = tf.keras.Sequential([ tf.keras.layers.Dense(1, input_shape=(1,)) ]) model.compile(optimizer='sgd', loss='mse') x = tf.constant([[1.0], [2.0], [3.0], [4.0]]) y = tf.constant([[2.0], [4.0], [6.0], [8.0]]) model.fit(x, y, epochs=2, verbose=0, callbacks=[callback])

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

"Start of epoch 0" "End of epoch 0, loss: 30.0000" "Start of epoch 1" "End of epoch 1, loss: 15.0000"

Пример

Давайте создадим колбэк, который отслеживает завершение каждого пакета данных и выводит его номер:

import tensorflow as tf tf.random.set_seed(0) # Define callback function def on_batch_end(batch, logs): print(f"Batch {batch} finished, loss: {logs['loss']:.4f}") # Create LambdaCallback callback = tf.keras.callbacks.LambdaCallback( on_batch_end=on_batch_end ) # Create and train a simple model model = tf.keras.Sequential([ tf.keras.layers.Dense(1, input_shape=(1,)) ]) model.compile(optimizer='sgd', loss='mse') x = tf.constant([[1.0], [2.0], [3.0], [4.0]]) y = tf.constant([[2.0], [4.0], [6.0], [8.0]]) model.fit(x, y, epochs=1, batch_size=2, verbose=0, callbacks=[callback])

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

"Batch 0 finished, loss: 30.0000" "Batch 1 finished, loss: 15.0000"

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

  • класс ModelCheckpoint,
    который сохраняет модель во время обучения
  • класс EarlyStopping,
    который останавливает обучение при отсутствии улучшений
  • класс ReduceLROnPlateau,
    который уменьшает скорость обучения при застое
  • класс CSVLogger,
    который записывает историю обучения в CSV-файл
← →
↑
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить