tf.estimator.Head
| Просмотреть исходный код на GitHub |
Интерфейс для головной части модели.
Головная часть расположена сверху сети модели и обрабатывает вычисление выходных данных сети. Учитывая логиты (или выход скрытого слоя), головная часть знает, как вычислить прогнозы, потерю, оператор обучения, метрики и выходные данные экспорта. Она предназначена для:
- Упрощения написания model_fn и повышения конфигурируемости model_fn для Estimator.
- Упрощения создания потерь и метрик для циклов обучения и тестирования в режиме Eager execution.
- Поддержки широкого спектра моделей машинного обучения. Поскольку большинство головных частей могут работать с логгитами, они могут поддерживать DNN, RNN, Wide, Wide&Deep, глобальные цели, деревья градиентного бустинга и многие другие типы моделей машинного обучения.
Общее использование:
Вот упрощенная модель_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. |
© 2022 The TensorFlow Authors. All rights reserved.
Licensed under the Creative Commons Attribution License 4.0.
Code samples licensed under the Apache 2.0 License.
https://www.tensorflow.org/versions/r2.9/api_docs/python/tf/estimator/Head