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 times
  • batch — independent samples (usually 1 for fine-tuning)
  • lat, lon — geographical grid coordinates
  • level (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:

  1. The name field must exactly match the DataArray name in your xarray.Dataset
  2. Choose normalization statistics representative of your dataset's climatology
  3. 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.py with appropriate normalization
  • Define a Task in weathernext/utils/task.py mapping inputs, targets, and forcings
  • Instantiate with is_training=True using fgn.FGN or gencast.GenCast
  • Train with JAX + Optax, leveraging the built-in loss() method from Predictor
  • Adjust noise schedules for GenCast diffusion models when variable statistics differ significantly
  • Save checkpoints via weathernext/utils/checkpoint.py for 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:

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 →