РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
773 of 824 menu

Функция custom_gradient

Функция custom_gradient позволяет определить собственную формулу градиента для операции. Это бывает нужно, когда автоматическое дифференцирование TensorFlow работает некорректно или когда известна более эффективная формула производной. Первым параметром передается функция, вычисляющая прямое значение. Вторым параметром передается функция градиента. Внутри функции градиента первым аргументом идет градиент от выходной величины, а далее - исходные аргументы прямой функции.

Синтаксис

tf.custom_gradient(f, [grad])

Пример

Давайте создадим функцию, которая вычисляет квадрат числа, и зададим для нее собственный градиент 2x:

import tensorflow as tf @tf.custom_gradient def square(x): def grad(dy): return dy * 2 * x return x * x, grad t = tf.constant(3.0) with tf.GradientTape() as tape: tape.watch(t) res = square(t) print(res) print(tape.gradient(res, t))

Результат выполнения кода:

tf.Tensor(9.0, shape=(), dtype=float32) tf.Tensor(6.0, shape=(), dtype=float32)

Пример

Давайте применим custom_gradient к функции, которая возвращает значение с обрезанным градиентом:

import tensorflow as tf @tf.custom_gradient def clip_grad(x): def grad(dy): return tf.clip_by_value(dy, -1.0, 1.0) return x, grad t = tf.constant([1.0, 2.0, 3.0]) with tf.GradientTape() as tape: tape.watch(t) res = tf.reduce_sum(clip_grad(t) * 5.0) print(res) print(tape.gradient(res, t))

Результат выполнения кода:

tf.Tensor(6.0, shape=(), dtype=float32) tf.Tensor([1. 1. 1.], shape=(3,), dtype=float32)

Пример

Давайте зададим градиент для функции, которая масштабирует входной тензор:

import tensorflow as tf @tf.custom_gradient def scale(x, factor): def grad(dy): return dy * factor, tf.reduce_sum(dy * x) return x * factor, grad t = tf.constant([1.0, 2.0, 3.0, 4.0, 5.0]) f = tf.constant(2.0) with tf.GradientTape() as tape: tape.watch(t) tape.watch(f) res = tf.reduce_sum(scale(t, f)) print(res) print(tape.gradient(res, t)) print(tape.gradient(res, f))

Результат выполнения кода:

tf.Tensor(30.0, shape=(), dtype=float32) tf.Tensor([2. 2. 2. 2. 2.], shape=(5,), dtype=float32) tf.Tensor(15.0, shape=(), dtype=float32)

Смотрите также

  • функцию tf_function,
    которая компилирует функцию в граф TensorFlow
  • функцию recompute_grad,
    которая пересчитывает градиент для экономии памяти
  • функцию device,
    которая задает устройство для выполнения операции
  • функцию print,
    которая выводит значение тензора при выполнении графа
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить