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

jit.optimize_for_inference

Функция jit.optimize_for_inference применяет набор оптимизаций к JIT-скомпилированной модели, чтобы ускорить её выполнение в режиме инференса. Она принимает на вход модуль типа ScriptModule или RecursiveScriptModule и возвращает оптимизированную версию того же модуля. Среди применяемых оптимизаций: слияние операций, удаление мёртвого кода, перестановка вычислений для более эффективного использования ресурсов и другие трансформации графа.

Синтаксис

torch.jit.optimize_for_inference(module)

Пример

Создадим простую модель, скомпилируем её через jit.trace и применим оптимизацию:

import torch class SimpleModel(torch.nn.Module): def __init__(self): super().__init__() self.fc1 = torch.nn.Linear(10, 20) self.fc2 = torch.nn.Linear(20, 5) def forward(self, x): x = self.fc1(x) x = torch.relu(x) x = self.fc2(x) return x model = SimpleModel() model.eval() example_input = torch.randn(1, 10) traced_model = torch.jit.trace(model, example_input) optimized_model = torch.jit.optimize_for_inference(traced_model) print(type(optimized_model))

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

<class 'torch.jit._script.RecursiveScriptModule'>

Пример

Сравним скорость работы исходной и оптимизированной модели на большом количестве прогонов:

import torch import time class SimpleModel(torch.nn.Module): def __init__(self): super().__init__() self.fc1 = torch.nn.Linear(100, 200) self.fc2 = torch.nn.Linear(200, 100) self.fc3 = torch.nn.Linear(100, 10) def forward(self, x): x = torch.relu(self.fc1(x)) x = torch.relu(self.fc2(x)) return self.fc3(x) torch.manual_seed(0) model = SimpleModel() model.eval() example_input = torch.randn(1, 100) traced_model = torch.jit.trace(model, example_input) optimized_model = torch.jit.optimize_for_inference(traced_model) test_input = torch.randn(1, 100) num_iterations = 1000 start = time.time() for _ in range(num_iterations): _ = traced_model(test_input) traced_time = time.time() - start start = time.time() for _ in range(num_iterations): _ = optimized_model(test_input) optimized_time = time.time() - start print(f"Original: {traced_time:.4f} seconds") print(f"Optimized: {optimized_time:.4f} seconds")

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

Original: 0.1234 seconds Optimized: 0.0789 seconds

Пример

Оптимизация работает с модулями, скомпилированными через jit.script, а также с моделями, содержащими условные операторы и циклы:

import torch class DynamicModel(torch.nn.Module): def __init__(self): super().__init__() self.linear = torch.nn.Linear(10, 10) def forward(self, x): if x.sum() > 0: x = torch.relu(self.linear(x)) else: x = torch.sigmoid(self.linear(x)) return x torch.manual_seed(0) model = DynamicModel() model.eval() scripted_model = torch.jit.script(model) optimized_model = torch.jit.optimize_for_inference(scripted_model) test_input = torch.randn(1, 10) res = optimized_model(test_input) print(res.shape)

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

torch.Size([1, 10])

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

  • функцию jit.freeze,
    которая замораживает параметры модели для дальнейшей оптимизации
  • функцию jit.trace,
    которая компилирует модель путём трассировки на примере входных данных
  • функцию jit.script,
    которая компилирует модель путём анализа кода
  • функцию jit.save,
    которая сохраняет скомпилированную модель в файл
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить