tf.contrib.learn.train
Обучить модель. (устарело)
tf.contrib.learn.train(
graph, output_dir, train_op, loss_op, global_step_tensor=None, init_op=None,
init_feed_dict=None, init_fn=None, log_every_steps=10, supervisor_is_chief=True,
supervisor_master='', supervisor_save_model_secs=600, keep_checkpoint_max=5,
supervisor_save_summaries_steps=100, feed_fn=None, steps=None,
fail_on_nan_loss=True, monitors=None, max_steps=None
)
Принимая graph, каталог для записи выходных данных (output_dir) и некоторые операции, запустить цикл обучения. Переданная train_op выполняет один шаг обучения модели. loss_op представляет целевую функцию обучения. Ожидается, что она будет увеличивать global_step_tensor, скалярный целочисленный тензор, считающий шаги обучения. Эта функция использует Supervisor для инициализации графа (из контрольной точки, если она доступна в output_dir), записи сводок, определённых в графе, и записи регулярных контрольных точек, как определено supervisor_save_model_secs.
Обучение продолжается до тех пор, пока global_step_tensor не примет значение max_steps, или, если fail_on_nan_loss, пока loss_op не примет значение NaN. В этом случае программа завершается с кодом выхода 1.
| Аргументы | |
|---|---|
graph | Граф для обучения. Ожидается, что этот граф не используется где-либо ещё. |
output_dir | Каталог для записи выходных данных. |
train_op | Операция, выполняющая один шаг обучения при запуске. |
loss_op | Скалярный тензор потерь. |
global_step_tensor | Тензор, представляющий глобальный шаг. Если он не задан, он извлекается из графа с использованием той же логики, что и в Supervisor. |
init_op | Операция, инициализирующая граф. Если None, используйте значения по умолчанию Supervisor. |
init_feed_dict | Словарь, сопоставляющий объекты Tensor со значениями подстановки. Этот словарь подстановки будет использоваться при вычислении init_op. |
init_fn | Необязательная функция, передаваемая в Supervisor для инициализации модели. |
log_every_steps | Регулярно выводит журналы. Журналы содержат данные о времени и текущую потерю. |
supervisor_is_chief | Является ли текущий процесс главным диспетчером, ответственным за восстановление модели и запуск стандартных служб. |
supervisor_master | Строка мастера, используемая при подготовке сессии. |
supervisor_save_model_secs | Сохранять контрольную точку каждые supervisor_save_model_secs секунд во время обучения. |
keep_checkpoint_max | Максимальное количество последних файлов контрольных точек для сохранения. По мере создания новых файлов старые файлы удаляются. Если None или 0, все файлы контрольных точек сохраняются. Это просто передаётся как параметр max_to_keep конструктору tf.compat.v1.train.Saver. |
supervisor_save_summaries_steps | Сохранять сводки каждые supervisor_save_summaries_steps секунд во время обучения. |
feed_fn | Функция, вызываемая на каждой итерации для создания feed_dict, передаваемого в вызовы session.run. Необязательно. |
steps | Обучать в течение этого количества шагов (например, текущий глобальный шаг + steps). |
fail_on_nan_loss | Если True, генерирует исключение NanLossDuringTrainingError, если loss_op принимает значение NaN. Если False, продолжить обучение как если бы ничего не произошло. |
monitors | Список экземпляров подклассов BaseMonitor. Используется для обратных вызовов внутри цикла обучения. |
max_steps | Общее количество шагов обучения модели. Если None, обучаться бесконечно. Два вызова fit(steps=100) означают 200 итераций обучения. С другой стороны, два вызова fit(max_steps=100) означают, что второй вызов не выполнит ни одной итерации, так как первый вызов выполнил все 100 шагов. |
| Возвращаемое значение | |
|---|---|
| Конечное значение потери. |
| Исключения | |
|---|---|
ValueError | Если output_dir, train_op, loss_op, или global_step_tensor не предоставлены. См. tf.contrib.framework.get_global_step, чтобы узнать, как мы ищем последнюю величину, если она не предоставлена явно. |
NanLossDuringTrainingError | Если fail_on_nan_loss равно True, и значение потерь когда-либо примет значение NaN. |
ValueError | Если и steps, и max_steps не являются None. |
© 2020 The TensorFlow Authors. All rights reserved.
Licensed under the Creative Commons Attribution License 3.0.
Code samples licensed under the Apache 2.0 License.
https://www.tensorflow.org/versions/r1.15/api_docs/python/tf/contrib/learn/train