credit.preblock#

Submodules#

Attributes#

Functions#

__getattr__(name)

register_preblock(block_type)

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

build_preblocks(→ torch.nn.ModuleDict)

Instantiate preblocks for a single phase from the full CREDIT config.

attach_channel_schema(→ None)

Attach a ChannelSchema to every ConcatToTensor block in a preblock group.

apply_preblocks_before_scaler(preblocks, batch[, ...])

Build the (unscaled) batch a BridgeScalerTransform is fit on during credit preprocess.

apply_preblocks(→ dict)

Apply a preblock group built by build_preblocks.

Package Contents#

credit.preblock.logger#
credit.preblock.__getattr__(name)#
credit.preblock.register_preblock(block_type)#

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

The class must inherit from credit.preblock.base.BasePreblock so that the correct forward(batch: dict) -> dict signature is enforced.

Parameters:

block_type – Key used in the config preblocks.<phase>.<block>.type field.

Example:

from credit.preblock import register_preblock
from credit.preblock.base import BasePreblock

@register_preblock("my_preblock")
class MyPreBlock(BasePreblock):
    def forward(self, batch: dict) -> dict:
        ...
credit.preblock.build_preblocks(conf: dict, phase: str = 'per_step') → torch.nn.ModuleDict#

Instantiate preblocks for a single phase from the full CREDIT config.

Config format:

preblocks:
  ic_only:          # run once at t=0 on the raw batch (e.g. static regrid)
    regrid_static:
      type: regrid
      args: ...
  per_step:         # run every rollout step (e.g. log_transform, concat)
    log_transform:
      type: log_transform
    concat:
      type: concat

Typical usage — build once per phase, store separately:

ic_preblocks   = build_preblocks(conf, phase="ic_only")
step_preblocks = build_preblocks(conf, phase="per_step")

# t=0: run both in sequence
ic_preprocessed    = apply_preblocks(ic_preblocks, batch, device=device)
preprocessed_batch = apply_preblocks(step_preblocks, ic_preprocessed, device=device)

# t>0: run per_step only
preprocessed_batch = apply_preblocks(step_preblocks, rollout_batch, device=device)
Parameters:
  • conf – the full CREDIT config dict.

  • phase – which section to build — "ic_only" or "per_step".

Returns:

nn.ModuleDict of instantiated blocks for the requested phase.

Raises:

ValueError – if the config contains keys other than "ic_only" / "per_step", if phase is not one of those values, or if a block’s type value is not in _PREBLOCK_REGISTRY.

credit.preblock.attach_channel_schema(preblocks: torch.nn.ModuleDict, schema) → None#

Attach a ChannelSchema to every ConcatToTensor block in a preblock group.

The schema is runtime state (built from config / loaded from save_loc), not a config arg, so it is injected after build_preblocks rather than through the registry. No-op for groups without a concat block or when schema is None.

credit.preblock.apply_preblocks_before_scaler(preblocks: torch.nn.ModuleDict, batch: dict, device=None, stop_key=None)#

Build the (unscaled) batch a BridgeScalerTransform is fit on during credit preprocess.

Applies the non-scaler preblocks (log, regrid, rename, …) that precede the target scaler and returns the transformed batch. Scaler blocks are always skipped: during preprocessing their scalers are still being fit and cannot transform, and because scalers operate on disjoint variable sets an earlier scaler never touches a later scaler’s variables — skipping it leaves each scaler’s fit batch with the correct raw values.

Preblocks never mutate the caller’s dict (they shallow-copy via _copy_batch and never write tensors in place), so this may be called repeatedly on the same raw batch — once per scaler — without the calls interfering.

Parameters:

stop_key – name of the scaler block to stop before; preblocks up to (but not including) this key are applied. When None, stop at the first BridgeScalerTransform encountered (single-scaler legacy behavior).

credit.preblock.apply_preblocks(preblocks: torch.nn.ModuleDict, batch: dict, device=None) → dict#

Apply a preblock group built by build_preblocks.

Parameters:
  • preblocks – nn.ModuleDict built by build_preblocks for a single phase.

  • batch – nested variable dict from the dataset (or a prior preblock pass).

  • device – move output tensors here after concat.

Returns:

{"x": tensor} plus optional "y" (if a target was produced) and "metadata" (if metadata was produced) keys. Otherwise: the transformed nested batch dict (pre-concat).

Return type:

When concat has run