CoolFace
Apppublic

Timechils/sapiens-mvp

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