"""Array-backed ACE dataset and NPZ loader.""" from __future__ import annotations from pathlib import Path import numpy as np import torch from torch.utils.data import Dataset from ACE.model.variables import ( DIAGNOSTIC_CHANNELS, FORCING_CHANNELS, INPUT_CHANNELS, OUTPUT_CHANNELS, PROGNOSTIC_CHANNELS, INPUT_UNITS, OUTPUT_UNITS, ) def gaussian_latitudes(height: int) -> np.ndarray: """Return south-to-north Gaussian latitudes in radians.""" nodes, _ = np.polynomial.legendre.leggauss(height) return np.arcsin(nodes).astype(np.float32) class ArrayPairDataset(Dataset): def __init__(self, inputs: np.ndarray, targets: np.ndarray) -> None: self.inputs = np.asarray(inputs, dtype=np.float32) self.targets = np.asarray(targets, dtype=np.float32) if self.inputs.ndim != 4 or self.targets.ndim != 4: raise ValueError("inputs and targets must be [N,C,H,W]") if self.inputs.shape[0] != self.targets.shape[0]: raise ValueError("inputs and targets must have equal sample counts") if self.inputs.shape[1] != len(INPUT_CHANNELS): raise ValueError(f"inputs require {len(INPUT_CHANNELS)} channels") if self.targets.shape[1] != len(OUTPUT_CHANNELS): raise ValueError(f"targets require {len(OUTPUT_CHANNELS)} channels") if self.inputs.shape[-2:] != self.targets.shape[-2:]: raise ValueError("inputs and targets must share spatial shape") def __len__(self) -> int: return self.inputs.shape[0] def __getitem__(self, index: int) -> tuple[torch.Tensor, torch.Tensor]: return torch.from_numpy(self.inputs[index]), torch.from_numpy(self.targets[index]) def load_npz(path: str | Path) -> ArrayPairDataset: data = np.load(path) if "inputs" not in data or "targets" not in data: raise KeyError("NPZ must contain inputs and targets arrays") return ArrayPairDataset(data["inputs"], data["targets"]) def make_fake_pairs( num_samples: int = 8, height: int = 180, width: int = 360, seed: int = 0, ) -> tuple[np.ndarray, np.ndarray]: """Create ACE-shaped smooth global fields with the paper's variable semantics. The official training samples are 6-hour pairs on a 1-degree Gaussian grid: ``inputs[N,40,180,360]`` and ``targets[N,44,180,360]``. This generator keeps that contract while using analytic spherical modes instead of white noise. ``height`` and ``width`` remain configurable for small CPU smoke tests. """ if min(num_samples, height, width) <= 0: raise ValueError("num_samples, height and width must be positive") rng = np.random.default_rng(seed) lat = gaussian_latitudes(height) lon = np.arange(width, dtype=np.float32) * (2 * np.pi / width) latitude, longitude = np.meshgrid(lat, lon, indexing="ij") cos_lat = np.cos(latitude) sin_lat = np.sin(latitude) topography = 1200.0 * ( 0.55 * np.sin(2.0 * latitude) ** 2 * np.cos(2.0 * longitude) + 0.25 * np.cos(3.0 * latitude + longitude) ) land_fraction = np.clip( 0.5 + 0.35 * np.sin(2.0 * latitude) * np.cos(longitude) + 0.15 * np.cos(3.0 * longitude), 0.0, 1.0, ) sea_ice_fraction = np.clip((np.abs(latitude) - np.deg2rad(55.0)) / np.deg2rad(30.0), 0.0, 0.9) ocean_fraction = np.clip(1.0 - land_fraction - sea_ice_fraction, 0.0, 1.0) inputs = np.empty((num_samples, len(INPUT_CHANNELS), height, width), dtype=np.float32) targets = np.empty((num_samples, len(OUTPUT_CHANNELS), height, width), dtype=np.float32) for sample in range(num_samples): phase = float(rng.uniform(0.0, 2.0 * np.pi)) seasonal = float(rng.normal(0.0, 0.08)) wave = ( cos_lat * np.cos(longitude + phase) + 0.35 * np.sin(2.0 * latitude - 0.13 * sample) * np.sin(2.0 * longitude - phase) + 0.15 * np.cos(3.0 * latitude + 0.07 * sample) ).astype(np.float32) wave_next = ( cos_lat * np.cos(longitude + phase + 0.08) + 0.35 * np.sin(2.0 * latitude - 0.13 * sample - 0.03) * np.sin(2.0 * longitude - phase + 0.05) + 0.15 * np.cos(3.0 * latitude + 0.07 * sample + 0.02) ).astype(np.float32) prognostic = np.empty((len(PROGNOSTIC_CHANNELS), height, width), dtype=np.float32) prognostic_next = np.empty_like(prognostic) for level in range(8): height_factor = 1.0 - 0.06 * level humidity_factor = np.exp(-0.28 * level) offset = 4 * level prognostic[offset] = 215.0 + 48.0 * height_factor + 7.0 * wave prognostic[offset + 1] = np.maximum(1.0e-6, 0.00015 + 0.012 * humidity_factor * (1.0 + 0.25 * wave)) prognostic[offset + 2] = 18.0 * np.sin(2.0 * latitude) * np.cos(longitude + phase) * height_factor prognostic[offset + 3] = 12.0 * np.sin(longitude - phase) * cos_lat * height_factor prognostic_next[offset] = prognostic[offset] + 0.35 * (wave_next - wave) prognostic_next[offset + 1] = np.maximum(1.0e-6, prognostic[offset + 1] * (1.0 + 0.015 * (wave_next - wave))) prognostic_next[offset + 2] = prognostic[offset + 2] + 0.08 * (wave_next - wave) prognostic_next[offset + 3] = prognostic[offset + 3] - 0.05 * (wave_next - wave) prognostic[32] = 276.0 + 11.0 * wave + 2.0 * land_fraction prognostic[33] = 101325.0 + 1800.0 * wave - 0.18 * topography prognostic_next[32] = prognostic[32] + 0.4 * (wave_next - wave) prognostic_next[33] = prognostic[33] + 30.0 * (wave_next - wave) forcing = np.stack( [ 340.0 * np.maximum(cos_lat, 0.0) * (1.0 + seasonal), 288.0 + 7.0 * wave + 0.5 * seasonal, topography, land_fraction, ocean_fraction, sea_ice_fraction, ] ).astype(np.float32) diagnostics = np.stack( [ 70.0 * np.maximum(cos_lat, 0.0) * (1.0 + 0.05 * wave_next), 235.0 + 10.0 * wave_next, 35.0 * np.maximum(cos_lat, 0.0) * (1.0 - land_fraction), 310.0 + 8.0 * wave_next, 250.0 * np.maximum(cos_lat, 0.0), 280.0 + 5.0 * wave_next, np.maximum(0.0, 1.5e-5 * (1.0 + wave_next) * (1.0 - 0.4 * sea_ice_fraction)), 2.0e-6 * (wave_next - wave), 70.0 * (1.0 - sea_ice_fraction) * (1.0 + 0.1 * wave_next), 12.0 * land_fraction * (1.0 + 0.1 * wave_next), ] ).astype(np.float32) inputs[sample] = np.concatenate([prognostic, forcing], axis=0) targets[sample] = np.concatenate([prognostic_next, diagnostics], axis=0) return inputs, targets def save_fake_pairs(path: str | Path, **kwargs: int) -> Path: """Save :func:`make_fake_pairs` with grid and variable metadata.""" output = Path(path) output.parent.mkdir(parents=True, exist_ok=True) inputs, targets = make_fake_pairs(**kwargs) height, width = inputs.shape[-2:] np.savez( output, inputs=inputs, targets=targets, lat=np.rad2deg(gaussian_latitudes(height)), lon=np.linspace(0.0, 360.0, width, endpoint=False, dtype=np.float32), input_channels=np.asarray(INPUT_CHANNELS), output_channels=np.asarray(OUTPUT_CHANNELS), input_units=np.asarray(INPUT_UNITS), output_units=np.asarray(OUTPUT_UNITS), prognostic_channels=np.asarray(PROGNOSTIC_CHANNELS), forcing_channels=np.asarray(FORCING_CHANNELS), diagnostic_channels=np.asarray(DIAGNOSTIC_CHANNELS), time_step_hours=np.asarray(6, dtype=np.int32), source=np.asarray("synthetic FV3GFS-compatible analytic fixture"), ) return output