File size: 5,347 Bytes
8bfc737
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from __future__ import annotations

import random
import sys
from pathlib import Path

import numpy as np
import torch
from torch.utils.data import DataLoader, Dataset, DistributedSampler

sys.path.insert(0, str(Path(__file__).resolve().parents[1]))

from script.fake_data import SPLIT_OFFSETS, generate_sequence


class SyntheticSEVIRDataset(Dataset):
    def __init__(self, size: int, config: dict, seed: int):
        data = config["data"]
        self.size = size
        self.input_length = int(data["input_length"])
        self.output_length = int(data["output_length"])
        self.height = int(data["height"])
        self.width = int(data["width"])
        self.seed = seed

    def __len__(self) -> int:
        return self.size

    def __getitem__(self, index: int) -> tuple[torch.Tensor, torch.Tensor]:
        sequence = generate_sequence(
            self.height, self.width, self.input_length, self.output_length, self.seed + index
        )
        tensor = torch.from_numpy(sequence)
        return tensor[: self.input_length], tensor[self.input_length :]


class NPZSequenceDataset(Dataset):
    def __init__(self, path: str | Path, config: dict):
        data = config["data"]
        with np.load(path) as payload:
            if "inputs" not in payload or "targets" not in payload:
                raise ValueError("NPZ must contain 'inputs' and 'targets'")
            inputs, targets = payload["inputs"], payload["targets"]
        expected_input = (int(data["input_length"]), int(data["height"]), int(data["width"]), int(data["channels"]))
        expected_target = (int(data["output_length"]), int(data["height"]), int(data["width"]), int(data["channels"]))
        if inputs.ndim != 5 or tuple(inputs.shape[1:]) != expected_input:
            raise ValueError(f"inputs must have shape [N,{','.join(map(str, expected_input))}], got {inputs.shape}")
        if targets.ndim != 5 or tuple(targets.shape[1:]) != expected_target:
            raise ValueError(f"targets must have shape [N,{','.join(map(str, expected_target))}], got {targets.shape}")
        if len(inputs) != len(targets) or len(inputs) == 0:
            raise ValueError("inputs and targets must have the same non-zero sample count")
        normalization = data.get("normalization", "unit")
        if normalization == "uint8_255":
            if inputs.dtype != np.uint8 or targets.dtype != np.uint8:
                raise ValueError("uint8_255 normalization requires uint8 NPZ arrays")
            inputs, targets = inputs.astype(np.float32) / 255.0, targets.astype(np.float32) / 255.0
        else:
            if not np.issubdtype(inputs.dtype, np.floating) or not np.issubdtype(targets.dtype, np.floating):
                raise ValueError("unit normalization requires floating-point NPZ arrays")
            inputs, targets = inputs.astype(np.float32), targets.astype(np.float32)
            if not np.isfinite(inputs).all() or not np.isfinite(targets).all():
                raise ValueError("NPZ arrays contain non-finite values")
            if inputs.min() < 0 or inputs.max() > 1 or targets.min() < 0 or targets.max() > 1:
                raise ValueError("unit-normalized NPZ arrays must be within [0,1]; float values are never implicitly divided by 255")
        self.inputs = inputs
        self.targets = targets

    def __len__(self) -> int:
        return len(self.inputs)

    def __getitem__(self, index: int) -> tuple[torch.Tensor, torch.Tensor]:
        return torch.from_numpy(self.inputs[index]), torch.from_numpy(self.targets[index])


def _seed_worker(worker_id: int) -> None:
    del worker_id
    worker_seed = torch.initial_seed() % 2**32
    np.random.seed(worker_seed)
    random.seed(worker_seed)


def make_loader(
    config: dict,
    split: str,
    distributed: bool = False,
    rank: int = 0,
    world_size: int = 1,
    shuffle: bool | None = None,
) -> tuple[DataLoader, DistributedSampler | None]:
    if split not in SPLIT_OFFSETS:
        raise ValueError(f"unknown split: {split}")
    data, train = config["data"], config["train"]
    path_value = data.get(f"{split}_npz")
    path = Path(path_value) if path_value else None
    if path is not None and path.is_file():
        dataset: Dataset = NPZSequenceDataset(path, config)
    elif bool(data.get("fallback_if_missing", True)):
        dataset = SyntheticSEVIRDataset(
            int(data[f"{split}_samples"]), config, int(train["seed"]) + SPLIT_OFFSETS[split]
        )
    else:
        raise FileNotFoundError(f"configured {split} NPZ does not exist: {path}")
    should_shuffle = split == "train" if shuffle is None else shuffle
    sampler = None
    if distributed:
        sampler = DistributedSampler(
            dataset, num_replicas=world_size, rank=rank, shuffle=should_shuffle, seed=int(train["seed"]), drop_last=False
        )
    generator = torch.Generator().manual_seed(int(train["seed"]) + SPLIT_OFFSETS[split] + rank)
    options = config["dataloader"]
    loader = DataLoader(
        dataset,
        batch_size=int(train["batch_size"]),
        shuffle=should_shuffle and sampler is None,
        sampler=sampler,
        num_workers=int(options.get("num_workers", 0)),
        pin_memory=bool(options.get("pin_memory", False)),
        worker_init_fn=_seed_worker,
        generator=generator,
    )
    return loader, sampler