CoolFace
Apppublic

Timechils/sapiens-mvp

sourceHugging Faceupdated 10mo agoView on Hugging Face
0likes
train_improved.py204 linesDownload Raw Back to scripts
1import torch2import torch.nn as nn3import torch.optim as optim4from torch.utils.data import DataLoader, WeightedRandomSampler5from sklearn.metrics import accuracy_score, classification_report6from tensorboardX import SummaryWriter7from torchvision import transforms8import sys9import os10import time11import numpy as np12from collections import Counter13 14sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), '..')))15from scripts.model_setup import PVFaultDataset16 17print("=" * 70)18print("IMPROVED Thermal PV Fault Detection Training")19print("=" * 70)20 21# Configuration22device = 'cuda' if torch.cuda.is_available() else 'cpu'23DATA_DIR_TRAIN = 'data/thermal_final/train'24DATA_DIR_VAL = 'data/thermal_final/val'25MODEL_SAVE_PATH = 'models/thermal_improved.pth'26NUM_EPOCHS = 5027BATCH_SIZE = 16  # Smaller batch for better gradients28LEARNING_RATE = 3e-4  # Higher learning rate29 30print(f"Using device: {device}\n")31 32# Enhanced data augmentation33train_transform = transforms.Compose([34    transforms.Resize((256, 256)),35    transforms.RandomCrop(224),36    transforms.RandomHorizontalFlip(p=0.5),37    transforms.RandomVerticalFlip(p=0.5),38    transforms.RandomRotation(20),39    transforms.ColorJitter(brightness=0.3, contrast=0.3),40    transforms.RandomAffine(degrees=0, translate=(0.1, 0.1)),41    transforms.ToTensor(),42    transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])43])44 45val_transform = transforms.Compose([46    transforms.Resize((224, 224)),47    transforms.ToTensor(),48    transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])49])50 51# Load datasets52train_dataset = PVFaultDataset(root_dir=DATA_DIR_TRAIN, transform=train_transform)53val_dataset = PVFaultDataset(root_dir=DATA_DIR_VAL, transform=val_transform)54 55num_classes = len(train_dataset.classes)56print(f"Classes: {train_dataset.classes}")57print(f"Train: {len(train_dataset)}, Val: {len(val_dataset)}\n")58 59# Calculate class weights for imbalanced dataset60class_counts = Counter([label for _, label in train_dataset.samples])61class_weights = torch.FloatTensor([1.0 / class_counts[i] for i in range(num_classes)])62class_weights = class_weights / class_weights.sum() * num_classes63class_weights = class_weights.to(device)64 65print("Class distribution (train):")66for i, cls in enumerate(train_dataset.classes):67    print(f"  {cls:15s}: {class_counts[i]:4d} images (weight: {class_weights[i]:.4f})")68 69# Weighted sampler for balanced batches70sample_weights = [class_weights[label] for _, label in train_dataset.samples]71sampler = WeightedRandomSampler(sample_weights, len(sample_weights), replacement=True)72 73train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, sampler=sampler, num_workers=4)74val_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=4)75 76# Create model - UNFREEZE MORE LAYERS77from torchvision import models78 79model = models.resnet50(weights='IMAGENET1K_V1')80 81# Unfreeze layer2, layer3, layer4 (more layers)82for param in model.parameters():83    param.requires_grad = False84 85for param in model.layer2.parameters():86    param.requires_grad = True87for param in model.layer3.parameters():88    param.requires_grad = True89for param in model.layer4.parameters():90    param.requires_grad = True91 92# Custom classifier93num_ftrs = model.fc.in_features94model.fc = nn.Sequential(95    nn.Dropout(0.5),96    nn.Linear(num_ftrs, 512),97    nn.ReLU(),98    nn.BatchNorm1d(512),99    nn.Dropout(0.3),100    nn.Linear(512, 256),101    nn.ReLU(),102    nn.BatchNorm1d(256),103    nn.Dropout(0.2),104    nn.Linear(256, num_classes)105)106 107model = model.to(device)108 109# Weighted loss for class imbalance110criterion = nn.CrossEntropyLoss(weight=class_weights, label_smoothing=0.1)111optimizer = optim.AdamW(model.parameters(), lr=LEARNING_RATE, weight_decay=0.01)112scheduler = optim.lr_scheduler.OneCycleLR(optimizer, max_lr=LEARNING_RATE, 113                                         steps_per_epoch=len(train_loader), 114                                         epochs=NUM_EPOCHS)115 116writer = SummaryWriter('runs/thermal_improved')117best_val_acc = 0.0118patience, patience_counter = 15, 0119 120print("\n" + "=" * 70)121print("Starting Training")122print("=" * 70 + "\n")123 124start_time = time.time()125 126for epoch in range(NUM_EPOCHS):127    # Training128    model.train()129    running_loss = 0.0130    train_preds, train_labels = [], []131    132    for inputs, labels in train_loader:133        inputs, labels = inputs.to(device), labels.to(device)134        135        optimizer.zero_grad()136        outputs = model(inputs)137        loss = criterion(outputs, labels)138        loss.backward()139        optimizer.step()140        scheduler.step()141        142        running_loss += loss.item() * inputs.size(0)143        _, preds = torch.max(outputs, 1)144        train_preds.extend(preds.cpu().numpy())145        train_labels.extend(labels.cpu().numpy())146    147    train_loss = running_loss / len(train_dataset)148    train_acc = accuracy_score(train_labels, train_preds)149    150    # Validation151    model.eval()152    val_preds, val_labels = [], []153    val_loss = 0.0154    155    with torch.no_grad():156        for inputs, labels in val_loader:157            inputs, labels = inputs.to(device), labels.to(device)158            outputs = model(inputs)159            loss = criterion(outputs, labels)160            val_loss += loss.item() * inputs.size(0)161            _, preds = torch.max(outputs, 1)162            val_preds.extend(preds.cpu().numpy())163            val_labels.extend(labels.cpu().numpy())164    165    val_loss = val_loss / len(val_dataset)166    val_acc = accuracy_score(val_labels, val_preds)167    168    print(f"Epoch {epoch+1:02d}/{NUM_EPOCHS} | "169          f"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f} | "170          f"Val Loss: {val_loss:.4f} | Val Acc: {val_acc:.4f}")171    172    writer.add_scalar('Loss/train', train_loss, epoch)173    writer.add_scalar('Accuracy/train', train_acc, epoch)174    writer.add_scalar('Loss/val', val_loss, epoch)175    writer.add_scalar('Accuracy/val', val_acc, epoch)176    177    if val_acc > best_val_acc:178        best_val_acc = val_acc179        torch.save({180            'epoch': epoch,181            'model_state_dict': model.state_dict(),182            'val_acc': val_acc,183            'classes': train_dataset.classes184        }, MODEL_SAVE_PATH)185        print(f"  ✓ New best: {best_val_acc:.4f}")186        patience_counter = 0187    else:188        patience_counter += 1189    190    if patience_counter >= patience:191        print(f"\nEarly stopping at epoch {epoch+1}")192        break193 194writer.close()195time_elapsed = time.time() - start_time196 197print("\n" + "=" * 70)198print(f"Training Complete! Time: {time_elapsed//60:.0f}m {time_elapsed%60:.0f}s")199print(f"Best Val Acc: {best_val_acc:.4f} ({best_val_acc*100:.2f}%)")200print("=" * 70 + "\n")201 202print("Per-Class Performance:")203print(classification_report(val_labels, val_preds, target_names=train_dataset.classes, digits=4))204