Spec-Zone.ru › TensorFlow 1.15

tf.estimator.Head

Просмотреть исходный код на GitHub

Интерфейс для головной части модели.

Просмотр псевдонимов

Псевдонимы совместимости для миграции

См. Руководство по миграции для получения дополнительной информации.

tf.compat.v1.estimator.Head, `tf.compat.v2.estimator.Head`

Головная часть расположена на вершине сети модели и обрабатывает вычисление выходных данных сети. Учитывая логиты (или выход скрытого слоя), головная часть знает, как вычислить предсказания, потерю, операцию обучения, метрики и выходные данные экспорта. Она предназначена для:

  1. Упрощения написания model_fn и повышения его конфигурируемости для Estimator.
  2. Упрощения создания потерь и метрик для циклов обучения и тестирования в режиме Eager execution.
  3. Поддержки широкого спектра моделей машинного обучения. Поскольку большинство головных частей могут работать с логитами, они могут поддерживать DNN, RNN, Wide, Wide&Deep, глобальные цели, деревья градиентного бустинга и многие другие типы моделей машинного обучения.

Примеры использования:

Вот упрощенный model_fn для построения модели регрессии DNN.

def _my_dnn_model_fn(features, labels, mode, params, config=None):
  # Optionally your callers can pass head to model_fn as a param.
  head = tf.estimator.RegressionHead(...)


  inputs = tf.feature_column.input_layer(features, ...)

  # Compute logits with tf.keras.layers API
  hidden_layer0 = tf.keras.layers.Dense(
      units=1000, activation="relu")(inputs)
  hidden_layer1 = tf.keras.layers.Dense(
      units=500, activation="relu")(hidden_layer0)
  logits = tf.keras.layers.Dense(
      units=head.logits_dimension, activation=None)(hidden_layer1)

  # Or use Keras model for logits computation
  model = tf.keras.Sequential()
  model.add(tf.keras.layers.Dense(units=1000, activation="relu"))
  model.add(tf.keras.layers.Dense(units=500, activation="relu"))
  model.add(tf.keras.layers.Dense(
     units=head.logits_dimension, activation=None))
  logits = model(inputs)

  return head.create_estimator_spec(
      features=features,
      labels=labels,
      mode=mode,
      logits=logits,
      optimizer=optimizer)
Атрибуты
logits_dimension Размер последнего измерения логитов Tensor.

Часто это количество классов, меток или действительных значений, которые необходимо предсказать. Как правило, логиты имеют форму [batch_size, logits_dimension].

loss_reduction Один из tf.losses.Reduction.

Описывает, как уменьшить потерю обучения по пакету, например, среднее или сумму.

name Название этой головной части.

Методы

create_estimator_spec

Просмотреть исходный код

create_estimator_spec(
    features, mode, logits, labels=None, optimizer=None, trainable_variables=None,
    train_op_fn=None, update_ops=None, regularization_losses=None
)

Возвращает EstimatorSpec , который может вернуть model_fn.

Рекомендуется передавать все аргументы через имя.

