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])