OSCR

Distilling population specific expertise into a unified model for generalizable brain tumor segmentation.

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 · 5 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 › Data preprocessing ↔ AI/data/make_dataset.py, lines 61–98 · score 0.86 · depth axis, aspect ratio, 0.65–1.5, flipping, cropped, foreground
  2. [2] § Materials and methods › Proposed MTSS-KDNet › Loss function › Logits level distillation loss ↔ AI/kd_modules/framework.py, the whole file · a weak match · score 0.85 · Binary Cross Entropy, teacher logits, distillation loss, BCE loss, Deep supervised, softens
  3. [3] § Materials and methods › Proposed MTSS-KDNet › Loss function › Supervised student loss ↔ AI/losses/loss.py, the whole file · a weak match · score 0.84 · Binary Cross Entropy, BCE loss, ground truth, Dice loss, segmentation loss, sum
  4. [4] § Materials and methods › Proposed MTSS-KDNet › Loss function › Supervised student loss ↔ AI/kd_modules/framework.py, the whole file · a weak match · score 0.84 · Binary Cross Entropy, ground truth, BCE loss, segmentation loss, deep supervision, component
  5. [5] § Materials and methods › BraTS datasets ↔ AI/losses/loss.py, the whole file · a weak match · score 0.81 · peritumoral edema, Tumor Core, Ground truth, enhancing tumor, background, TC
  6. [6] § Materials and methods › Data preprocessing ↔ AI/notebooks/Ablation/ABL_(BCE+SEG).ipynb, lines 33–180 · score 0.80 · aspect ratio, 0.65–1.5, axis, cropped, foreground, gamma
  7. [7] § Experiments and setup › Implementation details ↔ AI/notebooks/Ablation/ABL_(BCE+SEG).ipynb, lines 868–871 · score 0.79 · ReduceLROnPlateau, AdamW, weight decay, cooldown, patience, plateaus
  8. [8] § Experiments and setup › Implementation details ↔ AI/notebooks/Ablation/ABL_(KL+SEG).ipynb, lines 896–899 · score 0.79 · ReduceLROnPlateau, AdamW, weight decay, cooldown, patience, plateaus
  9. [9] § Materials and methods › Proposed MTSS-KDNet › Loss function › Latent space distillation loss ↔ AI/kd_modules/cbam_attention.py, lines 4–55 · score 0.72 · Convolutional Block Attention, feature map, emphasized, CBAM, network, Module
  10. [10] § Materials and methods › Inference ↔ AI/inference/postprocess.py, lines 5–29 · score 0.72 · sliding window inference, segmentation mask, tumor regions, prediction, model
  11. [11] § Results and discussion › Teachers progression evaluation ↔ AI/config/default_config.py, the whole file · a weak match · score 0.70 · SSA teacher, PED teacher, MEN teacher, MET teacher, teacher models, GLI
  12. [12] § Materials and methods › Inference ↔ AI/inference/predict.py, lines 37–102 · score 0.67 · sliding window inference, Gaussian, overlapping, volumes, prediction, mask
  13. [13] § Materials and methods › BraTS datasets ↔ AI/inference/preprocess.py, lines 8–49 · score 0.65 · T1 weighted, T2 weighted, FLAIR, preprocessed, volume
  14. [14] § Experiments and setup › Backbone architecture ↔ AI/models/dyn_unet.py, lines 6–20 · score 0.65 · DynUNet, medical imaging, Deep supervised, knowledge distillation, Dynamic, architecture
  15. [15] § Experiments and setup › Backbone architecture ↔ AI/models/blocks.py, lines 82–156 · score 0.63 · skip connections, medical imaging, resolutions, encoder, UNet, decoder
  16. [16] § Materials and methods › Proposed MTSS-KDNet › Teacher models ↔ AI/notebooks/Ablation/ABL_(BCE+SEG).ipynb, lines 868–871 · score 0.62 · AdamW, weight decay, Plateau, scheduler, optimized, model
  17. [17] § Results and discussion › Ablation study › Visualizing the effect of CBAM module ↔ AI/notebooks/Ablation/ABL_(KL_no_CBAM+SEG).ipynb, lines 725–784 · score 0.56 · bottleneck layer, bottleneck features, CBAM, decoder, map, ablation

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 · 1,068 lines · 41 KB · no license · 3 matches

  1. # %% [markdown]
  2. # # Importing & Loading Dependencies
  3. # %%
  4. !pip install monai
  5. import nibabel as nib
  6. from monai.transforms import LoadImage, Compose, NormalizeIntensityd, RandFlipd, RandAdjustContrastd, Resized, CropForegroundd, SpatialPadd
  7. import matplotlib.pyplot as plt
  8. from torch.utils.data import DataLoader, Dataset
  9. import torch
  10. import torch.nn as nn
  11. import torch.nn.functional as F
  12. import numpy as np
  13. from typing import Optional, Sequence, Tuple, Union
  14. from torch.nn.functional import interpolate
  15. from monai.networks.blocks.convolutions import Convolution
  16. from monai.networks.layers.factories import Act, Norm
  17. from monai.networks.layers.utils import get_act_layer, get_norm_layer
  18. from monai.metrics import DiceMetric, HausdorffDistanceMetric
  19. from torch import nn, optim, amp
  20. from itertools import chain
  21. from monai.losses import DiceLoss
  22. from tqdm import tqdm
  23. from pathlib import Path
  24. import math
  25. import os
  26. import random
  27. # %% [markdown]
  28. # # Creating Dataset with Preprocessing
  29. # %%
  30. class CustomDataset3D(Dataset):
  31. def __init__(self, data_dirs, patient_lists, mode):
  32. self.data_dirs = data_dirs
  33. self.patient_lists = patient_lists
  34. self.mode = mode
  35. @staticmethod
  36. def resize_with_aspect_ratio(keys, target_size):
  37. def transform(data):
  38. for key in keys:
  39. volume = data[key]
  40. original_shape = volume.shape[-3:]
  41. scaling_factor = min(
  42. target_size[0] / original_shape[0],
  43. target_size[1] / original_shape[1],
  44. target_size[2] / original_shape[2]
  45. )
  46. # Computing the intermediate size while preserving aspect ratio
  47. new_shape = [
  48. int(dim * scaling_factor) for dim in original_shape
  49. ]
  50. # Resizing to the intermediate shape
  51. resize_transform = Resized(keys=[key], spatial_size=new_shape, mode="trilinear" if key == "imgs" else "nearest-exact")
  52. data = resize_transform(data)
  53. # Padding to the final target size
  54. pad_transform = SpatialPadd(keys=[key], spatial_size=target_size, mode="constant")
  55. data = pad_transform(data)
  56. return data
  57. return transform
  58. def preprocess(cls, data, mode):
  59. if mode == 'training':
  60. transform = Compose([
  61. CropForegroundd(keys=["imgs", "masks"], source_key="imgs"),
  62. cls.resize_with_aspect_ratio(keys=["imgs", "masks"], target_size=[128, 128, 128]),
  63. NormalizeIntensityd( keys=['imgs'], nonzero=False, channel_wise=True),
  64. RandFlipd(keys=["imgs", "masks"],
  65. prob=0.5,
  66. spatial_axis=2,
  67. ),
  68. RandAdjustContrastd(
  69. keys=["imgs"],
  70. prob=0.15,
  71. gamma=(0.65, 1.5),
  72. ),
  73. ])
  74. elif mode == 'validation':
  75. transform = Compose([
  76. CropForegroundd(keys=["imgs", "masks"], source_key="imgs"),
  77. cls.resize_with_aspect_ratio(keys=["imgs", "masks"], target_size=[128, 128, 128]),
  78. NormalizeIntensityd( keys=['imgs'], nonzero=False, channel_wise=True)
  79. ])
  80. else: # 'testing'
  81. transform = Compose([
  82. CropForegroundd(keys=["imgs", "masks"], source_key="imgs"),
  83. cls.resize_with_aspect_ratio(keys=["imgs", "masks"], target_size=[128, 128, 128]),
  84. NormalizeIntensityd( keys=['imgs'], nonzero=False, channel_wise=True)
  85. ])
  86. augmented_data = transform(data)
  87. return augmented_data
  88. def __len__(self):
  89. return len(self.patient_lists)
  90. def __getitem__(self, idx):
  91. patient_id = self.patient_lists[idx]
  92. loadimage = LoadImage(reader='NibabelReader', image_only=True)
  93. data_type=patient_id.split('-')[1]
  94. if data_type == 'GLI':
  95. patient_folder_path = os.path.join('/kaggle/input/bratsglioma/Training', patient_id)
  96. elif data_type == 'SSA':
  97. patient_folder_path = os.path.join('/kaggle/input/bratsafrica24', patient_id)
  98. elif data_type == 'PED':
  99. patient_folder_path = os.path.join('/kaggle/input/bratsped/Training', patient_id)
  100. elif data_type == 'MEN':
  101. patient_folder_path = os.path.join('/kaggle/input/bratsmen', patient_id)
  102. else:
  103. patient_folder_path = os.path.join('/kaggle/input/bratsmet24', patient_id)
  104. def resolve_file_path(folder, name):
  105. file_path = os.path.join(folder, name)
  106. # Check if the given path is a directory (case with 4 subdirs)
  107. if os.path.isdir(file_path):
  108. # Find the first file inside the directory that ends with .nii
  109. for root, _, files in os.walk(file_path):
  110. for file in files:
  111. if file.endswith(".nii"):
  112. return os.path.join(root, file)
  113. return file_path
  114. # Resolve paths for all required image types
  115. t1c_path = resolve_file_path(patient_folder_path, patient_id + '-t1c.nii')
  116. t1n_path = resolve_file_path(patient_folder_path, patient_id + '-t1n.nii')
  117. t2f_path = resolve_file_path(patient_folder_path, patient_id + '-t2f.nii')
  118. t2w_path = resolve_file_path(patient_folder_path, patient_id + '-t2w.nii')
  119. seg_path = os.path.join(patient_folder_path, patient_id + '-seg.nii')
  120. t1c_loader = loadimage( t1c_path )
  121. t1n_loader = loadimage( t1n_path )
  122. t2f_loader = loadimage( t2f_path )
  123. t2w_loader = loadimage( t2w_path )
  124. masks_loader = loadimage( seg_path )
  125. # Make the dimension of channel
  126. t1c_tensor = torch.Tensor(t1c_loader).unsqueeze(0)
  127. t1n_tensor = torch.Tensor(t1n_loader).unsqueeze(0)
  128. t2f_tensor = torch.Tensor(t2f_loader).unsqueeze(0)
  129. t2w_tensor = torch.Tensor(t2w_loader).unsqueeze(0)
  130. masks_tensor = torch.Tensor(masks_loader).unsqueeze(0)
  131. concat_tensor = torch.cat( (t1c_tensor, t1n_tensor, t2f_tensor, t2w_tensor, masks_tensor), 0 )
  132. data = {
  133. 'imgs' : np.array(concat_tensor[0:4,:,:,:]),
  134. 'masks' : np.array(concat_tensor[4:,:,:,:])
  135. }
  136. augmented_imgs_masks = self.preprocess(data, self.mode)
  137. imgs = np.array(augmented_imgs_masks['imgs'])
  138. masks = np.array(augmented_imgs_masks['masks'])
  139. y = {
  140. 'imgs' : torch.from_numpy(imgs).type(torch.FloatTensor),
  141. 'masks' : torch.from_numpy(masks).type(torch.FloatTensor),
  142. 'patient_id' : patient_id,
  143. 'data_type' : data_type
  144. }
  145. return y
  146. # %% [markdown]
  147. # # Data Loaders
  148. # %%
  149. def combine_datasets(dataset_lists, batch_size=3):
  150. max_len = max(len(dataset) for dataset in dataset_lists)
  151. # Ensure batch_size matches the number of datasets
  152. if batch_size != len(dataset_lists):
  153. raise ValueError("Batch size must equal the number of datasets for this function.")
  154. combined_paths = []
  155. for i in range(0, max_len, batch_size):
  156. for j in range(batch_size):
  157. index = (i + j) % max_len
  158. batch = [dataset[index % len(dataset)] for dataset in dataset_lists]
  159. combined_paths.extend(batch)
  160. # if j == 0:
  161. # print(f"Batch {(i // batch_size) + 1}: {batch}")
  162. return combined_paths
  163. # %%
  164. def prepare_data_loaders(args):
  165. train_datasets, val_datasets, test_datasets = [], [], []
  166. split_ratio = {'training': 0.71, 'validation': 0.09, 'testing': 0.2}
  167. for i, data_dir in enumerate(args['data_dirs']):
  168. patient_lists = os.listdir( data_dir )
  169. patient_lists.sort()
  170. total_patients = len(patient_lists)
  171. random.seed(5)
  172. random.shuffle(patient_lists)
  173. train_split = int(split_ratio['training'] * total_patients)
  174. val_split = int(split_ratio['validation'] * total_patients)
  175. train_patient_lists = patient_lists[:train_split]
  176. val_patient_lists = patient_lists[train_split : train_split + val_split]
  177. test_patient_lists = patient_lists[train_split + val_split :]
  178. train_patient_lists.sort()
  179. val_patient_lists.sort()
  180. test_patient_lists.sort()
  181. print(f'Number of training samples in {data_dir.split("/")[3]} DataSet: {len(train_patient_lists)}')
  182. print(f'Number of validation samples in {data_dir.split("/")[3]} DataSet: {len(val_patient_lists)}')
  183. print(f'Number of testing samples in {data_dir.split("/")[3]} DataSet: {len(test_patient_lists)} ')
  184. train_datasets.append(train_patient_lists)
  185. val_datasets.append(val_patient_lists)
  186. test_datasets.append(test_patient_lists)
  187. combined_trainDataset = combine_datasets(train_datasets, batch_size=args['train_batch_size'])
  188. combined_valDataset = list(chain.from_iterable(val_datasets))
  189. combined_testDataset = list(chain.from_iterable(test_datasets))
  190. print(f'Number of combined training samples', len(combined_trainDataset))
  191. print(f'Number of combined validation samples', len(combined_valDataset))
  192. print(f'Number of combined testing samples', len(combined_testDataset))
  193. trainDataset = CustomDataset3D( args['data_dirs'], combined_trainDataset, mode='training')
  194. valDataset = CustomDataset3D( args['data_dirs'], combined_valDataset, mode='validation')
  195. testDataset = CustomDataset3D( args['data_dirs'], combined_testDataset, mode='testing')
  196. trainLoader = DataLoader(
  197. trainDataset, batch_size=args['train_batch_size'], num_workers=args['workers'], prefetch_factor=2,
  198. pin_memory=True, shuffle=False)
  199. valLoader = DataLoader(
  200. valDataset, batch_size=args['val_batch_size'], num_workers=args['workers'], prefetch_factor=2,
  201. pin_memory=True, shuffle=False)
  202. testLoader = DataLoader(
  203. testDataset, batch_size=args['test_batch_size'], num_workers=args['workers'], prefetch_factor=2,
  204. pin_memory=True, shuffle=False)
  205. return trainLoader, valLoader, testLoader
  206. # %% [markdown]
  207. # # Visualizing Data
  208. # %%
  209. # args = {
  210. # 'workers': 2,
  211. # 'epochs': 10,
  212. # 'train_batch_size': 2,
  213. # 'val_batch_size': 2,
  214. # 'test_batch_size': 2,
  215. # 'learning_rate': 1e-3,
  216. # 'weight_decay': 1e-5,
  217. # 'lambd': 0.0051,
  218. # 'data_dir': '/kaggle/input/bratsafrica24/',
  219. # 'in_checkpoint_dir': Path('/kaggle/input/adultgliomamodel-45epochs'),
  220. # 'out_checkpoint_dir': Path('/kaggle/working/')
  221. # }
  222. # trainLoader, valLoader, testLoader = prepare_data_loaders(args)
  223. # for step, y in enumerate( trainLoader ):
  224. # print(y['imgs'].shape)
  225. # print(y['patient_id'])
  226. # fig, axes = plt.subplots(1, 4, figsize=(16, 4))
  227. # for sequence in range(4):
  228. # sequence_data = y['imgs'][0][sequence, :, :, :].cpu().detach().numpy()
  229. # slice_index = sequence_data.shape[2] // 2
  230. # axes[sequence].imshow(np.rot90(sequence_data[:, :, slice_index]), cmap='gray', origin='lower')
  231. # axes[sequence].set_title(f'Sequence {sequence + 1}')
  232. # plt.show()
  233. # %% [markdown]
  234. # # DynUNet Model
  235. # %%
  236. class UnetBasicBlock(nn.Module):
  237. """
  238. A CNN module module that can be used for DynUNet, based on:
  239. `Automated Design of Deep Learning Methods for Biomedical Image Segmentation <https://arxiv.org/abs/1904.08128>`_.
  240. `nnU-Net: Self-adapting Framework for U-Net-Based Medical Image Segmentation <https://arxiv.org/abs/1809.10486>`_.
  241. Args:
  242. spatial_dims: number of spatial dimensions.
  243. in_channels: number of input channels.
  244. out_channels: number of output channels.
  245. kernel_size: convolution kernel size.
  246. stride: convolution stride.
  247. norm_name: feature normalization type and arguments.
  248. act_name: activation layer type and arguments.
  249. dropout: dropout probability.
  250. """
  251. def __init__(
  252. self,
  253. spatial_dims: int,
  254. in_channels: int,
  255. out_channels: int,
  256. kernel_size: Union[Sequence[int], int],
  257. stride: Union[Sequence[int], int],
  258. norm_name: Union[Tuple, str] = ("INSTANCE", {"affine": True}),
  259. act_name: Union[Tuple, str] = ("leakyrelu", {"inplace": True, "negative_slope": 0.01}),
  260. dropout: Optional[Union[Tuple, str, float]] = None,
  261. ):
  262. super().__init__()
  263. self.conv1 = get_conv_layer(
  264. spatial_dims,
  265. in_channels,
  266. out_channels,
  267. kernel_size=kernel_size,
  268. stride=stride,
  269. dropout=dropout,
  270. conv_only=True,
  271. )
  272. self.conv2 = get_conv_layer(
  273. spatial_dims,
  274. out_channels,
  275. out_channels,
  276. kernel_size=kernel_size,
  277. stride=1,
  278. dropout=dropout,
  279. conv_only=True
  280. )
  281. self.lrelu = get_act_layer(name=act_name)
  282. self.norm1 = get_norm_layer(name=norm_name, spatial_dims=spatial_dims, channels=out_channels)
  283. self.norm2 = get_norm_layer(name=norm_name, spatial_dims=spatial_dims, channels=out_channels)
  284. def forward(self, inp):
  285. out = self.conv1(inp)
  286. out = self.norm1(out)
  287. out = self.lrelu(out)
  288. out = self.conv2(out)
  289. out = self.norm2(out)
  290. out = self.lrelu(out)
  291. return out
  292. class UnetUpBlock(nn.Module):
  293. """
  294. An upsampling module that can be used for DynUNet, based on:
  295. `Automated Design of Deep Learning Methods for Biomedical Image Segmentation <https://arxiv.org/abs/1904.08128>`_.
  296. `nnU-Net: Self-adapting Framework for U-Net-Based Medical Image Segmentation <https://arxiv.org/abs/1809.10486>`_.
  297. Args:
  298. spatial_dims: number of spatial dimensions.
  299. in_channels: number of input channels.
  300. out_channels: number of output channels.
  301. kernel_size: convolution kernel size.
  302. stride: convolution stride.
  303. upsample_kernel_size: convolution kernel size for transposed convolution layers.
  304. norm_name: feature normalization type and arguments.
  305. act_name: activation layer type and arguments.
  306. dropout: dropout probability.
  307. trans_bias: transposed convolution bias.
  308. """
  309. def __init__(
  310. self,
  311. spatial_dims: int,
  312. in_channels: int,
  313. out_channels: int,
  314. kernel_size: Union[Sequence[int], int],
  315. upsample_kernel_size: Union[Sequence[int], int],
  316. norm_name: Union[Tuple, str] = ("INSTANCE", {"affine": True}),
  317. act_name: Union[Tuple, str] = ("leakyrelu", {"inplace": True, "negative_slope": 0.01}),
  318. dropout: Optional[Union[Tuple, str, float]] = None,
  319. trans_bias: bool = False,
  320. ):
  321. super().__init__()
  322. upsample_stride = upsample_kernel_size
  323. # ( a purple arrow in the paper )
  324. self.transp_conv = get_conv_layer(
  325. spatial_dims,
  326. in_channels,
  327. out_channels,
  328. kernel_size=upsample_kernel_size,
  329. stride=upsample_stride,
  330. dropout=dropout,
  331. bias=trans_bias,
  332. conv_only=True,
  333. is_transposed=True,
  334. )
  335. # A light blue conv blocks in the decoder of nnUNet
  336. self.conv_block = UnetBasicBlock(
  337. spatial_dims,
  338. out_channels + out_channels,
  339. out_channels,
  340. kernel_size=kernel_size,
  341. stride=1,
  342. dropout=dropout,
  343. norm_name=norm_name,
  344. act_name=act_name,
  345. )
  346. def forward(self, inp, skip):
  347. # number of channels for skip should equals to out_channels
  348. out = self.transp_conv(inp)
  349. out = torch.cat((out, skip), dim=1)
  350. out = self.conv_block(out)
  351. return out
  352. class UnetOutBlock(nn.Module):
  353. def __init__(
  354. self, spatial_dims: int, in_channels: int, out_channels: int, dropout: Optional[Union[Tuple, str, float]] = None
  355. ):
  356. super().__init__()
  357. self.conv = get_conv_layer(
  358. spatial_dims, in_channels, out_channels, kernel_size=1, stride=1, dropout=dropout, bias=True, conv_only=True
  359. )
  360. def forward(self, inp):
  361. return self.conv(inp)
  362. def get_conv_layer(
  363. spatial_dims: int,
  364. in_channels: int,
  365. out_channels: int,
  366. kernel_size: Union[Sequence[int], int] = 3,
  367. stride: Union[Sequence[int], int] = 1,
  368. act: Optional[Union[Tuple, str]] = Act.PRELU,
  369. norm: Union[Tuple, str] = Norm.INSTANCE,
  370. dropout: Optional[Union[Tuple, str, float]] = None,
  371. bias: bool = False,
  372. conv_only: bool = True,
  373. is_transposed: bool = False,
  374. ):
  375. padding = get_padding(kernel_size, stride)
  376. output_padding = None
  377. if is_transposed:
  378. output_padding = get_output_padding(kernel_size, stride, padding)
  379. return Convolution(
  380. spatial_dims,
  381. in_channels,
  382. out_channels,
  383. strides=stride,
  384. kernel_size=kernel_size,
  385. act=act,
  386. norm=norm,
  387. dropout=dropout,
  388. bias=bias,
  389. conv_only=conv_only,
  390. is_transposed=is_transposed,
  391. padding=padding,
  392. output_padding=output_padding,
  393. )
  394. def get_padding(
  395. kernel_size: Union[Sequence[int], int], stride: Union[Sequence[int], int]
  396. ) -> Union[Tuple[int, ...], int]:
  397. kernel_size_np = np.atleast_1d(kernel_size)
  398. stride_np = np.atleast_1d(stride)
  399. padding_np = (kernel_size_np - stride_np + 1) / 2
  400. if np.min(padding_np) < 0:
  401. raise AssertionError("padding value should not be negative, please change the kernel size and/or stride.")
  402. padding = tuple(int(p) for p in padding_np)
  403. return padding if len(padding) > 1 else padding[0]
  404. def get_output_padding(
  405. kernel_size: Union[Sequence[int], int], stride: Union[Sequence[int], int], padding: Union[Sequence[int], int]
  406. ) -> Union[Tuple[int, ...], int]:
  407. kernel_size_np = np.atleast_1d(kernel_size)
  408. stride_np = np.atleast_1d(stride)
  409. padding_np = np.atleast_1d(padding)
  410. out_padding_np = 2 * padding_np + stride_np - kernel_size_np
  411. if np.min(out_padding_np) < 0:
  412. raise AssertionError("out_padding value should not be negative, please change the kernel size and/or stride.")
  413. out_padding = tuple(int(p) for p in out_padding_np)
  414. return out_padding if len(out_padding) > 1 else out_padding[0]
  415. def set_requires_grad(nets, requires_grad=False):
  416. if not isinstance(nets, list):
  417. nets = [nets]
  418. for net in nets:
  419. if net is not None:
  420. for param in net.parameters():
  421. param.requires_grad = requires_grad
  422. # %%
  423. class DynUNet(nn.Module):
  424. def __init__(
  425. self,
  426. spatial_dims: int,
  427. in_channels: int,
  428. out_channels: int,
  429. deep_supervision: bool,
  430. KD: bool = False
  431. ):
  432. super().__init__()
  433. self.spatial_dims = spatial_dims
  434. self.in_channels = in_channels
  435. self.out_channels = out_channels
  436. self.deep_supervision = deep_supervision
  437. self.KD_enabled = KD
  438. self.input_conv = UnetBasicBlock( spatial_dims=self.spatial_dims,
  439. in_channels=self.in_channels,
  440. out_channels=64,
  441. kernel_size=3,
  442. stride=1
  443. )
  444. self.down1 = UnetBasicBlock( spatial_dims=self.spatial_dims,
  445. in_channels=64,
  446. out_channels=96,
  447. kernel_size=3,
  448. stride=2 # Reduces spatial dims by 2
  449. )
  450. self.down2 = UnetBasicBlock( spatial_dims=self.spatial_dims,
  451. in_channels=96,
  452. out_channels=128,
  453. kernel_size=3,
  454. stride=2
  455. )
  456. self.down3 = UnetBasicBlock( spatial_dims=self.spatial_dims,
  457. in_channels=128,
  458. out_channels=192,
  459. kernel_size=3,
  460. stride=2
  461. )
  462. self.down4 = UnetBasicBlock( spatial_dims=self.spatial_dims,
  463. in_channels=192,
  464. out_channels=256,
  465. kernel_size=3,
  466. stride=2
  467. )
  468. self.down5 = UnetBasicBlock( spatial_dims=self.spatial_dims,
  469. in_channels=256,
  470. out_channels=384,
  471. kernel_size=3,
  472. stride=2
  473. )
  474. self.bottleneck = UnetBasicBlock( spatial_dims=self.spatial_dims,
  475. in_channels=384,
  476. out_channels=512,
  477. kernel_size=3,
  478. stride=2
  479. )
  480. self.up1 = UnetUpBlock( spatial_dims=self.spatial_dims,
  481. in_channels=512,
  482. out_channels=384,
  483. kernel_size=3,
  484. upsample_kernel_size=2
  485. )
  486. self.up2 = UnetUpBlock( spatial_dims=self.spatial_dims,
  487. in_channels=384,
  488. out_channels=256,
  489. kernel_size=3,
  490. upsample_kernel_size=2
  491. )
  492. self.up3 = UnetUpBlock( spatial_dims=self.spatial_dims,
  493. in_channels=256,
  494. out_channels=192,
  495. kernel_size=3,
  496. upsample_kernel_size=2
  497. )
  498. self.up4 = UnetUpBlock( spatial_dims=self.spatial_dims,
  499. in_channels=192,
  500. out_channels=128,
  501. kernel_size=3,
  502. upsample_kernel_size=2
  503. )
  504. self.up5 = UnetUpBlock( spatial_dims=self.spatial_dims,
  505. in_channels=128,
  506. out_channels=96,
  507. kernel_size=3,
  508. upsample_kernel_size=2
  509. )
  510. self.up6 = UnetUpBlock( spatial_dims=self.spatial_dims,
  511. in_channels=96,
  512. out_channels=64,
  513. kernel_size=3,
  514. upsample_kernel_size=2
  515. )
  516. self.out1 = UnetOutBlock( spatial_dims=self.spatial_dims,
  517. in_channels=64,
  518. out_channels=self.out_channels,
  519. )
  520. self.out2 = UnetOutBlock( spatial_dims=self.spatial_dims,
  521. in_channels=96,
  522. out_channels=self.out_channels,
  523. )
  524. self.out3 = UnetOutBlock( spatial_dims=self.spatial_dims,
  525. in_channels=128,
  526. out_channels=self.out_channels,
  527. )
  528. def forward( self, input ):
  529. # Input
  530. x0 = self.input_conv( input ) # x0.shape = (B x 64 x 128 x 128 x 128)
  531. # Encoder
  532. x1 = self.down1( x0 ) # x1.shape = (B x 96 x 64 x 64 x 64)
  533. x2 = self.down2( x1 ) # x2.shape = (B x 128 x 32 x 32 x 32)
  534. x3 = self.down3( x2 ) # x3.shape = (B x 192 x 16 x 16 x 16)
  535. x4 = self.down4( x3 ) # x4.shape = (B x 256 x 8 x 8 x 8)
  536. x5 = self.down5( x4 ) # x5.shape = (B x 384 x 4 x 4 x 4)
  537. # Bottleneck
  538. x6 = self.bottleneck( x5 ) # x6.shape = (B x 512 x 2 x 2 x 2)
  539. # Decoder
  540. x7 = self.up1( x6, x5 ) # x7.shape = (B x 384 x 4 x 4 x 4)
  541. x8 = self.up2( x7, x4 ) # x8.shape = (B x 256 x 8 x 8 x 8)
  542. x9 = self.up3( x8, x3 ) # x9.shape = (B x 192 x 16 x 16 x 16)
  543. x10 = self.up4( x9, x2 ) # x10.shape = (B x 128 x 32 x 32 x 32)
  544. x11 = self.up5( x10, x1 ) # x11.shape = (B x 96 x 64 x 64 x 64)
  545. x12 = self.up6( x11, x0 ) # x12.shape = (B x 64 x 128 x 128 x 128)
  546. # Output
  547. output1 = self.out1( x12 )
  548. if (self.training and self.deep_supervision) or self.KD_enabled:
  549. # output['pred'].shape = B x 3 x 4 x 128 x 128 x 128
  550. output2 = interpolate( self.out2( x11 ), output1.shape[2:])
  551. output3 = interpolate( self.out3( x10 ), output1.shape[2:])
  552. output_all = [ output1, output2, output3 ]
  553. return { 'pred' : torch.stack(output_all, dim=1),
  554. 'bottleneck_feature_map' : x6 }
  555. return { 'pred' : output1 }
  556. # %% [markdown]
  557. # # Visualizing Model Instance
  558. # %%
  559. # !pip install torchsummary
  560. # from torchsummary import summary
  561. # # Initialize your DynUNet model
  562. # model = DynUNet(spatial_dims=3, in_channels=4, out_channels=4, deep_supervision=True, KD=True)
  563. # # Print model summary
  564. # summary(model, input_size=(4, 128, 128, 128)) # Adjust input_size according to your needs
  565. # %% [markdown]
  566. # # ClearML
  567. # %%
  568. !pip install clearml
  569. from clearml import Task
  570. %env CLEARML_WEB_HOST=https://app.clear.ml/
  571. %env CLEARML_API_HOST=https://api.clear.ml
  572. %env CLEARML_FILES_HOST=https://files.clear.ml
  573. %env CLEARML_API_ACCESS_KEY=CLEARML_API_ACCESS_KEY
  574. %env CLEARML_API_SECRET_KEY=CLEARML_API_SECRET_KEY
  575. # %% [markdown]
  576. # # GPUs Check
  577. # %%
  578. if torch.cuda.is_available():
  579. num_gpus = torch.cuda.device_count()
  580. print(f"Number of GPUs available: {num_gpus}")
  581. for i in range(num_gpus):
  582. print(f"GPU {i}: {torch.cuda.get_device_name(i)}")
  583. else:
  584. print("No GPU available. Running on CPU.")
  585. # %%
  586. # # For freeing gpu
  587. # import gc; gc.collect(); torch.cuda.empty_cache()
  588. # %% [markdown]
  589. # # Loss Function
  590. # %%
  591. class LossFunction(nn.Module):
  592. def __init__(self):
  593. super(LossFunction, self).__init__()
  594. self.dice = DiceLoss(sigmoid=True, batch=True, smooth_nr=1e-05, smooth_dr=1e-05)
  595. self.ce = nn.BCEWithLogitsLoss()
  596. def _loss(self, p, y):
  597. return self.dice(p, y) + self.ce(p, y.float())
  598. def forward(self, p, y):
  599. y_wt, y_tc, y_et = y > 0, ((y == 1) + (y == 3)) > 0, y == 3
  600. p_wt, p_tc, p_et = p[:, 1].unsqueeze(1), p[:, 2].unsqueeze(1), p[:, 3].unsqueeze(1)
  601. l_wt, l_tc, l_et = self._loss(p_wt, y_wt), self._loss(p_tc, y_tc), self._loss(p_et, y_et)
  602. return l_wt + l_tc + l_et
  603. # %% [markdown]
  604. # # Student KD Model
  605. # %%
  606. class Student_KD_loss(nn.Module):
  607. def __init__(self):
  608. super().__init__()
  609. self.student = DynUNet( spatial_dims=3, in_channels=4, out_channels=4, deep_supervision=True)
  610. self.loss_fn = LossFunction()
  611. self.temperature = 5.0
  612. self.bce_loss = nn.BCEWithLogitsLoss()
  613. def forward(self, teacher_outputs, y):
  614. with amp.autocast('cuda:1'):
  615. student_outputs = self.student( y['imgs'] )
  616. # Student loss with Deep supervision -> (Dice loss)
  617. segloss_s_decoder_1 = self.loss_fn( student_outputs['pred'][:,0], y['masks'] ) # student_outputs['pred'].shape = B x 3 x 4 x 128 x 128 x 128
  618. segloss_s_decoder_2 = self.loss_fn( student_outputs['pred'][:,1], y['masks'] )
  619. segloss_s_decoder_3 = self.loss_fn( student_outputs['pred'][:,2], y['masks'] )
  620. student_seg_loss = segloss_s_decoder_1 + 0.5*segloss_s_decoder_2 + 0.25*segloss_s_decoder_3
  621. # KD loss between prediction layers -> (BCE Loss on logits with DS)
  622. bce_loss_with_teacher = 0
  623. decoder_weights = [1, 0.5, 0.25]
  624. for decoder_idx, weight in enumerate(decoder_weights):
  625. teacher_logits = teacher_outputs['pred'][:, decoder_idx] / self.temperature # Shape: (B, 4, 128, 128, 128)
  626. student_logits = student_outputs['pred'][:, decoder_idx] / self.temperature # Shape: (B, 4, 128, 128, 128)
  627. # Compute hard labels for WT, TC, ET
  628. teacher_probs = torch.sigmoid(teacher_logits)
  629. teacher_hard_labels = [
  630. (teacher_probs[:, channel_idx] > 0.5).float().unsqueeze(1) for channel_idx in range(1, 4)
  631. ]
  632. # Extract student logits for WT, TC, ET
  633. student_logits_channels = [
  634. student_logits[:, channel_idx].unsqueeze(1) for channel_idx in range(1, 4)
  635. ]
  636. # Compute BCE losses for WT, TC, ET
  637. bce_losses = [
  638. self.bce_loss(student_logits_channels[i], teacher_hard_labels[i]) for i in range(3)
  639. ]
  640. # Aggregate loss for this decoder with the corresponding weight
  641. bce_loss_with_teacher += sum(bce_losses) * (self.temperature ** 2) * weight
  642. #-----------------------------------------------------------------------------------#
  643. alpha, zeta = 1.0, 0.1
  644. print("Seg loss: ", student_seg_loss)
  645. print("BCE loss with teacher: ", bce_loss_with_teacher)
  646. print("Seg loss weighted: ", alpha*student_seg_loss)
  647. print("BCE loss with teacher weighted: ", zeta*bce_loss_with_teacher)
  648. batch_total_student_loss = alpha*student_seg_loss + zeta*bce_loss_with_teacher
  649. print("-------------Final student loss-------------")
  650. print(batch_total_student_loss)
  651. print("-------------Final student loss-------------")
  652. KD_output = {
  653. 'batch_total_student_loss' : batch_total_student_loss,
  654. 'seg_weighted' : alpha*student_seg_loss,
  655. 'bce_weighted' : zeta*bce_loss_with_teacher,
  656. }
  657. return KD_output
  658. # %% [markdown]
  659. # # Training & Validation
  660. # %%
  661. def evaluate(model, loader, epoch, task):
  662. torch.manual_seed(0)
  663. model.eval()
  664. loss_fn = LossFunction()
  665. n_val_batches = len(loader)
  666. tumors_val_losses, running_loss = validate_model(model, loader, loss_fn)
  667. epoch_val_loss = running_loss / n_val_batches
  668. log_val_epoch_losses(tumors_val_losses, epoch, task, epoch_val_loss)
  669. print(f"------Final validation dice loss after epoch {epoch + 1}: {epoch_val_loss}-------")
  670. model.student.to('cuda:1')
  671. model.train()
  672. return epoch_val_loss
  673. def validate_model(model, loader, loss_fn):
  674. tumors_val_losses = {'GLI': [], 'PED': [], 'SSA': [], 'MEN':[], 'MET':[]}
  675. running_loss = 0
  676. n_val_batches = len(loader)
  677. with tqdm(total=n_val_batches, desc='Validating', unit='batch', leave=False) as pbar:
  678. with torch.no_grad():
  679. for y in loader:
  680. val_loss, data_type = process_batch(model, y, loss_fn)
  681. tumors_val_losses[data_type].append(val_loss.item())
  682. running_loss += val_loss
  683. pbar.update(1)
  684. return tumors_val_losses, running_loss
  685. def process_batch(model, y, loss_fn):
  686. y['imgs'], y['masks'] = y['imgs'].to('cuda'), y['masks'].to('cuda')
  687. data_type = y['data_type'][0]
  688. with torch.amp.autocast('cuda'):
  689. output = model.student.to('cuda')(y['imgs'])
  690. val_loss = loss_fn(output['pred'], y['masks'])
  691. print(f"Validation dice loss per batch: {val_loss}")
  692. return val_loss, data_type
  693. def log_val_epoch_losses(tumors_val_losses, epoch, task, epoch_val_loss):
  694. for tumor_type, losses in tumors_val_losses.items():
  695. avg_loss = sum(losses) / len(losses) if losses else 0
  696. task.get_logger().report_scalar(
  697. title=f"{tumor_type} Losses over Epochs",
  698. series=f"{tumor_type} Epoch valLoss",
  699. iteration=epoch + 1,
  700. value=avg_loss
  701. )
  702. task.get_logger().report_scalar("KD Losses over Epochs", "val_loss", iteration=epoch+1, value=epoch_val_loss)
  703. # %%
  704. def setup_environment(args):
  705. torch.manual_seed(0)
  706. args['out_checkpoint_dir'].mkdir(parents=True, exist_ok=True)
  707. def initialize_models():
  708. teacher_model = DynUNet(spatial_dims=3, in_channels=4, out_channels=4, deep_supervision=True, KD=True).to('cuda:0')
  709. student_model = Student_KD_loss().to('cuda:1')
  710. return teacher_model, student_model
  711. def initialize_optimizer_scheduler(student_model, args):
  712. optimizer = optim.AdamW(student_model.parameters(), lr=args['learning_rate'], weight_decay=args['weight_decay'], eps=1e-4)
  713. scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.1, patience=10, cooldown=1, threshold=0.001, min_lr=1e-6)
  714. return optimizer, scheduler
  715. def load_teacher_model(teacher_model, data_type, teacher_model_paths):
  716. teacher_model_path = teacher_model_paths.get(data_type)
  717. if teacher_model_path and Path(teacher_model_path).is_file():
  718. ckpt = torch.load(teacher_model_path, map_location='cuda:0', weights_only=True)
  719. teacher_model.load_state_dict(ckpt['teacher_model'])
  720. print(f"Loaded model: {teacher_model_path}")
  721. def load_student_checkpoint(student_model, optimizer, scaler, scheduler, args):
  722. checkpoint_path = args['in_checkpoint_dir'] / 'Student_model_after_epoch_58_trainLoss_0.8541_valLoss_0.3335.pth'
  723. if checkpoint_path.is_file():
  724. print(f"Found model {checkpoint_path}")
  725. ckpt = torch.load(checkpoint_path, map_location='cuda:1', weights_only=True)
  726. student_model.student.load_state_dict(ckpt['student_model'])
  727. optimizer.load_state_dict(ckpt['optimizer_student'])
  728. scaler.load_state_dict(ckpt['grad_scaler_state'])
  729. scheduler.load_state_dict(ckpt['scheduler_state_dict'])
  730. print(f"Loaded student model: {checkpoint_path} with lr: {optimizer.param_groups[0]['lr']}")
  731. return ckpt['epoch'] + 1
  732. return 0
  733. def train_epoch(epoch, trainLoader, train_config, start_ep):
  734. student_model = train_config['student_model']
  735. teacher_model = train_config['teacher_model']
  736. optimizer = train_config['optimizer']
  737. scaler = train_config['scaler']
  738. accumulation_steps = train_config['accumulation_steps']
  739. teacher_model_paths = train_config['teacher_model_paths']
  740. task = train_config['task']
  741. student_model.train()
  742. teacher_model.eval()
  743. epoch_losses = {'total': 0, 'seg': 0, 'bce': 0}
  744. tumors_losses = {'GLI': [], 'PED': [], 'SSA': [], 'MEN': [], 'MET': []}
  745. with tqdm(total=len(trainLoader), desc=f"(Epoch {epoch + 1}/{start_ep + train_config['epochs']})", unit='batch') as pbar:
  746. optimizer.zero_grad()
  747. for step, y in enumerate(trainLoader):
  748. batch_loss = 0
  749. for sub_step, data_type in enumerate(y['data_type']):
  750. imgs = y['imgs'][sub_step].unsqueeze(0).to('cuda:0')
  751. masks = y['masks'][sub_step].unsqueeze(0).to('cuda:0')
  752. load_teacher_model(teacher_model, data_type, teacher_model_paths)
  753. with amp.autocast('cuda:0'):
  754. teacher_outputs = teacher_model(imgs)
  755. detached_teacher_output = {k: v.detach().to('cuda:1') for k, v in teacher_outputs.items()}
  756. imgs, masks = imgs.to('cuda:1'), masks.to('cuda:1')
  757. with amp.autocast('cuda:1'):
  758. student_outputs = student_model(detached_teacher_output, {'imgs': imgs, 'masks': masks})
  759. loss = (student_outputs['batch_total_student_loss'] / accumulation_steps)
  760. batch_loss += loss.item()
  761. tumors_losses[data_type].append(loss.item())
  762. scaler.scale(loss).backward()
  763. task.get_logger().report_scalar(
  764. title=f"Tumors training losses per epoch {epoch+1}",
  765. series=f"{data_type} loss",
  766. iteration=len(tumors_losses[data_type]),
  767. value=float(loss.item())
  768. )
  769. for key in epoch_losses:
  770. if key != 'total':
  771. epoch_losses[key] += (student_outputs.get(f'{key}_weighted', 0) / accumulation_steps)
  772. if (sub_step + 1) % accumulation_steps == 0 or (sub_step + 1) == len(y['data_type']):
  773. scaler.step(optimizer)
  774. scaler.update()
  775. optimizer.zero_grad()
  776. epoch_losses['total'] += batch_loss
  777. pbar.update(1)
  778. for key in epoch_losses:
  779. epoch_losses[key] /= len(trainLoader)
  780. return epoch_losses, tumors_losses
  781. def log_KD_losses_over_epochs(epoch, epoch_losses, tumors_losses, task):
  782. for loss_type, val in epoch_losses.items():
  783. task.get_logger().report_scalar(
  784. title="KD Losses over Epochs",
  785. series=f"{loss_type} loss",
  786. iteration=epoch + 1,
  787. value=epoch_losses[loss_type]
  788. )
  789. for tumor_type, losses in tumors_losses.items():
  790. task.get_logger().report_scalar(
  791. title=f"{tumor_type} Losses over Epochs",
  792. series=f"{tumor_type} Epoch trainLoss",
  793. iteration=epoch + 1,
  794. value=sum(losses) / len(losses) if losses else 0
  795. )
  796. def validate_and_save(epoch, valLoader, train_config, epoch_losses):
  797. student_model = train_config['student_model']
  798. scheduler = train_config['scheduler']
  799. optimizer = train_config['optimizer']
  800. scaler = train_config['scaler']
  801. out_checkpoint_dir = train_config['out_checkpoint_dir']
  802. task = train_config['task']
  803. val_loss = evaluate(student_model, valLoader, epoch, task)
  804. scheduler.step(val_loss)
  805. task.get_logger().report_scalar("LR", "learning_rate", iteration=epoch+1, value=optimizer.param_groups[0]['lr'])
  806. print(f"Learning rate after epoch {epoch + 1}: {optimizer.param_groups[0]['lr']}")
  807. state = {
  808. 'epoch': epoch,
  809. 'student_model': student_model.student.state_dict(),
  810. 'optimizer_student': optimizer.state_dict(),
  811. 'lr': optimizer.param_groups[0]['lr'],
  812. 'grad_scaler_state': scaler.state_dict(),
  813. 'scheduler_state_dict': scheduler.state_dict(),
  814. 'val_dice_loss': val_loss
  815. }
  816. checkpoint_path = out_checkpoint_dir / f'Student_model_after_epoch_{epoch + 1}_trainLoss_{epoch_losses["total"]:.4f}_valLoss_{val_loss:.4f}.pth'
  817. torch.save(state, checkpoint_path)
  818. print(f"Model saved after epoch {epoch + 1}")
  819. def run_KD(trainLoader, valLoader, args):
  820. setup_environment(args)
  821. teacher_model, student_model = initialize_models()
  822. optimizer, scheduler = initialize_optimizer_scheduler(student_model, args)
  823. scaler = amp.GradScaler('cuda:1')
  824. teacher_model_paths = {
  825. 'GLI': '/kaggle/input/gliomateachernewlabels/Teacher_model_after_epoch_100_trainLoss_0.5972_valLoss_0.3019.pth',
  826. 'SSA': '/kaggle/input/africanewlabels/Teacher_model_after_epoch_67_trainLoss_1.1080_valLoss_0.5561.pth',
  827. 'PED': '/kaggle/input/pednewlabel/Teacher_model_after_epoch_99_trainLoss_1.4512_valLoss_1.0042.pth',
  828. 'MEN': '/kaggle/input/meningiomateachernewlabels/Teacher_model_after_epoch_85_trainLoss_0.5824_valLoss_0.3318.pth',
  829. 'MET': '/kaggle/input/met-teacher-new-labels/Teacher_model_after_epoch_100_trainLoss_1.6278_valLoss_0.7199.pth'
  830. }
  831. start_epoch = load_student_checkpoint(student_model, optimizer, scaler, scheduler, args)
  832. task = Task.init(project_name="Fairness KD 5 Tumors Models", task_name=f"ABL Study Fairness KD 5 Tumors with BCE+SEG NO KL", reuse_last_task_id=True)
  833. task.connect(args)
  834. task.add_tags(['ABL Study', 'BCE+SEG without KL', "Ahmed Pro"])
  835. print(f'''Starting Knowledge Distillation:
  836. Epochs: From {start_epoch + 1} to {start_epoch + args['epochs']}
  837. Batch size: 5 (effective through gradient accumulation)
  838. Learning rate: {args['learning_rate']}
  839. Training data coming from: {args['data_dirs']}
  840. ''')
  841. train_config = {
  842. 'teacher_model': teacher_model,
  843. 'student_model': student_model,
  844. 'optimizer': optimizer,
  845. 'scheduler': scheduler,
  846. 'scaler': scaler,
  847. 'accumulation_steps': 5,
  848. 'teacher_model_paths': teacher_model_paths,
  849. 'out_checkpoint_dir': args['out_checkpoint_dir'],
  850. 'task': task,
  851. 'epochs': args['epochs']
  852. }
  853. for epoch in range(start_epoch, start_epoch + args['epochs']):
  854. epoch_losses, tumors_losses = train_epoch(epoch, trainLoader, train_config, start_epoch)
  855. log_KD_losses_over_epochs(epoch, epoch_losses, tumors_losses, task)
  856. validate_and_save(epoch, valLoader, train_config, epoch_losses)
  857. print("Training completed.")
  858. task.close()
  859. # %%
  860. args = {
  861. 'workers': 2,
  862. 'epochs': 3,
  863. 'train_batch_size': 5,
  864. 'val_batch_size': 2,
  865. 'test_batch_size': 1,
  866. 'learning_rate': 1e-3,
  867. 'weight_decay': 1e-5,
  868. 'lambd': 0.0051,
  869. 'data_dirs': ["/kaggle/input/bratsglioma/Training/", "/kaggle/input/bratsafrica24/", "/kaggle/input/bratsped/Training/", "/kaggle/input/bratsmen/", "/kaggle/input/bratsmet24/"],
  870. 'in_checkpoint_dir': Path('/kaggle/input/data-abl-study-fairness-5-tumors-bce-seg-no-kl/'),
  871. 'out_checkpoint_dir': Path('/kaggle/working/')
  872. }
  873. trainLoader, valLoader, testLoader = prepare_data_loaders(args)
  874. run_KD(trainLoader, valLoader, args)
  875. # %% [markdown]
  876. # # Press here

ABL_(BCE+SEG).ipynb at commit 248789f, no license · at the source

Overview

Authors: Ahmed Elzayat1, Nourhan Hanafy1, Mariem Magdy1, Hazem Zakaria1, Mina Tayeh1, Abdulkhalek Al-Fakih2, Yeong Hyeon Gu2, Mohammed A. Al-masni2, Meena M. Makary1,3
  1. Systems and Biomedical Engineering Department, Faculty of Engineering, Cairo University,Cairo, Egypt
  2. Department of Artificial Intelligence and Data Science, College of Artificial Intelligence Convergence, Sejong University,Seoul, Republic of Korea
  3. Department of Radiology, Massachusetts General Hospital, Athinoula A. Martinos Center for Biomedical Imaging, Harvard Medical School,Charlestown, MA USA
Institutions: Cairo University (Egypt); Sejong University (South Korea); Harvard University (United States)
Journal: Scientific reports, volume 16, issue 1, article 12969
Dates: received 15 July 2025; accepted 7 January 2026; published online 10 March 2026
Type: Research article · Language: English
License: CC BY-NC-ND
Identifiers: DOI 10.1038/s41598-026-35627-x · PMID 41807430 · PMCID PMC13096204 · OpenAlex W7134947088
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: structural MRI / diffusion (modality), human (organism), other condition (population), methods / tools (subfield)
Methods: Connectivity, Machine learning
Keywords: Generalizability, Foundation model, Knowledge distillation, Brain tumor segmentation, Cancer, Computational biology and bioinformatics, Mathematics and computing, Neuroscience, Oncology
MeSH: Brain Neoplasms*, Glioma*, Image Processing, Computer-Assisted*, Humans, Magnetic Resonance Imaging (* major topic)
Topic: Advanced Neural Network Applications (Computer Vision and Pattern Recognition, Computer Science), according to OpenAlex
Funding: IITP (Institute of Information & Communications Technology Planning & Evaluation) - ITRC (Information Technology Research Center) grant funded by the Korea government (Ministry of Science and ICT) (IITP-2025-RS-2024-00437191); National Research Foundation of Korea (NRF) funded by the Korean government (MSIT) (No. RS-2023-00243034 and No. RS-2026-25487511)
Citations: not cited yet (Europe PMC); 53 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.

