Timechils/sapiens-mvp
0
1import torch2import torch.nn as nn3import torch.optim as optim4from torch.utils.data import DataLoader5from sklearn.metrics import accuracy_score, classification_report6from torchvision import transforms, models7import sys, os, time8 9sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), '..')))10from scripts.model_setup import PVFaultDataset11 12print("=" * 70)13print("Training ResNet-50 FROM SCRATCH on Thermal Images")14print("=" * 70)15 16device = 'cuda' if torch.cuda.is_available() else 'cpu'17 18# Use patches dataset (already classification format!)19DATA_DIR_TRAIN = 'data/patches/train'20DATA_DIR_VAL = 'data/patches/val'21MODEL_SAVE_PATH = 'models/thermal_from_scratch.pth'22 23NUM_EPOCHS = 10024BATCH_SIZE = 3225LEARNING_RATE = 1e-326 27# Aggressive augmentation for small dataset28train_transform = transforms.Compose([29 transforms.Resize((256, 256)),30 transforms.RandomResizedCrop(224, scale=(0.8, 1.0)),31 transforms.RandomHorizontalFlip(),32 transforms.RandomVerticalFlip(),33 transforms.RandomRotation(30),34 transforms.ColorJitter(brightness=0.4, contrast=0.4, saturation=0.2),35 transforms.RandomAffine(degrees=0, translate=(0.2, 0.2), scale=(0.8, 1.2)),36 transforms.ToTensor(),37 transforms.Normalize([0.5, 0.5, 0.5], [0.5, 0.5, 0.5])38])39 40val_transform = transforms.Compose([41 transforms.Resize((224, 224)),42 transforms.ToTensor(),43 transforms.Normalize([0.5, 0.5, 0.5], [0.5, 0.5, 0.5])44])45 46train_dataset = PVFaultDataset(DATA_DIR_TRAIN, train_transform)47val_dataset = PVFaultDataset(DATA_DIR_VAL, val_transform)48 49train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=4, pin_memory=True)50val_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=4, pin_memory=True)51 52num_classes = len(train_dataset.classes)53print(f"\nClasses: {train_dataset.classes}")54print(f"Train: {len(train_dataset)}, Val: {len(val_dataset)}\n")55 56# Create ResNet-50 WITHOUT pretrained weights57model = models.resnet50(weights=None)58 59# Initialize weights properly for thermal images60def init_weights(m):61 if isinstance(m, nn.Conv2d):62 nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')63 elif isinstance(m, nn.BatchNorm2d):64 nn.init.constant_(m.weight, 1)65 nn.init.constant_(m.bias, 0)66 67model.apply(init_weights)68 69# Replace final layer70num_ftrs = model.fc.in_features71model.fc = nn.Sequential(72 nn.Dropout(0.5),73 nn.Linear(num_ftrs, 256),74 nn.ReLU(),75 nn.Dropout(0.3),76 nn.Linear(256, num_classes)77)78 79model = model.to(device)80 81criterion = nn.CrossEntropyLoss(label_smoothing=0.1)82optimizer = optim.SGD(model.parameters(), lr=LEARNING_RATE, momentum=0.9, weight_decay=1e-4)83scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='max', factor=0.5, patience=10)84print("=" * 70)85print("Starting Training")86print("=" * 70)87 88best_val_acc = 0.089patience_counter = 090start_time = time.time()91 92for epoch in range(NUM_EPOCHS):93 # Train94 model.train()95 running_loss = 0.096 train_preds, train_labels = [], []97 98 for inputs, labels in train_loader:99 inputs, labels = inputs.to(device), labels.to(device)100 101 optimizer.zero_grad()102 outputs = model(inputs)103 loss = criterion(outputs, labels)104 loss.backward()105 optimizer.step()106 107 running_loss += loss.item() * inputs.size(0)108 _, preds = torch.max(outputs, 1)109 train_preds.extend(preds.cpu().numpy())110 train_labels.extend(labels.cpu().numpy())111 112 train_loss = running_loss / len(train_dataset)113 train_acc = accuracy_score(train_labels, train_preds)114 115 # Validate116 model.eval()117 val_preds, val_labels = [], []118 val_loss = 0.0119 120 with torch.no_grad():121 for inputs, labels in val_loader:122 inputs, labels = inputs.to(device), labels.to(device)123 outputs = model(inputs)124 loss = criterion(outputs, labels)125 val_loss += loss.item() * inputs.size(0)126 _, preds = torch.max(outputs, 1)127 val_preds.extend(preds.cpu().numpy())128 val_labels.extend(labels.cpu().numpy())129 130 val_loss = val_loss / len(val_dataset)131 val_acc = accuracy_score(val_labels, val_preds)132 133 print(f"Epoch {epoch+1:03d}/{NUM_EPOCHS} | "134 f"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f} | "135 f"Val Loss: {val_loss:.4f} | Val Acc: {val_acc:.4f}")136 137 scheduler.step(val_acc)138 139 if val_acc > best_val_acc:140 best_val_acc = val_acc141 torch.save(model.state_dict(), MODEL_SAVE_PATH)142 print(f" ✓ Best: {best_val_acc:.4f} ({best_val_acc*100:.1f}%)")143 patience_counter = 0144 else:145 patience_counter += 1146 147 if patience_counter >= 20:148 print(f"\nEarly stopping at epoch {epoch+1}")149 break150 151time_elapsed = time.time() - start_time152print(f"\n{'='*70}")153print(f"Complete! Time: {time_elapsed//60:.0f}m {time_elapsed%60:.0f}s")154print(f"Best Validation Accuracy: {best_val_acc:.4f} ({best_val_acc*100:.2f}%)")155print(f"{'='*70}\n")156 157print(classification_report(val_labels, val_preds, target_names=train_dataset.classes, digits=4))158 