Spec-Zone.ru › TensorFlow 2.9

tf.estimator.RegressionHead

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

Создаёт Head для регрессии с использованием mean_squared_error функции потерь.

Наследуется от: Head

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

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

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

tf.compat.v1.estimator.RegressionHead

tf.estimator.RegressionHead(
    label_dimension=1,
    weight_column=None,
    loss_reduction=tf.losses.Reduction.SUM_OVER_BATCH_SIZE,
    loss_fn=None,
    inverse_link_fn=None,
    name=None
)

Функция потерь — это взвешенная сумма по всем измерениям входных данных. Иными словами, если входные метки имеют форму [batch_size, label_dimension], функция потерь — это взвешенная сумма по batch_size и label_dimension.

Функция ожидает logits с формой [D0, D1, ... DN, label_dimension]. Во многих приложениях форма имеет вид [batch_size, label_dimension].

Форма labels должна соответствовать форме logits, а именно [D0, D1, ... DN, label_dimension]. Если label_dimension=1, то также поддерживается форма [D0, D1, ... DN].

Если weight_column указано, веса должны иметь форму [D0, D1, ... DN], [D0, D1, ... DN, 1] или [D0, D1, ... DN, label_dimension].

Поддерживает пользовательские loss_fn. loss_fn принимает (labels, logits) или (labels, logits, features, loss_reduction) в качестве аргументов и возвращает нередуцированную функцию потерь с формой [D0, D1, ... DN, label_dimension].

Также поддерживает пользовательские inverse_link_fn, также известные как функция «среднего значения». inverse_link_fn используется только в режиме PREDICT. Она принимает logits в качестве аргумента и возвращает предсказанные значения. Эта функция является обратной связью к функции связи, определённой в https://en.wikipedia.org/wiki/Generalized_linear_model#Link_function. Например, для регрессии Пуассона установите inverse_link_fn=tf.exp.

Использование:

head = tf.estimator.RegressionHead()
logits = np.array(((45,), (41,),), dtype=np.float32)
labels = np.array(((43,), (44,),), dtype=np.int32)
features = {'x': np.array(((42,),), dtype=np.float32)}
# expected_loss = weighted_loss / batch_size
#               = (43-45)^2 + (44-41)^2 / 2 = 6.50
loss = head.loss(labels, logits, features=features)
print('{:.2f}'.format(loss.numpy()))
6.50
eval_metrics = head.metrics()
updated_metrics = head.update_metrics(
  eval_metrics, features, logits, labels)
for k in sorted(updated_metrics):
 print('{} : {:.2f}'.format(k, updated_metrics[k].result().numpy()))
  average_loss : 6.50
  label/mean : 43.50
  prediction/mean : 43.00
preds = head.predictions(logits)
print(preds['predictions'])
tf.Tensor(
  [[45.]
   [41.]], shape=(2, 1), dtype=float32)

Использование с готовым оценщиком:

my_head = tf.estimator.RegressionHead()
my_estimator = tf.estimator.DNNEstimator(
    head=my_head,
    hidden_units=...,
    feature_columns=...)

Также может быть использовано с пользовательским model_fn. Пример:

def _my_model_fn(features, labels, mode):
  my_head = tf.estimator.RegressionHead()
  logits = tf.keras.Model(...)(features)

  return my_head.create_estimator_spec(
      features=features,
      mode=mode,
      labels=labels,
      optimizer=tf.keras.optimizers.Adagrad(lr=0.1),
      logits=logits)

my_estimator = tf.estimator.Estimator(model_fn=_my_model_fn)
Аргументы
weight_column Строка или NumericColumn, созданная с помощью tf.feature_column.numeric_column, определяющая столбец признаков, представляющий веса. Он используется для уменьшения или увеличения примеров во время обучения. Он будет умножаться на функцию потерь примера.
label_dimension Количество меток регрессии на пример. Это размер последнего измерения меток Tensor (обычно, эта форма [batch_size, label_dimension]).
loss_reduction Один из tf.losses.Reduction за исключением NONE. Определяет, как уменьшить функцию потерь обучения по пакету и измерению меток. По умолчанию SUM_OVER_BATCH_SIZE, а именно взвешенная сумма потерь, делённая на batch_size * label_dimension.
loss_fn Необязательная функция потерь. По умолчанию mean_squared_error.
inverse_link_fn Необязательная обратная функция связи, также известная как функция «среднего значения». По умолчанию — идентичность.
name Имя головки. Если указано, ключи сводки и метрик будут дополнены суффиксом "/" + name. Также используется как name_scope при создании операций.
Атрибуты
logits_dimension См. base_head.Head для получения подробной информации.
loss_reduction См. base_head.Head для получения подробной информации.
name См. base_head.Head для получения подробной информации.

Методы

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 ровно один из 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

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

loss(
    labels, logits, features=None, mode=None, regularization_losses=None
)

Возвращает прогнозы на основе ключей. См. base_head.Head для получения подробной информации.

metrics

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

metrics(
    regularization_losses=None
)

Создаёт метрики. См. base_head.Head для получения подробной информации.

predictions

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

predictions(
    logits
)

Возвращает прогнозы на основе ключей.

См. base_head.Head для получения подробной информации.

Аргументы
logits Логиты Tensor с формой [D0, D1, ... DN, logits_dimension]. Для многих приложений форма равна [batch_size, logits_dimension].
Возвращаемое значение
Словарь прогнозов.

update_metrics

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

update_metrics(
    eval_metrics, features, logits, labels, regularization_losses=None
)

Обновляет метрики оценки. См. base_head.Head для получения подробной информации.

© 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/RegressionHead

Spec-Zone.ru

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