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(...)
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. Часто это количество классов, меток или действительных значений, которые необходимо предсказать. Как правило, логиты имеют форму |
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