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:

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_forcings in data_utils.py
  • Compute losses via predictor.loss() from predictor_base.py
  • JIT-compile gradients with jax.value_and_grad for efficient training
  • Update parameters with any Optax optimizer, following the wn2_demo.ipynb pattern

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:

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 →