From f1fee867bbbffdd2686cc2f92313496f77d08b2f Mon Sep 17 00:00:00 2001 From: gdamms Date: Fri, 6 Feb 2026 11:07:37 +0100 Subject: add patience for early stopping or infinit epochs --- src/dataloader.py | 74 +++++++++++++++++++++++++++ src/train_autoencoder.py | 57 +++++++++++++++++---- src/train_diffusion.py | 130 +++++++++++++++++++++++++++++++++++++---------- 3 files changed, 223 insertions(+), 38 deletions(-) (limited to 'src') diff --git a/src/dataloader.py b/src/dataloader.py index 64ce7f7..ebe8d3a 100644 --- a/src/dataloader.py +++ b/src/dataloader.py @@ -185,6 +185,80 @@ def get_diffusion_dataloader( ) +def get_diffusion_dataloaders( + predict_x0: bool = True, + batch_size: int = 64, + shuffle: bool = True, + num_workers: int = 4, + autoencoder: AEModule | None = None, + val_split: float = 0.1, + test_split: float = 0.1, + seed: int = 42, +) -> tuple[DataLoader, DataLoader, DataLoader]: + """ + Create DataLoaders for diffusion training, validation, and testing. + + Args: + predict_x0: If True, model predicts x0. Otherwise predicts noise. + batch_size: Batch size + shuffle: Whether to shuffle training data + num_workers: Number of data loading workers + autoencoder: Optional autoencoder for latent diffusion + 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) + + if predict_x0: + dataset = DiffusionDatasetX0(mnist, autoencoder) + else: + dataset = DiffusionDataset(mnist, autoencoder) + + 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 + + def get_autoencoder_dataloader( batch_size: int = 64, shuffle: bool = True, diff --git a/src/train_autoencoder.py b/src/train_autoencoder.py index 04803c1..7dcc193 100644 --- a/src/train_autoencoder.py +++ b/src/train_autoencoder.py @@ -5,7 +5,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_dataloaders -from src.config import DEVICE, BATCH_SIZE, NUM_WORKERS, CHECKPOINT_DIR, PLOTS_DIR +from src.config import DEVICE, BATCH_SIZE, NUM_WORKERS import os import torch import torch.nn as nn @@ -19,12 +19,13 @@ sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) def train_autoencoder( - epochs: int = 10, + epochs: int | None = None, learning_rate: float = 1e-3, batch_size: int = BATCH_SIZE, latent_channels: int = 1, val_split: float = 0.1, test_split: float = 0.1, + patience: int | None = 5, checkpoint_path: str | None = None, run_name: str | None = None, ): @@ -38,9 +39,15 @@ def train_autoencoder( latent_channels: Number of channels in latent space val_split: Fraction of data for validation test_split: Fraction of data for testing + patience: Number of epochs to wait for validation loss improvement before stopping checkpoint_path: Path to checkpoint to resume training from run_name: Name for this training run (for logging) """ + if epochs is None and patience is None: + raise ValueError("Must specify either epochs or patience for training") + if patience is not None and patience <= 0: + raise ValueError("Patience must be a positive integer") + ensure_dirs() # Initialize model @@ -61,8 +68,6 @@ def train_autoencoder( # Setup training optimizer = torch.optim.Adam(model.parameters(), lr=learning_rate) - # criterion = nn.BCELoss() - # criterion = nn.MSELoss() criterion = nn.functional.binary_cross_entropy # Get dataloaders @@ -73,12 +78,26 @@ def train_autoencoder( test_split=test_split, ) + # Early stopping tracking + best_val_loss = float('inf') + epochs_without_improvement = 0 + # Training loop - for epoch in range(1, epochs + 1): + epoch = 0 + while epochs is None or epoch < epochs: + epoch += 1 + model.train() epoch_loss = 0.0 - for batch_idx, (x, target) in enumerate(track(train_loader, description=f"Epoch {epoch}/{epochs}")): + if epochs: + description = f"Epoch {epoch}/{epochs}" + else: + description = f"Epoch {epoch}" + if patience: + description += f" (Patience: {patience-epochs_without_improvement})" + + for batch_idx, (x, target) in enumerate(track(train_loader, description=description)): optimizer.zero_grad() # Forward pass @@ -93,18 +112,32 @@ def train_autoencoder( epoch_loss += loss.item() avg_loss = epoch_loss / len(train_loader) - mlflow.log_metric("train_loss", avg_loss, step=epoch) + 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) + mlflow.log_metric("val/loss", val_loss, step=epoch) + mlflow.log_metric("test/loss", test_loss, step=epoch) + + # Early stopping check + if val_loss < best_val_loss: + best_val_loss = val_loss + epochs_without_improvement = 0 + # Save best checkpoint + save_checkpoint(model, "autoencoder_best.pt") + else: + epochs_without_improvement += 1 - # Save checkpoint + # Save regular checkpoints save_checkpoint(model, f"autoencoder_epoch_{epoch:03d}.pt") save_checkpoint(model, "autoencoder_latest.pt") + # Stop if no improvement + if patience and epochs_without_improvement >= patience: + print(f"Early stopping: No improvement for {patience} epochs") + break + # Visualize results train_fig = visualize_reconstructions(model, train_loader) val_fig = visualize_reconstructions(model, val_loader) @@ -205,12 +238,13 @@ if __name__ == "__main__": import argparse parser = argparse.ArgumentParser(description="Train MNIST autoencoder") - parser.add_argument("--epochs", type=int, default=10, help="Number of epochs") + parser.add_argument("--epochs", type=int, default=None, help="Number of epochs") 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("--patience", type=int, default=5, help="Early stopping patience") parser.add_argument("--checkpoint", type=str, default=None, help="Resume from checkpoint") args = parser.parse_args() @@ -224,5 +258,6 @@ if __name__ == "__main__": latent_channels=args.latent_channels, val_split=args.val_split, test_split=args.test_split, + patience=args.patience, checkpoint_path=args.checkpoint, ) diff --git a/src/train_diffusion.py b/src/train_diffusion.py index 2ec4d60..d0777bc 100644 --- a/src/train_diffusion.py +++ b/src/train_diffusion.py @@ -3,16 +3,13 @@ Training script for MNIST diffusion model. """ from models import UNetMNIST -from src.utils import ( - ensure_dirs, save_checkpoint, tensor_to_image, - figure_to_image, -) +from src.utils import ensure_dirs, save_checkpoint from src.metrics import fid, kl_divergence, jsd from src.diffusion import p_xt_1_xt_x0_pred -from src.dataloader import get_diffusion_dataloader, get_mnist_dataset +from src.dataloader import get_diffusion_dataloaders, get_mnist_dataset from src.config import ( DEVICE, EPOCHS, LEARNING_RATE, BATCH_SIZE, NUM_WORKERS, - DIFFU_STEPS, NB_CHANNEL, IMG_SIZE, NB_LABEL, CHECKPOINT_DIR, PLOTS_DIR + DIFFU_STEPS, NB_CHANNEL, IMG_SIZE, NB_LABEL ) import os import torch @@ -28,11 +25,14 @@ sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) def train_diffusion( - epochs: int = EPOCHS, + epochs: int | None = None, learning_rate: float = LEARNING_RATE, batch_size: int = BATCH_SIZE, predict_x0: bool = True, use_attention: bool = False, + val_split: float = 0.1, + test_split: float = 0.1, + patience: int | None = 5, checkpoint_path: str | None = None, run_name: str | None = None, ): @@ -45,9 +45,17 @@ def train_diffusion( batch_size: Training batch size predict_x0: If True, model predicts x0. Otherwise predicts noise. use_attention: If True, use self-attention in UNet + val_split: Fraction of data for validation + test_split: Fraction of data for testing + patience: Number of epochs to wait for validation loss improvement before stopping checkpoint_path: Path to checkpoint to resume training from run_name: Name for this training run (for logging) """ + if epochs is None and patience is None: + raise ValueError("Must specify either epochs or patience for training") + if patience is not None and patience <= 0: + raise ValueError("Patience must be a positive integer") + ensure_dirs() # Setup run name and logging @@ -69,20 +77,36 @@ def train_diffusion( optimizer = torch.optim.Adam(model.parameters(), lr=learning_rate) criterion = nn.MSELoss() - # Get dataloader - dataloader = get_diffusion_dataloader( + # Get dataloaders + train_loader, val_loader, test_loader = get_diffusion_dataloaders( predict_x0=predict_x0, batch_size=batch_size, num_workers=NUM_WORKERS, + val_split=val_split, + test_split=test_split, ) + # Early stopping tracking + best_val_loss = float('inf') + epochs_without_improvement = 0 + # Training loop global_step = 0 - for epoch in range(1, epochs + 1): + epoch = 0 + while epochs is None or epoch < epochs: + epoch += 1 + model.train() epoch_loss = 0.0 - for batch_idx, (xt, t, vec, target) in enumerate(track(dataloader, description=f"Epoch {epoch}/{epochs}")): + if epochs: + description = f"Epoch {epoch}/{epochs}" + else: + description = f"Epoch {epoch}" + if patience: + description += f" (Patience: {patience-epochs_without_improvement})" + + for batch_idx, (xt, t, vec, target) in enumerate(track(train_loader, description=description)): optimizer.zero_grad() # Forward pass @@ -99,25 +123,68 @@ def train_diffusion( # Log training loss every 100 steps if global_step % 100 == 0: avg_loss = epoch_loss / (batch_idx + 1) - mlflow.log_metric("train_loss", avg_loss, step=global_step) + mlflow.log_metric("train/loss_step", avg_loss, step=global_step) - 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) - # Save checkpoint every epoch + val_loss = evaluate_diffusion_loss(model, val_loader, criterion) + test_loss = evaluate_diffusion_loss(model, test_loader, criterion) + + mlflow.log_metric("val/loss", val_loss, step=epoch) + mlflow.log_metric("test/loss", test_loss, step=epoch) + + # Early stopping check + if val_loss < best_val_loss: + best_val_loss = val_loss + epochs_without_improvement = 0 + # Save best checkpoint + save_checkpoint(model, "diffusion_best.pt") + else: + epochs_without_improvement += 1 + + # Save regular checkpoints save_checkpoint(model, f"diffusion_epoch_{epoch:03d}.pt") save_checkpoint(model, "diffusion_latest.pt") - # Generate and log sample images + # Stop if no improvement + if patience and epochs_without_improvement >= patience: + print(f"Early stopping: No improvement for {patience} epochs") + break + + # Generate and log sample images and metrics on test split if epoch % 1 == 0: - evaluate_and_log(model, epoch, predict_x0) + evaluate_and_log(model, epoch, test_loader, predict_x0) mlflow.end_run() return model -def evaluate_and_log(model: nn.Module, epoch: int, predict_x0: bool = True): - """Generate samples and log metrics.""" +def evaluate_diffusion_loss( + model: nn.Module, + dataloader, + criterion, +) -> float: + """Evaluate diffusion model and return average loss.""" + model.eval() + total_loss = 0.0 + + with torch.no_grad(): + for xt, t, vec, target in dataloader: + pred = model(xt, t, vec) + loss = criterion(pred, target) + total_loss += loss.item() + + return total_loss / len(dataloader) + + +def evaluate_and_log( + model: nn.Module, + epoch: int, + test_loader, + predict_x0: bool = True, +): + """Generate samples and log metrics evaluated on test split.""" model.eval() with torch.no_grad(): @@ -142,10 +209,13 @@ def evaluate_and_log(model: nn.Module, epoch: int, predict_x0: bool = True): fakes = np.concatenate(fakes) - # Get real samples for comparison - dataset = get_mnist_dataset(train=True) + # Get real samples from test split n_samples = len(fakes) - reals = torch.stack([dataset[i][0] for i in range(n_samples)]).numpy() + base_dataset = getattr(test_loader.dataset, "dataset", None) + if base_dataset is None: + base_dataset = get_mnist_dataset(train=True) + n_samples = min(n_samples, len(base_dataset)) + reals = torch.stack([base_dataset[i][0] for i in range(n_samples)]).cpu().numpy() reals = reals * 2 - 1 # Log metrics @@ -153,9 +223,9 @@ def evaluate_and_log(model: nn.Module, epoch: int, predict_x0: bool = True): kl_score = kl_divergence(reals, fakes) jsd_score = jsd(reals, fakes) - mlflow.log_metric("FID", fid_score, step=epoch) - mlflow.log_metric("KL Divergence", kl_score, step=epoch) - mlflow.log_metric("JSD", jsd_score, step=epoch) + mlflow.log_metric("test/FID", fid_score, step=epoch) + mlflow.log_metric("test/KL Divergence", kl_score, step=epoch) + mlflow.log_metric("test/JSD", jsd_score, step=epoch) # Log sample images using plotly fig = make_subplots(rows=4, cols=8, horizontal_spacing=0.01, vertical_spacing=0.02) @@ -176,17 +246,20 @@ def evaluate_and_log(model: nn.Module, epoch: int, predict_x0: bool = True): fig.update_xaxes(showticklabels=False, showgrid=False, zeroline=False) fig.update_yaxes(showticklabels=False, showgrid=False, zeroline=False) - mlflow.log_figure(fig, f"samples/epoch_{epoch:03d}.png") + mlflow.log_figure(fig, f"epoch_{epoch:03d}.png") if __name__ == "__main__": import argparse parser = argparse.ArgumentParser(description="Train MNIST diffusion model") - parser.add_argument("--epochs", type=int, default=EPOCHS, help="Number of epochs") + parser.add_argument("--epochs", type=int, default=None, help="Number of epochs") parser.add_argument("--lr", type=float, default=LEARNING_RATE, help="Learning rate") parser.add_argument("--batch-size", type=int, default=BATCH_SIZE, help="Batch size") parser.add_argument("--attention", action="store_true", help="Use self-attention") + 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("--patience", type=int, default=5, help="Early stopping patience") parser.add_argument("--checkpoint", type=str, default=None, help="Resume from checkpoint") parser.add_argument("--name", type=str, default=None, help="Run name") @@ -199,6 +272,9 @@ if __name__ == "__main__": learning_rate=args.lr, batch_size=args.batch_size, use_attention=args.attention, + val_split=args.val_split, + test_split=args.test_split, + patience=args.patience, checkpoint_path=args.checkpoint, run_name=args.name, ) -- cgit v1.3.1