# How to Fine-Tune the WeatherNext Model: A Complete Guide for GraphCast and GenCast

> Learn how to fine-tune WeatherNext models like GraphCast and GenCast. Adapt the data pipeline and resume training on your custom dataset for improved weather forecasting.

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

---

**Fine-tune any WeatherNext model by loading a pretrained checkpoint, adapting the data pipeline via the `Task` API, and resuming training with `is_training=True` on your target dataset.**

The WeatherNext repository from Google DeepMind provides production-ready neural weather models—**GraphCast** (deterministic) and **GenCast** (generative)—that you can adapt to custom datasets without training from scratch. Both are implemented in JAX/Flax and ship with pretrained weights trained on ERA5 reanalysis data. This guide walks through the exact steps, files, and code patterns used in the source.

---

## Overview of the Fine-Tuning Pipeline

WeatherNext's architecture separates model definitions, checkpoint management, and data handling into clean, reusable modules. The fine-tuning workflow follows three conceptual phases:

1. **Load pretrained weights** using checkpoint utilities
2. **Adapt to your data** via the `Task` abstraction
3. **Resume training** with standard JAX optimization loops

The same pattern applies to both model families—only import paths and minor configuration details differ.

---

## Core Components for Fine-Tuning

### Model Definitions

The `GraphCast` and `GenCast` classes accept an `is_training` flag that controls dropout, batch normalization, and other training-specific behaviors.

