diff options
| author | gdamms <damguillotin@gmail.com> | 2026-02-05 19:55:58 +0100 |
|---|---|---|
| committer | gdamms <damguillotin@gmail.com> | 2026-02-05 19:55:58 +0100 |
| commit | f1be02af8c33d4136ad7c90dd43158a8ff6e5af1 (patch) | |
| tree | 3d7facac2732046ce1447156afe7dacc3b64d585 /src/dataloader.py | |
| parent | a5d5f30fbd9c6c7c78834072401932c84bddaf14 (diff) | |
| download | diffusion-mnist-f1be02af8c33d4136ad7c90dd43158a8ff6e5af1.tar.gz diffusion-mnist-f1be02af8c33d4136ad7c90dd43158a8ff6e5af1.zip | |
better plots, semi working autoencoder
Diffstat (limited to 'src/dataloader.py')
| -rw-r--r-- | src/dataloader.py | 7 |
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. |
