tf.RegisterGradient
Декоратор для регистрации функции градиента для типа операции.
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: _T
) -> _T
Регистрирует функцию f в качестве функции градиента для op_type.
© 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/RegisterGradient