credit.datasets.gen_2.channel_utils#

Everything about the channel/tensor layout contract for gen2: the canonical field-type rank, per-group channel slicing for rollout, and ChannelSchema (the frozen input/target layout saved at training time and reused at inference).

FIELD_TYPE_RANK / build_channel_layout / update_x#

Derive per-group channel slices directly from the variables config and provide a standalone update function for multi-step rollout.

Both the input tensor x and the target tensor y_pred are laid out source-major by ConcatToTensor — every channel of the first source precedes every channel of the second. Within a source the two differ:

  • x orders field types by FIELD_TYPE_RANK (prognostic < static < dynamic_forcing); diagnostics are never inputs.

  • y_pred is prognostic then diagnostic.

So with more than one source a field type’s channels are not contiguous across the whole tensor, and a source’s prognostic block in y_pred is offset by the diagnostics of every source before it. Each group therefore carries both of its positions: x_slice (where it lives in x) and src_slice (where its replacement comes from in y_pred or x_dynfrc). See ChannelGroup.

Usage:

from credit.datasets.gen_2.channel_utils import build_channel_layout, update_x

# once at trainer / rollout init
groups, n_pred = build_channel_layout(conf)

# at every t > 1 step — y_pred is the FULL prediction, diagnostics included
x = update_x(x, x_dynfrc, y_pred, groups)

ChannelSchema#

The explicit channel-layout contract between the data config, the flat tensors built by ConcatToTensor, and the nested dict rebuilt by Reconstruct.

Why this exists#

The channel order of the model’s flat input/output tensors was previously implicit, re-derived per batch from whatever tensors happened to be present. That broke at inference: with return_target=False the batch carries no diagnostic variables, so the fallback channel map covered prognostics only and Reconstruct silently dropped every diagnostic channel from y_pred.

ChannelSchema freezes the layout once, from the config, at training time:

  • input layout — mirrors ConcatToTensor’s sort: sources in config order; within each source, field types by FIELD_TYPE_RANK (prognostic < static < dynamic_forcing); 3d before 2d; config list order within. Diagnostics are never model inputs.

  • target layout — mirrors the dataset’s target insertion order (BaseDataset.__getitem__): sources in config order; prognostic then diagnostic; vars_3D then vars_2D; config list order within. This is the channel order of y and therefore of y_pred.

Lifecycle#

Training saves the schema to {save_loc}/channel_schema.yaml (rank 0) and validates it once against the target channel map actually built from data. Inference loads that file (falling back to config derivation with a warning) and hands it to ConcatToTensor, which uses the schema’s target map whenever the batch has no target — so Reconstruct recovers every predicted variable, diagnostics included.

Attributes#

Classes#

ChannelGroup

One (source, field_type) run of channels in the model input tensor.

ChannelSchema

Frozen channel layout for the flat model input and output tensors.

Functions#

resolve_num_levels(→ int)

Resolve a source's vertical level count.

resolve_level_ids(→ list)

Resolve the coordinate ids for a source's vertical levels.

build_channel_layout(conf)

Return (groups, n_pred) derived from the variables config.

update_x(x_prev, x_dynfrc, y_pred, groups)

Build the next-step input tensor for autoregressive rollout.

Module Contents#

credit.datasets.gen_2.channel_utils.logger#
credit.datasets.gen_2.channel_utils.FIELD_TYPE_RANK: dict[str, int]#
credit.datasets.gen_2.channel_utils.resolve_num_levels(src_conf: dict, conf: dict) → int#

Resolve a source’s vertical level count.

A source’s explicit levels list wins when non-empty; otherwise the count falls back to model.levels. That fallback covers the levels: null and omitted-levels cases, where the dataset loads every vertical level available in the data (see LocalDataset). An empty list is treated the same as absent, matching build_channel_layout / ChannelSchema.from_config. Returns 0 when neither is resolvable.

Parameters:
  • src_conf – A single data.source.<name> config block.

  • conf – The full config dict (for the model.levels fallback).

credit.datasets.gen_2.channel_utils.resolve_level_ids(src_conf: dict, conf: dict) → list#

Resolve the coordinate ids for a source’s vertical levels.

Returns the explicit levels list when set; otherwise integer indices 0 .. model.levels - 1 (the levels: null / omitted case). The physical level values are unknown from the config alone in that case, so positional indices stand in for the output coordinate. Mirrors resolve_num_levels() for the “all levels” fallback.

Parameters:
  • src_conf – A single data.source.<name> config block.

  • conf – The full config dict (for the model.levels fallback).

class credit.datasets.gen_2.channel_utils.ChannelGroup#

One (source, field_type) run of channels in the model input tensor.