AhmeddEmad7/Brain-Tumor-Segmentation-Advancing-Generalizability

License: none: the authors keep all their rights
State: the link answers, verified on 30 September 2026
Evidence: files inventoried
Commit: 248789f547fba152cc4f61c1f4026871246c762a, 15 April 2026
Languages: Python (100), TypeScript (85), Jupyter (25), JavaScript (8), Shell (2)
Size: 482 files, 220 scripts
Software Heritage: not archived
Found in: “Data availability”
Holds: README, environment (AI/requirements.txt, software/docker-compose.yaml, software/backend/requirements.txt, software/frontend/Dockerfile, software/backend/ApiGateway/Dockerfile, software/backend/ApiGateway/requirements.txt, software/backend/Services/NiftiStorage/Dockerfile, software/backend/Services/NiftiStorage/requirements.txt, software/backend/Services/Orthanc/Dockerfile, software/backend/Services/Redis/dockerfile, software/backend/Services/Reporting/Dockerfile, software/backend/Services/Reporting/requirements.txt), documentation, 25 notebooks
Not found: license file, CITATION.cff, tests, continuous integration
Tools: PyTorch (15 files), MONAI (9 files), NumPy (8 files), NiBabel (7 files), Matplotlib (4 files)
Availability: 1 check, the latest on 30 September 2026: the link answers
  • 30 September 2026: the link answers
