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