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("๐ฅ ADVANCED TRAINING - Target: 85%+ ๐ฅ")14print("=" * 80)15 16# Configuration17EPOCHS = 10018WARMUP_EPOCHS = 10 # Warmup period19LEARNING_RATE = 0.000520MIN_LR = 1e-621BATCH_SIZE = 20 # Smaller for better gradient estimates22DATA_DIR = 'master_dataset'23MODEL_SAVE_PATH = 'thermal_model_advanced.pth'24PATIENCE = 2525SEED = 4226 27# Setup28torch.manual_seed(SEED)29np.random.seed(SEED)30random.seed(SEED)31 32device = torch.device("cuda" if torch.cuda.is_available() else "cpu")33print(f"\nโ Device: {device}")34 35base_dir = os.path.abspath(os.path.join(os.path.dirname(__file__), '..'))36data_path = os.path.join(base_dir, DATA_DIR)37 38if not os.path.exists(data_path):39 print(f"โ ERROR: {data_path} not found!")40 exit(1)41 42# Augmentation - Even Stronger43train_transforms = transforms.Compose([44 transforms.Resize((256, 256)),45 transforms.RandomResizedCrop(224, scale=(0.65, 1.0)), # More aggressive46 transforms.RandomHorizontalFlip(p=0.5),47 transforms.RandomVerticalFlip(p=0.5),48 transforms.RandomRotation(45), # Increased49 transforms.ColorJitter(brightness=0.5, contrast=0.5, saturation=0.4, hue=0.15),50 transforms.RandomAffine(degrees=0, translate=(0.2, 0.2), scale=(0.85, 1.15)),51 transforms.RandomPerspective(distortion_scale=0.3, p=0.5),52 transforms.ToTensor(),53 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),54 transforms.RandomErasing(p=0.3, scale=(0.02, 0.2)) # Added55])56 57val_transforms = transforms.Compose([58 transforms.Resize((224, 224)),59 transforms.ToTensor(),60 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])61])62 63# Load dataset64full_dataset = datasets.ImageFolder(data_path)65class_names = full_dataset.classes66num_classes = len(class_names)67 68print(f"โ Total: {len(full_dataset)} images, {num_classes} classes")69 70# Split71train_size = int(0.8 * len(full_dataset))72val_size = len(full_dataset) - train_size73train_split, val_split = random_split(full_dataset, [train_size, val_size],74 generator=torch.Generator().manual_seed(SEED))75 76# Apply transforms77class TransformedDataset(torch.utils.data.Dataset):78 def __init__(self, subset, transform=None):79 self.subset = subset80 self.transform = transform81 def __getitem__(self, index):82 x, y = self.subset[index]83 if self.transform:84 x = self.transform(x)85 return x, y86 def __len__(self):87 return len(self.subset)88 89train_dataset = TransformedDataset(train_split, train_transforms)90val_dataset = TransformedDataset(val_split, val_transforms)91 92# Weighted sampler93train_indices = train_split.indices94train_targets = [full_dataset.targets[i] for i in train_indices]95class_counts = np.array([len(np.where(np.array(train_targets) == t)[0]) for t in np.unique(train_targets)])96 97print("\nClass distribution (train):")98for i, count in enumerate(class_counts):99 print(f" {class_names[i]:15s}: {count:5d}")100 101weight = 1. / class_counts102sample_weights = np.array([weight[t] for t in train_targets])103sample_weights = torch.from_numpy(sample_weights).double()104 105sampler = WeightedRandomSampler(sample_weights, len(sample_weights), replacement=True)106 107train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, sampler=sampler,108 num_workers=4, pin_memory=True, persistent_workers=True)109val_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False,110 num_workers=4, pin_memory=True, persistent_workers=True)111 112# Model - Unfreeze MORE layers113model = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V2)114 115for param in model.parameters():116 param.requires_grad = False117 118# Unfreeze layer2, layer3, layer4 (DEEPER fine-tuning)119for param in model.layer2.parameters():120 param.requires_grad = True121for param in model.layer3.parameters():122 param.requires_grad = True123for param in model.layer4.parameters():124 param.requires_grad = True125 126# Enhanced classifier127num_ftrs = model.fc.in_features128model.fc = nn.Sequential(129 nn.Dropout(0.6),130 nn.Linear(num_ftrs, 512),131 nn.ReLU(),132 nn.BatchNorm1d(512),133 nn.Dropout(0.5),134 nn.Linear(512, 256),135 nn.ReLU(),136 nn.BatchNorm1d(256),137 nn.Dropout(0.4),138 nn.Linear(256, num_classes)139)140 141model = model.to(device)142 143trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)144total = sum(p.numel() for p in model.parameters())145print(f"\nโ Trainable: {trainable:,} / {total:,} params")146 147# Loss148criterion = nn.CrossEntropyLoss(label_smoothing=0.15) # Increased smoothing149 150# Optimizer with discriminative learning rates151optimizer = optim.AdamW([152 {'params': model.layer2.parameters(), 'lr': LEARNING_RATE / 30},153 {'params': model.layer3.parameters(), 'lr': LEARNING_RATE / 15},154 {'params': model.layer4.parameters(), 'lr': LEARNING_RATE / 10},155 {'params': model.fc.parameters(), 'lr': LEARNING_RATE}156], weight_decay=0.02)157 158# Warmup + Cosine scheduler159def get_lr_lambda(epoch):160 if epoch < WARMUP_EPOCHS:161 return (epoch + 1) / WARMUP_EPOCHS # Linear warmup162 else:163 # Cosine annealing after warmup164 progress = (epoch - WARMUP_EPOCHS) / (EPOCHS - WARMUP_EPOCHS)165 return 0.5 * (1 + np.cos(np.pi * progress))166 167scheduler = optim.lr_scheduler.LambdaLR(optimizer, lr_lambda=get_lr_lambda)168 169print(f"\nโ Loss: CrossEntropyLoss (smoothing=0.15)")170print(f"โ Optimizer: AdamW (discriminative LR)")171print(f"โ Scheduler: Warmup ({WARMUP_EPOCHS}) + CosineAnnealing")172 173# Training174print("\n" + "=" * 80)175print(f"๐ TRAINING FOR {EPOCHS} EPOCHS ๐")176print("=" * 80 + "\n")177 178best_val_acc = 0.0179patience_counter = 0180start_time = time.time()181 182scaler = torch.cuda.amp.GradScaler(enabled=torch.cuda.is_available())183 184for epoch in range(EPOCHS):185 # Train186 model.train()187 running_loss = 0.0188 correct = 0189 total_samples = 0190 191 for inputs, labels in train_loader:192 inputs, labels = inputs.to(device, non_blocking=True), labels.to(device, non_blocking=True)193 194 with torch.cuda.amp.autocast(enabled=torch.cuda.is_available()):195 outputs = model(inputs)196 loss = criterion(outputs, labels)197 198 scaler.scale(loss).backward()199 scaler.step(optimizer)200 scaler.update()201 optimizer.zero_grad(set_to_none=True)202 203 running_loss += loss.item() * inputs.size(0)204 _, preds = torch.max(outputs, 1)205 correct += (preds == labels).sum().item()206 total_samples += labels.size(0)207 208 train_loss = running_loss / total_samples209 train_acc = correct / total_samples210 211 # Validate212 model.eval()213 val_loss = 0.0214 correct = 0215 total_samples = 0216 217 with torch.no_grad():218 for inputs, labels in val_loader:219 inputs, labels = inputs.to(device, non_blocking=True), labels.to(device, non_blocking=True)220 221 with torch.cuda.amp.autocast(enabled=torch.cuda.is_available()):222 outputs = model(inputs)223 loss = criterion(outputs, labels)224 225 val_loss += loss.item() * inputs.size(0)226 _, preds = torch.max(outputs, 1)227 correct += (preds == labels).sum().item()228 total_samples += labels.size(0)229 230 val_loss = val_loss / total_samples231 val_acc = correct / total_samples232 233 current_lr = optimizer.param_groups[-1]['lr'] # Get FC layer LR234 gap = train_acc - val_acc235 236 print(f"Epoch {epoch+1:03d}/{EPOCHS} | "237 f"Loss: {train_loss:.4f}/{val_loss:.4f} | "238 f"Acc: {train_acc:.4f}/{val_acc:.4f} | "239 f"Gap: {gap:.4f} | LR: {current_lr:.6f}")240 241 scheduler.step()242 243 if val_acc > best_val_acc:244 best_val_acc = val_acc245 torch.save({246 'epoch': epoch,247 'model_state_dict': model.state_dict(),248 'val_acc': val_acc,249 'train_acc': train_acc,250 'class_names': class_names251 }, os.path.join(base_dir, MODEL_SAVE_PATH))252 print(f" โจ BEST: {val_acc:.4f} ({val_acc*100:.2f}%) โจ")253 patience_counter = 0254 else:255 patience_counter += 1256 257 if patience_counter >= PATIENCE:258 print(f"\nโณ Early stop at epoch {epoch+1}")259 break260 261time_elapsed = time.time() - start_time262 263print("\n" + "=" * 80)264print("๐ COMPLETE!")265print("=" * 80)266print(f"\nโ Time: {time_elapsed//60:.0f}m {time_elapsed%60:.0f}s")267print(f"๐ Best: {best_val_acc:.4f} ({best_val_acc*100:.2f}%)")268 269if best_val_acc >= 0.85:270 print("\n๐ SUCCESS! 85%+ achieved!")271elif best_val_acc >= 0.82:272 print("\nโ
VERY CLOSE! Consider Approach 3 (Ensemble)")273else:274 print("\nโ ๏ธ Need more data or different architecture")275 