"""ERA5-style HDF5 loader used by training and inference.""" from __future__ import annotations from datetime import datetime, timedelta from pathlib import Path from typing import Dict, List, Sequence, Tuple import h5py import numpy as np import torch from torch.utils.data import DataLoader, Dataset from torch.utils.data.distributed import DistributedSampler def _decode(value: object) -> str: return value.decode() if isinstance(value, bytes) else str(value) def resolve_data_dir(path: str | Path, project_root: Path | None = None) -> Path: root = project_root or Path(__file__).resolve().parents[1] candidate = Path(path).expanduser() return candidate if candidate.is_absolute() else (root / candidate).resolve() def read_metadata(data_dir: str | Path, channels: Sequence[str]) -> Dict[str, np.ndarray | int | List[str]]: data_dir = Path(data_dir) files = sorted((data_dir / "data").glob("*.h5")) if not files: raise FileNotFoundError(f"No yearly HDF5 files found under {data_dir / 'data'}") with h5py.File(files[0], "r") as source: fields = source["fields"] variables = [_decode(item) for item in fields.attrs["variables"]] time_step = int(fields.attrs.get("time_step", 6)) means = np.asarray(source["global_means"][:], dtype=np.float32) stds = np.asarray(source["global_stds"][:], dtype=np.float32) shape = tuple(int(item) for item in fields.shape) missing = [name for name in channels if name not in variables] if missing: raise ValueError(f"Variables missing from synthetic/ERA5 data: {missing}") indices = np.asarray([variables.index(name) for name in channels], dtype=np.int64) return { "variables": variables, "indices": indices, "time_step": time_step, "means": means[:, indices, :, :], "stds": np.maximum(stds[:, indices, :, :], 1.0e-6), "shape": shape, } class ERA5WindowDataset(Dataset): def __init__( self, data_dir: str | Path, years: Sequence[int], channels: Sequence[str], input_steps: int = 2, rollout_steps: int = 1, normalize: bool = True, ) -> None: self.data_dir = Path(data_dir) self.years = [int(year) for year in years] self.channels = list(channels) self.input_steps = max(1, int(input_steps)) self.rollout_steps = max(1, int(rollout_steps)) self.normalize = bool(normalize) metadata = read_metadata(self.data_dir, self.channels) self.channel_indices = metadata["indices"] self.time_step = int(metadata["time_step"]) self.means = torch.from_numpy(metadata["means"]) self.stds = torch.from_numpy(metadata["stds"]) self.shape = metadata["shape"] self.samples: List[Tuple[int, int]] = [] for year in self.years: path = self.data_dir / "data" / f"{year}.h5" if not path.exists(): raise FileNotFoundError(f"Missing year file: {path}") with h5py.File(path, "r") as source: timesteps = int(source["fields"].shape[0]) count = timesteps - self.input_steps - self.rollout_steps + 1 if count <= 0: raise ValueError( f"Year {year} has {timesteps} steps, but input={self.input_steps} " f"and rollout={self.rollout_steps} require at least {self.input_steps + self.rollout_steps}" ) self.samples.extend((year, index) for index in range(count)) def __len__(self) -> int: return len(self.samples) def _timestamp(self, year: int, index: int) -> str: value = datetime(year, 1, 1) + timedelta(hours=index * self.time_step) return value.strftime("%Y%m%d%H") def __getitem__(self, item: int): year, start = self.samples[item] path = self.data_dir / "data" / f"{year}.h5" with h5py.File(path, "r") as source: fields = source["fields"] input_data = np.asarray( fields[start : start + self.input_steps, self.channel_indices, :, :], dtype=np.float32 ) target_data = np.asarray( fields[ start + self.input_steps : start + self.input_steps + self.rollout_steps, self.channel_indices, :, :, ], dtype=np.float32, ) input_tensor = torch.from_numpy(input_data) target_tensor = torch.from_numpy(target_data) if self.normalize: input_tensor = (input_tensor - self.means) / self.stds target_tensor = (target_tensor - self.means) / self.stds timestamp = self._timestamp(year, start + self.input_steps) return input_tensor, target_tensor, timestamp def make_dataloader( data_dir: str | Path, years: Sequence[int], channels: Sequence[str], input_steps: int, rollout_steps: int, batch_size: int, num_workers: int = 0, distributed: bool = False, train: bool = False, pin_memory: bool = False, ): dataset = ERA5WindowDataset( data_dir=data_dir, years=years, channels=channels, input_steps=input_steps, rollout_steps=rollout_steps, ) sampler = DistributedSampler(dataset, shuffle=train) if distributed else None loader = DataLoader( dataset, batch_size=max(1, int(batch_size)), shuffle=train and sampler is None, sampler=sampler, num_workers=max(0, int(num_workers)), pin_memory=bool(pin_memory), drop_last=False, ) return loader, sampler