# How to Load Pre-Trained Model Weights from a Google Cloud Storage Bucket in WeatherNext

> Learn how to load pre-trained WeatherNext model weights from a Google Cloud Storage bucket using tfio gfile.Reconstruct your model parameters efficiently and speed up your development.

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

---

**Load WeatherNext model weights from GCS by opening a binary stream with `tf.io.gfile.GFile` and passing it to `weathernext.utils.checkpoint.load()` to reconstruct the dataclass-structured parameters.**

WeatherNext by Google DeepMind uses **serialized NumPy `.npz` files** for model checkpointing. These files can be streamed directly from Google Cloud Storage (GCS) without downloading to local disk. The `weathernext.utils.checkpoint` module provides the core deserialization logic that reconstructs typed weight structures from any file-like object.

## Prerequisites: Understanding the Checkpoint System

Before loading weights from a GCS bucket, you need to understand how WeatherNext structures its checkpoints.

### The Dataclass Schema Requirement

WeatherNext checkpoints encode **model parameters as nested dataclasses**, not raw arrays. When you call `checkpoint.load()`, you must provide the **exact dataclass type** used during training. This type-safe approach prevents silent shape mismatches.

Common dataclass patterns in the codebase include:

- `UnifiedModelParams` — for the main WeatherNext architecture
- `GraphForecasterParams` — for graph-based forecasters
- Custom dataclasses defined in model-specific modules