Variables:
  • source – Source name as it appears under data.source.

  • field_type – "prognostic", "static" or "dynamic_forcing".

  • x_slice – Where these channels sit in the model input tensor x.

  • src_slice – Where their replacement comes from at the next rollout step — into y_pred for prognostic groups, into x_dynfrc for dynamic-forcing groups, and None for fixed groups (static), which are carried forward unchanged.

source: str#
field_type: str#
x_slice: slice#
src_slice: slice | None#
credit.datasets.gen_2.channel_utils.build_channel_layout(conf)#

Return (groups, n_pred) derived from the variables config.

Handles any number of sources. ConcatToTensor lays both x and y_pred out source-major, so a field type’s channels are contiguous only within a source — hence one entry per (source, field_type) pair rather than one per field type.

Each source resolves its own level count (data.source.<name>.levels, falling back to model.levels), so sources may differ in vertical extent.

Parameters:

conf (dict) – Full CREDIT config dict.

Returns:

  • groups (dict[str, ChannelGroup]) – Keyed "{source}/{field_type}", in x channel order. Diagnostic groups are excluded (never inputs), but their width is still accounted for when computing prognostic offsets into y_pred.

  • n_pred (int) – Total number of predicted (prognostic) channels across all sources.

Raises:

ValueError – If a source declares 3D variables but no level count is resolvable, or if history_len is greater than 1 — the flat rollout paths overwrite channels in place and cannot shift a history window.

credit.datasets.gen_2.channel_utils.update_x(x_prev, x_dynfrc, y_pred, groups)#

Build the next-step input tensor for autoregressive rollout.

Replaces dataset-driven channels (dynamic_forcing) and model-predicted channels (prognostic) in x_prev. Fixed channels (static, etc.) are carried forward unchanged via clone.

Each group carries an explicit source slice, so this is correct however the sources interleave — it never assumes that a field type’s channels are contiguous or that they arrive in the same order as the destination.

Parameters:
  • x_prev (Tensor [B, C, ...]) – Full input tensor from the previous step.

  • x_dynfrc (Tensor [B, C_dyn, ...]) – New dynamic-forcing channels from the dataset, source-major.

  • y_pred (Tensor [B, C_out, ...]) – The full model output — prognostic and diagnostic channels, exactly as the model emits them. Do not pre-slice it down to the prognostics: with several sources those are not a leading contiguous block, and the group’s src_slice already indexes them.

  • groups (dict[str, ChannelGroup]) – From build_channel_layout().

Returns:

Updated input tensor ready for the next forward pass.

Return type:

Tensor [B, C, …]

credit.datasets.gen_2.channel_utils.DEFAULT_SCHEMA_FILENAME = 'channel_schema.yaml'#
class credit.datasets.gen_2.channel_utils.ChannelSchema(input_layout: list[dict], target_layout: list[dict])#

Frozen channel layout for the flat model input and output tensors.

Parameters:
  • input_layout – ordered list of {"var_key", "n_levels", "n_time"} dicts describing the concatenated input tensor.

  • target_layout – same, for the target / y_pred tensor.

input_layout#
target_layout#
classmethod from_config(conf: dict) → ChannelSchema#

Derive the schema from the full config dict.

n_levels for 3D variables comes from the source’s levels list when present, else from conf["model"]["levels"]. Sources that define 3D variables but resolve neither raise ValueError — for those configs the schema must come from a saved channel_schema.yaml.

Raises:

ValueError – if 3D variables exist but no level count is resolvable.

save(path: str) → None#

Write the schema as YAML (atomically: temp file + rename).

The staging file is per-process: concurrent writers (multiple ranks saving the schema at init) would otherwise interleave their dumps into one shared <path>.tmp and rename a corrupted file into place.

classmethod load(path: str) → ChannelSchema#
classmethod load_or_from_config(conf: dict, save_loc: str | None = None) → ChannelSchema | None#

Load channel_schema.yaml from save_loc, else derive from config.

Returns None (with a warning) when neither is possible, so callers can fall back to legacy per-batch behavior.

input_channel_map() → dict#

Channel map for the concatenated input tensor.

target_channel_map() → dict#

Channel map for the target / y_pred tensor (prognostic + diagnostic).

validate_channel_map(actual_map: dict, which: str = 'target') → None#

Check a data-derived channel map against the schema.

Compares key order, slices, and per-variable shapes. Raises with a side-by-side diff on mismatch — a mismatch means the schema (and any model trained against it) disagrees with what the dataset produces.

Parameters:
  • actual_map – {var_key: {"slice", "orig_shape"}} from ConcatToTensor.

  • which – "target" or "input" — selects the schema layout.

Raises:

ValueError – on any ordering, slicing, or shape difference.