Spec-Zone.ru › TensorFlow 2.3

tf.estimator.Head

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

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

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

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

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

tf.compat.v1.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(...)

  feature_columns = tf.feature_column.numeric_column(...)
  feature_layer = tf.keras.layers.DenseFeatures(feature_columns)
  inputs = feature_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_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 и возвращает операцию оптимизации модели с потерей в режиме обучения. Используется, если optimizer имеет значение None. Ровно один из train_op_fn и optimizer должен быть установлен в режиме обучения. По умолчанию это 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 Список дополнительных скалярных потерь, которые нужно добавить к тренировочной и оценочной потерям, например, потери регуляризации. Обратите внимание, что аргумент mode не используется в tf.estimator.*Head. Если обновление метрик не зависит от mode, его можно безопасно игнорировать в сигнатуре метода.
Возвращает
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/r2.3/api_docs/python/tf/estimator/Head

Spec-Zone.ru

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