Функция 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,
которая устанавливает режим вычисления градиентов глобально