From f1be02af8c33d4136ad7c90dd43158a8ff6e5af1 Mon Sep 17 00:00:00 2001 From: gdamms Date: Thu, 5 Feb 2026 19:55:58 +0100 Subject: better plots, semi working autoencoder --- src/dataloader.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) (limited to 'src/dataloader.py') 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. -- cgit v1.3.1