OSCR

Incorporating normal periventricular changes for enhanced pathological white matter hyperintensity segmentation: on multiclass deep learning approaches.

Code ↔ Paper

17 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 17 matches · 3 of them tie a paragraph to a whole file, not to given lines: weak matches, whose lines are not tinted
  1. [1] § Materials and methods › Deep learning architectures ↔ src/models/transunet.py, the whole file · a weak match · score 0.85 · multi head attention, Positional encoding, transformer bottleneck, sequences, networks, CNN
  2. [2] § Materials and methods › Deep learning architectures ↔ src/models/deeplabv3plus.py, lines 11–54 · score 0.83 · atrous spatial pyramid, ResNet, encoder decoder, backbone, ASPP, convolutions
  3. [3] § Materials and methods › Implementation details ↔ src/inference/wmh_leverage_normal_inference.py, lines 274–310 · score 0.70 · Morphological operations, binary opening, Post processing, disk, kernel, prediction
  4. [4] § Materials and methods › Implementation details ↔ src/training/wmh_leverage_normal_training.py, lines 274–310 · score 0.70 · Morphological operations, binary opening, Post processing, disk, kernel, prediction
  5. [5] § Materials and methods › Deep learning architectures ↔ src/models/transunet_L.py, the whole file · a weak match · score 0.69 · multi head attention, transformer bottleneck, CNN, layers, dropout, modeling
  6. [6] § Materials and methods › Evaluation metrics and statistical framework ↔ src/inference/wmh_leverage_normal_inference.py, lines 773–839 · score 0.66 · Wilcoxon signed rank, confidence intervals, Cohen, metric
  7. [7] § Materials and methods › Evaluation metrics and statistical framework ↔ src/training/wmh_leverage_normal_training.py, lines 773–839 · score 0.66 · Wilcoxon signed rank, confidence intervals, Cohen, metric
  8. [8] § Materials and methods › Deep learning architectures ↔ src/models/deeplabv3plus.py, lines 11–54 · score 0.64 · ReLU, encoder decoder, blocks, activation, convolutions, Deep
  9. [9] § Materials and methods › Training configuration and loss functions ↔ src/training/wmh_leverage_normal_training.py, lines 1122–1164 · score 0.64 · ReduceLROnPlateau, Adam, patience, optimizer, configuration, weight
  10. [10] § Materials and methods › Training configuration and loss functions ↔ src/inference/wmh_leverage_normal_inference.py, lines 1122–1164 · score 0.64 · ReduceLROnPlateau, Adam, patience, optimizer, configuration, weight
  11. [11] § Materials and methods › Local dataset ↔ src/inference/wmh_leverage_normal_inference.py, lines 274–310 · score 0.63 · Morphological opening, morphological operation, disk, pixels, masks, WMH
  12. [12] § Materials and methods › Local dataset ↔ src/training/wmh_leverage_normal_training.py, lines 274–310 · score 0.63 · Morphological opening, morphological operation, disk, pixels, masks, WMH
  13. [13] § Materials and methods › Evaluation metrics and statistical framework ↔ src/inference/wmh_leverage_normal_inference.py, lines 862–944 · score 0.61 · percentile Hausdorff Distance, indicating better boundary, HD95
  14. [14] § Materials and methods › Evaluation metrics and statistical framework ↔ src/training/wmh_leverage_normal_training.py, lines 862–944 · score 0.61 · percentile Hausdorff Distance, indicating better boundary, HD95
  15. [15] § Materials and methods › Implementation details ↔ src/training/wmh_leverage_normal_training.py, lines 98–177 · score 0.57 · random_state, experimental configurations, logged, timestamp, training, model
  16. [16] § Materials and methods › Implementation details ↔ src/inference/wmh_leverage_normal_inference.py, lines 98–177 · score 0.57 · random_state, experimental configurations, logged, timestamp, model, training
  17. [17] § Materials and methods › Deep learning architectures ↔ src/models/transunet.py, the whole file · a weak match · score 0.54 · transposed, network, blocks, max, ReLU, dropout

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

