Download scripts/data_loader.py from OneScience-Group/StormCast: direct link, hf CLI and curl.
- Browser
- Download file 7.11 kB
-
https://huggingface.co/OneScience-Group/StormCast/resolve/main/scripts/data_loader.py
- Command line
-
hf download hf://OneScience-Group/StormCast/scripts/data_loader.py
-
curl -L -o data_loader.py https://huggingface.co/OneScience-Group/StormCast/resolve/main/scripts/data_loader.py
7.11 kB
| from __future__ import annotations | |
| from pathlib import Path | |
| from typing import Any | |
| import h5py | |
| import torch | |
| from onescience.datapipes.climate.era5 import ERA5Dataset | |
| from torch.utils.data import Dataset | |
| from grid import lambert_grid | |
| class StormCastDataset(Dataset): | |
| """Pair OneScience ERA5 backgrounds with synchronized local state targets.""" | |
| def __init__( | |
| self, | |
| data_root: str | Path, | |
| years: list[int], | |
| era5_variables: list[str], | |
| state_variables: list[str], | |
| invariant_variables: list[str], | |
| image_size: list[int] | tuple[int, int], | |
| input_steps: int = 1, | |
| output_steps: int = 1, | |
| normalize: bool = True, | |
| ) -> None: | |
| if input_steps != 1 or output_steps != 1: | |
| raise ValueError("StormCast pairing currently requires one input and one target step") | |
| self.data_root = Path(data_root) | |
| self.years = years | |
| self.era5_variables = era5_variables | |
| self.state_variables = state_variables | |
| self.invariant_variables = invariant_variables | |
| self.image_size = tuple(image_size) | |
| self.normalize = normalize | |
| self.era5 = ERA5Dataset( | |
| dataset_dir=str(self.data_root / "era5"), | |
| used_years=years, | |
| used_variables=era5_variables, | |
| input_steps=input_steps, | |
| output_steps=output_steps, | |
| normalize=normalize, | |
| ) | |
| self.samples_per_year = self.era5.samples_per_year | |
| self._validate_era5_grid() | |
| self._validate_local_files() | |
| self.invariants = self._load_invariants() | |
| self._initialize_background_regrid() | |
| def _validate_era5_grid(self) -> None: | |
| if self.era5.H < 2 or self.era5.W < 2: | |
| raise ValueError("ERA5 grid must have at least two points per dimension") | |
| expected = (721, 1440) | |
| if (self.era5.H, self.era5.W) != expected: | |
| raise ValueError( | |
| f"StormCast expects ERA5 on the global {expected} grid, " | |
| f"got {(self.era5.H, self.era5.W)}" | |
| ) | |
| def _validate_local_files(self) -> None: | |
| for year in self.years: | |
| path = self.data_root / "hrrr" / "data" / f"{year}.h5" | |
| if not path.is_file(): | |
| raise FileNotFoundError(f"Missing local state file: {path}") | |
| with h5py.File(path, "r") as handle: | |
| fields = handle["fields"] | |
| variables = [ | |
| value.decode() if isinstance(value, bytes) else str(value) | |
| for value in fields.attrs["variables"] | |
| ] | |
| if variables != self.state_variables: | |
| raise ValueError( | |
| "Local state channel order differs from data.state_variables" | |
| ) | |
| expected_steps = self.samples_per_year + 1 | |
| if fields.shape[0] != expected_steps: | |
| raise ValueError( | |
| f"{path} has {fields.shape[0]} steps, expected {expected_steps}" | |
| ) | |
| if tuple(fields.shape[-2:]) != self.image_size: | |
| raise ValueError( | |
| f"Local state grid is {tuple(fields.shape[-2:])}, " | |
| f"expected regional grid {self.image_size}" | |
| ) | |
| def _load_invariants(self) -> torch.Tensor: | |
| path = self.data_root / "hrrr" / "invariants.h5" | |
| with h5py.File(path, "r") as handle: | |
| fields = handle["fields"] | |
| variables = [ | |
| value.decode() if isinstance(value, bytes) else str(value) | |
| for value in fields.attrs["variables"] | |
| ] | |
| if variables != self.invariant_variables: | |
| raise ValueError( | |
| "Invariant channel order differs from data.invariant_variables" | |
| ) | |
| invariants = torch.as_tensor(fields[:], dtype=torch.float32) | |
| if tuple(invariants.shape[-2:]) != self.image_size: | |
| raise ValueError( | |
| f"Invariant grid is {tuple(invariants.shape[-2:])}, " | |
| f"expected {self.image_size}" | |
| ) | |
| return invariants | |
| def _initialize_background_regrid(self) -> None: | |
| with h5py.File(self.data_root / "hrrr" / "invariants.h5", "r") as handle: | |
| if "lat" in handle and "lon" in handle: | |
| target_lat = torch.as_tensor(handle["lat"][:], dtype=torch.float32) | |
| target_lon = torch.as_tensor(handle["lon"][:], dtype=torch.float32) | |
| else: | |
| target_lat_np, target_lon_np = lambert_grid(self.image_size) | |
| target_lat = torch.from_numpy(target_lat_np) | |
| target_lon = torch.from_numpy(target_lon_np) | |
| if target_lat.shape != self.image_size or target_lon.shape != self.image_size: | |
| raise ValueError("StormCast target latitude/longitude grid has wrong shape") | |
| lat_position = (90.0 - target_lat) / (180.0 / (self.era5.H - 1)) | |
| lon_position = torch.remainder(target_lon, 360.0) / (360.0 / self.era5.W) | |
| self.lat0 = lat_position.floor().long().clamp(0, self.era5.H - 2) | |
| self.lat1 = self.lat0 + 1 | |
| self.lon0 = lon_position.floor().long().remainder(self.era5.W) | |
| self.lon1 = (self.lon0 + 1).remainder(self.era5.W) | |
| self.lat_weight = lat_position - self.lat0 | |
| self.lon_weight = lon_position - lon_position.floor() | |
| def _regrid_background(self, background: torch.Tensor) -> torch.Tensor: | |
| f00 = background[..., self.lat0, self.lon0] | |
| f01 = background[..., self.lat0, self.lon1] | |
| f10 = background[..., self.lat1, self.lon0] | |
| f11 = background[..., self.lat1, self.lon1] | |
| lon_weight = self.lon_weight.to(background.dtype) | |
| lat_weight = self.lat_weight.to(background.dtype) | |
| top = torch.lerp(f00, f01, lon_weight) | |
| bottom = torch.lerp(f10, f11, lon_weight) | |
| return torch.lerp(top, bottom, lat_weight) | |
| def __len__(self) -> int: | |
| return len(self.era5) | |
| def __getitem__(self, index: int) -> dict[str, Any]: | |
| background, _, _, step_index, time_index = self.era5[index] | |
| background = self._regrid_background(background) | |
| year_index = index // self.samples_per_year | |
| year = self.years[year_index] | |
| path = self.data_root / "hrrr" / "data" / f"{year}.h5" | |
| with h5py.File(path, "r") as handle: | |
| state = torch.as_tensor( | |
| handle["fields"][step_index : step_index + 2], dtype=torch.float32 | |
| ) | |
| if self.normalize: | |
| means = torch.as_tensor(handle["global_means"][:], dtype=torch.float32) | |
| stds = torch.as_tensor(handle["global_stds"][:], dtype=torch.float32) | |
| state = (state - means) / stds | |
| return { | |
| "background": background, | |
| "state": (state[0], state[1]), | |
| "invariant": self.invariants, | |
| "step_index": step_index, | |
| "time_index": time_index, | |
| } | |