Функция onnx.export
Функция onnx.export экспортирует модель PyTorch в формат ONNX (Open Neural Network Exchange).
Первым параметром передается модель, вторым - кортеж с примерами входных данных,
третьим - путь для сохранения файла.
Дополнительные параметры позволяют задать имя модели, тип данных, opset версию и другие настройки.
Синтаксис
torch.onnx.export(
model,
args,
f,
export_params=True,
verbose=False,
training=TrainingMode.EVAL,
input_names=None,
output_names=None,
opset_version=None,
dynamic_axes=None
)
Пример
Экспортируем простую модель с одним линейным слоем:
import torch
import torch.nn as nn
model = nn.Linear(10, 5)
model.eval()
dummy_input = torch.randn(1, 10)
torch.onnx.export(
model,
dummy_input,
"model.onnx",
export_params=True,
opset_version=11,
input_names=["input"],
output_names=["output"]
)
После выполнения в текущей директории появится файл model.onnx.
Пример
Экспортируем модель с динамической размерностью батча:
import torch
import torch.nn as nn
class SimpleModel(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(10, 5)
def forward(self, x):
return self.fc(x)
model = SimpleModel()
model.eval()
dummy_input = torch.randn(1, 10)
torch.onnx.export(
model,
dummy_input,
"dynamic_model.onnx",
input_names=["input"],
output_names=["output"],
dynamic_axes={
"input": {0: "batch_size"},
"output": {0: "batch_size"}
}
)
Пример
Экспортируем модель с несколькими входами и выходами:
import torch
import torch.nn as nn
class MultiInputModel(nn.Module):
def __init__(self):
super().__init__()
self.fc1 = nn.Linear(5, 3)
self.fc2 = nn.Linear(3, 2)
def forward(self, x1, x2):
out = self.fc1(x1) + x2
return self.fc2(out)
model = MultiInputModel()
model.eval()
x1 = torch.randn(2, 5)
x2 = torch.randn(2, 3)
torch.onnx.export(
model,
(x1, x2),
"multi_input.onnx",
input_names=["input1", "input2"],
output_names=["output"],
opset_version=11
)
Результатом будет файл ONNX, поддерживающий несколько входных тензоров.
Пример
Экспортируем модель с явным указанием типа данных и verbose режимом:
import torch
import torch.nn as nn
model = nn.Sequential(
nn.Linear(8, 16),
nn.ReLU(),
nn.Linear(16, 4)
)
model.eval()
dummy_input = torch.randn(1, 8, dtype=torch.float32)
torch.onnx.export(
model,
dummy_input,
"verbose_model.onnx",
verbose=True,
opset_version=12,
input_names=["features"],
output_names=["logits"]
)
При включенном verbose в консоль будет выведена отладочная информация о процессе экспорта.
Смотрите также
-
функцию
load,
которая загружает сохраненную модель из файла -
функцию
save,
которая сохраняет модель или тензор в формате PyTorch -
функцию
jit.trace,
которая создает TorchScript модуль трассировкой -
функцию
jit.script,
которая компилирует модель в TorchScript