Understanding the Autoregressive Rollout Process and PredictorFn Protocol in WeatherNext
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 and the PredictorFn protocol in 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 and 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). 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:
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:
- Reconstruct forcing — Unflattens the time-step-specific forcing data
- Merge constants — Combines static inputs with current dynamic state
- Call inner predictor — Invokes
self._predictor(...)for the forward pass - Extract predictions — Retrieves generated fields for the next time step
- Update rolling window — Advances
inputsvia_update_inputsfor 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:
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. The PredictorFn protocol enforces stateless predictors that enable JAX transformations for scaling.
Protocol Definition
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
rngparameter — Enables deterministic, reproducible stochasticity for ensemble forecasts - Pure function signature — No side effects or internal state, allowing
pmapandvmaptransforms - xarray-native I/O — Preserves labeled coordinates throughout the pipeline
Rollout Utilities Built on PredictorFn
The 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
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
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
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— Abstract base defining thePredictorinterface thatautoregressive.Predictorextendsweathernext/weathernext1_gen/transformer.py— Concrete one-step predictor implementable assingle_stepabove- Training loops — The loss aggregation in
autoregressive.pyintegrates 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
Predictorwrapper to iterate one-step models, implemented inweathernext/utils/autoregressive.py - The seven-stage pipeline validates inputs, flattens forcings, scan-loops with
hk.scan, and reconstructsxarray.Datasetoutputs PredictorFnprotocol inweathernext/utils/rollout.pyenforces stateless, RNG-explicit predictors enablingpmap/vmapscaling- 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 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 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 utilities.
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 →