diff options
| -rw-r--r-- | src/dataloader.py | 68 | ||||
| -rw-r--r-- | src/train_autoencoder.py | 54 |
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, ) |
