aboutsummaryrefslogtreecommitdiff
path: root/src/dataloader.py
diff options
context:
space:
mode:
authorgdamms <damguillotin@gmail.com>2026-02-06 11:07:37 +0100
committergdamms <damguillotin@gmail.com>2026-02-06 11:07:37 +0100
commitf1fee867bbbffdd2686cc2f92313496f77d08b2f (patch)
treee4a2cd033c8d9fdc372c6fbe033436db68183a2d /src/dataloader.py
parent8ee90daeedef75135ab50dd2bd5c5349929b53dd (diff)
downloaddiffusion-mnist-f1fee867bbbffdd2686cc2f92313496f77d08b2f.tar.gz
diffusion-mnist-f1fee867bbbffdd2686cc2f92313496f77d08b2f.zip
add patience for early stopping or infinit epochs
Diffstat (limited to 'src/dataloader.py')
-rw-r--r--src/dataloader.py74
1 files changed, 74 insertions, 0 deletions
diff --git a/src/dataloader.py b/src/dataloader.py
index 64ce7f7..ebe8d3a 100644
--- a/src/dataloader.py
+++ b/src/dataloader.py
@@ -185,6 +185,80 @@ def get_diffusion_dataloader(
)
+def get_diffusion_dataloaders(
+ predict_x0: bool = True,
+ batch_size: int = 64,
+ shuffle: bool = True,
+ num_workers: int = 4,
+ autoencoder: AEModule | None = None,
+ val_split: float = 0.1,
+ test_split: float = 0.1,
+ seed: int = 42,
+) -> tuple[DataLoader, DataLoader, DataLoader]:
+ """
+ Create DataLoaders for diffusion training, validation, and testing.
+
+ Args:
+ predict_x0: If True, model predicts x0. Otherwise predicts noise.
+ batch_size: Batch size
+ shuffle: Whether to shuffle training data
+ num_workers: Number of data loading workers
+ autoencoder: Optional autoencoder for latent diffusion
+ 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)
+
+ if predict_x0:
+ dataset = DiffusionDatasetX0(mnist, autoencoder)
+ else:
+ dataset = DiffusionDataset(mnist, autoencoder)
+
+ 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
+
+
def get_autoencoder_dataloader(
batch_size: int = 64,
shuffle: bool = True,