Download scripts/data.py from OneScience-Group/FourCastNet_v2: direct link, hf CLI and curl.
- Browser
- Download file 4.2 kB
-
https://huggingface.co/OneScience-Group/FourCastNet_v2/resolve/main/scripts/data.py
- Command line
-
hf download hf://OneScience-Group/FourCastNet_v2/scripts/data.py
-
curl -L -o data.py https://huggingface.co/OneScience-Group/FourCastNet_v2/resolve/main/scripts/data.py
4.2 kB
| 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]) | |