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 architectureGraphForecasterParams— 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 —.npzfiles are binary archives checkpoint.load()accepts any file-like object with aread()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
Default Credentials (Recommended)
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:
numpy.load()to read the.npzarchivetree_unflattenwith the provided dataclass to reconstruct nested structure- Validation that all expected fields are present in the archive
Performance Considerations
- Streaming vs. download:
tf.io.gfile.GFilestreams bytes on demand; for multi-GB checkpoints, consider local SSD caching withgsutil cpfollowed 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
.npzcollections
Summary
- WeatherNext stores weights as
.npzarchives with dataclass-structured schemas - Use
tf.io.gfile.GFileto 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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →