CoolFace
Apppublic

Droid210/FleetVision

sourceHugging Faceupdated 5mo agoView on Hugging Face
0likes
train.py190 linesDownload Raw Back to model_b
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