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:
xorders field types by FIELD_TYPE_RANK (prognostic < static < dynamic_forcing); diagnostics are never inputs.y_predis 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 byFIELD_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;prognosticthendiagnostic;vars_3Dthenvars_2D; config list order within. This is the channel order ofyand therefore ofy_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#
One |
|
Frozen channel layout for the flat model input and output tensors. |
Functions#
|
Resolve a source's vertical level count. |
|
Resolve the coordinate ids for a source's vertical levels. |
|
Return (groups, n_pred) derived from the variables config. |
|
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
levelslist wins when non-empty; otherwise the count falls back tomodel.levels. That fallback covers thelevels: nulland omitted-levelscases, where the dataset loads every vertical level available in the data (seeLocalDataset). An empty list is treated the same as absent, matchingbuild_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.levelsfallback).
- 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
levelslist when set; otherwise integer indices0 .. model.levels - 1(thelevels: 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. Mirrorsresolve_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.levelsfallback).
- 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_predfor prognostic groups, intox_dynfrcfor dynamic-forcing groups, andNonefor 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.
ConcatToTensorlays bothxandy_predout 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 tomodel.levels), so sources may differ in vertical extent.- Parameters:
conf (dict) – Full CREDIT config dict.
- Returns:
groups (dict[str, ChannelGroup]) – Keyed
"{source}/{field_type}", inxchannel order. Diagnostic groups are excluded (never inputs), but their width is still accounted for when computing prognostic offsets intoy_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_lenis 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_slicealready 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_predtensor.
- input_layout#
- target_layout#
- classmethod from_config(conf: dict) ChannelSchema#
Derive the schema from the full config dict.
n_levelsfor 3D variables comes from the source’slevelslist when present, else fromconf["model"]["levels"]. Sources that define 3D variables but resolve neither raiseValueError— for those configs the schema must come from a savedchannel_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>.tmpand 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.yamlfromsave_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_predtensor (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.