# How to Implement Custom Loss Functions in the WeatherNext Training Pipeline

> Learn to implement custom loss functions in WeatherNext. Create a callable or subclass a predictor to easily integrate your custom loss for improved model training.

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

---

**Implement custom loss functions in WeatherNext by creating a callable that returns a `(loss: xarray.DataArray, diagnostics: dict)` tuple and passing it to a predictor's constructor, or subclass a predictor and override its `loss` or `loss_and_predictions` methods.**

WeatherNext's training infrastructure is built on a clean, extensible loss-function interface that makes swapping in custom objectives straightforward. Whether you're experimenting with novel regularization terms or task-specific metrics, the `google-deepmind/weathernext` repository provides two primary integration paths grounded in the `PredictorBase` abstraction and the `LossAndDiagnostics` contract.

## Understanding the WeatherNext Loss Function Interface

At the core of WeatherNext's training pipeline lies **`PredictorBase`** in [`weathernext/utils/predictor_base.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/predictor_base.py). This base class defines the standard API that all trainable models must implement. The training loop interacts with models through two key methods:

- **`loss(prediction, target, **kwargs)`**
- **`loss_and_predictions(inputs, target, **kwargs)`**

Both methods must return a **`LossAndDiagnostics`** pair: a tuple containing a 1-D `xarray.DataArray` (loss per batch element) and a dictionary of diagnostic metrics for logging.

The contract is formally defined in **[`weathernext/utils/losses.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/losses.py)**, which also houses the built-in loss implementations like **`weighted_mse_per_level`** and **`weighted_mae_per_level`**. These are thin wrappers around the core helper **`weighted_loss_per_level`**, which handles latitude and pressure-level weighting.

## Method 1: Pass a Custom Loss Callable to a Predictor

The simplest approach is to implement a standalone loss function and inject it at model construction. This works seamlessly with GraphCast-style models that accept a `loss_fn` argument.

### Required Signature

Your custom loss must match this interface:

```python
def my_loss(
    prediction: xarray.DataArray,
    target: xarray.DataArray,
    **kwargs
) -> Tuple[xarray.DataArray, dict]:
    ...
    return loss, diagnostics

```

### Example: Custom L1 + Regularization Loss

```python

# my_custom_loss.py

import xarray
import jax.numpy as jnp
from weathernext.utils import losses

def my_weighted_l1_plus_reg(
    prediction: xarray.DataArray,
    target: xarray.DataArray,
    *,
    reg_coeff: float = 0.1,
    **kwargs
) -> losses.LossAndDiagnostics:
    """L1 loss weighted by latitude/pressure + L2 regularization on predictions."""
    # Core L1 term using WeatherNext's built-in helper

    mae, diagnostics = losses.weighted_mae_per_level(prediction, target, **kwargs)
    
    # Additional regularization on prediction magnitudes

    reg = reg_coeff * jnp.mean(jnp.square(prediction))
    reg_per_batch = xarray.DataArray(jnp.full_like(mae, reg), dims=mae.dims)
    
    total_loss = mae + reg_per_batch
    diagnostics["reg_l2"] = reg
    return total_loss, diagnostics

```

### Integration with GraphCast

```python

# train_script.py

from weathernext.weathernext1_graph import graphcast
from my_custom_loss import my_weighted_l1_plus_reg

model = graphcast.GraphCast(
    mesh=my_mesh,
    loss_fn=my_weighted_l1_plus_reg,  # Inject custom loss here

    # ... other hyperparameters

)

```

As implemented in [`weathernext/weathernext1_graph/graphcast.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/weathernext1_graph/graphcast.py), the `loss_and_predictions` method delegates to the injected `loss_fn`, making this the most non-invasive integration path.

## Method 2: Subclass a Predictor and Override Loss Methods

For finer control, subclass an existing predictor and override `loss` or `loss_and_predictions` directly. This pattern is useful when you need to modify preprocessing, add auxiliary outputs, or change how diagnostics are computed.

### Example: Customizing the FGN Predictor

```python

# custom_fgn.py

from weathernext.weathernext2 import fgn
from my_custom_loss import my_weighted_l1_plus_reg

class CustomFgn(fgn.Fgn):
    def __init__(self, *args, **kwargs):
        super().__init__(*args, **kwargs)
        # Replace the default MAE with custom implementation

        self._mae_function = my_weighted_l1_plus_reg

```

