Функция stop_gradient
Функция stop_gradient останавливает
вычисление градиента для переданного тензора.
Первым параметром функция принимает тензор,
для которого нужно прекратить отслеживание
градиента. Функция возвращает тензор с теми же
значениями, что и входной, но при вычислении
градиентов он рассматривается как константа.
Это полезно, когда нужно исключить часть
вычислений из процесса обучения.
Синтаксис
tf.stop_gradient(input, [name])
Пример
Давайте создадим тензор и остановим для него вычисление градиента:
import tensorflow as tf
t = tf.constant([1, 2, 3, 4, 5])
res = tf.stop_gradient(t)
print(res)
Результат выполнения кода:
tf.Tensor([1 2 3 4 5], shape=(5,), dtype=int32)
Пример
Давайте проверим, что градиент не вычисляется
для тензора, обработанного функцией
stop_gradient:
import tensorflow as tf
t = tf.constant([1, 2, 3, 4, 5], dtype=tf.float32)
with tf.GradientTape() as tape:
tape.watch(t)
res = tf.stop_gradient(t) * 2
grad = tape.gradient(res, t)
print(grad)
Результат выполнения кода:
None
Пример
Давайте сравним поведение обычного тензора и тензора с остановленным градиентом:
import tensorflow as tf
t = tf.constant([1, 2, 3, 4, 5], dtype=tf.float32)
with tf.GradientTape() as tape:
tape.watch(t)
res1 = t * 2
res2 = tf.stop_gradient(t) * 2
grad1 = tape.gradient(res1, t)
grad2 = tape.gradient(res2, t)
print(grad1)
print(grad2)
Результат выполнения кода:
tf.Tensor([2. 2. 2. 2. 2.], shape=(5,), dtype=float32)
None
Смотрите также
-
функцию
constant,
которая создает тензор из переданных данных -
функцию
cast,
которая преобразует тип данных тензора -
функцию
convert_to_tensor,
которая преобразует данные в тензор -
функцию
identity,
которая возвращает тензор с теми же значениями