credit.models.wxformer.wxformer_next#
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#
Classes#
Pre-norm channel MLP with Linear projections (channels-last GEMM). |
|
Short/long window attention with Linear q/k/v/out projections. |
|
Base class for all neural network modules. |
|
Learned per-level bias broadcast-added to atmospheric input channels. |
|
Multi-head attention across pressure levels at each (H, W) location. |
|
Grid-agnostic global bottleneck via a learned spectral graph convolution. |
|
CrossFormer U-Net with spectral GNN bottleneck, column attention, |
Functions#
|
Convert a pre-Linear-refactor NextGenWXFormer state dict to the new format. |
Module Contents#
- credit.models.wxformer.wxformer_next.logger#
- class credit.models.wxformer.wxformer_next.FeedForward(dim, mult=4, dropout=0.0, tp_plan=None)#
Bases:
torch.nn.ModulePre-norm channel MLP with Linear projections (channels-last GEMM).
Keeps the exact parameter FQNs of the conv version (
layers.0norm,layers.1up-projection,layers.4down-projection) so checkpoint remapping only reshapes weights, never renames keys.- layers#
- forward(x)#
- class credit.models.wxformer.wxformer_next.Attention(dim, attn_type, window_size, dim_head=32, dropout=0.0, tp_plan=None)#
Bases:
torch.nn.ModuleShort/long window attention with Linear q/k/v/out projections.
Same math as crossformer’s conv Attention, but the fused
to_qkv1x1 conv is split into separateto_q/to_k/to_vLinears (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).- Parameters:
dim (int) – Input dimension.
attn_type (str) – Type of attention, either “short” or “long”.
window_size (int) – Size of the attention window.
dim_head (int, optional) – Dimension of each attention head. Defaults to 32.
dropout (float, optional) – Dropout rate. Defaults to 0.0.
- heads#
- dim_head = 32#
- scale = 0.1767766952966369#
- attn_type#
- window_size#
- norm#
- dropout#
- to_q#
- to_k#
- to_v#
- to_out#
- dpb#
- forward(x)#
Forward pass of the Attention module.
- Parameters:
x (torch.Tensor) – Input tensor of shape (batch, dim, height, width).
- Returns:
Output tensor of the same shape as input.
- Return type:
torch.Tensor
- class credit.models.wxformer.wxformer_next.Transformer(dim, *, local_window_size, global_window_size, depth=4, dim_head=32, attn_dropout=0.0, ff_dropout=0.0)#
Bases:
torch.nn.ModuleBase 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
to(), etc.Note
As per the example above, an
__init__()call to the parent class must be made before assignment on the child.- Variables:
training (bool) – Boolean represents whether this module is in training or evaluation mode.
- layers#
- forward(x)#
- credit.models.wxformer.wxformer_next.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.weightFFN up/down (O, I, 1, 1)
plus the
weight_orig/weight_u/weight_vvariants when spectral norm was active. The Linear refactor keeps every other key unchanged.Remap rules:
fused
to_qkvtensors are split intoto_q/to_k/to_v(weight_vis 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.
- Parameters:
state_dict – model state dict (old conv format or new Linear format).
- Returns:
A new state dict in the Linear-projection format.
- class credit.models.wxformer.wxformer_next.LevelEmbedding(channels: int, levels: int)#
Bases:
torch.nn.ModuleLearned per-level bias broadcast-added to atmospheric input channels.
- channels#
- levels#
- bias#
- forward(x_atmos: torch.Tensor) torch.Tensor#
- class credit.models.wxformer.wxformer_next.ColumnAttention(channels: int, levels: int, num_heads: int = 4, spatial_stride: int = 1)#
Bases:
torch.nn.ModuleMulti-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.
- channels#
- levels#
- spatial_stride = 1#
- norm#
- attn#
- proj#
- forward(x_atmos: torch.Tensor) torch.Tensor#
- class credit.models.wxformer.wxformer_next.SpectralGNNBottleneck(dim: int, nlat: int, nlon: int, num_spectral_nodes: int = 64, mlp_ratio: float = 2.0)#
Bases:
torch.nn.ModuleGrid-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.
- Parameters:
dim – Channel dimension.
nlat – Bottleneck height (spatial nodes = nlat × nlon).
nlon – Bottleneck width.
num_spectral_nodes – Number of virtual spectral nodes K. Default 64.
mlp_ratio – Hidden-dim multiplier inside the spectral MLP.
- N#
- K = 64#
- norm#
- agg_w#
- scatter_w#
- mlp#
- forward(x: torch.Tensor) torch.Tensor#
- class credit.models.wxformer.wxformer_next.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:
credit.models.base_model.BaseModelCrossFormer U-Net with spectral GNN bottleneck, column attention, and pressure-level embeddings.
- Parameters:
image_height – Grid cells in south-north direction.
image_width – Grid cells in west-east direction.
frames – Number of input time steps.
channels – Number of 3D (pressure-level) variables.
surface_channels – Number of surface (single-level) variables.
input_only_channels – Forcing/static variables that are input-only.
output_only_channels – Diagnostic variables that are output-only.
levels – Number of vertical pressure levels.
dim – Hidden dimension at each of the 4 encoder stages.
depth – Transformer block depth at each encoder stage.
dim_head – Attention head dimension.
global_window_size – Long-range attention stride at each stage.
local_window_size – Short-range window size (scalar, shared across stages).
cross_embed_kernel_sizes – CrossEmbedLayer kernel sizes at each stage.
cross_embed_strides – CrossEmbedLayer strides at each stage.
col_attn_heads – Number of heads in column attention.
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.
decoder_col_attn – Apply column attention to atmospheric output channels before the residual add. Only atmospheric channels are processed; surface channels pass through.
num_spectral_nodes – Number of virtual spectral nodes K in the GNN bottleneck. Higher K = more global context, more memory. Default 64.
use_spectral_norm – Apply spectral normalization to conv/linear layers.
- image_height = 640#
- image_width = 1280#
- frames = 2#
- channels = 4#
- surface_channels = 7#
- input_only_channels = 3#
- levels = 15#
- use_spectral_norm = True#
- decoder_col_attn = False#
- output_channels = 67#
- level_embedding#
- col_attn#
- layers#
- spectral_bottleneck#
- up_block1#
- up_block2#
- up_block3#
- up_block4#
- forward(x: torch.Tensor) torch.Tensor#
- classmethod load_model(conf)#
- classmethod load_model_name(conf, model_name)#
- credit.models.wxformer.wxformer_next.device#