How to Implement Custom Loss Functions in the WeatherNext Training Pipeline
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. 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, 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:
def my_loss(
prediction: xarray.DataArray,
target: xarray.DataArray,
**kwargs
) -> Tuple[xarray.DataArray, dict]:
...
return loss, diagnostics
Example: Custom L1 + Regularization Loss
# 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
# 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, 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
# 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, 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-Dxarray.DataArraywith dimensions matching the batch axis. The training pipeline aggregates this across batch elements.diagnostics: Adictof 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:
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
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
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:
# 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 theLossAndDiagnosticscontract inweathernext/utils/losses.py. - Injection method: Pass to constructors like
GraphCast(loss_fn=my_loss)—cleanest for most use cases. - Override method: Subclass predictors like
Fgnand replacelossor stored loss functions for deeper customization. - Reuse utilities: Leverage
weighted_loss_per_level,normalized_latitude_weights, andnormalized_level_weightsfor physical consistency. - No training loop changes:
PredictorBaseabstracts 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, which contains the LossFn protocol, LossAndDiagnostics dataclass, and built-in weighted MAE/MSE implementations. The PredictorBase class in 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.
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 →