Метод reset класса GradientTape
Метод reset класса GradientTape очищает все записанные
на ленте операции и вычисленные градиенты. После вызова этого метода
лента возвращается в исходное состояние и может использоваться заново
для записи новых вычислений. Метод не принимает параметров и не
возвращает значения. Он особенно полезен, когда нужно выполнить
несколько циклов обучения в рамках одной ленты, не создавая новый
объект GradientTape каждый раз.
Синтаксис
tape.reset()
Пример
Давайте создадим ленту, вычислим градиент, а затем
сбросим состояние ленты методом reset:
import tensorflow as tf
x = tf.constant(3.0)
with tf.GradientTape() as tape:
tape.watch(x)
y = x * x
grad = tape.gradient(y, x)
print(grad)
tape.reset()
print(tape.watched_variables())
Результат выполнения кода:
tf.Tensor(6.0, shape=(), dtype=float32)
[]
Пример
Давайте покажем, что после сброса ленту можно использовать повторно для вычисления нового градиента:
import tensorflow as tf
x = tf.constant(2.0)
with tf.GradientTape() as tape:
tape.watch(x)
y = x * x * x
grad1 = tape.gradient(y, x)
print(grad1)
tape.reset()
with tape:
tape.watch(x)
z = x + x
grad2 = tape.gradient(z, x)
print(grad2)
Результат выполнения кода:
tf.Tensor(12.0, shape=(), dtype=float32)
tf.Tensor(2.0, shape=(), dtype=float32)
Смотрите также
-
класс
GradientTape,
который записывает операции для автоматического дифференцирования -
метод
watch,
который начинает отслеживание тензора на ленте -
метод
gradient,
который вычисляет градиент записанных операций -
метод
stop_recording,
который временно приостанавливает запись операций