16 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;
  • 16 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

Datasets cited

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:

Read it in the paper: doi.org/10.1038/s41598-026-35627-x.

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

Recorded: type, language, journal, volume, issue, pages, dates, 9 authors, 9 keywords, 5 MeSH terms, 2 funders, 18 references.

Cite

This paper

Elzayat, A., Hanafy, N., Magdy, M., Zakaria, H., Tayeh, M., Al-Fakih, A., Gu, Y. H., Al-masni, M. A., & Makary, M. M. (2026). Distilling population specific expertise into a unified model for generalizable brain tumor segmentation. Scientific reports, 16(1), 12969. https://doi.org/10.1038/s41598-026-35627-x

BibTeX

@article{elzayat2026distilling,
author = {Elzayat, Ahmed and Hanafy, Nourhan and Magdy, Mariem and Zakaria, Hazem and Tayeh, Mina and Al-Fakih, Abdulkhalek and Gu, Yeong Hyeon and Al-masni, Mohammed A. and Makary, Meena M.},
title = {{Distilling population specific expertise into a unified model for generalizable brain tumor segmentation}},
journal = {Scientific reports},
year = {2026},
month = mar,
volume = {16},
number = {1},
pages = {12969},
publisher = {Nature Publishing Group},
issn = {2045-2322},
doi = {10.1038/s41598-026-35627-x},
url = {https://doi.org/10.1038/s41598-026-35627-x},
pmid = {41807430},
pmcid = {PMC13096204}
}

RIS

TY - JOUR
AU - Elzayat, Ahmed
AU - Hanafy, Nourhan
AU - Magdy, Mariem
AU - Zakaria, Hazem
AU - Tayeh, Mina
AU - Al-Fakih, Abdulkhalek
AU - Gu, Yeong Hyeon
AU - Al-masni, Mohammed A.
AU - Makary, Meena M.
TI - Distilling population specific expertise into a unified model for generalizable brain tumor segmentation
T2 - Scientific reports
J2 - Sci Rep
PY - 2026
DA - 2026/03/10
VL - 16
IS - 1
SP - 12969
SN - 2045-2322
PB - Nature Publishing Group
DO - 10.1038/s41598-026-35627-x
UR - https://doi.org/10.1038/s41598-026-35627-x
LA - en
ER -

CSL-JSON

{
"id": "10.1038/s41598-026-35627-x",
"type": "article-journal",
"title": "Distilling population specific expertise into a unified model for generalizable brain tumor segmentation",
"container-title": "Scientific reports",
"author": [
{
"family": "Elzayat",
"given": "Ahmed"
},
{
"family": "Hanafy",
"given": "Nourhan"
},
{
"family": "Magdy",
"given": "Mariem"
},
{
"family": "Zakaria",
"given": "Hazem"
},
{
"family": "Tayeh",
"given": "Mina"
},
{
"family": "Al-Fakih",
"given": "Abdulkhalek"
},
{
"family": "Gu",
"given": "Yeong Hyeon"
},
{
"family": "Al-masni",
"given": "Mohammed A."
},
{
"family": "Makary",
"given": "Meena M."
}
],
"container-title-short": "Sci Rep",
"volume": "16",
"issue": "1",
"page": "12969",
"DOI": "10.1038/s41598-026-35627-x",
"PMID": "41807430",
"PMCID": "PMC13096204",
"ISSN": "2045-2322",
"publisher": "Nature Publishing Group",
"URL": "https://doi.org/10.1038/s41598-026-35627-x",
"language": "en",
"issued": {
"date-parts": [
[
2026,
3,
10
]
]
}
}

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.1016/j.patter.2026.101538 [code]
A multi-modal foundation model for brain disease diagnosis and medical imaging.
Journal: Patterns (New York, N.Y.)
In common: MONAI, NiBabel, PyTorch, 2 other tools, 3 references
[2] doi:10.1371/journal.pone.0354511 [code]
TokenUNet: A new case for transformers integration in efficient and interpretable 3D UNets for brain imaging segmentation.
Journal: PloS one
In common: MONAI, NiBabel, PyTorch, 2 other tools, methods / tools, other condition, 2 references
[3] doi:10.1038/s41598-026-48496-1 [code]
A unified FLAIR hyperintensity segmentation model for various CNS tumor types and acquisition time points.
Journal: Scientific reports
In common: NiBabel, Matplotlib, NumPy, methods / tools, structural MRI / diffusion, other condition, 4 references
[4] doi:10.21037/qims-2026-0792 [code]
An nnU-Net-based framework with adaptive feature representation for 3D brain tumor segmentation.
Journal: Quantitative imaging in medicine and surgery
In common: NiBabel, PyTorch, Matplotlib, 1 other tool, methods / tools, structural MRI / diffusion, other condition, 3 references
[5] doi:10.1158/2767-9764.crc-25-0710 [code]
MRI Deep Learning for Differentiating Glioblastoma, IDH Wild-type from Central Nervous System Diffuse Large B-cell Lymphoma.
Journal: Cancer research communications
In common: MONAI, NiBabel, PyTorch, 2 other tools, structural MRI / diffusion, other condition, 1 reference
[6] doi:10.1371/journal.pdig.0001316 [code]
Reliability of a convolutional neural network in segmenting multiple sclerosis lesions from MRI: Impact of data augmentation, image modality and tolerance with U-Net architecture.
Journal: PLOS digital health
In common: MONAI, NiBabel, PyTorch, 2 other tools, methods / tools, structural MRI / diffusion, 1 reference
[7] doi:10.3390/diagnostics16172806
Brain Tumor Segmentation and Grading on MRI Using Deep Learning: A Systematic Literature Review and Benchmark-Driven Comparative Analysis.
Journal: Diagnostics (Basel, Switzerland)
In common: methods / tools, structural MRI / diffusion, other condition, 5 references
[8] doi:10.1038/s41598-026-54446-8 [code]
Deep learning-based Desikan-Killiany parcellation of the brain using diffusion MRI.
Journal: Scientific reports
In common: MONAI, NiBabel, PyTorch, 1 other tool, methods / tools, structural MRI / diffusion, 1 reference
[9] doi:10.3389/fnins.2026.1870124 [code]
An end-to-end pipeline for automated fetal brain segmentation and biometry from 3D SSFP MRI.
Journal: Frontiers in neuroscience
In common: MONAI, NiBabel, PyTorch, 2 other tools, structural MRI / diffusion, 1 reference
[10] doi:10.1371/journal.pone.0351405 [code]
Unimodal vs. multimodal deep learning for non-invasive MGMT promoter methylation prediction in glioblastoma: A systematic evaluation on the BraTS 2021 dataset.
Journal: PloS one
In common: PyTorch, NumPy, methods / tools, structural MRI / diffusion, other condition, 3 references

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.