In [`weathernext/weathernext2/fgn.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/weathernext2/fgn.py), the `Fgn` class stores its loss function as `self._mae_function` and invokes it during training. By overriding this attribute, your custom logic integrates automatically with the existing `PredictorBase` call chain—no training loop modifications required.

## Key Implementation Requirements

### 1. Return Shape and Type

- **`loss`**: Must be a 1-D `xarray.DataArray` with dimensions matching the batch axis. The training pipeline aggregates this across batch elements.
- **`diagnostics`**: A `dict` of scalars or per-variable metrics. Keys appear in TensorBoard/MLflow logs.

### 2. Leverage Built-in Weighting Utilities

WeatherNext provides latitude and level weighting functions you should reuse for physical consistency:

```python
from weathernext.utils.losses import (
    normalized_latitude_weights,
    normalized_level_weights,
    weighted_loss_per_level
)

# Use weighted_loss_per_level with your own elemental loss

my_loss, diagnostics = weighted_loss_per_level(
    prediction,
    target,
    loss_fn=lambda p, t: jnp.abs(p - t),  # Your elemental loss

    **kwargs
)

```

### 3. Handle Missing Data Properly

Follow the pattern in existing losses: specify `loss_for_nan_targets=0.0` to ensure missing ground-truth values don't corrupt gradients. The `weighted_loss_per_level` helper handles this via the `loss_for_nan_targets` parameter.

### 4. Minimal Custom Loss Example

```python
def l2_loss(prediction, target, **kwargs):
    """Simple unweighted L2 loss per batch element."""
    loss = ((prediction - target) ** 2).mean(dim=("level", "lat", "lon"))
    return loss, {"l2": float(loss.mean())}

```

### 5. Per-Variable Weighting Example

```python
def weighted_variable_loss(prediction, target, var_weights, **kwargs):
    """Apply custom weights to different atmospheric variables."""
    diff = jnp.abs(prediction - target)
    weighted = sum(
        var_weights.get(v, 1.0) * diff.sel(var=v)
        for v in diff.var.values
    )
    loss = weighted.mean(dim=("level", "lat", "lon"))
    return loss, {"variable_weighted": float(loss.mean())}

```

## Training Loop Integration

Once your custom loss is injected, the standard training scripts work unchanged. The `PredictorBase` abstraction ensures the training loop in `*_train.py` scripts calls your loss implementation through the generic interface:

```python

# Standard training pattern—no modifications needed

from weathernext.training import train_loop

train_loop(model, train_dataset, val_dataset)

```

The pipeline automatically handles:
- Gradient computation via JAX's autodiff
- Per-batch loss aggregation
- Diagnostic logging to configured backends

## Summary

- **WeatherNext custom loss functions** must return `(xarray.DataArray, dict)` matching the `LossAndDiagnostics` contract in [`weathernext/utils/losses.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/losses.py).
- **Injection method**: Pass to constructors like `GraphCast(loss_fn=my_loss)`—cleanest for most use cases.
- **Override method**: Subclass predictors like `Fgn` and replace `loss` or stored loss functions for deeper customization.
- **Reuse utilities**: Leverage `weighted_loss_per_level`, `normalized_latitude_weights`, and `normalized_level_weights` for physical consistency.
- **No training loop changes**: `PredictorBase` abstracts loss invocation, so standard scripts work immediately.

## Frequently Asked Questions

### What file defines the loss function interface in WeatherNext?

The interface is defined in **[`weathernext/utils/losses.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/losses.py)**, which contains the `LossFn` protocol, `LossAndDiagnostics` dataclass, and built-in weighted MAE/MSE implementations. The `PredictorBase` class in [`weathernext/utils/predictor_base.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/predictor_base.py) enforces this contract across all models.

### Can I use losses that don't use latitude or level weighting?

Yes. The weighting helpers are optional convenience functions. Your custom loss can perform any reduction over spatial and vertical dimensions, provided the returned loss is 1-D over the batch dimension. Skip `weighted_loss_per_level` and implement your own reduction logic directly.

### How do I add logging metrics beyond the base loss?

Include additional entries in the `diagnostics` dictionary returned by your loss function. These automatically propagate to the training logs. For example: `diagnostics["custom_metric"] = float(my_computed_value)`. The `PredictorBase` implementation passes these through to the configured logger.

### Do custom losses work with distributed training?

Yes. The `PredictorBase` loss methods return per-batch-element losses as `xarray.DataArray` objects, which the training pipeline aggregates appropriately for distributed gradient computation. Ensure your loss maintains array shapes compatible with JAX's `pmap` or `pjit` transforms—avoid Python scalar reductions that break batch parallelism.