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

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) is_training=True
GenCast [weathernext/weathernext1_gen/gencast.py](https://github.com/google-deepmind/weathernext/blob/main/weathernext/weathernext1_gen/gencast.py) is_training=True

In 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 for the generative variant.


Checkpoint Utilities in [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
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)

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

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)

The Task class encapsulates everything about your dataset:

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.

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

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

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:

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) with pjit:

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

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 End-to-end GraphCast loading, inference, and fine-tuning
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) 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) Generative diffusion model architecture
[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) align_meshes(), remap_variables()
[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) Distributed training with pjit

Summary

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/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/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/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.

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 →