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

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

pip install tensorflow jax jaxlib flax

Step 2: Import Required Modules

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

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:

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:

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

Complete Working Examples

Loading Flax WeatherNext Parameters

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

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

On Google Cloud infrastructure, authentication is automatic:


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

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


# 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 Serialize/deserialize model trees load(file_obj, cls), dump(file_obj, obj)
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) 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.

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:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →