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

Метод on_batch_end класса Callback

Метод on_batch_end класса Callback вызывается автоматически в конце каждого батча во время обучения, оценки или предсказания модели. Метод принимает один параметр batch - номер текущего батча (целое число, начиная с 0), а также может принимать параметр logs - словарь с метриками текущего батча. Метод не возвращает значения и используется для реализации пользовательской логики: логирования, изменения скорости обучения, сохранения контрольных точек и других действий.

Синтаксис

class MyCallback(tf.keras.callbacks.Callback): def on_batch_end(self, batch, logs=None): # logic at the end of each batch

Пример

Давайте создадим простой колбэк, который выводит номер завершенного батча и значение функции потерь:

import tensorflow as tf import numpy as np tf.random.set_seed(0) class BatchLogger(tf.keras.callbacks.Callback): def on_batch_end(self, batch, logs=None): logs = logs or {} print(f"batch {batch} ended, loss: {logs.get('loss')}") x = np.random.rand(100, 4) y = np.random.rand(100, 1) model = tf.keras.Sequential([ tf.keras.layers.Dense(1, input_shape=(4,)) ]) model.compile(optimizer='sgd', loss='mse') model.fit(x, y, epochs=1, batch_size=32, callbacks=[BatchLogger()], verbose=0)

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

"batch 0 ended, loss: ..." "batch 1 ended, loss: ..." "batch 2 ended, loss: ..." "batch 3 ended, loss: ..."

Пример

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

import tensorflow as tf import numpy as np tf.random.set_seed(0) class LrReducer(tf.keras.callbacks.Callback): def on_batch_end(self, batch, logs=None): if batch == 2: lr = float(tf.keras.backend.get_value(self.model.optimizer.learning_rate)) tf.keras.backend.set_value(self.model.optimizer.learning_rate, lr * 0.5) print(f"learning rate reduced to {lr * 0.5}") x = np.random.rand(100, 4) y = np.random.rand(100, 1) model = tf.keras.Sequential([ tf.keras.layers.Dense(1, input_shape=(4,)) ]) model.compile(optimizer=tf.keras.optimizers.SGD(learning_rate=0.1), loss='mse') model.fit(x, y, epochs=1, batch_size=32, callbacks=[LrReducer()], verbose=0)

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

"learning rate reduced to 0.05"

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

  • класс Callback,
    который является базовым классом для всех колбэков
  • метод on_batch_begin,
    который вызывается в начале каждого батча
  • метод on_epoch_end,
    который вызывается в конце каждой эпохи
  • метод set_model,
    который устанавливает модель для колбэка
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить