yzt15806542928's picture
Upload folder using huggingface_hub
5c365c5 verified
Raw
History Blame Contribute Delete
7.89 kB
"""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