Spec-Zone.ru › TensorFlow 2.4

tf.estimator.Head

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

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

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

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

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

tf.compat.v1.estimator.Head

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

  1. Упрощения написания функции model_fn и настройки функции 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. А именно, устанавливает 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. По умолчанию он используется в других режимах. Если вы хотите оптимизировать потерю самостоятельно, вы можете передать 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.4/api_docs/python/tf/estimator/Head

Spec-Zone.ru

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