Python · 1,621 lines · 74 KB · MIT · 6 matches

  1. """
  2. Enhanced WMH Segmentation with U-Net - Journal Paper Implementation
  3. Three-class segmentation: Background vs Normal WMH vs Abnormal WMH
  4. Professional results saving and visualization for publication
  5. This relates to our article:
  6. "Incorporating Normal Periventricular Changes for Enhanced Pathological
  7. White Matter Hyperintensity Segmentation: On Multi-Class Deep Learning Approaches"
  8. Authors:
  9. "Mahdi Bashiri Bawil, Mousa Shamsi, Ali Fahmi Jafargholkhanloo, Abolhassan Shakeri Bavil"
  10. Developer:
  11. "Mahdi Bashiri Bawil"
  12. """
  13. ###################### Libraries ######################
  14. # General Utilities
  15. import numpy as np
  16. import matplotlib.pyplot as plt
  17. import matplotlib.patches as patches
  18. import seaborn as sns
  19. import cv2 as cv
  20. import os
  21. import pandas as pd
  22. from datetime import datetime
  23. from tqdm import tqdm
  24. import json
  25. import pickle
  26. from pathlib import Path
  27. from skimage.morphology import remove_small_objects, binary_opening, disk
  28. from skimage.measure import label
  29. # Deep Learning
  30. import tensorflow as tf
  31. import keras
  32. from keras.models import Model, load_model
  33. from keras.layers import Input, Conv2D, MaxPooling2D, Conv2DTranspose, concatenate
  34. from keras import backend as K
  35. from tensorflow.keras import layers, optimizers, callbacks
  36. from keras.utils import to_categorical
  37. # Analysis and Statistics
  38. from sklearn.model_selection import train_test_split
  39. from sklearn.metrics import confusion_matrix, accuracy_score, precision_score, recall_score, f1_score, classification_report
  40. from scipy import stats
  41. import warnings
  42. warnings.filterwarnings('ignore')
  43. # Models
  44. from unet import build_unet_3class
  45. from attention_unet import build_attention_unet_3class
  46. from transunet import build_trans_unet_3class
  47. from deeplabv3plus import build_deeplabv3_unet_3class
  48. # Loss Functions
  49. from losses import *
  50. # Metrics Functions
  51. from metrics import *
  52. # Check for GPU assistance
  53. import tensorflow as tf
  54. print("TensorFlow version:", tf.__version__)
  55. print("GPU Available: ", tf.test.is_gpu_available())
  56. print("Built with CUDA: ", tf.test.is_built_with_cuda())
  57. print("Physical devices: ", tf.config.list_physical_devices())
  58. # Force GPU if available
  59. if tf.config.list_physical_devices('GPU'):
  60. print("\n\n\t\t\tUsing GPU\n\n")
  61. else:
  62. print("\n\n\t\t\tUsing CPU\n\n")
  63. # os.environ['CUDA_VISIBLE_DEVICES'] = '-1'
  64. # Set publication-ready matplotlib settings
  65. plt.rcParams.update({
  66. 'font.size': 12,
  67. 'font.family': 'serif',
  68. 'axes.labelsize': 12,
  69. 'axes.titlesize': 14,
  70. 'xtick.labelsize': 10,
  71. 'ytick.labelsize': 10,
  72. 'legend.fontsize': 11,
  73. 'figure.titlesize': 16,
  74. 'figure.dpi': 300,
  75. 'savefig.dpi': 300,
  76. 'savefig.format': 'png',
  77. 'savefig.bbox': 'tight',
  78. 'savefig.pad_inches': 0.1
  79. })
  80. ###################### Configuration and Setup ######################
  81. class Config:
  82. """Configuration class for the experiment"""
  83. def __init__(self):
  84. # Model Name
  85. self.model_name = 'attn_unet' # 'unet', 'attn_unet', 'trans_unet', 'deepl3_unet'
  86. # Paths
  87. self.train_dir = "Leverage_Article_Data/train_3L_wmh_local_public/"
  88. self.test_dir = "Leverage_Article_Data/test_3L_wmh_local_public/"
  89. self.intended_study_dir = "leverage_results_20251124_160949_attn_unet" # for inference intentions
  90. # Model parameters
  91. self.input_shape = (256, 256, 1)
  92. self.target_size = (256, 256)
  93. self.num_classes_3 = 3
  94. self.num_classes_binary = 1
  95. # Training parameters
  96. self.mode = 'no-training'
  97. self.epochs = 50 # Increased for better convergence
  98. self.batch_size = 8
  99. self.learning_rate = 1e-4
  100. self.validation_split = 0.1
  101. self.random_state = 42
  102. # Loss function options
  103. self.loss_options = {
  104. 'scenario1': 'weighted_bce', # weighted_bce, focal, combined, dice
  105. 'scenario2': 'weighted_categorical' # weighted_categorical, multiclass_dice, categorical
  106. }
  107. # Choose a model to train or inference
  108. if self.model_name == 'unet':
  109. self.build_unet_variant = build_unet_3class
  110. elif self.model_name == 'attn_unet':
  111. self.build_unet_variant = build_attention_unet_3class
  112. elif self.model_name == 'trans_unet':
  113. self.build_unet_variant = build_trans_unet_3class
  114. elif self.model_name == 'deepl3_unet':
  115. self.build_unet_variant = build_deeplabv3_unet_3class
  116. # Create results directory with timestamp
  117. self.timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
  118. if self.mode != 'training':
  119. self.results_dir = Path(f"leverage_results_{self.timestamp}_{self.model_name}_no_training")
  120. else:
  121. self.results_dir = Path(f"leverage_results_{self.timestamp}_{self.model_name}")
  122. self.create_directory_structure()
  123. def create_directory_structure(self):
  124. """Create professional directory structure for results"""
  125. subdirs = [
  126. 'models',
  127. 'figures',
  128. 'tables',
  129. 'statistics',
  130. 'predictions',
  131. 'logs',
  132. 'config'
  133. ]
  134. self.results_dir.mkdir(exist_ok=True)
  135. for subdir in subdirs:
  136. (self.results_dir / subdir).mkdir(exist_ok=True)
  137. # Save experiment configuration
  138. config_dict = {
  139. 'timestamp': self.timestamp,
  140. 'input_shape': self.input_shape,
  141. 'target_size': self.target_size,
  142. 'epochs': self.epochs,
  143. 'batch_size': self.batch_size,
  144. 'learning_rate': self.learning_rate,
  145. 'validation_split': self.validation_split,
  146. 'random_state': self.random_state,
  147. 'loss_options': self.loss_options
  148. }
  149. with open(self.results_dir / 'config' / 'experiment_config.json', 'w') as f:
  150. json.dump(config_dict, f, indent=2)
  151. config = Config()
  152. ###################### Data Loading Functions ######################
  153. def extract_number(filename):
  154. """Extract patient ID and slice number for proper sorting"""
  155. return int(''.join(filter(str.isdigit, filename.split('_')[0])))
  156. def load_wmh_dataset(data_dir, target_size=(256, 256), save_info=True):
  157. """
  158. Load dataset with specific format: 256x512 images (FLAIR + GT mask concatenated)
  159. """
  160. images, masks_3class, masks_binary = [], [], []
  161. # image_files = sorted(os.listdir(data_dir), key=extract_number)
  162. image_files = [f for f in os.listdir(data_dir)]
  163. dataset_info = {
  164. 'total_files': len(image_files),
  165. 'loaded_files': 0,
  166. 'skipped_files': [],
  167. 'image_shapes': [],
  168. 'class_distributions': {'background': [], 'normal_wmh': [], 'abnormal_wmh': []}
  169. }
  170. for img_name in tqdm(image_files, desc=f"Loading from {os.path.basename(data_dir)}"):
  171. # Load concatenated image
  172. full_img = cv.imread(os.path.join(data_dir, img_name), cv.IMREAD_ANYDEPTH | cv.IMREAD_GRAYSCALE).astype(np.float32)
  173. if full_img is None or full_img.shape[1] != 512:
  174. dataset_info['skipped_files'].append(img_name)
  175. continue
  176. # Split into FLAIR and GT
  177. flair_img = full_img[:, :256]
  178. gt_mask = full_img[:, 256:]
  179. # Resize if needed
  180. if target_size != (256, 256):
  181. flair_img = cv.resize(flair_img, target_size)
  182. gt_mask = cv.resize(gt_mask, target_size)
  183. dataset_info['image_shapes'].append(flair_img.shape)
  184. # Normalize FLAIR image
  185. flair_img = flair_img.astype(np.float32)
  186. flair_img = (flair_img - np.mean(flair_img)) / (np.std(flair_img) + 1e-7)
  187. flair_img = np.expand_dims(flair_img, axis=-1)
  188. # Process ground truth masks
  189. gt_mask = gt_mask.astype(np.float32)
  190. # Create 3-class mask
  191. mask_3class = np.zeros_like(gt_mask, dtype=np.uint8)
  192. threshold_1 = 32767 // 2
  193. threshold_2 = 32767 + 1000
  194. threshold_3 = 65535 - 32767 // 2
  195. mask_3class[gt_mask < threshold_1] = 0
  196. mask_3class[(gt_mask >= threshold_1) & (gt_mask < threshold_2)] = 1
  197. mask_3class[gt_mask >= threshold_3] = 2
  198. # Create binary mask
  199. mask_binary = np.zeros_like(gt_mask, dtype=np.uint8)
  200. mask_binary[gt_mask >= threshold_3] = 1
  201. # Record class distributions
  202. unique, counts = np.unique(mask_3class, return_counts=True)
  203. class_dist = dict(zip(unique, counts))
  204. dataset_info['class_distributions']['background'].append(class_dist.get(0, 0))
  205. dataset_info['class_distributions']['normal_wmh'].append(class_dist.get(1, 0))
  206. dataset_info['class_distributions']['abnormal_wmh'].append(class_dist.get(2, 0))
  207. images.append(flair_img)
  208. masks_3class.append(mask_3class)
  209. masks_binary.append(mask_binary)
  210. dataset_info['loaded_files'] += 1
  211. # Save dataset information
  212. if save_info:
  213. dataset_info['class_distributions'] = {k: np.array(v) for k, v in dataset_info['class_distributions'].items()}
  214. with open(config.results_dir / 'logs' / f'dataset_info_{os.path.basename(data_dir)}.pkl', 'wb') as f:
  215. pickle.dump(dataset_info, f)
  216. return np.array(images), np.array(masks_3class), np.array(masks_binary), dataset_info
  217. ###################### U-Net Architecture ######################
  218. # callable from related functions saved in the main directory
  219. ###################### Loss Functions ######################
  220. # callable from the related function saved in the main directory
  221. ###################### Metrics and Evaluation ######################
  222. # callable from related functions saved in the main directory
  223. ###################### Post Processing ######################
  224. def post_process_predictions(predictions, min_object_size=5, apply_opening=True, kernel_size=3):
  225. """
  226. Post-process binary predictions to remove small objects and apply morphological operations
  227. Args:
  228. predictions: Binary prediction masks (numpy array)
  229. min_object_size: Minimum object size in pixels (objects smaller than this are removed)
  230. apply_opening: Whether to apply morphological opening operation
  231. kernel_size: Size of morphological kernel for opening operation
  232. Returns:
  233. post_processed_predictions: Cleaned binary masks
  234. """
  235. from skimage.morphology import remove_small_objects, binary_opening, disk
  236. from skimage.measure import label
  237. post_processed = np.zeros_like(predictions, dtype=np.uint8)
  238. for i in range(predictions.shape[0]):
  239. mask = predictions[i].astype(bool)
  240. # Remove small objects
  241. if min_object_size > 0:
  242. mask = remove_small_objects(mask, min_size=min_object_size)
  243. # Apply morphological opening
  244. if apply_opening:
  245. kernel = disk(kernel_size)
  246. mask = binary_opening(mask, kernel)
  247. # Remove small objects
  248. if min_object_size > 0:
  249. mask = remove_small_objects(mask, min_size=min_object_size)
  250. post_processed[i] = mask.astype(np.uint8)
  251. return post_processed
  252. ###################### Professional Visualization Functions ######################
  253. class PublicationPlotter:
  254. """Professional plotting class for publication-quality figures"""
  255. def __init__(self, results_dir):
  256. self.results_dir = Path(results_dir)
  257. self.figures_dir = self.results_dir / 'figures'
  258. def plot_training_curves(self, history_s1, history_s2, save_name='training_curves'):
  259. """Plot publication-quality training curves"""
  260. fig, axes = plt.subplots(2, 2, figsize=(12, 10))
  261. # Handle both History objects (from training) and dicts (from loading)
  262. if hasattr(history_s1, 'history'):
  263. hist_s1 = history_s1.history # From training
  264. else:
  265. hist_s1 = history_s1 # From loading (already a dict)
  266. if hasattr(history_s2, 'history'):
  267. hist_s2 = history_s2.history # From training
  268. else:
  269. hist_s2 = history_s2 # From loading (already a dict)
  270. # Scenario 1
  271. axes[0, 0].plot(hist_s1['loss'], 'b-', linewidth=2, label='Training')
  272. axes[0, 0].plot(hist_s1['val_loss'], 'r-', linewidth=2, label='Validation')
  273. axes[0, 0].set_title('(a) Binary Classification Loss')
  274. axes[0, 0].set_xlabel('Epoch')
  275. axes[0, 0].set_ylabel('Loss')
  276. axes[0, 0].legend()
  277. axes[0, 0].grid(True, alpha=0.3)
  278. axes[0, 1].plot(hist_s1['accuracy'], 'b-', linewidth=2, label='Training')
  279. axes[0, 1].plot(hist_s1['val_accuracy'], 'r-', linewidth=2, label='Validation')
  280. axes[0, 1].set_title('(b) Binary Classification Accuracy')
  281. axes[0, 1].set_xlabel('Epoch')
  282. axes[0, 1].set_ylabel('Accuracy')
  283. axes[0, 1].legend()
  284. axes[0, 1].grid(True, alpha=0.3)
  285. # Scenario 2
  286. axes[1, 0].plot(hist_s2['loss'], 'g-', linewidth=2, label='Training')
  287. axes[1, 0].plot(hist_s2['val_loss'], 'orange', linewidth=2, label='Validation')
  288. axes[1, 0].set_title('(c) Three-class Classification Loss')
  289. axes[1, 0].set_xlabel('Epoch')
  290. axes[1, 0].set_ylabel('Loss')
  291. axes[1, 0].legend()
  292. axes[1, 0].grid(True, alpha=0.3)
  293. axes[1, 1].plot(hist_s2['accuracy'], 'g-', linewidth=2, label='Training')
  294. axes[1, 1].plot(hist_s2['val_accuracy'], 'orange', linewidth=2, label='Validation')
  295. axes[1, 1].set_title('(d) Three-class Classification Accuracy')
  296. axes[1, 1].set_xlabel('Epoch')
  297. axes[1, 1].set_ylabel('Accuracy')
  298. axes[1, 1].legend()
  299. axes[1, 1].grid(True, alpha=0.3)
  300. plt.tight_layout()
  301. plt.savefig(self.figures_dir / f'{save_name}.png')
  302. plt.savefig(self.figures_dir / f'{save_name}.pdf') # For LaTeX
  303. # plt.show()
  304. def plot_comparison_visualization(self, images, gt_3class, gt_binary, pred_s1, pred_s2,
  305. indices=None, save_name='comparison_visualization'):
  306. """Create publication-quality comparison visualization in single column format"""
  307. # Use random selection:
  308. if indices is None:
  309. indices = np.random.choice(len(images), 3, replace=False)
  310. # or Use manual selection:
  311. indices = np.array([50, 51, 62, 74]) # our chosen indices
  312. # indices = np.array([44]) # our chosen indices
  313. # Create single column layout: 6 rows, 1 column per sample
  314. n_samples = len(indices)
  315. n_plots = 6 # Number of different visualizations
  316. # Adjust figure size for single column format
  317. fig, axes = plt.subplots(n_plots * n_samples, 1, figsize=(8, 3 * n_plots * n_samples))
  318. # If only one sample, ensure axes is iterable
  319. if n_samples == 1:
  320. axes = np.array(axes).reshape(-1)
  321. titles = ['FLAIR Image', 'GT (3-class)', 'GT (Abnormal)',
  322. 'Scenario 1 Performance', 'Scenario 2 Performance', 'Legend']
  323. for sample_idx, idx in enumerate(indices):
  324. base_row = sample_idx * n_plots
  325. # Add sample identifier if multiple samples
  326. if n_samples > 1:
  327. sample_title = f" - Sample {sample_idx + 1}"
  328. else:
  329. sample_title = ""
  330. # FLAIR Image
  331. axes[base_row + 0].imshow(images[idx].squeeze(), cmap='gray')
  332. axes[base_row + 0].set_title(titles[0] + sample_title, fontsize=12, pad=10)
  333. axes[base_row + 0].axis('off')
  334. # GT 3-class - Using grayscale
  335. axes[base_row + 1].imshow(gt_3class[idx], cmap='gray')
  336. axes[base_row + 1].set_title(titles[1] + sample_title, fontsize=12, pad=10)
  337. axes[base_row + 1].axis('off')
  338. # GT Binary (Abnormal only) - Black and white binary
  339. axes[base_row + 2].imshow(gt_binary[idx], cmap='gray', vmin=0, vmax=1)
  340. axes[base_row + 2].set_title(titles[2] + sample_title, fontsize=12, pad=10)
  341. axes[base_row + 2].axis('off')
  342. # Create RGB version of FLAIR image for overlays
  343. flair_rgb = np.stack([images[idx].squeeze()] * 3, axis=-1)
  344. # Normalize to 0-1 range if needed
  345. flair_rgb = (flair_rgb - flair_rgb.min()) / (flair_rgb.max() - flair_rgb.min())
  346. # Scenario 1 Performance Analysis
  347. # Convert prediction to binary (assuming abnormal class)
  348. pred_s1_binary = (pred_s1[idx] > 0).astype(np.uint8)
  349. # Calculate TP, FP, FN
  350. tp_s1 = (gt_binary[idx] == 1) & (pred_s1_binary == 1)
  351. fp_s1 = (gt_binary[idx] == 0) & (pred_s1_binary == 1)
  352. fn_s1 = (gt_binary[idx] == 1) & (pred_s1_binary == 0)
  353. # Create overlay image for Scenario 1
  354. overlay_s1 = flair_rgb.copy()
  355. overlay_s1[tp_s1, :] = [0, 1, 0] # Green for TP
  356. overlay_s1[fp_s1, :] = [1, 0, 0] # Red for FP
  357. overlay_s1[fn_s1, :] = [1, 1, 0] # Yellow for FN
  358. axes[base_row + 3].imshow(overlay_s1)
  359. axes[base_row + 3].set_title(titles[3] + sample_title, fontsize=12, pad=10)
  360. axes[base_row + 3].axis('off')
  361. # Scenario 2 Performance Analysis
  362. # Convert prediction to binary (abnormal only)
  363. if np.max(pred_s2) == 2:
  364. pred_s2_binary = (pred_s2[idx] == 2).astype(np.uint8)
  365. else:
  366. pred_s2_binary = (pred_s2[idx] > 0).astype(np.uint8)
  367. # Calculate TP, FP, FN
  368. tp_s2 = (gt_binary[idx] == 1) & (pred_s2_binary == 1)
  369. fp_s2 = (gt_binary[idx] == 0) & (pred_s2_binary == 1)
  370. fn_s2 = (gt_binary[idx] == 1) & (pred_s2_binary == 0)
  371. # Create overlay image for Scenario 2
  372. overlay_s2 = flair_rgb.copy()
  373. overlay_s2[tp_s2, :] = [0, 1, 0] # Green for TP
  374. overlay_s2[fp_s2, :] = [1, 0, 0] # Red for FP
  375. overlay_s2[fn_s2, :] = [0, 0, 1] # Blue for FN
  376. axes[base_row + 4].imshow(overlay_s2)
  377. axes[base_row + 4].set_title(titles[4] + sample_title, fontsize=12, pad=10)
  378. axes[base_row + 4].axis('off')
  379. # Legend plot (replacing the overlay comparison)
  380. axes[base_row + 5].axis('off')
  381. from matplotlib.patches import Rectangle
  382. from matplotlib.lines import Line2D
  383. # Create legend elements
  384. legend_elements = [
  385. Line2D([0], [0], marker='s', color='w', markerfacecolor='green',
  386. markersize=15, label='True Positive (TP)'),
  387. Line2D([0], [0], marker='s', color='w', markerfacecolor='red',
  388. markersize=15, label='False Positive (FP)'),
  389. Line2D([0], [0], marker='s', color='w', markerfacecolor='blue',
  390. markersize=15, label='False Negative (FN)')
  391. ]
  392. legend = axes[base_row + 5].legend(handles=legend_elements,
  393. loc='center', fontsize=12,
  394. title='Performance Metrics',
  395. title_fontsize=14)
  396. legend.get_title().set_fontweight('bold')
  397. # Adjust layout with more spacing for better readability
  398. plt.tight_layout(pad=2.0)
  399. # Save with high DPI for publication quality
  400. plt.savefig(self.figures_dir / f'{save_name}.png', dpi=300, bbox_inches='tight')
  401. plt.savefig(self.figures_dir / f'{save_name}.pdf', bbox_inches='tight')
  402. # plt.show()
  403. def plot_metrics_comparison(self, metrics_s1, metrics_s2, save_name='metrics_comparison'):
  404. """Create professional metrics comparison plot including IoU"""
  405. metrics_to_plot = ['Accuracy', 'Precision', 'Recall', 'Dice', 'IoU']
  406. s1_values = [metrics_s1[metric] for metric in metrics_to_plot]
  407. s2_values = [metrics_s2[metric] for metric in metrics_to_plot]
  408. x = np.arange(len(metrics_to_plot))
  409. width = 0.35
  410. fig, ax = plt.subplots(figsize=(12, 6))
  411. bars1 = ax.bar(x - width/2, s1_values, width, label='Binary Classification',
  412. color='skyblue', alpha=0.8)
  413. bars2 = ax.bar(x + width/2, s2_values, width, label='Three-class Classification',
  414. color='lightcoral', alpha=0.8)
  415. ax.set_xlabel('Metrics')
  416. ax.set_ylabel('Score')
  417. ax.set_title('Performance Comparison: Binary vs Three-class Classification')
  418. ax.set_xticks(x)
  419. ax.set_xticklabels(metrics_to_plot)
  420. ax.legend()
  421. ax.grid(True, alpha=0.3)
  422. ax.set_ylim(0, 1.0)
  423. # Add value labels on bars
  424. def autolabel(bars):
  425. for bar in bars:
  426. height = bar.get_height()
  427. ax.annotate(f'{height:.3f}',
  428. xy=(bar.get_x() + bar.get_width() / 2, height),
  429. xytext=(0, 3),
  430. textcoords="offset points",
  431. ha='center', va='bottom', fontsize=10)
  432. autolabel(bars1)
  433. autolabel(bars2)
  434. plt.tight_layout()
  435. plt.savefig(self.figures_dir / f'{save_name}.png')
  436. plt.savefig(self.figures_dir / f'{save_name}.pdf')
  437. def plot_dice_distribution(self, dice_s1, dice_s2, save_name='dice_distribution'):
  438. """Plot Dice coefficient distributions"""
  439. fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(15, 6))
  440. # Box plot comparison
  441. ax1.boxplot([dice_s1, dice_s2], labels=['Binary\nClassification', 'Three-class\nClassification'])
  442. ax1.set_ylabel('Dice Coefficient')
  443. ax1.set_title('(a) Dice Coefficient Distribution')
  444. ax1.grid(True, alpha=0.3)
  445. # Histogram overlay
  446. ax2.hist(dice_s1, alpha=0.6, bins=20, label='Binary Classification', color='skyblue')
  447. ax2.hist(dice_s2, alpha=0.6, bins=20, label='Three-class Classification', color='lightcoral')
  448. ax2.set_xlabel('Dice Coefficient')
  449. ax2.set_ylabel('Frequency')
  450. ax2.set_title('(b) Dice Coefficient Histogram')
  451. ax2.legend()
  452. ax2.grid(True, alpha=0.3)
  453. plt.tight_layout()
  454. plt.savefig(self.figures_dir / f'{save_name}.png')
  455. plt.savefig(self.figures_dir / f'{save_name}.pdf')
  456. # plt.show()
  457. def plot_dice_iou_distribution(self, dice_s1, dice_s2, iou_s1, iou_s2, save_name='dice_iou_distribution'):
  458. """Plot both Dice and IoU coefficient distributions"""
  459. fig, axes = plt.subplots(2, 2, figsize=(15, 10))
  460. # Dice box plots
  461. axes[0,0].boxplot([dice_s1, dice_s2], labels=['Binary\nClassification', 'Three-class\nClassification'])
  462. axes[0,0].set_ylabel('Dice Coefficient')
  463. axes[0,0].set_title('(a) Dice Coefficient Distribution')
  464. axes[0,0].grid(True, alpha=0.3)
  465. # Dice histograms
  466. axes[0,1].hist(dice_s1, alpha=0.6, bins=20, label='Binary Classification', color='skyblue')
  467. axes[0,1].hist(dice_s2, alpha=0.6, bins=20, label='Three-class Classification', color='lightcoral')
  468. axes[0,1].set_xlabel('Dice Coefficient')
  469. axes[0,1].set_ylabel('Frequency')
  470. axes[0,1].set_title('(b) Dice Coefficient Histogram')
  471. axes[0,1].legend()
  472. axes[0,1].grid(True, alpha=0.3)
  473. # IoU box plots
  474. axes[1,0].boxplot([iou_s1, iou_s2], labels=['Binary\nClassification', 'Three-class\nClassification'])
  475. axes[1,0].set_ylabel('IoU Coefficient')
  476. axes[1,0].set_title('(c) IoU Coefficient Distribution')
  477. axes[1,0].grid(True, alpha=0.3)
  478. # IoU histograms
  479. axes[1,1].hist(iou_s1, alpha=0.6, bins=20, label='Binary Classification', color='skyblue')
  480. axes[1,1].hist(iou_s2, alpha=0.6, bins=20, label='Three-class Classification', color='lightcoral')
  481. axes[1,1].set_xlabel('IoU Coefficient')
  482. axes[1,1].set_ylabel('Frequency')
  483. axes[1,1].set_title('(d) IoU Coefficient Histogram')
  484. axes[1,1].legend()
  485. axes[1,1].grid(True, alpha=0.3)
  486. plt.tight_layout()
  487. plt.savefig(self.figures_dir / f'{save_name}.png')
  488. plt.savefig(self.figures_dir / f'{save_name}.pdf')
  489. ###################### Results Saving Functions ######################
  490. class ResultsSaver:
  491. """Professional results saving and documentation"""
  492. def __init__(self, results_dir):
  493. self.results_dir = Path(results_dir)
  494. def save_models(self, model_s1, model_s2, history_s1, history_s2):
  495. """Save trained models and training histories"""
  496. # Save models
  497. model_s1.save(self.results_dir / 'models' / 'scenario1_binary_model.h5')
  498. model_s2.save(self.results_dir / 'models' / 'scenario2_multiclass_model.h5')
  499. # Save training histories
  500. with open(self.results_dir / 'models' / 'training_history_s1.pkl', 'wb') as f:
  501. pickle.dump(history_s1.history, f)
  502. with open(self.results_dir / 'models' / 'training_history_s2.pkl', 'wb') as f:
  503. pickle.dump(history_s2.history, f)
  504. def load_models(self, study_dir, loss_s1_func, loss_s2_func):
  505. """Load saved models and training histories with proper custom loss functions
  506. Args:
  507. study_dir: Directory containing the saved models
  508. """
  509. try:
  510. # Convert study_dir to Path object if it's a string
  511. study_dir = Path(study_dir)
  512. # Load models with their respective custom objects
  513. model_s1 = keras.models.load_model(
  514. study_dir / 'models' / 'scenario1_binary_model.h5',
  515. compile=False
  516. # custom_objects=loss_s1_func
  517. )
  518. model_s2 = keras.models.load_model(
  519. study_dir / 'models' / 'scenario2_multiclass_model.h5',
  520. compile=False
  521. # custom_objects=loss_s2_func
  522. )
  523. # Load training histories
  524. with open(study_dir / 'models' / 'training_history_s1.pkl', 'rb') as f:
  525. history_s1 = pickle.load(f)
  526. with open(study_dir / 'models' / 'training_history_s2.pkl', 'rb') as f:
  527. history_s2 = pickle.load(f)
  528. print("Models and histories loaded successfully!")
  529. return model_s1, model_s2, history_s1, history_s2
  530. except FileNotFoundError as e:
  531. print(f"Error: Could not find saved models. {e}")
  532. print("Make sure you have saved models using save_models() first.")
  533. print(f"Looking in directory: {Path(study_dir) / 'models'}")
  534. return None, None, None, None
  535. except Exception as e:
  536. print(f"Error loading models: {e}")
  537. return None, None, None, None
  538. def save_predictions(self, test_images, test_masks_3class, test_masks_binary,
  539. pred_s1, pred_s2, dataset_info):
  540. """Save predictions and test data"""
  541. predictions_dir = self.results_dir / 'predictions'
  542. # Save raw predictions
  543. np.save(predictions_dir / 'test_images.npy', test_images)
  544. np.save(predictions_dir / 'test_masks_3class.npy', test_masks_3class)
  545. np.save(predictions_dir / 'test_masks_binary.npy', test_masks_binary)
  546. np.save(predictions_dir / 'predictions_scenario1.npy', pred_s1)
  547. np.save(predictions_dir / 'predictions_scenario2.npy', pred_s2)
  548. # Save dataset information
  549. with open(predictions_dir / 'dataset_info.pkl', 'wb') as f:
  550. pickle.dump(dataset_info, f)
  551. def save_metrics_table(self, metrics_s1, metrics_s2, dice_stats):
  552. """Save comprehensive metrics table including HD95 and ASSD"""
  553. # Create comprehensive results table
  554. results_table = pd.DataFrame([metrics_s1, metrics_s2])
  555. # Add statistical information - updated for all metrics including surface-based
  556. stats_row = {
  557. 'Scenario': 'Statistical Analysis',
  558. 'Accuracy': f"Dice p={dice_stats['dice_p_value']:.4f}",
  559. 'Precision': f"Dice t={dice_stats['dice_t_statistic']:.4f}",
  560. 'Recall': f"Dice Δ={dice_stats['dice_improvement']:.4f}",
  561. 'Specificity': f"Dice ES={dice_stats['dice_effect_size']:.4f}",
  562. 'Dice': f"IoU p={dice_stats['iou_p_value']:.4f}",
  563. 'IoU': f"IoU Δ={dice_stats['iou_improvement']:.4f}"
  564. }
  565. results_table = pd.concat([results_table, pd.DataFrame([stats_row])], ignore_index=True)
  566. # Save as CSV and Excel
  567. results_table.to_csv(self.results_dir / 'tables' / 'comprehensive_results.csv', index=False)
  568. results_table.to_excel(self.results_dir / 'tables' / 'comprehensive_results.xlsx', index=False)
  569. # Create separate surface metrics table
  570. surface_metrics_data = {
  571. 'Scenario': ['Binary (S1)', 'Three-class (S2)', 'Statistical Analysis'],
  572. 'HD95_Mean': [
  573. dice_stats['hd95_scenario1_mean'],
  574. dice_stats['hd95_scenario2_mean'],
  575. None
  576. ],
  577. 'HD95_Std': [
  578. dice_stats['hd95_scenario1_std'],
  579. dice_stats['hd95_scenario2_std'],
  580. None
  581. ],
  582. 'HD95_Median': [
  583. dice_stats['hd95_scenario1_median'],
  584. dice_stats['hd95_scenario2_median'],
  585. None
  586. ],
  587. 'ASSD_Mean': [
  588. dice_stats['assd_scenario1_mean'],
  589. dice_stats['assd_scenario2_mean'],
  590. None
  591. ],
  592. 'ASSD_Std': [
  593. dice_stats['assd_scenario1_std'],
  594. dice_stats['assd_scenario2_std'],
  595. None
  596. ],
  597. 'ASSD_Median': [
  598. dice_stats['assd_scenario1_median'],
  599. dice_stats['assd_scenario2_median'],
  600. None
  601. ],
  602. 'HD95_Stats': [
  603. None,
  604. None,
  605. f"Δ={dice_stats['hd95_improvement']:.4f}px, p={dice_stats['hd95_p_value']:.4f}"
  606. ],
  607. 'ASSD_Stats': [
  608. None,
  609. None,
  610. f"Δ={dice_stats['assd_improvement']:.4f}px, p={dice_stats['assd_p_value']:.4f}"
  611. ]
  612. }
  613. surface_table = pd.DataFrame(surface_metrics_data)
  614. surface_table.to_csv(self.results_dir / 'tables' / 'surface_metrics.csv', index=False)
  615. surface_table.to_excel(self.results_dir / 'tables' / 'surface_metrics.xlsx', index=False)
  616. # Create comprehensive LaTeX table with all metrics
  617. latex_table = results_table.iloc[:-1].to_latex(
  618. index=False,
  619. float_format="%.4f",
  620. caption="Performance comparison between binary and three-class segmentation approaches",
  621. label="tab:performance_comparison"
  622. )
  623. with open(self.results_dir / 'tables' / 'latex_table.tex', 'w') as f:
  624. f.write(latex_table)
  625. # Create LaTeX table for surface metrics
  626. latex_surface_table = surface_table.iloc[:-1].to_latex(
  627. index=False,
  628. float_format="%.4f",
  629. caption="Surface-based metrics (HD95 and ASSD) comparison in pixels",
  630. label="tab:surface_metrics"
  631. )
  632. with open(self.results_dir / 'tables' / 'latex_surface_table.tex', 'w') as f:
  633. f.write(latex_surface_table)
  634. return results_table, surface_table
  635. def save_statistical_analysis(self, dice_s1, dice_s2, iou_s1, iou_s2,
  636. hd95_s1, hd95_s2, assd_s1, assd_s2,
  637. metrics_s1, metrics_s2):
  638. """Comprehensive statistical analysis for multiple metrics including surface-based metrics"""
  639. from scipy.stats import ttest_rel, wilcoxon, normaltest, levene
  640. import scipy.stats as stats
  641. def analyze_metric(metric1, metric2, metric_name, lower_is_better=False):
  642. """Analyze a single metric pair"""
  643. # Normality tests
  644. _, p_normal_1 = normaltest(metric1)
  645. _, p_normal_2 = normaltest(metric2)
  646. # Paired t-test
  647. # For lower_is_better metrics (HD95, ASSD), we reverse the comparison
  648. if lower_is_better:
  649. t_stat, p_ttest = ttest_rel(metric1, metric2) # Test if metric1 > metric2
  650. else:
  651. t_stat, p_ttest = ttest_rel(metric2, metric1) # Test if metric2 > metric1
  652. # Wilcoxon signed-rank test (non-parametric alternative)
  653. if lower_is_better:
  654. w_stat, p_wilcoxon = wilcoxon(metric1, metric2, alternative='two-sided')
  655. else:
  656. w_stat, p_wilcoxon = wilcoxon(metric2, metric1, alternative='two-sided')
  657. # Effect size (Cohen's d)
  658. pooled_std = np.sqrt((np.var(metric1, ddof=1) + np.var(metric2, ddof=1)) / 2)
  659. if lower_is_better:
  660. cohens_d = (np.mean(metric1) - np.mean(metric2)) / pooled_std if pooled_std > 0 else 0
  661. else:
  662. cohens_d = (np.mean(metric2) - np.mean(metric1)) / pooled_std if pooled_std > 0 else 0
  663. # Confidence interval for difference
  664. diff = metric2 - metric1 if not lower_is_better else metric1 - metric2
  665. ci_lower, ci_upper = stats.t.interval(0.95, len(diff)-1,
  666. loc=np.mean(diff),
  667. scale=stats.sem(diff))
  668. # Calculate improvement
  669. if lower_is_better:
  670. improvement = np.mean(metric1) - np.mean(metric2) # Reduction is good
  671. improvement_percent = (improvement / np.mean(metric1)) * 100 if np.mean(metric1) > 0 else 0
  672. else:
  673. improvement = np.mean(metric2) - np.mean(metric1) # Increase is good
  674. improvement_percent = (improvement / np.mean(metric1)) * 100 if np.mean(metric1) > 0 else 0
  675. return {
  676. f'{metric_name}_scenario1_mean': np.mean(metric1),
  677. f'{metric_name}_scenario1_std': np.std(metric1),
  678. f'{metric_name}_scenario1_median': np.median(metric1),
  679. f'{metric_name}_scenario2_mean': np.mean(metric2),
  680. f'{metric_name}_scenario2_std': np.std(metric2),
  681. f'{metric_name}_scenario2_median': np.median(metric2),
  682. f'{metric_name}_improvement': improvement,
  683. f'{metric_name}_improvement_percent': improvement_percent,
  684. f'{metric_name}_t_statistic': t_stat,
  685. f'{metric_name}_p_value': p_ttest,
  686. f'{metric_name}_wilcoxon_statistic': w_stat,
  687. f'{metric_name}_wilcoxon_p_value': p_wilcoxon,
  688. f'{metric_name}_effect_size': cohens_d,
  689. f'{metric_name}_ci_lower': ci_lower,
  690. f'{metric_name}_ci_upper': ci_upper,
  691. f'{metric_name}_normality_s1_p': p_normal_1,
  692. f'{metric_name}_normality_s2_p': p_normal_2,
  693. f'{metric_name}_significant': p_ttest < 0.05
  694. }
  695. # Analyze all metrics
  696. dice_results = analyze_metric(dice_s1, dice_s2, 'dice')
  697. iou_results = analyze_metric(iou_s1, iou_s2, 'iou')
  698. hd95_results = analyze_metric(hd95_s1, hd95_s2, 'hd95', lower_is_better=True)
  699. assd_results = analyze_metric(assd_s1, assd_s2, 'assd', lower_is_better=True)
  700. # Combine all results
  701. statistical_results = {
  702. 'sample_size': len(dice_s1),
  703. 'sample_size_hd95': len(hd95_s1), # May be different due to filtering
  704. 'sample_size_assd': len(assd_s1),
  705. **dice_results,
  706. **iou_results,
  707. **hd95_results,
  708. **assd_results
  709. }
  710. # Save statistical results
  711. with open(self.results_dir / 'statistics' / 'statistical_analysis.json', 'w') as f:
  712. json.dump(statistical_results, f, indent=2, default=str)
  713. # Create comprehensive statistical report
  714. report = f"""
  715. COMPREHENSIVE STATISTICAL ANALYSIS REPORT
  716. ==========================================
  717. Sample Sizes:
  718. - Dice/IoU: {len(dice_s1)} test images
  719. - HD95: {len(hd95_s1)} test images (after filtering invalid cases)
  720. - ASSD: {len(assd_s1)} test images (after filtering invalid cases)
  721. DICE COEFFICIENT ANALYSIS:
  722. --------------------------
  723. Scenario 1 (Binary): {dice_results['dice_scenario1_mean']:.4f} ± {dice_results['dice_scenario1_std']:.4f} (median: {dice_results['dice_scenario1_median']:.4f})
  724. Scenario 2 (Three-class): {dice_results['dice_scenario2_mean']:.4f} ± {dice_results['dice_scenario2_std']:.4f} (median: {dice_results['dice_scenario2_median']:.4f})
  725. Improvement: {dice_results['dice_improvement']:.4f} ({dice_results['dice_improvement_percent']:.2f}%)
  726. Paired t-test: t = {dice_results['dice_t_statistic']:.4f}, p = {dice_results['dice_p_value']:.4f}
  727. Wilcoxon signed-rank: W = {dice_results['dice_wilcoxon_statistic']:.4f}, p = {dice_results['dice_wilcoxon_p_value']:.4f}
  728. Effect Size (Cohen's d): {dice_results['dice_effect_size']:.4f}
  729. 95% CI for difference: [{dice_results['dice_ci_lower']:.4f}, {dice_results['dice_ci_upper']:.4f}]
  730. Result: {'SIGNIFICANT' if dice_results['dice_significant'] else 'NOT SIGNIFICANT'}
  731. IoU COEFFICIENT ANALYSIS:
  732. -------------------------
  733. Scenario 1 (Binary): {iou_results['iou_scenario1_mean']:.4f} ± {iou_results['iou_scenario1_std']:.4f} (median: {iou_results['iou_scenario1_median']:.4f})
  734. Scenario 2 (Three-class): {iou_results['iou_scenario2_mean']:.4f} ± {iou_results['iou_scenario2_std']:.4f} (median: {iou_results['iou_scenario2_median']:.4f})
  735. Improvement: {iou_results['iou_improvement']:.4f} ({iou_results['iou_improvement_percent']:.2f}%)
  736. Paired t-test: t = {iou_results['iou_t_statistic']:.4f}, p = {iou_results['iou_p_value']:.4f}
  737. Wilcoxon signed-rank: W = {iou_results['iou_wilcoxon_statistic']:.4f}, p = {iou_results['iou_wilcoxon_p_value']:.4f}
  738. Effect Size (Cohen's d): {iou_results['iou_effect_size']:.4f}
  739. 95% CI for difference: [{iou_results['iou_ci_lower']:.4f}, {iou_results['iou_ci_upper']:.4f}]
  740. Result: {'SIGNIFICANT' if iou_results['iou_significant'] else 'NOT SIGNIFICANT'}
  741. HD95 (95th Percentile Hausdorff Distance) ANALYSIS:
  742. --------------------------------------------------
  743. Scenario 1 (Binary): {hd95_results['hd95_scenario1_mean']:.4f} ± {hd95_results['hd95_scenario1_std']:.4f} pixels (median: {hd95_results['hd95_scenario1_median']:.4f})
  744. Scenario 2 (Three-class): {hd95_results['hd95_scenario2_mean']:.4f} ± {hd95_results['hd95_scenario2_std']:.4f} pixels (median: {hd95_results['hd95_scenario2_median']:.4f})
  745. Improvement: {hd95_results['hd95_improvement']:.4f} pixels ({hd95_results['hd95_improvement_percent']:.2f}% reduction)
  746. Paired t-test: t = {hd95_results['hd95_t_statistic']:.4f}, p = {hd95_results['hd95_p_value']:.4f}
  747. Wilcoxon signed-rank: W = {hd95_results['hd95_wilcoxon_statistic']:.4f}, p = {hd95_results['hd95_wilcoxon_p_value']:.4f}
  748. Effect Size (Cohen's d): {hd95_results['hd95_effect_size']:.4f}
  749. 95% CI for difference: [{hd95_results['hd95_ci_lower']:.4f}, {hd95_results['hd95_ci_upper']:.4f}]
  750. Result: {'SIGNIFICANT' if hd95_results['hd95_significant'] else 'NOT SIGNIFICANT'}
  751. ASSD (Average Symmetric Surface Distance) ANALYSIS:
  752. --------------------------------------------------
  753. Scenario 1 (Binary): {assd_results['assd_scenario1_mean']:.4f} ± {assd_results['assd_scenario1_std']:.4f} pixels (median: {assd_results['assd_scenario1_median']:.4f})
  754. Scenario 2 (Three-class): {assd_results['assd_scenario2_mean']:.4f} ± {assd_results['assd_scenario2_std']:.4f} pixels (median: {assd_results['assd_scenario2_median']:.4f})
  755. Improvement: {assd_results['assd_improvement']:.4f} pixels ({assd_results['assd_improvement_percent']:.2f}% reduction)
  756. Paired t-test: t = {assd_results['assd_t_statistic']:.4f}, p = {assd_results['assd_p_value']:.4f}
  757. Wilcoxon signed-rank: W = {assd_results['assd_wilcoxon_statistic']:.4f}, p = {assd_results['assd_wilcoxon_p_value']:.4f}
  758. Effect Size (Cohen's d): {assd_results['assd_effect_size']:.4f}
  759. 95% CI for difference: [{assd_results['assd_ci_lower']:.4f}, {assd_results['assd_ci_upper']:.4f}]
  760. Result: {'SIGNIFICANT' if assd_results['assd_significant'] else 'NOT SIGNIFICANT'}
  761. NORMALITY TESTS:
  762. ----------------
  763. Dice - Scenario 1 p-value: {dice_results['dice_normality_s1_p']:.4f} {'(Normal)' if dice_results['dice_normality_s1_p'] > 0.05 else '(Non-normal)'}
  764. Dice - Scenario 2 p-value: {dice_results['dice_normality_s2_p']:.4f} {'(Normal)' if dice_results['dice_normality_s2_p'] > 0.05 else '(Non-normal)'}
  765. IoU - Scenario 1 p-value: {iou_results['iou_normality_s1_p']:.4f} {'(Normal)' if iou_results['iou_normality_s1_p'] > 0.05 else '(Non-normal)'}
  766. IoU - Scenario 2 p-value: {iou_results['iou_normality_s2_p']:.4f} {'(Normal)' if iou_results['iou_normality_s2_p'] > 0.05 else '(Non-normal)'}
  767. HD95 - Scenario 1 p-value: {hd95_results['hd95_normality_s1_p']:.4f} {'(Normal)' if hd95_results['hd95_normality_s1_p'] > 0.05 else '(Non-normal)'}
  768. HD95 - Scenario 2 p-value: {hd95_results['hd95_normality_s2_p']:.4f} {'(Normal)' if hd95_results['hd95_normality_s2_p'] > 0.05 else '(Non-normal)'}
  769. ASSD - Scenario 1 p-value: {assd_results['assd_normality_s1_p']:.4f} {'(Normal)' if assd_results['assd_normality_s1_p'] > 0.05 else '(Non-normal)'}
  770. ASSD - Scenario 2 p-value: {assd_results['assd_normality_s2_p']:.4f} {'(Normal)' if assd_results['assd_normality_s2_p'] > 0.05 else '(Non-normal)'}
  771. OVERALL CONCLUSIONS:
  772. -------------------
  773. Dice Improvement: {'STATISTICALLY SIGNIFICANT' if dice_results['dice_significant'] else 'NOT SIGNIFICANT'}
  774. IoU Improvement: {'STATISTICALLY SIGNIFICANT' if iou_results['iou_significant'] else 'NOT SIGNIFICANT'}
  775. HD95 Improvement: {'STATISTICALLY SIGNIFICANT' if hd95_results['hd95_significant'] else 'NOT SIGNIFICANT'}
  776. ASSD Improvement: {'STATISTICALLY SIGNIFICANT' if assd_results['assd_significant'] else 'NOT SIGNIFICANT'}
  777. Note: For HD95 and ASSD, lower values indicate better boundary accuracy.
  778. """
  779. with open(self.results_dir / 'statistics' / 'statistical_report.txt', 'w') as f:
  780. f.write(report)
  781. return statistical_results
  782. def generate_leverage_summary(self, config, dataset_info_train, dataset_info_test,
  783. metrics_s1, metrics_s2, statistical_results):
  784. """Generate comprehensive leverage paper summary with all metrics"""
  785. summary = f"""
  786. LEVERAGE PAPER RESULTS SUMMARY
  787. ================================
  788. Experiment Timestamp: {config.timestamp}
  789. Model Architecture: {config.model_name.upper()}
  790. WMH Segmentation: Binary vs Three-class Classification Comparison
  791. DATASET INFORMATION:
  792. --------------------
  793. Training Images: {dataset_info_train['loaded_files']}
  794. Test Images: {dataset_info_test['loaded_files']}
  795. Image Size: {config.target_size}
  796. Classes: Background (0), Normal WMH (1), Abnormal WMH (2)
  797. METHODOLOGY:
  798. ------------
  799. Architecture: {config.model_name.upper()}
  800. Loss Functions:
  801. - Scenario 1: {config.loss_options['scenario1']}
  802. - Scenario 2: {config.loss_options['scenario2']}
  803. Training Epochs: {config.epochs}
  804. Batch Size: {config.batch_size}
  805. Learning Rate: {config.learning_rate}
  806. PERFORMANCE RESULTS:
  807. --------------------
  808. OVERLAP-BASED METRICS:
  809. | Scenario 1 (Binary) | Scenario 2 (3-class) | Improvement
  810. --------------------|---------------------|----------------------|------------
  811. Accuracy | {metrics_s1['Accuracy']:.4f} | {metrics_s2['Accuracy']:.4f} | {metrics_s2['Accuracy'] - metrics_s1['Accuracy']:+.4f}
  812. Precision | {metrics_s1['Precision']:.4f} | {metrics_s2['Precision']:.4f} | {metrics_s2['Precision'] - metrics_s1['Precision']:+.4f}
  813. Recall | {metrics_s1['Recall']:.4f} | {metrics_s2['Recall']:.4f} | {metrics_s2['Recall'] - metrics_s1['Recall']:+.4f}
  814. Specificity | {metrics_s1['Specificity']:.4f} | {metrics_s2['Specificity']:.4f} | {metrics_s2['Specificity'] - metrics_s1['Specificity']:+.4f}
  815. Dice Coefficient | {metrics_s1['Dice']:.4f} | {metrics_s2['Dice']:.4f} | {metrics_s2['Dice'] - metrics_s1['Dice']:+.4f}
  816. IoU Coefficient | {metrics_s1['IoU']:.4f} | {metrics_s2['IoU']:.4f} | {metrics_s2['IoU'] - metrics_s1['IoU']:+.4f}
  817. SURFACE-BASED METRICS (lower is better):
  818. | Scenario 1 (Binary) | Scenario 2 (3-class) | Improvement
  819. --------------------|---------------------|----------------------|------------
  820. HD95 (pixels) | {statistical_results['hd95_scenario1_mean']:.4f} ± {statistical_results['hd95_scenario1_std']:.4f} | {statistical_results['hd95_scenario2_mean']:.4f} ± {statistical_results['hd95_scenario2_std']:.4f} | {statistical_results['hd95_improvement']:+.4f}
  821. ASSD (pixels) | {statistical_results['assd_scenario1_mean']:.4f} ± {statistical_results['assd_scenario1_std']:.4f} | {statistical_results['assd_scenario2_mean']:.4f} ± {statistical_results['assd_scenario2_std']:.4f} | {statistical_results['assd_improvement']:+.4f}
  822. Note: For HD95 and ASSD, positive improvement means reduction (better boundary accuracy)
  823. Valid samples: HD95={statistical_results['sample_size_hd95']}/{statistical_results['sample_size']}, ASSD={statistical_results['sample_size_assd']}/{statistical_results['sample_size']}
  824. STATISTICAL SIGNIFICANCE:
  825. -------------------------
  826. DICE COEFFICIENT:
  827. Test: Paired t-test
  828. t-statistic: {statistical_results['dice_t_statistic']:.4f}
  829. p-value: {statistical_results['dice_p_value']:.4f}
  830. Effect Size (Cohen's d): {statistical_results['dice_effect_size']:.4f}
  831. 95% Confidence Interval: [{statistical_results['dice_ci_lower']:.4f}, {statistical_results['dice_ci_upper']:.4f}]
  832. Result: {'SIGNIFICANT' if statistical_results['dice_significant'] else 'NOT SIGNIFICANT'} improvement
  833. IoU COEFFICIENT:
  834. Test: Paired t-test
  835. t-statistic: {statistical_results['iou_t_statistic']:.4f}
  836. p-value: {statistical_results['iou_p_value']:.4f}
  837. Effect Size (Cohen's d): {statistical_results['iou_effect_size']:.4f}
  838. 95% Confidence Interval: [{statistical_results['iou_ci_lower']:.4f}, {statistical_results['iou_ci_upper']:.4f}]
  839. Result: {'SIGNIFICANT' if statistical_results['iou_significant'] else 'NOT SIGNIFICANT'} improvement
  840. HD95 (95th Percentile Hausdorff Distance):
  841. Test: Paired t-test
  842. t-statistic: {statistical_results['hd95_t_statistic']:.4f}
  843. p-value: {statistical_results['hd95_p_value']:.4f}
  844. Effect Size (Cohen's d): {statistical_results['hd95_effect_size']:.4f}
  845. 95% Confidence Interval: [{statistical_results['hd95_ci_lower']:.4f}, {statistical_results['hd95_ci_upper']:.4f}] pixels
  846. Result: {'SIGNIFICANT' if statistical_results['hd95_significant'] else 'NOT SIGNIFICANT'} improvement
  847. ASSD (Average Symmetric Surface Distance):
  848. Test: Paired t-test
  849. t-statistic: {statistical_results['assd_t_statistic']:.4f}
  850. p-value: {statistical_results['assd_p_value']:.4f}
  851. Effect Size (Cohen's d): {statistical_results['assd_effect_size']:.4f}
  852. 95% Confidence Interval: [{statistical_results['assd_ci_lower']:.4f}, {statistical_results['assd_ci_upper']:.4f}] pixels
  853. Result: {'SIGNIFICANT' if statistical_results['assd_significant'] else 'NOT SIGNIFICANT'} improvement
  854. KEY FINDINGS:
  855. -------------
  856. OVERLAP-BASED METRICS:
  857. 1. Three-class segmentation shows {statistical_results['dice_improvement_percent']:.2f}% improvement in Dice coefficient
  858. 2. Three-class segmentation shows {statistical_results['iou_improvement_percent']:.2f}% improvement in IoU coefficient
  859. 3. Dice improvement is {'statistically significant (p<0.05)' if statistical_results['dice_significant'] else 'not statistically significant'}
  860. 4. IoU improvement is {'statistically significant (p<0.05)' if statistical_results['iou_significant'] else 'not statistically significant'}
  861. SURFACE-BASED METRICS:
  862. 5. HD95 shows {abs(statistical_results['hd95_improvement_percent']):.2f}% {'reduction' if statistical_results['hd95_improvement'] > 0 else 'increase'} (lower is better)
  863. 6. ASSD shows {abs(statistical_results['assd_improvement_percent']):.2f}% {'reduction' if statistical_results['assd_improvement'] > 0 else 'increase'} (lower is better)
  864. 7. HD95 improvement is {'statistically significant (p<0.05)' if statistical_results['hd95_significant'] else 'not statistically significant'}
  865. 8. ASSD improvement is {'statistically significant (p<0.05)' if statistical_results['assd_significant'] else 'not statistically significant'}
  866. OVERALL ASSESSMENT:
  867. 9. Post-processing provided substantial improvements in both scenarios
  868. 10. Three-class approach shows {'consistent advantages' if statistical_results['dice_significant'] and statistical_results['iou_significant'] else 'mixed results'} across multiple metrics
  869. 11. Boundary accuracy (HD95/ASSD) {'improved significantly' if statistical_results['hd95_significant'] or statistical_results['assd_significant'] else 'showed no significant improvement'}
  870. FILES GENERATED:
  871. ----------------
  872. - Models: scenario1_binary_model.h5, scenario2_multiclass_model.h5
  873. - Figures: training_curves.png/.pdf, comparison_visualization.png/.pdf, metrics_comparison.png/.pdf
  874. - Tables: comprehensive_results.csv/.xlsx, surface_metrics.csv/.xlsx, latex_table.tex, latex_surface_table.tex
  875. - Statistics: statistical_analysis.json, statistical_report.txt
  876. - Predictions: All test predictions and ground truth data saved
  877. PUBLICATION READINESS:
  878. ----------------------
  879. ✓ High-resolution figures (300 DPI, PNG/PDF)
  880. ✓ LaTeX-formatted tables (overlap and surface metrics)
  881. ✓ Comprehensive statistical analysis (Dice, IoU, HD95, ASSD)
  882. ✓ Post-processing impact analysis
  883. ✓ Reproducible results with saved models
  884. ✓ Professional documentation
  885. ✓ Surface-based metrics for boundary accuracy assessment
  886. """
  887. with open(self.results_dir / 'leverage_summary.txt', 'w') as f:
  888. f.write(summary)
  889. print("="*80)
  890. print("LEVERAGE RESULTS SUMMARY GENERATED")
  891. print("="*80)
  892. print(summary)
  893. ###################### Main Experiment Function ######################
  894. def run_leverage_experiment():
  895. """Main function to run the complete leverage experiment"""
  896. print("="*80)
  897. print("STARTING LEVERAGE PAPER EXPERIMENT")
  898. print("="*80)
  899. # Initialize components
  900. plotter = PublicationPlotter(config.results_dir)
  901. saver = ResultsSaver(config.results_dir)
  902. # Load datasets
  903. print("\nLoading datasets...")
  904. train_images, train_masks_3class, train_masks_binary, dataset_info_train = load_wmh_dataset(
  905. config.train_dir, config.target_size
  906. )
  907. test_images, test_masks_3class, test_masks_binary, dataset_info_test = load_wmh_dataset(
  908. config.test_dir, config.target_size
  909. )
  910. # Split training data
  911. x_train, x_val, y_train_3class, y_val_3class, y_train_binary, y_val_binary = train_test_split(
  912. train_images, train_masks_3class, train_masks_binary,
  913. test_size=config.validation_split, random_state=config.random_state
  914. )
  915. print(f"Training: {x_train.shape[0]}, Validation: {x_val.shape[0]}, Test: {test_images.shape[0]}")
  916. # Calculate class weights
  917. binary_weights = calculate_class_weights(y_train_binary, 2)
  918. multiclass_weights = calculate_class_weights(y_train_3class, 3)
  919. print(f"Binary class weights: {binary_weights}")
  920. print(f"Multi-class weights: {multiclass_weights}")
  921. # Scenario 1: Binary Classification
  922. print("\n" + "="*60)
  923. print("TRAINING SCENARIO 1: BINARY CLASSIFICATION")
  924. print("="*60)
  925. model_s1 = config.build_unet_variant(config.input_shape, num_classes=1)
  926. print(f"Scenario 1 Model Parameters: {model_s1.count_params():,}")
  927. model_s1.summary() # Optional: for detailed architecture view
  928. # Configure loss function
  929. if config.loss_options['scenario1'] == 'weighted_bce':
  930. loss_s1 = weighted_binary_crossentropy(pos_weight=binary_weights[1])
  931. elif config.loss_options['scenario1'] == 'focal':
  932. loss_s1 = focal_loss(alpha=0.75, gamma=2.0)
  933. elif config.loss_options['scenario1'] == 'combined':
  934. loss_s1 = combined_loss(alpha=0.5, pos_weight=binary_weights[1])
  935. else:
  936. loss_s1 = 'binary_crossentropy'
  937. model_s1.compile(
  938. optimizer=optimizers.legacy.Adam(config.learning_rate),
  939. loss=loss_s1,
  940. metrics=['accuracy']
  941. )
  942. # Callbacks
  943. callbacks_s1 = [
  944. callbacks.EarlyStopping(patience=15, restore_best_weights=True),
  945. callbacks.ReduceLROnPlateau(patience=10, factor=0.5, min_lr=1e-7)
  946. ]
  947. y_train_binary_expanded = np.expand_dims(y_train_binary, axis=-1)
  948. y_val_binary_expanded = np.expand_dims(y_val_binary, axis=-1)
  949. if config.mode == 'training':
  950. history_s1 = model_s1.fit(
  951. x_train, y_train_binary_expanded,
  952. validation_data=(x_val, y_val_binary_expanded),
  953. epochs=config.epochs,
  954. batch_size=config.batch_size,
  955. callbacks=callbacks_s1,
  956. verbose=1
  957. )
  958. # Scenario 2: Three-class Classification
  959. print("\n" + "="*60)
  960. print("TRAINING SCENARIO 2: THREE-CLASS CLASSIFICATION")
  961. print("="*60)
  962. model_s2 = config.build_unet_variant(config.input_shape, num_classes=3)
  963. print(f"Scenario 2 Model Parameters: {model_s2.count_params():,}")
  964. model_s2.summary() # Optional: for detailed architecture view
  965. # Configure loss function
  966. if config.loss_options['scenario2'] == 'weighted_categorical':
  967. loss_s2 = weighted_categorical_crossentropy(multiclass_weights)
  968. elif config.loss_options['scenario2'] == 'multiclass_dice':
  969. loss_s2 = multiclass_dice_loss(num_classes=3)
  970. else:
  971. loss_s2 = 'categorical_crossentropy'
  972. model_s2.compile(
  973. optimizer=optimizers.legacy.Adam(config.learning_rate),
  974. loss=loss_s2,
  975. metrics=['accuracy']
  976. )
  977. callbacks_s2 = [
  978. callbacks.EarlyStopping(patience=15, restore_best_weights=True),
  979. callbacks.ReduceLROnPlateau(patience=10, factor=0.5, min_lr=1e-7)
  980. ]
  981. y_train_3class_categorical = to_categorical(y_train_3class, num_classes=3)
  982. y_val_3class_categorical = to_categorical(y_val_3class, num_classes=3)
  983. if config.mode == 'training':
  984. history_s2 = model_s2.fit(
  985. x_train, y_train_3class_categorical,
  986. validation_data=(x_val, y_val_3class_categorical),
  987. epochs=config.epochs,
  988. batch_size=config.batch_size,
  989. callbacks=callbacks_s2,
  990. verbose=1
  991. )
  992. # Save models
  993. if config.mode == 'training':
  994. saver.save_models(model_s1, model_s2, history_s1, history_s2)
  995. else:
  996. # or Load models
  997. model_s1, model_s2, history_s1, history_s2 = saver.load_models(config.intended_study_dir, loss_s1_func=loss_s1, loss_s2_func=loss_s2)
  998. # Check if models loaded successfully before using them
  999. if model_s1 is None or model_s2 is None:
  1000. print("Failed to load models. Cannot proceed with predictions.")
  1001. exit(1)
  1002. # Generate predictions
  1003. print("\n" + "="*60)
  1004. print("GENERATING PREDICTIONS AND EVALUATION")
  1005. print("="*60)
  1006. test_pred_s1 = model_s1.predict(test_images, batch_size=config.batch_size)
  1007. test_pred_s1_binary = (test_pred_s1.squeeze() > 0.5).astype(np.uint8)
  1008. test_pred_s2 = model_s2.predict(test_images, batch_size=config.batch_size)
  1009. test_pred_s2_classes = np.argmax(test_pred_s2, axis=-1)
  1010. # Post-process predictions
  1011. print("Applying post-processing...")
  1012. test_pred_s1_binary_processed = post_process_predictions(
  1013. test_pred_s1_binary,
  1014. min_object_size=5,
  1015. apply_opening=True,
  1016. kernel_size=2
  1017. )
  1018. abnormal_pred_from_3class = (test_pred_s2_classes == 2).astype(np.uint8)
  1019. abnormal_pred_from_3class_processed = post_process_predictions(
  1020. abnormal_pred_from_3class,
  1021. min_object_size=5,
  1022. apply_opening=True,
  1023. kernel_size=2
  1024. )
  1025. # Calculate metrics for raw predictions (without surface metrics for aggregate)
  1026. print("Calculating metrics for raw predictions...")
  1027. metrics_s1_raw = calculate_comprehensive_metrics(
  1028. test_masks_binary.flatten(),
  1029. test_pred_s1_binary.flatten(),
  1030. "Binary Classification (Raw)"
  1031. )
  1032. abnormal_true_from_3class = (test_masks_3class == 2).astype(np.uint8)
  1033. metrics_s2_raw = calculate_comprehensive_metrics(
  1034. abnormal_true_from_3class.flatten(),
  1035. abnormal_pred_from_3class.flatten(),
  1036. "Three-class Classification (Raw)"
  1037. )
  1038. # Calculate metrics for post-processed predictions (without surface metrics for aggregate)
  1039. print("Calculating metrics for post-processed predictions...")
  1040. metrics_s1 = calculate_comprehensive_metrics(
  1041. test_masks_binary.flatten(),
  1042. test_pred_s1_binary_processed.flatten(),
  1043. "Binary Classification (Processed)"
  1044. )
  1045. metrics_s2 = calculate_comprehensive_metrics(
  1046. abnormal_true_from_3class.flatten(),
  1047. abnormal_pred_from_3class_processed.flatten(),
  1048. "Three-class Classification (Processed)"
  1049. )
  1050. # Calculate per-image metrics for statistical analysis (including HD95 and ASSD)
  1051. dice_scores_s1_raw, dice_scores_s2_raw = [], []
  1052. dice_scores_s1, dice_scores_s2 = [], []
  1053. iou_scores_s1_raw, iou_scores_s2_raw = [], []
  1054. iou_scores_s1, iou_scores_s2 = [], []
  1055. hd95_scores_s1_raw, hd95_scores_s2_raw = [], []
  1056. hd95_scores_s1, hd95_scores_s2 = [], []
  1057. assd_scores_s1_raw, assd_scores_s2_raw = [], []
  1058. assd_scores_s1, assd_scores_s2 = [], []
  1059. for i in range(test_images.shape[0]):
  1060. # Raw predictions
  1061. dice_s1_raw = dice_coefficient_multiclass(test_masks_binary[i].flatten(), test_pred_s1_binary[i].flatten(), 1)
  1062. dice_s2_raw = dice_coefficient_multiclass(abnormal_true_from_3class[i].flatten(), abnormal_pred_from_3class[i].flatten(), 1)
  1063. iou_s1_raw = iou_coefficient_multiclass(test_masks_binary[i].flatten(), test_pred_s1_binary[i].flatten(), 1)
  1064. iou_s2_raw = iou_coefficient_multiclass(abnormal_true_from_3class[i].flatten(), abnormal_pred_from_3class[i].flatten(), 1)
  1065. hd95_s1_raw = hausdorff_distance_95(test_masks_binary[i], test_pred_s1_binary[i])
  1066. hd95_s2_raw = hausdorff_distance_95(abnormal_true_from_3class[i], abnormal_pred_from_3class[i])
  1067. assd_s1_raw = average_symmetric_surface_distance(test_masks_binary[i], test_pred_s1_binary[i])
  1068. assd_s2_raw = average_symmetric_surface_distance(abnormal_true_from_3class[i], abnormal_pred_from_3class[i])
  1069. # Post-processed predictions
  1070. dice_s1 = dice_coefficient_multiclass(test_masks_binary[i].flatten(), test_pred_s1_binary_processed[i].flatten(), 1)
  1071. dice_s2 = dice_coefficient_multiclass(abnormal_true_from_3class[i].flatten(), abnormal_pred_from_3class_processed[i].flatten(), 1)
  1072. iou_s1 = iou_coefficient_multiclass(test_masks_binary[i].flatten(), test_pred_s1_binary_processed[i].flatten(), 1)
  1073. iou_s2 = iou_coefficient_multiclass(abnormal_true_from_3class[i].flatten(), abnormal_pred_from_3class_processed[i].flatten(), 1)
  1074. hd95_s1 = hausdorff_distance_95(test_masks_binary[i], test_pred_s1_binary_processed[i])
  1075. hd95_s2 = hausdorff_distance_95(abnormal_true_from_3class[i], abnormal_pred_from_3class_processed[i])
  1076. assd_s1 = average_symmetric_surface_distance(test_masks_binary[i], test_pred_s1_binary_processed[i])
  1077. assd_s2 = average_symmetric_surface_distance(abnormal_true_from_3class[i], abnormal_pred_from_3class_processed[i])
  1078. dice_scores_s1_raw.append(dice_s1_raw)
  1079. dice_scores_s2_raw.append(dice_s2_raw)
  1080. dice_scores_s1.append(dice_s1)
  1081. dice_scores_s2.append(dice_s2)
  1082. iou_scores_s1_raw.append(iou_s1_raw)
  1083. iou_scores_s2_raw.append(iou_s2_raw)
  1084. iou_scores_s1.append(iou_s1)
  1085. iou_scores_s2.append(iou_s2)
  1086. hd95_scores_s1_raw.append(hd95_s1_raw)
  1087. hd95_scores_s2_raw.append(hd95_s2_raw)
  1088. hd95_scores_s1.append(hd95_s1)
  1089. hd95_scores_s2.append(hd95_s2)
  1090. assd_scores_s1_raw.append(assd_s1_raw)
  1091. assd_scores_s2_raw.append(assd_s2_raw)
  1092. assd_scores_s1.append(assd_s1)
  1093. assd_scores_s2.append(assd_s2)
  1094. dice_scores_s1_raw = np.array(dice_scores_s1_raw)
  1095. dice_scores_s2_raw = np.array(dice_scores_s2_raw)
  1096. dice_scores_s1 = np.array(dice_scores_s1)
  1097. dice_scores_s2 = np.array(dice_scores_s2)
  1098. iou_scores_s1_raw = np.array(iou_scores_s1_raw)
  1099. iou_scores_s2_raw = np.array(iou_scores_s2_raw)
  1100. iou_scores_s1 = np.array(iou_scores_s1)
  1101. iou_scores_s2 = np.array(iou_scores_s2)
  1102. hd95_scores_s1_raw = np.array(hd95_scores_s1_raw)
  1103. hd95_scores_s2_raw = np.array(hd95_scores_s2_raw)
  1104. hd95_scores_s1 = np.array(hd95_scores_s1)
  1105. hd95_scores_s2 = np.array(hd95_scores_s2)
  1106. assd_scores_s1_raw = np.array(assd_scores_s1_raw)
  1107. assd_scores_s2_raw = np.array(assd_scores_s2_raw)
  1108. assd_scores_s1 = np.array(assd_scores_s1)
  1109. assd_scores_s2 = np.array(assd_scores_s2)
  1110. # Filter out inf values for HD95 and ASSD while maintaining pairing
  1111. # Create masks for valid (finite) values in both scenarios
  1112. hd95_valid_mask = np.isfinite(hd95_scores_s1) & np.isfinite(hd95_scores_s2)
  1113. assd_valid_mask = np.isfinite(assd_scores_s1) & np.isfinite(assd_scores_s2)
  1114. hd95_valid_mask_raw = np.isfinite(hd95_scores_s1_raw) & np.isfinite(hd95_scores_s2_raw)
  1115. assd_valid_mask_raw = np.isfinite(assd_scores_s1_raw) & np.isfinite(assd_scores_s2_raw)
  1116. # Apply masks to both scenarios to keep them paired
  1117. hd95_scores_s1 = hd95_scores_s1[hd95_valid_mask]
  1118. hd95_scores_s2 = hd95_scores_s2[hd95_valid_mask]
  1119. assd_scores_s1 = assd_scores_s1[assd_valid_mask]
  1120. assd_scores_s2 = assd_scores_s2[assd_valid_mask]
  1121. hd95_scores_s1_raw = hd95_scores_s1_raw[hd95_valid_mask_raw]
  1122. hd95_scores_s2_raw = hd95_scores_s2_raw[hd95_valid_mask_raw]
  1123. assd_scores_s1_raw = assd_scores_s1_raw[assd_valid_mask_raw]
  1124. assd_scores_s2_raw = assd_scores_s2_raw[assd_valid_mask_raw]
  1125. print(f"\nValid samples after filtering infinite values:")
  1126. print(f"HD95: {len(hd95_scores_s1)} / {len(dice_scores_s1)} images")
  1127. print(f"ASSD: {len(assd_scores_s1)} / {len(dice_scores_s1)} images")
  1128. # Print comparison of raw vs processed results
  1129. print(f"\nPost-processing Impact:")
  1130. print(f"Scenario 1 - Dice improvement: {np.mean(dice_scores_s1) - np.mean(dice_scores_s1_raw):+.4f}")
  1131. print(f"Scenario 1 - IoU improvement: {np.mean(iou_scores_s1) - np.mean(iou_scores_s1_raw):+.4f}")
  1132. print(f"Scenario 1 - HD95 improvement: {np.mean(hd95_scores_s1_raw) - np.mean(hd95_scores_s1):+.4f} pixels (lower is better)")
  1133. print(f"Scenario 1 - ASSD improvement: {np.mean(assd_scores_s1_raw) - np.mean(assd_scores_s1):+.4f} pixels (lower is better)")
  1134. print(f"Scenario 2 - Dice improvement: {np.mean(dice_scores_s2) - np.mean(dice_scores_s2_raw):+.4f}")
  1135. print(f"Scenario 2 - IoU improvement: {np.mean(iou_scores_s2) - np.mean(iou_scores_s2_raw):+.4f}")
  1136. print(f"Scenario 2 - HD95 improvement: {np.mean(hd95_scores_s2_raw) - np.mean(hd95_scores_s2):+.4f} pixels (lower is better)")
  1137. print(f"Scenario 2 - ASSD improvement: {np.mean(assd_scores_s2_raw) - np.mean(assd_scores_s2):+.4f} pixels (lower is better)")
  1138. # Statistical analysis
  1139. print("\nPerforming statistical analysis...")
  1140. statistical_results = saver.save_statistical_analysis(
  1141. dice_scores_s1, dice_scores_s2, iou_scores_s1, iou_scores_s2,
  1142. hd95_scores_s1, hd95_scores_s2, assd_scores_s1, assd_scores_s2,
  1143. metrics_s1, metrics_s2
  1144. )
  1145. # Save all results
  1146. print("\nSaving results...")
  1147. saver.save_predictions(
  1148. test_images, test_masks_3class, test_masks_binary,
  1149. test_pred_s1_binary, test_pred_s2_classes,
  1150. {'train': dataset_info_train, 'test': dataset_info_test}
  1151. )
  1152. results_table, surface_table = saver.save_metrics_table(metrics_s1, metrics_s2, statistical_results)
  1153. # Generate visualizations
  1154. print("\nGenerating publication-quality figures...")
  1155. plotter.plot_training_curves(history_s1, history_s2)
  1156. plotter.plot_comparison_visualization(
  1157. test_images, test_masks_3class, test_masks_binary,
  1158. test_pred_s1_binary_processed, abnormal_pred_from_3class_processed
  1159. )
  1160. plotter.plot_metrics_comparison(metrics_s1, metrics_s2)
  1161. plotter.plot_dice_distribution(dice_scores_s1, dice_scores_s2)
  1162. plotter.plot_dice_iou_distribution(dice_scores_s1, dice_scores_s2, iou_scores_s1, iou_scores_s2)
  1163. # Generate final summary
  1164. saver.generate_leverage_summary(
  1165. config, dataset_info_train, dataset_info_test,
  1166. metrics_s1, metrics_s2, statistical_results
  1167. )
  1168. return {
  1169. 'config': config,
  1170. 'models': {'scenario1': model_s1, 'scenario2': model_s2},
  1171. 'histories': {'scenario1': history_s1, 'scenario2': history_s2},
  1172. 'metrics': {'scenario1': metrics_s1, 'scenario2': metrics_s2},
  1173. 'statistical_results': statistical_results,
  1174. 'results_table': results_table,
  1175. 'surface_table': surface_table
  1176. }
  1177. def create_comparative_analysis(all_results, models_to_test):
  1178. """Create a comparative analysis across all models"""
  1179. import pandas as pd
  1180. from pathlib import Path
  1181. # Create a comparative results directory
  1182. comparative_dir = Path(f"leverage_comparative_results_{datetime.now().strftime('%Y%m%d_%H%M%S')}")
  1183. comparative_dir.mkdir(exist_ok=True)
  1184. # Collect data for comparison
  1185. comparison_data = []
  1186. for model_name in models_to_test:
  1187. if all_results[model_name] is not None:
  1188. result = all_results[model_name]
  1189. metrics_s1 = result['metrics']['scenario1']
  1190. metrics_s2 = result['metrics']['scenario2']
  1191. stats = result['statistical_results']
  1192. # Scenario 1
  1193. comparison_data.append({
  1194. 'Model': model_name,
  1195. 'Scenario': 'Binary (S1)',
  1196. 'Accuracy': metrics_s1['Accuracy'],
  1197. 'Precision': metrics_s1['Precision'],
  1198. 'Recall': metrics_s1['Recall'],
  1199. 'Dice': metrics_s1['Dice'],
  1200. 'IoU': metrics_s1['IoU']
  1201. })
  1202. # Scenario 2
  1203. comparison_data.append({
  1204. 'Model': model_name,
  1205. 'Scenario': 'Three-class (S2)',
  1206. 'Accuracy': metrics_s2['Accuracy'],
  1207. 'Precision': metrics_s2['Precision'],
  1208. 'Recall': metrics_s2['Recall'],
  1209. 'Dice': metrics_s2['Dice'],
  1210. 'IoU': metrics_s2['IoU']
  1211. })
  1212. # Create DataFrame
  1213. df_comparison = pd.DataFrame(comparison_data)
  1214. # Save to CSV and Excel
  1215. df_comparison.to_csv(comparative_dir / 'model_comparison.csv', index=False)
  1216. df_comparison.to_excel(comparative_dir / 'model_comparison.xlsx', index=False)
  1217. # Create improvement summary
  1218. improvement_data = []
  1219. for model_name in models_to_test:
  1220. if all_results[model_name] is not None:
  1221. stats = all_results[model_name]['statistical_results']
  1222. improvement_data.append({
  1223. 'Model': model_name,
  1224. 'Dice_Improvement': stats['dice_improvement'],
  1225. 'Dice_p_value': stats['dice_p_value'],
  1226. 'Dice_Significant': stats['dice_significant'],
  1227. 'IoU_Improvement': stats['iou_improvement'],
  1228. 'IoU_p_value': stats['iou_p_value'],
  1229. 'IoU_Significant': stats['iou_significant'],
  1230. 'HD95_Improvement': stats['hd95_improvement'],
  1231. 'HD95_p_value': stats['hd95_p_value'],
  1232. 'ASSD_Improvement': stats['assd_improvement'],
  1233. 'ASSD_p_value': stats['assd_p_value']
  1234. })
  1235. df_improvement = pd.DataFrame(improvement_data)
  1236. df_improvement.to_csv(comparative_dir / 'improvement_comparison.csv', index=False)
  1237. df_improvement.to_excel(comparative_dir / 'improvement_comparison.xlsx', index=False)
  1238. # Create a summary report
  1239. summary_report = f"""
  1240. COMPARATIVE ANALYSIS ACROSS ALL MODELS
  1241. ========================================
  1242. Generated: {datetime.now().strftime("%Y-%m-%d %H:%M:%S")}
  1243. MODELS TESTED:
  1244. {', '.join([m.upper() for m in models_to_test if all_results[m] is not None])}
  1245. PERFORMANCE COMPARISON:
  1246. -----------------------
  1247. {df_comparison.to_string(index=False)}
  1248. IMPROVEMENT ANALYSIS:
  1249. ---------------------
  1250. {df_improvement.to_string(index=False)}
  1251. BEST PERFORMING MODEL:
  1252. ----------------------
  1253. Best Dice (S2): {df_comparison[df_comparison['Scenario'] == 'Three-class (S2)'].nlargest(1, 'Dice')['Model'].values[0].upper()}
  1254. Best IoU (S2): {df_comparison[df_comparison['Scenario'] == 'Three-class (S2)'].nlargest(1, 'IoU')['Model'].values[0].upper()}
  1255. Largest Dice Improvement: {df_improvement.nlargest(1, 'Dice_Improvement')['Model'].values[0].upper()}
  1256. Largest IoU Improvement: {df_improvement.nlargest(1, 'IoU_Improvement')['Model'].values[0].upper()}
  1257. Files saved in: {comparative_dir}
  1258. """
  1259. with open(comparative_dir / 'comparative_summary.txt', 'w') as f:
  1260. f.write(summary_report)
  1261. print(f"\n\nComparative analysis saved in: {comparative_dir}")
  1262. print(summary_report)
  1263. ###################### Execute Experiment ######################
  1264. if __name__ == "__main__":
  1265. # List of all models to test
  1266. models_to_test = ['unet', 'attn_unet', 'trans_unet', 'deepl3_unet']
  1267. # Dictionary to store results from all models
  1268. all_results = {}
  1269. print("\n" + "="*80)
  1270. print("STARTING MULTI-MODEL LEVERAGE PAPER EXPERIMENT")
  1271. print("="*80)
  1272. print(f"Models to test: {', '.join(models_to_test)}")
  1273. print("="*80)
  1274. # Run experiment for each model
  1275. for model_idx, model_name in enumerate(models_to_test, 1):
  1276. print("\n" + "="*80)
  1277. print(f"RUNNING EXPERIMENT {model_idx}/{len(models_to_test)}: {model_name.upper()}")
  1278. print("="*80)
  1279. # Create new config for this model
  1280. config = Config()
  1281. config.model_name = model_name
  1282. config.mode = 'no-training'
  1283. # Update build_unet_variant based on model name
  1284. if config.model_name == 'unet':
  1285. config.build_unet_variant = build_unet_3class
  1286. elif config.model_name == 'attn_unet':
  1287. config.build_unet_variant = build_attention_unet_3class
  1288. elif config.model_name == 'trans_unet':
  1289. config.build_unet_variant = build_trans_unet_3class
  1290. elif config.model_name == 'deepl3_unet':
  1291. config.build_unet_variant = build_deeplabv3_unet_3class
  1292. # Recreate results directory with correct model name
  1293. config.timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
  1294. if config.mode != 'training':
  1295. config.results_dir = Path(f"leverage_results_{config.timestamp}_{config.model_name}_no_training")
  1296. else:
  1297. config.results_dir = Path(f"leverage_results_{config.timestamp}_{config.model_name}")
  1298. config.create_directory_structure()
  1299. # if we work on non-training mode, load (use) the previously trained models
  1300. if config.model_name == 'unet':
  1301. config.intended_study_dir = "leverage_results_20251124_152044_unet"
  1302. elif config.model_name == 'attn_unet':
  1303. config.intended_study_dir = "leverage_results_20251125_133300_attn_unet"
  1304. elif config.model_name == 'trans_unet':
  1305. config.intended_study_dir = "leverage_results_20251124_171430_trans_unet"
  1306. elif config.model_name == 'deepl3_unet':
  1307. config.intended_study_dir = "leverage_results_20251124_180934_deepl3_unet"
  1308. # Set seeds for reproducibility for each model
  1309. np.random.seed(config.random_state)
  1310. tf.random.set_seed(config.random_state)
  1311. try:
  1312. # Run the complete experiment for this model
  1313. results = run_leverage_experiment()
  1314. all_results[model_name] = results
  1315. print("\n" + "="*80)
  1316. print(f"EXPERIMENT FOR {model_name.upper()} COMPLETED SUCCESSFULLY!")
  1317. print("="*80)
  1318. print(f"Results saved in: {config.results_dir}")
  1319. print("="*80)
  1320. except Exception as e:
  1321. print("\n" + "="*80)
  1322. print(f"ERROR: Experiment for {model_name.upper()} failed!")
  1323. print("="*80)
  1324. print(f"Error message: {str(e)}")
  1325. print("Continuing with next model...")
  1326. print("="*80)
  1327. all_results[model_name] = None
  1328. continue
  1329. # Generate comparative summary across all models
  1330. print("\n" + "="*80)
  1331. print("ALL EXPERIMENTS COMPLETED!")
  1332. print("="*80)
  1333. print("\nSUMMARY OF ALL MODELS:")
  1334. print("-" * 80)
  1335. for model_name in models_to_test:
  1336. if all_results[model_name] is not None:
  1337. result = all_results[model_name]
  1338. metrics_s1 = result['metrics']['scenario1']
  1339. metrics_s2 = result['metrics']['scenario2']
  1340. stats = result['statistical_results']
  1341. print(f"\n{model_name.upper()}:")
  1342. print(f" Results directory: {result['config'].results_dir}")
  1343. print(f" Scenario 1 - Dice: {metrics_s1['Dice']:.4f}, IoU: {metrics_s1['IoU']:.4f}")
  1344. print(f" Scenario 2 - Dice: {metrics_s2['Dice']:.4f}, IoU: {metrics_s2['IoU']:.4f}")
  1345. print(f" Dice Improvement: {stats['dice_improvement']:.4f} (p={stats['dice_p_value']:.4f})")
  1346. print(f" IoU Improvement: {stats['iou_improvement']:.4f} (p={stats['iou_p_value']:.4f})")
  1347. else:
  1348. print(f"\n{model_name.upper()}: FAILED")
  1349. print("\n" + "="*80)
  1350. print("All files are ready for leverage paper submission!")
  1351. print("="*80)
  1352. # Optionally: Create a comparative analysis file
  1353. create_comparative_analysis(all_results, models_to_test)

