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
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

# Example of tracing limitations

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

import torch

# ── 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

# More complex scripting example: a model with conditionals

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
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

# Export the complete ResNet model

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 numpy as np
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("&#x26a0; 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

import onnxruntime as ort
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

# ONNX model quantization

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

# Method 1: TorchScript mobile deployment
# 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

# Need to install coremltools
# 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

# Pre-export checklist

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

# Complete model export workflow 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!")
Other Extensions