How to Run Ensemble Predictions with Multiple Model Checkpoints in WeatherNext
To run ensemble predictions in WeatherNext, load multiple checkpoints into a predictor that stacks models along an ensemble dimension, broadcasts input data across all members, and aggregates the independent forecasts.
The WeatherNext framework from Google DeepMind is designed to support ensemble inference natively. Unlike running separate inference jobs for each checkpoint, WeatherNext treats the ensemble as a first-class axis that is merged with the batch dimension for efficient parallel execution. This guide explains the implementation details based on the google-deepmind/weathernext source code.
How WeatherNext Handles Ensembles Internally
The key insight is that ensemble members do not interact during inference. Each checkpoint produces its own forecast from identical input data, and the framework handles this by introducing an explicit ensemble axis into the tensor shapes.
In weathernext/utils/sharding.py, the sharding utilities define how batch and ensemble dimensions are combined. The model architecture in weathernext2/architecture.py automatically supports an extra leading dimension for ensemble members, enabling single-compilation inference across all checkpoints.
The weathernext/utils/rollout.py module provides higher-level utilities like ensemble_forecast that encapsulate this pattern for multi-step predictions.
Step-by-Step: Loading and Running Ensemble Checkpoints
1. Prepare Checkpoint Paths
Each checkpoint must be a complete serialized model containing weights, optimizer state, and configuration. The weathernext/utils/checkpoint.py module provides the loading interface.
import pathlib
from weathernext.utils import checkpoint
ckpt_paths = [
pathlib.Path("/path/to/checkpoint_001"),
pathlib.Path("/path/to/checkpoint_002"),
pathlib.Path("/path/to/checkpoint_003"),
]
checkpoints = [checkpoint.load(p, checkpoint.Checkpoint) for p in ckpt_paths]
2. Initialize the Predictor with Multiple Checkpoints
The predictor wrapper (conceptualized in weathernext/utils/predictor_base.py) accepts a list of checkpoints and stacks them internally. This creates a predictor where the model parameters have an additional ensemble dimension.
from weathernext.utils import predictor_base
predictor = predictor_base.Predictor.from_checkpoints(checkpoints)
Under the hood, from_checkpoints uses JAX's vmap or explicit parameter stacking to broadcast the forward pass across all ensemble members in a single compiled computation.
3. Run Inference with Broadcasted Inputs
The predictor automatically broadcasts your input batch across the ensemble dimension. The output shape becomes [ensemble, batch, ...] rather than the standard [batch, ...].
# inputs: xarray DataArray with weather fields
inputs = load_input_data(...) # your data loading logic
# raw_ensemble_output shape: [ensemble, batch, lead_time, lat, lon, vars]
raw_ensemble_output = predictor.predict(inputs)
4. Aggregate Ensemble Predictions
Compute statistics across the ensemble axis to obtain the final forecast. Common aggregations include mean (deterministic best estimate) or variance (uncertainty quantification).
# Deterministic ensemble mean
final_forecast = raw_ensemble_output.mean(axis=0)
# Ensemble spread (standard deviation)
ensemble_spread = raw_ensemble_output.std(axis=0)
Using Rollout Utilities for Multi-Step Ensembles
For longer-range predictions, weathernext/utils/rollout.py provides ensemble_forecast and MultipleRunsLooped helpers that handle chunked rollouts with proper ensemble handling.
from weathernext.utils import rollout
forecast = rollout.ensemble_forecast(
predictor=predictor,
init_inputs=inputs,
horizon=48, # 48 forecast steps
ensemble_axis=0, # ensemble dimension index
)
The MultipleRunsLooped class in the same file implements memory-efficient ensemble rollouts for large ensembles or long horizons.
Complete Working Example
import pathlib
from weathernext.utils import checkpoint, predictor_base, rollout
def run_ensemble_prediction(checkpoint_dirs, input_batch, forecast_steps):
"""
Run ensemble prediction across multiple WeatherNext checkpoints.
Args:
checkpoint_dirs: List of pathlib.Path objects pointing to checkpoints
input_batch: xarray DataArray with shape [batch, ...]
forecast_steps: Number of rollout steps
Returns:
xarray DataArray with ensemble predictions of shape
[ensemble, batch, forecast_steps, ...]
"""
# Load all checkpoints
checkpoints = [
checkpoint.load(p, checkpoint.Checkpoint)
for p in checkpoint_dirs
]
# Build ensemble predictor
predictor = predictor_base.Predictor.from_checkpoints(checkpoints)
# Run multi-step ensemble forecast
ensemble_output = rollout.ensemble_forecast(
predictor=predictor,
init_inputs=input_batch,
horizon=forecast_steps,
ensemble_axis=0,
)
return ensemble_output
# Example usage
ckpt_paths = [
pathlib.Path(f"/checkpoints/weathernext_v2_seed{i:03d}")
for i in range(4) # 4-member ensemble
]
results = run_ensemble_prediction(ckpt_paths, inputs, forecast_steps=48)
Key Implementation Files
These files define the ensemble functionality in the google-deepmind/weathernext repository:
weathernext/utils/sharding.py— Dimension combining for parallel ensemble executionweathernext/utils/rollout.py—ensemble_forecast()andMultipleRunsLoopedfor multi-step predictionsweathernext/weathernext2/architecture.py— Core model that accepts ensemble-dimension parametersweathernext/utils/checkpoint.py— Checkpoint serialization/deserialization API
Summary
- Ensemble predictions require loading multiple complete checkpoints via
checkpoint.load() - The predictor stacks models along an ensemble axis using
Predictor.from_checkpoints() - Input data broadcasts automatically; output shape is
[ensemble, batch, ...] - Aggregate results with standard numpy/xarray operations (mean, std, percentiles)
- Use
rollout.ensemble_forecast()for efficient multi-step ensemble rollouts
Frequently Asked Questions
How many checkpoints can I use in an ensemble?
There is no hard limit in the WeatherNext framework. Practical limits depend on available accelerator memory, since all ensemble members reside in device memory simultaneously. The sharding utilities in weathernext/utils/sharding.py allow ensemble and batch dimensions to be partitioned across multiple devices for scaling.
Do ensemble checkpoints need identical architectures?
Yes. All checkpoints must originate from the same model configuration. The Predictor.from_checkpoints() method assumes compatible parameter structures and will raise errors if architectures mismatch. This is verified during the stacking operation in predictor_base.py.
Can I mix deterministic and stochastic checkpoints in one ensemble?
Yes. The framework treats all checkpoints uniformly—each produces independent predictions. Stochastic checkpoints that include random dropout or sampling will generate varied outputs naturally. The rollout.py utilities handle this mixing without special configuration.
How do I compute probabilistic forecasts from ensemble outputs?
After obtaining raw ensemble predictions with shape [ensemble, batch, ...], apply xarray or numpy statistical operations: mean() for deterministic estimates, std() for spread, quantile() for prediction intervals, or kernel density estimation for full probability distributions.
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 →