tf.experimental.dispatch_for_api
Декоратор, который переопределяет стандартную реализацию API TensorFlow.
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