How to export PyTorch models to ONNX and run inference with ONNX Runtime?

Asked 22 days ago Updated 17 hours ago 132 views

0

Deploying PyTorch models directly into production environments often introduces heavy Python runtime dependencies and hardware abstraction overhead. Converting PyTorch neural networks into Open Neural Network Exchange format allows optimized execution across heterogeneous platforms using ONNX Runtime.

Model Serialization and Execution

Below is a working workflow illustrating how to export a computer vision or classification PyTorch module into ONNX format and execute inference sessions.

import torch
import torch.nn as nn
import onnxruntime as ort
import numpy as np

# Define a simple linear model module
class SimpleModel(nn.Module):
    def __init__(self):
        super().__init__()
        self.linear = nn.Linear(10, 2)
        
    def forward(self, x):
        return self.linear(x)

# Instantiate model and create dummy input tensor matching input shape
pytorch_model = SimpleModel()
pytorch_model.eval()
dummy_input = torch.randn(1, 10)

# Export PyTorch model graph to ONNX file
torch.onnx.export(
    pytorch_model,
    dummy_input,
    "model.onnx",
    input_names=["input"],
    output_names=["output"],
    dynamic_axes={"input": {0: "batch_size"}, "output": {0: "batch_size"}}
)

# Load exported ONNX model into ONNX Runtime session
ort_session = ort.InferenceSession("model.onnx")

# Run inference using numpy array format
ort_inputs = {"input": dummy_input.numpy()}
ort_outputs = ort_session.run(None, ort_inputs)
print("ONNX Inference output shape:", ort_outputs[0].shape)

1 Answer


0

You can export a trained PyTorch model to ONNX, then load the resulting file with ONNX Runtime for inference. The example below assumes model is your trained model and example_input is a CPU tensor shaped like a real model input.

import torch
import onnxruntime as ort

# Use evaluation behavior, such as disabling dropout.
model.eval()
example_input = example_input.detach().cpu()

# Export the model; the named input and output make the ONNX interface explicit.
with torch.no_grad():
    torch.onnx.export(
        model,
        example_input,
        "model.onnx",
        input_names=["input"],
        output_names=["output"],
        dynamic_axes={
            "input": {0: "batch_size"},
            "output": {0: "batch_size"},
        },
        opset_version=17,
        dynamo=False,
    )

# Create a CPU inference session and pass input data as a NumPy array.
session = ort.InferenceSession("model.onnx", providers=["CPUExecutionProvider"])
input_array = example_input.numpy()
onnoutput = session.run(["output"], {"input": input_array})[0]

# Compare against PyTorch to check that the exported model behaves as expected.
with torch.no_grad():
    torch_output = model(example_input).cpu().numpy()

print("ONNX output shape:", onnoutput.shape)
print("Outputs close:", torch.allclose(
    torch.from_numpy(onnoutput),
    torch.from_numpy(torch_output),
    rtol=1e-3,
    atol=1e-5,
))

Install the runtime with pip install onnxruntime; install onnx as well if you want to inspect or validate the exported file. For an NVIDIA GPU, use the compatible onnxruntime-gpu package and select an available execution provider instead of CPUExecutionProvider.

Things to check

  • Input shape and type: The example tensor should match the model’s expected layout and dtype. For image models, that often means float32 data in NCHW order.
  • Dynamic batch dimensions: The export marks dimension 0 as variable, so inference can use a different batch size. Other dimensions remain fixed.
  • Exporter version: This uses the legacy exporter because it supports the dynamic_axes argument shown. If your PyTorch version does not accept dynamo=False, remove that argument or follow the dynamic-shape syntax for your installed exporter.
  • Numerical differences: Small differences between PyTorch and ONNX Runtime can be normal. Larger differences may point to an unsupported operation, a preprocessing mismatch, or an incorrect input name.

You can check the actual interface with session.get_inputs() and session.get_outputs(). If those names differ from the ones used in the inference call, use the names reported by the session.

Write Your Answer