Efficient Serialization Formats for AI Model Inference
Saving a trained neural network requires converting its architecture and learned parameters into a durable format suitable for storage and later restoration. This process, known as model serialization, is a critical first step in model deployment. This article examines serialization strategies within inference engines, covering both generic and custom approaches, and provides a detailed look at two widely used binary formats: Protocol Buffers and FlatBuffers.
Fundamentals of Model Persistence
Data held in memory is ephemeral. To reuse a trained model, its state must be written to persistent storage. The corresponding processes are:
- Serialization: Translating an in-memory model object into a byte stream for storage.
- Deserialization: Reconstructing the model object in memory from a stored byte stream.
Categories of Serialization Approaches
1. Standard Cross-Platform Formats Several established formats enable interoperability across different systems and programming languages:
- Text-based: XML and JSON are human-readable but generally less efficient for large numerical arrays.
- Binary-based: Protocol Buffers (protobuf) and FlatBuffers provide compact, high-performance alternatives. These are the preferred choices for production deployment, as seen in ONNX's use of Protobuf for exchange and Apple's CoreML for on-device execution.
2. Framework-Specific Methods Machine learning frameworks often implement custom serialization optimized for their internal structures. These can be text or binary.
- Binary examples: TensorFlow's checkpoint files and PyTorch's saved state dictionaries are proprietary formats designed for rapid I/O of large parameter sets.
- Portability: While some framework formats are language-specific (like Scikit-learn's
.pklfiles), others aim for cross-platform use, with ONNX being a prime example of an interoperable standard.
3. Native Language Serialization General-purpose programming languages provide built-in mechanisms:
- Python: The
picklemodule can serialize arbitrary objects, whilejoblibis optimized for NumPy arrays. - R: Functions like
saveproduce.rdafiles for persisting R objects.
4. Fully Custom Serialization When standard methods fall short, particularly for edge deployments with severe resource constraints, engineers may design a bespoke format. Key design drivers include minimizing parsing time, maximizing compression, and ensuring runtime compatibility. The major trade-off is the added burden of maintaining format versioning and backward compatibility.
PyTorch Model Export Techniques
Internal State Dicts
The native PyTorch method uses torch.save to store the model's state_dict, a Python dictionary mapping layers to their parameter tensors. This captures weights and optimizer states but not the computational graph.
# Saving weights for inference
model_path = "./saved_weights.pth"
torch.save(trained_model.state_dict(), model_path)
# Loading weights into a model skeleton
loaded_model = MyModelClass(*init_args, **init_kwargs)
state = torch.load(model_path, map_location=torch.device('cpu'))
loaded_model.load_state_dict(state)
loaded_model.eval()
When moving a model trained on a GPU to a CPU-only environmentt, the device mapping must be handled explicitly:
# Transferring a pre-trained model from GPU to CPU
model = torchvision.models.resnet50(pretrained=True)
model.to('cpu')
ONNX Export
PyTorch can also export a model to the ONNX format via torch.onnx.export. This traces the model's execution graph and bundles it with the parameters, enabling deployment on various runtimes.
# Export to ONNX
demo_input = torch.randn(1, 3, 256, 256).cuda()
onnx_filename = "model_graph.onnx"
torch.onnx.export(
my_model, demo_input, onnx_filename,
input_names=['pixel_input'],
output_names=['class_scores'],
dynamic_axes={'pixel_input': {0: 'batch_size'}}
)
Deep Dive into Binary Formats
Protocol Buffers (Protobuf)
Protobuf is a language-neutral, platform-neutral mechanism for serializing structured data. A schema is defined in a .proto file, from which data access classes are generated.
Schema Definition
A message definition includes fields with specific rules (optional, required, repeated), data types, names, and unique field numbers.
// Example message structure
message LayerConfig {
required string layer_type = 1;
optional int32 kernel_size = 2;
repeated float weight_data = 3;
}
Encoding Mechanism
Protobuf uses a binary TLV (Tag-Length-Value) encoding scheme. The Tag is a variant-encoded integer combining the field number and its wire type (tag = (field_number << 3) | wire_type). The Value is the encoded data. This structure alows parsers to skip unknown fields, enabling schema evolution.
In TensorFlow, decoding a protobuf could look like this:
# Deserializing a protobuf message from a tensor
decoded_data = tf.io.decode_proto(
bytes=serialized_tensor,
message_type="model.LayerConfig",
field_names=["layer_type", "kernel_size", "weight_data"],
output_types=[tf.string, tf.int32, tf.float32]
)
FlatBuffers
FlatBuffers represents a zero-copy evolution of serialization. Unlike Protobuf, which requires parsing into a separate in-memory object, FlatBuffers allows direct access to fields on the raw binary buffer.
Schema and Usage
A structure is described in a .fbs file. The key advantage is that no unpacking step is needed before reading data.
table InferenceSettings {
model_name: string;
batch_size: int = 1;
input_dims: [int];
}
root_type InferenceSettings;
AI Framework Adoption Several mobile and edge-centric inference engines utilize FlatBuffers for its low-overhead memory access:
- MNN: Alibaba's lightweight deep learning engine leverages FlatBuffers for its model file format to minimize loading latency.
- MindSpore Lite: Huawei's framework supports FlatBuffers-based models and provides a registry mechanism for extending node parsing and graph optimization during model conversion.
Comparative Analysis
| Feature | Protocol Buffers | FlatBuffers |
|---|---|---|
| Language Support | C/C++, Java, Python, Go, C#, etc. | C/C++, C#, Java, JavaScript, Rust, Python, etc. |
| Schema File | .proto | .fbs |
| Access Method | Parse to object first | Direct read on buffer (zero-copy) |
| Memory Allocation | Requires for decoding | Minimal (reads from the buffer) |
| Generated Code | Larger footprint | Compact, often a single header |