Класс grad_mode
Класс grad_mode является базовым для всех контекстных менеджеров,
которые управляют включением и отключением автоматического
дифференцирования в PyTorch. От него наследуются классы
no_grad, enable_grad и inference_mode.
Этот класс предоставляет общий интерфейс для работы с режимами
градиентов и не предназначен для непосредственного использования.
Основная задача класса - определить единый способ переключения
состояния вычисления градиентов через контекстные менеджеры.
Все дочерние классы переопределяют методы __enter__ и
__exit__ для установки и восстановления нужного режима.
Иерархия наследования
Класс grad_mode находится в модуле torch.autograd.
Его иерархия выглядит следующим образом:
from torch.autograd import grad_mode
# Базовый класс для всех контекстных менеджеров
print(grad_mode.__name__)
print(grad_mode.__bases__)
Результат выполнения кода:
"grad_mode"
"(<class 'object'>,)"
Все конкретные реализации наследуются от grad_mode:
import torch
print(issubclass(torch.no_grad, torch.autograd.grad_mode))
print(issubclass(torch.enable_grad, torch.autograd.grad_mode))
print(issubclass(torch.inference_mode, torch.autograd.grad_mode))
Результат выполнения кода:
"True"
"True"
"True"
Назначение дочерних классов
Каждый дочерний класс grad_mode решает свою задачу:
-
no_grad- отключает вычисление градиентов для экономии памяти -
enable_grad- включает вычисление градиентов внутри блока, где оно было отключено -
inference_mode- самый строгий режим, полностью отключающий градиенты и некоторые проверки
Рассмотрим пример использования всех трёх режимов:
import torch
t = torch.tensor([1., 2., 3.], requires_grad=True)
# Обычный режим - градиенты вычисляются
res = t.sum()
print(res.requires_grad)
# Режим no_grad - градиенты отключены
with torch.no_grad():
res = t.sum()
print(res.requires_grad)
# Режим enable_grad - принудительно включает градиенты
with torch.enable_grad():
res = t.sum()
print(res.requires_grad)
# Режим inference_mode - полное отключение
with torch.inference_mode():
res = t.sum()
print(res.requires_grad)
Результат выполнения кода:
"True"
"False"
"True"
"False"
Обратите внимание, что в режиме inference_mode тензор
res не только не требует градиентов, но и вообще не
поддерживает операции, связанные с autograd.
Внутреннее устройство
Класс grad_mode определяет метод __call__, который
позволяет использовать его как декоратор для функций. Это удобно,
когда нужно применить режим к целой функции:
import torch
@torch.no_grad()
def no_grad_func(x):
return x.sum()
t = torch.tensor([1., 2., 3.], requires_grad=True)
res = no_grad_func(t)
print(res.requires_grad)
Результат выполнения кода:
"False"
Этот механизм работает для всех наследников grad_mode,
предоставляя единообразный интерфейс.
Пример вложенных режимов
Класс grad_mode определяет правила вложения контекстных
менеджеров. Например, режим inference_mode имеет наивысший
приоритет и отменяет действие других режимов:
import torch
t = torch.tensor([1., 2., 3.], requires_grad=True)
# Внутри inference_mode не работает даже enable_grad
with torch.inference_mode():
with torch.enable_grad():
res = t.sum()
print("In inference_mode:", res.requires_grad)
# А внутри no_grad работает enable_grad
with torch.no_grad():
with torch.enable_grad():
res = t.sum()
print("In no_grad:", res.requires_grad)
Результат выполнения кода:
"In inference_mode: False"
"In no_grad: True"
Это поведение чётко определено в реализации класса grad_mode
и его потомков.
Когда использовать класс grad_mode
Непосредственно класс grad_mode не используется в коде.
Вместо этого программист работает с его наследниками:
-
Используйте
torch.no_gradдля оценки модели или вывода результатов -
Используйте
torch.enable_gradдля принудительного включения градиентов -
Используйте
torch.inference_modeдля максимальной производительности при инференсе
Знание о классе grad_mode полезно для понимания архитектуры
PyTorch и написания собственных контекстных менеджеров с аналогичным
поведением.
Смотрите также
-
класс
no_grad,
который отключает вычисление градиентов в контекстном блоке -
класс
enable_grad,
который принудительно включает вычисление градиентов -
класс
inference_mode,
который полностью отключает autograd для инференса -
функцию
is_grad_enabled,
которая проверяет текущее состояние режима градиентов