Spec-Zone.ru › TensorFlow 2.9

tf.recompute_grad

Просмотреть исходный код на GitHub

Определяет функцию как точку восстановления вычислений для автоматической дифференциации с помощью ленты.

Просмотр псевдонимов

Псевдонимы для миграции

См. Руководство по миграции для получения дополнительной информации.

tf.compat.v1.recompute_grad

tf.recompute_grad(
    f
)

Вычисление по контрольным точкам ленты — это метод уменьшения потребления памяти лентой автоматической дифференциации:

  • Без вычисления по контрольным точкам ленты операции и промежуточные значения записываются в ленту для использования в обратном проходе.

  • При вычислении по контрольным точкам ленты записывается только вызов функции и её входные данные. Во время обратного распространения градиента recompute_grad пользовательский градиент (tf.custom_gradient) перевычисляет функцию в локализованном объекте ленты. Такое перевычисление функции во время обратного распространения производит избыточные вычисления, но уменьшает общее потребление памяти лентой.

y = tf.Variable(1.0)
def my_function(x):
  tf.print('running')
  z = x*y
  return z
my_function_recompute = tf.recompute_grad(my_function)
with tf.GradientTape() as tape:
  r = tf.constant(1.0)
  for i in range(4):
    r = my_function_recompute(r)
running
running
running
running
grad = tape.gradient(r, [y])
running
running
running
running

Без recompute_grad, лента содержит все промежуточные шаги, и перевычисление не выполняется.

with tf.GradientTape() as tape:
  r = tf.constant(1.0)
  for i in range(4):
    r = my_function(r)
running
running
running
running
grad = tape.gradient(r, [y])

Если f был tf.keras Model или Layer объектом, методы и атрибуты, такие как f.variables, недоступны для возвращаемой функции g. Сохраните ссылку на f или используйте g.__wrapped__ для доступа к этим переменным и методам.

def print_running_and_return(x):
  tf.print("running")
  return x
model = tf.keras.Sequential([
  tf.keras.layers.Lambda(print_running_and_return),
  tf.keras.layers.Dense(2)
])
model_recompute = tf.recompute_grad(model)
with tf.GradientTape(persistent=True) as tape:
  r = tf.constant([[1,2]])
  for i in range(4):
    r = model_recompute(r)
running
running
running
running
grad = tape.gradient(r, model.variables)
running
running
running
running

В качестве альтернативы используйте атрибут __wrapped__, чтобы получить доступ к исходному объекту модели.

grad = tape.gradient(r, model_recompute.__wrapped__.variables)
running
running
running
running
Аргументы
f функция f(*x), возвращающая Tensor или последовательность Tensor выходных значений.
Возвращаемое значение
Функция g — обертка вокруг f, которая определяет пользовательский градиент, перевычисляющий f в обратном проходе вызова градиента.

© 2022 The TensorFlow Authors. All rights reserved.
Licensed under the Creative Commons Attribution License 4.0.
Code samples licensed under the Apache 2.0 License.
https://www.tensorflow.org/versions/r2.9/api_docs/python/tf/recompute_grad

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API