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 weight1 / sigma_v ** scale_power, withsigma_vread from the fitted bridgescaler atscaler_path.scale_poweris 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 by1 / sigma^2would over-weight low-variance variables quadratically for the linear metrics and is meaningless for the dimensionless ones.
manual: weights come fromvariable_weightsalone.
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#
Abstract base class for a univariate per-variable verification metric. |
|
Container metric combining several |
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.ABCAbstract base class for a univariate per-variable verification metric.
forwardtakes the trainer’sfull_data_dictand scoresfull_data_dict["y_processed"]againstfull_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). Unlikecredit.losses.base.BaseLoss, metrics run undertorch.no_gradand return detached Python floats for logging; there is nolearnableweighting mode and no gradient/detach check.Subclasses implement
compute_variable()(elementwise error tensor) and optionally overridereduce()(default identity; e.g. RMSE takessqrt). 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_processedbut 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
latitudecoordinate (required whenuse_latitude_weightsis True).channel_schema – optional
ChannelSchemafixing 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.ModuleContainer metric combining several
BaseVariableMetricsubclasses.Holds one
BaseVariableMetricper registry name listed inmetricsand shares the cross-metric weighting configuration (variable weighting, latitude weighting, computed-diagnostics policy, channel schema) across all of them.forwardreturns the union of each child’s scores, namespaced by each child’smetric_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
latitudecoordinate.channel_schema – optional
ChannelSchemafixing 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#