aboutsummaryrefslogtreecommitdiff
path: root/src
diff options
context:
space:
mode:
authorgdamms <damguillotin@gmail.com>2026-02-06 09:51:13 +0100
committergdamms <damguillotin@gmail.com>2026-02-06 09:51:13 +0100
commitdb4b7ac7f04c2100b266fe8546a25d72fe409e64 (patch)
tree829fbabd63b97b77dfb5d2382a62a0b0c2cf7370 /src
parentf1be02af8c33d4136ad7c90dd43158a8ff6e5af1 (diff)
downloaddiffusion-mnist-db4b7ac7f04c2100b266fe8546a25d72fe409e64.tar.gz
diffusion-mnist-db4b7ac7f04c2100b266fe8546a25d72fe409e64.zip
train val test ae
Diffstat (limited to 'src')
-rw-r--r--src/dataloader.py68
-rw-r--r--src/train_autoencoder.py54
2 files changed, 113 insertions, 9 deletions
diff --git a/src/dataloader.py b/src/dataloader.py
index d1e7fcb..44cb6f5 100644
--- a/src/dataloader.py
+++ b/src/dataloader.py
@@ -3,7 +3,7 @@ Data loading utilities for MNIST diffusion training.
"""
import torch
-from torch.utils.data import Dataset, DataLoader
+from torch.utils.data import Dataset, DataLoader, random_split
from torchvision import datasets, transforms
from .config import DEVICE, DIFFU_STEPS, NB_LABEL, DATA_DIR
@@ -209,3 +209,69 @@ def get_autoencoder_dataloader(
num_workers=num_workers,
persistent_workers=True if num_workers > 0 else False,
)
+
+
+def get_autoencoder_dataloaders(
+ batch_size: int = 64,
+ shuffle: bool = True,
+ num_workers: int = 4,
+ val_split: float = 0.1,
+ test_split: float = 0.1,
+ seed: int = 42,
+) -> tuple[DataLoader, DataLoader, DataLoader]:
+ """
+ Create DataLoaders for autoencoder training, validation, and testing.
+
+ Args:
+ batch_size: Batch size
+ shuffle: Whether to shuffle training data
+ num_workers: Number of data loading workers
+ val_split: Fraction of data for validation
+ test_split: Fraction of data for testing
+ seed: Random seed for deterministic split
+
+ Returns:
+ Tuple of (train_loader, val_loader, test_loader)
+ """
+ if val_split < 0 or test_split < 0 or (val_split + test_split) >= 1:
+ raise ValueError("val_split and test_split must be >= 0 and sum to < 1")
+
+ mnist = get_mnist_dataset(train=True)
+ dataset = AutoencoderDataset(mnist)
+
+ total_len = len(dataset)
+ val_len = int(total_len * val_split)
+ test_len = int(total_len * test_split)
+ train_len = total_len - val_len - test_len
+
+ if train_len <= 0:
+ raise ValueError("Split sizes result in empty training set")
+
+ generator = torch.Generator().manual_seed(seed)
+ train_set, val_set, test_set = random_split(dataset, [train_len, val_len, test_len], generator=generator)
+
+ train_loader = DataLoader(
+ train_set,
+ batch_size=batch_size,
+ shuffle=shuffle,
+ num_workers=num_workers,
+ persistent_workers=True if num_workers > 0 else False,
+ )
+
+ val_loader = DataLoader(
+ val_set,
+ batch_size=batch_size,
+ shuffle=False,
+ num_workers=num_workers,
+ persistent_workers=True if num_workers > 0 else False,
+ )
+
+ test_loader = DataLoader(
+ test_set,
+ batch_size=batch_size,
+ shuffle=False,
+ num_workers=num_workers,
+ persistent_workers=True if num_workers > 0 else False,
+ )
+
+ return train_loader, val_loader, test_loader
diff --git a/src/train_autoencoder.py b/src/train_autoencoder.py
index 6af1fc5..04803c1 100644
--- a/src/train_autoencoder.py
+++ b/src/train_autoencoder.py
@@ -4,7 +4,7 @@ Training script for MNIST autoencoder.
from models import Autoencoder, AEModule
from src.utils import ensure_dirs, save_checkpoint
-from src.dataloader import get_autoencoder_dataloader
+from src.dataloader import get_autoencoder_dataloaders
from src.config import DEVICE, BATCH_SIZE, NUM_WORKERS, CHECKPOINT_DIR, PLOTS_DIR
import os
import torch
@@ -23,6 +23,8 @@ def train_autoencoder(
learning_rate: float = 1e-3,
batch_size: int = BATCH_SIZE,
latent_channels: int = 1,
+ val_split: float = 0.1,
+ test_split: float = 0.1,
checkpoint_path: str | None = None,
run_name: str | None = None,
):
@@ -34,6 +36,8 @@ def train_autoencoder(
learning_rate: Learning rate for optimizer
batch_size: Training batch size
latent_channels: Number of channels in latent space
+ val_split: Fraction of data for validation
+ test_split: Fraction of data for testing
checkpoint_path: Path to checkpoint to resume training from
run_name: Name for this training run (for logging)
"""
@@ -61,10 +65,12 @@ def train_autoencoder(
# criterion = nn.MSELoss()
criterion = nn.functional.binary_cross_entropy
- # Get dataloader
- dataloader = get_autoencoder_dataloader(
+ # Get dataloaders
+ train_loader, val_loader, test_loader = get_autoencoder_dataloaders(
batch_size=batch_size,
num_workers=NUM_WORKERS,
+ val_split=val_split,
+ test_split=test_split,
)
# Training loop
@@ -72,7 +78,7 @@ def train_autoencoder(
model.train()
epoch_loss = 0.0
- for batch_idx, (x, target) in enumerate(track(dataloader, description=f"Epoch {epoch}/{epochs}")):
+ for batch_idx, (x, target) in enumerate(track(train_loader, description=f"Epoch {epoch}/{epochs}")):
optimizer.zero_grad()
# Forward pass
@@ -86,22 +92,50 @@ def train_autoencoder(
epoch_loss += loss.item()
- avg_loss = epoch_loss / len(dataloader)
- mlflow.log_metric("epoch_loss", avg_loss, step=epoch)
+ avg_loss = epoch_loss / len(train_loader)
+ mlflow.log_metric("train_loss", avg_loss, step=epoch)
+
+ val_loss = evaluate_autoencoder(model, val_loader, criterion)
+ test_loss = evaluate_autoencoder(model, test_loader, criterion)
+
+ mlflow.log_metric("val_loss", val_loss, step=epoch)
+ mlflow.log_metric("test_loss", test_loss, step=epoch)
# Save checkpoint
save_checkpoint(model, f"autoencoder_epoch_{epoch:03d}.pt")
save_checkpoint(model, "autoencoder_latest.pt")
# Visualize results
- fig = visualize_reconstructions(model, dataloader)
+ train_fig = visualize_reconstructions(model, train_loader)
+ val_fig = visualize_reconstructions(model, val_loader)
+ test_fig = visualize_reconstructions(model, test_loader)
- mlflow.log_figure(fig, f"reconstructions/epoch_{epoch:03d}.png")
+ mlflow.log_figure(train_fig, f"train/epoch_{epoch:03d}.png")
+ mlflow.log_figure(val_fig, f"val/epoch_{epoch:03d}.png")
+ mlflow.log_figure(test_fig, f"test/epoch_{epoch:03d}.png")
mlflow.end_run()
return model
+def evaluate_autoencoder(
+ model: AEModule,
+ dataloader,
+ criterion,
+) -> float:
+ """Evaluate autoencoder and return average loss."""
+ model.eval()
+ total_loss = 0.0
+
+ with torch.no_grad():
+ for x, target in dataloader:
+ x_recon = model(x)
+ loss = criterion(x_recon, target)
+ total_loss += loss.item()
+
+ return total_loss / len(dataloader)
+
+
def visualize_reconstructions(model: AEModule, dataloader, n_samples: int = 10) -> go.Figure:
"""Visualize original, latent, and reconstructed images."""
model.eval()
@@ -175,6 +209,8 @@ if __name__ == "__main__":
parser.add_argument("--lr", type=float, default=1e-3, help="Learning rate")
parser.add_argument("--batch-size", type=int, default=BATCH_SIZE, help="Batch size")
parser.add_argument("--latent-channels", type=int, default=1, help="Latent channels")
+ parser.add_argument("--val-split", type=float, default=0.1, help="Validation split fraction")
+ parser.add_argument("--test-split", type=float, default=0.1, help="Test split fraction")
parser.add_argument("--checkpoint", type=str, default=None, help="Resume from checkpoint")
args = parser.parse_args()
@@ -186,5 +222,7 @@ if __name__ == "__main__":
learning_rate=args.lr,
batch_size=args.batch_size,
latent_channels=args.latent_channels,
+ val_split=args.val_split,
+ test_split=args.test_split,
checkpoint_path=args.checkpoint,
)