CoolFace
Apppublic

Timechils/sapiens-mvp

sourceHugging Faceupdated 10mo agoView on Hugging Face
0likes
train_from_scratch_thermal.py158 linesDownload Raw Back to scripts
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