How to Use the Checkpoint Module for Saving and Loading Models in WeatherNext
Use weathernext.utils.checkpoint.dump() to serialize model state to a .npz file and load() to restore it with type-safe schema validation.
The checkpoint module in the WeatherNext repository provides a lightweight, portable way to persist and recover tree-like structures containing NumPy arrays, scalars, and nested dataclasses. Unlike standard pickling, this implementation produces pure NumPy archives that are safe to share and load without executing arbitrary code.
Core API: dump() and load()
The module exposes two public functions in weathernext/utils/checkpoint.py:
dump(dest, value)(lines 26–38): Flattens any nested structure and writes it as a NumPy.npzarchive.load(source, typ)(lines 42–55): Reads an.npzarchive and reconstructs the original structure using a provided type schema.
Both functions operate on binary file-like objects, making them compatible with open(), io.BytesIO, or cloud storage streams.
How Checkpoint Serialization Works
Step 1: Flattening the Tree
The internal _flatten() function (lines 60–81) recursively traverses dataclasses, dictionaries, lists, and tuples. It concatenates nested keys using a colon separator (:) to produce a flat dictionary:
# A nested dataclass becomes flat keys
{"layer:weights": array(...), "layer:bias": 0.5}
Leaves are stored directly as NumPy arrays or scalar values.
Step 2: Unflattening and Type Conversion
On load, _unflatten() (lines 84–95) splits keys on : and rebuilds the nested hierarchy. Then _convert_types() (lines 98–120) walks the provided type schema to cast each leaf to the correct Python or NumPy type.
This schema-driven deserialization guarantees that the restored object matches the original structure exactly—including optional fields in dataclasses, which default to None when missing.
Saving a Model Checkpoint
Here is how to save a simple dataclass-based model using the checkpoint module:
import io
import numpy as np
from dataclasses import dataclass
from weathernext.utils import checkpoint
@dataclass
class SimpleModel:
weights: np.ndarray
bias: float
description: str | None = None
# Create and populate a model
model = SimpleModel(
weights=np.random.randn(10, 5),
bias=0.42,
description="Demo checkpoint"
)
# Serialize to disk
with open("model.ckpt.npz", "wb") as f:
checkpoint.dump(f, model)
The dump() call handles the optional description field automatically—optional fields are fully supported without extra configuration.
Loading a Model Checkpoint
To restore a checkpoint, provide the same type as a schema to load():
from weathernext.utils import checkpoint
with open("model.ckpt.npz", "rb") as f:
restored = checkpoint.load(f, SimpleModel)
assert isinstance(restored, SimpleModel)
print(restored.weights.shape) # (10, 5)
print(restored.bias) # 0.42
The type schema ensures that restored is not just a dictionary with the right keys, but an actual SimpleModel instance with properly typed fields.
Working with Nested Containers
The checkpoint module handles arbitrarily nested structures. This example mixes dictionaries, lists, and dataclasses:
from dataclasses import dataclass
from typing import Dict, List, Union
import numpy as np
from weathernext.utils import checkpoint
@dataclass
class Layer:
kernel: np.ndarray
bias: np.ndarray
# Nested structure: dict → list → dataclass
network = {
"encoder": [
Layer(kernel=np.eye(64), bias=np.zeros(64)),
Layer(kernel=np.ones((64, 64)), bias=np.zeros(64)),
],
"decoder": {"final": Layer(kernel=np.eye(10), bias=np.zeros(10))},
}
# Define the type schema for loading
NetworkType = Dict[
str,
Union[List[Layer], Dict[str, Layer]]
]
# Save and load
with open("network.ckpt.npz", "wb") as f:
checkpoint.dump(f, network)
with open("network.ckpt.npz", "rb") as f:
restored_network = checkpoint.load(f, NetworkType)
print(restored_network["encoder"][0].kernel.shape) # (64, 64)
Note that _unflatten() preserves list and tuple ordering by sorting flat keys numerically.
Key Advantages Over Standard Pickling
| Feature | weathernext.utils.checkpoint |
pickle |
|---|---|---|
| Security | Pure NumPy .npz format, no code execution |
Arbitrary code execution risk |
| Portability | Cross-language compatible | Python-only |
| Type safety | Schema-enforced deserialization on load | No structural validation |
| Optional fields | Automatic None handling for missing dataclass fields |
Manual handling required |
Summary
- Use
checkpoint.dump(f, obj)to serialize any tree of NumPy arrays, scalars, dataclasses, and containers to a.npzfile. - Use
checkpoint.load(f, Type)to restore with guaranteed type safety via schema-driven deserialization. - The implementation in
weathernext/utils/checkpoint.pyavoids pickling entirely, producing portable archives suitable for production ML pipelines. - Optional dataclass fields, nested lists, and mixed containers are all handled automatically.
Frequently Asked Questions
What file format does the checkpoint module use?
The checkpoint module writes standard NumPy .npz archives. These are ZIP files containing individual .npy arrays, readable by any language with NumPy support. No Python-specific pickle protocol is used.
Can I checkpoint JAX or PyTorch tensors directly?
No—the module only accepts NumPy arrays and Python scalars. Convert JAX arrays with np.array(x) and PyTorch tensors with .numpy() before calling dump(). On load, you can convert back to your framework's tensor type in post-processing.
Why does load() require a type parameter?
The type parameter enables schema-driven deserialization. Without it, the module could only return generic dictionaries. By providing the original dataclass or container types, _convert_types() reconstructs the exact original structure with proper field types and optional field handling.
How does the module handle large models that exceed memory?
The current implementation loads the entire .npz into memory. For out-of-core checkpoints, you would need to extend weathernext/utils/checkpoint.py with chunked array support using numpy.memmap or similar techniques.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →