Функция 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,
которая выводит значение тензора при выполнении графа