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