РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
624 of 769 menu

Класс 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,
    которая проверяет текущее состояние режима градиентов
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить