CoolFace
Modelpublic

kikogazda/Efficient_NetV2_Edition

sourceHugging Faceapache-2.0updated 1y agoView on Hugging Face
1likes
last_model.py741 linesDownload Raw Back to root
1# -*- coding: utf-8 -*-2"""Last_model.ipynb3 4Automatically generated by Colab.5 6Original file is located at7    https://colab.research.google.com/drive/1AdRILP1oqdiVuRSQr2dZZy0QgU8insn_8 9🚗 TwinCar Project: SOTA Training, Full Visuals, and Advanced Reporting10 11 12---13 14 15---16 171. Environment Setup and Imports18Explanation:19We start by importing all necessary libraries and prepping our working environment for advanced data handling and visualization.20 21---22"""23 24# Block 1: Environment Setup and Imports25import os26import zipfile27import numpy as np28import pandas as pd29import matplotlib.pyplot as plt30import seaborn as sns31from PIL import Image32from tqdm import tqdm33 34import torch35import torch.nn as nn36import torch.optim as optim37from torch.utils.data import Dataset, DataLoader, WeightedRandomSampler38from torchvision import transforms39 40from sklearn.model_selection import train_test_split41from sklearn.utils.class_weight import compute_class_weight42from sklearn.metrics import (43    accuracy_score, precision_score, recall_score, f1_score, hamming_loss,44    cohen_kappa_score, matthews_corrcoef, jaccard_score,45    confusion_matrix, classification_report46)47 48import timm49import scipy.io50 51"""2. Data Extraction and Preparation52Explanation:53We extract and organize the Stanford Cars dataset, parse .mat files to CSV for class and label mapping, and prepare all paths.54 55---56 57 58"""59 60# Block 2: Data Extraction and Preparation61from google.colab import drive62drive.mount('/content/drive')63 64zip_path = '/content/drive/MyDrive/stanford_cars.zip'65extract_dir = '/content/stanford_cars'66if not os.path.exists(extract_dir):67    with zipfile.ZipFile(zip_path, 'r') as zip_ref:68        zip_ref.extractall(extract_dir)69print("✅ Dataset extracted at", extract_dir)70 71meta = scipy.io.loadmat(f"{extract_dir}/car_devkit/devkit/cars_meta.mat")72class_names = [x[0] for x in meta['class_names'][0]]73NUM_CLASSES = len(class_names)74 75train_annos = scipy.io.loadmat(f"{extract_dir}/car_devkit/devkit/cars_train_annos.mat")['annotations'][0]76train_rows = [[x[5][0], int(x[4][0]) - 1] for x in train_annos]77df_train = pd.DataFrame(train_rows, columns=["filename", "label"])78df_train.to_csv('/content/train_labels.csv', index=False)79 80test_annos = scipy.io.loadmat(f"{extract_dir}/car_devkit/devkit/cars_test_annos.mat")['annotations'][0]81test_rows = [[x[4][0]] for x in test_annos]82df_test = pd.DataFrame(test_rows, columns=["filename"])83df_test.to_csv('/content/test_labels.csv', index=False)84 85train_root = f"{extract_dir}/cars_train/cars_train"86test_root = f"{extract_dir}/cars_test/cars_test"87 88"""3. Advanced Dataset and Augmentations89Explanation:90We build a flexible dataset class, apply advanced augmentations, and lay the foundation for Mixup/CutMix later.91 92---93 94 95"""96 97# Block 3: Dataset and Advanced Augmentations98 99class StanfordCarsFromCSV(Dataset):100    def __init__(self, root_dir, csv_file, transform=None, has_labels=True):101        self.root_dir = root_dir102        self.data = pd.read_csv(csv_file)103        self.transform = transform104        self.has_labels = has_labels105    def __len__(self):106        return len(self.data)107    def __getitem__(self, idx):108        row = self.data.iloc[idx]109        img_path = os.path.join(self.root_dir, row['filename'])110        image = Image.open(img_path).convert('RGB')111        if self.transform:112            image = self.transform(image)113        if self.has_labels:114            return image, int(row['label'])115        return image, row['filename']116 117imagenet_mean = [0.485, 0.456, 0.406]118imagenet_std = [0.229, 0.224, 0.225]119train_transform = transforms.Compose([120    transforms.RandomResizedCrop(224, scale=(0.7, 1.0)),121    transforms.RandomHorizontalFlip(),122    transforms.RandomRotation(15),123    transforms.ColorJitter(0.4, 0.4, 0.4, 0.2),124    transforms.RandomApply([transforms.GaussianBlur(3)], p=0.15),125    transforms.ToTensor(),126    transforms.Normalize(mean=imagenet_mean, std=imagenet_std)127])128val_transform = transforms.Compose([129    transforms.Resize(256),130    transforms.CenterCrop(224),131    transforms.ToTensor(),132    transforms.Normalize(mean=imagenet_mean, std=imagenet_std)133])134 135"""4. Data Splitting, Weighted Sampling, and DataLoader136Explanation:137We split the data into train and validation sets with stratification for balanced classes,138use class weighting to counter imbalance, and create PyTorch DataLoaders for efficient training and evaluation.139 140---141 142 143"""144 145# 4. Data Splitting, Loader Setup, and Weighted Sampling146 147from torch.utils.data import DataLoader, WeightedRandomSampler148from sklearn.model_selection import train_test_split149from sklearn.utils.class_weight import compute_class_weight150 151# --- Settings ---152BATCH_SIZE = 32153VAL_RATIO = 0.1154RANDOM_SEED = 42155 156# --- Stratified Split for Balanced Classes ---157df_all = pd.read_csv('/content/train_labels.csv')158df_train, df_val = train_test_split(159    df_all,160    test_size=VAL_RATIO,161    stratify=df_all['label'],162    random_state=RANDOM_SEED163)164df_train.to_csv('/content/train_split.csv', index=False)165df_val.to_csv('/content/val_split.csv', index=False)166 167# --- Datasets ---168train_dataset = StanfordCarsFromCSV(train_root, '/content/train_split.csv', train_transform)169val_dataset = StanfordCarsFromCSV(train_root, '/content/val_split.csv', val_transform)170test_dataset = StanfordCarsFromCSV(test_root, '/content/test_labels.csv', val_transform, has_labels=False)171 172# --- Weighted Sampler for Balanced Training ---173labels = [label for _, label in train_dataset]174class_weights = compute_class_weight(class_weight='balanced', classes=np.unique(labels), y=labels)175sample_weights = [class_weights[label] for label in labels]176sampler = WeightedRandomSampler(sample_weights, len(sample_weights), replacement=True)177 178# --- DataLoaders (drop_last=True for Mixup/CutMix compatibility) ---179train_loader = DataLoader(180    train_dataset,181    batch_size=BATCH_SIZE,182    sampler=sampler,183    num_workers=2,184    pin_memory=True,185    drop_last=True186)187val_loader = DataLoader(188    val_dataset,189    batch_size=BATCH_SIZE,190    shuffle=False,191    num_workers=2,192    pin_memory=True,193    drop_last=False194)195test_loader = DataLoader(196    test_dataset,197    batch_size=BATCH_SIZE,198    shuffle=False,199    num_workers=2,200    pin_memory=True,201    drop_last=False202)203 204print(f"Train samples: {len(train_dataset)} | Val samples: {len(val_dataset)} | Test samples: {len(test_dataset)}")205print(f"Train loader batches (per epoch): {len(train_loader)} (should be integer and even-sized)")206 207"""5. Model Initialization: EfficientNetV2 + Mixup/CutMix Ready208Explanation:209We load EfficientNetV2 with ImageNet weights for best transfer learning,210set up optimizer, scheduler, and prepare for Mixup/CutMix advanced augmentation.211 212---213 214 215"""216 217# Block 5: Model Initialization (EfficientNetV2 + Mixup/CutMix)218 219from timm.data import Mixup220 221device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')222model = timm.create_model('efficientnetv2_rw_s', pretrained=True, num_classes=NUM_CLASSES, drop_rate=0.3)223model = model.to(device)224 225optimizer = optim.AdamW(model.parameters(), lr=3e-4, weight_decay=1e-5)226scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=25)227criterion = nn.CrossEntropyLoss(label_smoothing=0.0)228 229mixup_fn = Mixup(230    mixup_alpha=0.4, cutmix_alpha=1.0, cutmix_minmax=None,231    prob=1.0, switch_prob=0.5, mode='batch',232    label_smoothing=0.1, num_classes=NUM_CLASSES233)234 235"""6. Advanced Training Loop: Full Metrics, Early Stopping, and Mixup236Explanation:237This loop supports Mixup/CutMix, logs all advanced metrics, and uses early stopping with automatic best model saving.238Ready for real production—and all your plots and reporting.239 240---241 242 243"""244 245# Block 6: Advanced Training Loop246 247EPOCHS = 25248patience, counter = 7, 0249best_val_f1 = 0250 251metrics_dict = {252    'train_loss': [], 'train_acc': [],253    'val_loss': [], 'val_acc': [],254    'val_precision_macro': [], 'val_precision_weighted': [],255    'val_recall_macro': [], 'val_recall_weighted': [],256    'val_f1_macro': [], 'val_f1_weighted': [],257    'val_hamming': [], 'val_cohen_kappa': [],258    'val_mcc': [], 'val_jaccard_macro': [],259    'val_top3': [], 'val_top5': [],260}261 262for epoch in range(EPOCHS):263    # TRAIN264    model.train()265    total_loss, correct, total = 0, 0, 0266    for imgs, labels in tqdm(train_loader, desc=f"Train Epoch {epoch+1}"):267        imgs, labels = imgs.to(device), labels.to(device)268        optimizer.zero_grad()269        imgs, labels = mixup_fn(imgs, labels)270        outputs = model(imgs)271        loss = criterion(outputs, labels)272        loss.backward()273        optimizer.step()274        total_loss += loss.item() * imgs.size(0)275        correct += (outputs.argmax(1) == labels.argmax(1)).sum().item()276        total += labels.size(0)277    train_loss = total_loss / total278    train_acc = correct / total279    metrics_dict['train_loss'].append(train_loss)280    metrics_dict['train_acc'].append(train_acc)281 282    # VALIDATION283    model.eval()284    val_loss, val_correct, val_total = 0, 0, 0285    val_probs, val_preds, val_targets = [], [], []286    with torch.no_grad():287        for imgs, labels in tqdm(val_loader, desc=f"Val Epoch {epoch+1}"):288            imgs, labels = imgs.to(device), labels.to(device)289            outputs = model(imgs)290            v_loss = criterion(outputs, labels)291            val_loss += v_loss.item() * imgs.size(0)292            probs = torch.softmax(outputs, dim=1)293            preds = outputs.argmax(1)294            val_correct += (preds == labels).sum().item()295            val_total += labels.size(0)296            val_probs.extend(probs.cpu().numpy())297            val_preds.extend(preds.cpu().numpy())298            val_targets.extend(labels.cpu().numpy())299    val_loss /= val_total300    val_acc = val_correct / val_total301    val_preds_np = np.array(val_preds)302    val_targets_np = np.array(val_targets)303    val_probs_np = np.array(val_probs)304 305    # Metrics306    val_precision_macro = precision_score(val_targets_np, val_preds_np, average='macro', zero_division=0)307    val_precision_weighted = precision_score(val_targets_np, val_preds_np, average='weighted', zero_division=0)308    val_recall_macro = recall_score(val_targets_np, val_preds_np, average='macro', zero_division=0)309    val_recall_weighted = recall_score(val_targets_np, val_preds_np, average='weighted', zero_division=0)310    val_f1_macro = f1_score(val_targets_np, val_preds_np, average='macro', zero_division=0)311    val_f1_weighted = f1_score(val_targets_np, val_preds_np, average='weighted', zero_division=0)312    top3_acc = np.mean([313        label in np.argsort(prob)[-3:] for prob, label in zip(val_probs_np, val_targets_np)314    ])315    top5_acc = np.mean([316        label in np.argsort(prob)[-5:] for prob, label in zip(val_probs_np, val_targets_np)317    ])318    val_hamming = hamming_loss(val_targets_np, val_preds_np)319    val_cohen_kappa = cohen_kappa_score(val_targets_np, val_preds_np)320    val_mcc = matthews_corrcoef(val_targets_np, val_preds_np)321    val_jaccard_macro = jaccard_score(val_targets_np, val_preds_np, average='macro', zero_division=0)322 323    # Log metrics324    metrics_dict['val_loss'].append(val_loss)325    metrics_dict['val_acc'].append(val_acc)326    metrics_dict['val_precision_macro'].append(val_precision_macro)327    metrics_dict['val_precision_weighted'].append(val_precision_weighted)328    metrics_dict['val_recall_macro'].append(val_recall_macro)329    metrics_dict['val_recall_weighted'].append(val_recall_weighted)330    metrics_dict['val_f1_macro'].append(val_f1_macro)331    metrics_dict['val_f1_weighted'].append(val_f1_weighted)332    metrics_dict['val_hamming'].append(val_hamming)333    metrics_dict['val_cohen_kappa'].append(val_cohen_kappa)334    metrics_dict['val_mcc'].append(val_mcc)335    metrics_dict['val_jaccard_macro'].append(val_jaccard_macro)336    metrics_dict['val_top3'].append(top3_acc)337    metrics_dict['val_top5'].append(top5_acc)338 339    scheduler.step()340    print(f"Epoch {epoch+1:2d} | Train Acc: {train_acc:.4f} | Val Acc: {val_acc:.4f} | F1(macro): {val_f1_macro:.4f} | Top3: {top3_acc:.3f} | Top5: {top5_acc:.3f}")341 342    # Early Stopping343    if val_f1_macro > best_val_f1:344        best_val_f1 = val_f1_macro345        torch.save(model.state_dict(), '/content/drive/MyDrive/efficientnetv2_best_model.pth')346        counter = 0347    else:348        counter += 1349        if counter >= patience:350            print("⏹️ Early stopping triggered.")351            break352 353print("✅ Training complete. Best model saved.")354 355"""7.Explanation356After training, all metrics (accuracy, loss, precision, recall, F1, top-k, etc.) are saved as a CSV for analysis and reporting.357 358We plot core metrics (accuracy, F1, loss, precision/recall, top-3/top-5 accuracy) with:359 360Large, clear fonts361 362Annotations for best epoch363 364Colorful, pro-style Seaborn plots365 366Publication-ready grid and tight layouts367 368---369 370 371"""372 373# 7. Metrics Export & Advanced Visualizations374 375import seaborn as sns376 377# --- Save all metrics for reproducibility and later analysis378metrics_df = pd.DataFrame(metrics_dict)379metrics_df.to_csv('/content/drive/MyDrive/metrics_log.csv', index_label='epoch')380print("✅ metrics_log.csv saved.")381 382sns.set(style='whitegrid', font_scale=1.3)383 384# 1. Accuracy & Macro F1385plt.figure(figsize=(12,7))386plt.plot(metrics_df['train_acc'], label='Train Acc', lw=2)387plt.plot(metrics_df['val_acc'], label='Val Acc', lw=2)388plt.plot(metrics_df['val_f1_macro'], label='Val F1 (macro)', lw=2)389plt.xlabel('Epoch', fontsize=16)390plt.ylabel('Score', fontsize=16)391plt.title('Accuracy and Macro F1 per Epoch', fontsize=18)392plt.legend(loc='lower right')393plt.grid(True, alpha=0.3)394best_epoch = metrics_df['val_f1_macro'].idxmax()395plt.scatter(best_epoch, metrics_df['val_f1_macro'][best_epoch], c='red', s=90, label='Best Epoch')396plt.annotate(f'Best\n{metrics_df["val_f1_macro"][best_epoch]:.2f}',397             (best_epoch, metrics_df["val_f1_macro"][best_epoch]),398             textcoords="offset points", xytext=(-5,10), ha='right', fontsize=14, color='red')399plt.tight_layout()400plt.savefig('/content/drive/MyDrive/metrics_acc_f1_beautiful.png')401plt.show()402 403# 2. Loss Curves404plt.figure(figsize=(12,7))405plt.plot(metrics_df['train_loss'], label='Train Loss', lw=2)406plt.plot(metrics_df['val_loss'], label='Val Loss', lw=2)407plt.xlabel('Epoch', fontsize=16)408plt.ylabel('Loss', fontsize=16)409plt.title('Train & Validation Loss per Epoch', fontsize=18)410plt.legend(loc='upper right')411plt.grid(True, alpha=0.3)412plt.tight_layout()413plt.savefig('/content/drive/MyDrive/metrics_loss_beautiful.png')414plt.show()415 416# 3. Precision & Recall (Macro & Weighted)417plt.figure(figsize=(12,7))418plt.plot(metrics_df['val_precision_macro'], label='Val Precision (macro)', lw=2)419plt.plot(metrics_df['val_recall_macro'], label='Val Recall (macro)', lw=2)420plt.plot(metrics_df['val_precision_weighted'], label='Val Precision (weighted)', lw=2)421plt.plot(metrics_df['val_recall_weighted'], label='Val Recall (weighted)', lw=2)422plt.xlabel('Epoch', fontsize=16)423plt.ylabel('Score', fontsize=16)424plt.title('Validation Precision & Recall per Epoch', fontsize=18)425plt.legend(loc='lower right')426plt.grid(True, alpha=0.3)427plt.tight_layout()428plt.savefig('/content/drive/MyDrive/metrics_precision_recall_beautiful.png')429plt.show()430 431# 4. Top-3 and Top-5 Validation Accuracy as Area Plot432plt.figure(figsize=(12,7))433plt.fill_between(metrics_df.index, metrics_df['val_top3'], alpha=0.3, label='Val Top-3 Acc')434plt.fill_between(metrics_df.index, metrics_df['val_top5'], alpha=0.2, label='Val Top-5 Acc', color='orange')435plt.plot(metrics_df['val_top3'], lw=2, color='blue')436plt.plot(metrics_df['val_top5'], lw=2, color='orange')437plt.xlabel('Epoch', fontsize=16)438plt.ylabel('Accuracy', fontsize=16)439plt.title('Top-3 and Top-5 Validation Accuracy per Epoch', fontsize=18)440plt.legend(loc='lower right')441plt.grid(True, alpha=0.3)442plt.tight_layout()443plt.savefig('/content/drive/MyDrive/metrics_topk_beautiful.png')444plt.show()445 446"""8.Confusion Matrix & Per-Class Analysis with Advanced Visuals447Explanation448After training, it's crucial to understand not just overall metrics, but where your model succeeds and fails.449We:450 451Save a detailed classification report (per-class precision/recall/F1).452 453Draw a high-contrast confusion matrix with large ticks, tight color scaling, and readable value overlays.454 455Plot Top 20 Most Confused Classes for targeted debugging.456 457Show Top 20 Most Accurate Classes with horizontal barplots (values on bars, sorted).458 459 460 461---462 463 464"""465 466# 8. Confusion Matrix & Per-Class Analysis (Advanced Visuals)467 468from sklearn.metrics import classification_report, confusion_matrix469import seaborn as sns470 471# Reload best model for evaluation472model.load_state_dict(torch.load('/content/drive/MyDrive/efficientnetv2_best_model.pth', map_location=device))473model.eval()474 475# Collect all validation predictions and true labels476all_preds, all_labels = [], []477with torch.no_grad():478    for imgs, labels in val_loader:479        imgs, labels = imgs.to(device), labels.to(device)480        outputs = model(imgs)481        preds = outputs.argmax(1)482        all_preds.extend(preds.cpu().numpy())483        all_labels.extend(labels.cpu().numpy())484all_preds = np.array(all_preds)485all_labels = np.array(all_labels)486 487# Save detailed classification report (per-class)488report = classification_report(489    all_labels, all_preds, target_names=class_names, output_dict=True490)491pd.DataFrame(report).transpose().to_csv('/content/drive/MyDrive/classification_report.csv')492print("✅ classification_report.csv saved.")493 494# Confusion Matrix (full, high-res)495cm = confusion_matrix(all_labels, all_preds)496plt.figure(figsize=(18,18))497sns.heatmap(498    cm,499    cmap="Blues",500    xticklabels=class_names,501    yticklabels=class_names,502    square=True,503    cbar_kws={"shrink": 0.5, "label": "Count"},504    linewidths=.2505)506plt.title('Confusion Matrix', fontsize=20)507plt.xlabel('Predicted label', fontsize=16)508plt.ylabel('True label', fontsize=16)509plt.xticks(fontsize=8, rotation=90)510plt.yticks(fontsize=8)511plt.tight_layout()512plt.savefig('/content/drive/MyDrive/confusion_matrix_beautiful.png', dpi=300)513plt.show()514 515# Most Confused Classes (Top 20, value overlays)516off_diag = cm.copy()517np.fill_diagonal(off_diag, 0)518most_confused = np.argsort(off_diag.sum(axis=1))[::-1][:20]519cm_top = cm[np.ix_(most_confused, most_confused)]520labels_top = [class_names[i] for i in most_confused]521 522plt.figure(figsize=(12,10))523sns.heatmap(524    cm_top,525    annot=True, fmt='d', cmap="Blues",526    xticklabels=labels_top, yticklabels=labels_top,527    linewidths=.2, cbar=False, annot_kws={"size":14}528)529plt.title('Most Confused Classes (Top 20)', fontsize=18)530plt.xlabel('Predicted label', fontsize=15)531plt.ylabel('True label', fontsize=15)532plt.xticks(fontsize=11, rotation=90)533plt.yticks(fontsize=11)534plt.tight_layout()535plt.savefig('/content/drive/MyDrive/confused_top20_beautiful.png', dpi=300)536plt.show()537 538# Top-20 Most Accurate Classes (barplot, values on bars)539acc_per_class = cm.diagonal() / (cm.sum(axis=1) + 1e-8)540df_acc = pd.DataFrame({'class': class_names, 'accuracy': acc_per_class})541top_acc = df_acc.sort_values('accuracy', ascending=False).head(20)542plt.figure(figsize=(10,8))543sns.barplot(544    data=top_acc, y='class', x='accuracy', palette='Blues_d', orient='h'545)546plt.title('Top 20 Classes by Accuracy', fontsize=18)547plt.xlabel('Accuracy', fontsize=15)548plt.ylabel('Class', fontsize=15)549for i, v in enumerate(top_acc['accuracy']):550    plt.text(v + 0.01, i, f"{v:.2f}", color='blue', va='center', fontsize=13)551plt.tight_layout()552plt.savefig('/content/drive/MyDrive/top20_accuracy_beautiful.png', dpi=300)553plt.show()554 555"""9. Test-Time Augmentation (TTA) & Batch Prediction556Explanation557Test-Time Augmentation boosts prediction robustness by averaging predictions over multiple random transformations of each test image.558Batch Prediction allows you to efficiently label a folder of test images with class names—production style.559"""560 561# 9. Test-Time Augmentation (TTA) for Validation562 563tta_transforms = [564    val_transform,565    transforms.Compose([566        transforms.Resize(256),567        transforms.RandomHorizontalFlip(p=1.0),568        transforms.CenterCrop(224),569        transforms.ToTensor(),570        transforms.Normalize(mean=imagenet_mean, std=imagenet_std)571    ]),572    transforms.Compose([573        transforms.Resize(256),574        transforms.RandomRotation(10),575        transforms.CenterCrop(224),576        transforms.ToTensor(),577        transforms.Normalize(mean=imagenet_mean, std=imagenet_std)578    ]),579    transforms.Compose([580        transforms.Resize(256),581        transforms.ColorJitter(0.2, 0.2, 0.2, 0.1),582        transforms.CenterCrop(224),583        transforms.ToTensor(),584        transforms.Normalize(mean=imagenet_mean, std=imagenet_std)585    ])586]587 588def tta_predict(model, img_pil, tta_transforms, device='cuda'):589    model.eval()590    logits = []591    for tform in tta_transforms:592        img = tform(img_pil).unsqueeze(0).to(device)593        with torch.no_grad():594            logit = model(img)595            logits.append(logit)596    avg_logits = torch.stack(logits).mean(0)597    return avg_logits598 599# Apply TTA to validation set600tta_val_preds, tta_val_labels = [], []601for imgs, labels in tqdm(val_loader, desc="TTA Validation"):602    batch_preds = []603    for i in range(imgs.size(0)):604        img_pil = transforms.ToPILImage()(imgs[i].cpu())605        avg_logits = tta_predict(model, img_pil, tta_transforms, device)606        pred = avg_logits.argmax(dim=1).cpu().item()607        batch_preds.append(pred)608    tta_val_preds.extend(batch_preds)609    tta_val_labels.extend(labels.cpu().numpy())610 611tta_val_preds = np.array(tta_val_preds)612tta_val_labels = np.array(tta_val_labels)613 614# Metrics for TTA615tta_f1_macro = f1_score(tta_val_labels, tta_val_preds, average='macro', zero_division=0)616tta_acc = accuracy_score(tta_val_labels, tta_val_preds)617tta_precision = precision_score(tta_val_labels, tta_val_preds, average='macro', zero_division=0)618tta_recall = recall_score(tta_val_labels, tta_val_preds, average='macro', zero_division=0)619print(f"TTA Validation Accuracy: {tta_acc:.4f}")620print(f"TTA Validation F1 (macro): {tta_f1_macro:.4f}")621print(f"TTA Validation Precision (macro): {tta_precision:.4f}")622print(f"TTA Validation Recall (macro): {tta_recall:.4f}")623 624# TTA Confusion matrix (optional)625cm_tta = confusion_matrix(tta_val_labels, tta_val_preds)626plt.figure(figsize=(18,18))627sns.heatmap(628    cm_tta,629    cmap="Purples",630    xticklabels=class_names,631    yticklabels=class_names,632    square=True,633    cbar_kws={"shrink": 0.5, "label": "Count"},634    linewidths=.2635)636plt.title('TTA Confusion Matrix (Validation)', fontsize=20)637plt.xlabel('Predicted label', fontsize=16)638plt.ylabel('True label', fontsize=16)639plt.xticks(fontsize=8, rotation=90)640plt.yticks(fontsize=8)641plt.tight_layout()642plt.savefig('/content/drive/MyDrive/tta_confusion_matrix_beautiful.png', dpi=300)643plt.show()644 645"""10. Extraordinary Grad-CAM++ Overlays (Grid)646Explanation647We generate Grad-CAM++ visualizations for a set of sample images.648Each visualization shows:The input image,The Grad-CAM++ heatmap overlay,The true and predicted class for easy comparison.649All visualizations are saved both individually and as a large, labeled grid.650 651 652 653---654"""655 656# Grad-CAM++ Explanations: Multi-Image Grid (Fixed for latest grad-cam)657 658from pytorch_grad_cam import GradCAMPlusPlus659from pytorch_grad_cam.utils.image import show_cam_on_image660from pytorch_grad_cam.utils.model_targets import ClassifierOutputTarget661 662import os663 664os.makedirs('/content/drive/MyDrive/gradcam_outputs', exist_ok=True)665 666# Make sure model is on the right device667model.eval()668model.to(device)669 670# Pick the right target layer for EfficientNetV2 (last block)671target_layer = model.blocks[-1] if hasattr(model, "blocks") else model.layer4[-1]672 673# No more use_cuda argument—just instantiate674cam = GradCAMPlusPlus(model=model, target_layers=[target_layer])675 676num_images = 12677fig, axes = plt.subplots(3, 4, figsize=(18, 14))678fig.suptitle('Grad-CAM++ Explanations: True vs. Predicted', fontsize=22, weight='bold')679 680for idx in range(num_images):681    img_tensor, label = val_dataset[idx]682    img_pil = transforms.ToPILImage()(img_tensor.cpu())683    input_tensor = img_tensor.unsqueeze(0).to(device)684    with torch.no_grad():685        output = model(input_tensor)686        pred = output.argmax(1).item()687    targets = [ClassifierOutputTarget(pred)]688    grayscale_cam = cam(input_tensor=input_tensor, targets=targets)[0]689    image_np = img_tensor.permute(1, 2, 0).cpu().numpy()690    image_np = (image_np * np.array(imagenet_std)) + np.array(imagenet_mean)691    image_np = np.clip(image_np, 0, 1)692    cam_image = show_cam_on_image(image_np, grayscale_cam, use_rgb=True)693 694    # Save each Grad-CAM overlay individually695    overlay_path = f"/content/drive/MyDrive/gradcam_outputs/val_{idx}_true_{class_names[label]}_pred_{class_names[pred]}.png"696    plt.imsave(overlay_path, cam_image)697 698    # Add to grid699    ax = axes[idx // 4, idx % 4]700    ax.imshow(cam_image)701    ax.set_title(702        f"True: {class_names[label][:18]}\nPred: {class_names[pred][:18]}",703        fontsize=12,704        color="green" if pred == label else "red",705        weight="bold"706    )707    ax.axis('off')708 709plt.tight_layout(rect=[0, 0.03, 1, 0.95])710plt.savefig('/content/drive/MyDrive/gradcam_outputs/gradcam_grid.png', dpi=250)711plt.show()712 713"""11. Gradio Interactive Demo: Model + Grad-CAM++"""714 715# 11. Gradio Interactive Demo: EfficientNetV2 + Grad-CAM++716 717import gradio as gr718from PIL import Image as PILImage719 720def predict_and_explain(img):721    image_pil = img.convert("RGB").resize((224, 224))722    input_tensor = val_transform(image_pil).unsqueeze(0).to(device)723    with torch.no_grad():724        output = model(input_tensor)725        pred_idx = output.argmax().item()726    targets = [ClassifierOutputTarget(pred_idx)]727    grayscale_cam = cam(input_tensor=input_tensor, targets=targets)[0]728    image_np = np.array(image_pil).astype(np.float32) / 255.0729    cam_image = show_cam_on_image(image_np, grayscale_cam, use_rgb=True)730    pred_name = class_names[pred_idx]731    return PILImage.fromarray(cam_image), f"Prediction: {pred_name} (class index {pred_idx})"732 733demo = gr.Interface(734    fn=predict_and_explain,735    inputs=gr.Image(type="pil", label="Upload Car Image"),736    outputs=[gr.Image(label="Grad-CAM++ Output"), gr.Text(label="Prediction")],737    title="🚗 TwinCar: Car Make/Model Classifier + Explainability Demo",738    description="Upload a car photo. See the prediction (make/model/year) and a Grad-CAM++ heatmap showing what influenced the model.",739    allow_flagging='never'740)741demo.launch(share=True)