Droid210/FleetVision
0
1"""Training loop with recall-focused optimization."""2import copy3from typing import Dict, List, Tuple4 5import torch6from sklearn.metrics import precision_recall_curve7from torch import nn, optim8from torch.utils.data import DataLoader9 10 11def calculate_class_weights(loader: DataLoader) -> torch.Tensor:12 """Calculate class weights to balance dataset.13 14 Args:15 loader: DataLoader to analyze.16 17 Returns:18 Class weights tensor.19 """20 class_counts = torch.zeros(2)21 for _, labels in loader:22 class_counts += torch.bincount(labels, minlength=2).float()23 24 weights = 1.0 / (class_counts + 1e-8)25 weights = weights / weights.sum()26 return weights27 28 29def run_epoch(30 model: nn.Module,31 loader: DataLoader,32 criterion: nn.Module,33 optimizer: optim.Optimizer | None,34 device: torch.device,35 recall_weight: float = 1.0,36) -> Tuple[float, float, float, float]:37 """Run one epoch with recall tracking.38 39 Args:40 model: ViT model.41 loader: DataLoader.42 criterion: Loss function.43 optimizer: Optimizer (None for validation).44 device: Device.45 recall_weight: Weight for recall loss.46 47 Returns:48 Tuple of (loss, accuracy, precision, recall).49 """50 is_train = optimizer is not None51 model.train(mode=is_train)52 53 running_loss = 0.054 all_preds = []55 all_labels = []56 57 for pixel_values, labels in loader:58 pixel_values = pixel_values.to(device)59 labels = labels.to(device)60 61 if is_train:62 optimizer.zero_grad()63 64 with torch.set_grad_enabled(is_train):65 outputs = model(pixel_values)66 logits = outputs.logits67 loss = criterion(logits, labels)68 69 if is_train:70 loss.backward()71 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)72 optimizer.step()73 74 running_loss += loss.item() * pixel_values.size(0)75 predictions = logits.argmax(dim=1)76 77 all_preds.extend(predictions.cpu().numpy().tolist())78 all_labels.extend(labels.cpu().numpy().tolist())79 80 epoch_loss = running_loss / len(loader.dataset)81 82 # Calculate metrics83 all_preds = torch.tensor(all_preds)84 all_labels = torch.tensor(all_labels)85 86 accuracy = (all_preds == all_labels).float().mean().item()87 88 # Calculate precision and recall89 tp = ((all_preds == 1) & (all_labels == 1)).sum().item()90 fp = ((all_preds == 1) & (all_labels == 0)).sum().item()91 fn = ((all_preds == 0) & (all_labels == 1)).sum().item()92 93 precision = tp / (tp + fp + 1e-8)94 recall = tp / (tp + fn + 1e-8)95 96 return epoch_loss, accuracy, precision, recall97 98 99def train_model(100 model: nn.Module,101 train_loader: DataLoader,102 val_loader: DataLoader,103 epochs: int,104 learning_rate: float,105 device: torch.device,106 output_path,107 recall_weight: float = 2.0,108) -> Dict[str, List[float]]:109 """Train model with recall optimization.110 111 Args:112 model: ViT model.113 train_loader: Training dataloader.114 val_loader: Validation dataloader.115 epochs: Number of epochs.116 learning_rate: Learning rate.117 device: Device.118 output_path: Path to save best model.119 recall_weight: Weight for recall vs accuracy trade-off.120 121 Returns:122 Training history.123 """124 criterion = nn.CrossEntropyLoss()125 optimizer = optim.AdamW(126 [p for p in model.parameters() if p.requires_grad],127 lr=learning_rate,128 weight_decay=1e-5,129 )130 131 history: Dict[str, List[float]] = {132 "train_loss": [],133 "train_acc": [],134 "train_recall": [],135 "val_loss": [],136 "val_acc": [],137 "val_recall": [],138 }139 140 best_val_recall = -1.0141 best_state = copy.deepcopy(model.state_dict())142 143 for epoch in range(1, epochs + 1):144 train_loss, train_acc, train_prec, train_recall = run_epoch(145 model=model,146 loader=train_loader,147 criterion=criterion,148 optimizer=optimizer,149 device=device,150 recall_weight=recall_weight,151 )152 153 val_loss, val_acc, val_prec, val_recall = run_epoch(154 model=model,155 loader=val_loader,156 criterion=criterion,157 optimizer=None,158 device=device,159 )160 161 history["train_loss"].append(train_loss)162 history["train_acc"].append(train_acc)163 history["train_recall"].append(train_recall)164 history["val_loss"].append(val_loss)165 history["val_acc"].append(val_acc)166 history["val_recall"].append(val_recall)167 168 print(169 f"Epoch {epoch:02d}/{epochs} | "170 f"Train Loss: {train_loss:.4f}, Acc: {train_acc:.4f}, Recall: {train_recall:.4f} | "171 f"Val Loss: {val_loss:.4f}, Acc: {val_acc:.4f}, Recall: {val_recall:.4f}"172 )173 174 # Save best model based on Recall (minimize False Negatives)175 if val_recall > best_val_recall:176 best_val_recall = val_recall177 best_state = copy.deepcopy(model.state_dict())178 checkpoint = {179 "model_state_dict": best_state,180 "best_val_recall": best_val_recall,181 "best_val_acc": val_acc,182 }183 torch.save(checkpoint, output_path)184 print(f"✓ Saved best model (Recall: {best_val_recall:.4f})")185 186 model.load_state_dict(best_state)187 print(f"\nBest validation recall: {best_val_recall:.4f}")188 print(f"Best model saved to: {output_path}")189 return history190 