# How to Fine-Tune WeatherNext on Custom Datasets with Different Variables

> Learn how to fine-tune WeatherNext on custom datasets. Extend variables, define new tasks, and train using JAX and Optax for advanced weather forecasting.

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

---

**Fine-tune WeatherNext on custom atmospheric datasets by extending the variable catalogue in [`weathernext/utils/variables.py`](https://github.com/google-deepmind/weathernext/blob/main/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`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/weathernext2/architecture.py) | JAX-compatible neural net implementing the forecast forward pass |
| **Predictor API** | [`weathernext/utils/predictor_base.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/predictor_base.py) | Converts `xarray` datasets to JAX arrays; provides `loss()` method |
| **Task definition** | [`weathernext/utils/task.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/task.py) | Maps which variables are inputs, targets, and forcings |
| **Variable catalogue** | [`weathernext/utils/variables.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/variables.py) | Registry of all variables with normalization schemes |
| **Configuration handling** | [`weathernext/utils/fiddle_config_io.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/fiddle_config_io.py) | Fiddle config utilities for hyperparameter changes |
| **Checkpointing** | [`weathernext/utils/checkpoint.py`](https://github.com/google-deepmind/weathernext/blob/main/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`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/variables.py):

```python

# 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:

```python
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`](https://github.com/google-deepmind/weathernext/blob/main/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:

```python
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)

```python
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)

```python
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`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/predictor_base.py) provides the `loss()` method that handles variable-wise preprocessing internally.

### Minimal Training Implementation

```python
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`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/weathernext1_gen/gencast.py):

```python

# 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`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/checkpoint.py) handle parameter serialization:

```python
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:

```python

# 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`](https://github.com/google-deepmind/weathernext/blob/main/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`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/variables.py) with appropriate normalization
- **Define a Task** in [`weathernext/utils/task.py`](https://github.com/google-deepmind/weathernext/blob/main/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`](https://github.com/google-deepmind/weathernext/blob/main/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`](https://github.com/google-deepmind/weathernext/blob/main/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.