Класс ProgbarLogger
Класс ProgbarLogger является колбэком,
который автоматически выводит в консоль
прогресс-бар во время обучения модели.
Он показывает номер эпохи, количество
обработанных шагов и значения метрик.
Первым параметром передается строка
count_mode, которая определяет,
что считать единицей прогресса: шаги
('steps') или примеры ('samples').
Вторым параметром передается режим
вывода verbose (0 или 1).
Дополнительно можно передать словарь
stateful_metrics с метриками,
которые не нужно усреднять по эпохе.
Синтаксис
tf.keras.callbacks.ProgbarLogger(count_mode='samples', [verbose=1], [stateful_metrics=None])
Пример
Давайте обучим простую модель и посмотрим,
как ProgbarLogger выводит прогресс
с режимом 'steps':
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')
x = tf.constant([[1.0], [2.0], [3.0], [4.0], [5.0]])
y = tf.constant([[1.0], [2.0], [3.0], [4.0], [5.0]])
logger = tf.keras.callbacks.ProgbarLogger(count_mode='steps')
model.fit(x, y, epochs=2, verbose=1, callbacks=[logger])
Результат выполнения кода:
Epoch 1/2
1/1 [==============================] - 0s 200ms/step - loss: 23.4567
Epoch 2/2
1/1 [==============================] - 0s 10ms/step - loss: 12.3456
<keras.src.callbacks.progbar_logger.ProgbarLogger object at 0x...>
Пример
Давайте используем stateful_metrics,
чтобы метрика не усреднялась по эпохе,
а выводилась как есть:
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', metrics=['mae'])
x = tf.constant([[1.0], [2.0], [3.0], [4.0], [5.0]])
y = tf.constant([[1.0], [2.0], [3.0], [4.0], [5.0]])
logger = tf.keras.callbacks.ProgbarLogger(
count_mode='samples',
stateful_metrics=['mae']
)
model.fit(x, y, epochs=2, verbose=1, callbacks=[logger])
Результат выполнения кода:
Смотрите также
-
класс
History,
который сохраняет историю обучения модели -
класс
CSVLogger,
который записывает логи обучения в CSV-файл -
класс
EarlyStopping,
который останавливает обучение при отсутствии улучшений -
класс
ModelCheckpoint,
который сохраняет модель во время обучения