diff options
| author | gdamms <damguillotin@gmail.com> | 2026-02-06 09:57:16 +0100 |
|---|---|---|
| committer | gdamms <damguillotin@gmail.com> | 2026-02-06 09:57:16 +0100 |
| commit | 8ee90daeedef75135ab50dd2bd5c5349929b53dd (patch) | |
| tree | 7820c4ae0fc9493d7cc73a639ec092f242943ae7 /src/dataloader.py | |
| parent | db4b7ac7f04c2100b266fe8546a25d72fe409e64 (diff) | |
| download | diffusion-mnist-8ee90daeedef75135ab50dd2bd5c5349929b53dd.tar.gz diffusion-mnist-8ee90daeedef75135ab50dd2bd5c5349929b53dd.zip | |
use AEModule for diffusion
Diffstat (limited to 'src/dataloader.py')
| -rw-r--r-- | src/dataloader.py | 8 |
1 files changed, 5 insertions, 3 deletions
diff --git a/src/dataloader.py b/src/dataloader.py index 44cb6f5..64ce7f7 100644 --- a/src/dataloader.py +++ b/src/dataloader.py @@ -6,6 +6,8 @@ import torch from torch.utils.data import Dataset, DataLoader, random_split from torchvision import datasets, transforms +from models.autoencoder import AEModule + from .config import DEVICE, DIFFU_STEPS, NB_LABEL, DATA_DIR from .diffusion import q_xt_x0 @@ -38,7 +40,7 @@ class DiffusionDataset(Dataset): autoencoder: Optional autoencoder for latent diffusion """ - def __init__(self, dataset: Dataset, autoencoder: torch.nn.Module | None = None): + def __init__(self, dataset: Dataset, autoencoder: AEModule | None = None): super().__init__() self.dataset = dataset self.autoencoder = autoencoder @@ -87,7 +89,7 @@ class DiffusionDatasetX0(Dataset): autoencoder: Optional autoencoder for latent diffusion """ - def __init__(self, dataset: Dataset, autoencoder: torch.nn.Module | None = None): + def __init__(self, dataset: Dataset, autoencoder: AEModule | None = None): super().__init__() self.dataset = dataset self.autoencoder = autoencoder @@ -152,7 +154,7 @@ def get_diffusion_dataloader( batch_size: int = 64, shuffle: bool = True, num_workers: int = 4, - autoencoder: torch.nn.Module | None = None, + autoencoder: AEModule | None = None, ) -> DataLoader: """ Create a DataLoader for diffusion training. |
