From db4b7ac7f04c2100b266fe8546a25d72fe409e64 Mon Sep 17 00:00:00 2001 From: gdamms Date: Fri, 6 Feb 2026 09:51:13 +0100 Subject: train val test ae --- src/dataloader.py | 68 ++++++++++++++++++++++++++++++++++++++++++++++++++++++- 1 file changed, 67 insertions(+), 1 deletion(-) (limited to 'src/dataloader.py') diff --git a/src/dataloader.py b/src/dataloader.py index d1e7fcb..44cb6f5 100644 --- a/src/dataloader.py +++ b/src/dataloader.py @@ -3,7 +3,7 @@ Data loading utilities for MNIST diffusion training. """ import torch -from torch.utils.data import Dataset, DataLoader +from torch.utils.data import Dataset, DataLoader, random_split from torchvision import datasets, transforms from .config import DEVICE, DIFFU_STEPS, NB_LABEL, DATA_DIR @@ -209,3 +209,69 @@ def get_autoencoder_dataloader( num_workers=num_workers, persistent_workers=True if num_workers > 0 else False, ) + + +def get_autoencoder_dataloaders( + batch_size: int = 64, + shuffle: bool = True, + num_workers: int = 4, + val_split: float = 0.1, + test_split: float = 0.1, + seed: int = 42, +) -> tuple[DataLoader, DataLoader, DataLoader]: + """ + Create DataLoaders for autoencoder training, validation, and testing. + + Args: + batch_size: Batch size + shuffle: Whether to shuffle training data + num_workers: Number of data loading workers + val_split: Fraction of data for validation + test_split: Fraction of data for testing + seed: Random seed for deterministic split + + Returns: + Tuple of (train_loader, val_loader, test_loader) + """ + if val_split < 0 or test_split < 0 or (val_split + test_split) >= 1: + raise ValueError("val_split and test_split must be >= 0 and sum to < 1") + + mnist = get_mnist_dataset(train=True) + dataset = AutoencoderDataset(mnist) + + total_len = len(dataset) + val_len = int(total_len * val_split) + test_len = int(total_len * test_split) + train_len = total_len - val_len - test_len + + if train_len <= 0: + raise ValueError("Split sizes result in empty training set") + + generator = torch.Generator().manual_seed(seed) + train_set, val_set, test_set = random_split(dataset, [train_len, val_len, test_len], generator=generator) + + train_loader = DataLoader( + train_set, + batch_size=batch_size, + shuffle=shuffle, + num_workers=num_workers, + persistent_workers=True if num_workers > 0 else False, + ) + + val_loader = DataLoader( + val_set, + batch_size=batch_size, + shuffle=False, + num_workers=num_workers, + persistent_workers=True if num_workers > 0 else False, + ) + + test_loader = DataLoader( + test_set, + batch_size=batch_size, + shuffle=False, + num_workers=num_workers, + persistent_workers=True if num_workers > 0 else False, + ) + + return train_loader, val_loader, test_loader -- cgit v1.3.1