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