aboutsummaryrefslogtreecommitdiff
path: root/src/dataloader.py
diff options
context:
space:
mode:
authorgdamms <damguillotin@gmail.com>2026-02-06 09:57:16 +0100
committergdamms <damguillotin@gmail.com>2026-02-06 09:57:16 +0100
commit8ee90daeedef75135ab50dd2bd5c5349929b53dd (patch)
tree7820c4ae0fc9493d7cc73a639ec092f242943ae7 /src/dataloader.py
parentdb4b7ac7f04c2100b266fe8546a25d72fe409e64 (diff)
downloaddiffusion-mnist-8ee90daeedef75135ab50dd2bd5c5349929b53dd.tar.gz
diffusion-mnist-8ee90daeedef75135ab50dd2bd5c5349929b53dd.zip
use AEModule for diffusion
Diffstat (limited to 'src/dataloader.py')
-rw-r--r--src/dataloader.py8
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.