import copy import time import torch import torch.nn as nn from torch.utils.data import DataLoader from sklearn.metrics import f1_score def train_model( model: nn.Module, train_loader: DataLoader, val_loader: DataLoader, criterion: nn.Module, optimizer: torch.optim.Optimizer, scheduler, device: torch.device, num_epochs: int = 30, patience: int = 5, model_name: str = "", ) -> dict: """Train model with mixed precision, early stopping on val macro F1. Returns a dict with: best_model_state, train_losses, val_losses, val_f1s (macro), val_weighted_f1s, train_time (seconds), epochs_trained, best_val_f1 (macro) """ scaler = torch.amp.GradScaler('cuda') if device.type == 'cuda' else None best_val_f1 = -1.0 best_state = None patience_counter = 0 train_losses, val_losses, val_f1s, val_weighted_f1s = [], [], [], [] start_time = time.time() for epoch in range(1, num_epochs + 1): # ── Training phase ────────────────────────────────────────────────── model.train() running_loss = 0.0 n_train = 0 for images, labels in train_loader: images = images.to(device, non_blocking=True) labels = labels.to(device, non_blocking=True) optimizer.zero_grad() if scaler is not None: with torch.amp.autocast('cuda'): outputs = model(images) loss = criterion(outputs, labels) scaler.scale(loss).backward() scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) scaler.step(optimizer) scaler.update() else: outputs = model(images) loss = criterion(outputs, labels) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) optimizer.step() running_loss += loss.item() * images.size(0) n_train += images.size(0) train_loss = running_loss / n_train train_losses.append(train_loss) # ── Validation phase ───────────────────────────────────────────────── model.eval() val_running_loss = 0.0 n_val = 0 all_preds, all_labels = [], [] with torch.no_grad(): for images, labels in val_loader: images = images.to(device, non_blocking=True) labels = labels.to(device, non_blocking=True) if scaler is not None: with torch.amp.autocast('cuda'): outputs = model(images) loss = criterion(outputs, labels) else: outputs = model(images) loss = criterion(outputs, labels) val_running_loss += loss.item() * images.size(0) n_val += images.size(0) preds = outputs.argmax(dim=1) all_preds.extend(preds.cpu().numpy()) all_labels.extend(labels.cpu().numpy()) val_loss = val_running_loss / n_val val_losses.append(val_loss) val_f1_macro = f1_score(all_labels, all_preds, average='macro', zero_division=0) val_f1_weighted = f1_score(all_labels, all_preds, average='weighted', zero_division=0) val_f1s.append(val_f1_macro) val_weighted_f1s.append(val_f1_weighted) current_lr = scheduler.get_last_lr()[0] if hasattr(scheduler, 'get_last_lr') else optimizer.param_groups[0]['lr'] print( f"[{model_name}] Epoch {epoch:3d}/{num_epochs} | " f"train loss: {train_loss:.4f} | val loss: {val_loss:.4f} | " f"macro F1: {val_f1_macro:.4f} | weighted F1: {val_f1_weighted:.4f} | lr: {current_lr:.2e}" ) scheduler.step() # ── Early stopping (monitored on macro F1) ─────────────────────────── if val_f1_macro > best_val_f1: best_val_f1 = val_f1_macro best_state = copy.deepcopy(model.state_dict()) patience_counter = 0 else: patience_counter += 1 if patience_counter >= patience: print(f" Early stopping at epoch {epoch} (no improvement for {patience} epochs).") break train_time = time.time() - start_time epochs_trained = len(train_losses) print(f" Training complete: {epochs_trained} epochs, {train_time:.1f}s, best val macro F1: {best_val_f1:.4f}") return { 'best_model_state': best_state, 'train_losses': train_losses, 'val_losses': val_losses, 'val_f1s': val_f1s, # macro F1 per epoch (used for early stopping) 'val_weighted_f1s': val_weighted_f1s, 'train_time': train_time, 'epochs_trained': epochs_trained, 'best_val_f1': best_val_f1, # best macro F1 }