credit.losses#

Submodules#

Attributes#

Functions#

is_crps_loss(loss_type)

effective_loss_name(loss_conf)

Resolve the registry loss name a config actually trains with.

__getattr__(name)

register_loss(loss_type)

Decorator that adds an external loss class to the loss registry.

load_loss(conf[, reduction, validation])

Load the appropriate loss function based on the configuration.

Package Contents#

credit.losses.logger#
credit.losses.CRPS_LOSSES#
credit.losses.is_crps_loss(loss_type)#
credit.losses.effective_loss_name(loss_conf)#

Resolve the registry loss name a config actually trains with.

For a new-style BaseLoss section (type: base) the univariate args.training_loss is what drives ensemble semantics (e.g. ring-crps needs shared batches and one member per dp rank), so it — not the literal "base" — is returned. Any other section resolves to the flat training_loss key or the new-style type.

Parameters:

loss_conf – the config’s loss: section (conf["loss"]).

Returns:

Registry loss name string, or None if the section is empty.

credit.losses.__getattr__(name)#
credit.losses.register_loss(loss_type)#

Decorator that adds an external loss class to the loss registry.

The class must inherit from torch.nn.Module so that it can be used as a loss function in PyTorch training.

Parameters:

loss_type – Key used in the config loss.training_loss or loss.validation_loss field.

Example:

import torch.nn as nn
from credit.losses import register_loss

@register_loss("my_loss")
class MyLoss(nn.Module):
    ...
credit.losses.load_loss(conf, reduction='none', validation=False)#

Load the appropriate loss function based on the configuration.

Determines whether to use a weighted custom loss wrapper (such as VariableTotalLoss2D) when latitude or variable weights are enabled, or to load a standard or custom loss directly from the registry.

If in validation mode and a separate validation loss is specified in the config, that loss type will be used. Otherwise, the training loss is used.

Parameters:
  • conf (dict) –

    Configuration dictionary. Must contain a ‘loss’ section in one of two formats:

    New-style (Gen 2, mirrors preblocks/postblocks)::
    loss:

    type: base # registry loss name args: {…} # constructor kwargs

    Legacy flat keys: - ‘training_loss’ (str): The primary loss function name. - ‘validation_loss’ (optional, str): An alternate loss for validation. - ‘use_latitude_weights’ (bool): Whether to use latitude-based weighting. - ‘use_variable_weights’ (bool): Whether to use variable-specific weighting. A downscaling config is detected by the presence of conf["data"]["datasets"] (the multi-source dataset key); when present, DownscalingLoss is returned regardless of the weight flags.

  • reduction (str, optional) – Reduction method to apply to the loss (‘mean’, ‘sum’, or ‘none’). Default is ‘none’.

  • validation (bool, optional) – Whether the loss is being used for validation. Defaults to False.

Returns:

A loss function instance: BaseLoss (or another type-dispatched

loss) for new-style sections, DownscalingLoss for downscaling configs, VariableTotalLoss2D when latitude/variable weights are enabled, or a standard/custom loss from the registry otherwise.

Return type:

torch.nn.Module

Raises:

ValueError – If the requested loss type is not recognized in the registry.