tf.compat.v1.estimator.tpu.TPUEstimator
Эстиматор с поддержкой TPU.
Наследуется от: Estimator
tf.compat.v1.estimator.tpu.TPUEstimator(
model_fn=None,
model_dir=None,
config=None,
params=None,
use_tpu=True,
train_batch_size=None,
eval_batch_size=None,
predict_batch_size=None,
batch_axis=None,
eval_on_tpu=True,
export_to_tpu=True,
export_to_cpu=True,
warm_start_from=None,
embedding_config_spec=None,
export_saved_model_api_version=ExportSavedModelApiVersion.V1
)
Перейти к TF2
TPU Estimator управляет собственной TensorFlow-графикой и сессией, поэтому он несовместим с поведением TF2. Мы рекомендуем перейти к новому tf.distribute.TPUStrategy. Подробнее см. в руководстве по TPU.
Описание
TPUEstimator также поддерживает обучение на CPU и GPU. Вам не нужно определять отдельный tf.estimator.Estimator.
TPUEstimator обрабатывает многие детали выполнения на устройствах TPU, такие как дублирование входных данных и моделей для каждого ядра и периодическое возвращение на хост для выполнения хуков.
TPUEstimator преобразует глобальный размер пакета в params в размер пакета на фрагмент при вызове input_fn и model_fn. Пользователи должны указать глобальный размер пакета в конструкторе, а затем получить размер пакета для каждого фрагмента в input_fn и model_fn с помощью params['batch_size'].
Для обучения
model_fnполучает размер пакета на ядро;input_fnможет получить размер пакета на ядро или на хост в зависимости отper_host_input_for_trainingвTPUConfig(подробнее см. в документации TPUConfig).Для оценки и предсказания
model_fnполучает размер пакета на ядро, аinput_fnполучает размер пакета на хост.
Оценка
model_fn должен возвращать TPUEstimatorSpec, который ожидает eval_metrics для оценки на TPU. Если eval_on_tpu равно False, оценка выполняется на CPU или GPU; в этом случае обсуждение оценки на TPU не применяется.
TPUEstimatorSpec.eval_metrics — это кортеж из metric_fn и tensors, где tensors может быть списком любых вложенных структур Tensor (подробнее см. TPUEstimatorSpec). metric_fn принимает tensors и возвращает словарь из имени метки метрики и результата вызова функции метрики, а именно кортеж (metric_tensor, update_op).
Можно установить use_tpu в False для тестирования. Все обучение, оценка и предсказание будут выполняться на CPU. input_fn и model_fn получат train_batch_size или eval_batch_size без изменений в виде params['batch_size'].
Текущие ограничения:
Оценка на TPU работает только на одном хосте (одном TPU-рабочем узле), за исключением режима BROADCAST.
input_fnдля оценки НЕ должно вызывать исключение конца входных данных (OutOfRangeErrorилиStopIteration). И все шаги оценки и все пакеты должны иметь одинаковый размер.
Пример (MNIST):
# The metric Fn which runs on CPU.
def metric_fn(labels, logits):
predictions = tf.argmax(logits, 1)
return {
'accuracy': tf.compat.v1.metrics.precision(
labels=labels, predictions=predictions),
}
# Your model Fn which runs on TPU (eval_metrics is list in this example)
def model_fn(features, labels, mode, config, params):
...
logits = ...
if mode = tf.estimator.ModeKeys.EVAL:
return tpu_estimator.TPUEstimatorSpec(
mode=mode,
loss=loss,
eval_metrics=(metric_fn, [labels, logits]))
# or specify the eval_metrics tensors as dict.
def model_fn(features, labels, mode, config, params):
...
final_layer_output = ...
if mode = tf.estimator.ModeKeys.EVAL:
return tpu_estimator.TPUEstimatorSpec(
mode=mode,
loss=loss,
eval_metrics=(metric_fn, {
'labels': labels,
'logits': final_layer_output,
}))
Предсказание
Предсказание на TPU — это экспериментальная функция для поддержки предсказания с большими пакетами. Она не предназначена для систем с критическим значением задержки. Кроме того, из-за проблем с удобством использования для предсказания с небольшим набором данных использование CPU .predict, то есть создание нового экземпляра TPUEstimator с use_tpu=False, может быть удобнее.
Примечание: В отличие от обучения/оценки на TPU,input_fnдля предсказания должно вызывать исключение конца входных данных (OutOfRangeErrorилиStopIteration), что служит сигналом остановки дляTPUEstimator. Точнее, операции, созданныеinput_fn, создают один пакет данных. APIpredict()обрабатывает один пакет за раз. При достижении конца источника данных должно быть вызвано исключение конца входных данных одной из этих операций. Обычно пользователю не нужно делать это вручную. До тех пор, пока набор данных не повторяется бесконечно, APItf.dataбудет автоматически вызывать исключение конца входных данных после создания последнего пакета.
Примечание: Estimator.predict возвращает Python-генератор. Пожалуйста, обработайте все данные из генератора, чтобы TPUEstimator мог должным образом остановить систему TPU для пользователя.
Текущие ограничения:
Предсказание на TPU работает только на одном хосте (одном TPU-рабочем узле).
input_fnдолжен возвращать экземплярDatasetвместоfeatures. Фактически, .train() и .evaluate() также поддерживают Dataset в качестве возвращаемого значения.
Пример (MNIST):
height = 32
width = 32
total_examples = 100
def predict_input_fn(params):
batch_size = params['batch_size']
images = tf.random.uniform(
[total_examples, height, width, 3], minval=-1, maxval=1)
dataset = tf.data.Dataset.from_tensor_slices(images)
dataset = dataset.map(lambda images: {'image': images})
dataset = dataset.batch(batch_size)
return dataset
def model_fn(features, labels, params, mode):
# Generate predictions, called 'output', from features['image']
if mode == tf.estimator.ModeKeys.PREDICT:
return tf.contrib.tpu.TPUEstimatorSpec(
mode=mode,
predictions={
'predictions': output,
'is_padding': features['is_padding']
})
tpu_est = TPUEstimator(
model_fn=model_fn,
...,
predict_batch_size=16)
# Fully consume the generator so that TPUEstimator can shutdown the TPU
# system.
for item in tpu_est.predict(input_fn=input_fn):
# Filter out item if the `is_padding` is 1.
# Process the 'predictions'
Экспорт
export_saved_model экспортирует 2 метаграфа, один с saved_model.SERVING, а другой с saved_model.SERVING и saved_model.TPU тегами. Во время предоставления эти теги используются для выбора подходящего метаграфа для загрузки.
Перед запуском графика на TPU необходимо инициализировать систему TPU. Если используется сервер TensorFlow Serving model-server, это происходит автоматически. В противном случае используйте session.run(tpu.initialize_system()).
Существуют две версии API: 1 или 2.
В версии 1 экспортированный CPU-график экспортируется без изменений. Экспортированный TPU-график оборачивает tpu.rewrite() и TPUPartitionedCallOp вокруг model_fn, поэтому model_fn по умолчанию выполняется на TPU. Чтобы разместить операции на CPU, можно использовать tpu.outside_compilation(host_call, logits).
Пример:
def model_fn(features, labels, mode, config, params):
...
logits = ...
export_outputs = {
'logits': export_output_lib.PredictOutput(
{'logits': logits})
}
def host_call(logits):
class_ids = math_ops.argmax(logits)
classes = string_ops.as_string(class_ids)
export_outputs['classes'] =
export_output_lib.ClassificationOutput(classes=classes)
tpu.outside_compilation(host_call, logits)
...
В версии 2 export_saved_model() устанавливает флаг params['use_tpu'], чтобы сообщить пользователю, экспортируется ли код на TPU (или нет). Когда params['use_tpu'] равно True, пользователям нужно вызвать tpu.rewrite(), TPUPartitionedCallOp и/или batch_function().
СОВЕТ: версия 2 рекомендуется, так как она более гибкая (например, пакетные операции и т. д.).
END_OF_DOCUMENT_MARKER| Аргументы | |
|---|---|
model_fn | Функция модели, как требуется Estimator, которая возвращает EstimatorSpec или TPUEstimatorSpec. training_hooks, 'evaluation_hooks', и prediction_hooks не должны захватывать какие-либо тензоры TPU внутри функции model_fn. |
model_dir | Директория для сохранения параметров модели, графа и т. д. Это также можно использовать для загрузки контрольных точек из директории в оценщик для продолжения обучения ранее сохраненной модели. Если None, будет использоваться model_dir в config, если он задан. Если оба заданы, они должны быть одинаковыми. Если оба None, будет использована временная директория. |
config | Объект конфигурации tpu_config.RunConfig. Не может быть None. |
params | Необязательный dict гиперпараметров, которые будут переданы в input_fn и model_fn. Ключи — имена параметров, значения — базовые типы Python. Существуют зарезервированные ключи для TPUEstimator, включая 'batch_size'. |
use_tpu | Булево значение, указывающее, включена ли поддержка TPU. В настоящее время - обучение и оценка TPU учитывают этот бит, но eval_on_tpu может переопределить выполнение оценки. См. ниже. |
train_batch_size | Целое число, представляющее глобальный размер пакета для обучения. TPUEstimator преобразует этот глобальный размер пакета в размер пакета на фрагмент как params['batch_size'] при вызове input_fn и model_fn. Не может быть None, если use_tpu равно True. Должно быть кратно общему числу реплик. |
eval_batch_size | Целое число, представляющее размер пакета для оценки. Должно быть кратно общему числу реплик. |
predict_batch_size | Целое число, представляющее размер пакета для предсказания. Должно быть кратно общему числу реплик. |
batch_axis | Кортеж Python из целых значений, описывающих, как каждый тензор, созданный оценщиком input_fn, должен быть разделен между фрагментами вычисления TPU. Например, если ваша функция input_fn возвращает (images, labels), где тензор images находится в формате HWCN, ваши размеры фрагментов будут [3, 0], где 3 соответствует размеру измерения N вашего тензора images, а 0 соответствует размеру измерения, по которому следует разделить метки, чтобы они соответствовали соответствующим изображениям. Если None, и per_host_input_for_training равно True, пакеты будут разделены по основному измерению. Если tpu_config.per_host_input_for_training равно False или PER_HOST_V2, ось batch_axis игнорируется. |
eval_on_tpu | Если False, оценка выполняется на CPU или GPU. В этом случае функция model_fn должна возвращать EstimatorSpec при вызове с mode как EVAL. |
export_to_tpu | Если True, export_saved_model() экспортирует метаграф для обслуживания на TPU. Обратите внимание, что неподдерживаемые режимы экспорта, такие как EVAL, будут проигнорированы. Для этих режимов будет экспортирована только модель CPU. В настоящее время export_to_tpu поддерживает только PREDICT. |
export_to_cpu | Если True, export_saved_model() экспортирует метаграф для обслуживания на CPU. |
warm_start_from | Необязательный строковый путь к файлу контрольной точки или SavedModel для начальной загрузки, или объект tf.estimator.WarmStartSettings, чтобы полностью настроить начальную загрузку. Если указан строковый путь к файлу вместо объекта WarmStartSettings, то все переменные будут загружены, и предполагается, что словари и имена тензоров не изменены. |
embedding_config_spec | Необязательный экземпляр EmbeddingConfigSpec для поддержки использования TPU embedding. |
export_saved_model_api_version | целое число: 1 или 2. 1 соответствует V1, 2 соответствует V2. (По умолчанию V1). При V1, export_saved_model() добавляет rewrite() и TPUPartitionedCallOp() для пользователя; в v2 ожидается, что пользователь добавит rewrite(), TPUPartitionedCallOp() и т. д. в свою функцию model_fn. |
| Исключения | |
|---|---|
ValueError | params уже имеет зарезервированные ключи. |
| Атрибуты | |
|---|---|
config | |
model_dir | |
model_fn | Возвращает model_fn, связанный с self.params. |
params | |
Методы
eval_dir
eval_dir(
name=None
)
Показывает имя директории, куда выгружаются метрики оценки.
| Аргументы | |
|---|---|
name | Название оценки, если пользователю необходимо выполнить несколько оценок на разных наборах данных, например, на обучающих и тестовых данных. Метрики для различных оценок сохраняются в отдельных папках и отображаются отдельно в TensorBoard. |
| Возвращаемое значение | |
|---|---|
| Строка, представляющая путь к директории, содержащей метрики оценки. |
evaluate
evaluate(
input_fn, steps=None, hooks=None, checkpoint_path=None, name=None
)
Оценивает модель, используя данные оценки input_fn.
На каждом шаге вызывается input_fn, которая возвращает одну порцию данных. Оценка выполняется до тех пор, пока:
-
stepsпакетов обработано, или -
input_fnвызывает исключение конца ввода (tf.errors.OutOfRangeErrorилиStopIteration).
| Аргументы | |
|---|---|
input_fn | Функция, которая создает входные данные для оценки. См. Предопределенные оценщики для получения дополнительной информации. Функция должна создать и вернуть один из следующих элементов:
|
steps | Количество шагов для оценки модели. Если None, оценивает до тех пор, пока input_fn не вызовет исключение конца ввода. |
hooks | Список экземпляров подклассов tf.train.SessionRunHook. Используется для обратных вызовов внутри вызова оценки. |
checkpoint_path | Путь к конкретной контрольной точке для оценки. Если None, используется последняя контрольная точка в model_dir. Если контрольных точек в model_dir нет, оценка выполняется с новыми инициализированными Variables вместо восстановленных из контрольной точки. |
name | Название оценки, если пользователю необходимо выполнить несколько оценок на разных наборах данных, например, на обучающих и тестовых данных. Метрики для различных оценок сохраняются в отдельных папках и отображаются отдельно в TensorBoard. |
| Возвращаемое значение | |
|---|---|
Словарь, содержащий метрики оценки, указанные в model_fn с ключами по имени, а также запись global_step, содержащую значение глобального шага, для которого была выполнена эта оценка. Для готовых оценщиков словарь содержит loss (средняя потеря на мини-пакет) и average_loss (средняя потеря на образец). Готовые классификаторы также возвращают accuracy. Готовые регрессоры также возвращают label/mean и prediction/mean. |
| Исключения | |
|---|---|
ValueError | Если steps <= 0. |
experimental_export_all_saved_models
experimental_export_all_saved_models(
export_dir_base,
input_receiver_fn_map,
assets_extra=None,
as_text=False,
checkpoint_path=None
)
Экспортирует SavedModel с tf.MetaGraphDefs для каждого запрошенного режима.
Для каждого режима, переданного через input_receiver_fn_map, этот метод создает новую графу, вызывая input_receiver_fn для получения признаков и меток Tensors. Далее, этот метод вызывает Estimator's model_fn в переданном режиме, чтобы сгенерировать граф модели на основе этих признаков и меток, и восстанавливает заданную контрольную точку (или, если её нет, последнюю контрольную точку) в граф. Только один из режимов используется для сохранения переменных в SavedModel (порядок предпочтения: tf.estimator.ModeKeys.TRAIN, tf.estimator.ModeKeys.EVAL, затем tf.estimator.ModeKeys.PREDICT), таким образом, до трёх tf.MetaGraphDefs сохраняются с одной группой переменных в одной директории экспорта SavedModel.
Для переменных и tf.MetaGraphDefs, создается временная директория экспорта ниже export_dir_base, и записывается SavedModel в неё, содержащий tf.MetaGraphDef для данного режима и связанных с ним сигнатур.
Для предсказания экспортированный MetaGraphDef предоставит по одному SignatureDef для каждого элемента словаря export_outputs, возвращенного из model_fn, с именами, использующими те же ключи. Один из этих ключей всегда tf.saved_model.signature_constants.DEFAULT_SERVING_SIGNATURE_DEF_KEY, указывая, какая сигнатура будет использоваться, когда запрос на обслуживание не указывает её. Для каждой сигнатуры выходы предоставляются соответствующими tf.estimator.export.ExportOutput, а входные данные всегда — приемники входных данных, предоставляемые serving_input_receiver_fn.
Для обучения и оценки train_op хранится в дополнительном наборе, а потери, метрики и предсказания включаются в SignatureDef для соответствующего режима.
Дополнительные ресурсы могут быть записаны в SavedModel через аргумент assets_extra. Это должен быть словарь, где каждый ключ указывает путь назначения (включая имя файла) относительно каталога assets.extra. Соответствующее значение указывает полный путь исходного файла, который нужно скопировать. Например, простой случай копирования одного файла без переименования задаётся как {'my_asset_file.txt': '/path/to/my_asset_file.txt'}.
| Аргументы | |
|---|---|
export_dir_base | Строка, содержащая каталог, в котором будут создаваться подкаталоги с отметкой времени, содержащие экспортированные SavedModels. |
input_receiver_fn_map | Словарь соответствий между ключами tf.estimator.ModeKeys и input_receiver_fn отображения, где input_receiver_fn — функция, которая не принимает аргументов и возвращает соответствующий подкласс InputReceiver. |
assets_extra | Словарь, определяющий, как заполнить каталог assets.extra внутри экспортированного SavedModel, или None если дополнительные ресурсы не нужны. |
as_text | нужно ли записывать прото SavedModel в текстовом формате. |
checkpoint_path | Путь к контрольной точке для экспорта. Если None (по умолчанию), выбирается самая последняя контрольная точка, найденная в каталоге модели. |
| Возвращаемые значения | |
|---|---|
| Путь к экспортированному каталогу в виде объекта типа bytes. |
| Исключения | |
|---|---|
ValueError | если какой-либо input_receiver_fn является None, не указаны export_outputs или не найдена контрольная точка. |
export_saved_model
export_saved_model(
export_dir_base,
serving_input_receiver_fn,
assets_extra=None,
as_text=False,
checkpoint_path=None,
experimental_mode=ModeKeys.PREDICT
)
Экспортирует граф вывода как SavedModel в указанный каталог.
Для подробного руководства по SavedModel см. Использование формата SavedModel.
Этот метод строит новую графу, сначала вызывая serving_input_receiver_fn для получения признаков Tensor, а затем вызывая Estimator этого model_fn для генерации графа модели на основе этих признаков. Он восстанавливает заданную контрольную точку (или, если таковой нет, последнюю контрольную точку) в этой графе в новой сессии. Наконец, он создаёт каталог экспорта с отметкой времени ниже заданного export_dir_base, и записывает в него SavedModel , содержащий одну tf.MetaGraphDef сохранённую из этой сессии.
Экспортированный MetaGraphDef предоставит по одному SignatureDef для каждого элемента словаря export_outputs, возвращенного из model_fn, с именами, использующими те же ключи. Один из этих ключей всегда tf.saved_model.signature_constants.DEFAULT_SERVING_SIGNATURE_DEF_KEY, указывая, какая сигнатура будет использоваться, когда запрос на обслуживание не указывает её. Для каждой сигнатуры выходы предоставляются соответствующими tf.estimator.export.ExportOutput, а входные данные всегда — приемники входных данных, предоставляемые serving_input_receiver_fn.
Дополнительные ресурсы могут быть записаны в SavedModel через аргумент assets_extra. Это должен быть словарь, где каждый ключ указывает путь назначения (включая имя файла) относительно каталога assets.extra. Соответствующее значение указывает полный путь исходного файла, который нужно скопировать. Например, простой случай копирования одного файла без переименования задаётся как {'my_asset_file.txt': '/path/to/my_asset_file.txt'}.
Параметр experimental_mode может использоваться для экспорта отдельной тренировочной/оценочной/предсказательной графы в виде SavedModel. См. experimental_export_all_saved_models для полной документации.
| Аргументы | |
|---|---|
export_dir_base | Строка, содержащая каталог, в котором будут создаваться подкаталоги с отметкой времени, содержащие экспортированные SavedModels. |
serving_input_receiver_fn | Функция, которая не принимает аргументов и возвращает tf.estimator.export.ServingInputReceiver или tf.estimator.export.TensorServingInputReceiver. |
assets_extra | Словарь, определяющий, как заполнить каталог assets.extra внутри экспортированного SavedModel, или None если дополнительные ресурсы не нужны. |
as_text | нужно ли записывать прото SavedModel в текстовом формате. |
checkpoint_path | Путь к контрольной точке для экспорта. Если None (по умолчанию), выбирается самая последняя контрольная точка, найденная в каталоге модели. |
experimental_mode | Значение tf.estimator.ModeKeys , указывающее, какой режим будет экспортирован. Обратите внимание, что эта функция экспериментальна. |
| Возвращаемые значения | |
|---|---|
| Путь к экспортированному каталогу в виде объекта типа bytes. |
| Исключения | |
|---|---|
ValueError | если не указан serving_input_receiver_fn, не указаны export_outputs или не найдена контрольная точка. |
export_savedmodel
export_savedmodel(
export_dir_base,
serving_input_receiver_fn,
assets_extra=None,
as_text=False,
checkpoint_path=None,
strip_default_attrs=False
)
Экспортирует граф вывода как SavedModel в указанный каталог. (устарело)
Для подробного руководства см. SavedModel из Estimators..
Этот метод строит новую графу, сначала вызывая serving_input_receiver_fn для получения признаков Tensor, а затем вызывая Estimator этого model_fn для генерации графа модели на основе этих признаков. Он восстанавливает заданную контрольную точку (или, если таковой нет, последнюю контрольную точку) в этой графе в новой сессии. Наконец, он создаёт каталог экспорта с отметкой времени ниже заданного export_dir_base, и записывает в него SavedModel , содержащий одну tf.MetaGraphDef сохранённую из этой сессии.
Экспортированный MetaGraphDef предоставит по одному SignatureDef для каждого элемента словаря export_outputs, возвращенного из model_fn, с именами, использующими те же ключи. Один из этих ключей всегда tf.saved_model.signature_constants.DEFAULT_SERVING_SIGNATURE_DEF_KEY, указывая, какая сигнатура будет использоваться, когда запрос на обслуживание не указывает её. Для каждой сигнатуры выходы предоставляются соответствующими tf.estimator.export.ExportOutput, а входные данные всегда — приемники входных данных, предоставляемые serving_input_receiver_fn.
Дополнительные ресурсы могут быть записаны в SavedModel через аргумент assets_extra. Это должен быть словарь, где каждый ключ указывает путь назначения (включая имя файла) относительно каталога assets.extra. Соответствующее значение указывает полный путь исходного файла, который нужно скопировать. Например, простой случай копирования одного файла без переименования задаётся как {'my_asset_file.txt': '/path/to/my_asset_file.txt'}.
| Аргументы | |
|---|---|
export_dir_base | Строка, содержащая каталог, в котором будут создаваться подкаталоги с отметкой времени, содержащие экспортированные SavedModels. |
serving_input_receiver_fn | Функция, которая не принимает аргументов и возвращает tf.estimator.export.ServingInputReceiver или tf.estimator.export.TensorServingInputReceiver. |
assets_extra | Словарь, определяющий, как заполнить каталог assets.extra внутри экспортированного SavedModel, или None если дополнительные ресурсы не нужны. |
as_text | нужно ли записывать прото SavedModel в текстовом формате. |
checkpoint_path | Путь к контрольной точке для экспорта. Если None (по умолчанию), выбирается самая последняя контрольная точка, найденная в каталоге модели. |
strip_default_attrs | Булево значение. Если True, атрибуты с значениями по умолчанию будут удалены из NodeDefs. Для подробного руководства см. Удаление атрибутов с значениями по умолчанию. |
| Возвращаемые значения | |
|---|---|
| Путь к экспортированному каталогу в виде объекта типа bytes. |
| Исключения | |
|---|---|
ValueError | если не указан serving_input_receiver_fn, не указаны export_outputs или не найдена контрольная точка. |
get_variable_names
get_variable_names()
Возвращает список всех имён переменных в этой модели.
| Возвращает | |
|---|---|
| Список имён. |
| Возбуждает | |
|---|---|
ValueError | Если Estimator ещё не создал контрольную точку. |
get_variable_value
get_variable_value(
name
)
Возвращает значение переменной, заданной именем.
| Аргументы | |
|---|---|
name | Строка или список строк, имя тензора. |
| Возвращает | |
|---|---|
| Массив NumPy - значение тензора. |
| Возбуждает | |
|---|---|
ValueError | Если Estimator ещё не создал контрольную точку. |
latest_checkpoint
latest_checkpoint()
Находит имя файла последней сохранённой контрольной точки в model_dir.
| Возвращает | |
|---|---|
Полный путь к последней контрольной точке или None , если контрольная точка не найдена. |
predict
predict(
input_fn,
predict_keys=None,
hooks=None,
checkpoint_path=None,
yield_single_examples=True
)
Возвращает прогнозы для заданных признаков.
Обратите внимание, что чередование двух выходов predict не работает. См.: issue/20506
| Аргументы | |
|---|---|
input_fn | Функция, которая строит признаки. Прогнозирование продолжается до тех пор, пока input_fn не сгенерирует исключение конца ввода (tf.errors.OutOfRangeError или StopIteration). См. Premade Estimators для получения дополнительной информации. Функция должна создать и вернуть одно из следующих:
|
predict_keys | Список строк, имена ключей для прогнозирования. Используется, если tf.estimator.EstimatorSpec.predictions является dict. Если predict_keys используется, то остальные прогнозы из словаря будут отфильтрованы. Если None, возвращает все. |
hooks | Список экземпляров подкласса tf.train.SessionRunHook. Используется для обратных вызовов внутри вызова прогнозирования. |
checkpoint_path | Путь к определенной контрольной точке для прогнозирования. Если None, используется последняя контрольная точка в model_dir. Если контрольных точек в model_dir нет, прогнозирование выполняется с вновь инициализированными Variables вместо восстановленных из контрольной точки. |
yield_single_examples | Если False, возвращает всю партию, как она возвращается model_fn вместо разделения партии на отдельные элементы. Это полезно, если model_fn возвращает некоторые тензоры, первая размерность которых не равна размеру партии. |
| Возвращает | |
|---|---|
Вычисленные значения тензоров predictions . |
| Возбуждает | |
|---|---|
ValueError | Если длина партии прогнозов не одинакова, и yield_single_examples равно True. |
ValueError | Если есть конфликт между predict_keys и predictions. Например, если predict_keys не None , но tf.estimator.EstimatorSpec.predictions не является dict. |
train
train(
input_fn, hooks=None, steps=None, max_steps=None, saving_listeners=None
)
Обучает модель с использованием обучающих данных input_fn.
| Аргументы | |
|---|---|
input_fn | Функция, которая предоставляет обучающие данные в виде мини-пакетов. См. Premade Estimators для получения дополнительной информации. Функция должна создать и вернуть одно из следующих:
|
hooks | Список экземпляров подкласса tf.train.SessionRunHook. Используется для обратных вызовов внутри цикла обучения. |
steps | Количество шагов для обучения модели. Если None, обучаться бесконечно или до тех пор, пока input_fn не сгенерирует ошибку tf.errors.OutOfRange или исключение StopIteration. steps работает инкрементно. Если вызвать train(steps=10) дважды, обучение произойдёт в общей сложности 20 шагов. Если OutOfRange или StopIteration произойдёт в середине, обучение остановится до 20 шагов. Если не нужно инкрементное поведение, установите max_steps. Если установлено, max_steps должно быть None. |
max_steps | Общее число шагов для обучения модели. Если None, обучаться бесконечно или до тех пор, пока input_fn не сгенерирует ошибку tf.errors.OutOfRange или исключение StopIteration. Если установлено, steps должно быть None. Если OutOfRange или StopIteration произойдёт в середине, обучение остановится до max_steps шагов. Два вызова train(steps=100) означают 200 итераций обучения. С другой стороны, два вызова train(max_steps=100) означают, что второй вызов не выполнит итераций, так как первый вызов выполнил все 100 шагов. |
saving_listeners | Список объектов CheckpointSaverListener. Используется для обратных вызовов, выполняемых непосредственно перед или после сохранения контрольной точки. |
| Возвращает | |
|---|---|
self, для цепочки вызовов. |
| Возбуждает | |
|---|---|
ValueError | Если оба steps и max_steps не None. |
ValueError | Если либо steps , либо max_steps <= 0. |
© 2022 The TensorFlow Authors. All rights reserved.
Licensed under the Creative Commons Attribution License 4.0.
Code samples licensed under the Apache 2.0 License.
https://www.tensorflow.org/versions/r2.9/api_docs/python/tf/compat/v1/estimator/tpu/TPUEstimator