tf.keras.ops.custom_gradient
Декоратор для определения функции с пользовательским градиентом.
tf.keras.ops.custom_gradient(
f
)
Этот декоратор позволяет точно управлять градиентами последовательности операций. Это может быть полезно по многим причинам, включая обеспечение более эффективного или численно стабильного градиента для последовательности операций.
| Аргументы | |
|---|---|
f | Функция f(*args), которая возвращает кортеж (output, grad_fn), где:
|
| Возвращаемое значение | |
|---|---|
Функция h(*args), которая возвращает то же значение, что и f(*args)[0], и градиент которой определяется f(*args)[1]. |
Примеры:
- Примеp, не зависящий от бэкэнда.
@ops.custom_gradient
def log1pexp(x):
e = ops.exp(x)
def grad(*args, upstream=None):
if upstream is None:
(upstream,) = args
return ops.multiply(upstream, 1.0 - 1.0 / ops.add(1, e))
return ops.log(1 + e), grad
Обратите внимание, что функция grad, возвращающая вычисление градиента, требует args, а также аргумент upstream, в зависимости от установленного бэкэнда. С бэкендами JAX и TensorFlow требуется только один аргумент, тогда как в случае бэкэнда PyTorch может использоваться аргумент upstream.
При работе с бэкендом TensorFlow/JAX достаточно grad(upstream). С PyTorch функция grad требует *args а также upstream, например, def grad(*args, upstream). Следуйте предыдущему примеру, чтобы использовать @ops.custom_gradient способом, совместимым со всеми бэкендами.
- Вот пример, специфичный для JAX и TensorFlow:
@ops.custom_gradient
def log1pexp(x):
e = ops.exp(x)
def grad(upstream):
return ops.multiply(upstream, 1.0 - 1.0 / ops.add(1, e))
return ops.log(1 + e), grad
- Наконец, вот пример, специфичный для PyTorch, использующий
*argsиupstream:
@ops.custom_gradient
def log1pexp(x):
e = ops.exp(x)
def grad(*args, upstream):
return ops.multiply(upstream, 1.0 - 1.0 / ops.add(1, e))
return ops.log(1 + e), grad
© 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/keras/ops/custom_gradient