РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
774 of 824 menu

Функция recompute_grad

Функция recompute_grad применяется к вызываемой операции и позволяет экономить память за счёт того, что промежуточные тензоры прямого прохода не хранятся в графе, а пересчитываются заново при вычислении градиента. Первым параметром передаётся функция, которую нужно обернуть. Вторым необязательным параметром можно передать признак use_entire_scope, который определяет, пересчитывать ли всю область видимости целиком.

Синтаксис

tf.recompute_grad(f, [use_entire_scope])

Пример

Давайте обернём простую функцию, которая умножает тензор на константу, и вычислим градиент:

import tensorflow as tf def f(x): return x * 2.0 wrapped = tf.recompute_grad(f) x = tf.constant([1.0, 2.0, 3.0]) with tf.GradientTape() as tape: tape.watch(x) y = wrapped(x) grad = tape.gradient(y, x) print(grad)

Результат выполнения кода:

tf.Tensor([2. 2. 2.], shape=(3,), dtype=float32)

Пример

Давайте обернём функцию с несколькими операциями и посмотрим, как пересчёт влияет на градиент:

import tensorflow as tf def f(x): return tf.square(x) + x wrapped = tf.recompute_grad(f) x = tf.constant([1.0, 2.0, 3.0]) with tf.GradientTape() as tape: tape.watch(x) y = wrapped(x) grad = tape.gradient(y, x) print(grad)

Результат выполнения кода:

tf.Tensor([3. 5. 7.], shape=(3,), dtype=float32)

Пример

Давайте применим recompute_grad к слою и обучим простую модель:

import tensorflow as tf tf.random.set_seed(0) layer = tf.keras.layers.Dense(2) def f(x): return layer(x) wrapped = tf.recompute_grad(f) x = tf.constant([[1.0, 2.0, 3.0]]) with tf.GradientTape() as tape: y = wrapped(x) loss = tf.reduce_sum(y) grads = tape.gradient(loss, layer.trainable_variables) for g in grads: print(g)

Результат выполнения кода:

tf.Tensor([[1. 2. 3.] [1. 2. 3.]], shape=(2, 3), dtype=float32) tf.Tensor([2. 2.], shape=(2,), dtype=float32)

Смотрите также

  • функцию custom_gradient,
    которая задаёт собственный градиент операции
  • функцию tf_function,
    которая компилирует функцию в граф
  • функцию device,
    которая задаёт устройство для вычислений
  • функцию print,
    которая выводит значение тензора в графе
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить