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.
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— implements the message-passing layers and graph convolution operations - High-level constructor:
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— graph-based transformer architecture - GenCast:
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 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 channelsforecast_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
Taskclass inutils/task.py - Model construction: Build F-Graph-Net via
weathernext2/architecture.pyor legacy models fromweathernext1_graph/orweathernext1_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=Falsemode 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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →