kikogazda/Efficient_NetV2_Edition
1
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)