PyTorch TorchScript / ONNX Export
After model training is complete, the model needs to be deployed to the production environment. PyTorch provides multiple model export methods, among whichTorchScriptandONNXare the two most commonly used formats.
This section details the principles, usage, and best practices of these two export methods.
Model export is the process of converting a PyTorch model into a format that can run on different platforms and frameworks. This is crucial for scenarios such as model deployment, mobile inference, and cross-framework migration.
1. TorchScript Basics
1.1 What is TorchScript
TorchScript is PyTorch's serialization format, which converts Python code into stand-alone C++ virtual machine code. TorchScript programs can run in environments without a Python interpreter.
Key features of TorchScript:
- Converts dynamic graphs to static graphs
- Supports a subset of Python syntax
- Can run in a C++ environment
- Preserves the model's structure and parameters
1.2 Two Conversion Methods
TorchScript provides two ways to convert models to TorchScript:
- TorchScript Tracing (Tracing): Records operations by executing the model and generates a static computation graph
- TorchScript Scripting (Scripting): Directly analyzes Python code and compiles it to TorchScript
2. TorchScript Tracing (Tracing)
2.1 Basic Tracing Methods
Example
import torch.nn as nn
# ── Define model ──────────────────────────────────────
class SimpleNet(nn.Module):
def __init__(self):
super().__init__()
self.conv1 = nn.Conv2d(3, 64, kernel_size=3, padding=1)
self.conv2 = nn.Conv2d(64, 128, kernel_size=3, padding=1)
self.pool = nn.MaxPool2d(2)
self.fc = nn.Linear(128 * 8 * 8, 10)
def forward(self, x):
x = self.pool(torch.relu(self.conv1(x)))
x = self.pool(torch.relu(self.conv2(x)))
x = x.view(x.size(0), -1)
x = self.fc(x)
return x
# Create model instance
model = SimpleNet()
model.eval()
# Example input
example_input = torch.randn(1, 3, 64, 64)
# ── Method 1: torch.jit.trace tracing ───────────────────
# Create TorchScript by running the model and recording operations
traced_model = torch.jit.trace(model, example_input)
print("Traced model:")
print(traced_model)
# Save the model
traced_model.save("simple_net_traced.pt")
# Load the model
loaded_model = torch.jit.load("simple_net_traced.pt")
# Perform inference with the loaded model
output = loaded_model(example_input)
print(f"Output shape: {output.shape}")
2.2 Limitations of Tracing
The tracing method has some limitations:
- Only records operations that are actually executed
- Control flow (e.g., if, for) is fixed
- Inputs with dynamic sizes may cause issues
Example
class DynamicModel(nn.Module):
def __init__(self):
super().__init__()
def forward(self, x):
# Control flow is fixed during tracing
if x.sum() > 0:
return x * 2
else:
return x / 2
model = DynamicModel()
model.eval()
# During tracing, the if branch is fixed to the path taken during tracing
example_input = torch.tensor([1.0])
traced = torch.jit.trace(model, example_input)
# Even with different input, the same branch is executed
print(traced(torch.tensor([1.0]))) # 2
print(traced(torch.tensor([-1.0]))) # The result is still 2, not -0.5
For models containing control flow, you should use TorchScript Scripting instead of tracing.
3. TorchScript Scripting (Scripting)
3.1 Basic Usage
Example
# ── Using the @torch.jit.script decorator ───────────────
@torch.jit.script
def scripted_function(x: torch.Tensor) -> torch.Tensor:
"""Convert function using scripting"""
if x.sum() > 0:
return x * 2
else:
return x / 2
# Test the scripted function
input1 = torch.tensor([1.0, 2.0])
input2 = torch.tensor([-1.0, -2.0])
print(scripted_function(input1)) # [2., 4.]
print(scripted_function(input2)) # [-0.5, -1.]
# ── Scripting the model ───────────────────────────────────
class ScriptableModel(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(10, 10)
@torch.jit.export
def forward(self, x: torch.Tensor) -> torch.Tensor:
return torch.relu(self.fc(x))
@torch.jit.export
def predict(self, x: torch.Tensor) -> torch.Tensor:
"""Additional export method"""
out = self.forward(x)
return torch.argmax(out, dim=1)
model = ScriptableModel()
# Script the model
scripted_model = torch.jit.script(model)
print(scripted_model)
# Save
scripted_model.save("scripted_model.pt")
3.2 Scripting Complex Models
Example
class ConditionModel(nn.Module):
def __init__(self, num_classes: int):
super().__init__()
self.num_classes = num_classes
self.features = nn.Sequential(
nn.Conv2d(3, 32, 3, padding=1),
nn.ReLU(),
nn.MaxPool2d(2),
nn.Conv2d(32, 64, 3, padding=1),
nn.ReLU(),
nn.AdaptiveAvgPool2d(1),
nn.Flatten()
)
self.classifier = nn.Linear(64, num_classes)
def forward(self, x: torch.Tensor, use_softmax: bool = False) -> torch.Tensor:
"""
A model that supports dynamic conditions
"""
features = self.features(x)
logits = self.classifier(features)
if use_softmax:
return torch.softmax(logits, dim=1)
else:
return logits
def get_prediction(self, x: torch.Tensor) -> torch.Tensor:
"""Helper method"""
logits = self.forward(x, use_softmax=False)
return torch.argmax(logits, dim=1)
# Scripting
model = ConditionModel(num_classes=10)
scripted_model = torch.jit.script(model, example_inputs=(torch.randn(1, 3, 32, 32),))
# Test
test_input = torch.randn(2, 3, 32, 32)
output1 = scripted_model(test_input, use_softmax=False)
output2 = scripted_model(test_input, use_softmax=True)
print(f"Logits output shape: {output1.shape}")
print(f"Softmax output shape: {output2.shape}")
4. ONNX Export
4.1 ONNX Basics
ONNX(Open Neural Network eXchange)ONNX is an open neural network exchange format that supports converting models between different deep learning frameworks.
Advantages of ONNX:
- Cross-framework: Supported by PyTorch, TensorFlow, Caffe2, etc.
- Cross-platform: Supports CPU, GPU, mobile, and other platforms
- Hardware optimization: Can leverage ONNX Runtime for efficient inference
- Rich tools: A wealth of optimization tools and deployment solutions
4.2 Basic ONNX Export
Example
import torch.nn as nn
import torchvision
# ── Define model ──────────────────────────────────────
class ImageClassifier(nn.Module):
def __init__(self):
super().__init__()
self.features = nn.Sequential(
nn.Conv2d(3, 32, 3, padding=1),
nn.ReLU(),
nn.MaxPool2d(2),
nn.Conv2d(32, 64, 3, padding=1),
nn.ReLU(),
nn.AdaptiveAvgPool2d(1),
nn.Flatten()
)
self.classifier = nn.Linear(64, 10)
def forward(self, x):
x = self.features(x)
x = self.classifier(x)
return x
model = ImageClassifier()
model.eval()
# Example input
example_input = torch.randn(1, 3, 32, 32)
# ── Export to ONNX ──────────────────────────────────
output_path = "image_classifier.onnx"
torch.onnx.export(
model,
example_input,
output_path,
export_params=True, # Export model parameters
opset_version=14, # ONNX version
do_constant_folding=True, # Constant folding optimization
input_names=['input'], # Input tensor names
output_names=['output'], # Output tensor names
dynamic_axes={
'input': {0: 'batch_size'}, # Dynamic batch dimension
'output': {0: 'batch_size'}
}
)
print(f"Model exported to: {output_path}")
# Validate the exported model
import onnx
onnx_model = onnx.load(output_path)
onnx.checker.check_model(onnx_model)
print("ONNX model validation passed!")
4.3 Exporting Complex Models
Example
import torchvision.models as models
# Load a pretrained model
model = models.resnet18(pretrained=True)
model.eval()
# Example input (standard ResNet 224x224)
example_input = torch.randn(1, 3, 224, 224)
# Export
torch.onnx.export(
model,
example_input,
"resnet18.onnx",
export_params=True,
opset_version=14,
input_names=['input'],
output_names=['output'],
dynamic_axes={
'input': {0: 'batch_size', 2: 'height', 3: 'width'},
'output': {0: 'batch_size'}
}
)
print("ResNet18 has been exported")
# Handle special layers during export
# For layers that need special handling, use a workaround
class ModelWithSpecialOps(nn.Module):
def __init__(self):
super().__init__()
self.conv = nn.Conv2d(3, 64, 3, padding=1)
def forward(self, x):
# Use F.interpolate instead of the nn.functional alias
x = self.conv(x)
x = torch.nn.functional.interpolate(x, scale_factor=2, mode='bilinear', align_corners=False)
return x
model = ModelWithSpecialOps()
model.eval()
torch.onnx.export(
model,
torch.randn(1, 3, 32, 32),
"special_ops.onnx",
input_names=['input'],
output_names=['output'],
opset_version=14
)
5. Export Validation and Optimization
5.1 Validating Exported Models
Example
import onnx
import onnxruntime as ort
def verify_onnx_model(onnx_path, pytorch_model, test_input):
"""
Verify the output consistency between the ONNX model and the PyTorch model
"""
# 1. Check the ONNX model
onnx_model = onnx.load(onnx_path)
onnx.checker.check_model(onnx_model)
print("✓ ONNX model structure validation passed")
# 2. PyTorch inference
pytorch_model.eval()
with torch.no_grad():
pytorch_output = pytorch_model(test_input)
# 3. ONNX Runtime inference
ort_session = ort.InferenceSession(onnx_path)
ort_output = ort_session.run(None, {'input': test_input.numpy()})[0]
# 4. Compare outputs
diff = np.abs(pytorch_output.numpy() - ort_output).max()
print(f"✓ Maximum output difference: {diff:.6f}")
if diff < 1e-5:
print("✓ Output validation passed")
return True
else:
print("⚠ Output difference is large")
return False
# Usage example
model = ImageClassifier()
model.eval()
# Export first
torch.onnx.export(
model,
torch.randn(1, 3, 32, 32),
"test_model.onnx",
input_names=['input'],
output_names=['output']
)
# Validate
test_input = torch.randn(2, 3, 32, 32)
verify_onnx_model("test_model.onnx", model, test_input)
5.2 ONNX Runtime Optimization
Example
from onnxruntime import GraphOptimizationLevel
# Create inference session and configure optimization
def create_optimized_session(onnx_path, providers=None):
if providers is None:
providers = ['CUDAExecutionProvider', 'CPUExecutionProvider']
# Configure session options
sess_options = ort.SessionOptions()
# Enable graph optimization
sess_options.graph_optimization_level = GraphOptimizationLevel.ORT_ENABLE_ALL
# Enable other optimizations
sess_options.intra_op_num_threads = 4
sess_options.inter_op_num_threads = 4
# Create session
session = ort.InferenceSession(onnx_path, sess_options, providers=providers)
return session
# Use the optimized session for inference
session = create_optimized_session("resnet18.onnx")
# Get input and output names
input_name = session.get_inputs()[0].name
output_name = session.get_outputs()[0].name
# Inference
input_data = np.random.randn(1, 3, 224, 224).astype(np.float32)
output = session.run([output_name], {input_name: input_data})[0]
print(f"Inference output shape: {output.shape}")
5.3 Model Quantization
Example
import onnx
from onnxruntime.quantization import quantize_dynamic, QuantType
def quantize_onnx_model(input_path, output_path):
"""
Dynamically quantize ONNX model
Reduce model size and accelerate inference
"""
# Dynamic quantization (no calibration data required)
quantize_dynamic(
input_path,
output_path,
# Specify the weight types to quantize
weight_type=QuantType.QInt8
)
# Get file size comparison
import os
original_size = os.path.getsize(input_path) / 1024 / 1024
quantized_size = os.path.getsize(output_path) / 1024 / 1024
print(f"Original model: {original_size:.2f} MB")
print(f"Quantized model: {quantized_size:.2f} MB")
print(f"Compression ratio: {original_size/quantized_size:.2f}x")
# Usage
quantize_onnx_model("resnet18.onnx", "resnet18_quantized.onnx")
6. Mobile Deployment
6.1 Exporting to Mobile Formats
Example
# 1. Trace the model
model = ImageClassifier()
model.eval()
example_input = torch.randn(1, 3, 32, 32)
traced_model = torch.jit.trace(model, example_input)
# 2. Optimize the model (mobile optimization)
optimized_model = torch.jit.optimize_for_inference(traced_model)
# 3. Save the mobile model
optimized_model.save("mobile_model.pt")
# Method 2: Use torchmobile (if available)
# torchmobile is used for lighter mobile deployment
# Please refer to the official documentation for cross-compilation
# Method 3: Export to ONNX and use a mobile runtime
# iOS: Use Core ML Tools
# Android: Use ONNX Runtime Mobile
print("Mobile model is ready")
6.2 Core ML Export (iOS)
For the iOS platform, you can use Core ML Tools to convert the model to Core ML format. You need to first convert it to ONNX format, then use Core ML Tools for conversion.
Example
# pip install coremltools
# Core ML export steps
# Note: This is only available on macOS
# Step 1: First export the PyTorch model to ONNX
# from pytorch_example import your_model
# model = your_model()
# model.eval()
# example_input = torch.randn(1, 3, 224, 224)
# torch.onnx.export(
# model,
# example_input,
# "model.onnx",
# input_names=['input'],
# output_names=['output']
# )
# Step 2: Use Core ML Tools to convert to Core ML format
# from coremltools.converters import onnx as onnx_coreml
# coreml_model = onnx_coreml.convert(
# model="model.onnx",
# minimum_deployment_target='13',
# image_input_names=['input']
# )
# coreml_model.save("model.mlmodel")
# Usage example
print("iOS deployment requires macOS environment")
print("Detailed steps: 1. pip install coremltools 2. Export ONNX 3. Convert using coremltools")
7. Best Practices and Common Issues
7.1 Export Troubleshooting
| Problem | Cause | Solution |
|---|---|---|
| Export failed | Unsupported operation | Use a higher opset_version, or replace the operation |
| Inconsistent output | Dynamic control flow | Use scripting instead of tracing, or fix the input size |
| ONNX Runtime error | Operator not supported | Check the list of operations supported by ONNX |
7.2 Comparison of Export Formats
| Format | Advantages | Disadvantages | Applicable scenarios |
|---|---|---|---|
| TorchScript (.pt) | Native PyTorch support | Poor cross-framework support | Desktop/server deployment |
| ONNX (.onnx) | Cross-framework, cross-platform | Some operations are not supported | General deployment |
7.3 Export Checklist
Example
def export_checklist(model, example_input):
"""Pre-export checks"""
model.eval()
# 1. Ensure the model is in inference mode
print(f"Model mode: {'eval' if not model.training else 'train'}")
# 2. Check input size
print(f"Example input shape: {example_input.shape}")
# 3. Validate forward pass
with torch.no_grad():
output = model(example_input)
print(f"Output shape: {output.shape}")
# 4. Check for non-serializable operations
# E.g., lambda functions
for name, module in model.named_modules():
if hasattr(module, '__call__'):
pass # Check custom modules
print("✓ Export check completed")
# Usage example
model = ImageClassifier()
example_input = torch.randn(1, 3, 32, 32)
export_checklist(model, example_input)
7.4 Complete Export Process
Example
class ProductionModel(nn.Module):
"""Production-grade model"""
def __init__(self):
super().__init__()
self.backbone = nn.Sequential(
nn.Conv2d(3, 64, 7, stride=2, padding=3),
nn.BatchNorm2d(64),
nn.ReLU(),
nn.MaxPool2d(3, stride=2, padding=1),
# ... more layers
)
self.head = nn.Linear(512, 10)
def forward(self, x, return_features=False):
features = self.backbone(x)
features = torch.nn.functional.adaptive_avg_pool2d(features, 1).flatten(1)
if return_features:
return features
return self.head(features)
# Complete export workflow
model = ProductionModel()
model.load_state_dict(torch.load("weights.pth")) # Load weights
model.eval()
# Prepare example input
example_input = torch.randn(1, 3, 224, 224)
# 1. Trace TorchScript
traced = torch.jit.trace(model, example_input)
traced.save("model_traced.pt")
# 2. Export ONNX
torch.onnx.export(
traced,
example_input,
"model.onnx",
input_names=['input'],
output_names=['output'],
dynamic_axes={'input': {0: 'batch'}, 'output': {0: 'batch'}},
opset_version=14
)
# 3. Optimize ONNX
# Use onnx-simplifier to remove redundant operations
# onnxsim model.onnx model_simplified.onnx
print("All formats exported successfully!")