Класс GradientTape
Класс GradientTape позволяет записывать операции, выполняемые над тензорами, для последующего автоматического вычисления градиентов. Это основной инструмент для реализации пользовательских циклов обучения и отладки моделей. Первым параметром конструктор принимает логическое значение persistent, которое определяет, можно ли вызывать метод gradient несколько раз. Вторым параметром можно передать watch_accessed_variables, чтобы автоматически отслеживать все обучаемые переменные.
Синтаксис
tf.GradientTape(persistent=False, watch_accessed_variables=True)
Пример
Давайте вычислим градиент функции y = x² в точке x = 3. Для этого создадим объект GradientTape и вызовем метод gradient:
import tensorflow as tf
x = tf.constant(3.0)
with tf.GradientTape() as tape:
tape.watch(x)
y = x * x
res = tape.gradient(y, x)
print(res)
Результат выполнения кода:
tf.Tensor(6.0, shape=(), dtype=float32)
Пример
Давайте вычислим градиенты для двух переменных одновременно. Для этого передадим список тензоров в метод gradient:
import tensorflow as tf
x = tf.constant(2.0)
y = tf.constant(4.0)
with tf.GradientTape() as tape:
tape.watch([x, y])
z = x * x + y * y
res = tape.gradient(z, [x, y])
print(res)
Результат выполнения кода:
[<tf.Tensor: shape=(), dtype=float32, numpy=4.0>, <tf.Tensor: shape=(), dtype=float32, numpy=8.0>]
Пример
Давайте используем параметр persistent=True, чтобы вычислить градиент дважды без повторной записи операций:
import tensorflow as tf
x = tf.constant(3.0)
with tf.GradientTape(persistent=True) as tape:
tape.watch(x)
y = x * x
res1 = tape.gradient(y, x)
res2 = tape.gradient(y, x)
print(res1)
print(res2)
del tape
Результат выполнения кода:
tf.Tensor(6.0, shape=(), dtype=float32)
tf.Tensor(6.0, shape=(), dtype=float32)
Смотрите также
-
метод
watch,
который начинает отслеживание тензора -
метод
gradient,
который вычисляет градиент записанных операций -
метод
reset,
который очищает записанные операции -
метод
stop_recording,
который временно останавливает запись операций