File size: 7,106 Bytes
006ea64 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 | 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,
}
|