The deserialization logic in [`weathernext/utils/checkpoint.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/checkpoint.py) (lines 42–55) uses `tree_unflatten` to reconstruct the original structure from flat `.npz` keys.

## Step-by-Step: Load Pre-Trained Weights from GCS

### Step 1: Install Dependencies

Ensure TensorFlow is installed for GFile GCS support:

```bash
pip install tensorflow jax jaxlib flax

```

### Step 2: Import Required Modules

```python
import tensorflow as tf
from weathernext.utils import checkpoint

```

TensorFlow's `tf.io.gfile` provides **native GCS access** via the `gs://` protocol. No explicit authentication is needed when running on Google Cloud VMs or with Application Default Credentials configured.

### Step 3: Open the GCS Path and Deserialize

```python
from my_model_config import MyModelParams  # Your dataclass definition

gcs_path = "gs://weathernext-public-bucket/checkpoints/v1.0/params.npz"

with tf.io.gfile.GFile(gcs_path, "rb") as f:
    params = checkpoint.load(f, MyModelParams)

```

Key points about this operation:

- **Mode `"rb"`** is required — `.npz` files are binary archives
- `checkpoint.load()` accepts any file-like object with a `read()` method
- The function returns a fully-initialized dataclass instance, not a raw dict

### Step 4: Inject Weights Into Your Model

The injection pattern depends on your framework:

**For Flax/JAX models:**

```python
from flax import linen as nn

class MyModel(nn.Module):
    @nn.compact
    def __call__(self, x):
        # model definition

        pass

# Initialize with dummy input to get structure

model = MyModel()
key = jax.random.PRNGKey(0)
variables = model.init(key, dummy_input)

# Replace parameters with loaded weights

new_variables = variables.copy({"params": params})

```

**For pure function interfaces:**

Many WeatherNext models expose functional APIs where `params` is passed explicitly:

```python
predictions = model.apply({"params": params}, inputs)

```

## Complete Working Examples

### Loading Flax WeatherNext Parameters

```python
"""Load pre-trained WeatherNext weights from GCS into a Flax model."""
import tensorflow as tf
import jax
from weathernext.utils import checkpoint
from weathernext.models import WeatherNextParams, WeatherNextModel

def load_pretrained_from_gcs(gcs_uri: str, seed: int = 42):
    """
    Load WeatherNext parameters from a GCS path.
    
    Args:
        gcs_uri: Full GCS path, e.g., "gs://bucket/path/params.npz"
        seed: Random seed for model initialization (structure only)
    
    Returns:
        Tuple of (model, initialized_variables_with_loaded_params)
    """
    # Deserialize from GCS stream

    with tf.io.gfile.GFile(gcs_uri, "rb") as f:
        params = checkpoint.load(f, WeatherNextParams)
        print(f"Loaded parameters with {len(jax.tree_leaves(params))} arrays")
    
    # Initialize model structure

    model = WeatherNextModel()
    rng = jax.random.PRNGKey(seed)
    dummy_input = jax.numpy.zeros((1, 36, 721, 1440))  # (batch, levels, lat, lon)

    
    # Create variable tree, then substitute loaded parameters

    init_vars = model.init(rng, dummy_input)
    loaded_vars = jax.tree_map(
        lambda new, old: new if isinstance(new, jax.Array) else old,
        {"params": params},
        init_vars
    )
    
    return model, loaded_vars

# Usage

model, variables = load_pretrained_from_gcs(
    "gs://weather-datasets/weathernext/v2.0-params.npz"
)

```

### Loading for TensorFlow/Keras Custom Training

When checkpoints store raw weight dictionaries:

```python
"""Load WeatherNext weights into TensorFlow variables."""
import tensorflow as tf
from weathernext.utils import checkpoint

def load_weights_to_keras(gcs_path: str, model: tf.keras.Model):
    """
    Load .npz checkpoint and assign to Keras model variables.
    
    Assumes checkpoint keys match TF variable names (without :0 suffix).
    """
    with tf.io.gfile.GFile(gcs_path, "rb") as f:
        weight_dict = checkpoint.load(f, dict)  # Load as plain dictionary

    
    # Build name-to-variable mapping

    var_map = {
        var.name.replace(":0", ""): var 
        for var in model.trainable_variables
    }
    
    # Assign with validation

    assigned, missing = 0, []
    for name, array in weight_dict.items():
        if name in var_map:
            var_map[name].assign(array)
            assigned += 1
        else:
            missing.append(name)
    
    print(f"Assigned {assigned}/{len(weight_dict)} weight arrays")
    if missing:
        print(f"Warning: unmatched checkpoint keys: {missing[:5]}...")
    
    return assigned

```

## GCS Authentication Patterns

### Default Credentials (Recommended)

On Google Cloud infrastructure, authentication is automatic:

```python

# Works on GCE, Vertex AI, Cloud Run with appropriate service account

with tf.io.gfile.GFile("gs://private-bucket/model.npz", "rb") as f:
    params = checkpoint.load(f, MyParams)

```

### Explicit Credentials

For local development or cross-cloud access:

```python
import os
from google.cloud import storage

# Set before TensorFlow initializes GFile

os.environ["GOOGLE_APPLICATION_CREDENTIALS"] = "/path/to/key.json"

# Or use explicit client with GFile shim

client = storage.Client.from_service_account_json("/path/to/key.json")

```

### Public Bucket Access

For publicly readable buckets (no authentication):

```python

# anonymous=True bypasses credential requirements

tf.io.gfile.exists("gs://public-bucket/checkpoint.npz")  # True

```

## Troubleshooting Common Issues

| Issue | Cause | Solution |
|-------|-------|----------|
| `UnpicklingError` or `BadZipFile` | File not opened in binary mode | Use `"rb"`, not `"r"` |
| `KeyError` during load | Dataclass schema mismatch | Verify dataclass matches checkpoint version |
| `Permission denied` on GCS | Missing IAM permissions | Check service account has `storage.objectViewer` |
| Array shape mismatches | Model architecture changed | Compare `jax.tree_map(lambda x: x.shape, params)` against expected |
| Slow loading | Large checkpoint, no streaming | Ensure using `GFile` directly, not download-then-load |

## Key Files in the Repository

| File | Purpose | Critical Functions |
|------|---------|------------------|
| [`weathernext/utils/checkpoint.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/checkpoint.py) | Serialize/deserialize model trees | `load(file_obj, cls)`, `dump(file_obj, obj)` |
| [`weathernext/utils/checkpoint_test.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/checkpoint_test.py) | Round-trip validation tests | Examples of dataclass schema usage |
| `docs/weathernext2/wn2_demo.ipynb` | End-to-end inference demo | Adapts checkpoint loading for prediction |

The `checkpoint.load()` implementation ([source](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/checkpoint.py#L42-L55)) uses:

1. `numpy.load()` to read the `.npz` archive
2. `tree_unflatten` with the provided dataclass to reconstruct nested structure
3. Validation that all expected fields are present in the archive

## Performance Considerations

- **Streaming vs. download**: `tf.io.gfile.GFile` streams bytes on demand; for multi-GB checkpoints, consider local SSD caching with `gsutil cp` followed by local file load
- **JAX device placement**: Loaded arrays are CPU-backed; use `jax.device_put(params, jax.devices('gpu')[0])` to transfer to accelerator
- **Checkpoint sharding**: Large models may use sharded checkpoints; WeatherNext's checkpoint format supports this via directory-based `.npz` collections

## Summary

- **WeatherNext stores weights as `.npz` archives** with dataclass-structured schemas
- **Use `tf.io.gfile.GFile` to stream directly from GCS** — no intermediate download required
- **Call `checkpoint.load(file_obj, YourDataclass)`** to reconstruct typed parameters
- **Match the exact dataclass** used during checkpoint creation to avoid deserialization failures
- **Inject loaded weights** into Flax via variable replacement or pass explicitly to functional APIs

## Frequently Asked Questions

### What if my checkpoint doesn't match the current dataclass schema?

If the repository has evolved, compare your checkpoint's keys against the expected schema. You can inspect raw `.npz` contents with `np.load()` to see available arrays, then construct a compatibility layer or request the matching model version from the authors.

### Can I use `gcsfs` or another GCS library instead of TensorFlow?

Yes. Any library providing a binary file-like object works. Replace `tf.io.gfile.GFile` with `gcsfs.GCSFileSystem().open(path, "rb")` or Google's `google-cloud-storage` blob reader wrapped in `io.BytesIO`.

### Why does the loaded params object fail when passed to my model?

Most commonly: **dtype or device mismatch**. WeatherNext checkpoints typically use `float32` or `bfloat16`. Verify with `jax.tree_map(lambda x: (x.dtype, x.device()), params)` and cast or transfer as needed before model execution.

### Are there pre-trained public checkpoints available?

The WeatherNext repository indicates public releases exist; check the documentation and `docs/` directory for official GCS paths. When available, these use the same `checkpoint.load()` pattern documented here.