credit.preblock.scaler#
Classes#
Scaling preblock using a dictionary of bridgescaler scalers to fit and transform CREDIT state dictionaries. |
Functions#
|
Recursively move all tensor statistics in a nested scaler dict to CPU. |
|
Combine a list of nested scaler dicts into a single merged scaler dict. |
Module Contents#
- credit.preblock.scaler.move_scaler_dict_to_cpu(scaler_dict)#
Recursively move all tensor statistics in a nested scaler dict to CPU.
Required before
gather_objectin multi-GPU runs: each rank’s scaler holds tensors on its own device, and combining scalers from different devices raises a device-mismatch error.
- credit.preblock.scaler.combine_scaler_dicts(scaler_dicts)#
Combine a list of nested scaler dicts into a single merged scaler dict.
- Parameters:
scaler_dicts (list) – Nested scaler dicts with identical structure, e.g. the per-rank results gathered in
credit.applications.preprocess.- Returns:
A single nested scaler dict merging the statistics of every input.
- Return type:
dict
- class credit.preblock.scaler.BridgeScalerTransform(scaler_path: str, variables: list[str], method: str = 'transform', scaler_type: str = 'standard', scaler_params=None, data_types: list[str] = None, spatial_variables: list[str] = None)#
Bases:
credit.preblock.base.BasePreblockScaling preblock using a dictionary of bridgescaler scalers to fit and transform CREDIT state dictionaries.
Applies per-variable z-score scaling (or its inverse) to tensors in a nested batch dict of the form
batch[data_type][source][var_key].The scaler dict mirrors the batch structure — it is organized as
{data_type: {source: {var_key: scaler}}}. Variables that only exist in one data type (e.g. statics and dynamic forcings are"input"only) are therefore only scaled in that data type; no special configuration is needed. The scaler dict is produced by runningcredit preprocess.variablesaccepts full variable keys (e.g."era5/prognostic/3d/T"), partial paths (e.g."era5/prognostic"— expands to all variables whose key starts with that prefix), or an empty list (expands to all variables). The expansion is resolved lazily against the actual batch on the first forward call.data_typesscopes which data types this scaler instance processes. Its main use case is multi-step rollout training, where the input is already scaled from the previous step and only the target needs to be scaled. Omit it (or pass["input", "target"]) to scale both sides as usual.Config examples:
# Scale specific variables in both input and target (default behaviour) type: "bridgescaler_transform" args: scaler_path: "/path/to/scaler.json" variables: - "era5/prognostic/3d/T" - "era5/prognostic/3d/U" method: "transform" # Scale all variables (input and target) by leaving variables empty type: "bridgescaler_transform" args: scaler_path: "/path/to/scaler.json" variables: [] method: "transform" # Scale all prognostic variables using a partial path type: "bridgescaler_transform" args: scaler_path: "/path/to/scaler.json" variables: - "era5/prognostic" method: "transform" # Multi-step rollout: scale only the target (input already scaled) type: "bridgescaler_transform" args: scaler_path: "/path/to/scaler.json" variables: [] method: "transform" data_types: - "target" # Grid-wise scaling: variables in `spatial_variables` are scaled per # gridpoint (one mean/variance per lat/lon cell) rather than per level. # Their scaler's x_columns_/mean_x_/var_x_ must have length == n_lat * n_lon. # `spatial_variables` entries must also be covered by `variables`. type: "bridgescaler_transform" args: scaler_path: "/path/to/scaler.json" variables: [] method: "transform" spatial_variables: - "cesm/prognostic/2d/some_var"
- variables#
- variables_expanded = False#
- method = 'transform'#
- scaler_path#
- data_types = None#
- spatial_variables = []#
- scaler_params = None#
- forward(batch: dict) dict#
- fit_scaler_batch(batch: dict) dict#
Fit the scaler dictionary to a batch of data. This will usually be called with credit preprocess. Repeated calls to fit_scaler_batch will perform a running average of the scaler parameters.
- Parameters:
batch (dict) – state dictionary
- Returns:
The accumulated nested dict of fitted scalers (scaler[data_type][source][var_key]) merged across every call so far. The same dict is stored on self.scaler.
- Return type:
dict