# Understanding the Autoregressive Rollout Process and PredictorFn Protocol in WeatherNext

> Explore the autoregressive rollout process in WeatherNext for multi-step forecasts. Learn how PredictorFn protocol and the Predictor wrapper enable iterative predictions using past outputs as new inputs.

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

---

**The autoregressive rollout process in WeatherNext generates multi-step weather forecasts by repeatedly feeding a one-step model's predictions back as inputs, controlled by the `Predictor` wrapper in [`weathernext/utils/autoregressive.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/autoregressive.py) and the `PredictorFn` protocol in [`weathernext/utils/rollout.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/rollout.py).**

WeatherNext, Google DeepMind's open-source weather forecasting system, achieves long-range predictions through an elegant autoregressive mechanism. Rather than training separate models for different lead times, it wraps a single-step predictor and iterates it forward in time. This design, implemented across [`weathernext/utils/autoregressive.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/autoregressive.py) and [`weathernext/utils/rollout.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/rollout.py), enables memory-efficient, parallelizable forecasts that scale to ensemble runs.

## How the Autoregressive Rollout Process Works in WeatherNext

The autoregressive rollout transforms any one-step predictor into a multi-step forecaster through a seven-stage pipeline. This process lives primarily in the `Predictor` class, which inherits from `predictor_base.Predictor` and orchestrates the time evolution.

### Stage 1: Input Validation and Separation

