Distilling population specific expertise into a unified model for generalizable brain tumor segmentation.
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] § 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] § 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] § 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] § 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] § 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] § 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] § 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] § 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] § 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] § Materials and methods › Inference ↔ AI/inference/postprocess.py, lines 5–29 · score 0.72 · sliding window inference, segmentation mask, tumor regions, prediction, model
- [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] § Materials and methods › Inference ↔ AI/inference/predict.py, lines 37–102 · score 0.67 · sliding window inference, Gaussian, overlapping, volumes, prediction, mask
- [13] § Materials and methods › BraTS datasets ↔ AI/inference/preprocess.py, lines 8–49 · score 0.65 · T1 weighted, T2 weighted, FLAIR, preprocessed, volume
- [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] § Experiments and setup › Backbone architecture ↔ AI/models/blocks.py, lines 82–156 · score 0.63 · skip connections, medical imaging, resolutions, encoder, UNet, decoder
- [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] § 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
- # %% [markdown]
- # # Importing & Loading Dependencies
- # %%
- !pip install monai
- import nibabel as nib
- from monai.transforms import LoadImage, Compose, NormalizeIntensityd, RandFlipd, RandAdjustContrastd, Resized, CropForegroundd, SpatialPadd
- import matplotlib.pyplot as plt
- from torch.utils.data import DataLoader, Dataset
- import torch
- import torch.nn as nn
- import torch.nn.functional as F
- import numpy as np
- from typing import Optional, Sequence, Tuple, Union
- from torch.nn.functional import interpolate
- from monai.networks.blocks.convolutions import Convolution
- from monai.networks.layers.factories import Act, Norm
- from monai.networks.layers.utils import get_act_layer, get_norm_layer
- from monai.metrics import DiceMetric, HausdorffDistanceMetric
- from torch import nn, optim, amp
- from itertools import chain
- from monai.losses import DiceLoss
- from tqdm import tqdm
- from pathlib import Path
- import math
- import os
- import random
- # %% [markdown]
- # # Creating Dataset with Preprocessing
- # %%
- class CustomDataset3D(Dataset):
- def __init__(self, data_dirs, patient_lists, mode):
- self.data_dirs = data_dirs
- self.patient_lists = patient_lists
- self.mode = mode
- @staticmethod
- def resize_with_aspect_ratio(keys, target_size):
- def transform(data):
- for key in keys:
- volume = data[key]
- original_shape = volume.shape[-3:]
- scaling_factor = min(
- target_size[0] / original_shape[0],
- target_size[1] / original_shape[1],
- target_size[2] / original_shape[2]
- )
- # Computing the intermediate size while preserving aspect ratio
- new_shape = [
- int(dim * scaling_factor) for dim in original_shape
- ]
- # Resizing to the intermediate shape
- resize_transform = Resized(keys=[key], spatial_size=new_shape, mode="trilinear" if key == "imgs" else "nearest-exact")
- data = resize_transform(data)
- # Padding to the final target size
- pad_transform = SpatialPadd(keys=[key], spatial_size=target_size, mode="constant")
- data = pad_transform(data)
- return data
- return transform
- def preprocess(cls, data, mode):
- if mode == 'training':
- transform = Compose([
- CropForegroundd(keys=["imgs", "masks"], source_key="imgs"),
- cls.resize_with_aspect_ratio(keys=["imgs", "masks"], target_size=[128, 128, 128]),
- NormalizeIntensityd( keys=['imgs'], nonzero=False, channel_wise=True),
- RandFlipd(keys=["imgs", "masks"],
- prob=0.5,
- spatial_axis=2,
- ),
- RandAdjustContrastd(
- keys=["imgs"],
- prob=0.15,
- gamma=(0.65, 1.5),
- ),
- ])
- elif mode == 'validation':
- transform = Compose([
- CropForegroundd(keys=["imgs", "masks"], source_key="imgs"),
- cls.resize_with_aspect_ratio(keys=["imgs", "masks"], target_size=[128, 128, 128]),
- NormalizeIntensityd( keys=['imgs'], nonzero=False, channel_wise=True)
- ])
- else: # 'testing'
- transform = Compose([
- CropForegroundd(keys=["imgs", "masks"], source_key="imgs"),
- cls.resize_with_aspect_ratio(keys=["imgs", "masks"], target_size=[128, 128, 128]),
- NormalizeIntensityd( keys=['imgs'], nonzero=False, channel_wise=True)
- ])
- augmented_data = transform(data)
- return augmented_data
- def __len__(self):
- return len(self.patient_lists)
- def __getitem__(self, idx):
- patient_id = self.patient_lists[idx]
- loadimage = LoadImage(reader='NibabelReader', image_only=True)
- data_type=patient_id.split('-')[1]
- if data_type == 'GLI':
- patient_folder_path = os.path.join('/kaggle/input/bratsglioma/Training', patient_id)
- elif data_type == 'SSA':
- patient_folder_path = os.path.join('/kaggle/input/bratsafrica24', patient_id)
- elif data_type == 'PED':
- patient_folder_path = os.path.join('/kaggle/input/bratsped/Training', patient_id)
- elif data_type == 'MEN':
- patient_folder_path = os.path.join('/kaggle/input/bratsmen', patient_id)
- else:
- patient_folder_path = os.path.join('/kaggle/input/bratsmet24', patient_id)
- def resolve_file_path(folder, name):
- file_path = os.path.join(folder, name)
- # Check if the given path is a directory (case with 4 subdirs)
- if os.path.isdir(file_path):
- # Find the first file inside the directory that ends with .nii
- for root, _, files in os.walk(file_path):
- for file in files:
- if file.endswith(".nii"):
- return os.path.join(root, file)
- return file_path
- # Resolve paths for all required image types
- t1c_path = resolve_file_path(patient_folder_path, patient_id + '-t1c.nii')
- t1n_path = resolve_file_path(patient_folder_path, patient_id + '-t1n.nii')
- t2f_path = resolve_file_path(patient_folder_path, patient_id + '-t2f.nii')
- t2w_path = resolve_file_path(patient_folder_path, patient_id + '-t2w.nii')
- seg_path = os.path.join(patient_folder_path, patient_id + '-seg.nii')
- t1c_loader = loadimage( t1c_path )
- t1n_loader = loadimage( t1n_path )
- t2f_loader = loadimage( t2f_path )
- t2w_loader = loadimage( t2w_path )
- masks_loader = loadimage( seg_path )
- # Make the dimension of channel
- t1c_tensor = torch.Tensor(t1c_loader).unsqueeze(0)
- t1n_tensor = torch.Tensor(t1n_loader).unsqueeze(0)
- t2f_tensor = torch.Tensor(t2f_loader).unsqueeze(0)
- t2w_tensor = torch.Tensor(t2w_loader).unsqueeze(0)
- masks_tensor = torch.Tensor(masks_loader).unsqueeze(0)
- concat_tensor = torch.cat( (t1c_tensor, t1n_tensor, t2f_tensor, t2w_tensor, masks_tensor), 0 )
- data = {
- 'imgs' : np.array(concat_tensor[0:4,:,:,:]),
- 'masks' : np.array(concat_tensor[4:,:,:,:])
- }
- augmented_imgs_masks = self.preprocess(data, self.mode)
- imgs = np.array(augmented_imgs_masks['imgs'])
- masks = np.array(augmented_imgs_masks['masks'])
- y = {
- 'imgs' : torch.from_numpy(imgs).type(torch.FloatTensor),
- 'masks' : torch.from_numpy(masks).type(torch.FloatTensor),
- 'patient_id' : patient_id,
- 'data_type' : data_type
- }
- return y
- # %% [markdown]
- # # Data Loaders
- # %%
- def combine_datasets(dataset_lists, batch_size=3):
- max_len = max(len(dataset) for dataset in dataset_lists)
- # Ensure batch_size matches the number of datasets
- if batch_size != len(dataset_lists):
- raise ValueError("Batch size must equal the number of datasets for this function.")
- combined_paths = []
- for i in range(0, max_len, batch_size):
- for j in range(batch_size):
- index = (i + j) % max_len
- batch = [dataset[index % len(dataset)] for dataset in dataset_lists]
- combined_paths.extend(batch)
- # if j == 0:
- # print(f"Batch {(i // batch_size) + 1}: {batch}")
- return combined_paths
- # %%
- def prepare_data_loaders(args):
- train_datasets, val_datasets, test_datasets = [], [], []
- split_ratio = {'training': 0.71, 'validation': 0.09, 'testing': 0.2}
- for i, data_dir in enumerate(args['data_dirs']):
- patient_lists = os.listdir( data_dir )
- patient_lists.sort()
- total_patients = len(patient_lists)
- random.seed(5)
- random.shuffle(patient_lists)
- train_split = int(split_ratio['training'] * total_patients)
- val_split = int(split_ratio['validation'] * total_patients)
- train_patient_lists = patient_lists[:train_split]
- val_patient_lists = patient_lists[train_split : train_split + val_split]
- test_patient_lists = patient_lists[train_split + val_split :]
- train_patient_lists.sort()
- val_patient_lists.sort()
- test_patient_lists.sort()
- print(f'Number of training samples in {data_dir.split("/")[3]} DataSet: {len(train_patient_lists)}')
- print(f'Number of validation samples in {data_dir.split("/")[3]} DataSet: {len(val_patient_lists)}')
- print(f'Number of testing samples in {data_dir.split("/")[3]} DataSet: {len(test_patient_lists)} ')
- train_datasets.append(train_patient_lists)
- val_datasets.append(val_patient_lists)
- test_datasets.append(test_patient_lists)
- combined_trainDataset = combine_datasets(train_datasets, batch_size=args['train_batch_size'])
- combined_valDataset = list(chain.from_iterable(val_datasets))
- combined_testDataset = list(chain.from_iterable(test_datasets))
- print(f'Number of combined training samples', len(combined_trainDataset))
- print(f'Number of combined validation samples', len(combined_valDataset))
- print(f'Number of combined testing samples', len(combined_testDataset))
- trainDataset = CustomDataset3D( args['data_dirs'], combined_trainDataset, mode='training')
- valDataset = CustomDataset3D( args['data_dirs'], combined_valDataset, mode='validation')
- testDataset = CustomDataset3D( args['data_dirs'], combined_testDataset, mode='testing')
- trainLoader = DataLoader(
- trainDataset, batch_size=args['train_batch_size'], num_workers=args['workers'], prefetch_factor=2,
- pin_memory=True, shuffle=False)
- valLoader = DataLoader(
- valDataset, batch_size=args['val_batch_size'], num_workers=args['workers'], prefetch_factor=2,
- pin_memory=True, shuffle=False)
- testLoader = DataLoader(
- testDataset, batch_size=args['test_batch_size'], num_workers=args['workers'], prefetch_factor=2,
- pin_memory=True, shuffle=False)
- return trainLoader, valLoader, testLoader
- # %% [markdown]
- # # Visualizing Data
- # %%
- # args = {
- # 'workers': 2,
- # 'epochs': 10,
- # 'train_batch_size': 2,
- # 'val_batch_size': 2,
- # 'test_batch_size': 2,
- # 'learning_rate': 1e-3,
- # 'weight_decay': 1e-5,
- # 'lambd': 0.0051,
- # 'data_dir': '/kaggle/input/bratsafrica24/',
- # 'in_checkpoint_dir': Path('/kaggle/input/adultgliomamodel-45epochs'),
- # 'out_checkpoint_dir': Path('/kaggle/working/')
- # }
- # trainLoader, valLoader, testLoader = prepare_data_loaders(args)
- # for step, y in enumerate( trainLoader ):
- # print(y['imgs'].shape)
- # print(y['patient_id'])
- # fig, axes = plt.subplots(1, 4, figsize=(16, 4))
- # for sequence in range(4):
- # sequence_data = y['imgs'][0][sequence, :, :, :].cpu().detach().numpy()
- # slice_index = sequence_data.shape[2] // 2
- # axes[sequence].imshow(np.rot90(sequence_data[:, :, slice_index]), cmap='gray', origin='lower')
- # axes[sequence].set_title(f'Sequence {sequence + 1}')
- # plt.show()
- # %% [markdown]
- # # DynUNet Model
- # %%
- class UnetBasicBlock(nn.Module):
- """
- A CNN module module that can be used for DynUNet, based on:
- `Automated Design of Deep Learning Methods for Biomedical Image Segmentation <https://arxiv.org/abs/1904.08128>`_.
- `nnU-Net: Self-adapting Framework for U-Net-Based Medical Image Segmentation <https://arxiv.org/abs/1809.10486>`_.
- Args:
- spatial_dims: number of spatial dimensions.
- in_channels: number of input channels.
- out_channels: number of output channels.
- kernel_size: convolution kernel size.
- stride: convolution stride.
- norm_name: feature normalization type and arguments.
- act_name: activation layer type and arguments.
- dropout: dropout probability.
- """
- def __init__(
- self,
- spatial_dims: int,
- in_channels: int,
- out_channels: int,
- kernel_size: Union[Sequence[int], int],
- stride: Union[Sequence[int], int],
- norm_name: Union[Tuple, str] = ("INSTANCE", {"affine": True}),
- act_name: Union[Tuple, str] = ("leakyrelu", {"inplace": True, "negative_slope": 0.01}),
- dropout: Optional[Union[Tuple, str, float]] = None,
- ):
- super().__init__()
- self.conv1 = get_conv_layer(
- spatial_dims,
- in_channels,
- out_channels,
- kernel_size=kernel_size,
- stride=stride,
- dropout=dropout,
- conv_only=True,
- )
- self.conv2 = get_conv_layer(
- spatial_dims,
- out_channels,
- out_channels,
- kernel_size=kernel_size,
- stride=1,
- dropout=dropout,
- conv_only=True
- )
- self.lrelu = get_act_layer(name=act_name)
- self.norm1 = get_norm_layer(name=norm_name, spatial_dims=spatial_dims, channels=out_channels)
- self.norm2 = get_norm_layer(name=norm_name, spatial_dims=spatial_dims, channels=out_channels)
- def forward(self, inp):
- out = self.conv1(inp)
- out = self.norm1(out)
- out = self.lrelu(out)
- out = self.conv2(out)
- out = self.norm2(out)
- out = self.lrelu(out)
- return out
- class UnetUpBlock(nn.Module):
- """
- An upsampling module that can be used for DynUNet, based on:
- `Automated Design of Deep Learning Methods for Biomedical Image Segmentation <https://arxiv.org/abs/1904.08128>`_.
- `nnU-Net: Self-adapting Framework for U-Net-Based Medical Image Segmentation <https://arxiv.org/abs/1809.10486>`_.
- Args:
- spatial_dims: number of spatial dimensions.
- in_channels: number of input channels.
- out_channels: number of output channels.
- kernel_size: convolution kernel size.
- stride: convolution stride.
- upsample_kernel_size: convolution kernel size for transposed convolution layers.
- norm_name: feature normalization type and arguments.
- act_name: activation layer type and arguments.
- dropout: dropout probability.
- trans_bias: transposed convolution bias.
- """
- def __init__(
- self,
- spatial_dims: int,
- in_channels: int,
- out_channels: int,
- kernel_size: Union[Sequence[int], int],
- upsample_kernel_size: Union[Sequence[int], int],
- norm_name: Union[Tuple, str] = ("INSTANCE", {"affine": True}),
- act_name: Union[Tuple, str] = ("leakyrelu", {"inplace": True, "negative_slope": 0.01}),
- dropout: Optional[Union[Tuple, str, float]] = None,
- trans_bias: bool = False,
- ):
- super().__init__()
- upsample_stride = upsample_kernel_size
- # ( a purple arrow in the paper )
- self.transp_conv = get_conv_layer(
- spatial_dims,
- in_channels,
- out_channels,
- kernel_size=upsample_kernel_size,
- stride=upsample_stride,
- dropout=dropout,
- bias=trans_bias,
- conv_only=True,
- is_transposed=True,
- )
- # A light blue conv blocks in the decoder of nnUNet
- self.conv_block = UnetBasicBlock(
- spatial_dims,
- out_channels + out_channels,
- out_channels,
- kernel_size=kernel_size,
- stride=1,
- dropout=dropout,
- norm_name=norm_name,
- act_name=act_name,
- )
- def forward(self, inp, skip):
- # number of channels for skip should equals to out_channels
- out = self.transp_conv(inp)
- out = torch.cat((out, skip), dim=1)
- out = self.conv_block(out)
- return out
- class UnetOutBlock(nn.Module):
- def __init__(
- self, spatial_dims: int, in_channels: int, out_channels: int, dropout: Optional[Union[Tuple, str, float]] = None
- ):
- super().__init__()
- self.conv = get_conv_layer(
- spatial_dims, in_channels, out_channels, kernel_size=1, stride=1, dropout=dropout, bias=True, conv_only=True
- )
- def forward(self, inp):
- return self.conv(inp)
- def get_conv_layer(
- spatial_dims: int,
- in_channels: int,
- out_channels: int,
- kernel_size: Union[Sequence[int], int] = 3,
- stride: Union[Sequence[int], int] = 1,
- act: Optional[Union[Tuple, str]] = Act.PRELU,
- norm: Union[Tuple, str] = Norm.INSTANCE,
- dropout: Optional[Union[Tuple, str, float]] = None,
- bias: bool = False,
- conv_only: bool = True,
- is_transposed: bool = False,
- ):
- padding = get_padding(kernel_size, stride)
- output_padding = None
- if is_transposed:
- output_padding = get_output_padding(kernel_size, stride, padding)
- return Convolution(
- spatial_dims,
- in_channels,
- out_channels,
- strides=stride,
- kernel_size=kernel_size,
- act=act,
- norm=norm,
- dropout=dropout,
- bias=bias,
- conv_only=conv_only,
- is_transposed=is_transposed,
- padding=padding,
- output_padding=output_padding,
- )
- def get_padding(
- kernel_size: Union[Sequence[int], int], stride: Union[Sequence[int], int]
- ) -> Union[Tuple[int, ...], int]:
- kernel_size_np = np.atleast_1d(kernel_size)
- stride_np = np.atleast_1d(stride)
- padding_np = (kernel_size_np - stride_np + 1) / 2
- if np.min(padding_np) < 0:
- raise AssertionError("padding value should not be negative, please change the kernel size and/or stride.")
- padding = tuple(int(p) for p in padding_np)
- return padding if len(padding) > 1 else padding[0]
- def get_output_padding(
- kernel_size: Union[Sequence[int], int], stride: Union[Sequence[int], int], padding: Union[Sequence[int], int]
- ) -> Union[Tuple[int, ...], int]:
- kernel_size_np = np.atleast_1d(kernel_size)
- stride_np = np.atleast_1d(stride)
- padding_np = np.atleast_1d(padding)
- out_padding_np = 2 * padding_np + stride_np - kernel_size_np
- if np.min(out_padding_np) < 0:
- raise AssertionError("out_padding value should not be negative, please change the kernel size and/or stride.")
- out_padding = tuple(int(p) for p in out_padding_np)
- return out_padding if len(out_padding) > 1 else out_padding[0]
- def set_requires_grad(nets, requires_grad=False):
- if not isinstance(nets, list):
- nets = [nets]
- for net in nets:
- if net is not None:
- for param in net.parameters():
- param.requires_grad = requires_grad
- # %%
- class DynUNet(nn.Module):
- def __init__(
- self,
- spatial_dims: int,
- in_channels: int,
- out_channels: int,
- deep_supervision: bool,
- KD: bool = False
- ):
- super().__init__()
- self.spatial_dims = spatial_dims
- self.in_channels = in_channels
- self.out_channels = out_channels
- self.deep_supervision = deep_supervision
- self.KD_enabled = KD
- self.input_conv = UnetBasicBlock( spatial_dims=self.spatial_dims,
- in_channels=self.in_channels,
- out_channels=64,
- kernel_size=3,
- stride=1
- )
- self.down1 = UnetBasicBlock( spatial_dims=self.spatial_dims,
- in_channels=64,
- out_channels=96,
- kernel_size=3,
- stride=2 # Reduces spatial dims by 2
- )
- self.down2 = UnetBasicBlock( spatial_dims=self.spatial_dims,
- in_channels=96,
- out_channels=128,
- kernel_size=3,
- stride=2
- )
- self.down3 = UnetBasicBlock( spatial_dims=self.spatial_dims,
- in_channels=128,
- out_channels=192,
- kernel_size=3,
- stride=2
- )
- self.down4 = UnetBasicBlock( spatial_dims=self.spatial_dims,
- in_channels=192,
- out_channels=256,
- kernel_size=3,
- stride=2
- )
- self.down5 = UnetBasicBlock( spatial_dims=self.spatial_dims,
- in_channels=256,
- out_channels=384,
- kernel_size=3,
- stride=2
- )
- self.bottleneck = UnetBasicBlock( spatial_dims=self.spatial_dims,
- in_channels=384,
- out_channels=512,
- kernel_size=3,
- stride=2
- )
- self.up1 = UnetUpBlock( spatial_dims=self.spatial_dims,
- in_channels=512,
- out_channels=384,
- kernel_size=3,
- upsample_kernel_size=2
- )
- self.up2 = UnetUpBlock( spatial_dims=self.spatial_dims,
- in_channels=384,
- out_channels=256,
- kernel_size=3,
- upsample_kernel_size=2
- )
- self.up3 = UnetUpBlock( spatial_dims=self.spatial_dims,
- in_channels=256,
- out_channels=192,
- kernel_size=3,
- upsample_kernel_size=2
- )
- self.up4 = UnetUpBlock( spatial_dims=self.spatial_dims,
- in_channels=192,
- out_channels=128,
- kernel_size=3,
- upsample_kernel_size=2
- )
- self.up5 = UnetUpBlock( spatial_dims=self.spatial_dims,
- in_channels=128,
- out_channels=96,
- kernel_size=3,
- upsample_kernel_size=2
- )
- self.up6 = UnetUpBlock( spatial_dims=self.spatial_dims,
- in_channels=96,
- out_channels=64,
- kernel_size=3,
- upsample_kernel_size=2
- )
- self.out1 = UnetOutBlock( spatial_dims=self.spatial_dims,
- in_channels=64,
- out_channels=self.out_channels,
- )
- self.out2 = UnetOutBlock( spatial_dims=self.spatial_dims,
- in_channels=96,
- out_channels=self.out_channels,
- )
- self.out3 = UnetOutBlock( spatial_dims=self.spatial_dims,
- in_channels=128,
- out_channels=self.out_channels,
- )
- def forward( self, input ):
- # Input
- x0 = self.input_conv( input ) # x0.shape = (B x 64 x 128 x 128 x 128)
- # Encoder
- x1 = self.down1( x0 ) # x1.shape = (B x 96 x 64 x 64 x 64)
- x2 = self.down2( x1 ) # x2.shape = (B x 128 x 32 x 32 x 32)
- x3 = self.down3( x2 ) # x3.shape = (B x 192 x 16 x 16 x 16)
- x4 = self.down4( x3 ) # x4.shape = (B x 256 x 8 x 8 x 8)
- x5 = self.down5( x4 ) # x5.shape = (B x 384 x 4 x 4 x 4)
- # Bottleneck
- x6 = self.bottleneck( x5 ) # x6.shape = (B x 512 x 2 x 2 x 2)
- # Decoder
- x7 = self.up1( x6, x5 ) # x7.shape = (B x 384 x 4 x 4 x 4)
- x8 = self.up2( x7, x4 ) # x8.shape = (B x 256 x 8 x 8 x 8)
- x9 = self.up3( x8, x3 ) # x9.shape = (B x 192 x 16 x 16 x 16)
- x10 = self.up4( x9, x2 ) # x10.shape = (B x 128 x 32 x 32 x 32)
- x11 = self.up5( x10, x1 ) # x11.shape = (B x 96 x 64 x 64 x 64)
- x12 = self.up6( x11, x0 ) # x12.shape = (B x 64 x 128 x 128 x 128)
- # Output
- output1 = self.out1( x12 )
- if (self.training and self.deep_supervision) or self.KD_enabled:
- # output['pred'].shape = B x 3 x 4 x 128 x 128 x 128
- output2 = interpolate( self.out2( x11 ), output1.shape[2:])
- output3 = interpolate( self.out3( x10 ), output1.shape[2:])
- output_all = [ output1, output2, output3 ]
- return { 'pred' : torch.stack(output_all, dim=1),
- 'bottleneck_feature_map' : x6 }
- return { 'pred' : output1 }
- # %% [markdown]
- # # Visualizing Model Instance
- # %%
- # !pip install torchsummary
- # from torchsummary import summary
- # # Initialize your DynUNet model
- # model = DynUNet(spatial_dims=3, in_channels=4, out_channels=4, deep_supervision=True, KD=True)
- # # Print model summary
- # summary(model, input_size=(4, 128, 128, 128)) # Adjust input_size according to your needs
- # %% [markdown]
- # # ClearML
- # %%
- !pip install clearml
- from clearml import Task
- %env CLEARML_WEB_HOST=https://app.clear.ml/
- %env CLEARML_API_HOST=https://api.clear.ml
- %env CLEARML_FILES_HOST=https://files.clear.ml
- %env CLEARML_API_ACCESS_KEY=CLEARML_API_ACCESS_KEY
- %env CLEARML_API_SECRET_KEY=CLEARML_API_SECRET_KEY
- # %% [markdown]
- # # GPUs Check
- # %%
- if torch.cuda.is_available():
- num_gpus = torch.cuda.device_count()
- print(f"Number of GPUs available: {num_gpus}")
- for i in range(num_gpus):
- print(f"GPU {i}: {torch.cuda.get_device_name(i)}")
- else:
- print("No GPU available. Running on CPU.")
- # %%
- # # For freeing gpu
- # import gc; gc.collect(); torch.cuda.empty_cache()
- # %% [markdown]
- # # Loss Function
- # %%
- class LossFunction(nn.Module):
- def __init__(self):
- super(LossFunction, self).__init__()
- self.dice = DiceLoss(sigmoid=True, batch=True, smooth_nr=1e-05, smooth_dr=1e-05)
- self.ce = nn.BCEWithLogitsLoss()
- def _loss(self, p, y):
- return self.dice(p, y) + self.ce(p, y.float())
- def forward(self, p, y):
- y_wt, y_tc, y_et = y > 0, ((y == 1) + (y == 3)) > 0, y == 3
- p_wt, p_tc, p_et = p[:, 1].unsqueeze(1), p[:, 2].unsqueeze(1), p[:, 3].unsqueeze(1)
- l_wt, l_tc, l_et = self._loss(p_wt, y_wt), self._loss(p_tc, y_tc), self._loss(p_et, y_et)
- return l_wt + l_tc + l_et
- # %% [markdown]
- # # Student KD Model
- # %%
- class Student_KD_loss(nn.Module):
- def __init__(self):
- super().__init__()
- self.student = DynUNet( spatial_dims=3, in_channels=4, out_channels=4, deep_supervision=True)
- self.loss_fn = LossFunction()
- self.temperature = 5.0
- self.bce_loss = nn.BCEWithLogitsLoss()
- def forward(self, teacher_outputs, y):
- with amp.autocast('cuda:1'):
- student_outputs = self.student( y['imgs'] )
- # Student loss with Deep supervision -> (Dice loss)
- 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
- segloss_s_decoder_2 = self.loss_fn( student_outputs['pred'][:,1], y['masks'] )
- segloss_s_decoder_3 = self.loss_fn( student_outputs['pred'][:,2], y['masks'] )
- student_seg_loss = segloss_s_decoder_1 + 0.5*segloss_s_decoder_2 + 0.25*segloss_s_decoder_3
- # KD loss between prediction layers -> (BCE Loss on logits with DS)
- bce_loss_with_teacher = 0
- decoder_weights = [1, 0.5, 0.25]
- for decoder_idx, weight in enumerate(decoder_weights):
- teacher_logits = teacher_outputs['pred'][:, decoder_idx] / self.temperature # Shape: (B, 4, 128, 128, 128)
- student_logits = student_outputs['pred'][:, decoder_idx] / self.temperature # Shape: (B, 4, 128, 128, 128)
- # Compute hard labels for WT, TC, ET
- teacher_probs = torch.sigmoid(teacher_logits)
- teacher_hard_labels = [
- (teacher_probs[:, channel_idx] > 0.5).float().unsqueeze(1) for channel_idx in range(1, 4)
- ]
- # Extract student logits for WT, TC, ET
- student_logits_channels = [
- student_logits[:, channel_idx].unsqueeze(1) for channel_idx in range(1, 4)
- ]
- # Compute BCE losses for WT, TC, ET
- bce_losses = [
- self.bce_loss(student_logits_channels[i], teacher_hard_labels[i]) for i in range(3)
- ]
- # Aggregate loss for this decoder with the corresponding weight
- bce_loss_with_teacher += sum(bce_losses) * (self.temperature ** 2) * weight
- #-----------------------------------------------------------------------------------#
- alpha, zeta = 1.0, 0.1
- print("Seg loss: ", student_seg_loss)
- print("BCE loss with teacher: ", bce_loss_with_teacher)
- print("Seg loss weighted: ", alpha*student_seg_loss)
- print("BCE loss with teacher weighted: ", zeta*bce_loss_with_teacher)
- batch_total_student_loss = alpha*student_seg_loss + zeta*bce_loss_with_teacher
- print("-------------Final student loss-------------")
- print(batch_total_student_loss)
- print("-------------Final student loss-------------")
- KD_output = {
- 'batch_total_student_loss' : batch_total_student_loss,
- 'seg_weighted' : alpha*student_seg_loss,
- 'bce_weighted' : zeta*bce_loss_with_teacher,
- }
- return KD_output
- # %% [markdown]
- # # Training & Validation
- # %%
- def evaluate(model, loader, epoch, task):
- torch.manual_seed(0)
- model.eval()
- loss_fn = LossFunction()
- n_val_batches = len(loader)
- tumors_val_losses, running_loss = validate_model(model, loader, loss_fn)
- epoch_val_loss = running_loss / n_val_batches
- log_val_epoch_losses(tumors_val_losses, epoch, task, epoch_val_loss)
- print(f"------Final validation dice loss after epoch {epoch + 1}: {epoch_val_loss}-------")
- model.student.to('cuda:1')
- model.train()
- return epoch_val_loss
- def validate_model(model, loader, loss_fn):
- tumors_val_losses = {'GLI': [], 'PED': [], 'SSA': [], 'MEN':[], 'MET':[]}
- running_loss = 0
- n_val_batches = len(loader)
- with tqdm(total=n_val_batches, desc='Validating', unit='batch', leave=False) as pbar:
- with torch.no_grad():
- for y in loader:
- val_loss, data_type = process_batch(model, y, loss_fn)
- tumors_val_losses[data_type].append(val_loss.item())
- running_loss += val_loss
- pbar.update(1)
- return tumors_val_losses, running_loss
- def process_batch(model, y, loss_fn):
- y['imgs'], y['masks'] = y['imgs'].to('cuda'), y['masks'].to('cuda')
- data_type = y['data_type'][0]
- with torch.amp.autocast('cuda'):
- output = model.student.to('cuda')(y['imgs'])
- val_loss = loss_fn(output['pred'], y['masks'])
- print(f"Validation dice loss per batch: {val_loss}")
- return val_loss, data_type
- def log_val_epoch_losses(tumors_val_losses, epoch, task, epoch_val_loss):
- for tumor_type, losses in tumors_val_losses.items():
- avg_loss = sum(losses) / len(losses) if losses else 0
- task.get_logger().report_scalar(
- title=f"{tumor_type} Losses over Epochs",
- series=f"{tumor_type} Epoch valLoss",
- iteration=epoch + 1,
- value=avg_loss
- )
- task.get_logger().report_scalar("KD Losses over Epochs", "val_loss", iteration=epoch+1, value=epoch_val_loss)
- # %%
- def setup_environment(args):
- torch.manual_seed(0)
- args['out_checkpoint_dir'].mkdir(parents=True, exist_ok=True)
- def initialize_models():
- teacher_model = DynUNet(spatial_dims=3, in_channels=4, out_channels=4, deep_supervision=True, KD=True).to('cuda:0')
- student_model = Student_KD_loss().to('cuda:1')
- return teacher_model, student_model
- def initialize_optimizer_scheduler(student_model, args):
- optimizer = optim.AdamW(student_model.parameters(), lr=args['learning_rate'], weight_decay=args['weight_decay'], eps=1e-4)
- scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.1, patience=10, cooldown=1, threshold=0.001, min_lr=1e-6)
- return optimizer, scheduler
- def load_teacher_model(teacher_model, data_type, teacher_model_paths):
- teacher_model_path = teacher_model_paths.get(data_type)
- if teacher_model_path and Path(teacher_model_path).is_file():
- ckpt = torch.load(teacher_model_path, map_location='cuda:0', weights_only=True)
- teacher_model.load_state_dict(ckpt['teacher_model'])
- print(f"Loaded model: {teacher_model_path}")
- def load_student_checkpoint(student_model, optimizer, scaler, scheduler, args):
- checkpoint_path = args['in_checkpoint_dir'] / 'Student_model_after_epoch_58_trainLoss_0.8541_valLoss_0.3335.pth'
- if checkpoint_path.is_file():
- print(f"Found model {checkpoint_path}")
- ckpt = torch.load(checkpoint_path, map_location='cuda:1', weights_only=True)
- student_model.student.load_state_dict(ckpt['student_model'])
- optimizer.load_state_dict(ckpt['optimizer_student'])
- scaler.load_state_dict(ckpt['grad_scaler_state'])
- scheduler.load_state_dict(ckpt['scheduler_state_dict'])
- print(f"Loaded student model: {checkpoint_path} with lr: {optimizer.param_groups[0]['lr']}")
- return ckpt['epoch'] + 1
- return 0
- def train_epoch(epoch, trainLoader, train_config, start_ep):
- student_model = train_config['student_model']
- teacher_model = train_config['teacher_model']
- optimizer = train_config['optimizer']
- scaler = train_config['scaler']
- accumulation_steps = train_config['accumulation_steps']
- teacher_model_paths = train_config['teacher_model_paths']
- task = train_config['task']
- student_model.train()
- teacher_model.eval()
- epoch_losses = {'total': 0, 'seg': 0, 'bce': 0}
- tumors_losses = {'GLI': [], 'PED': [], 'SSA': [], 'MEN': [], 'MET': []}
- with tqdm(total=len(trainLoader), desc=f"(Epoch {epoch + 1}/{start_ep + train_config['epochs']})", unit='batch') as pbar:
- optimizer.zero_grad()
- for step, y in enumerate(trainLoader):
- batch_loss = 0
- for sub_step, data_type in enumerate(y['data_type']):
- imgs = y['imgs'][sub_step].unsqueeze(0).to('cuda:0')
- masks = y['masks'][sub_step].unsqueeze(0).to('cuda:0')
- load_teacher_model(teacher_model, data_type, teacher_model_paths)
- with amp.autocast('cuda:0'):
- teacher_outputs = teacher_model(imgs)
- detached_teacher_output = {k: v.detach().to('cuda:1') for k, v in teacher_outputs.items()}
- imgs, masks = imgs.to('cuda:1'), masks.to('cuda:1')
- with amp.autocast('cuda:1'):
- student_outputs = student_model(detached_teacher_output, {'imgs': imgs, 'masks': masks})
- loss = (student_outputs['batch_total_student_loss'] / accumulation_steps)
- batch_loss += loss.item()
- tumors_losses[data_type].append(loss.item())
- scaler.scale(loss).backward()
- task.get_logger().report_scalar(
- title=f"Tumors training losses per epoch {epoch+1}",
- series=f"{data_type} loss",
- iteration=len(tumors_losses[data_type]),
- value=float(loss.item())
- )
- for key in epoch_losses:
- if key != 'total':
- epoch_losses[key] += (student_outputs.get(f'{key}_weighted', 0) / accumulation_steps)
- if (sub_step + 1) % accumulation_steps == 0 or (sub_step + 1) == len(y['data_type']):
- scaler.step(optimizer)
- scaler.update()
- optimizer.zero_grad()
- epoch_losses['total'] += batch_loss
- pbar.update(1)
- for key in epoch_losses:
- epoch_losses[key] /= len(trainLoader)
- return epoch_losses, tumors_losses
- def log_KD_losses_over_epochs(epoch, epoch_losses, tumors_losses, task):
- for loss_type, val in epoch_losses.items():
- task.get_logger().report_scalar(
- title="KD Losses over Epochs",
- series=f"{loss_type} loss",
- iteration=epoch + 1,
- value=epoch_losses[loss_type]
- )
- for tumor_type, losses in tumors_losses.items():
- task.get_logger().report_scalar(
- title=f"{tumor_type} Losses over Epochs",
- series=f"{tumor_type} Epoch trainLoss",
- iteration=epoch + 1,
- value=sum(losses) / len(losses) if losses else 0
- )
- def validate_and_save(epoch, valLoader, train_config, epoch_losses):
- student_model = train_config['student_model']
- scheduler = train_config['scheduler']
- optimizer = train_config['optimizer']
- scaler = train_config['scaler']
- out_checkpoint_dir = train_config['out_checkpoint_dir']
- task = train_config['task']
- val_loss = evaluate(student_model, valLoader, epoch, task)
- scheduler.step(val_loss)
- task.get_logger().report_scalar("LR", "learning_rate", iteration=epoch+1, value=optimizer.param_groups[0]['lr'])
- print(f"Learning rate after epoch {epoch + 1}: {optimizer.param_groups[0]['lr']}")
- state = {
- 'epoch': epoch,
- 'student_model': student_model.student.state_dict(),
- 'optimizer_student': optimizer.state_dict(),
- 'lr': optimizer.param_groups[0]['lr'],
- 'grad_scaler_state': scaler.state_dict(),
- 'scheduler_state_dict': scheduler.state_dict(),
- 'val_dice_loss': val_loss
- }
- checkpoint_path = out_checkpoint_dir / f'Student_model_after_epoch_{epoch + 1}_trainLoss_{epoch_losses["total"]:.4f}_valLoss_{val_loss:.4f}.pth'
- torch.save(state, checkpoint_path)
- print(f"Model saved after epoch {epoch + 1}")
- def run_KD(trainLoader, valLoader, args):
- setup_environment(args)
- teacher_model, student_model = initialize_models()
- optimizer, scheduler = initialize_optimizer_scheduler(student_model, args)
- scaler = amp.GradScaler('cuda:1')
- teacher_model_paths = {
- 'GLI': '/kaggle/input/gliomateachernewlabels/Teacher_model_after_epoch_100_trainLoss_0.5972_valLoss_0.3019.pth',
- 'SSA': '/kaggle/input/africanewlabels/Teacher_model_after_epoch_67_trainLoss_1.1080_valLoss_0.5561.pth',
- 'PED': '/kaggle/input/pednewlabel/Teacher_model_after_epoch_99_trainLoss_1.4512_valLoss_1.0042.pth',
- 'MEN': '/kaggle/input/meningiomateachernewlabels/Teacher_model_after_epoch_85_trainLoss_0.5824_valLoss_0.3318.pth',
- 'MET': '/kaggle/input/met-teacher-new-labels/Teacher_model_after_epoch_100_trainLoss_1.6278_valLoss_0.7199.pth'
- }
- start_epoch = load_student_checkpoint(student_model, optimizer, scaler, scheduler, args)
- 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)
- task.connect(args)
- task.add_tags(['ABL Study', 'BCE+SEG without KL', "Ahmed Pro"])
- print(f'''Starting Knowledge Distillation:
- Epochs: From {start_epoch + 1} to {start_epoch + args['epochs']}
- Batch size: 5 (effective through gradient accumulation)
- Learning rate: {args['learning_rate']}
- Training data coming from: {args['data_dirs']}
- ''')
- train_config = {
- 'teacher_model': teacher_model,
- 'student_model': student_model,
- 'optimizer': optimizer,
- 'scheduler': scheduler,
- 'scaler': scaler,
- 'accumulation_steps': 5,
- 'teacher_model_paths': teacher_model_paths,
- 'out_checkpoint_dir': args['out_checkpoint_dir'],
- 'task': task,
- 'epochs': args['epochs']
- }
- for epoch in range(start_epoch, start_epoch + args['epochs']):
- epoch_losses, tumors_losses = train_epoch(epoch, trainLoader, train_config, start_epoch)
- log_KD_losses_over_epochs(epoch, epoch_losses, tumors_losses, task)
- validate_and_save(epoch, valLoader, train_config, epoch_losses)
- print("Training completed.")
- task.close()
- # %%
- args = {
- 'workers': 2,
- 'epochs': 3,
- 'train_batch_size': 5,
- 'val_batch_size': 2,
- 'test_batch_size': 1,
- 'learning_rate': 1e-3,
- 'weight_decay': 1e-5,
- 'lambd': 0.0051,
- 'data_dirs': ["/kaggle/input/bratsglioma/Training/", "/kaggle/input/bratsafrica24/", "/kaggle/input/bratsped/Training/", "/kaggle/input/bratsmen/", "/kaggle/input/bratsmet24/"],
- 'in_checkpoint_dir': Path('/kaggle/input/data-abl-study-fairness-5-tumors-bce-seg-no-kl/'),
- 'out_checkpoint_dir': Path('/kaggle/working/')
- }
- trainLoader, valLoader, testLoader = prepare_data_loaders(args)
- run_KD(trainLoader, valLoader, args)
- # %% [markdown]
- # # Press here
ABL_(BCE+SEG).ipynb at commit 248789f, no license · at the source
Overview
- Systems and Biomedical Engineering Department, Faculty of Engineering, Cairo University,Cairo, Egypt
- Department of Artificial Intelligence and Data Science, College of Artificial Intelligence Convergence, Sejong University,Seoul, Republic of Korea
- Department of Radiology, Massachusetts General Hospital, Athinoula A. Martinos Center for Biomedical Imaging, Harvard Medical School,Charlestown, MA USA
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
248789f547fba152cc4f61c1f4026871246c762a, 15 April 2026Availability: 1 check, the latest on 30 September 2026: the link answers
- 30 September 2026: the link answers
16 files
- AI/
config/ , Python, 123 lines, 1 matchdefault_config.py - AI/
data/ , Python, 125 linesloaders.py - AI/
data/ , Python, 182 lines, 1 matchmake_dataset.py - AI/
inference/ , Python, 55 lines, 1 matchpostprocess.py - AI/
inference/ , Python, 102 lines, 1 matchpredict.py - AI/
inference/ , Python, 87 lines, 1 matchpreprocess.py - AI/
kd_modules/ , Python, 122 lines, 1 matchcbam_attention.py - AI/
kd_modules/ , Python, 105 lines, 2 matchesframework.py - AI/
losses/ , Python, 68 lines, 2 matchesloss.py - AI/
models/ , Python, 189 lines, 1 matchblocks.py - AI/
models/ , Python, 196 lines, 1 matchdyn_unet.py - AI/
models/ , Python, 132 linesmodel_utils.py - AI/
notebooks/ , Jupyter, 1,068 lines, 3 matchesAblation/ ABL_(BCE+SEG).ipynb - AI/
notebooks/ , Jupyter, 1,096 lines, 1 matchAblation/ ABL_(KL+SEG).ipynb - AI/
notebooks/ , Jupyter, 1,052 lines, 1 matchAblation/ ABL_(KL_no_CBAM+SEG).ipy nb - AI/
notebooks/ , Jupyter, 1,032 linesAblation/ ABL_(SEG only).ipynb - repository limit reached (2,000 files or 30 MB): the rest is at the source (205 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
- synapse.org/
synapse , at Synapse; found in “Data availability”
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 points to a dataset: synapse.org/
synapse - it points to the authors' code: AhmeddEmad7/
Brain-Tumor-Segmentation -Advancing-Generalizabil ity
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://
BibTeX
@article{elzayat2026dist
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/
url = {https://
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/
VL - 16
IS - 1
SP - 12969
SN - 2045-2322
PB - Nature Publishing Group
DO - 10.1038/
UR - https://
LA - en
ER -
CSL-JSON
{
"id": "10.1038/
"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":
"volume": "16",
"issue": "1",
"page": "12969",
"DOI": "10.1038/
"PMID": "41807430",
"PMCID": "PMC13096204",
"ISSN": "2045-2322",
"publisher": "Nature Publishing Group",
"URL": "https://
"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 oneIn 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 reportsIn 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 surgeryIn 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 communicationsIn 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 healthIn 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 reportsIn 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 neuroscienceIn 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 oneIn 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.
Claim this paper
Correct its record
Say what each link of this record is, remove the ones that are not the paper's, add the ones that are missing. The correction becomes a new version of the record, in its Versions section.
Validate its tracing map
You validate the map as this page shows it: 1 repository of the authors' code, each at its verified commit and with its license, 16 scripts, and 17 matches between paragraphs and code (see the Code and Map sections). It then receives a DOI on Zenodo, with you (your ORCID iD) and OSCR as its creators; the code itself is not deposited.
The map's fingerprint: sha256:78f5259d11ce048c…
Add the badge to its README
The badge links the code to this page. Copy one of these into the README of the paper's code: only you decide where it goes, and nothing is changed for you.
Markdown
[, paste the snippet at the top, then “Commit changes…” and, to review it first, “Create a new branch and start a pull request”. You open the pull request; OSCR asks for no permission.
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.
