| """ACE variable order and channel contract.""" |
|
|
| from __future__ import annotations |
|
|
| import torch |
|
|
| PROGNOSTIC_3D = tuple(f"{name}_{level}" for name in ("T", "qT", "u", "v") for level in range(8)) |
| PROGNOSTIC_SURFACE = ("Ts_land_or_seaice", "ps") |
| PROGNOSTIC_CHANNELS = PROGNOSTIC_3D + PROGNOSTIC_SURFACE |
| FORCING_CHANNELS = ("DSW_RF_toa", "Ts_open_ocean", "zs", "fl", "fo", "fsi") |
| DIAGNOSTIC_CHANNELS = ( |
| "USW_RF_toa", "ULW_RF_toa", "USW_RF_sfc", "ULW_RF_sfc", |
| "DSW_RF_sfc", "DLW_RF_sfc", "P", "dTWP_adv_dt", "LHF", "SHF", |
| ) |
| PROGNOSTIC_UNITS = ("K", "kg/kg", "m/s", "m/s") * 8 + ("K", "Pa") |
| FORCING_UNITS = ("W/m2", "K", "m", "1", "1", "1") |
| DIAGNOSTIC_UNITS = ("W/m2",) * 6 + ("kg/m2/s", "kg/m2/s", "W/m2", "W/m2") |
| INPUT_UNITS = PROGNOSTIC_UNITS + FORCING_UNITS |
| OUTPUT_UNITS = PROGNOSTIC_UNITS + DIAGNOSTIC_UNITS |
| INPUT_CHANNELS = PROGNOSTIC_CHANNELS + FORCING_CHANNELS |
| OUTPUT_CHANNELS = PROGNOSTIC_CHANNELS + DIAGNOSTIC_CHANNELS |
|
|
|
|
| def ledger() -> dict: |
| return { |
| "input_channels": list(INPUT_CHANNELS), |
| "output_channels": list(OUTPUT_CHANNELS), |
| "input_units": list(INPUT_UNITS), |
| "output_units": list(OUTPUT_UNITS), |
| "prognostic_channels": len(PROGNOSTIC_CHANNELS), |
| "forcing_channels": len(FORCING_CHANNELS), |
| "diagnostic_channels": len(DIAGNOSTIC_CHANNELS), |
| "input_total": len(INPUT_CHANNELS), |
| "output_total": len(OUTPUT_CHANNELS), |
| } |
|
|
|
|
| def validate_channels(x: torch.Tensor, expected: int, name: str = "tensor") -> None: |
| if x.ndim != 4: |
| raise ValueError(f"{name} must have shape [B,C,H,W], got {tuple(x.shape)}") |
| if x.shape[1] != expected: |
| raise ValueError(f"{name} must have {expected} channels, got {x.shape[1]}") |
|
|
|
|
| def split_input(x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: |
| validate_channels(x, len(INPUT_CHANNELS), "input") |
| n = len(PROGNOSTIC_CHANNELS) |
| return x[:, :n], x[:, n:] |
|
|
|
|
| def split_output(y: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: |
| validate_channels(y, len(OUTPUT_CHANNELS), "output") |
| n = len(PROGNOSTIC_CHANNELS) |
| return y[:, :n], y[:, n:] |
|
|