tf.estimator.Head
| Просмотреть исходный код на GitHub |
Интерфейс для головной части модели.
Головная часть расположена сверху модели и обрабатывает вычисление выходных данных сети. Учитывая логиты (или выходные данные скрытого слоя), головная часть знает, как вычислить прогнозы, потерю, операцию обучения, метрики и выходные данные экспорта. Она предназначена для:
- Упрощения написания функции
model_fnи повышения её конфигурируемости для Estimator. - Упрощения создания потерь и метрик для цикла обучения и тестирования в Eager execution.
- Поддержки широкого спектра моделей машинного обучения. Поскольку большинство головных частей могут работать с логитами, они могут поддерживать 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. Часто это количество классов, меток или вещественных значений, которые нужно спрогнозировать. Как правило, логиты имеют форму |
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