credit.postblock#

Submodules#

Attributes#

Functions#

__getattr__(name)

register_postblock(block_type)

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

build_postblocks(→ torch.nn.ModuleDict)

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

apply_postblocks(→ dict)

Apply a postblock group built by build_postblocks.

Package Contents#

credit.postblock.logger#
credit.postblock.__getattr__(name)#
credit.postblock.register_postblock(block_type)#

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

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

Parameters:

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

Example:

from credit.postblock import register_postblock
from credit.postblock.base import BasePostblock

@register_postblock("my_postblock")
class MyPostBlock(BasePostblock):
    def forward(self, batch: dict) -> dict:
        ...
credit.postblock.build_postblocks(conf: dict, phase: str = 'per_step') → torch.nn.ModuleDict#

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

Config format:

postblocks:
  per_step:          # run after every forward pass in the rollout loop
    reconstruct:
      type: reconstruct
    inverse_scale:
      type: bridgescaler_transform
      args:
        method: inverse_transform
        scaler_path: /path/to/scaler.json
  post_rollout:      # run once after all rollout steps complete
    mass_fixer:
      type: global_mass_fixer
      args: ...

Typical usage — build once per phase, store separately:

step_postblocks    = build_postblocks(conf, phase="per_step")
rollout_postblocks = build_postblocks(conf, phase="post_rollout")

# inside rollout loop, after each forward pass:
full_data_dict = apply_postblocks(step_postblocks, full_data_dict)

# once after rollout loop completes:
apply_postblocks(rollout_postblocks, full_data_dict)
Parameters:
  • conf – the full CREDIT config dict.

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

Returns:

nn.ModuleDict of instantiated blocks for the requested phase.

Raises:

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

credit.postblock.apply_postblocks(postblocks: torch.nn.ModuleDict, batch_dict: dict) → dict#

Apply a postblock group built by build_postblocks.

Parameters:
  • postblocks – nn.ModuleDict built by build_postblocks for a single phase.

  • batch_dict – dict containing at minimum "y_pred" and "metadata".

Returns:

The same batch_dict after all blocks in the group have run.