wmh_leverage_normal_inference.py at commit 1f819ef, under MIT · at the source

Overview

  1. Biomedical Engineering Faculty, Sahand University of Technology,Tabriz, Iran
  2. Department of Engineering Sciences, Faculty of Advanced Technologies, University of Mohaghegh Ardabili,Namin, Iran
  3. Radiology Department, Tabriz University of Medical Sciences,Tabriz, Iran
Journal: Biomedical engineering online, volume 25, issue 1, article 69
Dates: received 27 November 2025; accepted 9 March 2026; published online 16 April 2026
Type: Research article · Language: English
License: CC BY-NC-ND
Identifiers: DOI 10.1186/s12938-026-01555-0 · PMID 41992229 · PMCID PMC13202883 · OpenAlex W7154576188
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: structural MRI / diffusion (modality), human (organism), stroke (population), methods / tools (subfield)
Methods: Connectivity, Statistics, Preprocessing, Machine learning
Keywords: White matter hyperintensities (WMH), Deep learning, Medical image segmentation, FLAIR MRI, U-Net, Pathological segmentation, Neuroimaging
MeSH: Deep Learning*, Image Processing, Computer-Assisted*, White Matter*, Classification Algorithms, Humans, Magnetic Resonance Imaging (* major topic)
Topic: Dementia and Cognitive Impairment Research (Psychiatry and Mental health, Medicine), according to OpenAlex
Citations: not cited yet (Europe PMC); 26 references in the paper

