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

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