CoolFace
Apppublic

Timechils/sapiens-mvp

sourceHugging Faceupdated 10mo agoView on Hugging Face
0likes
train.py197 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("๐Ÿš€ 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