tf.RegisterGradient
| Просмотреть исходный код на GitHub |
Декоратор для регистрации функции градиента для типа операции.
tf.RegisterGradient(
op_type
)
Этот декоратор используется только при определении нового типа операции. Для операции с m входами и n выходами функция градиента — это функция, которая принимает исходные Operation и n Tensor объекты (представляющие градиенты по отношению к каждому выходу операции) и возвращает m Tensor объекты (представляющие частные градиенты по отношению к каждому входу операции).
Например, предположим, что операции типа "Sub" принимают два входа x и y, и возвращают один выход x - y, следующая функция градиента будет зарегистрирована:
@tf.RegisterGradient("Sub")
def _sub_grad(unused_op, grad):
return grad, tf.negative(grad)
Аргумент декоратора op_type — это строковый тип операции. Это соответствует полю OpDef.name для протокола, определяющего операцию.
| Аргументы | |
|---|---|
op_type | Строковый тип операции. Это соответствует полю OpDef.name для протокола, определяющего операцию. |
| Исключения | |
|---|---|
TypeError | Если op_type не является строкой. |
Методы
__call__
__call__(
f
)
Регистрирует функцию f как функцию градиента для op_type.
© 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/r2.3/api_docs/python/tf/RegisterGradient