tf.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 zmy_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
runninggrad = 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
runninggrad = 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 xmodel = 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
runninggrad = 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/api_docs/python/tf/recompute_grad