Spec-Zone.ru › TensorFlow

tf.experimental.dispatch_for_api

Декоратор, который переопределяет стандартную реализацию API TensorFlow.

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

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

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

tf.compat.v1.experimental.dispatch_for_api

tf.experimental.dispatch_for_api(
    api, *signatures
)

Использование в блокнотах

Используется в руководстве
  • Типы расширений

Декорированная функция (известная как "цель диспетчеризации") переопределит стандартную реализацию API, когда API вызывается с параметрами, соответствующими заданной сигнатуре типа. Сигнатуры задаются с помощью словарей, которые сопоставляют имена параметров с аннотациями типов. Например, в следующем примере masked_add будет вызвана для tf.add, если оба x и y являются MaskedTensor:

class MaskedTensor(tf.experimental.ExtensionType):
  values: tf.Tensor
  mask: tf.Tensor
@dispatch_for_api(tf.math.add, {'x': MaskedTensor, 'y': MaskedTensor})
def masked_add(x, y, name=None):
  return MaskedTensor(x.values + y.values, x.mask & y.mask)
mt = tf.add(MaskedTensor([1, 2], [True, False]), MaskedTensor(10, True))
print(f"values={mt.values.numpy()}, mask={mt.mask.numpy()}")
values=[11 12], mask=[ True False]

Если указано несколько сигнатур типов, цель диспетчеризации будет вызвана, если совпадёт любая из сигнатур. Например, следующий код регистрирует masked_add для вызова, если x является MaskedTensor или y является MaskedTensor.

@dispatch_for_api(tf.math.add, {'x': MaskedTensor}, {'y':MaskedTensor})
def masked_add(x, y):
  x_values = x.values if isinstance(x, MaskedTensor) else x
  x_mask = x.mask if isinstance(x, MaskedTensor) else True
  y_values = y.values if isinstance(y, MaskedTensor) else y
  y_mask = y.mask if isinstance(y, MaskedTensor) else True
  return MaskedTensor(x_values + y_values, x_mask & y_mask)

Аннотации типов в сигнатурах типов могут быть объектами типов (например, MaskedTensor), значениями typing.List или значениями typing.Union. Например, следующее зарегистрирует masked_concat для вызова, если values представляет собой список значений MaskedTensor:

@dispatch_for_api(tf.concat, {'values': typing.List[MaskedTensor]})
def masked_concat(values, axis):
  return MaskedTensor(tf.concat([v.values for v in values], axis),
                      tf.concat([v.mask for v in values], axis))

Каждая сигнатура типа должна содержать по крайней мере один подкласс tf.CompositeTensor (включая подклассы tf.ExtensionType), и диспетчеризация будет выполнена только в том случае, если хотя бы один параметр с аннотацией типа содержит значение CompositeTensor. Это правило предотвращает вызов диспетчеризации в вырожденных случаях, таких как следующие примеры:

  • @dispatch_for_api(tf.concat, {'values': List[MaskedTensor]}): Не будет отправлена диспетчеризация на декорируемую цель диспетчеризации, когда пользователь вызовет tf.concat([]).

  • @dispatch_for_api(tf.add, {'x': Union[MaskedTensor, Tensor], 'y': Union[MaskedTensor, Tensor]}): Не будет отправлена диспетчеризация на декорируемую цель диспетчеризации, когда пользователь вызовет tf.add(tf.constant(1), tf.constant(2)).

Подпись целевой функции диспетчеризации должна соответствовать подписи API, которая переопределяется. В частности, параметры должны иметь одинаковые имена и должны следовать в том же порядке. Цель диспетчеризации может необязательно опустить параметр "имя", в этом случае он будет обернут вызовом tf.name_scope при необходимости.

Аргументы
api API TensorFlow для переопределения.
*signatures Словари, сопоставляющие имена или индексы параметров с аннотациями типов, определяющими, когда необходимо вызвать целевую функцию диспетчеризации. В частности, целевая функция диспетчеризации будет вызвана, если любая из сигнатур совпадёт; и сигнатура совпадает, если типы всех указанных параметров соответствуют указанным аннотациям типов. Если сигнатуры не указаны, сигнатура будет считываться из аннотаций типа функции-цели диспетчеризации.
Возвращаемое значение
Декоратор, который переопределяет стандартную реализацию api.

Зарегистрированные API

API TensorFlow, которые могут быть переопределены с помощью @dispatch_for_api:

<>

© 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/api_docs/python/tf/experimental/dispatch_for_api

Spec-Zone.ru

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