How to Fine-Tune WeatherNext on Custom Datasets with Different Variables
Fine-tune WeatherNext on custom atmospheric datasets by extending the variable catalogue in weathernext/utils/variables.py, defining a new Task mapping inputs/targets/forcings, and training with JAX + Optax using the model's built-in loss method.
WeatherNext—Google DeepMind's state-of-the-art weather forecasting system encompassing both the FGN-based WeatherNext 2 and the diffusion-based GenCast models—was built with a modular architecture that simplifies adaptation to custom datasets. This guide walks through fine-tuning WeatherNext when your data contains atmospheric variables not present in the original training distribution, drawing directly from the google-deepmind/weathernext source code.
Understanding WeatherNext's Fine-Tuning Architecture
The repository cleanly separates concerns across six core components. Understanding each lets you isolate exactly where changes are needed for custom variables:
| Component | Location | Purpose |
|---|---|---|
| Model definition | weathernext/weathernext2/architecture.py |
JAX-compatible neural net implementing the forecast forward pass |
| Predictor API | weathernext/utils/predictor_base.py |
Converts xarray datasets to JAX arrays; provides loss() method |
| Task definition | weathernext/utils/task.py |
Maps which variables are inputs, targets, and forcings |
| Variable catalogue | weathernext/utils/variables.py |
Registry of all variables with normalization schemes |
| Configuration handling | weathernext/utils/fiddle_config_io.py |
Fiddle config utilities for hyperparameter changes |
| Checkpointing | weathernext/utils/checkpoint.py |
Save and restore model parameters |
This separation means you typically only need to modify the variable catalogue and task definition when adding new atmospheric fields.
Preparing Your Custom Xarray Dataset
WeatherNext expects data as an xarray.Dataset with these mandatory dimensions (flexible ordering):
time— forecast lead times or analysis timesbatch— independent samples (usually 1 for fine-tuning)lat,lon— geographical grid coordinateslevel(optional) — vertical coordinates for pressure-level data
Each atmospheric field is a DataArray spanning these dimensions.
Adding a New Variable to the Catalogue
When your dataset contains variables absent from the original training (e.g., custom diagnostics, novel satellite retrievals, or derived quantities like humidity-minus-temperature), you must register them in weathernext/utils/variables.py:
# In weathernext/utils/variables.py
from dataclasses import dataclass
@dataclass
class VariableInfo:
name: str
units: str
normaliser: Normalizer # StandardScaler or MinMaxScaler
# Example: adding a custom surface radiation variable
CUSTOM_SSR = VariableInfo(
name='custom_ssr',
units='W m-2',
normaliser=StandardScaler(mean=200.0, std=150.0)
)
Critical requirements:
- The
namefield must exactly match theDataArrayname in yourxarray.Dataset - Choose normalization statistics representative of your dataset's climatology
- For variables that should not be predicted but supplied to the model (prescribed SST, solar forcing), skip adding to target lists
Dataset Structure Validation
Verify your dataset structure before training:
import xarray as xr
ds = xr.open_dataset('my_custom_data.nc')
required_dims = {'time', 'batch', 'lat', 'lon'}
assert required_dims.issubset(set(ds.dims)), f"Missing: {required_dims - set(ds.dims)}"
# Confirm custom variable exists with correct name
assert 'custom_ssr' in ds.data_vars, "Variable name mismatch with catalogue"
Defining the Training Task
The Task class in weathernext/utils/task.py is the central mapping between your raw dataset and the model's expected tensor structure. Instantiate it with your variable selections:
from weathernext.utils import task, variables
# Extend the built-in variable lists with your custom additions
my_input_vars = variables.INPUT_VARIABLES + [variables.CUSTOM_SSR]
my_target_vars = variables.TARGET_VARIABLES # or add custom_ssr if predicting it
my_forcing_vars = variables.FORCING_VARIABLES # or move variables here
# Build the task
train_task = task.Task(
inputs=ds[my_input_vars],
targets=ds[my_target_vars],
forcings=ds[my_forcing_vars]
)
Variable role guidelines:
- Inputs: Variables the model sees at initialization time (typically analysis fields)
- Targets: Variables the model learns to predict (can include custom variables)
- Forcings: Variables supplied at each lead time but not predicted (boundary conditions, prescribed forcings)
Loading the Model for Training
Instantiate your architecture with is_training=True to enable dropout, batch norm updates, and gradient computation.
WeatherNext 2 (FGN)
from weathernext.weathernext2 import architecture, fgn
# Load default configuration; modify via Fiddle if needed
model_cfg = architecture.DefaultConfig()
# Initialize trainable model
model = fgn.FGN(
config=model_cfg,
is_training=True # Critical: enables training mode
)
GenCast (Diffusion Model)
from weathernext.weathernext1_gen import gencast
# For diffusion-based GenCast, adjust noise schedule for new variable statistics
model = gencast.GenCast(
config=gencast.DefaultConfig(),
is_training=True
)
Implementing the Training Loop
WeatherNext uses JAX for automatic differentiation and Optax for optimization. The Predictor base class in weathernext/utils/predictor_base.py provides the loss() method that handles variable-wise preprocessing internally.
Minimal Training Implementation
import jax
import optax
from weathernext.utils import checkpoint
# 1. Optimizer configuration
LEARNING_RATE = 1e-4
optimizer = optax.adam(learning_rate=LEARNING_RATE)
opt_state = optimizer.init(model.parameters)
# 2. Loss function using Predictor API
def compute_loss(params, batch):
"""Wrapper around predictor.loss for gradient computation."""
loss, diagnostics = model.loss(
inputs=batch.inputs,
targets=batch.targets,
forcings=batch.forcings
)
# Mean across batch and spatial dimensions
return loss.mean(), diagnostics
# 3. JIT-compiled training step
@jax.jit
def train_step(params, opt_state, batch):
(loss, diagnostics), grads = jax.value_and_grad(
compute_loss,
has_aux=True
)(params, batch)
updates, opt_state = optimizer.update(grads, opt_state, params)
params = optax.apply_updates(params, updates)
return params, opt_state, loss, diagnostics
# 4. Training loop with checkpointing
NUM_EPOCHS = 100
for epoch in range(NUM_EPOCHS):
epoch_losses = []
for batch in train_task.iter_batches(batch_size=1):
model.parameters, opt_state, loss, diag = train_step(
model.parameters,
opt_state,
batch
)
epoch_losses.append(float(loss))
avg_loss = sum(epoch_losses) / len(epoch_losses)
print(f'Epoch {epoch:03d}: avg_loss={avg_loss:.6f}')
# Periodic checkpointing every 10 epochs
if epoch % 10 == 0:
checkpoint.save_checkpoint(
model.parameters,
f'/tmp/weathernext_finetuned_epoch{epoch}.ckpt'
)
The Task.iter_batches() helper yields properly-structured minibatches that respect WeatherNext's dimensional requirements, applying any necessary masking for missing values.
Adjusting the Noise Schedule (GenCast Only)
When fine-tuning the diffusion model on variables with substantially different variance characteristics, update the noise schedule in weathernext/weathernext1_gen/gencast.py:
# In your training script, after model initialization
from weathernext.weathernext1_gen import gencast
# Access and modify noise configuration
noise_cfg = model._noise_config # Internal configuration object
# Set based on your custom variable statistics:
noise_cfg.training_min_noise_level = 0.1 # Increased for high-variance variables
noise_cfg.training_max_noise_level = 1.0 # Upper bound
noise_cfg.training_noise_level_rho = 7.0 # Density of noise levels
These values control the diffusion corruption process during training. Variables with larger natural variability typically benefit from higher minimum noise levels to prevent mode collapse.
Saving and Loading Fine-Tuned Checkpoints
The checkpoint utilities in weathernext/utils/checkpoint.py handle parameter serialization:
from weathernext.utils import checkpoint
# Save final fine-tuned model
final_path = '/path/to/weathernext_custom_vars.ckpt'
checkpoint.save_checkpoint(model.parameters, final_path)
# Restore for inference or continued training
restored_params = checkpoint.restore_checkpoint(final_path)
# Verify restoration
loaded_model = fgn.FGN(config=model_cfg, is_training=False)
loaded_model.parameters = restored_params
Validating Fine-Tuned Performance
Run inference on a held-out validation slice to confirm proper handling of custom variables:
# Select validation time slice
val_inputs = train_task.inputs.isel(time=slice(-10, None))
val_targets = train_task.targets.isel(time=slice(-10, None))
val_forcings = train_task.forcings.isel(time=slice(-10, None))
# Generate predictions
predictions = loaded_model(
inputs=val_inputs,
targets_template=val_targets, # Provides target structure without values
forcings=val_forcings
)
# Validate custom variable output
assert 'custom_ssr' in predictions.data_vars
assert predictions['custom_ssr'].shape == val_targets['custom_ssr'].shape
# Quick sanity: verify predictions are not constant
pred_std = float(predictions['custom_ssr'].std())
print(f'Custom variable prediction std: {pred_std:.4f}')
assert pred_std > 0, "Predictions collapsed to constant value"
Troubleshooting Custom Variable Integration
| Symptom | Cause | Solution |
|---|---|---|
KeyError on variable name |
Name mismatch between catalogue and dataset | Verify exact string match in variables.py |
| NaN loss values | Improper normalization statistics | Recompute mean/std on your dataset's training split |
| Poor prediction quality for new variable | Variable not in TARGET_VARIABLES |
Add to target list in Task definition |
| GenCast training instability | Noise schedule inappropriate for variance | Adjust training_min_noise_level upward |
Summary
Fine-tuning WeatherNext on custom datasets with different variables requires:
- Register new variables in
weathernext/utils/variables.pywith appropriate normalization - Define a Task in
weathernext/utils/task.pymapping inputs, targets, and forcings - Instantiate with
is_training=Trueusingfgn.FGNorgencast.GenCast - Train with JAX + Optax, leveraging the built-in
loss()method fromPredictor - Adjust noise schedules for GenCast diffusion models when variable statistics differ significantly
- Save checkpoints via
weathernext/utils/checkpoint.pyfor production deployment
The modular design preserves all training infrastructure—simply extend the variable catalogue and reconfigure the task.
Frequently Asked Questions
Can I fine-tune WeatherNext on variables with different spatial resolutions?
WeatherNext expects consistent lat/lon grids across all variables. Resample your custom data to match the model's native resolution (typically 0.25° or 1.0°) using xarray's interpolation methods before task construction. The architecture itself does not handle multi-resolution inputs natively.
How do I handle missing values in custom variables?
The Task class applies automatic masking for NaN values during batch generation. Ensure your xarray.Dataset uses np.nan for missing data rather than sentinel values. For variables with systematic gaps (e.g., satellite coverage holes), consider adding a binary mask variable to the inputs.
What's the minimum dataset size for effective fine-tuning?
While the original WeatherNext models trained on decades of ERA5 data, fine-tuning can succeed with 6–12 months of representative samples if you freeze early layers and use aggressive data augmentation. For entirely new variables, 2+ years of data helps the model learn stable climatological patterns.
Can I mix pretrained and newly initialized variables in the same model?
Yes. The FGN architecture in weathernext/weathernext2/architecture.py supports partial checkpoint restoration. Load pretrained parameters for standard variables, then let JAX initialize new variable embeddings randomly. The first training epochs will rapidly adapt the new embeddings while preserving atmospheric knowledge in frozen or low-learning-rate layers.
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 →