# How to Fine-Tune WeatherNext Models: A Complete Guide for Custom Weather Forecasting

> Learn how to fine-tune WeatherNext models with this guide. Load checkpoints, prepare data, and perform gradient updates using JAX and Optax for custom weather forecasting.

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

---

**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`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/weathernext2/fgn.py) — FGN model constructors and checkpoint handling
- [`weathernext/utils/data_utils.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/data_utils.py) — Data preprocessing utilities
- [`weathernext/utils/predictor_base.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/predictor_base.py) — Base `Predictor` class with the `loss` method
- `docs/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`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/weathernext2/fgn.py), the `CheckPoint` dataclass and `construct_predictor` function handle checkpoint loading【124†L124-L147】:

```python
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`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/data_utils.py) to split raw NetCDF data into model-ready tensors【62†L62-L70】:

```python
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`](https://github.com/google-deepmind/weathernext/blob/main/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】:

```python
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】:

```python

# 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`](https://github.com/google-deepmind/weathernext/blob/main/weathernext2/fgn.py)
- **Prepare** HRES forecast data using `extract_inputs_targets_forcings` in [`data_utils.py`](https://github.com/google-deepmind/weathernext/blob/main/data_utils.py)
- **Compute** losses via `predictor.loss()` from [`predictor_base.py`](https://github.com/google-deepmind/weathernext/blob/main/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`](https://github.com/google-deepmind/weathernext/blob/main/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`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/checkpoint.py), then reload on restart:

```python

# 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`](https://github.com/google-deepmind/weathernext/blob/main/fgn.py) replaces GraphCast's direct `GraphCast` class instantiation. The demo notebook handles this abstraction automatically.