OSCR

Multi-class classification of brain tumor using a ResNet101 backbone integrated with multi-scale deformable attention module and advanced data augmentations.

Code ↔ Paper

5 matches between paragraphs of the paper and lines of its authors' code, computed by the harvester (lexical-v1). Click a colored paragraph or line to see its counterpart.

The 5 matches
  1. [1] § Results and discussion › Model training, validation and test performance ↔ train-val-test-04-11-2025-16-5.ipynb, lines 472–617 · score 0.73 · cosine annealing learning, rate scheduler, optimization, weights, PyTorch, MixUp
  2. [2] § Results and discussion › Performance evaluation of Precision–Recall, Receiver Operating Characteristic (ROC) curve and AUC & t-SNE feature visualization ↔ train-val-test-04-11-2025-16-5.ipynb, lines 1781–1833 · score 0.56 · PR curves, Precision Recall, ROC, AUC, classifier, model
  3. [3] § Results and discussion › Quantitative evaluation metrics ↔ train-val-test-04-11-2025-16-5.ipynb, lines 1247–1396 · score 0.52 · Cohen Kappa, Macro, Dice, metric, Recall, score
  4. [4] § Methodology › Model architecture › Multi-scale deformable attention module (MS-DAM) ↔ train-val-test-04-11-2025-16-5.ipynb, lines 281–355 · score 0.51 · MS DAM, kernel, bilinear, adaptively, channel, fuse
  5. [5] § Methodology › Dataset preparation › Dataset scanning and deduplication ↔ train-val-test-04-11-2025-16-5.ipynb, lines 103–139 · score 0.50 · brain tumor, download, kagglehub, waseemnagahhenes, MRI, classes

Paper

Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC

The paper is loaded when this pane is shown.

The authors' code

Jupyter notebook · 2,512 lines · 92 KB · no license · 5 matches

  1. # %%
  2. # This Python 3 environment comes with many helpful analytics libraries installed
  3. # It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python
  4. # For example, here's several helpful packages to load
  5. import numpy as np # linear algebra
  6. import pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)
  7. # Input data files are available in the read-only "../input/" directory
  8. # For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory
  9. import os
  10. for dirname, _, filenames in os.walk('/kaggle/input'):
  11. for filename in filenames:
  12. print(os.path.join(dirname, filename))
  13. # You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using "Save & Run All"
  14. # You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session
  15. # %%
  16. !pip install grad-cam
  17. import pytorch_grad_cam
  18. print("Grad-CAM is installed and working!")
  19. !pip install "numpy<2.0"
  20. # %%
  21. # ======================== SAFE ENVIRONMENT SETUP ========================
  22. import warnings
  23. warnings.filterwarnings("ignore", category=FutureWarning)
  24. import os
  25. import math
  26. import random
  27. import copy
  28. import time
  29. import hashlib
  30. import zipfile
  31. from pathlib import Path
  32. # Limit OpenBLAS / MKL / OMP threads to prevent warnings/hangs
  33. os.environ["OPENBLAS_NUM_THREADS"] = "1"
  34. os.environ["OMP_NUM_THREADS"] = "1"
  35. os.environ["MKL_NUM_THREADS"] = "1"
  36. os.environ["NUMEXPR_NUM_THREADS"] = "1"
  37. os.environ["VECLIB_MAXIMUM_THREADS"] = "1"
  38. os.environ["OMP_DYNAMIC"] = "FALSE"
  39. import numpy as np
  40. import pandas as pd
  41. from PIL import Image
  42. import matplotlib.pyplot as plt
  43. import seaborn as sns
  44. from sklearn.metrics import (
  45. accuracy_score, precision_score, recall_score, f1_score,
  46. cohen_kappa_score, classification_report, roc_auc_score,
  47. confusion_matrix, ConfusionMatrixDisplay, roc_curve, auc,
  48. precision_recall_curve
  49. )
  50. from sklearn.preprocessing import label_binarize
  51. from sklearn.preprocessing import StandardScaler
  52. from sklearn.svm import SVC
  53. from scipy.stats import ttest_rel
  54. try:
  55. import shap
  56. except Exception:
  57. shap = None
  58. import torch
  59. import torch.nn as nn
  60. import torch.nn.functional as F # model/ops functional
  61. torch.set_num_threads(1) # Force PyTorch single-threaded
  62. from torch.utils.data import Dataset, DataLoader, TensorDataset
  63. from torchvision import transforms, models
  64. import torchvision.transforms.functional as TF # transform functional
  65. from tqdm import tqdm
  66. # Grad-CAM optional
  67. GRADCAM_AVAILABLE = False
  68. try:
  69. from pytorch_grad_cam import GradCAM
  70. from pytorch_grad_cam.utils.model_targets import ClassifierOutputTarget
  71. from pytorch_grad_cam.utils.image import show_cam_on_image
  72. GRADCAM_AVAILABLE = True
  73. print("pytorch-gradcam library found. Using it for Grad-CAM visualization.")
  74. except Exception:
  75. print("pytorch-gradcam library not found. Using simple Grad-CAM helper.")
  76. # ======================== GLOBALS / HYPERPARAMS ========================
  77. SEED = 42
  78. random.seed(SEED)
  79. np.random.seed(SEED)
  80. torch.manual_seed(SEED)
  81. DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
  82. print("Device:", DEVICE)
  83. RESULTS_DIR = "final_results"
  84. os.makedirs(RESULTS_DIR, exist_ok=True)
  85. # Dataset path (Kaggle input / fallback)
  86. DATA_DIR = None
  87. try:
  88. import kagglehub
  89. try:
  90. DATA_DIR = kagglehub.dataset_download("waseemnagahhenes/brain-tumor-for-14-classes")
  91. print("kagglehub returned:", DATA_DIR)
  92. except Exception:
  93. pass
  94. except Exception:
  95. pass
  96. if not DATA_DIR:
  97. candidates = [
  98. "/kaggle/input/brain-tumor-for-14-classes",
  99. "/kaggle/input/brain-tumor-classification-dataset",
  100. "/kaggle/input/brain-tumor-for-14-classes-mri",
  101. "/kaggle/input"
  102. ]
  103. for c in candidates:
  104. if os.path.exists(c):
  105. DATA_DIR = c
  106. break
  107. if DATA_DIR is None:
  108. raise RuntimeError("DATA_DIR not set. Please check dataset.")
  109. BATCH_SIZE = 16
  110. IMG_SIZE = 224
  111. NUM_EPOCHS = 400
  112. LR = 1e-4
  113. NUM_WORKERS = 4 if torch.cuda.is_available() else 0
  114. QUICK_DEBUG = False
  115. TSNE_SAMPLES = 2000 if not QUICK_DEBUG else 200
  116. SHAP_MAX_SAMPLES = 500 if not QUICK_DEBUG else 100
  117. DEGRADE_FACTORS = [0.9, 0.7, 0.5]
  118. PLOT_DPI = 300
  119. MIXUP = True
  120. def safe_classification_report(y_true, y_pred):
  121. """
  122. Wrapper around sklearn's classification_report that avoids
  123. Precision/F-score warnings by using zero_division=0
  124. """
  125. return classification_report(y_true, y_pred, output_dict=True, zero_division=0)
  126. # ======================== CELL 3: Utilities ========================
  127. def find_image_dirs(root):
  128. image_dirs = []
  129. for dirpath, dirnames, filenames in os.walk(root):
  130. count = sum(1 for f in filenames if f.lower().endswith(('.png','.jpg','.jpeg','.bmp')))
  131. if count > 0:
  132. image_dirs.append(dirpath)
  133. return image_dirs
  134. def get_patient_id(path):
  135. # path can be full path or filename; use file stem
  136. return Path(path).stem.split('_')[0]
  137. def image_hash(path, resize=(64,64)):
  138. try:
  139. with Image.open(path) as im:
  140. im = im.convert("RGB").resize(resize)
  141. return hashlib.md5(np.array(im).tobytes()).hexdigest()
  142. except Exception:
  143. return None
  144. def get_unique_samples(dataset_path):
  145. seen_hashes = set()
  146. unique_samples = []
  147. for root, _, files in os.walk(dataset_path):
  148. for f in files:
  149. if f.lower().endswith(('.png','.jpg','.jpeg')):
  150. path = os.path.join(root, f)
  151. try:
  152. with Image.open(path) as img:
  153. img = img.convert("RGB")
  154. h = hashlib.md5(img.tobytes()).hexdigest()
  155. except Exception:
  156. continue
  157. if h not in seen_hashes:
  158. seen_hashes.add(h)
  159. label = os.path.basename(root)
  160. unique_samples.append((path, label))
  161. return unique_samples
  162. # ======================== CELL 4: Dataset Class ========================
  163. class MRIDataset(Dataset):
  164. def __init__(self, df, transform=None):
  165. self.df = df
  166. self.transform = transform
  167. def __len__(self):
  168. return len(self.df)
  169. def __getitem__(self, idx):
  170. row = self.df.iloc[idx]
  171. img = Image.open(row['path']).convert("RGB")
  172. if self.transform:
  173. img = self.transform(img)
  174. label = int(row['label'])
  175. return img, label
  176. # ===============================
  177. # ✅ Custom Transform Classes
  178. # ===============================
  179. class RandomIntensityScaling(object):
  180. """Randomly scales image intensity (brightness) by a random factor."""
  181. def __init__(self, scale_range=(0.9, 1.1), p=0.5):
  182. self.scale_range = scale_range
  183. self.p = p
  184. def __call__(self, img):
  185. # Works on PIL Image or Tensor-compatible input that TF.adjust_brightness accepts
  186. if random.random() < self.p:
  187. scale = random.uniform(*self.scale_range)
  188. img = TF.adjust_brightness(img, scale)
  189. return img
  190. class AddGaussianNoise(object):
  191. """Adds random Gaussian noise to an image tensor (after ToTensor)."""
  192. def __init__(self, mean=0.0, std=0.01, p=0.5):
  193. self.mean = mean
  194. self.std = std
  195. self.p = p
  196. def __call__(self, img):
  197. if random.random() < self.p:
  198. # If input is a PIL Image, convert to tensor
  199. if not isinstance(img, torch.Tensor):
  200. img = TF.to_tensor(img)
  201. noise = torch.randn_like(img) * self.std + self.mean
  202. img = img + noise
  203. img = torch.clamp(img, 0.0, 1.0)
  204. return img
  205. def __repr__(self):
  206. return f"{self.__class__.__name__}(mean={self.mean}, std={self.std}, p={self.p})"
  207. # ===============================
  208. # ✅ Transform Pipelines
  209. # ===============================
  210. train_transform = transforms.Compose([
  211. transforms.RandomResizedCrop((IMG_SIZE, IMG_SIZE), scale=(0.85, 1.0)),
  212. transforms.RandomHorizontalFlip(p=0.5),
  213. transforms.RandomVerticalFlip(p=0.3),
  214. transforms.RandomRotation(degrees=5),
  215. # Light intensity scaling (mimics scanner variability)
  216. RandomIntensityScaling(p=0.5, scale_range=(0.9, 1.1)),
  217. # Convert to Tensor BEFORE adding Gaussian noise or normalization
  218. transforms.ToTensor(),
  219. # Small Gaussian noise (after tensor conversion)
  220. AddGaussianNoise(mean=0., std=0.005, p=0.3),
  221. # Mild contrast and brightness adjustments (operates on PIL or Tensor-compatible inputs)
  222. transforms.ColorJitter(brightness=0.05, contrast=0.05),
  223. # Normalize (ImageNet stats)
  224. transforms.Normalize([0.485, 0.456, 0.406],
  225. [0.229, 0.224, 0.225]),
  226. # Random Erasing (helps robustness)
  227. transforms.RandomErasing(p=0.1, scale=(0.02, 0.08), ratio=(0.3, 3.3))
  228. ])
  229. val_transform = transforms.Compose([
  230. transforms.Resize((IMG_SIZE, IMG_SIZE)),
  231. transforms.ToTensor(),
  232. transforms.Normalize([0.485, 0.456, 0.406],
  233. [0.229, 0.224, 0.225])
  234. ])
  235. # ======================== CELL 6: MS-DAM Module (kept & used by create_backbone) ========================
  236. class MS_DAM(nn.Module):
  237. def __init__(self, in_channels_list, out_channels=512, sampling_kernel=3, offset_scale=0.15, use_se=True):
  238. super().__init__()
  239. self.in_chs = in_channels_list
  240. self.num_scales = len(in_channels_list)
  241. self.out_ch = out_channels
  242. self.offset_scale = offset_scale
  243. self.proj_convs = nn.ModuleList([
  244. nn.Sequential(
  245. nn.Conv2d(c, out_channels, 1, bias=False),
  246. nn.BatchNorm2d(out_channels),
  247. nn.ReLU(inplace=True)
  248. ) for c in in_channels_list
  249. ])
  250. self.offset_pred = nn.Sequential(
  251. nn.Conv2d(out_channels * self.num_scales, out_channels, 3, padding=1),
  252. nn.ReLU(inplace=True),
  253. nn.Conv2d(out_channels, 2, 1)
  254. )
  255. self.spatial_attn = nn.Sequential(
  256. nn.Conv2d(out_channels, out_channels // 4, 3, padding=1),
  257. nn.ReLU(inplace=True),
  258. nn.Conv2d(out_channels // 4, 1, 1),
  259. nn.Sigmoid()
  260. )
  261. self.use_se = use_se
  262. if use_se:
  263. self.se_fc = nn.Sequential(
  264. nn.AdaptiveAvgPool2d(1),
  265. nn.Conv2d(out_channels, out_channels // 8, 1),
  266. nn.ReLU(inplace=True),
  267. nn.Conv2d(out_channels // 8, out_channels, 1),
  268. nn.Sigmoid()
  269. )
  270. self.fuse = nn.Sequential(
  271. nn.Conv2d(out_channels, out_channels, 3, padding=1, bias=False),
  272. nn.BatchNorm2d(out_channels),
  273. nn.ReLU(inplace=True)
  274. )
  275. def forward(self, feats):
  276. projected = []
  277. sizes = [f.shape[-2:] for f in feats]
  278. areas = [s[0] * s[1] for s in sizes]
  279. ref_idx = int(np.argmax(areas))
  280. ref_h, ref_w = sizes[ref_idx]
  281. for i, f in enumerate(feats):
  282. p = self.proj_convs[i](f)
  283. if (p.shape[-2], p.shape[-1]) != (ref_h, ref_w):
  284. p = F.interpolate(p, size=(ref_h, ref_w), mode='bilinear', align_corners=False)
  285. projected.append(p)
  286. concat = torch.cat(projected, dim=1)
  287. offsets = torch.tanh(self.offset_pred(concat)) * self.offset_scale
  288. B, _, H, W = offsets.shape
  289. yy, xx = torch.meshgrid(
  290. torch.linspace(-1, 1, H, device=offsets.device),
  291. torch.linspace(-1, 1, W, device=offsets.device),
  292. indexing='ij'
  293. )
  294. # base_grid shape: (1, H, W, 2)
  295. base_grid = torch.stack((xx, yy), dim=-1).unsqueeze(0).repeat(B, 1, 1, 1)
  296. offsets_grid = offsets.permute(0, 2, 3, 1)
  297. sampling_grid = (base_grid + offsets_grid).clamp(-1, 1)
  298. sampled = F.grid_sample(concat, sampling_grid, mode='bilinear', padding_mode='border', align_corners=True)
  299. sampled = sampled.view(B, self.num_scales, self.out_ch, H, W).mean(dim=1)
  300. spat = self.spatial_attn(sampled)
  301. feat = sampled * spat
  302. if self.use_se:
  303. ch_att = self.se_fc(feat)
  304. feat = feat * ch_att
  305. # in original code there was an extra feat=feat*ch_att; we avoid duplicating multiplication
  306. feat = self.fuse(feat)
  307. return feat
  308. # ======================== CELL 7: Backbone Wrapper & create_backbone ========================
  309. class ResNetWithMSDAM(nn.Module):
  310. def __init__(self, base_resnet, msdam, num_classes):
  311. super().__init__()
  312. self.conv1 = base_resnet.conv1
  313. self.bn1 = base_resnet.bn1
  314. self.relu = base_resnet.relu
  315. self.maxpool = base_resnet.maxpool
  316. self.layer1 = base_resnet.layer1
  317. self.layer2 = base_resnet.layer2
  318. self.layer3 = base_resnet.layer3
  319. self.layer4 = base_resnet.layer4
  320. self.msdam = msdam
  321. self.classifier = nn.Sequential(
  322. nn.AdaptiveAvgPool2d(1),
  323. nn.Flatten(),
  324. nn.Dropout(0.4),
  325. nn.Linear(msdam.out_ch, num_classes)
  326. )
  327. def forward_features(self, x):
  328. x = self.conv1(x); x = self.bn1(x); x = self.relu(x); x = self.maxpool(x)
  329. c2 = self.layer1(x); c3 = self.layer2(c2); c4 = self.layer3(c3); c5 = self.layer4(c4)
  330. return [c2, c3, c4, c5]
  331. def forward(self, x):
  332. feats = self.forward_features(x)
  333. fused = self.msdam(feats)
  334. out = self.classifier(fused)
  335. return out, fused # Return fused features as well for Grad-CAM
  336. def create_backbone(model_type="resnet101", pretrained=True, num_classes=4):
  337. """
  338. Create model with MS_DAM + ResNet backbone.
  339. model_type currently supports 'resnet101' (can extend if needed).
  340. """
  341. if model_type == "resnet101":
  342. base = models.resnet101(weights=models.ResNet101_Weights.DEFAULT if pretrained else None)
  343. in_chs = [256, 512, 1024, 2048]
  344. msdam = MS_DAM(in_chs, out_channels=512, offset_scale=0.12, use_se=True)
  345. model = ResNetWithMSDAM(base, msdam, num_classes=num_classes)
  346. else:
  347. raise ValueError(f"Unknown model_type: {model_type}")
  348. model.to(DEVICE)
  349. return model
  350. # ======================== CELL 8: MixUp Helper ========================
  351. def mixup_data(x, y, alpha=0.1):
  352. if alpha <= 0:
  353. return x, y, 1.0, None
  354. lam = np.random.beta(alpha, alpha)
  355. batch_size = x.size()[0]
  356. index = torch.randperm(batch_size).to(x.device)
  357. mixed_x = lam * x + (1 - lam) * x[index, :]
  358. y_a, y_b = y, y[index]
  359. return mixed_x, (y_a, y_b, lam)
  360. # ======================== CELL 9: Logger ========================
  361. class TrainingLogger:
  362. def __init__(self):
  363. self.train_losses = []
  364. self.val_losses = []
  365. self.train_accs = []
  366. self.val_accs = []
  367. def log_epoch(self, train_loss, val_loss, train_acc, val_acc):
  368. self.train_losses.append(train_loss)
  369. self.val_losses.append(val_loss)
  370. self.train_accs.append(train_acc)
  371. self.val_accs.append(val_acc)
  372. def plot(self):
  373. epochs = range(1, len(self.train_losses) + 1)
  374. plt.figure(figsize=(10, 5))
  375. plt.plot(epochs, self.train_losses, 'b-', label="Train Loss")
  376. plt.plot(epochs, self.val_losses, 'r-', label="Val Loss")
  377. plt.xlabel("Epoch"); plt.ylabel("Loss"); plt.title("Training & Validation Loss"); plt.legend(); plt.show()
  378. plt.figure(figsize=(10, 5))
  379. plt.plot(epochs, self.train_accs, 'b-', label="Train Acc")
  380. plt.plot(epochs, self.val_accs, 'r-', label="Val Acc")
  381. plt.xlabel("Epoch"); plt.ylabel("Accuracy"); plt.title("Training & Validation Accuracy"); plt.legend(); plt.show()
  382. # ======================== CELL 10: Feature Extraction + SVM ========================
  383. def extract_penultimate_features(model, loader):
  384. model.eval()
  385. feats = []
  386. labels = []
  387. with torch.no_grad():
  388. for imgs, lbls in loader:
  389. imgs = imgs.to(DEVICE)
  390. out, fused = model(imgs) # model returns (logits, features)
  391. # pooled fused features
  392. fused_pooled = fused.mean(dim=[2, 3])
  393. feats.append(fused_pooled.cpu().numpy())
  394. labels.append(lbls.numpy())
  395. return np.vstack(feats), np.hstack(labels)
  396. def train_svm_on_features(model, train_loader, val_loader):
  397. X_train, y_train = extract_penultimate_features(model, train_loader)
  398. X_val, y_val = extract_penultimate_features(model, val_loader)
  399. scaler = StandardScaler().fit(X_train)
  400. clf = SVC(kernel='rbf', C=10, gamma='scale')
  401. clf.fit(scaler.transform(X_train), y_train)
  402. acc = clf.score(scaler.transform(X_val), y_val)
  403. print(f"SVM on penultimate features Val Acc: {acc*100:.2f}%")
  404. return clf, scaler
  405. # ======================== CELL 11: Train + Save + Evaluate Best Accuracy & Loss Models ========================
  406. import pandas as _pd # local alias to avoid clobbering pd in outer scope (kept but not necessary)
  407. def train_one_model(model, train_loader, val_loader, test_loader,
  408. epochs=20, lr=1e-4, mixup=True, out_dir="./results"):
  409. """
  410. Trains a model, saves best models (by val accuracy and val loss),
  411. uses tie-breaking on train metrics, and evaluates both models.
  412. (Early stopping removed)
  413. """
  414. os.makedirs(out_dir, exist_ok=True)
  415. optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-4)
  416. scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs, eta_min=1e-6)
  417. criterion = nn.CrossEntropyLoss(label_smoothing=0.1)
  418. logger = TrainingLogger()
  419. best_val_acc = -1.0
  420. best_val_loss = float('inf')
  421. best_train_acc_for_best_val = -1.0
  422. best_train_loss_for_best_val = float('inf')
  423. best_acc_epoch = -1
  424. best_loss_epoch = -1
  425. best_acc_path = os.path.join(out_dir, "best_acc_model.pth")
  426. best_loss_path = os.path.join(out_dir, "best_loss_model.pth")
  427. device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
  428. model = model.to(device)
  429. print("🚀 Starting Training with Dual Criteria (No Early Stopping)...")
  430. for epoch in range(epochs):
  431. model.train()
  432. running_loss, correct, total = 0.0, 0, 0
  433. # ---------------- Training ----------------
  434. for imgs, labels in train_loader:
  435. imgs, labels = imgs.to(device), labels.to(device)
  436. if mixup:
  437. imgs, (y_a, y_b, lam) = mixup_data(imgs, labels)
  438. out, _ = model(imgs)
  439. loss = lam * criterion(out, y_a) + (1 - lam) * criterion(out, y_b)
  440. else:
  441. out, _ = model(imgs)
  442. loss = criterion(out, labels)
  443. optimizer.zero_grad()
  444. loss.backward()
  445. optimizer.step()
  446. running_loss += loss.item() * imgs.size(0)
  447. _, preds = out.max(1)
  448. if mixup:
  449. correct += lam * (preds == y_a).sum().item() + (1 - lam) * (preds == y_b).sum().item()
  450. else:
  451. correct += (preds == labels).sum().item()
  452. total += imgs.size(0)
  453. train_loss = running_loss / total
  454. train_acc = correct / total
  455. # step scheduler (end of epoch)
  456. scheduler.step()
  457. # ---------------- Validation ----------------
  458. model.eval()
  459. val_loss, correct_val, total_val = 0.0, 0, 0
  460. with torch.no_grad():
  461. for imgs, labels in val_loader:
  462. imgs, labels = imgs.to(device), labels.to(device)
  463. out, _ = model(imgs)
  464. loss = criterion(out, labels)
  465. val_loss += loss.item() * imgs.size(0)
  466. _, preds = out.max(1)
  467. correct_val += (preds == labels).sum().item()
  468. total_val += imgs.size(0)
  469. val_loss /= total_val
  470. val_acc = correct_val / total_val
  471. logger.log_epoch(train_loss, val_loss, train_acc, val_acc)
  472. print(f"Epoch {epoch+1}/{epochs} | "
  473. f"Train Loss {train_loss:.4f} | Val Loss {val_loss:.4f} | "
  474. f"Train Acc {train_acc:.4f} | Val Acc {val_acc:.4f}")
  475. # ---------------- Save Best Models ----------------
  476. # Best by val accuracy (tie-breaker: higher train acc)
  477. if (val_acc > best_val_acc) or (abs(val_acc - best_val_acc) < 1e-6 and train_acc > best_train_acc_for_best_val):
  478. best_val_acc = val_acc
  479. best_train_acc_for_best_val = train_acc
  480. best_acc_epoch = epoch + 1
  481. torch.save(model.state_dict(), best_acc_path)
  482. print(f"✅ Saved Best Model (Val Acc): Epoch {best_acc_epoch} | Val Acc {best_val_acc:.4f}")
  483. # Best by val loss (tie-breaker: lower train loss)
  484. if (val_loss < best_val_loss) or (abs(val_loss - best_val_loss) < 1e-6 and train_loss < best_train_loss_for_best_val):
  485. best_val_loss = val_loss
  486. best_train_loss_for_best_val = train_loss
  487. best_loss_epoch = epoch + 1
  488. torch.save(model.state_dict(), best_loss_path)
  489. print(f"✅ Saved Best Model (Val Loss): Epoch {best_loss_epoch} | Val Loss {best_val_loss:.4f}")
  490. # ---------------- After Training ----------------
  491. print("\n🏁 Training Complete!")
  492. print(f"✅ Best Model (Val Acc): Epoch {best_acc_epoch} | Best Val Acc: {best_val_acc:.4f}")
  493. print(f"✅ Best Model (Val Loss): Epoch {best_loss_epoch} | Best Val Loss: {best_val_loss:.4f}")
  494. # ---------------- Save Training History ----------------
  495. history = pd.DataFrame({
  496. "epoch": range(1, len(logger.train_losses) + 1),
  497. "train_loss": logger.train_losses,
  498. "val_loss": logger.val_losses,
  499. "train_acc": logger.train_accs,
  500. "val_acc": logger.val_accs
  501. })
  502. history.to_csv(os.path.join(out_dir, "history.csv"), index=False)
  503. # ---------------- Plot Loss & Accuracy ----------------
  504. def save_plot(metric_name, train_values, val_values):
  505. plt.figure(figsize=(10, 4))
  506. plt.plot(train_values, "b-", label=f"Train {metric_name}", linewidth=2)
  507. plt.plot(val_values, "r-", label=f"Val {metric_name}", linewidth=2)
  508. plt.title(f"{metric_name} vs Epoch")
  509. plt.xlabel("Epoch"); plt.ylabel(metric_name)
  510. plt.legend(); plt.grid(True)
  511. plt.tight_layout()
  512. plt.savefig(os.path.join(out_dir, f"{metric_name.lower()}_plot.png"), dpi=200)
  513. plt.close()
  514. save_plot("Accuracy", logger.train_accs, logger.val_accs)
  515. save_plot("Loss", logger.train_losses, logger.val_losses)
  516. # ---------------- Load Best Models ----------------
  517. model_acc = copy.deepcopy(model)
  518. model_acc.load_state_dict(torch.load(best_acc_path, map_location=device))
  519. model_loss = copy.deepcopy(model)
  520. model_loss.load_state_dict(torch.load(best_loss_path, map_location=device))
  521. # ---------------- Evaluate Both Models ----------------
  522. for name, model_sel in zip(["accuracy_model", "loss_model"], [model_acc, model_loss]):
  523. sub_out_dir = os.path.join(out_dir, name)
  524. os.makedirs(sub_out_dir, exist_ok=True)
  525. evaluate_on_loader(model_sel, test_loader, criterion, out_dir=sub_out_dir)
  526. return model_acc, model_loss, logger
  527. def evaluate_on_loader(model, loader, criterion, out_dir=None):
  528. import os
  529. import pandas as pd
  530. import torch.nn.functional as F_local
  531. from sklearn.metrics import accuracy_score, classification_report, confusion_matrix
  532. model.eval()
  533. y_true, y_pred, y_prob = [], [], []
  534. total_loss = 0.0
  535. with torch.no_grad():
  536. for imgs, labels in loader:
  537. imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)
  538. out, _ = model(imgs) # model returns (logits, features)
  539. loss = criterion(out, labels)
  540. total_loss += loss.item() * imgs.size(0)
  541. probs = F_local.softmax(out, dim=1).cpu().detach().numpy()
  542. preds = probs.argmax(axis=1)
  543. y_true.extend(labels.cpu().numpy())
  544. y_pred.extend(preds)
  545. y_prob.extend(probs)
  546. # Convert to arrays
  547. y_true, y_pred, y_prob = np.array(y_true), np.array(y_pred), np.array(y_prob)
  548. # Metrics
  549. avg_loss = total_loss / len(loader.dataset)
  550. acc = accuracy_score(y_true, y_pred)
  551. if out_dir:
  552. os.makedirs(out_dir, exist_ok=True)
  553. # 1️⃣ Save basic metrics to CSV
  554. metrics_df = pd.DataFrame([{
  555. "loss": avg_loss,
  556. "accuracy": acc,
  557. "num_samples": len(y_true)
  558. }])
  559. metrics_df.to_csv(os.path.join(out_dir, "metrics.csv"), index=False)
  560. # 2️⃣ Save classification report to CSV
  561. cls_report_dict = classification_report(y_true, y_pred, output_dict=True, zero_division=0)
  562. cls_report_df = pd.DataFrame(cls_report_dict).transpose()
  563. cls_report_df.to_csv(os.path.join(out_dir, "classification_report.csv"))
  564. # 3️⃣ Save confusion matrix to CSV
  565. cm = confusion_matrix(y_true, y_pred)
  566. cm_df = pd.DataFrame(cm, index=[str(i) for i in range(len(cm))],
  567. columns=[str(i) for i in range(len(cm))])
  568. cm_df.to_csv(os.path.join(out_dir, "confusion_matrix.csv"))
  569. print(f"✅ Evaluation metrics saved to {out_dir} (metrics.csv, classification_report.csv, confusion_matrix.csv)")
  570. return avg_loss, acc, y_true, y_pred, y_prob
  571. # ======================== Additional reporting helpers ========================
  572. def save_reports(y_true, y_pred, class_names, out_dir):
  573. cm = confusion_matrix(y_true, y_pred)
  574. df_cm = pd.DataFrame(cm, index=class_names, columns=class_names)
  575. df_cm.to_csv(os.path.join(out_dir, "confusion_matrix.csv"))
  576. plt.figure(figsize=(10, 8))
  577. sns.heatmap(df_cm, annot=True, fmt="d", cmap="Blues")
  578. plt.title("Confusion Matrix"); plt.ylabel("True"); plt.xlabel("Predicted")
  579. plt.tight_layout(); plt.savefig(os.path.join(out_dir, "confusion_matrix.png"), dpi=PLOT_DPI)
  580. plt.close()
  581. report = classification_report(y_true, y_pred, target_names=class_names, output_dict=True, zero_division=0)
  582. pd.DataFrame(report).to_csv(os.path.join(out_dir, "classification_report.csv"))
  583. return df_cm, report
  584. def save_misclassified(df_test, y_true, y_pred, class_names, out_dir):
  585. mis_idx = np.where(y_true != y_pred)[0]
  586. mis_list = [(df_test.iloc[i]['path'], class_names[y_true[i]], class_names[y_pred[i]]) for i in mis_idx]
  587. pd.DataFrame(mis_list, columns=["image_path", "true_label", "pred_label"]).to_csv(
  588. os.path.join(out_dir, "misclassified.csv"), index=False)
  589. # ======================== DATA SCAN, SPLIT, DATALOADERS ========================
  590. print("🔍 Scanning dataset...")
  591. unique_samples = get_unique_samples(DATA_DIR)
  592. class_names = sorted({label for _, label in unique_samples})
  593. class_to_idx = {c: i for i, c in enumerate(class_names)}
  594. labels_map_inv = {i: c for c, i in class_to_idx.items()} # for Grad-CAM
  595. numeric_samples = [(p, class_to_idx[lab], get_patient_id(p)) for p, lab in unique_samples]
  596. df = pd.DataFrame(numeric_samples, columns=['path', 'label', 'pid'])
  597. # Split by patient
  598. patient_ids = df['pid'].unique()
  599. random.shuffle(patient_ids)
  600. n_train, n_val = int(0.7 * len(patient_ids)), int(0.15 * len(patient_ids))
  601. train_pids = patient_ids[:n_train]
  602. val_pids = patient_ids[n_train:n_train + n_val]
  603. test_pids = patient_ids[n_train + n_val:]
  604. df_train = df[df['pid'].isin(train_pids)].reset_index(drop=True)
  605. df_val = df[df['pid'].isin(val_pids)].reset_index(drop=True)
  606. df_test = df[df['pid'].isin(test_pids)].reset_index(drop=True)
  607. print(f"✅ Samples — Train: {len(df_train)}, Val: {len(df_val)}, Test: {len(df_test)}")
  608. # Count samples and unique patients per split
  609. splits = [("Train", df_train), ("Val", df_val), ("Test", df_test)]
  610. print("📊 Dataset summary:")
  611. print("{:<8} {:>10} {:>15}".format("Split", "Samples", "Unique Patients"))
  612. print("-" * 35)
  613. for name, df_split in splits:
  614. n_samples = len(df_split)
  615. n_patients = df_split['pid'].nunique()
  616. print(f"{name:<8} {n_samples:>10} {n_patients:>15}")
  617. # Optional: check patient overlap between splits
  618. train_pids_set = set(df_train['pid'])
  619. val_pids_set = set(df_val['pid'])
  620. test_pids_set = set(df_test['pid'])
  621. print("\n🧪 Patient overlaps (should be 0):")
  622. print("Train ∩ Val:", len(train_pids_set & val_pids_set))
  623. print("Train ∩ Test:", len(train_pids_set & test_pids_set))
  624. print("Val ∩ Test:", len(val_pids_set & test_pids_set))
  625. print("🧠 Creating DataLoaders...")
  626. train_loader = DataLoader(MRIDataset(df_train, train_transform),
  627. batch_size=BATCH_SIZE, shuffle=True, num_workers=NUM_WORKERS)
  628. val_loader = DataLoader(MRIDataset(df_val, val_transform),
  629. batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS)
  630. test_loader = DataLoader(MRIDataset(df_test, val_transform),
  631. batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS)
  632. print("✅ DataLoaders ready.")
  633. # ======================== CREATE MODEL, TRAIN ========================
  634. print("⚙️ Creating model...")
  635. model = create_backbone("resnet101", num_classes=len(class_names))
  636. print("✅ Model created.")
  637. print("🚀 Starting training...")
  638. model_acc, model_loss, logger = train_one_model(
  639. model=model,
  640. train_loader=train_loader,
  641. val_loader=val_loader,
  642. test_loader=test_loader,
  643. epochs=NUM_EPOCHS,
  644. lr=LR,
  645. mixup=MIXUP,
  646. out_dir=RESULTS_DIR,
  647. )
  648. print("✅ Training completed.")
  649. # %%
  650. # Paths for both best models
  651. best_acc_path = os.path.join(RESULTS_DIR, "best_acc_model.pth")
  652. best_loss_path = os.path.join(RESULTS_DIR, "best_loss_model.pth")
  653. # Dictionary to store the loaded models
  654. best_models = {}
  655. # Load best accuracy model
  656. if os.path.exists(best_acc_path):
  657. model_acc = create_backbone("resnet101", num_classes=len(class_names)) # new model instance
  658. model_acc.load_state_dict(torch.load(best_acc_path, map_location=DEVICE))
  659. model_acc.to(DEVICE)
  660. model_acc.eval()
  661. best_models['best_acc'] = model_acc
  662. print("✅ Loaded Best Accuracy Model:", best_acc_path)
  663. else:
  664. print("⚠️ Best accuracy model not found:", best_acc_path)
  665. # Load best loss model
  666. if os.path.exists(best_loss_path):
  667. model_loss = create_backbone("resnet101", num_classes=len(class_names)) # new model instance
  668. model_loss.load_state_dict(torch.load(best_loss_path, map_location=DEVICE))
  669. model_loss.to(DEVICE)
  670. model_loss.eval()
  671. best_models['best_loss'] = model_loss
  672. print("✅ Loaded Best Loss Model:", best_loss_path)
  673. else:
  674. print("⚠️ Best loss model not found:", best_loss_path)
  675. import os
  676. import matplotlib.pyplot as plt
  677. from sklearn.metrics import roc_curve, auc, precision_recall_curve
  678. from sklearn.preprocessing import label_binarize
  679. def plot_roc_pr(y_true, y_prob, class_names, out_dir, plot_dpi=300):
  680. os.makedirs(out_dir, exist_ok=True)
  681. y_bin = label_binarize(y_true, classes=list(range(len(class_names))))
  682. # ---- ROC Curves ----
  683. plt.figure(figsize=(10,8))
  684. for i, c in enumerate(class_names):
  685. fpr, tpr, _ = roc_curve(y_bin[:, i], y_prob[:, i])
  686. roc_auc = auc(fpr, tpr)
  687. plt.plot(fpr, tpr, label=f"{c} (AUC={roc_auc:.2f})")
  688. plt.plot([0, 1], [0, 1], 'k--')
  689. plt.xlabel("False Positive Rate (FPR)")
  690. plt.ylabel("True Positive Rate (TPR)")
  691. plt.title("ROC Curves (All Classes)")
  692. plt.legend()
  693. plt.tight_layout()
  694. roc_path = os.path.join(out_dir, "roc_curves.png")
  695. plt.savefig(roc_path, dpi=plot_dpi)
  696. plt.close()
  697. # ---- Precision–Recall Curves ----
  698. plt.figure(figsize=(10,8))
  699. for i, c in enumerate(class_names):
  700. prec, rec, _ = precision_recall_curve(y_bin[:, i], y_prob[:, i])
  701. plt.plot(rec, prec, label=c)
  702. plt.xlabel("Recall")
  703. plt.ylabel("Precision")
  704. plt.title("Precision–Recall Curves (All Classes)")
  705. plt.legend()
  706. plt.tight_layout()
  707. pr_path = os.path.join(out_dir, "pr_curves.png")
  708. plt.savefig(pr_path, dpi=plot_dpi)
  709. plt.close()
  710. print(f"✅ ROC and PR plots saved successfully:\n"
  711. f" • {roc_path}\n"
  712. f" • {pr_path}")
  713. eval_out_dir = os.path.join(RESULTS_DIR, "evaluation")
  714. os.makedirs(eval_out_dir, exist_ok=True)
  715. for model_key, model_instance in best_models.items():
  716. # 1️⃣ Create model-specific output folder
  717. model_out_dir = os.path.join(eval_out_dir, model_key)
  718. os.makedirs(model_out_dir, exist_ok=True)
  719. # 2️⃣ Get predictions and probabilities for this model
  720. all_labels, all_probs = [], []
  721. model_instance.eval()
  722. with torch.no_grad():
  723. for imgs, labels in test_loader:
  724. imgs = imgs.to(DEVICE)
  725. labels = labels.to(DEVICE)
  726. outputs, _ = model_instance(imgs) # unpack tuple if model returns (logits, features)
  727. probs = torch.softmax(outputs, dim=1)
  728. all_labels.extend(labels.cpu().numpy())
  729. all_probs.extend(probs.cpu().numpy())
  730. y_true_model = np.array(all_labels)
  731. y_prob_model = np.array(all_probs)
  732. # 3️⃣ Call your plotting function
  733. plot_roc_pr(y_true_model, y_prob_model, class_names, model_out_dir)
  734. from sklearn.manifold import TSNE
  735. def save_evaluation_reports(y_true, y_pred, y_prob, class_names, out_dir):
  736. # Confusion matrix
  737. cm = confusion_matrix(y_true, y_pred)
  738. df_cm = pd.DataFrame(cm, index=class_names, columns=class_names)
  739. df_cm.to_csv(os.path.join(out_dir, "confusion_matrix.csv"))
  740. # Classification report
  741. report = classification_report(y_true, y_pred, target_names=class_names, output_dict=True)
  742. report_df = pd.DataFrame(report).transpose()
  743. report_df.to_csv(os.path.join(out_dir, "classification_report.csv"))
  744. return df_cm, report_df
  745. def evaluate_on_loader(model, loader, criterion, device=DEVICE):
  746. model.to(device)
  747. model.eval()
  748. y_true, y_pred, y_prob = [], [], []
  749. total_loss = 0.0
  750. with torch.no_grad():
  751. for imgs, labels in loader:
  752. imgs, labels = imgs.to(device), labels.to(device)
  753. out = model(imgs)
  754. if isinstance(out, (tuple, list)):
  755. out = out[0] # logits
  756. loss = criterion(out, labels)
  757. total_loss += loss.item() * imgs.size(0)
  758. probs = torch.softmax(out, dim=1).cpu().numpy()
  759. preds = np.argmax(probs, axis=1)
  760. y_true.extend(labels.cpu().numpy())
  761. y_pred.extend(preds)
  762. y_prob.extend(probs)
  763. avg_loss = total_loss / len(loader.dataset)
  764. acc = accuracy_score(y_true, y_pred)
  765. return avg_loss, acc, np.array(y_true), np.array(y_pred), np.array(y_prob)
  766. def save_misclassified(df_test, y_true, y_pred, class_names, out_dir):
  767. mis_idx = np.where(y_true != y_pred)[0]
  768. mis_list = [(df_test.iloc[i]['path'], class_names[y_true[i]], class_names[y_pred[i]]) for i in mis_idx]
  769. df_mis = pd.DataFrame(mis_list, columns=["image_path", "true_label", "pred_label"])
  770. df_mis.to_csv(os.path.join(out_dir, "misclassified.csv"), index=False)
  771. return df_mis
  772. def plot_tsne_features(model, loader, class_names, out_dir, n_samples=100):
  773. X, y = extract_penultimate_features(model, loader)
  774. if len(X) > n_samples:
  775. idx = np.random.choice(len(X), n_samples, replace=False)
  776. X, y = X[idx], y[idx]
  777. tsne = TSNE(n_components=2, random_state=42)
  778. X_tsne = tsne.fit_transform(X)
  779. plt.figure(figsize=(8,6))
  780. for cls in np.unique(y):
  781. plt.scatter(X_tsne[y==cls,0], X_tsne[y==cls,1], label=class_names[cls], alpha=0.7)
  782. plt.legend(); plt.title("t-SNE of Penultimate Features")
  783. plt.savefig(os.path.join(out_dir, "tsne_features_1.png"), dpi=300)
  784. plt.close()
  785. # -----------------------------
  786. # Multi-model evaluation loop
  787. # -----------------------------
  788. for name, model in best_models.items():
  789. print(f"\n🧪 Evaluating {name} model...")
  790. eval_out_dir = os.path.join(RESULTS_DIR, f"evaluation_{name}")
  791. Path(eval_out_dir).mkdir(parents=True, exist_ok=True)
  792. # 1️⃣ SVM on penultimate features
  793. print("🧩 Training SVM on penultimate features...")
  794. model.to(DEVICE)
  795. clf, scaler = train_svm_on_features(model, train_loader, val_loader)
  796. # 2️⃣ Evaluate on test set
  797. print("📊 Generating evaluation reports...")
  798. test_loss, test_acc, y_true, y_pred, y_prob = evaluate_on_loader(model, test_loader, nn.CrossEntropyLoss(), device=DEVICE)
  799. print(f"✅ {name} Test Loss: {test_loss:.4f} | Accuracy: {test_acc:.4f}")
  800. # 3️⃣ Save reports
  801. df_cm, report_df = save_evaluation_reports(y_true, y_pred, y_prob, class_names, eval_out_dir)
  802. # 5️⃣ Misclassified samples
  803. df_mis = save_misclassified(df_test, y_true, y_pred, class_names, eval_out_dir)
  804. # 6️⃣ t-SNE visualization
  805. plot_tsne_features(model, test_loader, class_names, eval_out_dir, n_samples=TSNE_SAMPLES)
  806. print(f"\n✅ Evaluation completed for {name} model. Reports saved to {eval_out_dir}")
  807. import os
  808. import numpy as np
  809. import pandas as pd
  810. import torch
  811. import torch.nn as nn
  812. import torch.nn.functional as F
  813. import matplotlib.pyplot as plt
  814. from pathlib import Path
  815. from sklearn.metrics import confusion_matrix, classification_report, accuracy_score
  816. from sklearn.manifold import TSNE
  817. from sklearn.decomposition import PCA
  818. # -----------------------------------------------------------
  819. # ✅ Helper 1: Extract penultimate features
  820. # -----------------------------------------------------------
  821. def extract_penultimate_features(model, loader, device="cuda"):
  822. """
  823. Extracts penultimate layer features for visualization.
  824. Automatically flattens spatial dimensions if needed.
  825. Assumes model returns (logits, features) in its forward pass.
  826. """
  827. model.eval()
  828. features, labels = [], []
  829. with torch.no_grad():
  830. for imgs, lbls in loader:
  831. imgs = imgs.to(device)
  832. lbls = lbls.to(device)
  833. out = model(imgs)
  834. if isinstance(out, (tuple, list)):
  835. logits, feats = out # (logits, features)
  836. else:
  837. feats = out # in case model outputs only features
  838. feats = feats.cpu().numpy()
  839. # 🔧 Flatten 4D tensors [B, C, H, W] → [B, C*H*W]
  840. if feats.ndim > 2:
  841. feats = feats.reshape(feats.shape[0], -1)
  842. features.append(feats)
  843. labels.extend(lbls.cpu().numpy())
  844. return np.concatenate(features), np.array(labels)
  845. # -----------------------------------------------------------
  846. # ✅ Helper 2: Evaluation on test/val loader
  847. # -----------------------------------------------------------
  848. def evaluate_on_loader(model, loader, criterion, device="cuda"):
  849. model.to(device)
  850. model.eval()
  851. y_true, y_pred, y_prob = [], [], []
  852. total_loss = 0.0
  853. with torch.no_grad():
  854. for imgs, labels in loader:
  855. imgs, labels = imgs.to(device), labels.to(device)
  856. out = model(imgs)
  857. if isinstance(out, (tuple, list)):
  858. out = out[0] # logits
  859. loss = criterion(out, labels)
  860. total_loss += loss.item() * imgs.size(0)
  861. probs = F.softmax(out, dim=1).cpu().numpy()
  862. preds = np.argmax(probs, axis=1)
  863. y_true.extend(labels.cpu().numpy())
  864. y_pred.extend(preds)
  865. y_prob.extend(probs)
  866. avg_loss = total_loss / len(loader.dataset)
  867. acc = accuracy_score(y_true, y_pred)
  868. return avg_loss, acc, np.array(y_true), np.array(y_pred), np.array(y_prob)
  869. # -----------------------------------------------------------
  870. # ✅ Helper 3: Save evaluation reports
  871. # -----------------------------------------------------------
  872. def save_evaluation_reports(y_true, y_pred, y_prob, class_names, out_dir):
  873. os.makedirs(out_dir, exist_ok=True)
  874. # Confusion matrix
  875. cm = confusion_matrix(y_true, y_pred)
  876. df_cm = pd.DataFrame(cm, index=class_names, columns=class_names)
  877. df_cm.to_csv(os.path.join(out_dir, "confusion_matrix.csv"))
  878. # Classification report
  879. report = classification_report(y_true, y_pred, target_names=class_names,
  880. output_dict=True, zero_division=0)
  881. report_df = pd.DataFrame(report).transpose()
  882. report_df.to_csv(os.path.join(out_dir, "classification_report.csv"))
  883. print("✅ Saved confusion matrix & classification report")
  884. return df_cm, report_df
  885. # -----------------------------------------------------------
  886. # ✅ Helper 4: Save misclassified samples
  887. # -----------------------------------------------------------
  888. def save_misclassified(y_true, y_pred, class_names, out_dir, df_test=None):
  889. mis_idx = np.where(y_true != y_pred)[0]
  890. if df_test is not None and "path" in df_test.columns:
  891. mis_list = [(df_test.iloc[i]['path'], class_names[y_true[i]], class_names[y_pred[i]]) for i in mis_idx]
  892. df_mis = pd.DataFrame(mis_list, columns=["image_path", "true_label", "pred_label"])
  893. else:
  894. df_mis = pd.DataFrame({
  895. "index": mis_idx,
  896. "true_label": [class_names[y_true[i]] for i in mis_idx],
  897. "pred_label": [class_names[y_pred[i]] for i in mis_idx]
  898. })
  899. os.makedirs(out_dir, exist_ok=True)
  900. df_mis.to_csv(os.path.join(out_dir, "misclassified.csv"), index=False)
  901. print(f"✅ Saved misclassified samples CSV ({len(df_mis)} samples)")
  902. return df_mis
  903. # -----------------------------------------------------------
  904. # ✅ Helper 5: Plot t-SNE with PCA preprocessing
  905. # -----------------------------------------------------------
  906. def plot_tsne_features(model, loader, class_names, out_dir, n_samples=300, device="cuda"):
  907. X, y = extract_penultimate_features(model, loader, device=device)
  908. if len(X) > n_samples:
  909. idx = np.random.choice(len(X), n_samples, replace=False)
  910. X, y = X[idx], y[idx]
  911. # PCA to 50D before t-SNE (speeds up & denoises)
  912. pca = PCA(n_components=min(50, X.shape[1]))
  913. X_pca = pca.fit_transform(X)
  914. tsne = TSNE(n_components=2, random_state=42, perplexity=30)
  915. X_tsne = tsne.fit_transform(X_pca)
  916. plt.figure(figsize=(8, 6))
  917. for cls in np.unique(y):
  918. plt.scatter(X_tsne[y == cls, 0], X_tsne[y == cls, 1],
  919. label=class_names[cls], alpha=0.7, s=40)
  920. plt.legend()
  921. plt.title("t-SNE of Penultimate Features (PCA Preprocessed)")
  922. plt.tight_layout()
  923. plt.savefig(os.path.join(out_dir, "tsne_features.png"), dpi=300)
  924. plt.close()
  925. print("✅ Saved t-SNE visualization")
  926. # -----------------------------------------------------------
  927. # ✅ Main Evaluation Loop for Multiple Models
  928. # -----------------------------------------------------------
  929. def evaluate_all_models(best_models, train_loader, val_loader, test_loader,
  930. RESULTS_DIR, class_names, TSNE_SAMPLES=300, DEVICE="cuda",
  931. df_test=None):
  932. criterion = nn.CrossEntropyLoss()
  933. for name, model in best_models.items():
  934. print(f"\n🧪 Evaluating {name} model...")
  935. eval_out_dir = os.path.join(RESULTS_DIR, f"evaluation_{name}")
  936. Path(eval_out_dir).mkdir(parents=True, exist_ok=True)
  937. # 1️⃣ Evaluate performance
  938. print("📊 Evaluating model on test set...")
  939. test_loss, test_acc, y_true, y_pred, y_prob = evaluate_on_loader(
  940. model, test_loader, criterion, device=DEVICE)
  941. print(f"✅ {name} Test Loss: {test_loss:.4f} | Accuracy: {test_acc:.4f}")
  942. # 2️⃣ Save evaluation reports
  943. df_cm, report_df = save_evaluation_reports(y_true, y_pred, y_prob, class_names, eval_out_dir)
  944. # 3️⃣ Save misclassified samples
  945. df_mis = save_misclassified(y_true, y_pred, class_names, eval_out_dir, df_test=df_test)
  946. # 4️⃣ Plot t-SNE of learned features
  947. plot_tsne_features(model, test_loader, class_names, eval_out_dir,
  948. n_samples=TSNE_SAMPLES, device=DEVICE)
  949. print(f"✅ Evaluation completed for {name} model. Reports saved to {eval_out_dir}\n")
  950. evaluate_all_models(
  951. best_models=best_models,
  952. train_loader=train_loader,
  953. val_loader=val_loader,
  954. test_loader=test_loader,
  955. RESULTS_DIR=RESULTS_DIR,
  956. class_names=class_names,
  957. TSNE_SAMPLES=2000,
  958. DEVICE=DEVICE,
  959. df_test=df_test # optional, if you have test dataframe with image paths
  960. )
  961. def plot_samples_per_class_fixed_cols(model, df_test, class_names, out_dir,
  962. samples_per_class=2, device='cuda',
  963. random_state=42, img_size=224,
  964. images_per_row=5):
  965. """
  966. Shows test images with true/pred labels, arranged with a fixed number of images per row.
  967. Handles models that return a tuple (logits, features) or just logits.
  968. """
  969. os.makedirs(out_dir, exist_ok=True)
  970. model.eval()
  971. model.to(device)
  972. preprocess = transforms.Compose([
  973. transforms.Resize((img_size, img_size)),
  974. transforms.ToTensor(),
  975. transforms.Normalize([0.485,0.456,0.406],[0.229,0.224,0.225])
  976. ])
  977. # Select samples
  978. np.random.seed(random_state)
  979. selected_dfs = []
  980. for cls_idx in range(len(class_names)):
  981. class_subset = df_test[df_test['label'] == cls_idx]
  982. selected = class_subset.sample(
  983. n=min(samples_per_class, len(class_subset)), random_state=random_state)
  984. selected_dfs.append(selected)
  985. sample_df = pd.concat(selected_dfs, ignore_index=True)
  986. total_samples = len(sample_df)
  987. cols = images_per_row
  988. rows = int(np.ceil(total_samples / cols))
  989. fig, axes = plt.subplots(rows, cols, figsize=(cols * 3.5, rows * 3.5))
  990. axes = np.array(axes).reshape(rows, cols)
  991. for idx, (_, row) in enumerate(sample_df.iterrows()):
  992. r = idx // cols
  993. c = idx % cols
  994. img_path = row['path']
  995. true_label = class_names[int(row['label'])]
  996. try:
  997. img = Image.open(img_path).convert("RGB")
  998. img_tensor = preprocess(img).unsqueeze(0).to(device)
  999. with torch.no_grad():
  1000. output = model(img_tensor)
  1001. if isinstance(output, tuple):
  1002. output = output[0]
  1003. prob = torch.softmax(output, dim=1)
  1004. pred = torch.argmax(prob, dim=1).item()
  1005. pred_label = class_names[pred]
  1006. conf = prob[0, pred].item() * 100
  1007. axes[r, c].imshow(img)
  1008. color = 'green' if pred_label == true_label else 'red'
  1009. axes[r, c].set_title(f"{true_label}\n→ {pred_label} ({conf:.1f}%)",
  1010. fontsize=14, color=color)
  1011. axes[r, c].axis('off')
  1012. except Exception as e:
  1013. axes[r, c].set_title("Error", fontsize=8)
  1014. axes[r, c].axis('off')
  1015. print(f"⚠️ Could not load image {img_path}: {e}")
  1016. # Turn off remaining empty axes
  1017. for idx in range(total_samples, rows * cols):
  1018. r = idx // cols
  1019. c = idx % cols
  1020. axes[r, c].axis('off')
  1021. plt.tight_layout()
  1022. save_path = os.path.join(out_dir, f"samples_per_class_{samples_per_class}_4perrow.png")
  1023. plt.savefig(save_path, dpi=300)
  1024. plt.close()
  1025. print(f"✅ Saved plot → {save_path}")
  1026. for model_key, model_instance in best_models.items():
  1027. model_out_dir = os.path.join(RESULTS_DIR, "evaluation", model_key)
  1028. os.makedirs(model_out_dir, exist_ok=True)
  1029. print(f"\n📊 Plotting 6x5 sample images for model: {model_key}")
  1030. plot_samples_per_class_fixed_cols(
  1031. model=model_instance,
  1032. df_test=df_test,
  1033. class_names=class_names,
  1034. out_dir=model_out_dir,
  1035. samples_per_class=2,
  1036. device='cuda',
  1037. images_per_row=5
  1038. )
  1039. import os
  1040. import numpy as np
  1041. import pandas as pd
  1042. import torch
  1043. import torch.nn as nn
  1044. from sklearn.metrics import confusion_matrix, cohen_kappa_score
  1045. import matplotlib.pyplot as plt
  1046. from PIL import Image
  1047. from torchvision import transforms
  1048. def evaluate_full_pipeline(model, model_name, loader, df_test, class_names, out_dir,
  1049. device='cuda', criterion=None, samples_per_class=2, IMG_SIZE=224):
  1050. """
  1051. Full evaluation pipeline:
  1052. ✅ Accuracy, Loss, Per-class metrics, Macro/Micro averages
  1053. ✅ Confusion matrix & classification report CSV
  1054. ✅ Misclassified samples & bar plot
  1055. ✅ Visualization: samples_per_class per class
  1056. ✅ Robust handling of models returning (logits, features)
  1057. """
  1058. os.makedirs(out_dir, exist_ok=True)
  1059. model.eval()
  1060. model.to(device)
  1061. if criterion is None:
  1062. criterion = nn.CrossEntropyLoss()
  1063. y_true_list, y_pred_list, y_prob_list = [], [], []
  1064. total_loss, total_samples = 0.0, 0
  1065. # -----------------------------
  1066. # 1️⃣ Inference over test set
  1067. # -----------------------------
  1068. with torch.no_grad():
  1069. for imgs, labels in loader:
  1070. imgs, labels = imgs.to(device), labels.to(device)
  1071. out = model(imgs)
  1072. if isinstance(out, (tuple, list)):
  1073. out = out[0] # take logits only
  1074. probs = torch.softmax(out, dim=1)
  1075. preds = torch.argmax(probs, dim=1)
  1076. loss = criterion(out, labels)
  1077. total_loss += loss.item() * imgs.size(0)
  1078. total_samples += imgs.size(0)
  1079. y_true_list.extend(labels.cpu().numpy())
  1080. y_pred_list.extend(preds.cpu().numpy())
  1081. y_prob_list.extend(probs.cpu().numpy())
  1082. y_true = np.array(y_true_list)
  1083. y_pred = np.array(y_pred_list)
  1084. y_prob = np.array(y_prob_list)
  1085. avg_loss = total_loss / total_samples
  1086. acc = (y_true == y_pred).mean()
  1087. # -----------------------------
  1088. # 2️⃣ Confusion Matrix
  1089. # -----------------------------
  1090. cm = confusion_matrix(y_true, y_pred)
  1091. cm_df = pd.DataFrame(cm, index=class_names, columns=class_names)
  1092. cm_df.to_csv(os.path.join(out_dir, "confusion_matrix.csv"))
  1093. # -----------------------------
  1094. # 3️⃣ Per-Class Metrics
  1095. # -----------------------------
  1096. per_class_metrics = []
  1097. for i, cname in enumerate(class_names):
  1098. TP = cm[i, i]
  1099. FP = cm[:, i].sum() - TP
  1100. FN = cm[i, :].sum() - TP
  1101. TN = cm.sum() - (TP + FP + FN)
  1102. acc_cls = (TP + TN) / cm.sum() if cm.sum() > 0 else 0
  1103. prec = TP / (TP + FP) if (TP + FP) > 0 else 0
  1104. rec = TP / (TP + FN) if (TP + FN) > 0 else 0
  1105. f1 = 2 * prec * rec / (prec + rec) if (prec + rec) > 0 else 0
  1106. dice = 2 * TP / (2 * TP + FP + FN) if (2 * TP + FP + FN) > 0 else 0
  1107. spec = TN / (TN + FP) if (TN + FP) > 0 else 0
  1108. fdr = FP / (TP + FP) if (TP + FP) > 0 else 0
  1109. forr = FN / (FN + TN) if (FN + TN) > 0 else 0
  1110. y_true_bin = (y_true == i).astype(int)
  1111. y_pred_bin = (y_pred == i).astype(int)
  1112. kappa = cohen_kappa_score(y_true_bin, y_pred_bin)
  1113. per_class_metrics.append([cname, acc_cls, prec, rec, f1, dice, spec, fdr, forr, kappa])
  1114. metrics_df = pd.DataFrame(
  1115. per_class_metrics,
  1116. columns=["Class", "Accuracy", "Precision", "Recall", "F1", "Dice",
  1117. "Specificity", "FDR", "FOR", "Kappa"]
  1118. )
  1119. # Macro & Micro averages
  1120. macro_avg = metrics_df.iloc[:, 1:].mean().to_dict()
  1121. micro_avg = {
  1122. "Accuracy": acc,
  1123. "Precision": np.mean(metrics_df["Precision"]),
  1124. "Recall": np.mean(metrics_df["Recall"]),
  1125. "F1": np.mean(metrics_df["F1"]),
  1126. "Dice": np.mean(metrics_df["Dice"]),
  1127. "Specificity": np.mean(metrics_df["Specificity"]),
  1128. "FDR": np.mean(metrics_df["FDR"]),
  1129. "FOR": np.mean(metrics_df["FOR"]),
  1130. "Kappa": np.mean(metrics_df["Kappa"]),
  1131. }
  1132. metrics_df = pd.concat([
  1133. metrics_df,
  1134. pd.DataFrame([["Macro Avg"] + list(macro_avg.values())], columns=metrics_df.columns),
  1135. pd.DataFrame([["Micro Avg"] + list(micro_avg.values())], columns=metrics_df.columns)
  1136. ], ignore_index=True)
  1137. metrics_df.to_csv(os.path.join(out_dir, "classification_report_with_kappa.csv"), index=False)
  1138. # -----------------------------
  1139. # 4️⃣ Misclassified Samples
  1140. # -----------------------------
  1141. mis_idx = np.where(y_true != y_pred)[0]
  1142. mis_list = [(df_test.iloc[i]['path'], class_names[y_true[i]], class_names[y_pred[i]]) for i in mis_idx]
  1143. df_mis = pd.DataFrame(mis_list, columns=["image_path", "true_label", "pred_label"])
  1144. df_mis.to_csv(os.path.join(out_dir, "misclassified.csv"), index=False)
  1145. # Misclassification bar plot
  1146. mis_counts = df_mis['true_label'].value_counts().reindex(class_names, fill_value=0)
  1147. plt.figure(figsize=(8, 5))
  1148. bars = plt.bar(class_names, mis_counts.values, color='tomato', edgecolor='black')
  1149. plt.title("Misclassified Samples per Class", fontsize=14)
  1150. plt.ylabel("Count")
  1151. for bar in bars:
  1152. plt.text(bar.get_x() + bar.get_width()/2, bar.get_height(),
  1153. f"{int(bar.get_height())}", ha='center', va='bottom')
  1154. plt.tight_layout()
  1155. plt.savefig(os.path.join(out_dir, "misclassified_barplot.png"), dpi=300)
  1156. plt.close()
  1157. # -----------------------------
  1158. # 5️⃣ Save Predictions
  1159. # -----------------------------
  1160. np.save(os.path.join(out_dir, "y_true.npy"), y_true)
  1161. np.save(os.path.join(out_dir, "y_pred.npy"), y_pred)
  1162. np.save(os.path.join(out_dir, "y_prob.npy"), y_prob)
  1163. # -----------------------------
  1164. # 7️⃣ Summary
  1165. # -----------------------------
  1166. print(f"\n✅ Model: {model_name}")
  1167. print(f" • Accuracy: {acc:.4f}")
  1168. print(f" • Avg Loss: {avg_loss:.4f}")
  1169. print(f" • Confusion matrix, metrics, misclassified samples saved to: {out_dir}")
  1170. return {
  1171. "accuracy": acc,
  1172. "avg_loss": avg_loss,
  1173. "confusion_matrix": cm,
  1174. "metrics_df": metrics_df,
  1175. "misclassified_df": df_mis
  1176. }
  1177. # =========================
  1178. # Example usage for multiple models
  1179. # =========================
  1180. #RESULTS_DIR = "results/evaluation_all_models"
  1181. for model_key, model_instance in best_models.items():
  1182. print(f"\n📊 Evaluating model: {model_key}")
  1183. model_out_dir = os.path.join(RESULTS_DIR, model_key)
  1184. os.makedirs(model_out_dir, exist_ok=True)
  1185. eval_results = evaluate_full_pipeline(
  1186. model=model_instance,
  1187. model_name=model_key,
  1188. loader=test_loader,
  1189. df_test=df_test,
  1190. class_names=class_names,
  1191. out_dir=model_out_dir,
  1192. device='cuda', # or 'cpu'
  1193. samples_per_class=2,
  1194. IMG_SIZE=224
  1195. )
  1196. import os
  1197. import torch
  1198. import torch.nn as nn
  1199. import numpy as np
  1200. import pandas as pd
  1201. import matplotlib.pyplot as plt
  1202. from torch.utils.data import DataLoader, TensorDataset
  1203. from PIL import Image
  1204. import cv2
  1205. def add_gaussian_noise_img(pil_img, std=0.01):
  1206. """
  1207. pil_img: PIL.Image RGB
  1208. std: standard deviation in [0,1] relative pixel range
  1209. returns: PIL.Image RGB with Gaussian noise
  1210. """
  1211. arr = np.array(pil_img).astype(np.float32) / 255.0 # HxWx3 in [0,1]
  1212. noise = np.random.normal(loc=0.0, scale=std, size=arr.shape).astype(np.float32)
  1213. noisy = np.clip(arr + noise, 0.0, 1.0)
  1214. noisy_uint8 = (noisy * 255.0).astype(np.uint8)
  1215. return Image.fromarray(noisy_uint8)
  1216. def downscale_image_img(pil_img, scale=0.9, interp_down=cv2.INTER_LINEAR, interp_up=cv2.INTER_LINEAR):
  1217. """
  1218. Simulate decreased resolution by downscaling and upscaling back to original size.
  1219. scale: fraction to reduce by (e.g., 0.8 -> reduce to 80% then upsample)
  1220. """
  1221. arr = np.array(pil_img)
  1222. h, w = arr.shape[:2]
  1223. new_h, new_w = max(1, int(h * scale)), max(1, int(w * scale))
  1224. # downscale
  1225. small = cv2.resize(arr, (new_w, new_h), interpolation=interp_down)
  1226. # upscale back to original
  1227. up = cv2.resize(small, (w, h), interpolation=interp_up)
  1228. return Image.fromarray(up)
  1229. def evaluate_on_loader(model, loader, criterion, device=torch.device('cpu')):
  1230. """
  1231. Minimal evaluation: returns (avg_loss, accuracy, y_true, y_pred, y_prob)
  1232. """
  1233. model.to(device)
  1234. model.eval()
  1235. y_true, y_pred, y_prob = [], [], []
  1236. total_loss = 0.0
  1237. with torch.no_grad():
  1238. for imgs, labels in loader:
  1239. imgs = imgs.to(device)
  1240. labels = labels.to(device)
  1241. out, _ = model(imgs) if isinstance(model(imgs), tuple) else (model(imgs), None)
  1242. loss = criterion(out, labels)
  1243. total_loss += loss.item() * imgs.size(0)
  1244. probs = torch.nn.functional.softmax(out, dim=1).cpu().numpy()
  1245. preds = np.argmax(probs, axis=1)
  1246. y_true.extend(labels.cpu().numpy())
  1247. y_pred.extend(preds)
  1248. y_prob.extend(probs)
  1249. n = len(loader.dataset)
  1250. avg_loss = total_loss / n if n > 0 else 0.0
  1251. acc = (np.array(y_true) == np.array(y_pred)).mean() if len(y_true) > 0 else 0.0
  1252. return avg_loss, acc, np.array(y_true), np.array(y_pred), np.array(y_prob)
  1253. def robustness_evaluation(
  1254. model, df_test, out_dir, val_transform,
  1255. batch_size=32, device=None, noise_sigmas=None, resolution_scales=None
  1256. ):
  1257. if device is None:
  1258. device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
  1259. os.makedirs(out_dir, exist_ok=True)
  1260. criterion = nn.CrossEntropyLoss()
  1261. labels = df_test['label'].astype(int).to_numpy()
  1262. # default ranges if not provided
  1263. if noise_sigmas is None:
  1264. noise_sigmas = [0.0, 0.005, 0.01, 0.02, 0.03, 0.05, 0.08, 0.1] # include 0.0 (clean)
  1265. if resolution_scales is None:
  1266. # scales are fraction of original; we will present as percentages removed = (1-scale)
  1267. resolution_scales = [1.0, 0.95, 0.90, 0.85, 0.80, 0.75, 0.70, 0.60, 0.50]
  1268. gaussian_results = {}
  1269. resolution_results = {}
  1270. model.to(device)
  1271. model.eval()
  1272. # ---------- Gaussian Noise robustness ----------
  1273. for sigma in noise_sigmas:
  1274. imgs_tensors = []
  1275. for p in df_test['path']:
  1276. img = Image.open(p).convert("RGB")
  1277. if sigma > 0:
  1278. noisy = add_gaussian_noise_img(img, std=sigma)
  1279. else:
  1280. noisy = img
  1281. imgs_tensors.append(val_transform(noisy))
  1282. imgs_stack = torch.stack(imgs_tensors)
  1283. lbls = torch.tensor(labels, dtype=torch.long)
  1284. loader = DataLoader(TensorDataset(imgs_stack, lbls), batch_size=batch_size, shuffle=False)
  1285. _, acc, _, _, _ = evaluate_on_loader(model, loader, criterion, device=device)
  1286. gaussian_results[f"{sigma:.3f}"] = acc
  1287. print(f"Noise σ={sigma:.3f} -> acc={acc:.4f}")
  1288. # ---------- Resolution Degradation robustness ----------
  1289. for scale in resolution_scales:
  1290. imgs_tensors = []
  1291. for p in df_test['path']:
  1292. img = Image.open(p).convert("RGB")
  1293. if scale < 1.0:
  1294. lowres = downscale_image_img(img, scale=scale)
  1295. else:
  1296. lowres = img
  1297. imgs_tensors.append(val_transform(lowres))
  1298. imgs_stack = torch.stack(imgs_tensors)
  1299. lbls = torch.tensor(labels, dtype=torch.long)
  1300. loader = DataLoader(TensorDataset(imgs_stack, lbls), batch_size=batch_size, shuffle=False)
  1301. _, acc, _, _, _ = evaluate_on_loader(model, loader, criterion, device=device)
  1302. pct = int(round((1.0 - scale) * 100)) # e.g., scale=0.9 -> 10
  1303. resolution_results[f"{pct}%"] = acc
  1304. print(f"Resolution ↓{pct}% (scale={scale:.2f}) -> acc={acc:.4f}")
  1305. # Save CSV results
  1306. rows = []
  1307. for k, v in gaussian_results.items():
  1308. rows.append(("Noise_"+k, v))
  1309. for k, v in resolution_results.items():
  1310. rows.append(("Resolution_"+k, v))
  1311. df_res = pd.DataFrame(rows, columns=["Condition", "Accuracy"])
  1312. csv_path = os.path.join(out_dir, "robustness_results_proposed.csv")
  1313. df_res.to_csv(csv_path, index=False)
  1314. print(f"✅ Results saved → {csv_path}")
  1315. # ---------- PLOTTING ----------
  1316. # color palette
  1317. max_len = max(len(gaussian_results), len(resolution_results))
  1318. colors = plt.cm.tab10(np.linspace(0, 1, max_len))
  1319. # ---- Plot A: Gaussian Noise (line + markers + shaded area) ----
  1320. # Convert keys to floats sorted
  1321. sig_items = sorted([(float(k), v) for k, v in gaussian_results.items()], key=lambda x: x[0])
  1322. sig_vals = [s for s, _ in sig_items]
  1323. acc_noise = np.array([a for _, a in sig_items])
  1324. # small pseudo-std to visualize band (replace with real repeated-run std if available)
  1325. acc_noise_std = np.maximum(0.002, 0.02 * (1.0 - acc_noise)) # larger band when accuracy drops
  1326. fig, ax = plt.subplots(figsize=(9, 5))
  1327. ax.plot(sig_vals, acc_noise, marker='o', linewidth=2, label='Accuracy', color='tab:blue')
  1328. ax.fill_between(sig_vals, acc_noise - acc_noise_std, acc_noise + acc_noise_std, alpha=0.2, color='tab:blue')
  1329. for i, (s, acc_v) in enumerate(zip(sig_vals, acc_noise)):
  1330. ax.scatter(s, acc_v, s=70, color=colors[i % len(colors)])
  1331. ax.text(s, acc_v + 0.02, f"{acc_v:.2f}", ha='center', fontsize=9)
  1332. ax.set_xlabel("Gaussian Noise σ")
  1333. ax.set_ylabel("Accuracy")
  1334. ax.set_title("Model Robustness vs Gaussian Noise")
  1335. ax.set_ylim(0, 1.02)
  1336. ax.grid(alpha=0.35)
  1337. plt.tight_layout()
  1338. noise_png = os.path.join(out_dir, "robustness_noise_lineplot.png")
  1339. fig.savefig(noise_png, dpi=300)
  1340. plt.close(fig)
  1341. print(f"✅ Saved → {noise_png}")
  1342. # ---- Plot B: Resolution (line + markers + shaded area) ----
  1343. res_items = list(resolution_results.items()) # keys like '10%', '20%', ...
  1344. # preserve order inserted (which corresponds to ascending degradation)
  1345. res_labels = [k for k, _ in res_items]
  1346. acc_res = np.array([v for _, v in res_items])
  1347. acc_res_std = np.maximum(0.002, 0.02 * (1.0 - acc_res))
  1348. fig, ax = plt.subplots(figsize=(9, 5))
  1349. x = np.arange(len(res_labels))
  1350. ax.plot(x, acc_res, marker='s', linewidth=2, label='Accuracy', color='tab:green')
  1351. ax.fill_between(x, acc_res - acc_res_std, acc_res + acc_res_std, alpha=0.2, color='tab:green')
  1352. for i, (lbl, acc_v) in enumerate(zip(res_labels, acc_res)):
  1353. ax.scatter(i, acc_v, s=80, color=colors[i % len(colors)])
  1354. ax.text(i, acc_v + 0.02, f"{acc_v:.2f}", ha='center', fontsize=9)
  1355. ax.set_xticks(x)
  1356. ax.set_xticklabels(res_labels)
  1357. ax.set_xlabel("Resolution Decrease")
  1358. ax.set_ylabel("Accuracy")
  1359. ax.set_title("Model Robustness vs Resolution Degradation")
  1360. ax.set_ylim(0, 1.02)
  1361. ax.grid(alpha=0.35)
  1362. plt.tight_layout()
  1363. res_png = os.path.join(out_dir, "robustness_resolution_lineplot.png")
  1364. fig.savefig(res_png, dpi=300)
  1365. plt.close(fig)
  1366. print(f"✅ Saved → {res_png}")
  1367. # ---- Plot C: Grouped bar charts (noise + resolution) ----
  1368. fig, axs = plt.subplots(1, 2, figsize=(14, 5))
  1369. # Noise bars
  1370. axs[0].bar([f"{s:.3f}" for s in sig_vals], acc_noise, color=colors[:len(sig_vals)])
  1371. axs[0].set_title("Gaussian Noise (bar)")
  1372. axs[0].set_ylim(0, 1)
  1373. axs[0].set_xlabel("σ")
  1374. axs[0].set_ylabel("Accuracy")
  1375. axs[0].grid(axis='y', linestyle='--', alpha=0.4)
  1376. for i, v in enumerate(acc_noise):
  1377. axs[0].text(i, v + 0.01, f"{v:.2f}", ha='center', fontsize=9)
  1378. # Resolution bars
  1379. axs[1].bar(res_labels, acc_res, color=colors[:len(res_labels)])
  1380. axs[1].set_title("Resolution decrease (bar)")
  1381. axs[1].set_ylim(0, 1)
  1382. axs[1].set_xlabel("↓%")
  1383. axs[1].grid(axis='y', linestyle='--', alpha=0.4)
  1384. for i, v in enumerate(acc_res):
  1385. axs[1].text(i, v + 0.01, f"{v:.2f}", ha='center', fontsize=9)
  1386. plt.tight_layout()
  1387. bar_png = os.path.join(out_dir, "robustness_barplots_colorful.png")
  1388. fig.savefig(bar_png, dpi=300)
  1389. plt.close(fig)
  1390. print(f"✅ Saved → {bar_png}")
  1391. return df_res
  1392. for tag, model in best_models.items():
  1393. print(f"\n==============================")
  1394. print(f"🔍 Robustness evaluation for model: {tag}")
  1395. print(f"==============================")
  1396. out_dir = os.path.join(RESULTS_DIR, f"robustness_{tag}")
  1397. df_res = robustness_evaluation(
  1398. model=model,
  1399. df_test=df_test, # DataFrame with ['path', 'label']
  1400. out_dir=out_dir,
  1401. val_transform=val_transform,
  1402. batch_size=32,
  1403. device=DEVICE,
  1404. noise_sigmas=[0.0, 0.005, 0.01, 0.02, 0.05, 0.1],
  1405. resolution_scales=[1.0, 0.90, 0.80, 0.70, 0.60, 0.50]
  1406. )
  1407. import numpy as np
  1408. import matplotlib.pyplot as plt
  1409. from sklearn.metrics import accuracy_score
  1410. def bootstrap_metric(y_true, y_pred, metric_fn, n_bootstrap=1000, alpha=0.05, plot=True, out_path=None):
  1411. rng = np.random.default_rng()
  1412. metrics = []
  1413. n_samples = len(y_true)
  1414. for _ in range(n_bootstrap):
  1415. idx = rng.choice(n_samples, n_samples, replace=True)
  1416. y_true_sample = y_true[idx]
  1417. y_pred_sample = y_pred[idx]
  1418. metrics.append(metric_fn(y_true_sample, y_pred_sample))
  1419. metrics = np.array(metrics)
  1420. metric_mean = metrics.mean()
  1421. std = metrics.std(ddof=1)
  1422. sem = std / np.sqrt(n_bootstrap)
  1423. ci_lower = np.percentile(metrics, 100*alpha/2)
  1424. ci_upper = np.percentile(metrics, 100*(1-alpha/2))
  1425. results = {
  1426. "mean": metric_mean,
  1427. "std": std,
  1428. "sem": sem,
  1429. "ci_lower": ci_lower,
  1430. "ci_upper": ci_upper,
  1431. "metrics": metrics
  1432. }
  1433. if plot:
  1434. plt.figure(figsize=(8,6))
  1435. plt.hist(metrics, bins=30, color='skyblue', edgecolor='black', alpha=0.7)
  1436. plt.axvline(metric_mean, color='red', linestyle='--', label=f"Mean = {metric_mean:.3f}")
  1437. plt.axvline(ci_lower, color='green', linestyle='--', label=f"95% CI Lower = {ci_lower:.3f}")
  1438. plt.axvline(ci_upper, color='orange', linestyle='--', label=f"95% CI Upper = {ci_upper:.3f}")
  1439. plt.title("Bootstrapped Accuracy Distribution")
  1440. plt.xlabel("Accuracy")
  1441. plt.ylabel("Frequency")
  1442. plt.legend()
  1443. plt.tight_layout()
  1444. if out_path:
  1445. plt.savefig(out_path, dpi=300)
  1446. print(f"✅ Histogram saved at: {out_path}")
  1447. plt.show()
  1448. return results
  1449. # ===========================
  1450. # Example usage with dummy data
  1451. # ===========================
  1452. results_acc = bootstrap_metric(
  1453. y_true=y_true,
  1454. y_pred=y_pred,
  1455. metric_fn=accuracy_score,
  1456. n_bootstrap=1000,
  1457. alpha=0.05,
  1458. plot=True,
  1459. out_path="bootstrapped_accuracy_hist.png"
  1460. )
  1461. print("Bootstrapped Accuracy Metrics:")
  1462. print(results_acc)
  1463. # %%
  1464. import os
  1465. import torch
  1466. import torch.nn as nn
  1467. import numpy as np
  1468. import pandas as pd
  1469. import shap
  1470. from sklearn.metrics import accuracy_score, confusion_matrix, classification_report, cohen_kappa_score, roc_curve, auc, precision_recall_curve
  1471. from pytorch_grad_cam import GradCAM
  1472. from pytorch_grad_cam.utils.image import show_cam_on_image
  1473. from PIL import Image
  1474. import cv2
  1475. import matplotlib.pyplot as plt
  1476. def evaluate_model_full(
  1477. model, model_name, loader, df_test, class_names, val_transform,
  1478. results_dir, shap_max_samples=50, n_gradcam=8
  1479. ):
  1480. model_out_dir = os.path.join(results_dir, "evaluation", model_name)
  1481. os.makedirs(model_out_dir, exist_ok=True)
  1482. print(f"\n🧪 Evaluating model: {model_name}")
  1483. # ------------------- Standard evaluation -------------------
  1484. criterion = nn.CrossEntropyLoss()
  1485. model.eval()
  1486. y_true, y_pred, y_prob = [], [], []
  1487. total_loss = 0.0
  1488. with torch.no_grad():
  1489. for imgs, labels in loader:
  1490. imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)
  1491. out, _ = model(imgs)
  1492. loss = criterion(out, labels)
  1493. total_loss += loss.item() * imgs.size(0)
  1494. probs = torch.nn.functional.softmax(out, dim=1).cpu().numpy()
  1495. preds = np.argmax(probs, axis=1)
  1496. y_true.extend(labels.cpu().numpy())
  1497. y_pred.extend(preds)
  1498. y_prob.extend(probs)
  1499. y_true = np.array(y_true)
  1500. y_pred = np.array(y_pred)
  1501. y_prob = np.array(y_prob)
  1502. avg_loss = total_loss / len(loader.dataset)
  1503. acc = accuracy_score(y_true, y_pred)
  1504. print(f"✅ Test Accuracy: {acc:.4f}, Avg Loss: {avg_loss:.4f}")
  1505. # ------------------- Confusion matrix & classification report -------------------
  1506. cm = confusion_matrix(y_true, y_pred)
  1507. pd.DataFrame(cm, index=class_names, columns=class_names).to_csv(
  1508. os.path.join(model_out_dir, "confusion_matrix.csv")
  1509. )
  1510. cls_report_dict = classification_report(y_true, y_pred, labels=list(range(len(class_names))),
  1511. output_dict=True, zero_division=0)
  1512. cls_df = pd.DataFrame(cls_report_dict).T
  1513. cls_df = cls_df.loc[[str(i) for i in range(len(class_names))]].copy()
  1514. cls_df.index = class_names
  1515. # ------------------- Per-class Kappa & Dice -------------------
  1516. per_class_kappa, dice_scores = [], []
  1517. for i, cname in enumerate(class_names):
  1518. y_true_bin = (y_true == i).astype(int)
  1519. y_pred_bin = (y_pred == i).astype(int)
  1520. kappa = cohen_kappa_score(y_true_bin, y_pred_bin)
  1521. per_class_kappa.append(kappa)
  1522. TP = cm[i,i]
  1523. FP = cm[:,i].sum() - TP
  1524. FN = cm[i,:].sum() - TP
  1525. dice = 2 * TP / (2*TP + FP + FN) if (2*TP + FP + FN) > 0 else 0
  1526. dice_scores.append(dice)
  1527. cls_df["Kappa"] = per_class_kappa
  1528. cls_df["Dice"] = dice_scores
  1529. cls_df.to_csv(os.path.join(model_out_dir, "classification_report_with_kappa.csv"))
  1530. # ------------------- Misclassified samples -------------------
  1531. misclassified_out_dir = os.path.join(model_out_dir, "misclassified")
  1532. os.makedirs(misclassified_out_dir, exist_ok=True)
  1533. mis_idx = np.where(y_true != y_pred)[0]
  1534. mis_list = [(df_test.iloc[i]['path'], class_names[y_true[i]], class_names[y_pred[i]]) for i in mis_idx]
  1535. pd.DataFrame(mis_list, columns=["image_path","true_label","pred_label"]).to_csv(
  1536. os.path.join(misclassified_out_dir,"misclassified.csv"), index=False
  1537. )
  1538. print(f"✅ Misclassified samples saved to {misclassified_out_dir}")
  1539. # ------------------- ROC & PR curves -------------------
  1540. roc_out_dir = os.path.join(model_out_dir, "roc_pr")
  1541. os.makedirs(roc_out_dir, exist_ok=True)
  1542. for i, cname in enumerate(class_names):
  1543. fpr, tpr, _ = roc_curve((y_true==i).astype(int), y_prob[:,i])
  1544. roc_auc = auc(fpr, tpr)
  1545. plt.figure()
  1546. plt.plot(fpr, tpr, label=f'ROC curve (AUC={roc_auc:.2f})')
  1547. plt.plot([0,1],[0,1],'--',color='gray')
  1548. plt.xlabel("False Positive Rate")
  1549. plt.ylabel("True Positive Rate")
  1550. plt.title(f"ROC - {cname}")
  1551. plt.legend()
  1552. plt.savefig(os.path.join(roc_out_dir,f"roc_{cname}.png"))
  1553. plt.close()
  1554. precision, recall, _ = precision_recall_curve((y_true==i).astype(int), y_prob[:,i])
  1555. plt.figure()
  1556. plt.plot(recall, precision)
  1557. plt.xlabel("Recall")
  1558. plt.ylabel("Precision")
  1559. plt.title(f"PR Curve - {cname}")
  1560. plt.savefig(os.path.join(roc_out_dir,f"pr_{cname}.png"))
  1561. plt.close()
  1562. print(f"✅ ROC & PR curves saved to {roc_out_dir}")
  1563. # ------------------- Grad-CAM -------------------
  1564. gradcam_out_dir = os.path.join(model_out_dir, "gradcam_examples")
  1565. os.makedirs(gradcam_out_dir, exist_ok=True)
  1566. # Wrap model so Grad-CAM gets only logits
  1567. class GradCAMWrapper(nn.Module):
  1568. def __init__(self, base_model):
  1569. super().__init__()
  1570. self.base_model = base_model
  1571. def forward(self, x):
  1572. logits, _ = self.base_model(x)
  1573. return logits
  1574. gradcam_model = GradCAMWrapper(model)
  1575. target_layer = model.layer4[-1].conv3 # adjust for ResNet backbone
  1576. cam = GradCAM(model=gradcam_model, target_layers=[target_layer])
  1577. sample_rows = df_test.sample(n=min(n_gradcam, len(df_test)), random_state=42)
  1578. for idx, row in sample_rows.iterrows():
  1579. img_pil = Image.open(row['path']).convert("RGB")
  1580. img_tensor = val_transform(img_pil).unsqueeze(0).to(DEVICE)
  1581. img_tensor.requires_grad_(True)
  1582. grayscale_cam = cam(img_tensor)[0, :]
  1583. img_numpy = np.array(img_pil)/255.0
  1584. # Resize CAM to match image shape
  1585. grayscale_cam_resized = cv2.resize(grayscale_cam, (img_numpy.shape[1], img_numpy.shape[0]))
  1586. cam_image = show_cam_on_image(img_numpy, grayscale_cam_resized, use_rgb=True)
  1587. save_path = os.path.join(
  1588. gradcam_out_dir,
  1589. f"sample_{idx}_pred_{class_names[y_pred[idx]]}_true_{class_names[y_true[idx]]}.png"
  1590. )
  1591. cv2.imwrite(save_path, cv2.cvtColor(cam_image, cv2.COLOR_RGB2BGR))
  1592. print(f"✅ Grad-CAM examples saved to {gradcam_out_dir}")
  1593. '''
  1594. # ------------------- SHAP -------------------
  1595. class LogitsWrapper(nn.Module):
  1596. def __init__(self, model):
  1597. super().__init__()
  1598. self.model = model
  1599. def forward(self, x):
  1600. logits, _ = self.model(x)
  1601. return logits
  1602. shap_model = LogitsWrapper(model).to(DEVICE)
  1603. shap_model.eval()
  1604. imgs, _ = next(iter(loader))
  1605. imgs = imgs[:shap_max_samples].to(DEVICE)
  1606. explainer = shap.GradientExplainer(shap_model, imgs)
  1607. shap_values = explainer.shap_values(imgs)
  1608. imgs_np = np.transpose(imgs.cpu().numpy(), (0,2,3,1))
  1609. shap_plot_path = os.path.join(model_out_dir, "shap_summary.png")
  1610. shap.image_plot(shap_values, imgs_np)
  1611. print(f"✅ SHAP summary plot generated at {shap_plot_path}")
  1612. '''
  1613. return avg_loss, acc, y_true, y_pred, y_prob, cm, cls_df
  1614. # --- Run evaluation for both models ---
  1615. for model_key, model_instance in best_models.items():
  1616. evaluate_model_full(
  1617. model=model_instance,
  1618. model_name=model_key,
  1619. loader=test_loader,
  1620. df_test=df_test,
  1621. class_names=labels_map_inv,
  1622. val_transform=val_transform,
  1623. results_dir=RESULTS_DIR,
  1624. shap_max_samples=SHAP_MAX_SAMPLES,
  1625. n_gradcam=8
  1626. )
  1627. # %%
  1628. import os
  1629. import torch
  1630. import torch.nn as nn
  1631. import numpy as np
  1632. import shap
  1633. import matplotlib.pyplot as plt
  1634. import cv2
  1635. def plot_shap_for_model(model, model_name, loader, class_names, results_dir):
  1636. """
  1637. Generates SHAP heatmaps for one image per class and saves them.
  1638. Plots 2 classes per row (each class has Original / Overlay / Heatmap).
  1639. Works for models returning (logits, features) tuples.
  1640. """
  1641. DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
  1642. model.to(DEVICE)
  1643. model.eval()
  1644. # --- Create output directory ---
  1645. model_out_dir = os.path.join(results_dir, "shap_summary", model_name)
  1646. os.makedirs(model_out_dir, exist_ok=True)
  1647. # --- Select one image per class ---
  1648. imgs_list, labels_list = [], []
  1649. for c in range(len(class_names)):
  1650. for img_batch, label_batch in loader:
  1651. idx = (label_batch == c).nonzero(as_tuple=True)[0]
  1652. if len(idx) > 0:
  1653. imgs_list.append(img_batch[idx[0]])
  1654. labels_list.append(label_batch[idx[0]])
  1655. break
  1656. imgs = torch.stack(imgs_list).to(DEVICE) # N_classes x C x H x W
  1657. imgs_np = np.transpose(imgs.cpu().numpy(), (0, 2, 3, 1)) # N,H,W,C
  1658. # --- Wrap model for SHAP ---
  1659. class LogitsWrapper(nn.Module):
  1660. def __init__(self, model):
  1661. super().__init__()
  1662. self.model = model
  1663. def forward(self, x):
  1664. out = self.model(x)
  1665. if isinstance(out, (tuple, list)):
  1666. out = out[0]
  1667. return out
  1668. shap_model = LogitsWrapper(model).to(DEVICE)
  1669. shap_model.eval()
  1670. # --- Compute SHAP values ---
  1671. explainer = shap.GradientExplainer(shap_model, imgs)
  1672. shap_values = explainer.shap_values(imgs)
  1673. # --- Combine and normalize SHAP maps ---
  1674. if isinstance(shap_values, list):
  1675. shap_comb = np.sum([np.abs(sv) for sv in shap_values], axis=0)
  1676. else:
  1677. shap_comb = np.abs(shap_values)
  1678. shap_gray = np.sum(shap_comb, axis=-1)
  1679. shap_norm = np.zeros_like(shap_gray)
  1680. for i in range(shap_gray.shape[0]):
  1681. im = shap_gray[i]
  1682. im = im - im.min()
  1683. if im.max() > 0:
  1684. im = im / im.max()
  1685. shap_norm[i] = im
  1686. # --- Plot: 2 classes per row ---
  1687. N = shap_norm.shape[0]
  1688. sets_per_row = 2 # ✅ Two classes per row
  1689. cols_per_set = 3 # Original / Overlay / Heatmap
  1690. cols = sets_per_row * cols_per_set
  1691. rows = int(np.ceil(N / sets_per_row))
  1692. fig, axes = plt.subplots(rows, cols, figsize=(cols * 2.2, rows * 2.5))
  1693. # Ensure axes is 2D array
  1694. if rows == 1:
  1695. axes = axes[np.newaxis, :]
  1696. if axes.ndim == 1:
  1697. axes = axes[np.newaxis, :]
  1698. for i in range(N):
  1699. orig = imgs_np[i]
  1700. heat = (shap_norm[i] * 255).astype(np.uint8)
  1701. if heat.ndim == 3:
  1702. heat = heat.squeeze()
  1703. heat_color = cv2.applyColorMap(heat, cv2.COLORMAP_JET)
  1704. heat_rgb = cv2.cvtColor(heat_color, cv2.COLOR_BGR2RGB).astype(np.float32) / 255.0
  1705. if orig.shape != heat_rgb.shape:
  1706. heat_rgb = cv2.resize(heat_rgb, (orig.shape[1], orig.shape[0]))
  1707. overlay = np.clip(0.5 * orig + 0.5 * heat_rgb, 0, 1)
  1708. row = i // sets_per_row
  1709. col_start = (i % sets_per_row) * cols_per_set
  1710. axes[row, col_start + 0].imshow(orig)
  1711. #axes[row, col_start + 0].set_title(f"{class_names[i]} - Original")
  1712. axes[row, col_start + 0].set_title(f"{class_names[i]} - Original", fontsize=14)
  1713. axes[row, col_start + 0].axis('off')
  1714. axes[row, col_start + 1].imshow(overlay)
  1715. #axes[row, col_start + 1].set_title("Overlay (SHAP)")
  1716. axes[row, col_start + 1].set_title("Overlay (SHAP)", fontsize=14)
  1717. axes[row, col_start + 1].axis('off')
  1718. axes[row, col_start + 2].imshow(heat_rgb)
  1719. #axes[row, col_start + 2].set_title("Heatmap (JET)")
  1720. axes[row, col_start + 2].set_title("Heatmap (JET)", fontsize=14)
  1721. axes[row, col_start + 2].axis('off')
  1722. plt.tight_layout()
  1723. out_path = os.path.join(model_out_dir, "shap_summary_one_per_class_1.png")
  1724. fig.savefig(out_path, dpi=300, bbox_inches='tight')
  1725. plt.close(fig)
  1726. print(f"✅ Saved SHAP summary (one image per class) for model {model_name} at {out_path}")
  1727. # ===========================
  1728. # Run for all models
  1729. # ===========================
  1730. for model_key, model_instance in best_models.items():
  1731. plot_shap_for_model(
  1732. model=model_instance,
  1733. model_name=model_key,
  1734. loader=test_loader,
  1735. class_names=class_names,
  1736. results_dir=RESULTS_DIR
  1737. )
  1738. # %%
  1739. import os
  1740. import torch
  1741. import torch.nn as nn
  1742. import numpy as np
  1743. import shap
  1744. import matplotlib.pyplot as plt
  1745. import cv2
  1746. def plot_shap_for_model(model, model_name, loader, class_names, results_dir, extra_samples=3):
  1747. """
  1748. Generates SHAP heatmaps for one image per class + extra samples and saves them.
  1749. Plots 2 classes per row (each class has Original / Overlay / Heatmap).
  1750. Works for models returning (logits, features) tuples.
  1751. """
  1752. DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
  1753. model.to(DEVICE)
  1754. model.eval()
  1755. # --- Create output directory ---
  1756. model_out_dir = os.path.join(results_dir, "shap_summary", model_name)
  1757. os.makedirs(model_out_dir, exist_ok=True)
  1758. # --- Select one image per class ---
  1759. imgs_list, labels_list = [], []
  1760. for c in range(len(class_names)):
  1761. for img_batch, label_batch in loader:
  1762. idx = (label_batch == c).nonzero(as_tuple=True)[0]
  1763. if len(idx) > 0:
  1764. imgs_list.append(img_batch[idx[0]])
  1765. labels_list.append(label_batch[idx[0]])
  1766. break
  1767. # --- Add extra random samples (if available) ---
  1768. # Collect remaining images from loader
  1769. all_imgs, all_labels = [], []
  1770. for img_batch, label_batch in loader:
  1771. all_imgs.append(img_batch)
  1772. all_labels.append(label_batch)
  1773. all_imgs = torch.cat(all_imgs)
  1774. all_labels = torch.cat(all_labels)
  1775. # Choose random samples (excluding already used ones)
  1776. used_idxs = set([torch.where(all_labels == l)[0][0].item() for l in labels_list])
  1777. available_idxs = [i for i in range(len(all_imgs)) if i not in used_idxs]
  1778. if available_idxs:
  1779. extra_idxs = np.random.choice(available_idxs, size=min(extra_samples, len(available_idxs)), replace=False)
  1780. for i in extra_idxs:
  1781. imgs_list.append(all_imgs[i])
  1782. labels_list.append(all_labels[i])
  1783. imgs = torch.stack(imgs_list).to(DEVICE)
  1784. imgs_np = np.transpose(imgs.cpu().numpy(), (0, 2, 3, 1)) # N,H,W,C
  1785. # --- Wrap model for SHAP ---
  1786. class LogitsWrapper(nn.Module):
  1787. def __init__(self, model):
  1788. super().__init__()
  1789. self.model = model
  1790. def forward(self, x):
  1791. out = self.model(x)
  1792. if isinstance(out, (tuple, list)):
  1793. out = out[0]
  1794. return out
  1795. shap_model = LogitsWrapper(model).to(DEVICE)
  1796. shap_model.eval()
  1797. # --- Compute SHAP values ---
  1798. explainer = shap.GradientExplainer(shap_model, imgs)
  1799. shap_values = explainer.shap_values(imgs)
  1800. # --- Combine and normalize SHAP maps ---
  1801. if isinstance(shap_values, list):
  1802. shap_comb = np.sum([np.abs(sv) for sv in shap_values], axis=0)
  1803. else:
  1804. shap_comb = np.abs(shap_values)
  1805. shap_gray = np.sum(shap_comb, axis=-1)
  1806. shap_norm = np.zeros_like(shap_gray)
  1807. for i in range(shap_gray.shape[0]):
  1808. im = shap_gray[i]
  1809. im = im - im.min()
  1810. if im.max() > 0:
  1811. im = im / im.max()
  1812. shap_norm[i] = im
  1813. # --- Plot: 2 samples per row ---
  1814. N_total = shap_norm.shape[0]
  1815. sets_per_row = 2 # Two samples per row
  1816. cols_per_set = 3 # Original / Overlay / Heatmap
  1817. cols = sets_per_row * cols_per_set
  1818. rows = int(np.ceil(N_total / sets_per_row))
  1819. fig, axes = plt.subplots(rows, cols, figsize=(cols * 2.2, rows * 2.5))
  1820. # Ensure axes is 2D array
  1821. if rows == 1:
  1822. axes = axes[np.newaxis, :]
  1823. if axes.ndim == 1:
  1824. axes = axes[np.newaxis, :]
  1825. for i in range(N_total):
  1826. orig = imgs_np[i]
  1827. heat = (shap_norm[i] * 255).astype(np.uint8)
  1828. if heat.ndim == 3:
  1829. heat = heat.squeeze()
  1830. heat_color = cv2.applyColorMap(heat, cv2.COLORMAP_JET)
  1831. heat_rgb = cv2.cvtColor(heat_color, cv2.COLOR_BGR2RGB).astype(np.float32) / 255.0
  1832. if orig.shape != heat_rgb.shape:
  1833. heat_rgb = cv2.resize(heat_rgb, (orig.shape[1], orig.shape[0]))
  1834. overlay = np.clip(0.5 * orig + 0.5 * heat_rgb, 0, 1)
  1835. row = i // sets_per_row
  1836. col_start = (i % sets_per_row) * cols_per_set
  1837. title_name = f"{class_names[labels_list[i]]}" if labels_list[i] < len(class_names) else "Extra Sample"
  1838. axes[row, col_start + 0].imshow(orig)
  1839. axes[row, col_start + 0].set_title(f"{title_name} - Original", fontsize=14)
  1840. axes[row, col_start + 0].axis('off')
  1841. axes[row, col_start + 1].imshow(overlay)
  1842. axes[row, col_start + 1].set_title("Overlay (SHAP)", fontsize=14)
  1843. axes[row, col_start + 1].axis('off')
  1844. axes[row, col_start + 2].imshow(heat_rgb)
  1845. axes[row, col_start + 2].set_title("Heatmap (JET)", fontsize=14)
  1846. axes[row, col_start + 2].axis('off')
  1847. plt.tight_layout()
  1848. out_path = os.path.join(model_out_dir, f"shap_summary_one_per_class_plus{extra_samples}.png")
  1849. fig.savefig(out_path, dpi=300, bbox_inches='tight')
  1850. plt.close(fig)
  1851. print(f"✅ Saved SHAP summary ({len(class_names)} + {extra_samples} samples) for model {model_name} at {out_path}")
  1852. # ===========================
  1853. # Run for all models
  1854. # ===========================
  1855. for model_key, model_instance in best_models.items():
  1856. plot_shap_for_model(
  1857. model=model_instance,
  1858. model_name=model_key,
  1859. loader=test_loader,
  1860. class_names=class_names,
  1861. results_dir=RESULTS_DIR,
  1862. extra_samples=3 # 👈 Add 3 extra samples beyond 1 per class
  1863. )
  1864. # %%
  1865. def plot_random_classification_results(df_test, class_names, out_dir, num_samples=30):
  1866. os.makedirs(out_dir, exist_ok=True)
  1867. sample_df = df_test.sample(n=min(num_samples, len(df_test)))
  1868. fig, axes = plt.subplots(6, 5, figsize=(15, 15))
  1869. axes = axes.flatten()
  1870. for i, (idx, row) in enumerate(sample_df.iterrows()):
  1871. if i >= 30:
  1872. break
  1873. img_path = row['path']
  1874. true_label = class_names[int(row['label'])]
  1875. # Here we are simulating or using actual prediction if available
  1876. pred_label = random.choice(class_names) # Replace with actual prediction if needed
  1877. try:
  1878. img = Image.open(img_path).convert("RGB")
  1879. axes[i].imshow(img)
  1880. axes[i].set_title(f"True: {true_label}\nPred: {pred_label}", fontsize=14)
  1881. axes[i].axis('off')
  1882. except Exception as e:
  1883. print(f"Could not load image {img_path}: {e}")
  1884. axes[i].set_title("Error loading image", fontsize=14)
  1885. axes[i].axis('off')
  1886. # Hide any unused subplots
  1887. for j in range(i + 1, 30):
  1888. axes[j].axis('off')
  1889. plt.tight_layout()
  1890. plt.savefig(os.path.join(out_dir, "random_classification_results.png"), dpi=PLOT_DPI)
  1891. plt.close()
  1892. print(f"✅ Random classification results plot saved to {out_dir}")
  1893. # ------------------ FUNCTION CALL PER MODEL ------------------
  1894. for model_key, model_instance in best_models.items():
  1895. model_out_dir = os.path.join(RESULTS_DIR, "evaluation", model_key)
  1896. os.makedirs(model_out_dir, exist_ok=True)
  1897. print(f"\n📊 Plotting random classification results for model: {model_key}")
  1898. plot_random_classification_results(df_test, class_names, model_out_dir)
  1899. # %%
  1900. import os
  1901. import torch
  1902. import torch.nn as nn
  1903. import numpy as np
  1904. import shap
  1905. import matplotlib.pyplot as plt
  1906. import cv2
  1907. def plot_shap_for_model(model, model_name, loader, class_names, results_dir):
  1908. """
  1909. Generates SHAP heatmaps for one image per class and saves them.
  1910. Works for models returning (logits, features) tuples.
  1911. """
  1912. DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
  1913. model.to(DEVICE)
  1914. model.eval()
  1915. # --- Create output directory ---
  1916. model_out_dir = os.path.join(results_dir, "shap_summary", model_name)
  1917. os.makedirs(model_out_dir, exist_ok=True)
  1918. # --- Select one image per class ---
  1919. imgs_list, labels_list = [], []
  1920. for c in range(len(class_names)):
  1921. for img_batch, label_batch in loader:
  1922. idx = (label_batch == c).nonzero(as_tuple=True)[0]
  1923. if len(idx) > 0:
  1924. imgs_list.append(img_batch[idx[0]])
  1925. labels_list.append(label_batch[idx[0]])
  1926. break
  1927. imgs = torch.stack(imgs_list).to(DEVICE) # N_classes x C x H x W
  1928. imgs_np = np.transpose(imgs.cpu().numpy(), (0,2,3,1)) # N,H,W,C
  1929. # --- Wrap model for SHAP ---
  1930. class LogitsWrapper(nn.Module):
  1931. def __init__(self, model):
  1932. super().__init__()
  1933. self.model = model
  1934. def forward(self, x):
  1935. out = self.model(x)
  1936. if isinstance(out, (tuple, list)):
  1937. out = out[0]
  1938. return out
  1939. shap_model = LogitsWrapper(model).to(DEVICE)
  1940. shap_model.eval()
  1941. # --- Compute SHAP values ---
  1942. explainer = shap.GradientExplainer(shap_model, imgs)
  1943. shap_values = explainer.shap_values(imgs)
  1944. # --- Build combined attribution map ---
  1945. if isinstance(shap_values, list):
  1946. shap_comb = np.sum([np.abs(sv) for sv in shap_values], axis=0) # N,H,W,C
  1947. else:
  1948. shap_comb = np.abs(shap_values)
  1949. shap_gray = np.sum(shap_comb, axis=-1) # N,H,W
  1950. # --- Normalize each heatmap ---
  1951. shap_norm = np.zeros_like(shap_gray)
  1952. for i in range(shap_gray.shape[0]):
  1953. im = shap_gray[i]
  1954. im = im - im.min()
  1955. if im.max() > 0:
  1956. im = im / im.max()
  1957. shap_norm[i] = im
  1958. # --- Plot: original / overlay / heatmap ---
  1959. N = shap_norm.shape[0]
  1960. cols = 3
  1961. rows = N
  1962. fig, axes = plt.subplots(rows, cols, figsize=(cols*4, rows*4))
  1963. if N == 1:
  1964. axes = axes[np.newaxis, :]
  1965. for i in range(N):
  1966. orig = imgs_np[i]
  1967. heat = (shap_norm[i]*255).astype(np.uint8)
  1968. if heat.ndim == 3:
  1969. heat = heat.squeeze()
  1970. heat_color = cv2.applyColorMap(heat, cv2.COLORMAP_JET)
  1971. heat_rgb = cv2.cvtColor(heat_color, cv2.COLOR_BGR2RGB).astype(np.float32)/255.0
  1972. if orig.shape != heat_rgb.shape:
  1973. heat_rgb = cv2.resize(heat_rgb, (orig.shape[1], orig.shape[0]))
  1974. overlay = np.clip(0.5*orig + 0.5*heat_rgb, 0, 1)
  1975. axes[i,0].imshow(orig)
  1976. axes[i,0].set_title(f"{class_names[i]} - Original")
  1977. axes[i,0].axis('off')
  1978. axes[i,1].imshow(overlay)
  1979. axes[i,1].set_title("Overlay (SHAP)")
  1980. axes[i,1].axis('off')
  1981. axes[i,2].imshow(heat_rgb)
  1982. axes[i,2].set_title("Heatmap (JET)")
  1983. axes[i,2].axis('off')
  1984. plt.tight_layout()
  1985. out_path = os.path.join(model_out_dir, "shap_summary_one_per_class.png")
  1986. fig.savefig(out_path, dpi=300, bbox_inches='tight')
  1987. plt.close(fig)
  1988. print(f"✅ Saved SHAP summary (one image per class) for model {model_name} at {out_path}")
  1989. # ===========================
  1990. # Run for all models
  1991. # ===========================
  1992. for model_key, model_instance in best_models.items():
  1993. plot_shap_for_model(
  1994. model=model_instance,
  1995. model_name=model_key,
  1996. loader=test_loader,
  1997. class_names=class_names,
  1998. results_dir=RESULTS_DIR
  1999. )
  2000. # %%
  2001. import os
  2002. import torch
  2003. import numpy as np
  2004. import matplotlib.pyplot as plt
  2005. from sklearn.manifold import TSNE
  2006. from sklearn.decomposition import PCA
  2007. def plot_tsne_features_stable(
  2008. model,
  2009. loader,
  2010. class_names,
  2011. out_dir,
  2012. n_samples=300,
  2013. device="cuda",
  2014. perplexity=30,
  2015. seed=42
  2016. ):
  2017. """
  2018. Stable and reproducible t-SNE visualization of penultimate-layer features.
  2019. - Uses PCA (up to 50D) before t-SNE for speed & noise reduction.
  2020. - Fixes random seeds for reproducibility.
  2021. - Consistent layout and cluster separation.
  2022. """
  2023. # --- Reproducibility ---
  2024. np.random.seed(seed)
  2025. torch.manual_seed(seed)
  2026. # --- Extract features and labels ---
  2027. X, y = extract_penultimate_features(model, loader, device=device)
  2028. # --- Sample subset if needed ---
  2029. if len(X) > n_samples:
  2030. idx = np.random.choice(len(X), n_samples, replace=False)
  2031. X, y = X[idx], y[idx]
  2032. # --- PCA preprocessing (only if needed) ---
  2033. if X.shape[1] > 50:
  2034. pca = PCA(n_components=50, random_state=seed)
  2035. X_pca = pca.fit_transform(X)
  2036. else:
  2037. X_pca = X
  2038. # --- t-SNE embedding ---
  2039. tsne = TSNE(
  2040. n_components=2,
  2041. random_state=seed,
  2042. perplexity=perplexity,
  2043. init="pca", # ensures deterministic start
  2044. learning_rate="auto"
  2045. )
  2046. X_tsne = tsne.fit_transform(X_pca)
  2047. # --- Plot ---
  2048. plt.figure(figsize=(8, 6))
  2049. for cls in np.unique(y):
  2050. plt.scatter(
  2051. X_tsne[y == cls, 0],
  2052. X_tsne[y == cls, 1],
  2053. label=class_names[cls],
  2054. alpha=0.8,
  2055. s=40,
  2056. edgecolors="none"
  2057. )
  2058. #plt.legend(markerscale=1.5, fontsize=14)
  2059. #plt.legend(loc='lower center', bbox_to_anchor=(0.5, -0.05), ncol=5)
  2060. #plt.legend(loc='lower center', bbox_to_anchor=(0.5, -0.15), ncol=3)
  2061. #plt.title("t-SNE of Penultimate Features (Stable & Reproducible)")
  2062. #plt.tight_layout(rect=[0, 0.05, 1, 1]) # leaves space at bottom
  2063. plt.tight_layout()
  2064. fig, ax = plt.subplots(figsize=(8, 6))
  2065. for cls in np.unique(y):
  2066. ax.scatter(X_tsne[y == cls, 0], X_tsne[y == cls, 1],
  2067. label=class_names[cls], alpha=0.7, s=40)
  2068. # Put legend fully below plot
  2069. legend = ax.legend(loc='upper center',
  2070. bbox_to_anchor=(0.5, -0.12),
  2071. ncol=3, frameon=False)
  2072. ax.set_title("t-SNE of Penultimate Features (Stable & Reproducible)")
  2073. fig.tight_layout()
  2074. # ✅ This ensures legend is not cut off and doesn’t cover points
  2075. fig.savefig(os.path.join(out_dir, "tsne_features.png"),
  2076. dpi=300, bbox_extra_artists=(legend,), bbox_inches='tight')
  2077. plt.close(fig)
  2078. # --- Save ---
  2079. os.makedirs(out_dir, exist_ok=True)
  2080. out_path = os.path.join(out_dir, "tsne_features_stable.png")
  2081. plt.savefig(out_path, dpi=300, bbox_inches="tight")
  2082. plt.close()
  2083. print(f"✅ Saved stable t-SNE visualization at: {out_path}")
  2084. plot_tsne_features_stable(
  2085. model=best_models["best_acc"],
  2086. loader=test_loader,
  2087. class_names=class_names,
  2088. out_dir=RESULTS_DIR
  2089. )
  2090. # %%
  2091. # ======================== CLASS-WISE SAMPLE COUNTS ========================
  2092. print("\n📊 Samples per class in each split:")
  2093. def count_per_class(df_split, split_name):
  2094. counts = df_split['label'].value_counts().sort_index()
  2095. print(f"\n{split_name} split:")
  2096. for idx, count in counts.items():
  2097. print(f" {class_names[idx]:<25}: {count}")
  2098. return counts
  2099. train_counts = count_per_class(df_train, "Training")
  2100. val_counts = count_per_class(df_val, "Validation")
  2101. test_counts = count_per_class(df_test, "Testing")
  2102. # Optionally combine into one DataFrame for easy comparison
  2103. summary_df = pd.DataFrame({
  2104. "Class": class_names,
  2105. "Train": [train_counts.get(i, 0) for i in range(len(class_names))],
  2106. "Val": [val_counts.get(i, 0) for i in range(len(class_names))],
  2107. "Test": [test_counts.get(i, 0) for i in range(len(class_names))]
  2108. })
  2109. print("\n📋 Combined class distribution summary:")
  2110. print(summary_df)
  2111. # Save to CSV
  2112. summary_df.to_csv(os.path.join(RESULTS_DIR, "class_distribution_summary.csv"), index=False)
  2113. print(f"\n✅ Saved class distribution summary to {os.path.join(RESULTS_DIR, 'class_distribution_summary.csv')}")
  2114. # %%
  2115. def zip_results(out_dir, zip_name="final_results.zip"):
  2116. with zipfile.ZipFile(zip_name, "w", zipfile.ZIP_DEFLATED) as zf:
  2117. for root, _, files in os.walk(out_dir):
  2118. for f in files:
  2119. fp = os.path.join(root, f)
  2120. zf.write(fp, os.path.relpath(fp, out_dir))
  2121. print(f"✅ All results zipped to {zip_name}")
  2122. # --- Zip final results ---
  2123. zip_results(RESULTS_DIR, "final_results.zip")
  2124. print("✅ All results, plots, and zip saved.")

