# How to Train a WeatherNext Model: A Complete Guide to Data-Driven Weather Forecasting

> Learn to train a WeatherNext model with our complete guide. Prepare data, build architecture, and run a JAX training loop for advanced weather forecasting.

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

---

**Train a WeatherNext model by preparing xarray datasets, instantiating a `Task` object, building the architecture from [`weathernext2/architecture.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext2/architecture.py) or legacy GraphCast/GenCast modules, and running a JAX-based training loop with Optax.**

WeatherNext is DeepMind's research-grade framework for learning data-driven weather forecasting models on large-scale satellite and reanalysis data. This guide walks through the complete training pipeline for both the state-of-the-art **WeatherNext 2 (F-Graph-Net)** family and the earlier **WeatherNext 1 (GraphCast/GenCast)** models, referencing actual source files and runnable code from the [google-deepmind/weathernext](https://github.com/google-deepmind/weathernext) repository.

## WeatherNext Model Families: Choose Your Architecture

The repository provides two distinct model families with different compute characteristics and forecasting approaches.

### WeatherNext 2: F-Graph-Net (Recommended)

**F-Graph-Net** represents the current state-of-the-art, using fully-connected learnable mesh-based graph networks for global forecasting.

- **Core module**: [`weathernext/weathernext2/fgn.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/weathernext2/fgn.py) — implements the message-passing layers and graph convolution operations
- **High-level constructor**: [`weathernext/weathernext2/architecture.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/weathernext2/architecture.py) — wraps FGN components into a trainable predictor

### WeatherNext 1: Legacy Transformers and Diffusion Models

For compatibility with published research or specific use cases:

- **GraphCast**: [`weathernext/weathernext1_graph/graphcast.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/weathernext1_graph/graphcast.py) — graph-based transformer architecture
- **GenCast**: [`weathernext/weathernext1_gen/gencast.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/weathernext1_gen/gencast.py) — diffusion-based generative forecaster

Both families share the same training infrastructure in `utils/`, enabling consistent workflows regardless of model choice.

## Step 1: Prepare Your Dataset

WeatherNext expects data in **xarray format** (typically NetCDF). The [`utils/data_utils.py`](https://github.com/google-deepmind/weathernext/blob/main/utils/data_utils.py) module handles conversion from raw ERA5 or satellite files into the canonical "inputs/targets" convention.

```python
import xarray as xr
from weathernext.utils import data_utils

# Load raw reanalysis or satellite data

ds = xr.open_dataset('era5_sample.nc')

# Convert to WeatherNext's expected format

inputs, targets = data_utils.prepare_inputs_and_targets(ds)

```

The `prepare_inputs_and_targets()` function normalizes variables, handles temporal alignment, and ensures compatible coordinate systems for spherical harmonic operations.

## Step 2: Define a Training Task

The `Task` class in [`utils/task.py`](https://github.com/google-deepmind/weathernext/blob/main/utils/task.py) bundles all training metadata: input tensors, target tensors, variable lists, forecast horizon, and optional conditioning fields.

```python
from weathernext.utils import task

train_task = task.Task(
    inputs=inputs,
    targets=targets,
    variables=['t2m', 'u10', 'v10', 'sp', 'z500'],  # 2m temp, 10m wind, surface pressure, 500hPa geopotential

    forecast_steps=6)                                # 6-hour lead time

```

Key `Task` attributes include:
- `inputs`: Tensor of shape `[batch, time_history, levels, lat, lon, variables]`
- `targets`: Tensor of shape `[batch, forecast_steps, levels, lat, lon, variables]`
- `variables`: List of prognostic variable names for output channels
- `forecast_steps`: Number of autoregressive steps to unroll during training

## Step 3: Build the Model Architecture

### WeatherNext 2 (F-Graph-Net) Construction

```python
from weathernext.weathernext2 import architecture as wn2_arch
from weathernext.utils import mesh_transformer as mesh_utils

# Load icosahedral mesh (resolution=1 ≈ 1° or ~110km; resolution=0.25 ≈ 25km)

mesh = mesh_utils.load_mesh(resolution=1)

model = wn2_arch.build_model(
    mesh=mesh,
    variables=train_task.variables,
    latent_size=256,        # Width of latent channels in message passing

    num_message_passing=16, # Graph convolution depth

    is_training=True)       # Enables dropout and stochastic layers

```

The `build_model()` function in [`architecture.py`](https://github.com/google-deepmind/weathernext/blob/main/architecture.py) integrates the mesh geometry with the FGN encoder-processor-decoder stack. The **latent_size** parameter controls model capacity; published WeatherNext 2 models use 256 or 512.

### Alternative: WeatherNext 1 GraphCast

```python
from weathernext.weathernext1_graph import graphcast

model = graphcast.GraphCast(
    mesh=mesh,
    variables=train_task.variables,
    latent_size=512,
    num_attention_heads=8,
    is_training=True)

```

## Step 4: Configure the Optimizer

WeatherNext uses **Optax** for gradient-based optimization, fully compatible with JAX's functional transformation model.

```python
import optax

# AdamW with weight decay for regularization

optimizer = optax.adamw(
    learning_rate=5e-5,
    weight_decay=1e-6,
    b1=0.9,
    b2=0.95)

# Initialize optimizer state from model parameters

opt_state = optimizer.init(model.parameters())

```

Common learning rate schedules include:
- **Warmup + cosine decay**: `optax.warmup_cosine_decay_schedule()`
- **Exponential decay**: `optax.exponential_decay()` with decay rate 0.95 every 10,000 steps

## Step 5: Implement the Training Loop

The [`predictor_base.py`](https://github.com/google-deepmind/weathernext/blob/main/predictor_base.py) module provides standardized loss computation across all WeatherNext predictors. A minimal JAX training step with JIT compilation:

```python
import jax
import jax.numpy as jnp
from weathernext.utils import predictor_base

@jax.jit
def train_step(params, opt_state, batch):
    """Single training iteration with gradient computation."""
    
    def loss_fn(p):
        # Forward pass in training mode (dropout active)

        preds = model.apply(p, batch.inputs, is_training=True)
        
        # Compute loss (MSE or custom spectral loss)

        return predictor_base.compute_loss(preds, batch.targets)
    
    # Gradient via automatic differentiation

    grads = jax.grad(loss_fn)(params)
    
    # Optax optimizer update

    updates, new_opt_state = optimizer.update(grads, opt_state, params)
    new_params = optax.apply_updates(params, updates)
    
    return new_params, new_opt_state

# Execute training step

new_params, new_opt_state = train_step(
    model.parameters(),
    opt_state,
    train_task)

```

For **multi-step (autoregressive) training**, unroll the model over `forecast_steps`:

```python
def unrolled_loss_fn(params, batch):
    total_loss = 0.0
    current_input = batch.inputs
    
    for step in range(batch.forecast_steps):
        pred = model.apply(params, current_input, is_training=True)
        target = batch.targets[:, step]
        total_loss += predictor_base.compute_loss(pred, target)
        
        # Autoregressive feedback: prediction becomes next input

        current_input = update_input_state(current_input, pred)
    
    return total_loss / batch.forecast_steps

```

## Step 6: Checkpoint and Resume Training

The [`utils/checkpoint.py`](https://github.com/google-deepmind/weathernext/blob/main/utils/checkpoint.py) module handles serialization of parameters and optimizer state.

```python
from weathernext.utils import checkpoint

# Save checkpoint every N steps

checkpoint.save(
    path='checkpoints/wn2_step_10000',
    params=model.parameters(),
    opt_state=opt_state,
    step=10000)

# Restore for fine-tuning or resumption

params, opt_state, step = checkpoint.restore('checkpoints/wn2_step_10000')

```

Checkpoint files contain:
- Model parameters (PyTree of JAX arrays)
- Optimizer state (including momentum buffers)
- Training step counter and metadata

## Step 7: Evaluation and Inference

Switch to evaluation mode by setting `is_training=False`, which disables dropout and stochastic layers:

```python
@jax.jit
def eval_step(params, batch):
    preds = model.apply(params, batch.inputs, is_training=False)
    loss = predictor_base.compute_loss(preds, batch.targets)
    return preds, loss

# Run on validation set

val_predictions, val_loss = eval_step(model.parameters(), val_task)

```

For probabilistic forecasts with GenCast (diffusion model), the forward pass includes denoising iterations:

```python
from weathernext.weathernext1_gen import gencast

# Sample requires random key for diffusion process

rng = jax.random.PRNGKey(42)
ensemble_preds = gencast.sample(params, inputs, rng, num_samples=50)

```

## Key Training Configuration Parameters

| Parameter | Typical Values | Location |
|-----------|---------------|----------|
| Mesh resolution | 1, 0.5, 0.25 (degrees) | `mesh_utils.load_mesh()` |
| Latent size | 128, 256, 512 | `architecture.build_model()` |
| Message passing steps | 8, 16, 24 | FGN processor config |
| Learning rate | 1e-4 to 5e-6 | Optax scheduler |
| Batch size | 1-4 (per device,gradient-accumulated to 32) | Data pipeline |
| Forecast steps (training) | 1-12 autoregressive steps | `Task` constructor |

## Complete Training Pipeline Example

Combine all components into a runnable script structure:

```python
import jax
from weathernext.weathernext2 import architecture as wn2_arch
from weathernext.utils import mesh_transformer, task, data_utils, checkpoint
import optax

# 1. Load data

ds = data_utils.load_dataset('era5_2020_2023.nc')
inputs, targets = data_utils.prepare_inputs_and_targets(ds)

# 2. Create Task

train_task = task.Task(inputs=inputs, targets=targets,
                       variables=['t2m', 'u10', 'v10', 'z500', 'q700'],
                       forecast_steps=6)

# 3. Build model

mesh = mesh_transformer.load_mesh(resolution=0.5)
model = wn2_arch.build_model(mesh=mesh, variables=train_task.variables,
                             latent_size=256, is_training=True)

# 4. Optimizer

schedule = optax.warmup_cosine_decay_schedule(
    init_value=0.0, peak_value=5e-5, warmup_steps=1000,
    decay_steps=100000, end_value=1e-6)
optimizer = optax.adamw(learning_rate=schedule, weight_decay=1e-6)
opt_state = optimizer.init(model.parameters())

# 5. Training loop

for step in range(100000):
    params, opt_state = train_step(model.parameters(), opt_state, train_task)
    
    if step % 1000 == 0:
        checkpoint.save(f'checkpoints/step_{step}', params, opt_state, step)

```

## Summary

Training a WeatherNext model involves seven core steps:

- **Data preparation**: Convert ERA5/satellite data to xarray format using [`utils/data_utils.py`](https://github.com/google-deepmind/weathernext/blob/main/utils/data_utils.py)
- **Task definition**: Bundle inputs/targets/metadata with the `Task` class in [`utils/task.py`](https://github.com/google-deepmind/weathernext/blob/main/utils/task.py)
- **Model construction**: Build F-Graph-Net via [`weathernext2/architecture.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext2/architecture.py) or legacy models from `weathernext1_graph/` or `weathernext1_gen/`
- **Optimization setup**: Configure Optax optimizers with appropriate learning rate schedules
- **Training execution**: Run JIT-compiled JAX loops using `predictor_base.compute_loss()` for gradient computation
- **Checkpointing**: Persist parameters and optimizer state with [`utils/checkpoint.py`](https://github.com/google-deepmind/weathernext/blob/main/utils/checkpoint.py)
- **Evaluation**: Switch to `is_training=False` mode for deterministic inference

The repository's `docs/weathernext2/wn2_demo.ipynb` provides an end-to-end runnable demonstration of this entire pipeline.

## Frequently Asked Questions

### What hardware is required to train WeatherNext models?

Training WeatherNext 2 at standard resolutions (0.5°-1°) requires TPU v4 or v5e pods with 16-64 chips for batch-parallel training. Single-device training is possible for development but impractical for converged global models. The JAX backend enables efficient data parallelism across TPU topologies with `jax.pmap` or `jax.experimental.pjit`.

### How does WeatherNext 2 differ from the original GraphCast?

WeatherNext 2 (F-Graph-Net) replaces GraphCast's transformer-based processor with a fully-connected graph network that learns message passing on an icosahedral mesh. According to the source code in [`fgn.py`](https://github.com/google-deepmind/weathernext/blob/main/fgn.py), this eliminates the O(N²) attention complexity, enabling higher-resolution meshes (0.25° vs. GraphCast's 0.5°) with comparable parameter counts. Both use the same encoder-decoder structure for spherical-to-mesh and mesh-to-spherical projections.

### Can I fine-tune a pre-trained WeatherNext checkpoint on regional data?

Yes. Use [`utils/checkpoint.py`](https://github.com/google-deepmind/weathernext/blob/main/utils/checkpoint.py) to load published weights, then continue training with `is_training=True` on your regional dataset. The [`model_utils.py`](https://github.com/google-deepmind/weathernext/blob/main/model_utils.py) module provides helpers for variable subsetting and mesh interpolation when target domains differ from the pre-training configuration. Set a lower learning rate (1e-5) for fine-tuning to preserve large-scale learned dynamics.

### Where is the loss function defined for custom training objectives?

The base implementation resides in [`utils/predictor_base.py`](https://github.com/google-deepmind/weathernext/blob/main/utils/predictor_base.py) as the `compute_loss()` method. Override this by subclassing the predictor or passing a custom loss function to your training step. Spectral losses (weighting by wavenumber) and precipitation-heavy losses for extreme events are common modifications, implemented by operating on `jnp.fft.rfft2`-transformed predictions and targets.