| Model | File Path | Key Flag |
|-------|-----------|----------|
| GraphCast | [[`weathernext/weathernext1_graph/graphcast.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/weathernext1_graph/graphcast.py)](https://github.com/google-deepmind/weathernext/blob/main/weathernext/weathernext1_graph/graphcast.py) | `is_training=True` |
| GenCast | [[`weathernext/weathernext1_gen/gencast.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/weathernext1_gen/gencast.py)](https://github.com/google-deepmind/weathernext/blob/main/weathernext/weathernext1_gen/gencast.py) | `is_training=True` |

In [`graphcast.py`](https://github.com/google-deepmind/weathernext/blob/main/graphcast.py), the model constructor signature includes `is_training: bool = False`, which you must override when fine-tuning. The same pattern appears in [`gencast.py`](https://github.com/google-deepmind/weathernext/blob/main/gencast.py) for the generative variant.

---

### Checkpoint Utilities in [[`weathernext/utils/checkpoint.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/checkpoint.py)](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/checkpoint.py)

The `load_checkpoint` function handles pretrained weight loading, including:

- Mapping parameter trees across different mesh resolutions
- Handling sharded checkpoints for distributed training
- Validating compatibility with current model architecture

```python
from weathernext.utils.checkpoint import load_checkpoint, save_checkpoint

# Load pretrained ERA5 weights

params = load_checkpoint("gs://weather-next-checkpoints/graphcast_operational.ckpt")

# Save fine-tuned checkpoint

save_checkpoint(params, "gs://my-bucket/fine_tuned_model.ckpt")

```

---

### Model Adaptation Utilities in [[`weathernext/utils/model_utils.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/model_utils.py)](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/model_utils.py)

When your target dataset uses a different mesh resolution than the pretrained model, use `align_meshes` to interpolate or project weights appropriately:

```python
from weathernext.utils.model_utils import align_meshes

params = align_meshes(
    params,
    source_mesh=pretrained_mesh,
    target_mesh=new_task_mesh
)

```

Additional helpers in this file handle:
- Variable name remapping when input/target fields differ
- Weight reshaping for modified edge lengths or node counts
- Forcing field injection/removal

---

### Data Pipeline via [[`weathernext/utils/task.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/task.py)](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/task.py)

The `Task` class encapsulates everything about your dataset:

```python
from weathernext.utils.task import Task

my_task = Task(
    input_variables=["u10m", "v10m", "t2m", "msl"],
    target_variables=["t2m", "u10m", "v10m"],
    dataset_path="gs://my-bucket/regional-hres/",
    batch_size=16,
    num_steps=12,  # Lead time steps for training targets

    preprocess_fn=custom_normalization  # Optional

)

```

Key `Task` parameters:
- `input_variables`: Fields fed into the model at initialization
- `target_variables`: Fields the model predicts (may overlap with inputs for autoregressive training)
- `num_steps`: Number of autoregressive rollout steps during training
- `mesh`: The spherical mesh object defining resolution

---

## Step-by-Step Fine-Tuning Workflow

### Step 1: Environment Setup

Install dependencies as specified in the repository's top-level documentation. WeatherNext requires JAX with TPU or GPU support, plus xarray for data handling.

```bash
pip install -e .
pip install jax[tpu] -f https://storage.googleapis.com/jax-releases/libtpu_releases.html

```

---

### Step 2: Download Pretrained Checkpoint

WeatherNext releases include operational checkpoints trained on ERA5. Download the appropriate base model for your use case:

| Checkpoint | Description | Typical Fine-Tuning Target |
|------------|-------------|---------------------------|
| `graphcast_operational` | 0.25° deterministic GraphCast | Regional high-res analysis, custom variables |
| `gencast_operational` | 0.25° generative GenCast | Probabilistic forecasting, ensemble generation |

---

### Step 3: Instantiate Model with Training Mode

```python
from weathernext.weathernext1_graph.graphcast import GraphCast

# For GenCast: from weathernext.weathernext1_gen.gencast import GenCast

model = GraphCast(
    is_training=True,        # Enable dropout, training stats

    mesh_size=6,             # Match or modify vs. pretrained

    latent_dims=512,         # Architecture width

    num_processors=16        # Message passing steps

)

```

The `is_training=True` flag is **critical**—without it, batch normalization uses running statistics and dropout is disabled, preventing proper gradient flow during fine-tuning.

---

### Step 4: Load and Adapt Pretrained Weights

```python
from weathernext.utils.checkpoint import load_checkpoint
from weathernext.utils.model_utils import align_meshes

# Load raw checkpoint

params = load_checkpoint(checkpoint_path)

# Optional: adapt to new mesh resolution

if new_mesh != pretrained_mesh:
    params = align_meshes(params, new_mesh=new_mesh)

# Optional: adapt to new variable set

if modified_variables:
    from weathernext.utils.model_utils import remap_variables
    params = remap_variables(params, variable_mapping)

```

---

### Step 5: Training Loop

WeatherNext models expose a `loss` method compatible with JAX automatic differentiation. Use standard optax optimizers:

```python
import jax
import optax

optimizer = optax.adamw(
    learning_rate=optax.exponential_decay(
        init_value=1e-4,
        transition_steps=1000,
        decay_rate=0.9
    ),
    weight_decay=1e-5
)

opt_state = optimizer.init(params)
grad_fn = jax.value_and_grad(model.loss)

for step in range(num_training_steps):
    batch = my_task.next_batch()  # (inputs, targets, forcings)

    
    loss, grads = grad_fn(params, batch, model)
    
    updates, opt_state = optimizer.update(grads, opt_state, params)
    params = optax.apply_updates(params, updates)
    
    if step % 100 == 0:
        print(f"Step {step}: loss = {loss:.6f}")

```

For distributed training across multiple TPU cores, use the sharding utilities in [[`weathernext/utils/sharding.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/sharding.py)](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/sharding.py) with `pjit`:

```python
from weathernext.utils.sharding import get_mesh_sharding_spec

# Define global mesh for data/model parallelism

sharding_spec = get_mesh_sharding_spec(params, mesh_axes=('data', 'model'))
params = jax.device_put(params, sharding_spec)

```

---

### Step 6: Evaluation and Checkpoint Export

```python
from weathernext.utils.checkpoint import save_checkpoint

# Validation on held-out period

model_eval = GraphCast(is_training=False)  # Disable training-specific layers

val_loss = model_eval.loss(params, val_batch)
print(f"Validation loss: {val_loss:.6f}")

# Export fine-tuned weights

save_checkpoint(
    params,
    "gs://my-bucket/weathernext_finetuned_v1.ckpt",
    metadata={"trained_on": "custom_hres", "steps": num_training_steps}
)

```

---

## Reference Notebooks

| Notebook | Purpose |
|----------|---------|
| [`docs/weathernext1_graph/graphcast_demo.ipynb`](https://github.com/google-deepmind/weathernext/blob/main/docs/weathernext1_graph/graphcast_demo.ipynb) | End-to-end GraphCast loading, inference, and fine-tuning |
| [`docs/weathernext1_gen/gencast_demo.ipynb`](https://github.com/google-deepmind/weathernext/blob/main/docs/weathernext1_gen/gencast_demo.ipynb) | GenCast generative model fine-tuning with diffusion sampling |

These notebooks demonstrate production patterns including gradient accumulation, mixed precision training, and checkpoint resumption.

---

## Key Files for Fine-Tuning Reference

| File | Function |
|------|----------|
| [[`weathernext/weathernext1_graph/graphcast.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/weathernext1_graph/graphcast.py)](https://github.com/google-deepmind/weathernext/blob/main/weathernext/weathernext1_graph/graphcast.py) | Core deterministic model with `is_training` flag |
| [[`weathernext/weathernext1_gen/gencast.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/weathernext1_gen/gencast.py)](https://github.com/google-deepmind/weathernext/blob/main/weathernext/weathernext1_gen/gencast.py) | Generative diffusion model architecture |
| [[`weathernext/utils/checkpoint.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/checkpoint.py)](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/checkpoint.py) | `load_checkpoint()`, `save_checkpoint()` |
| [[`weathernext/utils/model_utils.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/model_utils.py)](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/model_utils.py) | `align_meshes()`, `remap_variables()` |
| [[`weathernext/utils/task.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/task.py)](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/task.py) | `Task` data pipeline abstraction |
| [[`weathernext/utils/sharding.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/sharding.py)](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/sharding.py) | Distributed training with `pjit` |

---

## Summary

- **Load** pretrained weights with `load_checkpoint()` from [[`checkpoint.py`](https://github.com/google-deepmind/weathernext/blob/main/checkpoint.py)](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/checkpoint.py)
- **Adapt** to your dataset using the `Task` API in [[`task.py`](https://github.com/google-deepmind/weathernext/blob/main/task.py)](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/task.py)
- **Transform** parameters for new meshes via `align_meshes()` in [[`model_utils.py`](https://github.com/google-deepmind/weathernext/blob/main/model_utils.py)](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/model_utils.py)
- **Train** with `is_training=True` on `GraphCast` or `GenCast` models, using standard JAX/Optax optimizers
- **Distribute** across TPUs using utilities in [[`sharding.py`](https://github.com/google-deepmind/weathernext/blob/main/sharding.py)](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/sharding.py)

Both model families share identical fine-tuning APIs—only the import path and architecture-specific hyperparameters differ.

---

## Frequently Asked Questions

### What is the minimum dataset size for fine-tuning WeatherNext?

Fine-tuning can begin with as few as 1,000–2,000 samples (roughly 2–3 months of 6-hourly data) for regional adaptation, though 6–12 months yields more stable convergence. The pretrained ERA5 weights provide strong initialization, reducing data requirements versus training from scratch. For variable addition or mesh changes, more data may be needed to learn the new parameter mappings.

### Can I fine-tune on a different spatial resolution than the pretrained model?

Yes. Use `align_meshes()` from [[`model_utils.py`](https://github.com/google-deepmind/weathernext/blob/main/model_utils.py)](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/model_utils.py) to project pretrained spherical harmonic weights onto your target mesh. This performs interpolation in the spectral domain, preserving physical constraints better than naive spatial interpolation. Finer resolutions (higher mesh_size) require proportionally more compute but can leverage the same base weights.

### How do I add new input variables not present in the pretrained checkpoint?

Create a `Task` with your extended `input_variables` list, then use `remap_variables()` from [[`model_utils.py`](https://github.com/google-deepmind/weathernext/blob/main/model_utils.py)](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/model_utils.py) to initialize new variable embeddings. The function copies compatible weights from related variables (e.g., using 850 hPa temperature initialization for new 925 hPa temperature) and randomly initializes truly novel fields. Fine-tuning then learns appropriate representations.

### Is multi-GPU or TPU pod training supported for fine-tuning?

Yes. WeatherNext uses JAX's `pjit` for automatic parallelism across devices. Define a global mesh in your training script and apply `get_mesh_sharding_spec()` from [[`sharding.py`](https://github.com/google-deepmind/weathernext/blob/main/sharding.py)](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/sharding.py) to both parameters and data batches. The same code runs on single devices, 8-GPU servers, or full TPU pods without modification—only the `mesh_axes` configuration changes.