Функция checkpoint
Функция checkpoint (из модуля torch.utils.checkpoint) используется для уменьшения использования памяти графическим процессором во время обучения нейронных сетей. Вместо того чтобы сохранять все промежуточные активации для обратного распространения ошибки, эта функция сохраняет только входные данные, а активации пересчитывает заново на этапе backward. Это особенно полезно при работе с большими моделями или высоким разрешением изображений. Первым аргументом передаётся функция (или слой), которая будет выполнена, а последующие аргументы - это входные данные для этой функции.
Синтаксис
torch.utils.checkpoint.checkpoint(function, *args, use_reentrant=True)
Параметр use_reentrant управляет поведением при повторных вызовах. В новых версиях PyTorch рекомендуется устанавливать его в False для повышения производительности.
Пример
Рассмотрим простую модель и применим функцию checkpoint к одному из линейных слоёв:
import torch
import torch.nn as nn
from torch.utils.checkpoint import checkpoint
class Model(nn.Module):
def __init__(self):
super().__init__()
self.fc1 = nn.Linear(10, 20)
self.fc2 = nn.Linear(20, 30)
def forward(self, x):
x = self.fc1(x)
x = checkpoint(self.fc2, x, use_reentrant=False)
return x
model = Model()
x = torch.randn(5, 10, requires_grad=True)
y = model(x)
y.sum().backward()
print("Backward completed")
Результат выполнения кода:
"Backward completed"
При использовании чекпоинта градиенты всё равно вычисляются корректно, но память используется более экономно.
Пример
Использование чекпоинта внутри цикла или для кастомной функции:
import torch
from torch.utils.checkpoint import checkpoint
def my_computation(x, y):
return x * y
a = torch.randn(3, 4, requires_grad=True)
b = torch.randn(3, 4, requires_grad=False)
res = checkpoint(my_computation, a, b, use_reentrant=False)
res.sum().backward()
print(a.grad)
Результат выполнения кода:
tensor([[1., 1., 1., 1.],
[1., 1., 1., 1.],
[1., 1., 1., 1.]])
Обратите внимание, что grad для тензора a был вычислен без сохранения промежуточных результатов.
Пример
Использование чекпоинта с несколькими аргументами:
import torch
from torch.utils.checkpoint import checkpoint
def forward(x, w, b):
return x @ w + b
x = torch.randn(4, 3, requires_grad=True)
w = torch.randn(3, 2, requires_grad=True)
b = torch.randn(2, requires_grad=True)
res = checkpoint(forward, x, w, b, use_reentrant=False)
res.sum().backward()
print("Gradients computed")
Результат выполнения кода:
"Gradients computed"
Чекпоинт корректно работает с любым количеством входных тензоров.
Смотрите также
-
контекстный менеджер
no_grad,
который отключает вычисление градиентов на время выполнения блока -
контекстный менеджер
enable_grad,
который принудительно включает режим вычисления градиентов -
функцию
backward,
которая запускает обратное распространение ошибки -
функцию
grad,
которая вычисляет градиенты для заданных входных тензоров