eel4759_classification

Comparison of different ImageNet-based CNN models for classifying diabetic retinopathy images
Log | Files | Refs | README

train.py (5243B)


      1 import copy
      2 import time
      3 
      4 import torch
      5 import torch.nn as nn
      6 from torch.utils.data import DataLoader
      7 from sklearn.metrics import f1_score
      8 
      9 
     10 def train_model(
     11     model: nn.Module,
     12     train_loader: DataLoader,
     13     val_loader: DataLoader,
     14     criterion: nn.Module,
     15     optimizer: torch.optim.Optimizer,
     16     scheduler,
     17     device: torch.device,
     18     num_epochs: int = 30,
     19     patience: int = 5,
     20     model_name: str = "",
     21 ) -> dict:
     22     """Train model with mixed precision, early stopping on val macro F1.
     23 
     24     Returns a dict with:
     25       best_model_state, train_losses, val_losses, val_f1s (macro), val_weighted_f1s,
     26       train_time (seconds), epochs_trained, best_val_f1 (macro)
     27     """
     28     scaler = torch.amp.GradScaler('cuda') if device.type == 'cuda' else None
     29 
     30     best_val_f1 = -1.0
     31     best_state = None
     32     patience_counter = 0
     33 
     34     train_losses, val_losses, val_f1s, val_weighted_f1s = [], [], [], []
     35     start_time = time.time()
     36 
     37     for epoch in range(1, num_epochs + 1):
     38         # ── Training phase ──────────────────────────────────────────────────
     39         model.train()
     40         running_loss = 0.0
     41         n_train = 0
     42 
     43         for images, labels in train_loader:
     44             images = images.to(device, non_blocking=True)
     45             labels = labels.to(device, non_blocking=True)
     46 
     47             optimizer.zero_grad()
     48 
     49             if scaler is not None:
     50                 with torch.amp.autocast('cuda'):
     51                     outputs = model(images)
     52                     loss = criterion(outputs, labels)
     53                 scaler.scale(loss).backward()
     54                 scaler.unscale_(optimizer)
     55                 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
     56                 scaler.step(optimizer)
     57                 scaler.update()
     58             else:
     59                 outputs = model(images)
     60                 loss = criterion(outputs, labels)
     61                 loss.backward()
     62                 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
     63                 optimizer.step()
     64 
     65             running_loss += loss.item() * images.size(0)
     66             n_train += images.size(0)
     67 
     68         train_loss = running_loss / n_train
     69         train_losses.append(train_loss)
     70 
     71         # ── Validation phase ─────────────────────────────────────────────────
     72         model.eval()
     73         val_running_loss = 0.0
     74         n_val = 0
     75         all_preds, all_labels = [], []
     76 
     77         with torch.no_grad():
     78             for images, labels in val_loader:
     79                 images = images.to(device, non_blocking=True)
     80                 labels = labels.to(device, non_blocking=True)
     81 
     82                 if scaler is not None:
     83                     with torch.amp.autocast('cuda'):
     84                         outputs = model(images)
     85                         loss = criterion(outputs, labels)
     86                 else:
     87                     outputs = model(images)
     88                     loss = criterion(outputs, labels)
     89 
     90                 val_running_loss += loss.item() * images.size(0)
     91                 n_val += images.size(0)
     92 
     93                 preds = outputs.argmax(dim=1)
     94                 all_preds.extend(preds.cpu().numpy())
     95                 all_labels.extend(labels.cpu().numpy())
     96 
     97         val_loss = val_running_loss / n_val
     98         val_losses.append(val_loss)
     99 
    100         val_f1_macro = f1_score(all_labels, all_preds, average='macro', zero_division=0)
    101         val_f1_weighted = f1_score(all_labels, all_preds, average='weighted', zero_division=0)
    102         val_f1s.append(val_f1_macro)
    103         val_weighted_f1s.append(val_f1_weighted)
    104 
    105         current_lr = scheduler.get_last_lr()[0] if hasattr(scheduler, 'get_last_lr') else optimizer.param_groups[0]['lr']
    106         print(
    107             f"[{model_name}] Epoch {epoch:3d}/{num_epochs} | "
    108             f"train loss: {train_loss:.4f} | val loss: {val_loss:.4f} | "
    109             f"macro F1: {val_f1_macro:.4f} | weighted F1: {val_f1_weighted:.4f} | lr: {current_lr:.2e}"
    110         )
    111 
    112         scheduler.step()
    113 
    114         # ── Early stopping (monitored on macro F1) ───────────────────────────
    115         if val_f1_macro > best_val_f1:
    116             best_val_f1 = val_f1_macro
    117             best_state = copy.deepcopy(model.state_dict())
    118             patience_counter = 0
    119         else:
    120             patience_counter += 1
    121             if patience_counter >= patience:
    122                 print(f"  Early stopping at epoch {epoch} (no improvement for {patience} epochs).")
    123                 break
    124 
    125     train_time = time.time() - start_time
    126     epochs_trained = len(train_losses)
    127 
    128     print(f"  Training complete: {epochs_trained} epochs, {train_time:.1f}s, best val macro F1: {best_val_f1:.4f}")
    129 
    130     return {
    131         'best_model_state': best_state,
    132         'train_losses': train_losses,
    133         'val_losses': val_losses,
    134         'val_f1s': val_f1s,           # macro F1 per epoch (used for early stopping)
    135         'val_weighted_f1s': val_weighted_f1s,
    136         'train_time': train_time,
    137         'epochs_trained': epochs_trained,
    138         'best_val_f1': best_val_f1,   # best macro F1
    139     }