Функция is_grad_enabled
Функция is_grad_enabled проверяет,
включен ли в данный момент режим автоматического
дифференцирования (вычисления градиентов).
Она не принимает никаких параметров и возвращает
логическое значение. Функция полезна для проверки
текущего состояния системы, особенно когда используются
контекстные менеджеры, изменяющие режим работы
с градиентами.
Синтаксис
torch.is_grad_enabled()
Возвращаемое значение: bool
Пример
Давайте проверим состояние градиентов по умолчанию:
import torch
res = torch.is_grad_enabled()
print(res)
Результат выполнения кода:
True
По умолчанию вычисление градиентов включено.
Пример
Давайте проверим состояние внутри контекста no_grad:
import torch
with torch.no_grad():
res = torch.is_grad_enabled()
print(res)
Результат выполнения кода:
False
Внутри блока no_grad функция возвращает
False, так как вычисление градиентов отключено.
Пример
Давайте проверим состояние внутри контекста enable_grad:
import torch
with torch.enable_grad():
res = torch.is_grad_enabled()
print(res)
Результат выполнения кода:
True
Контекст enable_grad явно включает
вычисление градиентов, поэтому функция возвращает
True.
Пример
Давайте используем is_grad_enabled
в условном блоке для проверки возможности вычисления
градиентов:
import torch
x = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
y = x ** 2
if torch.is_grad_enabled():
print("Gradients are enabled, computing backward...")
y.sum().backward()
print(x.grad)
else:
print("Gradients are disabled")
Результат выполнения кода:
Gradients are enabled, computing backward...
tensor([2., 4., 6.])
Функция позволяет безопасно выполнять операции, зависящие от состояния градиентов.
Смотрите также
-
контекстный менеджер
no_grad,
который отключает вычисление градиентов -
контекстный менеджер
enable_grad,
который включает вычисление градиентов -
функцию
set_grad_enabled,
которая устанавливает режим вычисления градиентов -
контекстный менеджер
inference_mode,
который переключает режим инференса