Аргументы
features Входные dict , сопоставляющие имена признаков строкового типа с Tensor или SparseTensor объектами, содержащими значения для этого признака в мини-пакете. Часто используется для извлечения тензора веса примера.
mode ModeKeys оценщика.
logits Логиты Tensor , которые будут использоваться головной частью.
labels Метки Tensor, или dict сопоставление имен меток строкового типа с Tensor объектами значений меток.
optimizer Экземпляр tf.keras.optimizers.Optimizer для оптимизации потери в режиме TRAIN. В частности, устанавливает train_op = optimizer.get_updates(loss, trainable_variables), который обновляет переменные для минимизации loss.
trainable_variables Список или кортеж объектов Variable для обновления, чтобы минимизировать loss. В Tensorflow 1.x по умолчанию это список переменных, собранных в графе под ключом GraphKeys.TRAINABLE_VARIABLES. Так как Tensorflow 2.x не имеет коллекций и GraphKeys, trainable_variables необходимо передавать явно здесь.
train_op_fn Функция, которая принимает скалярную потерю Tensor и возвращает операцию для оптимизации модели с потерей в режиме TRAIN. Используется, если optimizer равно None. Точно один из train_op_fn и optimizer должен быть установлен в режиме TRAIN. По умолчанию он равен None в других режимах. Если вы хотите оптимизировать потерю самостоятельно, вы можете передать lambda _: tf.no_op() и затем использовать EstimatorSpec.loss для вычисления и применения градиентов.
update_ops Список или кортеж операций обновления, которые должны выполняться во время обучения. Например, слои, такие как BatchNormalization, создают операции обновления среднего и дисперсии, которые необходимо выполнить во время обучения. В Tensorflow 1.x они помещаются в коллекцию UPDATE_OPS. Так как Tensorflow 2.x не имеет коллекций, update_ops необходимо передавать явно здесь.
regularization_losses Список дополнительных скалярных потерь, которые необходимо добавить к тренировочной потере, таких как потери регуляризации.
Возвращаемые значения
EstimatorSpec.

loss

Просмотреть исходный код

@abc.abstractmethod
loss(
    labels, logits, features=None, mode=None, regularization_losses=None
)

Возвращает потерю Tensor из предоставленных аргументов.

Обратите внимание, что аргументы features и mode скорее всего не используются, но некоторые реализации Head могут их потребовать.

Аргументы
labels Метки Tensor, или dict сопоставление имен меток строкового типа с Tensor объектами значений меток.
logits Логиты Tensor , используемые для построения потери.
features Входные dict , сопоставляющие имена признаков строкового типа с Tensor или SparseTensor объектами, содержащими значения для этого признака в мини-пакете. Часто используется для извлечения тензора веса примера.
mode ModeKeys оценщика. Используется в случае, если вычисление потери отличается в режимах обучения и оценки.
regularization_losses Список дополнительных скалярных потерь, которые необходимо добавить к тренировочной потере, таких как потери регуляризации.
Возвращаемые значения
Скалярная Tensor , представляющая регулярную тренировочную потерю, используемую в режимах обучения и оценки.

metrics

Просмотреть исходный код

@abc.abstractmethod
metrics(
    regularization_losses=None
)

Возвращает dict объектов метрик.

Аргументы
regularization_losses Список дополнительных скалярных потерь, которые необходимо добавить к тренировочной потере, таких как потери регуляризации.
Возвращаемые значения
dict метрик с ключами строкового типа. Значением является экземпляр класса Metric.

predictions

Просмотреть исходный код

@abc.abstractmethod
predictions(
    logits, keys=None
)

Возвращает dict предсказаний из предоставленных логитов.

Аргументы
logits Логиты Tensor , которые используются для построения предсказаний.
keys Список string для ключей предсказаний. По умолчанию None, что означает, что если не указано, предсказания будут созданы для всех предварительно определенных допустимых ключей в головной части.
Возвращаемые значения
dict предсказанных Tensor с ключами имени предсказания.

update_metrics

Просмотреть исходный код

@abc.abstractmethod
update_metrics(
    eval_metrics, features, logits, labels, mode=None, regularization_losses=None
)

Обновляет объекты метрик и возвращает dict обновленных метрик.

Аргументы
eval_metrics dict метрик, которые необходимо обновить.
features Входные dict , сопоставляющие имена признаков строкового типа с Tensor или SparseTensor объектами, содержащими значения для этого признака в мини-пакете. Часто используется для извлечения тензора веса примера.
logits Логиты Tensor , которые используются для обновления метрик.
labels Метки Tensor, или dict сопоставление имен меток строкового типа с Tensor объектами значений меток.
mode ModeKeys оценщика.
regularization_losses Список дополнительных скалярных потерь, которые необходимо добавить к потерям обучения и оценки, таких как потери регуляризации.
Возвращаемые значения
dict обновленных метрик с ключами имени. Значением является экземпляр класса Metric .

© 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/estimator/Head

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API