Функция export.export
Функция export.export из модуля torch.export
предназначена для преобразования модели PyTorch
в оптимизированное представление, пригодное для
инференса (выполнения модели). Эта функция
выполняет статическую диспетчеризацию операций и
создаёт вычислительный граф, который впоследствии
может быть сохранён и загружен без необходимости
наличия исходного Python-кода. Основная цель -
достичь максимальной производительности при
выполнении модели на целевой платформе.
Синтаксис
torch.export.export(
model,
args,
kwargs=None,
constraints=None,
strict=True,
)
Основные параметры функции:
⁅i⁆model⁅/i⁆ (torch.nn.Module) - экспортируемая модель.
⁅i⁆args⁅/i⁆ (tuple) - кортеж позиционных аргументов для прогона модели.
⁅i⁆kwargs⁅/i⁆ (dict, опционально) - словарь именованных аргументов.
⁅i⁆constraints⁅/i⁆ (list, опционально) - список ограничений на размеры динамических измерений.
⁅i⁆strict⁅/i⁆ (bool, по умолчанию True) - если флаг включён, экспорт завершится ошибкой при обнаружении неподдерживаемых операций.
Пример
Давайте создадим простую линейную модель и
экспортируем её с помощью функции
export.export:
import torch
import torch.nn as nn
class SimpleModel(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(5, 2)
def forward(self, x):
return self.fc(x)
model = SimpleModel()
example_input = torch.randn(3, 5)
exported_program = torch.export.export(
model,
(example_input,),
)
print(exported_program)
Результат выполнения кода:
ExportedProgram(
graph_module=GraphModule(...),
graph=Graph(...),
)
Пример
Экспортируем модель с передачей именованных аргументов:
import torch
import torch.nn as nn
class CustomModel(nn.Module):
def forward(self, x, bias=None):
if bias is not None:
return x + bias
return x
model = CustomModel()
x = torch.ones(2, 3)
bias = torch.tensor([1.0, 2.0, 3.0])
exported_program = torch.export.export(
model,
(x,),
kwargs={"bias": bias},
)
print("Model exported successfully")
Результат выполнения кода:
"Model exported successfully"
Пример
Используем экспортированную программу для выполнения на новых данных:
import torch
import torch.nn as nn
model = nn.Linear(4, 2)
example_input = torch.randn(2, 4)
exported_program = torch.export.export(
model,
(example_input,),
)
new_input = torch.randn(3, 4)
res = exported_program.module()(new_input)
print(res.shape)
Результат выполнения кода:
torch.Size([3, 2])
Смотрите также
-
функцию
export.save,
которая сохраняет экспортированную программу в файл -
функцию
export.load,
которая загружает экспортированную программу из файла -
функцию
jit.script,
которая компилирует модель в TorchScript -
функцию
onnx.export,
которая экспортирует модель в формат ONNX