Функция export.save
Функция export.save сохраняет экспортированную программу, полученную с помощью export.export, в файл на диске. Первым параметром функция принимает экспортированную программу (объект ExportedProgram), вторым - путь к файлу для сохранения. Сохранённый файл можно впоследствии загрузить с помощью функции export.load и выполнить на целевой платформе без необходимости повторного экспорта.
Синтаксис
torch.export.save(exported_program, file_path)
Пример
Давайте экспортируем простую модель и сохраним её в файл:
import torch
class MyModel(torch.nn.Module):
def __init__(self):
super().__init__()
self.linear = torch.nn.Linear(10, 5)
def forward(self, x):
return self.linear(x)
model = MyModel()
example_inputs = (torch.randn(3, 10),)
exported_program = torch.export.export(model, example_inputs)
torch.export.save(exported_program, "model.pt")
Пример
Давайте сохраним экспортированную программу с указанием дополнительных опций:
import torch
class MyModel(torch.nn.Module):
def __init__(self):
super().__init__()
self.linear = torch.nn.Linear(5, 2)
def forward(self, x):
return self.linear(x)
model = MyModel()
example_inputs = (torch.randn(2, 5),)
exported_program = torch.export.export(model, example_inputs)
torch.export.save(exported_program, "model.pt")
print("Model saved successfully")
Результат выполнения кода:
"Model saved successfully"
Пример
Давайте сохраним экспортированную программу с динамическими размерами:
import torch
class DynamicModel(torch.nn.Module):
def forward(self, x, y):
return x + y
model = DynamicModel()
example_inputs = (torch.randn(3, 4), torch.randn(3, 4))
dynamic_shapes = {
"x": {0: torch.export.Dim("batch")},
"y": {0: torch.export.Dim("batch")},
}
exported_program = torch.export.export(
model,
example_inputs,
dynamic_shapes=dynamic_shapes
)
torch.export.save(exported_program, "dynamic_model.pt")
print("Dynamic model saved")
Результат выполнения кода:
"Dynamic model saved"
Смотрите также
-
функцию
export.load,
которая загружает ранее сохранённую экспортированную программу -
функцию
export.export,
которая выполняет экспорт модели PyTorch -
функцию
save,
которая сохраняет объекты PyTorch в файл -
функцию
load,
которая загружает объекты PyTorch из файла