Abstract

The abstract is not reproduced here: the paper's license (CC BY-NC-ND) does not allow it. Read it in the paper, at the publisher or on Europe PMC.

Repository

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

Mahdi-Bashiri/wmh-normal-abnormal-segmentation

License: MIT
State: the link answers, verified on 29 September 2026
Evidence: files inventoried
Commit: 1f819ef2ebcb32f9062cd1aa09c20430351b7277, 17 February 2026
Languages: Python (9), Shell (1)
Size: 137 files, 10 scripts
Software Heritage: not archived
Found in: the text, “Implementation details”
Holds: README, license file, environment (requirements.txt), documentation
Not found: CITATION.cff, tests, continuous integration
Tools: Keras (6 files), TensorFlow (4 files), Matplotlib (2 files), NumPy (2 files), OpenCV (2 files), pandas (2 files), scikit-image (2 files), scikit-learn (2 files), SciPy (2 files), seaborn (2 files)
Availability: 1 check, the latest on 29 September 2026: the link answers
  • 29 September 2026: the link answers
12 files

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;
  • 10 scripts, each with its path and the digest of its content;
  • 17 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

No dataset and no data link were found in the paper.

Code and data availability statement

The paper has a code and data availability statement. Its license (CC BY-NC-ND) does not allow reproducing it here; in short, from what the harvester recognized in it:

  • it says that the data are available on request

