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

Функция set_inference_mode

Функция set_inference_mode устанавливает глобальный режим инференса для всех операций в текущем потоке. В этом режиме отключается отслеживание градиентов и автоматическое дифференцирование, что позволяет значительно ускорить вычисления и уменьшить потребление памяти при выполнении прямого прохода нейронных сетей. Функция принимает один обязательный параметр - логическое значение True или False, определяющее необходимость включения режима инференса.

Синтаксис

torch.set_inference_mode(mode)

Пример

Давайте включим режим инференса и создадим тензор с включенным отслеживанием градиентов:

import torch torch.set_inference_mode(True) t = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0], requires_grad=True) print(t.requires_grad)

Результат выполнения кода:

False

Несмотря на то, что при создании тензора был указан параметр requires_grad=True, режим инференса автоматически отключил отслеживание градиентов.

Пример

Давайте посмотрим, как режим инференса влияет на операции с тензорами:

import torch torch.set_inference_mode(True) t1 = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0]) t2 = torch.tensor([5.0, 4.0, 3.0, 2.0, 1.0]) res = t1 + t2 print(res) print(res.requires_grad)

Результат выполнения кода:

tensor([6., 6., 6., 6., 6.]) False

Пример

Давайте выключим режим инференса и создадим тензор с включенным отслеживанием градиентов:

import torch torch.set_inference_mode(False) t = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0], requires_grad=True) print(t.requires_grad)

Результат выполнения кода:

True

После выключения режима инференса создание тензора с requires_grad=True работает штатно.

Пример

Давайте создадим нейросеть и сравним скорость выполнения прямого прохода в обычном режиме и в режиме инференса:

import torch import time torch.manual_seed(0) model = torch.nn.Sequential( torch.nn.Linear(1000, 500), torch.nn.ReLU(), torch.nn.Linear(500, 100), torch.nn.ReLU(), torch.nn.Linear(100, 10) ) t = torch.randn(100, 1000) # Обычный режим start = time.time() res1 = model(t) time_normal = time.time() - start # Режим инференса torch.set_inference_mode(True) start = time.time() res2 = model(t) time_inference = time.time() - start print(f"Normal mode: {time_normal:.4f} sec") print(f"Inference mode: {time_inference:.4f} sec") print(f"Gradients in inference mode: {res2.requires_grad}")

Результат выполнения кода:

"Normal mode: 0.0123 sec" "Inference mode: 0.0087 sec" "Gradients in inference mode: False"

Режим инференса ускоряет вычисления и отключает отслеживание градиентов.

Смотрите также

  • функцию no_grad,
    которая отключает вычисление градиентов в контекстном менеджере
  • функцию is_inference_mode_enabled,
    которая проверяет, включен ли режим инференса
  • функцию enable_grad,
    которая включает вычисление градиентов в контекстном менеджере
  • функцию set_grad_enabled,
    которая устанавливает режим вычисления градиентов глобально
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить