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:
- Load pretrained weights using checkpoint utilities
- Adapt to your data via the
Taskabstraction - 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 initializationtarget_variables: Fields the model predicts (may overlap with inputs for autoregressive training)num_steps: Number of autoregressive rollout steps during trainingmesh: 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
Summary
- Load pretrained weights with
load_checkpoint()from [checkpoint.py](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/checkpoint.py) - Adapt to your dataset using the
TaskAPI in [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/weathernext/utils/model_utils.py) - Train with
is_training=TrueonGraphCastorGenCastmodels, using standard JAX/Optax optimizers - Distribute across TPUs using utilities in [
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/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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →