credit.models.wxformer.wxformer_next
====================================

.. py:module:: credit.models.wxformer.wxformer_next

.. autoapi-nested-parse::

   NextGen WXFormer: CrossFormer U-Net + spectral GNN bottleneck + column attention + level embeddings.

   Design additions over WXFormer:
     - Learned pressure-level embeddings at input (Pangu/Aurora style)
     - Column attention at input for explicit vertical coupling (ArchesWeather inspired)
     - Spectral GNN bottleneck for global mixing: pools the spatial nodes to K learned
       virtual spectral nodes, applies a channel MLP, and scatters corrections back
       (learned spectral graph convolution, O(N*K); not an FFT/spherical-harmonic SFNO)
     - Spectral normalization throughout



Attributes
----------

.. autoapisummary::

   credit.models.wxformer.wxformer_next.logger
   credit.models.wxformer.wxformer_next.device


Classes
-------

.. autoapisummary::

   credit.models.wxformer.wxformer_next.FeedForward
   credit.models.wxformer.wxformer_next.Attention
   credit.models.wxformer.wxformer_next.Transformer
   credit.models.wxformer.wxformer_next.LevelEmbedding
   credit.models.wxformer.wxformer_next.ColumnAttention
   credit.models.wxformer.wxformer_next.SpectralGNNBottleneck
   credit.models.wxformer.wxformer_next.NextGenWXFormer


Functions
---------

.. autoapisummary::

   credit.models.wxformer.wxformer_next.remap_conv_state_dict


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

.. py:data:: logger

.. py:class:: FeedForward(dim, mult=4, dropout=0.0, tp_plan=None)

   Bases: :py:obj:`torch.nn.Module`


   Pre-norm channel MLP with Linear projections (channels-last GEMM).

   Keeps the exact parameter FQNs of the conv version (``layers.0`` norm,
   ``layers.1`` up-projection, ``layers.4`` down-projection) so checkpoint
   remapping only reshapes weights, never renames keys.


   .. py:attribute:: layers


   .. py:method:: forward(x)


.. py:class:: Attention(dim, attn_type, window_size, dim_head=32, dropout=0.0, tp_plan=None)

   Bases: :py:obj:`torch.nn.Module`


   Short/long window attention with Linear q/k/v/out projections.

   Same math as crossformer's conv Attention, but the fused ``to_qkv`` 1x1
   conv is split into separate ``to_q``/``to_k``/``to_v`` Linears (the
   torchtitan pattern) so native ColwiseParallel sharding keeps whole heads
   on each rank. The forward derives the head count from the local channel
   width, so it is correct both serially and under tensor parallelism
   (where each rank holds heads // tp heads).

   :param dim: Input dimension.
   :type dim: int
   :param attn_type: Type of attention, either "short" or "long".
   :type attn_type: str
   :param window_size: Size of the attention window.
   :type window_size: int
   :param dim_head: Dimension of each attention head. Defaults to 32.
   :type dim_head: int, optional
   :param dropout: Dropout rate. Defaults to 0.0.
   :type dropout: float, optional


   .. py:attribute:: heads


   .. py:attribute:: dim_head
      :value: 32



   .. py:attribute:: scale
      :value: 0.1767766952966369



   .. py:attribute:: attn_type


   .. py:attribute:: window_size


   .. py:attribute:: norm


   .. py:attribute:: dropout


   .. py:attribute:: to_q


   .. py:attribute:: to_k


   .. py:attribute:: to_v


   .. py:attribute:: to_out


   .. py:attribute:: dpb


   .. py:method:: forward(x)

      Forward pass of the Attention module.

      :param x: Input tensor of shape (batch, dim, height, width).
      :type x: torch.Tensor

      :returns: Output tensor of the same shape as input.
      :rtype: torch.Tensor



.. py:class:: Transformer(dim, *, local_window_size, global_window_size, depth=4, dim_head=32, attn_dropout=0.0, ff_dropout=0.0)

   Bases: :py:obj:`torch.nn.Module`


   Base class for all neural network modules.

   Your models should also subclass this class.

   Modules can also contain other Modules, allowing them to be nested in
   a tree structure. You can assign the submodules as regular attributes::

       import torch.nn as nn
       import torch.nn.functional as F


       class Model(nn.Module):
           def __init__(self) -> None:
               super().__init__()
               self.conv1 = nn.Conv2d(1, 20, 5)
               self.conv2 = nn.Conv2d(20, 20, 5)

           def forward(self, x):
               x = F.relu(self.conv1(x))
               return F.relu(self.conv2(x))

   Submodules assigned in this way will be registered, and will also have their
   parameters converted when you call :meth:`to`, etc.

   .. note::
       As per the example above, an ``__init__()`` call to the parent class
       must be made before assignment on the child.

   :ivar training: Boolean represents whether this module is in training or
                   evaluation mode.
   :vartype training: bool


   .. py:attribute:: layers


   .. py:method:: forward(x)


.. py:function:: remap_conv_state_dict(state_dict: dict) -> dict

   Convert a pre-Linear-refactor NextGenWXFormer state dict to the new format.

   Old checkpoints store the transformer projections as 1x1 Conv2d weights:

     - ``*.to_qkv.weight``        (3*inner, dim, 1, 1)  fused q/k/v
     - ``*.to_out.weight``        (dim, inner, 1, 1)
     - ``*.layers.1.weight`` / ``*.layers.4.weight``  FFN up/down (O, I, 1, 1)

   plus the ``weight_orig``/``weight_u``/``weight_v`` variants when spectral
   norm was active. The Linear refactor keeps every other key unchanged.

   Remap rules:

     - fused ``to_qkv`` tensors are split into ``to_q``/``to_k``/``to_v``
       (``weight_v`` is copied to all three: the input-side power-iteration
       vector has the same shape for each split, and u/v re-converge within
       a few steps anyway)
     - 1x1 conv kernels at the converted positions are viewed as (O, I)

   Idempotent: new-format state dicts pass through unchanged.

   Note: with spectral norm active, the old model normalized the fused qkv
   matrix jointly while the refactored model normalizes q/k/v separately —
   the loaded weights are identical, but the regularization differs slightly.

   :param state_dict: model state dict (old conv format or new Linear format).

   :returns: A new state dict in the Linear-projection format.


.. py:class:: LevelEmbedding(channels: int, levels: int)

   Bases: :py:obj:`torch.nn.Module`


   Learned per-level bias broadcast-added to atmospheric input channels.


   .. py:attribute:: channels


   .. py:attribute:: levels


   .. py:attribute:: bias


   .. py:method:: forward(x_atmos: torch.Tensor) -> torch.Tensor


.. py:class:: ColumnAttention(channels: int, levels: int, num_heads: int = 4, spatial_stride: int = 1)

   Bases: :py:obj:`torch.nn.Module`


   Multi-head attention across pressure levels at each (H, W) location.

   Operates only on the atmospheric variable channels; surface and forcing
   channels pass through unchanged.

   spatial_stride > 1 pools the spatial grid before attending and upsamples
   the correction back — reduces memory from O(B·H·W·L²) to O(B·H·W·L²/s²).
   Use stride=8 for full 640×1280 grids to keep attention weights under ~100 MB.


   .. py:attribute:: channels


   .. py:attribute:: levels


   .. py:attribute:: spatial_stride
      :value: 1



   .. py:attribute:: norm


   .. py:attribute:: attn


   .. py:attribute:: proj


   .. py:method:: forward(x_atmos: torch.Tensor) -> torch.Tensor


.. py:class:: SpectralGNNBottleneck(dim: int, nlat: int, nlon: int, num_spectral_nodes: int = 64, mlp_ratio: float = 2.0)

   Bases: :py:obj:`torch.nn.Module`


   Grid-agnostic global bottleneck via a learned spectral graph convolution.

   Works on any node layout (equiangular, Gaussian, HEALPix, unstructured).
   No FFT, no SHT, no grid assumption.

   Strategy: pool all spatial nodes to K learned "virtual" spectral nodes via a
   soft attention aggregation, apply a channel MLP, then broadcast corrections
   back to every spatial node. Cost is O(N·K) not O(N²), and K << N.

   :param dim: Channel dimension.
   :param nlat: Bottleneck height (spatial nodes = nlat × nlon).
   :param nlon: Bottleneck width.
   :param num_spectral_nodes: Number of virtual spectral nodes K. Default 64.
   :param mlp_ratio: Hidden-dim multiplier inside the spectral MLP.


   .. py:attribute:: N


   .. py:attribute:: K
      :value: 64



   .. py:attribute:: norm


   .. py:attribute:: agg_w


   .. py:attribute:: scatter_w


   .. py:attribute:: mlp


   .. py:method:: forward(x: torch.Tensor) -> torch.Tensor


.. py:class:: NextGenWXFormer(image_height: int = 640, image_width: int = 1280, frames: int = 2, channels: int = 4, surface_channels: int = 7, input_only_channels: int = 3, output_only_channels: int = 0, levels: int = 15, dim: tuple = (64, 128, 256, 512), depth: tuple = (2, 2, 8, 2), dim_head: int = 32, global_window_size: tuple = (5, 5, 2, 1), local_window_size: int = 10, cross_embed_kernel_sizes: tuple = ((4, 8, 16, 32), (2, 4), (2, 4), (2, 4)), cross_embed_strides: tuple = (4, 2, 2, 2), col_attn_heads: int = 4, col_attn_stride: int = 1, decoder_col_attn: bool = False, num_spectral_nodes: int = 64, use_spectral_norm: bool = True, **kwargs)

   Bases: :py:obj:`credit.models.base_model.BaseModel`


   CrossFormer U-Net with spectral GNN bottleneck, column attention,
   and pressure-level embeddings.

   :param image_height: Grid cells in south-north direction.
   :param image_width: Grid cells in west-east direction.
   :param frames: Number of input time steps.
   :param channels: Number of 3D (pressure-level) variables.
   :param surface_channels: Number of surface (single-level) variables.
   :param input_only_channels: Forcing/static variables that are input-only.
   :param output_only_channels: Diagnostic variables that are output-only.
   :param levels: Number of vertical pressure levels.
   :param dim: Hidden dimension at each of the 4 encoder stages.
   :param depth: Transformer block depth at each encoder stage.
   :param dim_head: Attention head dimension.
   :param global_window_size: Long-range attention stride at each stage.
   :param local_window_size: Short-range window size (scalar, shared across stages).
   :param cross_embed_kernel_sizes: CrossEmbedLayer kernel sizes at each stage.
   :param cross_embed_strides: CrossEmbedLayer strides at each stage.
   :param col_attn_heads: Number of heads in column attention.
   :param col_attn_stride: Spatial pooling stride before column attention (1 = full resolution).
                           Use 8 for 640×1280 grids to stay under ~100 MB attention memory.
   :param decoder_col_attn: Apply column attention to atmospheric output channels before the
                            residual add. Only atmospheric channels are processed; surface channels pass through.
   :param num_spectral_nodes: Number of virtual spectral nodes K in the GNN bottleneck.
                              Higher K = more global context, more memory. Default 64.
   :param use_spectral_norm: Apply spectral normalization to conv/linear layers.


   .. py:attribute:: image_height
      :value: 640



   .. py:attribute:: image_width
      :value: 1280



   .. py:attribute:: frames
      :value: 2



   .. py:attribute:: channels
      :value: 4



   .. py:attribute:: surface_channels
      :value: 7



   .. py:attribute:: input_only_channels
      :value: 3



   .. py:attribute:: levels
      :value: 15



   .. py:attribute:: use_spectral_norm
      :value: True



   .. py:attribute:: decoder_col_attn
      :value: False



   .. py:attribute:: output_channels
      :value: 67



   .. py:attribute:: level_embedding


   .. py:attribute:: col_attn


   .. py:attribute:: layers


   .. py:attribute:: spectral_bottleneck


   .. py:attribute:: up_block1


   .. py:attribute:: up_block2


   .. py:attribute:: up_block3


   .. py:attribute:: up_block4


   .. py:method:: forward(x: torch.Tensor) -> torch.Tensor


   .. py:method:: load_model(conf)
      :classmethod:



   .. py:method:: load_model_name(conf, model_name)
      :classmethod:



.. py:data:: device

