credit.datasets.gen_2.channel_utils
===================================

.. py:module:: credit.datasets.gen_2.channel_utils

.. autoapi-nested-parse::

   channel_utils.py
   ----------------
   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 :class:`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
----------

.. autoapisummary::

   credit.datasets.gen_2.channel_utils.logger
   credit.datasets.gen_2.channel_utils.FIELD_TYPE_RANK
   credit.datasets.gen_2.channel_utils.DEFAULT_SCHEMA_FILENAME


Classes
-------

.. autoapisummary::

   credit.datasets.gen_2.channel_utils.ChannelGroup
   credit.datasets.gen_2.channel_utils.ChannelSchema


Functions
---------

.. autoapisummary::

   credit.datasets.gen_2.channel_utils.resolve_num_levels
   credit.datasets.gen_2.channel_utils.resolve_level_ids
   credit.datasets.gen_2.channel_utils.build_channel_layout
   credit.datasets.gen_2.channel_utils.update_x


Module Contents
---------------

.. py:data:: logger

.. py:data:: FIELD_TYPE_RANK
   :type:  dict[str, int]

.. py:function:: 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.

   :param src_conf: A single ``data.source.<name>`` config block.
   :param conf: The full config dict (for the ``model.levels`` fallback).


.. py:function:: 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
   :func:`resolve_num_levels` for the "all levels" fallback.

   :param src_conf: A single ``data.source.<name>`` config block.
   :param conf: The full config dict (for the ``model.levels`` fallback).


.. py:class:: ChannelGroup

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

   :ivar source: Source name as it appears under ``data.source``.
   :ivar field_type: ``"prognostic"``, ``"static"`` or ``"dynamic_forcing"``.
   :ivar x_slice: Where these channels sit in the model input tensor ``x``.
   :ivar 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.



   .. py:attribute:: source
      :type:  str


   .. py:attribute:: field_type
      :type:  str


   .. py:attribute:: x_slice
      :type:  slice


   .. py:attribute:: src_slice
      :type:  slice | None


.. py:function:: 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.

   :param conf: Full CREDIT config dict.
   :type conf: 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.


.. py:function:: 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.

   :param x_prev: Full input tensor from the previous step.
   :type x_prev: Tensor  [B, C, ...]
   :param x_dynfrc: New dynamic-forcing channels from the dataset, source-major.
   :type x_dynfrc: Tensor  [B, C_dyn, ...]
   :param y_pred: 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.
   :type y_pred: Tensor  [B, C_out, ...]
   :param groups: From build_channel_layout().
   :type groups: dict[str, ChannelGroup]

   :returns: Updated input tensor ready for the next forward pass.
   :rtype: Tensor  [B, C, ...]


.. py:data:: DEFAULT_SCHEMA_FILENAME
   :value: 'channel_schema.yaml'


.. py:class:: ChannelSchema(input_layout: list[dict], target_layout: list[dict])

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

   :param input_layout: ordered list of ``{"var_key", "n_levels", "n_time"}``
                        dicts describing the concatenated input tensor.
   :param target_layout: same, for the target / ``y_pred`` tensor.


   .. py:attribute:: input_layout


   .. py:attribute:: target_layout


   .. py:method:: from_config(conf: dict) -> ChannelSchema
      :classmethod:


      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.



   .. py:method:: 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.



   .. py:method:: load(path: str) -> ChannelSchema
      :classmethod:



   .. py:method:: load_or_from_config(conf: dict, save_loc: str | None = None) -> ChannelSchema | None
      :classmethod:


      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.



   .. py:method:: input_channel_map() -> dict

      Channel map for the concatenated input tensor.



   .. py:method:: target_channel_map() -> dict

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



   .. py:method:: 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.

      :param actual_map: ``{var_key: {"slice", "orig_shape"}}`` from ConcatToTensor.
      :param which: ``"target"`` or ``"input"`` — selects the schema layout.

      :raises ValueError: on any ordering, slicing, or shape difference.