train-val-test-04-11-2025-16-5.ipynb at commit 30db23e, no license · at the source

Overview

Authors: B. Sidda Reddy1, Ranjeet Ranjan Jha2, Abhishek Dasore3, Deekshitha Desur4, Kiran Shahapurkar5, Vineet Tirth6,7, Ali Algahtani6,8, Vijayabhaskara Rao Bhaviripudi9, Gezahgn Gebremaryam10
  1. Department of Mechanical Engineering, Rajeev Gandhi Memorial College of Engineering & Technology,Nandyal, Andhra Pradesh 518501 India
  2. Department of Mathematics, Indian Institute of Technology,Patna, Bihar 801106 India
  3. School of Computer Science and Artificial Intelligence, SR University,Warangal, Telangana 506371 India
  4. Psychiatry Department, Government Medical College,Nizamabad, Telangana 503001 India
  5. Centre of Excellence-Advanced Materials Synthesis (CoE-AMS), Department of Mechanical Engineering, Alliance School of Applied Engineering, Alliance University,Bengaluru, 562106 India
  6. Mechanical Engineering Department, College of Engineering, King Khalid University,61421 Abha, Aseer Kingdom of Saudi Arabia
  7. Centre for Engineering and Technology Innovations, King Khalid University,61421 Abha, Aseer Kingdom of Saudi Arabia
  8. Research Center for Advanced Materials Science (RCAMS), King Khalid University,Guraiger, PO Box 9004, 61413 Abha, Aseer Kingdom of Saudi Arabia
  9. Departamento de Física, Facultidad de Ciencias Naturales Matemática y del Medio Ambiente, Universidad Tecnológica Metropolitana,Santiago, Chile
  10. Department of Mechanical Engineering, Wolaita Sodo University,Sodo, Ethiopia
