Функция jit.save
Функция jit.save сохраняет скомпилированную модель PyTorch в файл на диске.
Первым параметром функция принимает модель, созданную с помощью jit.script или jit.trace.
Вторым параметром передаётся путь к файлу для сохранения.
Функция сохраняет как архитектуру модели, так и её параметры, что позволяет загрузить модель без исходного кода.
Синтаксис
torch.jit.save(model, file_path)
Пример
Давайте создадим простую модель и сохраним её с помощью jit.save:
import torch
import torch.nn as nn
class SimpleModel(nn.Module):
def __init__(self):
super(SimpleModel, self).__init__()
self.fc = nn.Linear(10, 5)
def forward(self, x):
return self.fc(x)
model = SimpleModel()
scripted_model = torch.jit.script(model)
torch.jit.save(scripted_model, 'model.pt')
print("Model saved successfully")
Результат выполнения кода:
"Model saved successfully"
Пример
Давайте сохраним модель, созданную с помощью jit.trace:
import torch
import torch.nn as nn
class TraceModel(nn.Module):
def forward(self, x):
return x * 2 + 1
model = TraceModel()
example_input = torch.randn(1, 3)
traced_model = torch.jit.trace(model, example_input)
torch.jit.save(traced_model, 'traced_model.pt')
print("Traced model saved")
Результат выполнения кода:
"Traced model saved"
Пример
Давайте сохраним модель с использованием дополнительных параметров, например с включением дополнительной информации:
import torch
import torch.nn as nn
class ExtraModel(nn.Module):
def __init__(self):
super(ExtraModel, self).__init__()
self.linear = nn.Linear(5, 2)
def forward(self, x):
return self.linear(x)
model = ExtraModel()
scripted_model = torch.jit.script(model)
extra_files = {'metadata': 'model_info'}
torch.jit.save(scripted_model, 'model_with_meta.pt', _extra_files=extra_files)
print("Model with metadata saved")
Результат выполнения кода:
"Model with metadata saved"
Смотрите также
-
функцию
jit.load,
которая загружает сохранённую скомпилированную модель -
функцию
jit.script,
которая компилирует модель в представление TorchScript -
функцию
jit.trace,
которая трассирует модель на основе примера входных данных -
функцию
save,
которая сохраняет обычную модель PyTorch в формате state_dict