Метод watch класса GradientTape
Метод watch класса tf.GradientTape добавляет тензор в список отслеживаемых переменных. По умолчанию GradientTape автоматически отслеживает обучаемые переменные, созданные через tf.Variable. Однако если вы работаете с обычными тензорами, созданными через tf.constant, их необходимо явно передать в метод watch, чтобы лента могла вычислить градиент. Первым параметром метод принимает тензор, который нужно отслеживать. Вторым необязательным параметром можно указать тип отслеживания.
Синтаксис
GradientTape.watch(tensor)
Пример
Давайте создадим ленту градиентов и вычислим производную функции y = x² в точке x = 3. Поскольку x создан как константа, его нужно явно отследить:
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)
Пример
Давайте отследим несколько тензоров и вычислим частные производные функции z = x² + y³:
import tensorflow as tf
x = tf.constant(2.0)
y = tf.constant(3.0)
with tf.GradientTape(persistent=True) as tape:
tape.watch(x)
tape.watch(y)
z = x * x + y * y * y
dz_dx = tape.gradient(z, x)
dz_dy = tape.gradient(z, y)
print(dz_dx)
print(dz_dy)
del tape
Результат выполнения кода:
tf.Tensor(4.0, shape=(), dtype=float32)
tf.Tensor(27.0, shape=(), dtype=float32)
Пример
Давайте отследим тензор и вычислим градиент для линейной функции y = 2x + 1:
import tensorflow as tf
x = tf.constant(5.0)
with tf.GradientTape() as tape:
tape.watch(x)
y = 2 * x + 1
res = tape.gradient(y, x)
print(res)
Результат выполнения кода:
tf.Tensor(2.0, shape=(), dtype=float32)
Смотрите также
-
класс
GradientTape,
который записывает операции для автоматического дифференцирования -
метод
gradient,
который вычисляет градиент записанных операций -
метод
reset,
который очищает записанную информацию ленты -
метод
watched_variables,
который возвращает список отслеживаемых переменных