Метод stop_recording
Метод stop_recording класса GradientTape временно останавливает запись операций. Пока запись остановлена, все выполняемые операции не отслеживаются и не попадают в ленту для вычисления градиентов. Метод используется как контекстный менеджер и не принимает параметров. Это позволяет исключить из вычислений вспомогательные операции, которые не должны влиять на градиент.
Синтаксис
with tape.stop_recording():
operations
Пример
Давайте создадим ленту GradientTape и сравним, какие операции попадают в запись при обычном выполнении и внутри блока stop_recording:
import tensorflow as tf
x = tf.constant(3.0)
with tf.GradientTape() as tape:
tape.watch(x)
y = x * x
with tape.stop_recording():
z = x + x
w = x + x
print("gradient of y:", tape.gradient(y, x))
print("watched:", tape.watched_variables())
Результат выполнения кода:
gradient of y: tf.Tensor(6.0, shape=(), dtype=float32)
watched: ()
Пример
Давайте убедимся, что операции, выполненные внутри stop_recording, не отслеживаются и градиент по ним не вычисляется:
import tensorflow as tf
x = tf.constant(2.0)
with tf.GradientTape() as tape:
tape.watch(x)
with tape.stop_recording():
inner = x * x
outer = x * x
print("gradient of outer:", tape.gradient(outer, x))
print("gradient of inner:", tape.gradient(inner, x))
Результат выполнения кода:
gradient of outer: tf.Tensor(4.0, shape=(), dtype=float32)
gradient of inner: None
Пример
Давайте используем stop_recording для исключения вспомогательных вычислений из ленты при обучении модели:
import tensorflow as tf
tf.random.set_seed(0)
model = tf.keras.Sequential([
tf.keras.layers.Dense(1)
])
x = tf.constant([[1.0], [2.0], [3.0]])
y = tf.constant([[2.0], [4.0], [6.0]])
with tf.GradientTape() as tape:
pred = model(x)
loss = tf.reduce_mean(tf.square(pred - y))
with tape.stop_recording():
metric = tf.reduce_mean(tf.abs(pred - y))
grads = tape.gradient(loss, model.trainable_variables)
print("loss:", loss.numpy())
print("metric:", metric.numpy())
print("gradients:", [g.numpy() for g in grads])
Результат выполнения кода:
loss: 0.07195725
metric: 0.2175926
gradients: [array([[-0.11383986]], dtype=float32), array([-0.06893867], dtype=float32)]
Смотрите также
-
класс
GradientTape,
который записывает операции для вычисления градиентов -
метод
watch,
который начинает отслеживание тензора -
метод
gradient,
который вычисляет градиент записанных операций -
метод
reset,
который очищает записанную информацию ленты