Earthformer / script /data_loader.py
yzt15806542928's picture
Upload folder using huggingface_hub
8bfc737 verified
Raw
History Blame
5.35 kB
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