Функция enable_grad
Функция enable_grad является контекстным менеджером,
который включает режим вычисления градиентов для всех операций,
выполняемых внутри его блока. Это полезно, когда необходимо
временно активировать автоматическое дифференцирование в участках
кода, где оно было отключено с помощью no_grad или
set_grad_enabled(False). Функция не принимает никаких
параметров и всегда включает режим градиентов, независимо от
текущего состояния.
Синтаксис
with torch.enable_grad():
# Код, в котором вычисляются градиенты
pass
Пример
Давайте рассмотрим базовый пример использования enable_grad
для включения градиентов внутри блока кода, где они были отключены:
import torch
t = torch.tensor([1.0, 2.0, 3.0], requires_grad=False)
with torch.no_grad():
# Здесь градиенты отключены
res = t.sum()
print(res.requires_grad)
with torch.enable_grad():
# Здесь градиенты включены
t2 = t * 2
res2 = t2.sum()
print(res2.requires_grad)
Результат выполнения кода:
False
True
Пример
Теперь рассмотрим пример, где enable_grad используется
для вычисления градиентов внутри функции, которая по умолчанию
выполняется без градиентов:
import torch
def compute_without_grad():
t = torch.tensor([4.0, 5.0, 6.0], requires_grad=True)
res = t.sum()
return res
with torch.no_grad():
# Здесь градиенты отключены
res = compute_without_grad()
print(res.requires_grad)
with torch.enable_grad():
# Здесь градиенты включены
res2 = compute_without_grad()
print(res2.requires_grad)
Результат выполнения кода:
False
True
Пример
Покажем, как enable_grad позволяет вычислить градиенты
для тензора внутри блока, где они были отключены глобально:
import torch
torch.set_grad_enabled(False)
t = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
res = t.sum()
print(res.requires_grad)
with torch.enable_grad():
t2 = t * 2
res2 = t2.sum()
res2.backward()
print(t.grad)
Результат выполнения кода:
False
tensor([2., 2., 2.])
Смотрите также
-
функцию
no_grad,
которая отключает вычисление градиентов для экономии памяти -
функцию
set_grad_enabled,
которая устанавливает режим градиентов глобально -
функцию
is_grad_enabled,
которая возвращает текущий статус режима градиентов -
функцию
inference_mode,
которая обеспечивает более быстрый режим без градиентов