# How to Run Ensemble Predictions with Multiple Model Checkpoints in WeatherNext

> Learn to run ensemble predictions with multiple model checkpoints in WeatherNext. Load checkpoints, stack models, broadcast data, and aggregate forecasts for enhanced accuracy.

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

---

**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`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/sharding.py), the sharding utilities define how batch and ensemble dimensions are combined. The model architecture in [`weathernext2/architecture.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext2/architecture.py) automatically supports an extra leading dimension for ensemble members, enabling single-compilation inference across all checkpoints.

The [`weathernext/utils/rollout.py`](https://github.com/google-deepmind/weathernext/blob/main/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`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/checkpoint.py) module provides the loading interface.

```python
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`](https://github.com/google-deepmind/weathernext/blob/main/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.

```python
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, ...]`.

```python

# 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).

```python

# 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`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/rollout.py) provides `ensemble_forecast` and `MultipleRunsLooped` helpers that handle chunked rollouts with proper ensemble handling.

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

```python
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`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/sharding.py) — Dimension combining for parallel ensemble execution
- [`weathernext/utils/rollout.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/rollout.py) — `ensemble_forecast()` and `MultipleRunsLooped` for multi-step predictions
- [`weathernext/weathernext2/architecture.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/weathernext2/architecture.py) — Core model that accepts ensemble-dimension parameters
- [`weathernext/utils/checkpoint.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/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`](https://github.com/google-deepmind/weathernext/blob/main/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`](https://github.com/google-deepmind/weathernext/blob/main/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`](https://github.com/google-deepmind/weathernext/blob/main/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.