Timechils/sapiens-mvp
0
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 