The wrapper first validates and separates input types. Constant (time-independent) variables are extracted from time-dependent ones through `_get_and_validate_constant_inputs` (lines 88-98 in [`autoregressive.py`](https://github.com/google-deepmind/weathernext/blob/main/autoregressive.py)). Meanwhile, `_validate_targets_and_forcings` (lines 99-113) ensures consistency between inputs, target templates, and forcing fields.

This separation matters because constants like topography are applied at every step without modification, while dynamic variables evolve through the rollout.

### Stage 2: Template Preparation and Forcing Flattening

The target template is truncated to its first time slice:

```python
target_template = targets_template.isel(time=[0])

```

Forcing fields are flattened once via `_get_flat_arrays_and_single_timestep_treedef`. This preprocessing optimizes the subsequent `jax.lax.scan` operations by avoiding repeated tree manipulation inside the loop.

### Stage 3: The One-Step Prediction Function

The core iteration logic resides in `one_step_prediction` (lines 75-94). Each scan loop performs:

1. **Reconstruct forcing** — Unflattens the time-step-specific forcing data
2. **Merge constants** — Combines static inputs with current dynamic state
3. **Call inner predictor** — Invokes `self._predictor(...)` for the forward pass
4. **Extract predictions** — Retrieves generated fields for the next time step
5. **Update rolling window** — Advances `inputs` via `_update_inputs` for the next iteration

### Stage 4: JAX Scan for Efficient Iteration

Rather than Python loops, WeatherNext uses `hk.scan` (Haiku's scan) to unroll the autoregressive process:

```python
hk.scan(one_step_prediction, init_carry, none_sequence, length=num_steps)

```

Optional **gradient checkpointing** via `hk.remat` reduces memory consumption at the cost of recomputation—critical for long rollouts on accelerators with limited HBM.

### Stage 5: Result Reconstruction

After the scan completes, flat prediction arrays are unflattened and reassembled into an `xarray.Dataset` matching the original target template structure (lines 14-22). This preserves coordinate metadata and enables downstream analysis workflows.

### Stage 6: Loss Aggregation for Training

During training, the wrapper computes per-step losses from the inner predictor and averages them across the rollout horizon (lines 24-42). This ensures gradients flow through all autoregressive steps rather than only the final prediction.

## The PredictorFn Protocol: Stateless Forecasting Interface

The autoregressive rollout process relies on a strict functional protocol defined in [`weathernext/utils/rollout.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/rollout.py). The `PredictorFn` protocol enforces stateless predictors that enable JAX transformations for scaling.

### Protocol Definition

```python
class PredictorFn(Protocol):
    """Functional version of base.Predictor.__call__ with explicit rng."""
    def __call__(self,
                 rng: chex.PRNGKey,
                 inputs: xarray.Dataset,
                 targets_template: xarray.Dataset,
                 forcings: Optional[xarray.Dataset],
                 **optional_kwargs) -> xarray.Dataset:
        ...

```

Three design choices distinguish `PredictorFn` from object-oriented alternatives:

- **Explicit `rng` parameter** — Enables deterministic, reproducible stochasticity for ensemble forecasts
- **Pure function signature** — No side effects or internal state, allowing `pmap` and `vmap` transforms
- **xarray-native I/O** — Preserves labeled coordinates throughout the pipeline

### Rollout Utilities Built on PredictorFn

The [`rollout.py`](https://github.com/google-deepmind/weathernext/blob/main/rollout.py) module provides higher-level orchestration:

| Utility | Purpose |
|---------|---------|
| `chunked_prediction` | Generates long trajectories in memory-bounded segments |
| `chunked_prediction_generator` | Yields chunks lazily for streaming pipelines |
| `chunked_prediction_generator_multiple_runs` | Parallel ensemble generation with `pmap` |

These functions accept any `PredictorFn`-compliant callable, decoupling model architecture from rollout strategy.

## Practical Implementation Examples

### Wrapping a One-Step Predictor for Autoregression

```python
import weathernext.utils.autoregressive as ar
from weathernext.utils import predictor_base

# single_step: any Predictor implementation (e.g., transformer)

single_step = MyOneStepPredictor()
ar_predictor = ar.Predictor(single_step)

# Multi-step forecast

predictions = ar_predictor(inputs, targets_template, forcings)

```

The `ar_predictor` object now behaves like a multi-step model while internally managing the autoregressive loop, input window updates, and result packaging.

### Stateless Functional API with PredictorFn

```python
from weathernext.utils.rollout import PredictorFn, chunked_prediction

def my_predictor_fn(rng, inputs, targets_template, forcings, **kwargs):
    # Stateless implementation (e.g., Haiku-transformed model)

    net = hk.transform_with_state(forward_fn)
    params, state = load_checkpoint()
    predictions, _ = net.apply(params, state, rng, inputs, forcings)
    return align_to_template(predictions, targets_template)

my_fn: PredictorFn = my_predictor_fn  # Protocol verification

# 240-hour forecast in 6-hour chunks

forecast = chunked_prediction(
    predictor_fn=my_fn,
    rng=jax.random.PRNGKey(0),
    inputs=analysis_data,
    targets_template=target_240h,
    forcings=forcing_data,
    num_steps_per_chunk=6,
)

```

### Ensemble Parallelization with Multiple Runs

```python
from weathernext.utils.rollout import chunked_prediction_generator_multiple_runs

num_samples = 50
rngs = jax.random.split(jax.random.PRNGKey(0), num_samples)

for chunk in chunked_prediction_generator_multiple_runs(
        predictor_fn=my_fn,
        rngs=rngs,
        inputs=analysis_data,
        targets_template=target_structure,
        forcings=forcing_data,
        num_samples=num_samples,
        pmap_devices=jax.devices()):
    # chunk: xarray.Dataset with samples parallelized across devices

    save_chunk(chunk)

```

The `pmap_devices` argument distributes ensemble members across available accelerators, with results gathered transparently.

## Integration with WeatherNext's Architecture

The autoregressive components connect to broader system parts:

- **[`weathernext/utils/predictor_base.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/predictor_base.py)** — Abstract base defining the `Predictor` interface that `autoregressive.Predictor` extends
- **[`weathernext/weathernext1_gen/transformer.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/weathernext1_gen/transformer.py)** — Concrete one-step predictor implementable as `single_step` above
- **Training loops** — The loss aggregation in [`autoregressive.py`](https://github.com/google-deepmind/weathernext/blob/main/autoregressive.py) integrates with Optax optimizers for end-to-end differentiation through rollout steps

This composability allows researchers to swap model architectures without modifying rollout infrastructure.

## Summary

- **Autoregressive rollouts** in WeatherNext use a `Predictor` wrapper to iterate one-step models, implemented in [`weathernext/utils/autoregressive.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/autoregressive.py)
- The seven-stage pipeline validates inputs, flattens forcings, scan-loops with `hk.scan`, and reconstructs `xarray.Dataset` outputs
- **`PredictorFn`** protocol in [`weathernext/utils/rollout.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/rollout.py) enforces stateless, RNG-explicit predictors enabling `pmap`/`vmap` scaling
- **Chunked utilities** provide memory-bounded, parallelizable forecast generation for operational deployment
- Gradient checkpointing and lazy generators support long-horizon forecasts (240+ hours) and large ensembles (50+ members)

## Frequently Asked Questions

### What is the difference between `Predictor` and `PredictorFn` in WeatherNext?

`Predictor` is a class-based interface in [`predictor_base.py`](https://github.com/google-deepmind/weathernext/blob/main/predictor_base.py) that objects implement to become multi-step forecasters, with `autoregressive.Predictor` being one concrete wrapper. `PredictorFn` is a functional protocol requiring a pure function signature with explicit `rng` parameter. The functional form enables JAX transforms while the class form provides object-oriented convenience. Both support the same autoregressive rollout semantics but at different abstraction levels.

### How does WeatherNext handle memory during long autoregressive rollouts?

Memory management employs three strategies: **forcing flattening** outside the scan loop avoids tree overhead; **`hk.scan` computation graph optimization** fused across steps; and **optional `hk.remat`** which trades computation for memory by recomputing forward passes during backpropagation rather than storing all intermediate activations. For extreme lengths, `chunked_prediction` breaks the rollout into independently executed segments.

### Can I use a custom one-step model with WeatherNext's autoregressive system?

Yes. Any object implementing `predictor_base.Predictor` (with `__call__` accepting `inputs`, `targets_template`, `forcings`) can be wrapped by `autoregressive.Predictor`. For functional APIs, define a callable matching `PredictorFn` and pass directly to `chunked_prediction`. The [`transformer.py`](https://github.com/google-deepmind/weathernext/blob/main/transformer.py) implementation demonstrates the expected interface for neural network predictors.

### Why does `PredictorFn` require an explicit random key?

Explicit `rng: chex.PRNGKey` parameters enable deterministic, reproducible stochasticity essential for ensemble forecasting. Unlike hidden state or global randomness, explicit keys allow: splitting for independent ensemble members, `pmap` for device-parallel generation, andbit-exact reproduction of published results. This design aligns with JAX's functional purity principles and appears throughout the [`rollout.py`](https://github.com/google-deepmind/weathernext/blob/main/rollout.py) utilities.