| from __future__ import annotations |
|
|
| import sys |
| from pathlib import Path |
| from typing import Any |
|
|
| import torch |
| import torch.nn.functional as functional |
| from torch.utils.data import DataLoader, Dataset |
|
|
|
|
| def _import_era5_dataset(onescience_source_dir: str | None = None): |
| if onescience_source_dir: |
| source_dir = str(Path(onescience_source_dir).expanduser().resolve()) |
| if source_dir not in sys.path: |
| sys.path.insert(0, source_dir) |
| from onescience.datapipes.climate import ERA5Dataset |
|
|
| return ERA5Dataset |
|
|
|
|
| class SpatialAdapter(Dataset): |
| """Resize OneScience ERA5 samples only for reduced smoke profiles.""" |
|
|
| def __init__(self, dataset: Dataset, output_size: tuple[int, int]) -> None: |
| self.dataset = dataset |
| self.output_size = output_size |
|
|
| def __len__(self) -> int: |
| return len(self.dataset) |
|
|
| def _resize(self, tensor: torch.Tensor) -> torch.Tensor: |
| if tuple(tensor.shape[-2:]) == self.output_size: |
| return tensor |
| leading_shape = tensor.shape[:-2] |
| resized = functional.interpolate( |
| tensor.reshape(-1, 1, *tensor.shape[-2:]), |
| size=self.output_size, |
| mode="bilinear", |
| align_corners=False, |
| ) |
| return resized.reshape(*leading_shape, *self.output_size) |
|
|
| def __getitem__(self, index: int): |
| inputs, targets, cos_zenith, step_idx, time_index = self.dataset[index] |
| inputs = self._resize(inputs) |
| targets = self._resize(targets) |
| cos_zenith = self._resize(cos_zenith) |
| return inputs, targets, cos_zenith, step_idx, time_index |
|
|
|
|
| def build_dataset( |
| config: dict[str, Any], |
| years: list[int], |
| *, |
| output_steps: int = 1, |
| ) -> Dataset: |
| from common import active_model_config, resolve_path |
|
|
| era5_dataset = _import_era5_dataset(config["project"].get("onescience_source_dir")) |
| data_config = config["data"] |
| dataset = era5_dataset( |
| dataset_dir=str(resolve_path(config, data_config["dataset_dir"])), |
| used_years=years, |
| used_variables=data_config["variables"], |
| input_steps=data_config["input_steps"], |
| output_steps=output_steps, |
| normalize=data_config["normalize"], |
| ) |
| model_size = tuple(active_model_config(config)["img_size"]) |
| data_size = tuple(data_config["grid_shape"]) |
| if model_size != data_size: |
| dataset = SpatialAdapter(dataset, model_size) |
| return dataset |
|
|
|
|
| def build_loader( |
| config: dict[str, Any], |
| years: list[int], |
| *, |
| train: bool, |
| distributed: bool, |
| output_steps: int = 1, |
| ) -> tuple[DataLoader, torch.utils.data.Sampler | None]: |
| dataset = build_dataset(config, years, output_steps=output_steps) |
| sampler = None |
| if distributed: |
| sampler = torch.utils.data.distributed.DistributedSampler( |
| dataset, shuffle=train |
| ) |
| loader = DataLoader( |
| dataset, |
| batch_size=config["training"]["batch_size"], |
| shuffle=train and sampler is None, |
| sampler=sampler, |
| num_workers=config["training"]["num_workers"], |
| pin_memory=True, |
| drop_last=False, |
| ) |
| return loader, sampler |
|
|
|
|
| def load_statistics(config: dict[str, Any]) -> tuple[torch.Tensor, torch.Tensor]: |
| import h5py |
| import numpy as np |
|
|
| from common import resolve_path |
|
|
| data_config = config["data"] |
| year = data_config["test_years"][0] |
| path = resolve_path(config, data_config["dataset_dir"]) / "data" / f"{year}.h5" |
| with h5py.File(path, "r") as handle: |
| fields = handle["fields"] |
| all_variables = [ |
| item.decode() if isinstance(item, bytes) else str(item) |
| for item in fields.attrs["variables"] |
| ] |
| indices = [all_variables.index(name) for name in data_config["variables"]] |
| if "global_means" in handle: |
| means = handle["global_means"][:] |
| stds = handle["global_stds"][:] |
| else: |
| stats_dir = path.parents[1] / "stats" |
| means = np.load(stats_dir / "global_means.npy") |
| stds = np.load(stats_dir / "global_stds.npy") |
| return torch.from_numpy(means[:, indices]), torch.from_numpy(stds[:, indices]) |
|
|