credit.metrics.base#

Gen 2 per-variable verification metrics on the postblocks state dictionary.

The Gen 2 trainer’s per_step postblock chain turns the raw model output into full_data_dict["y_processed"] — a nested {source: {var_key: tensor}} dict in physical units (Reconstruct + inverse bridgescaler, plus any physics fixers). The metric classes in this module score each variable against the matching entry of full_data_dict["y_target_processed"] (the same chain applied to the flat target y) and combine the per-variable scores, mirroring credit.losses.base.BaseLoss but for evaluation rather than training: metrics run under torch.no_grad and return detached Python floats for logging.

Required postblocks config (per_step phase) — identical to BaseLoss:

postblocks:
  per_step:
    reconstruct:        {type: reconstruct, args: {detach: false}}
    scaler:             {type: bridgescaler_transform, args: {... method: inverse_transform}}
    reconstruct_target: {type: reconstruct, args: {in_key: "y", out_key: "y_target_processed"}}
    scaler_target:      {type: bridgescaler_transform, args: {... method: inverse_transform, key: "y_target_processed"}}

Metric config (conf["metrics"], mirroring the loss / preblocks {type, args} structure):

metrics:
  type: combined
  args:
    metrics:                         # registry names of univariate metrics
      rmse: {}
      mae: {}
      bias: {}
    var_weighting: "inverse_variance"  # inverse_variance | manual | none
    scaler_path: "/path/scaler.json"   # required for inverse_variance
    variable_weights: {}                # manual multipliers per var_key (all modes)
    normalize_weights: true            # rescale combination weights to mean 1
    include_computed_diagnostics: true  # score postblock-computed diagnostics
    use_latitude_weights: true          # cos(lat) spatial weighting per variable
    latitude_weights: "/path/static.zarr"

Scored variables: every variable in the data target layout (prognostic AND diagnostic variables from conf["data"]["source"]) is always scored. Variables that only appear in y_processed because a postblock computed them from prognostic variables (e.g. mslp_diagnostic) are “computed diagnostics”: they are scored only when include_computed_diagnostics: true, and then require a matching entry in y_target_processed.

var_weighting reuses the BaseLoss modes except learnable (metrics are not optimized):

  • inverse_variance (default): per-variable weight 1 / sigma_v ** scale_power, with sigma_v read from the fitted bridgescaler at scaler_path. scale_power is a class attribute of each metric giving the power of sigma its score carries, so the weighted aggregate is dimensionless regardless of the metric’s order: 2 for MSE, 1 for RMSE / MAE / bias / activity, and 0 for already-normalized scores such as R2 and ACC (which ignore the scaler entirely). Weighting every metric by 1 / sigma^2 would over-weight low-variance variables quadratically for the linear metrics and is meaningless for the dimensionless ones.

  • manual: weights come from variable_weights alone.

  • none: uniform combination; only sensible when variables are already on comparable scales.

variable_weights multipliers apply on top of every mode (default 1.0).

Code example — building a combined metric directly:

from credit.metrics.base import BaseCombinedMetric
from credit.datasets.gen_2.channel_utils import ChannelSchema

schema = ChannelSchema.from_config(conf)
metric = BaseCombinedMetric(
    channel_schema=schema,
    metrics={"rmse": {}, "mae": {}, "bias": {}},
    var_weighting="inverse_variance",
    scaler_path="/path/scaler.json",
    use_latitude_weights=True,
    latitude_weights="/path/static.zarr",
)
scores = metric(full_data_dict)   # {"rmse/ERA5/.../T": 1.2, "rmse": 0.9, ...}

Code example — registering a custom univariate metric:

from credit.metrics import register_metric
from credit.metrics.base import BaseVariableMetric

@register_metric("huber_metric")
class HuberMetric(BaseVariableMetric):
    def compute_variable(self, pred, target):
        return torch.nn.functional.huber_loss(pred, target, reduction="none")

And in config:

metrics:
  type: combined
  args:
    metrics: {rmse: {}, huber_metric: {delta: 1.0}}

Attributes#

Classes#

BaseVariableMetric

Abstract base class for a univariate per-variable verification metric.

BaseCombinedMetric

Container metric combining several BaseVariableMetric subclasses.

Module Contents#

credit.metrics.base.logger#
credit.metrics.base.METRIC_VAR_WEIGHTING_MODES = ('inverse_variance', 'manual', 'none')#
class credit.metrics.base.BaseVariableMetric(metric_name: str, var_weighting: str = 'inverse_variance', scaler_path: str | None = None, variable_weights: dict | None = None, normalize_weights: bool = True, include_computed_diagnostics: bool = True, use_latitude_weights: bool = False, latitude_weights: str | None = None, channel_schema=None, **kwargs)#

Bases: torch.nn.Module, abc.ABC

