How to Fine-Tune WeatherNext Models: A Complete Guide for Custom Weather Forecasting
Fine-tuning WeatherNext involves loading a pre-trained FGN checkpoint, preparing HRES-derived training data, and running JIT-compiled gradient updates using JAX and Optax, all orchestrated through the wn2_demo.ipynb notebook.
WeatherNext is an open-source weather forecasting system developed by Google DeepMind, built on a flexible Fully-Connected Graph Neural (FGN) architecture. Fine-tuning lets you adapt pre-trained models to specialized domains—such as tropical cyclone tracking or regional forecasting—using operational data like ECMWF HRES forecasts. This guide walks through the complete fine-tuning pipeline using the actual source code from the google-deepmind/weathernext repository.
Prerequisites and Architecture Overview
WeatherNext's design cleanly separates model definition from training logic, making fine-tuning straightforward. The core components reside in:
weathernext/weathernext2/fgn.py— FGN model constructors and checkpoint handlingweathernext/utils/data_utils.py— Data preprocessing utilitiesweathernext/utils/predictor_base.py— BasePredictorclass with thelossmethoddocs/weathernext2/wn2_demo.ipynb— End-to-end reference implementation
The codebase uses JAX + Haiku, enabling seamless execution on TPU, GPU, or CPU with only minor mesh-block adjustments.
Step 1: Load a Pre-Trained WeatherNext Checkpoint
Start by downloading a checkpoint and its matching Fiddle configuration. The repository provides several pre-trained variants (e.g., WeatherNextCyclones_Mini for tropical cyclones, WeatherNext2_<2025> for general forecasting).
In weathernext/weathernext2/fgn.py, the CheckPoint dataclass and construct_predictor function handle checkpoint loading【124†L124-L147】:
import dataclasses
from google.cloud import storage
from weathernext.weathernext2 import fgn
from weathernext.utils import checkpoint, fiddle_config_io
# Select model variant and split year
model_name = "WeatherNextCyclones_Mini"
split = "2024" # checkpoint suffix
config_name = f"weathernext2/configs/{model_name}"
weights_path = f"weathernext2/params/{model_name}_<{split}.npz"
# Download from public GCS bucket
gcs = storage.Client.create_anonymous_client()
bucket = gcs.get_bucket("dm_graphcast")
with bucket.blob(weights_path).open("rb") as f:
ckpt = checkpoint.load(f, fgn.CheckPoint)
config = fiddle_config_io.get_fiddle_config_by_name(config_name)
The config contains hyperparameters—mesh size, pressure levels, target lead times—that must match your checkpoint's training resolution.
Step 2: Prepare Fine-Tuning Data (HRES Format)
WeatherNext was originally trained on ERA5 reanalysis, but fine-tuning uses ECMWF HRES forecasts【57†L57-L60】. The README specifies required datasets in the "Training Data" section【154†L154-L162】.
Use extract_inputs_targets_forcings from weathernext/utils/data_utils.py to split raw NetCDF data into model-ready tensors【62†L62-L70】:
import xarray
from weathernext.utils import data_utils
# Load HRES forecast data
data_path = ("weathernext2/dataset/source-hres_forecast_init-2024-10-07"
"_00:00:00_res-1.0_levels-13_steps-20.nc")
with bucket.blob(data_path).open("rb") as f:
batch = xarray.load_dataset(f).compute()
# Extract components for single-step fine-tuning
task_cfg = config.task
train_inputs, train_targets, train_forcings = data_utils.extract_inputs_targets_forcings(
batch,
target_lead_times=slice("6h", "6h"), # 6-hour lead time
**dataclasses.asdict(task_cfg)
)
The target_lead_times parameter controls the prediction horizon. For multi-step fine-tuning, extend this slice (e.g., slice("6h", "48h")).
Step 3: Build JIT-Compiled Loss and Gradient Functions
The fine-tuning loss is computed via predictor.loss() in weathernext/utils/predictor_base.py【92†L92-L103】. The demo notebook (cells 30-33) shows how to JIT-compile this for efficient training【30†L30-L46】:
import haiku as hk
import jax
import optax
import xarray_jax
import xarray_tree
@hk.transform
def loss_fn(inputs, targets, forcings):
predictor = fgn.construct_predictor(config)
loss, diagnostics = predictor.loss(inputs, targets, forcings)
# Unwrap JAX arrays for logging
loss = xarray_tree.map_structure(
lambda x: xarray_jax.unwrap_data(x.mean(), require_jax=True),
loss
)
return loss, diagnostics
# JIT-compiled loss and gradient
grads_fn = jax.value_and_grad(
lambda params, inp, tgt, frc: loss_fn.apply(params, None, inp, tgt, frc),
has_aux=True
)
The has_aux=True flag ensures diagnostics (useful for monitoring) are returned alongside the loss.
Step 4: Execute the Fine-Tuning Step
Run forward-backward passes and update parameters with any Optax optimizer. The notebook demonstrates gradient norm tracking【70†L70-L78】:
# Initialize RNG and optimizer
rng = jax.random.PRNGKey(0)
opt = optax.adam(learning_rate=1e-4)
opt_state = opt.init(ckpt.params)
# Single optimization step
(loss, diagnostics), grads = grads_fn(
ckpt.params, train_inputs, train_targets, train_forcings
)
# Apply updates
updates, opt_state = opt.update(grads, opt_state, ckpt.params)
new_params = optax.apply_updates(ckpt.params, updates)
print(f"Fine-tuning loss: {float(loss):.4f}")
For full training loops, wrap this in a jax.lax.scan or standard Python for loop with checkpointing every N steps.
Hardware Configuration Tips
WeatherNext's JAX foundation provides device-agnostic execution. Adjust mesh-block sizes in cell 4 of wn2_demo.ipynb to match your hardware:
- TPU v3/v4: Use default mesh sizes (typically 8x8 or higher)
- NVIDIA A100: Reduce mesh blocks to fit memory (e.g., 4x4)
- CPU: Single-device mode with minimal mesh
Production Fine-Tuning Considerations
| Aspect | Recommendation |
|---|---|
| Learning rate | Start with 1e-4, reduce on plateau |
| Batch size | 1-4 samples per device (limited by HRES memory) |
| Lead times | Begin with 6h, progressively add 12h, 24h, 48h |
| Checkpointing | Save every 100 steps; validate on held-out storms |
| Mixed precision | Use jax.numpy default (full float32 for stability) |
Summary
- Load a pre-trained FGN checkpoint with matching Fiddle config from
weathernext2/fgn.py - Prepare HRES forecast data using
extract_inputs_targets_forcingsindata_utils.py - Compute losses via
predictor.loss()frompredictor_base.py - JIT-compile gradients with
jax.value_and_gradfor efficient training - Update parameters with any Optax optimizer, following the
wn2_demo.ipynbpattern
The modular architecture lets you fine-tune on arbitrary pressure levels, geographic regions, or forecast horizons without modifying core model code.
Frequently Asked Questions
What data format does WeatherNext fine-tuning require?
WeatherNext fine-tuning expects ECMWF HRES operational forecasts in NetCDF format with variables matching the checkpoint's pressure levels and surface fields. The extract_inputs_targets_forcings helper in weathernext/utils/data_utils.py converts these to the internal tensor format. For custom datasets, ensure your NetCDF includes coordinates for lat, lon, level, and time compatible with the FGN mesh.
Can I fine-tune WeatherNext on a single GPU?
Yes. WeatherNext uses JAX's device parallelism, so single-GPU fine-tuning works by setting appropriate mesh-block sizes. Reduce the spatial mesh in the Fiddle config (e.g., from 8x8 to 4x4) and use gradient accumulation if batch size is limited. The wn2_demo.ipynb notebook runs on Colab's T4 GPU without modification.
How do I resume fine-tuning from an interrupted run?
Serialize the new_params and opt_state using checkpoint.save() from weathernext/utils/checkpoint.py, then reload on restart:
# Save
with open("finetuned_step_500.pkl", "wb") as f:
checkpoint.save({"params": new_params, "opt_state": opt_state}, f)
# Resume
with open("finetuned_step_500.pkl", "rb") as f:
restored = checkpoint.load(f)
params, opt_state = restored["params"], restored["opt_state"]
What's the difference between WeatherNext and GraphCast fine-tuning?
WeatherNext uses the FGN architecture with decoupled mesh refinement, whereas GraphCast uses a learned mesh. The fine-tuning pipeline is structurally similar—both use JAX, Haiku, and HRES data—but WeatherNext's construct_predictor in fgn.py replaces GraphCast's direct GraphCast class instantiation. The demo notebook handles this abstraction automatically.
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 →