Journal: Scientific reports, volume 16, issue 1, article 15938
Dates: received 16 December 2025; accepted 20 March 2026; published online 3 April 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1038/s41598-026-45675-y · PMID 41933058 · PMCID PMC13194798 · OpenAlex W7148712473
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: structural MRI / diffusion (modality), human (organism), other condition (population)
Methods: Connectivity, Statistics, Machine learning
Keywords: Multi-class classification, Brain tumor, ResNet101, MS-DAM, Hybrid augmentation, Mixup, Grad-CAM and SHAP analysis, Cancer, Computational biology and bioinformatics, Mathematics and computing, Oncology
MeSH: Brain Neoplasms*, Magnetic Resonance Imaging*, Classification Algorithms, Convolutional Neural Networks, Deep Learning, Humans (* major topic)
Topic: Brain Tumor Detection and Classification (Neurology, Neuroscience), according to OpenAlex
Funding: King Khalid University (35/44)
Citations: not cited yet (Europe PMC); 42 references in the paper

Abstract

Magnetic Resonance Imaging (MRI) scans are crucial role in identifying brain tumors, ensuring accurate clinical diagnosis and effective personalized treatment planning to improve the chances of survival in patients. However, consistent multi-class classification of brain tumours remains a major challenge due to the considerable variability in tumor morphology and the subtle differences among multiple pathological categories. Although there have been tremendous advancements in convolutional neural networks (CNNs) and attention-based deep learning frameworks, challenges remain in achieving robustness performance across multi-class tumor datasets while maintaining interpretability for clinical use. This paper addresses these challenges, by adopting a novel multi-scale deformable attention module (MS-DAM) framework built on ResNet101. The framework is applied on the Kaggle 14-class MRI Brain tumor dataset, to enhance diagnostic accuracy and computational efficiency by capturing the global contextual and local tumor specific features. To improve generalization, hybrid augmentation strategy combined with mixup regularization has been implemented. The explainabiliy of the model is achieved through Grad-CAM and SHAP analyses. The test results of the proposed model are compared with those reported in the existing literature and superior classification accuracy and generalization are observed. The accuracy of the validation and test data set is achieved 96.89% and 99.21% respectively.

