Метод save класса ScriptModule
Метод save класса ScriptModule сохраняет
сериализованное представление скриптового модуля в файл.
Первым параметром метод принимает путь к файлу,
вторым параметром можно передать словарь с дополнительными
данными для сохранения.
Синтаксис
module.save(file_path, [extra_files])
Пример
Давайте создадим простой скриптовый модуль и сохраним его в файл:
import torch
class MyModule(torch.nn.Module):
def forward(self, x):
return x + 1
module = torch.jit.script(MyModule())
module.save('model.pt')
print("Module saved successfully")
Результат выполнения кода:
"Module saved successfully"
Пример
Сохраним модуль с дополнительными данными с помощью параметра extra_files:
import torch
class MyModule(torch.nn.Module):
def forward(self, x):
return x * 2
module = torch.jit.script(MyModule())
extra_files = {'metadata': 'This is important info'}
module.save('model_with_extra.pt', extra_files)
print("Module with extra files saved")
Результат выполнения кода:
"Module with extra files saved"
Пример
Сохраним более сложный модуль с параметрами и проверим, что сохранение прошло успешно:
import torch
class LinearModule(torch.nn.Module):
def __init__(self):
super().__init__()
self.linear = torch.nn.Linear(5, 3)
def forward(self, x):
return self.linear(x)
torch.manual_seed(0)
module = torch.jit.script(LinearModule())
module.save('linear_model.pt')
print("Linear module saved")
Результат выполнения кода:
"Linear module saved"
Смотрите также
-
класс
ScriptModule,
который представляет скриптовую версию модуля PyTorch -
метод
graph,
который возвращает граф вычислений скриптового модуля -
метод
code,
который возвращает код скриптового модуля в виде строки -
функцию
save,
которая сохраняет скриптовый модуль в файл