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

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