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

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

Метод on_train_begin класса Callback вызывается один раз в самом начале обучения модели, до обработки первой эпохи. Метод принимает один параметр logs - словарь, содержащий информацию о текущем состоянии обучения. Этот метод удобно использовать для инициализации переменных, сброса счетчиков, логирования времени старта или подготовки внешних ресурсов перед тренировкой.

Синтаксис

class MyCallback(tf.keras.callbacks.Callback): def on_train_begin(self, logs=None): pass

Пример

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

import tensorflow as tf tf.random.set_seed(0) class StartCallback(tf.keras.callbacks.Callback): def on_train_begin(self, logs=None): print("Training started") 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], [5.0]]) y = tf.constant([[2.0], [4.0], [6.0], [8.0], [10.0]]) model.fit(x, y, epochs=2, callbacks=[StartCallback()], verbose=0)

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

Training started

Пример

Давайте используем параметр logs для вывода доступной информации в начале обучения:

import tensorflow as tf tf.random.set_seed(0) class LogsCallback(tf.keras.callbacks.Callback): def on_train_begin(self, logs=None): print("logs:", logs) print("params:", self.params) 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], [5.0]]) y = tf.constant([[2.0], [4.0], [6.0], [8.0], [10.0]]) model.fit(x, y, epochs=2, callbacks=[LogsCallback()], verbose=0)

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

logs: {} params: {'verbose': 0, 'epochs': 2, 'steps': 5}

Пример

Давайте применим on_train_begin для сброса счетчика эпох перед обучением:

import tensorflow as tf tf.random.set_seed(0) class CounterCallback(tf.keras.callbacks.Callback): def __init__(self): super().__init__() self.epoch_count = 0 def on_train_begin(self, logs=None): self.epoch_count = 0 print("Counter reset to", self.epoch_count) def on_epoch_end(self, epoch, logs=None): self.epoch_count += 1 print("Epoch count:", self.epoch_count) 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], [5.0]]) y = tf.constant([[2.0], [4.0], [6.0], [8.0], [10.0]]) model.fit(x, y, epochs=3, callbacks=[CounterCallback()], verbose=0)

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

Counter reset to 0 Epoch count: 1 Epoch count: 2 Epoch count: 3

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

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