credit.preblock#
Submodules#
Attributes#
Functions#
|
|
|
Decorator that adds an external preblock class to the preblock registry. |
|
Instantiate preblocks for a single phase from the full CREDIT config. |
|
Attach a |
|
Build the (unscaled) batch a BridgeScalerTransform is fit on during |
|
Apply a preblock group built by |
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.BasePreblockso that the correctforward(batch: dict) -> dictsignature is enforced.- Parameters:
block_type – Key used in the config
preblocks.<phase>.<block>.typefield.
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.ModuleDictof instantiated blocks for the requested phase.- Raises:
ValueError – if the config contains keys other than
"ic_only"/"per_step", ifphaseis not one of those values, or if a block’stypevalue is not in_PREBLOCK_REGISTRY.
- credit.preblock.attach_channel_schema(preblocks: torch.nn.ModuleDict, schema) None#
Attach a
ChannelSchemato 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_preblocksrather 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_batchand never write tensors in place), so this may be called repeatedly on the same rawbatch— 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 firstBridgeScalerTransformencountered (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.ModuleDictbuilt bybuild_preblocksfor 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