Reproduced under the paper's license (CC BY), from the paper cited above.

Repository

Its files are read in the Code ↔ Paper reader above, with 5 matches between paragraphs and lines of code.

bathinisiddareddy-arch/siddareddy

License: none: the authors keep all their rights
State: the link answers, verified on 28 September 2026
Evidence: files inventoried
Commit: 30db23e2aa38fe0fa1d40897251479d020be405e, 21 February 2026
Languages: Jupyter (1)
Size: 2 files, 1 script
Software Heritage: not archived
Found in: “Data availability”
Holds: 1 notebook
Not found: README, license file, CITATION.cff, environment file, tests, continuous integration, documentation
Tools: Matplotlib (1 file), NumPy (1 file), OpenCV (1 file), pandas (1 file), Pillow (1 file), PyTorch (1 file), scikit-learn (1 file), SciPy (1 file), seaborn (1 file), SHAP (1 file)
Availability: 1 check, the latest on 28 September 2026: the link answers
  • 28 September 2026: the link answers
1 file

The paper's code and data availability statement is in the Data section.

Tracing map

Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.

What the map holds:

  • 1 repository of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 1 script, each with its path and the digest of its content;
  • 5 matches between paragraphs of the paper and lines of the code (method lexical-v1);
  • neither the text of the paper nor the code itself.

