Функция no_grad
Функция no_grad представляет собой контекстный менеджер, который отключает вычисление градиентов для всех операций внутри своего блока. Это позволяет существенно ускорить выполнение кода и уменьшить потребление памяти, так как PyTorch не сохраняет историю вычислений для градиентов. Функция не принимает параметров и используется в паре с оператором with.
Синтаксис
with torch.no_grad():
# операции без градиентов
t = тензор + 1
Пример
Давайте создадим тензор с градиентом и выполним операцию внутри контекста no_grad:
import torch
t = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
print(t.requires_grad)
with torch.no_grad():
res = t + 1
print(res.requires_grad)
Результат выполнения кода:
True
False
Пример
Давайте используем no_grad для вычислений без сохранения истории градиентов:
import torch
t = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
with torch.no_grad():
res = t * 2
res = res + 10
print(res)
print(res.requires_grad)
Результат выполнения кода:
tensor([12., 14., 16.])
False
Пример
Давайте сравним производительность вычислений с no_grad и без него:
import torch
import time
t = torch.randn(1000, 1000, requires_grad=True)
start = time.time()
for _ in range(100):
res = t * 2 + 1
end = time.time()
print(f"Without no_grad: {end - start:.4f} seconds")
start = time.time()
with torch.no_grad():
for _ in range(100):
res = t * 2 + 1
end = time.time()
print(f"With no_grad: {end - start:.4f} seconds")
Пример
Давайте используем no_grad для оценки модели на валидационных данных:
import torch
model = torch.nn.Linear(10, 1)
model.train()
data = torch.randn(32, 10)
target = torch.randn(32, 1)
# обучение
output = model(data)
loss = torch.nn.functional.mse_loss(output, target)
loss.backward()
# валидация без градиентов
model.eval()
with torch.no_grad():
val_data = torch.randn(32, 10)
val_output = model(val_data)
print(val_output)
Смотрите также
-
функцию
enable_grad,
которая включает вычисление градиентов в контексте -
функцию
set_grad_enabled,
которая включает или отключает градиенты по условию -
функцию
is_grad_enabled,
которая проверяет текущее состояние вычисления градиентов -
функцию
inference_mode,
которая является более быстрой альтернативой для инференса