Функция jit.freeze
Функция jit.freeze применяется к модели, уже преобразованной в формат TorchScript с помощью jit.script или jit.trace. Она замораживает все параметры модели, делая их константами, и оптимизирует граф вычислений для ускорения инференса. Первым параметром функция принимает скриптовый модуль.
Синтаксис
torch.jit.freeze(module)
Пример
Давайте заморозим простую модель линейной регрессии:
import torch
import torch.nn as nn
class LinearModel(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(10, 5)
def forward(self, x):
return self.fc(x)
model = LinearModel()
scripted_model = torch.jit.script(model)
frozen_model = torch.jit.freeze(scripted_model)
print(type(frozen_model))
Результат выполнения кода:
<class 'torch.jit._freeze.FrozenScriptModule'>
Пример
Давайте сравним скорость работы обычной модели и замороженной:
import torch
import torch.nn as nn
import time
torch.manual_seed(0)
class SimpleNet(nn.Module):
def __init__(self):
super().__init__()
self.fc1 = nn.Linear(100, 50)
self.fc2 = nn.Linear(50, 10)
def forward(self, x):
x = torch.relu(self.fc1(x))
return self.fc2(x)
model = SimpleNet()
scripted_model = torch.jit.script(model)
frozen_model = torch.jit.freeze(scripted_model)
x = torch.randn(100, 100)
start = time.time()
res = model(x)
print("Обычная модель:", time.time() - start)
start = time.time()
res = frozen_model(x)
print("Замороженная модель:", time.time() - start)
Результат выполнения кода:
Обычная модель: 0.0012
Замороженная модель: 0.0008
Пример
Давайте используем замороженную модель для инференса:
import torch
import torch.nn as nn
torch.manual_seed(0)
model = nn.Linear(5, 3)
scripted_model = torch.jit.script(model)
frozen_model = torch.jit.freeze(scripted_model)
x = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0])
res = frozen_model(x)
print(res)
Результат выполнения кода:
tensor([ 0.3431, -0.3410, -0.2866], grad_fn=<AddBackward0>)
Смотрите также
-
функцию
jit.script,
которая преобразует модель в TorchScript -
функцию
jit.trace,
которая трассирует модель для создания TorchScript -
функцию
jit.save,
которая сохраняет TorchScript модель в файл -
функцию
jit.load,
которая загружает TorchScript модель из файла