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,
которая сохраняет скомпилированную модель в файл