# How to Use the Checkpoint Module for Saving and Loading Models in WeatherNext

> Learn to save and load WeatherNext models using the checkpoint module. Easily serialize with dump() and restore with load() for efficient model management.

- Repository: [Google DeepMind/weathernext](https://github.com/google-deepmind/weathernext)
- Tags: how-to-guide
- Published: 2026-08-12

---

**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`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/checkpoint.py):

- **`dump(dest, value)`** (lines 26–38): Flattens any nested structure and writes it as a NumPy `.npz` archive.
- **`load(source, typ)`** (lines 42–55): Reads an `.npz` archive 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:

```python

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

```python
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()`:

```python
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:

```python
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 `.npz` file.
- Use `checkpoint.load(f, Type)` to restore with guaranteed type safety via schema-driven deserialization.
- The implementation in [`weathernext/utils/checkpoint.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/checkpoint.py) avoids 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`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/checkpoint.py) with chunked array support using `numpy.memmap` or similar techniques.