Функция set_grad_enabled
Функция set_grad_enabled управляет режимом вычисления градиентов
для всех операций с тензорами в PyTorch. Первым параметром функция
принимает булево значение True или False, которое
включает или отключает вычисление градиентов соответственно.
В отличие от контекстных менеджеров no_grad и enable_grad,
данная функция изменяет глобальное состояние автоматической
дифференциации и может использоваться внутри функций для управления
поведением в зависимости от условий.
Синтаксис
torch.set_grad_enabled(mode: bool) -> None
Пример
Давайте рассмотрим базовый пример использования функции для отключения градиентов при вычислениях:
import torch
x = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
torch.set_grad_enabled(False)
y = x * 2
print(y.requires_grad)
torch.set_grad_enabled(True)
Результат выполнения кода:
False
Пример
Рассмотрим использование функции внутри условной конструкции для гибкого управления режимом градиентов:
import torch
def compute_loss(x, use_grad=True):
torch.set_grad_enabled(use_grad)
res = x * x
torch.set_grad_enabled(True)
return res
x = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
loss1 = compute_loss(x, use_grad=True)
loss2 = compute_loss(x, use_grad=False)
print(loss1.requires_grad, loss2.requires_grad)
Результат выполнения кода:
True False
Пример
Покажем, что функция set_grad_enabled действует глобально
и переопределяет текущий режим:
import torch
x = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
with torch.no_grad():
y1 = x * 2
torch.set_grad_enabled(True)
y2 = x * 3
print(y1.requires_grad)
print(y2.requires_grad)
Результат выполнения кода:
False
True
Пример
Важно восстанавливать исходное состояние после использования функции, чтобы избежать неожиданного поведения в других частях кода:
import torch
torch.manual_seed(0)
x = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
original_state = torch.is_grad_enabled()
torch.set_grad_enabled(False)
res = x * x
torch.set_grad_enabled(original_state)
print(res.requires_grad)
Результат выполнения кода:
False
Смотрите также
-
функцию
no_grad,
которая создает контекстный менеджер для отключения градиентов -
функцию
enable_grad,
которая создает контекстный менеджер для включения градиентов -
функцию
is_grad_enabled,
которая возвращает текущее состояние режима градиентов -
функцию
inference_mode,
которая создает контекстный менеджер для инференса без градиентов