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,
        }