credit.losses#
Submodules#
Attributes#
Functions#
|
|
|
Resolve the registry loss name a config actually trains with. |
|
|
|
Decorator that adds an external loss class to the loss registry. |
|
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 univariateargs.training_lossis 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 flattraining_losskey or the new-styletype.- 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.Moduleso that it can be used as a loss function in PyTorch training.- Parameters:
loss_type – Key used in the config
loss.training_lossorloss.validation_lossfield.
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,DownscalingLossis 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 anothertype-dispatched loss) for new-style sections,
DownscalingLossfor downscaling configs,VariableTotalLoss2Dwhen latitude/variable weights are enabled, or a standard/custom loss from the registry otherwise.
- A loss function instance:
- Return type:
torch.nn.Module
- Raises:
ValueError – If the requested loss type is not recognized in the registry.