Spec-Zone.ru › TensorFlow

tf.keras.ops.custom_gradient

Декоратор для определения функции с пользовательским градиентом.

tf.keras.ops.custom_gradient(
    f
)

Этот декоратор позволяет точно управлять градиентами последовательности операций. Это может быть полезно по многим причинам, включая обеспечение более эффективного или численно стабильного градиента для последовательности операций.

Аргументы
f Функция f(*args), которая возвращает кортеж (output, grad_fn), где:
  • args — последовательность (вложенных структур) тензорных входных данных функции.
  • output — (вложенная структура) тензорный результат применения операций в forward_fn к args.
  • grad_fn — функция со сигнатурой grad_fn(*args, upstream), которая возвращает кортеж тензоров того же размера, что и (сглаженная) args: производные тензоров в output по тензорам в args. upstream — тензор или последовательность тензоров, содержащих начальное значение градиента для каждого тензора в output.
Возвращаемое значение
Функция h(*args), которая возвращает то же значение, что и f(*args)[0], и градиент которой определяется f(*args)[1].

Примеры:

  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 способом, совместимым со всеми бэкендами.

  1. Вот пример, специфичный для 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
  1. Наконец, вот пример, специфичный для 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

Spec-Zone.ru

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