Read it in the paper: doi.org/10.1186/s12938-026-01555-0.

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 1, 29 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 4 authors, 7 keywords, 6 MeSH terms, 26 references.

Cite

This paper

Bawil, M. B., Shamsi, M., Jafargholkhanloo, A. F., & Bavil, A. S. (2026). Incorporating normal periventricular changes for enhanced pathological white matter hyperintensity segmentation: on multiclass deep learning approaches. Biomedical engineering online, 25(1), 69. https://doi.org/10.1186/s12938-026-01555-0

BibTeX

@article{bawil2026incorporating,
author = {Bawil, Mahdi Bashiri and Shamsi, Mousa and Jafargholkhanloo, Ali Fahmi and Bavil, Abolhassan Shakeri},
title = {{Incorporating normal periventricular changes for enhanced pathological white matter hyperintensity segmentation: on multiclass deep learning approaches}},
journal = {Biomedical engineering online},
year = {2026},
month = apr,
volume = {25},
number = {1},
pages = {69},
publisher = {BMC},
issn = {1475-925X},
doi = {10.1186/s12938-026-01555-0},
url = {https://doi.org/10.1186/s12938-026-01555-0},
pmid = {41992229},
pmcid = {PMC13202883}
}

RIS

TY - JOUR
AU - Bawil, Mahdi Bashiri
AU - Shamsi, Mousa
AU - Jafargholkhanloo, Ali Fahmi
AU - Bavil, Abolhassan Shakeri
TI - Incorporating normal periventricular changes for enhanced pathological white matter hyperintensity segmentation: on multiclass deep learning approaches
T2 - Biomedical engineering online
J2 - Biomed Eng Online
PY - 2026
DA - 2026/04/16
VL - 25
IS - 1
SP - 69
SN - 1475-925X
PB - BMC
DO - 10.1186/s12938-026-01555-0
UR - https://doi.org/10.1186/s12938-026-01555-0
LA - en
ER -

CSL-JSON

{
"id": "10.1186/s12938-026-01555-0",
"type": "article-journal",
"title": "Incorporating normal periventricular changes for enhanced pathological white matter hyperintensity segmentation: on multiclass deep learning approaches",
"container-title": "Biomedical engineering online",
"author": [
{
"family": "Bawil",
"given": "Mahdi Bashiri"
},
{
"family": "Shamsi",
"given": "Mousa"
},
{
"family": "Jafargholkhanloo",
"given": "Ali Fahmi"
},
{
"family": "Bavil",
"given": "Abolhassan Shakeri"
}
],
"container-title-short": "Biomed Eng Online",
"volume": "25",
"issue": "1",
"page": "69",
"DOI": "10.1186/s12938-026-01555-0",
"PMID": "41992229",
"PMCID": "PMC13202883",
"ISSN": "1475-925X",
"publisher": "BMC",
"URL": "https://doi.org/10.1186/s12938-026-01555-0",
"language": "en",
"issued": {
"date-parts": [
[
2026,
4,
16
]
]
}
}

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.1038/s41597-026-07184-5 [code]
A Multiple Sclerosis MRI Dataset with Tri-Mask Annotations for Lesion Segmentation.
Journal: Scientific data
In common: Keras, TensorFlow, OpenCV, 7 other tools, methods / tools, structural MRI / diffusion, 3 references, 4 authors
[2] doi:10.1186/s12880-026-02481-2 [code]
Deep learning-based neuroanatomical profiling reveals population-specific brain changes in multiple sclerosis: a large-scale Middle Eastern study.
Journal: BMC medical imaging
In common: Keras, TensorFlow, OpenCV, 7 other tools, structural MRI / diffusion, 5 references, 3 authors
[3] doi:10.1016/j.nicl.2026.104001 [code]
Effect of vascular lesion preprocessing on Brain Intensity AbNormality Classification Algorithm (BIANCA) white matter hyperintensity segmentation.
Journal: NeuroImage. Clinical
In common: seaborn, scikit-learn, pandas, 3 other tools, stroke, methods / tools, structural MRI / diffusion, 5 references
[4] doi:10.1212/wnl.0000000000218472 [code]
Lesion-Level Subtypes of White Matter Hyperintensity Evolution Beyond Spatial Location.
Journal: Neurology
In common: scikit-learn, pandas, SciPy, 2 other tools, stroke, structural MRI / diffusion, 4 references
[5] doi:10.1038/s42003-026-10957-8 [code]
Brain defence by the extracellular matrix protein Cochlin.
Journal: Communications biology
In common: Keras, TensorFlow, OpenCV, 7 other tools
[6] doi:10.1371/journal.pcbi.1014571 [code]
SynAPSeg: A novel dataset and image analysis framework for deep learning-based synapse detection and quantification.
Journal: PLoS computational biology
In common: Keras, TensorFlow, OpenCV, 7 other tools
[7] doi:10.1016/j.xcrm.2026.102766 [code]
A longitudinal single-cell and spatial multiomic atlas of pediatric high-grade glioma.
Journal: Cell reports. Medicine
In common: Keras, TensorFlow, OpenCV, 7 other tools
[8] doi: [code]
Real-time closed-loop feedback system for mouse mesoscale cortical signal and movement control
Journal: eLife
In common: Keras, TensorFlow, OpenCV, 7 other tools
[9] doi:10.1364/boe.605322 [code]
Generalized plaque digitization framework for multi-dimensional mesoscopic images.
Journal: Biomedical optics express
In common: Keras, TensorFlow, OpenCV, 7 other tools
[10] doi:10.1038/s41467-026-72057-9 [code]
Sex-specific behavioral feedback modulates sensorimotor processing and drives flexible social behavior.
Journal: Nature communications
In common: Keras, TensorFlow, OpenCV, 7 other tools

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.