Abstract base class for a univariate per-variable verification metric.

forward takes the trainer’s full_data_dict and scores full_data_dict["y_processed"] against full_data_dict["y_target_processed"] variable by variable, then combines the per-variable scores across variables (see the module docstring for the full config format and weighting modes). Unlike credit.losses.base.BaseLoss, metrics run under torch.no_grad and return detached Python floats for logging; there is no learnable weighting mode and no gradient/detach check.

Subclasses implement compute_variable() (elementwise error tensor) and optionally override reduce() (default identity; e.g. RMSE takes sqrt). Latitude weighting is applied to the elementwise tensor before the spatial mean, exactly as in BaseLoss.

Parameters:
  • metric_name – name used to namespace the returned scores ("{metric_name}/{var_key}" and "{metric_name}" aggregate).

  • var_weighting – combination weighting — "inverse_variance" (default), "manual", or "none". "learnable" is rejected.

  • scaler_path – bridgescaler dict JSON used to read per-variable variances; required for inverse_variance.

  • variable_weights – manual multipliers per var_key (all modes, default 1.0).

  • normalize_weights – rescale static combination weights to mean 1.

  • include_computed_diagnostics – score postblock-computed diagnostics (variables present in y_processed but not in the data target layout). Data diagnostics are always scored.

  • use_latitude_weights – apply cos(lat) spatial weighting per variable.

  • latitude_weights – path to a dataset with a latitude coordinate (required when use_latitude_weights is True).

  • channel_schema – optional ChannelSchema fixing the data target variable layout; when None the scored variables are discovered from the state dict on the first forward pass.

Variables:
  • last_var_scores – {var_key: float} detached per-variable scores (pre-combination) from the most recent forward pass.

  • var_keys – scored variable list, resolved on the first forward pass.

  • scale_power – class attribute; see below.

scale_power: int = 2#
metric_name#
var_weighting = 'inverse_variance'#
manual_weights#
normalize_weights = True#
include_computed_diagnostics = True#
lat_weights = None#
data_var_keys = None#
var_keys = None#
last_var_scores#
abstractmethod compute_variable(pred: torch.Tensor, target: torch.Tensor) → torch.Tensor#

Return the elementwise (un-reduced) error tensor for one variable.

The tensor must broadcast against pred (typically the same shape); forward() applies latitude weighting and takes the spatial mean. For example, MSE returns (pred - target) ** 2.

reduce(score: torch.Tensor) → torch.Tensor#

Finalize the per-variable scalar after the spatial mean.

Default is identity; subclasses override (e.g. RMSE returns torch.sqrt).

forward(full_data_dict: dict) → dict#
class credit.metrics.base.BaseCombinedMetric(metrics: dict, var_weighting: str = 'inverse_variance', scaler_path: str | None = None, variable_weights: dict | None = None, normalize_weights: bool = True, include_computed_diagnostics: bool = True, use_latitude_weights: bool = False, latitude_weights: str | None = None, channel_schema=None, **kwargs)#

Bases: torch.nn.Module

Container metric combining several BaseVariableMetric subclasses.

Holds one BaseVariableMetric per registry name listed in metrics and shares the cross-metric weighting configuration (variable weighting, latitude weighting, computed-diagnostics policy, channel schema) across all of them. forward returns the union of each child’s scores, namespaced by each child’s metric_name (e.g. "rmse/ERA5/.../T", "rmse", "mae/ERA5/.../T", "mae").

Parameters:
  • metrics – {metric_name: per_metric_args} mapping. Each name must be in the metric registry (credit.metrics._METRIC_REGISTRY).

  • var_weighting – combination weighting passed to every child — "inverse_variance" (default), "manual", or "none".

  • scaler_path – bridgescaler dict JSON; required for inverse_variance.

  • variable_weights – manual multipliers per var_key (all modes, default 1.0).

  • normalize_weights – rescale static combination weights to mean 1.

  • include_computed_diagnostics – score postblock-computed diagnostics.

  • use_latitude_weights – apply cos(lat) spatial weighting per variable.

  • latitude_weights – path to a dataset with a latitude coordinate.

  • channel_schema – optional ChannelSchema fixing the data target variable layout.

Variables:

metric_modules – {metric_name: BaseVariableMetric} child instances.

Example (config):

metrics:
  type: combined
  args:
    metrics: {rmse: {}, mae: {}, bias: {}}
    var_weighting: "inverse_variance"
    scaler_path: "/path/scaler.json"
    use_latitude_weights: true
    latitude_weights: "/path/static.zarr"

Example (code):

metric = BaseCombinedMetric(
    channel_schema=schema,
    metrics={"rmse": {}, "mae": {}},
    var_weighting="none",
)
scores = metric(full_data_dict)
metric_modules#
forward(full_data_dict: dict) → dict#