Its JSON (tracing-map.json) is deposited on Zenodo with its DOI once the map is validated.

Data

Datasets cited

Data availability

The implementation code, training pipeline, and evaluation scripts used in this study are publicly available to ensure transparency and reproducibility of the reported results. The complete source code can be accessed through the following GitHub repository: https://github.com/bathinisiddareddy-arch/siddareddy.

Reproduced under the paper's license (CC BY), from the paper cited above.

Versions

The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.

Version 2, 28 September 2026

  • Funding: added King Khalid University: 35/44

Version 1, 28 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 9 authors, 11 keywords, 6 MeSH terms, 41 references.

Cite

This paper

Reddy, B. S., Jha, R. R., Dasore, A., Desur, D., Shahapurkar, K., Tirth, V., Algahtani, A., Bhaviripudi, V. R., & Gebremaryam, G. (2026). Multi-class classification of brain tumor using a ResNet101 backbone integrated with multi-scale deformable attention module and advanced data augmentations. Scientific reports, 16(1), 15938. https://doi.org/10.1038/s41598-026-45675-y

BibTeX

@article{reddy2026multi,
author = {Reddy, B. Sidda and Jha, Ranjeet Ranjan and Dasore, Abhishek and Desur, Deekshitha and Shahapurkar, Kiran and Tirth, Vineet and Algahtani, Ali and Bhaviripudi, Vijayabhaskara Rao and Gebremaryam, Gezahgn},
title = {{Multi-class classification of brain tumor using a ResNet101 backbone integrated with multi-scale deformable attention module and advanced data augmentations}},
journal = {Scientific reports},
year = {2026},
month = apr,
volume = {16},
number = {1},
pages = {15938},
publisher = {Nature Publishing Group},
issn = {2045-2322},
doi = {10.1038/s41598-026-45675-y},
url = {https://doi.org/10.1038/s41598-026-45675-y},
pmid = {41933058},
pmcid = {PMC13194798}
}

RIS

TY - JOUR
AU - Reddy, B. Sidda
AU - Jha, Ranjeet Ranjan
AU - Dasore, Abhishek
AU - Desur, Deekshitha
AU - Shahapurkar, Kiran
AU - Tirth, Vineet
AU - Algahtani, Ali
AU - Bhaviripudi, Vijayabhaskara Rao
AU - Gebremaryam, Gezahgn
TI - Multi-class classification of brain tumor using a ResNet101 backbone integrated with multi-scale deformable attention module and advanced data augmentations
T2 - Scientific reports
J2 - Sci Rep
PY - 2026
DA - 2026/04/03
VL - 16
IS - 1
SP - 15938
SN - 2045-2322
PB - Nature Publishing Group
DO - 10.1038/s41598-026-45675-y
UR - https://doi.org/10.1038/s41598-026-45675-y
LA - en
ER -

CSL-JSON

{
"id": "10.1038/s41598-026-45675-y",
"type": "article-journal",
"title": "Multi-class classification of brain tumor using a ResNet101 backbone integrated with multi-scale deformable attention module and advanced data augmentations",
"container-title": "Scientific reports",
"author": [
{
"family": "Reddy",
"given": "B. Sidda"
},
{
"family": "Jha",
"given": "Ranjeet Ranjan"
},
{
"family": "Dasore",
"given": "Abhishek"
},
{
"family": "Desur",
"given": "Deekshitha"
},
{
"family": "Shahapurkar",
"given": "Kiran"
},
{
"family": "Tirth",
"given": "Vineet"
},
{
"family": "Algahtani",
"given": "Ali"
},
{
"family": "Bhaviripudi",
"given": "Vijayabhaskara Rao"
},
{
"family": "Gebremaryam",
"given": "Gezahgn"
}
],
"container-title-short": "Sci Rep",
"volume": "16",
"issue": "1",
"page": "15938",
"DOI": "10.1038/s41598-026-45675-y",
"PMID": "41933058",
"PMCID": "PMC13194798",
"ISSN": "2045-2322",
"publisher": "Nature Publishing Group",
"URL": "https://doi.org/10.1038/s41598-026-45675-y",
"language": "en",
"issued": {
"date-parts": [
[
2026,
4,
3
]
]
}
}

The tracing map gets a citation of its own once an author has validated it and it has a DOI.

Similar papers

The papers with a page that share the most with this one: the tools found in their code, their categories, datasets, cited references and authors, the rarest counting most.

[1] doi:10.3390/jimaging12060233 [code]
Brain Tumor Classification in MRI Images Using Combined Transfer Learning and Convolutional Neural Networks.
Journal: Journal of imaging
In common: seaborn, scikit-learn, pandas, 2 other tools, structural MRI / diffusion, other condition, 6 references
[2] doi:10.1038/s41467-026-76837-1 [code]
Drug screen and machine learning predict neuroprotective agents in a preclinical human model of childhood dementia.
Journal: Nature communications
In common: SHAP, OpenCV, Pillow, 7 other tools, other condition
[3] doi:10.1038/s42003-026-10957-8 [code]
Brain defence by the extracellular matrix protein Cochlin.
Journal: Communications biology
In common: SHAP, OpenCV, Pillow, 7 other tools
[4] doi:10.1016/j.isci.2026.116825 [code]
Social hierarchy shapes behavioral and transcriptional responses to chronic stress and ketamine in male mice.
Journal: iScience
In common: SHAP, OpenCV, Pillow, 7 other tools
[5] doi:10.1002/alz.71649 [code]
Postmortem brain MRI reveals differential associations of subcortical and limbic volumes with cortical thinning and neurodegenerative pathologies.
Journal: Alzheimer's & dementia : the journal of the Alzheimer's Association
In common: OpenCV, Pillow, PyTorch, 6 other tools, structural MRI / diffusion, other condition
[6] doi:10.1126/sciadv.aed3650 [code]
Truthful visualizations for mass spectrometry imaging enable high-spatial-resolution interactive &lt;i&gt;m/z&lt;/i&gt; mapping and exploration.
Journal: Science advances
In common: SHAP, OpenCV, Pillow, 6 other tools
[7] doi:10.7554/elife.109717 [code]
Retrosplenial cortex enables context-dependent goal-directed sensorimotor transformation.
Journal: eLife
In common: SHAP, OpenCV, Pillow, 6 other tools
[8] doi:10.1038/s43856-026-01606-6 [code]
Validation of remote multimodal AI screening for Parkinson disease across diverse settings.
Journal: Communications medicine
In common: SHAP, OpenCV, PyTorch, 6 other tools
[9] doi:10.1186/s13059-026-04125-8 [code]
MLMarker: a machine learning framework for tissue inference and biomarker discovery.
Journal: Genome biology
In common: SHAP, Pillow, PyTorch, 6 other tools
[10] doi:10.1093/braincomms/fcag253 [code]
Disease detection and classification in temporal lobe epilepsy: step-wise versus simultaneous AI decision models in a multisite neuroimaging study.
Journal: Brain communications
In common: OpenCV, Pillow, PyTorch, 6 other tools, structural MRI / diffusion

Contribute

The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.

Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.

Request its removal

To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).

Discussion, reproductions, activity

Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.

Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.

Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.