File size: 4,199 Bytes
eca4864 | 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 | 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])
|