aboutsummaryrefslogtreecommitdiff
path: root/src/dataloader.py
diff options
context:
space:
mode:
authorgdamms <damguillotin@gmail.com>2026-02-05 19:55:58 +0100
committergdamms <damguillotin@gmail.com>2026-02-05 19:55:58 +0100
commitf1be02af8c33d4136ad7c90dd43158a8ff6e5af1 (patch)
tree3d7facac2732046ce1447156afe7dacc3b64d585 /src/dataloader.py
parenta5d5f30fbd9c6c7c78834072401932c84bddaf14 (diff)
downloaddiffusion-mnist-f1be02af8c33d4136ad7c90dd43158a8ff6e5af1.tar.gz
diffusion-mnist-f1be02af8c33d4136ad7c90dd43158a8ff6e5af1.zip
better plots, semi working autoencoder
Diffstat (limited to 'src/dataloader.py')
-rw-r--r--src/dataloader.py7
1 files changed, 4 insertions, 3 deletions
diff --git a/src/dataloader.py b/src/dataloader.py
index a08ae27..d1e7fcb 100644
--- a/src/dataloader.py
+++ b/src/dataloader.py
@@ -38,7 +38,7 @@ class DiffusionDataset(Dataset):
autoencoder: Optional autoencoder for latent diffusion
"""
- def __init__(self, dataset: Dataset, autoencoder: torch.nn.Module = None):
+ def __init__(self, dataset: Dataset, autoencoder: torch.nn.Module | None = None):
super().__init__()
self.dataset = dataset
self.autoencoder = autoencoder
@@ -87,7 +87,7 @@ class DiffusionDatasetX0(Dataset):
autoencoder: Optional autoencoder for latent diffusion
"""
- def __init__(self, dataset: Dataset, autoencoder: torch.nn.Module = None):
+ def __init__(self, dataset: Dataset, autoencoder: torch.nn.Module | None = None):
super().__init__()
self.dataset = dataset
self.autoencoder = autoencoder
@@ -143,6 +143,7 @@ class AutoencoderDataset(Dataset):
def __getitem__(self, idx: int) -> tuple[torch.Tensor, torch.Tensor]:
data = self.dataset[idx][0].to(DEVICE)
+ # data = data * 2 - 1 # Normalize to [-1, 1]
return data, data
@@ -151,7 +152,7 @@ def get_diffusion_dataloader(
batch_size: int = 64,
shuffle: bool = True,
num_workers: int = 4,
- autoencoder: torch.nn.Module = None,
+ autoencoder: torch.nn.Module | None = None,
) -> DataLoader:
"""
Create a DataLoader for diffusion training.