aboutsummaryrefslogtreecommitdiff
path: root/src
diff options
context:
space:
mode:
authorgdamms <damguillotin@gmail.com>2026-02-06 11:07:37 +0100
committergdamms <damguillotin@gmail.com>2026-02-06 11:07:37 +0100
commitf1fee867bbbffdd2686cc2f92313496f77d08b2f (patch)
treee4a2cd033c8d9fdc372c6fbe033436db68183a2d /src
parent8ee90daeedef75135ab50dd2bd5c5349929b53dd (diff)
downloaddiffusion-mnist-f1fee867bbbffdd2686cc2f92313496f77d08b2f.tar.gz
diffusion-mnist-f1fee867bbbffdd2686cc2f92313496f77d08b2f.zip
add patience for early stopping or infinit epochs
Diffstat (limited to 'src')
-rw-r--r--src/dataloader.py74
-rw-r--r--src/train_autoencoder.py57
-rw-r--r--src/train_diffusion.py130
3 files changed, 223 insertions, 38 deletions
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,
)