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

Train a WeatherNext model by preparing xarray datasets, instantiating a Task object, building the architecture from 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 repository.

WeatherNext Model Families: Choose Your Architecture

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

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

WeatherNext 1: Legacy Transformers and Diffusion Models

For compatibility with published research or specific use cases:

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 module handles conversion from raw ERA5 or satellite files into the canonical "inputs/targets" convention.

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 bundles all training metadata: input tensors, target tensors, variable lists, forecast horizon, and optional conditioning fields.

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

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

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.

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 module provides standardized loss computation across all WeatherNext predictors. A minimal JAX training step with JIT compilation:

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:

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 module handles serialization of parameters and optimizer state.

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:

@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:

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:

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
  • Task definition: Bundle inputs/targets/metadata with the Task class in utils/task.py
  • Model construction: Build F-Graph-Net via 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
  • 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, 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 to load published weights, then continue training with is_training=True on your regional dataset. The 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 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.

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:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →