Timechils/sapiens-mvp
0
1import torch2import torch.nn as nn3import torch.optim as optim4from torchvision import datasets, models, transforms5from torch.utils.data import DataLoader, random_split, WeightedRandomSampler6from collections import Counter7import os8import time9import numpy as np10import random11 12print("=" * 80)13print("๐ REVISITING EfficientNetV2-S (Lower LR) - Target: 90%+ Accuracy ๐")14print("=" * 80)15 16# --- 1. CONFIGURATION ---17EPOCHS = 75 # EfficientNet usually trains faster18LEARNING_RATE = 0.0003 # <<< LOWERED LR for EfficientNetV2-S19BATCH_SIZE = 32 # Adjust based on VRAM if needed20DATA_DIR = 'master_dataset'21MODEL_SAVE_PATH = 'thermal_model_effnetv2s_lowLR_best.pth' # New name22PATIENCE = 15 # Standard patience23NUM_WORKERS = 424SEED = 4225 26# --- 2. SETUP ---27torch.manual_seed(SEED)28np.random.seed(SEED)29random.seed(SEED)30device = torch.device("cuda" if torch.cuda.is_available() else "cpu")31print(f"\nโ Using device: {device}")32base_dir = os.path.abspath(os.path.join(os.path.dirname(__file__), '..'))33data_path = os.path.join(base_dir, DATA_DIR)34print(f"โ Data Path: {data_path}")35if not os.path.exists(data_path) or not os.listdir(data_path):36 print(f"\nโ ERROR: Data directory '{data_path}' missing/empty!")37 exit(1)38 39# --- 3. DATA AUGMENTATION & TRANSFORMS ---40# (Using the same successful augmentations)41print("\n--- Applying Data Augmentation ---")42train_transforms = transforms.Compose([43 transforms.Resize((256, 256)),44 transforms.RandomResizedCrop(224, scale=(0.7, 1.0)),45 transforms.RandomHorizontalFlip(p=0.5),46 transforms.RandomVerticalFlip(p=0.5),47 transforms.RandomRotation(35),48 transforms.ColorJitter(brightness=0.4, contrast=0.4, saturation=0.3, hue=0.1),49 transforms.RandomAffine(degrees=0, translate=(0.15, 0.15), scale=(0.9, 1.1)),50 transforms.ToTensor(),51 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])52])53val_transforms = transforms.Compose([54 transforms.Resize((224, 224)),55 transforms.ToTensor(),56 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])57])58print("โ Augmentations defined.")59 60# --- 4. LOAD DATASET & SPLIT ---61# (Same as before)62print("\n--- Loading and Splitting Dataset ---")63full_dataset = datasets.ImageFolder(data_path)64class_names = full_dataset.classes65num_classes = len(class_names)66print(f"โ Found {len(full_dataset)} total images.")67print(f"โ Classes ({num_classes}): {class_names}")68train_size = int(0.8 * len(full_dataset))69val_size = len(full_dataset) - train_size70train_dataset_split, val_dataset_split = random_split(full_dataset, [train_size, val_size],71 generator=torch.Generator().manual_seed(SEED))72class TransformedDataset(torch.utils.data.Dataset):73 def __init__(self, subset, transform=None): self.subset = subset; self.transform = transform74 def __getitem__(self, index): x, y = self.subset[index]; return self.transform(x) if self.transform else x, y75 def __len__(self): return len(self.subset)76train_dataset = TransformedDataset(train_dataset_split, transform=train_transforms)77val_dataset = TransformedDataset(val_dataset_split, transform=val_transforms)78print(f"โ Split complete: Train={len(train_dataset)}, Validation={len(val_dataset)}")79 80# --- 5. WEIGHTED RANDOM SAMPLER ---81# (Same successful sampler setup)82print("\n--- Implementing Weighted Random Sampler ---")83train_indices = train_dataset_split.indices84train_targets = [full_dataset.targets[i] for i in train_indices]85class_sample_count = np.array([len(np.where(train_targets == t)[0]) for t in np.unique(train_targets)])86print("Class distribution (Training Set):")87for i, count in enumerate(class_sample_count): print(f" {class_names[i]:15s}: {count:5d} images")88weight = 1. / np.maximum(class_sample_count, 1)89samples_weight = np.array([weight[t] for t in train_targets])90samples_weight = torch.from_numpy(samples_weight).double()91sampler = WeightedRandomSampler(samples_weight, len(samples_weight), replacement=True)92print("โ WeightedRandomSampler configured.")93train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, sampler=sampler, num_workers=NUM_WORKERS, pin_memory=True, drop_last=True)94val_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS, pin_memory=True)95print("โ DataLoaders created.")96 97# --- 6. MODEL DEFINITION (EfficientNetV2-S Fine-tuning) --- ### <<-- USING EFFICIENTNET -->> ###98print("\n--- Building EfficientNetV2-S Model ---")99model = models.efficientnet_v2_s(weights=models.EfficientNet_V2_S_Weights.IMAGENET1K_V1)100for param in model.parameters(): param.requires_grad = False101num_ftrs = model.classifier[1].in_features102model.classifier[1] = nn.Sequential(103 nn.Dropout(p=0.3, inplace=True), # Use dropout rate appropriate for EfficientNet104 nn.Linear(num_ftrs, num_classes)105)106print(f"โ Replaced final classifier layer for {num_classes} classes with Dropout.")107for param in model.classifier[1].parameters(): param.requires_grad = True # Only unfreeze the new layer108model = model.to(device)109trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)110total_params = sum(p.numel() for p in model.parameters())111print(f"โ Model ready. Trainable parameters: {trainable_params:,} / {total_params:,}")112 113# --- 7. LOSS, OPTIMIZER, SCHEDULER ---114print("\n--- Configuring Training Components ---")115criterion = nn.CrossEntropyLoss(label_smoothing=0.1)116print(f"โ Loss: CrossEntropyLoss (Label Smoothing=0.1)")117# Use AdamW for just the trainable parameters (the final layer)118optimizer = optim.AdamW(119 filter(lambda p: p.requires_grad, model.parameters()),120 lr=LEARNING_RATE, # <<< Use the lowered LR121 weight_decay=0.01122)123print(f"โ Optimizer: AdamW (LR={LEARNING_RATE}, Weight Decay=0.01)")124# Adjust T_max for the scheduler125scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=LEARNING_RATE / 100)126print(f"โ Scheduler: CosineAnnealingLR (T_max={EPOCHS})")127 128# --- 8. TRAINING LOOP ---129# (Remains the same structure)130print("\n" + "=" * 80)131print(f"๐ STARTING TRAINING FOR {EPOCHS} EPOCHS (PATIENCE={PATIENCE}) ๐")132print("=" * 80 + "\n")133best_val_acc = 0.0134patience_counter = 0135start_time = time.time()136scaler = torch.cuda.amp.GradScaler(enabled=torch.cuda.is_available())137 138for epoch in range(EPOCHS):139 model.train()140 running_train_loss = 0.0; running_train_corrects = 0; batch_count = 0141 for inputs, labels in train_loader:142 inputs, labels = inputs.to(device, non_blocking=True), labels.to(device, non_blocking=True)143 with torch.cuda.amp.autocast(enabled=torch.cuda.is_available()):144 outputs = model(inputs); loss = criterion(outputs, labels)145 scaler.scale(loss).backward(); scaler.step(optimizer); scaler.update()146 optimizer.zero_grad(set_to_none=True)147 _, preds = torch.max(outputs, 1)148 running_train_loss += loss.item() * inputs.size(0)149 running_train_corrects += torch.sum(preds == labels.data)150 batch_count += 1151 epoch_train_loss = running_train_loss / len(train_dataset) if len(train_dataset) > 0 else 0152 epoch_train_acc = running_train_corrects.double() / len(train_dataset) if len(train_dataset) > 0 else 0153 154 model.eval()155 running_val_loss = 0.0; running_val_corrects = 0156 with torch.no_grad():157 for inputs, labels in val_loader:158 inputs, labels = inputs.to(device, non_blocking=True), labels.to(device, non_blocking=True)159 with torch.cuda.amp.autocast(enabled=torch.cuda.is_available()):160 outputs = model(inputs); loss = criterion(outputs, labels)161 _, preds = torch.max(outputs, 1)162 running_val_loss += loss.item() * inputs.size(0)163 running_val_corrects += torch.sum(preds == labels.data)164 epoch_val_loss = running_val_loss / len(val_dataset) if len(val_dataset) > 0 else 0165 epoch_val_acc = running_val_corrects.double() / len(val_dataset) if len(val_dataset) > 0 else 0166 167 current_lr = optimizer.param_groups[0]['lr']168 print(f"Epoch {epoch+1:03d}/{EPOCHS} | Train Loss: {epoch_train_loss:.4f} Acc: {epoch_train_acc:.4f} | Val Loss: {epoch_val_loss:.4f} Acc: {epoch_val_acc:.4f} | LR: {current_lr:.6f}")169 scheduler.step()170 171 if epoch_val_acc > best_val_acc:172 best_val_acc = epoch_val_acc; save_path = os.path.join(base_dir, MODEL_SAVE_PATH)173 try:174 torch.save({'epoch': epoch, 'model_state_dict': model.state_dict(),175 'optimizer_state_dict': optimizer.state_dict(), 'scheduler_state_dict': scheduler.state_dict(),176 'val_acc': best_val_acc.item(), 'class_names': class_names}, save_path)177 print(f" โจ NEW BEST! Val Acc: {best_val_acc:.4f} ({best_val_acc*100:.2f}%) - Model Saved โจ")178 except Exception as save_e: print(f"โ ๏ธ Warning: Could not save model. Error: {save_e}")179 patience_counter = 0180 else:181 patience_counter += 1; print(f" (Patience: {patience_counter}/{PATIENCE})")182 if patience_counter >= PATIENCE:183 print(f"\nโณ Early stopping triggered at epoch {epoch+1}.")184 break185 186# --- 9. FINAL RESULTS ---187time_elapsed = time.time() - start_time188print("\n" + "=" * 80); print("๐ TRAINING COMPLETE! ๐"); print("=" * 80)189print(f"\nโ Total Training Time: {time_elapsed // 60:.0f}m {time_elapsed % 60:.0f}s")190print(f"๐ Best Validation Accuracy: {best_val_acc:.4f} ({best_val_acc*100:.2f}%)")191print(f"โ Best model saved to: {os.path.join(base_dir, MODEL_SAVE_PATH)}")192if best_val_acc >= 0.90: print("\n๐๐๐ EXCELLENT! Reached 90%+ target! Ready for Phase 3! ๐๐๐")193elif best_val_acc >= 0.85: print("\n๐ VERY GOOD! Achieved 85%+ accuracy! Strong candidate for Phase 3.")194elif best_val_acc >= 0.80: print("\n๐ SOLID RESULT! Achieved 80%+ accuracy.")195else: print(f"\nโ ๏ธ BELOW TARGET ({best_val_acc*100:.1f}%). Model might require more tuning or ResNet-50 was better.")196print("\n" + "=" * 80)197 