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

Класс TensorBoard

Класс TensorBoard создает колбэк, который сохраняет логи обучения модели для их последующей визуализации в TensorBoard. Первым параметром передается путь к каталогу для сохранения логов. Вторым параметром можно передать частоту записи гистограмм. Третьим параметром задается частота записи скалярных значений. Колбэк передается в метод fit при обучении модели.

Синтаксис

tf.keras.callbacks.TensorBoard(log_dir, [histogram_freq], [write_graph])

Пример

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

import tensorflow as tf tf.random.set_seed(0) model = tf.keras.Sequential([ tf.keras.layers.Dense(1, input_shape=(1,)) ]) model.compile(optimizer='sgd', loss='mse') tensorboard_cb = tf.keras.callbacks.TensorBoard(log_dir='logs') model.fit([1, 2, 3, 4, 5], [2, 4, 6, 8, 10], epochs=5, callbacks=[tensorboard_cb]) print("training finished")

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

"training finished"

Пример

Давайте зададим частоту записи гистограмм и скалярных значений через параметры histogram_freq и update_freq:

import tensorflow as tf tf.random.set_seed(0) model = tf.keras.Sequential([ tf.keras.layers.Dense(1, input_shape=(1,)) ]) model.compile(optimizer='sgd', loss='mse') tensorboard_cb = tf.keras.callbacks.TensorBoard( log_dir='logs', histogram_freq=1, update_freq='epoch' ) model.fit([1, 2, 3, 4, 5], [2, 4, 6, 8, 10], epochs=3, callbacks=[tensorboard_cb]) print("logs saved")

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

"logs saved"

Пример

Давайте отключим запись графа вычислений, передав write_graph=False:

<+python+> import tensorflow as tf tf.random.set_seed(0) model = tf.keras.Sequential([ tf.keras.layers.Dense(1, input_shape=(1,)) ]) model.compile(optimizer='sgd', loss='mse') tensorboard_cb = tf.keras.callbacks.TensorBoard( log_dir='logs', write_graph=False ) model.fit([1, 2, 3, 4, 5], [2, 4, 6, 8, 10], epochs=3, callbacks=[tensorboard_cb]) print("graph disabled") <-python+>

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

"graph disabled"

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

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