Класс 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-файл