diff options
| author | gdamms <damguillotin@gmail.com> | 2026-02-06 09:51:13 +0100 |
|---|---|---|
| committer | gdamms <damguillotin@gmail.com> | 2026-02-06 09:51:13 +0100 |
| commit | db4b7ac7f04c2100b266fe8546a25d72fe409e64 (patch) | |
| tree | 829fbabd63b97b77dfb5d2382a62a0b0c2cf7370 /src/dataloader.py | |
| parent | f1be02af8c33d4136ad7c90dd43158a8ff6e5af1 (diff) | |
| download | diffusion-mnist-db4b7ac7f04c2100b266fe8546a25d72fe409e64.tar.gz diffusion-mnist-db4b7ac7f04c2100b266fe8546a25d72fe409e64.zip | |
train val test ae
Diffstat (limited to 'src/dataloader.py')
| -rw-r--r-- | src/dataloader.py | 68 |
1 files changed, 67 insertions, 1 deletions
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 |
