OSCR

Adaptive multi-stage domain unlearning for white-matter lesion segmentation.

Code ↔ Paper

3 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 3 matches
  1. [1] § Methods › nnU-Net baseline ↔ nnunetv2/training/nnUNetTrainer/nnUNetTrainer.py, lines 690–782 · score 0.83 · Gaussian blurring, Gaussian noise, spatial transformation, low resolution, gamma, brightness
  2. [2] § Methods › nnU-Net baseline ↔ nnunetv2/training/nnUNetTrainer/variants/data_augmentation/nnUNetTrainerNoMirroring.py, lines 99–162 · score 0.77 · Gaussian blurring, Gaussian noise, low resolution, gamma, brightness, simulation
  3. [3] § Methods › nnU-Net with multi-stage unlearning ↔ nnunetv2/training/nnUNetTrainer/variants/unlearning/models.py, lines 8–96 · score 0.56 · Conv blocks, feature map, flatten, stride, unlearning, domain

Paper

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

The paper is loaded when this pane is shown.

The authors' code

Python · 1,306 lines · 69 KB · Apache-2.0 · 1 match

  1. import inspect
  2. import multiprocessing
  3. import os
  4. import shutil
  5. import sys
  6. import warnings
  7. from copy import deepcopy
  8. from datetime import datetime
  9. from time import time, sleep
  10. from typing import Union, Tuple, List
  11. import numpy as np
  12. import torch
  13. from batchgenerators.dataloading.multi_threaded_augmenter import MultiThreadedAugmenter
  14. from batchgenerators.dataloading.nondet_multi_threaded_augmenter import NonDetMultiThreadedAugmenter
  15. from batchgenerators.dataloading.single_threaded_augmenter import SingleThreadedAugmenter
  16. from batchgenerators.transforms.abstract_transforms import AbstractTransform, Compose
  17. from batchgenerators.transforms.color_transforms import BrightnessMultiplicativeTransform, \
  18. ContrastAugmentationTransform, GammaTransform
  19. from batchgenerators.transforms.noise_transforms import GaussianNoiseTransform, GaussianBlurTransform
  20. from batchgenerators.transforms.resample_transforms import SimulateLowResolutionTransform
  21. from batchgenerators.transforms.spatial_transforms import SpatialTransform, MirrorTransform
  22. from batchgenerators.transforms.utility_transforms import RemoveLabelTransform, RenameTransform, NumpyToTensor
  23. from batchgenerators.utilities.file_and_folder_operations import join, load_json, isfile, save_json, maybe_mkdir_p, isdir
  24. from torch._dynamo import OptimizedModule
  25. from nnunetv2.configuration import ANISO_THRESHOLD, default_num_processes
  26. from nnunetv2.evaluation.evaluate_predictions import compute_metrics_on_folder
  27. from nnunetv2.inference.export_prediction import export_prediction_from_logits, resample_and_save
  28. from nnunetv2.inference.predict_from_raw_data import nnUNetPredictor
  29. from nnunetv2.inference.sliding_window_prediction import compute_gaussian
  30. from nnunetv2.paths import nnUNet_preprocessed, nnUNet_results
  31. from nnunetv2.training.data_augmentation.compute_initial_patch_size import get_patch_size
  32. from nnunetv2.training.data_augmentation.custom_transforms.cascade_transforms import MoveSegAsOneHotToData, \
  33. ApplyRandomBinaryOperatorTransform, RemoveRandomConnectedComponentFromOneHotEncodingTransform
  34. from nnunetv2.training.data_augmentation.custom_transforms.deep_supervision_donwsampling import \
  35. DownsampleSegForDSTransform2
  36. from nnunetv2.training.data_augmentation.custom_transforms.limited_length_multithreaded_augmenter import \
  37. LimitedLenWrapper
  38. from nnunetv2.training.data_augmentation.custom_transforms.masking import MaskTransform
  39. from nnunetv2.training.data_augmentation.custom_transforms.region_based_training import \
  40. ConvertSegmentationToRegionsTransform
  41. from nnunetv2.training.data_augmentation.custom_transforms.transforms_for_dummy_2d import Convert2DTo3DTransform, \
  42. Convert3DTo2DTransform
  43. from nnunetv2.training.dataloading.data_loader_2d import nnUNetDataLoader2D
  44. from nnunetv2.training.dataloading.data_loader_3d import nnUNetDataLoader3D
  45. from nnunetv2.training.dataloading.nnunet_dataset import nnUNetDataset
  46. from nnunetv2.training.dataloading.utils import get_case_identifiers, unpack_dataset
  47. from nnunetv2.training.logging.nnunet_logger import nnUNetLogger
  48. from nnunetv2.training.loss.compound_losses import DC_and_CE_loss, DC_and_BCE_loss
  49. from nnunetv2.training.loss.deep_supervision import DeepSupervisionWrapper
  50. from nnunetv2.training.loss.dice import get_tp_fp_fn_tn, MemoryEfficientSoftDiceLoss
  51. from nnunetv2.training.lr_scheduler.polylr import PolyLRScheduler
  52. from nnunetv2.utilities.collate_outputs import collate_outputs
  53. from nnunetv2.utilities.crossval_split import generate_crossval_split
  54. from nnunetv2.utilities.default_n_proc_DA import get_allowed_n_proc_DA
  55. from nnunetv2.utilities.file_path_utilities import check_workers_alive_and_busy
  56. from nnunetv2.utilities.get_network_from_plans import get_network_from_plans
  57. from nnunetv2.utilities.helpers import empty_cache, dummy_context
  58. from nnunetv2.utilities.label_handling.label_handling import convert_labelmap_to_one_hot, determine_num_input_channels
  59. from nnunetv2.utilities.plans_handling.plans_handler import PlansManager, ConfigurationManager
  60. from torch import autocast, nn
  61. from torch import distributed as dist
  62. from torch.cuda import device_count
  63. from torch.cuda.amp import GradScaler
  64. from torch.nn.parallel import DistributedDataParallel as DDP
  65. class nnUNetTrainer(object):
  66. def __init__(self, plans: dict, configuration: str, fold: int, dataset_json: dict, experiment_identifier: str = "", unpack_dataset: bool = True,
  67. device: torch.device = torch.device('cuda'), dir_can_exist: bool = False):
  68. # From https://grugbrain.dev/. Worth a read ya big brains ;-)
  69. # apex predator of grug is complexity
  70. # complexity bad
  71. # say again:
  72. # complexity very bad
  73. # you say now:
  74. # complexity very, very bad
  75. # given choice between complexity or one on one against t-rex, grug take t-rex: at least grug see t-rex
  76. # complexity is spirit demon that enter codebase through well-meaning but ultimately very clubbable non grug-brain developers and project managers who not fear complexity spirit demon or even know about sometime
  77. # one day code base understandable and grug can get work done, everything good!
  78. # next day impossible: complexity demon spirit has entered code and very dangerous situation!
  79. # OK OK I am guilty. But I tried.
  80. # https://www.osnews.com/images/comics/wtfm.jpg
  81. # https://i.pinimg.com/originals/26/b2/50/26b250a738ea4abc7a5af4d42ad93af0.jpg
  82. self.is_ddp = dist.is_available() and dist.is_initialized()
  83. self.local_rank = 0 if not self.is_ddp else dist.get_rank()
  84. self.device = device
  85. # print what device we are using
  86. if self.is_ddp: # implicitly it's clear that we use cuda in this case
  87. print(f"I am local rank {self.local_rank}. {device_count()} GPUs are available. The world size is "
  88. f"{dist.get_world_size()}."
  89. f"Setting device to {self.device}")
  90. self.device = torch.device(type='cuda', index=self.local_rank)
  91. else:
  92. if self.device.type == 'cuda':
  93. # we might want to let the user pick this but for now please pick the correct GPU with CUDA_VISIBLE_DEVICES=X
  94. self.device = torch.device(type='cuda', index=0)
  95. print(f"Using device: {self.device}")
  96. # loading and saving this class for continuing from checkpoint should not happen based on pickling. This
  97. # would also pickle the network etc. Bad, bad. Instead we just reinstantiate and then load the checkpoint we
  98. # need. So let's save the init args
  99. self.my_init_kwargs = {}
  100. for k in inspect.signature(self.__init__).parameters.keys():
  101. self.my_init_kwargs[k] = locals()[k]
  102. ### Saving all the init args into class variables for later access
  103. self.plans_manager = PlansManager(plans)
  104. self.configuration_manager = self.plans_manager.get_configuration(configuration)
  105. self.configuration_name = configuration
  106. self.dataset_json = dataset_json
  107. self.fold = fold
  108. self.unpack_dataset = unpack_dataset
  109. ### Setting all the folder names. We need to make sure things don't crash in case we are just running
  110. # inference and some of the folders may not be defined!
  111. self.preprocessed_dataset_folder_base = join(nnUNet_preprocessed, self.plans_manager.dataset_name) \
  112. if nnUNet_preprocessed is not None else None
  113. dir_name = f"{self.__class__.__name__}__{self.plans_manager.plans_name}__{configuration}"
  114. self.output_folder_base = join(nnUNet_results, self.plans_manager.dataset_name, dir_name) if nnUNet_results is not None else None
  115. # added experiment_identifier to make easier experiment numbering
  116. # experiment identifier is appended to the back of output_folder_name
  117. self.experiment_identifier = experiment_identifier
  118. if experiment_identifier:
  119. self.output_folder_base += f"__{experiment_identifier}"
  120. self.output_folder = join(self.output_folder_base, f'fold_{fold}')
  121. if isdir(self.output_folder) and not dir_can_exist:
  122. raise ValueError(f"Folder {dir_name} exist. Safety check to prevent overwriting of experiments")
  123. self.preprocessed_dataset_folder = join(self.preprocessed_dataset_folder_base,
  124. self.configuration_manager.data_identifier)
  125. # unlike the previous nnunet folder_with_segs_from_previous_stage is now part of the plans. For now it has to
  126. # be a different configuration in the same plans
  127. # IMPORTANT! the mapping must be bijective, so lowres must point to fullres and vice versa (using
  128. # "previous_stage" and "next_stage"). Otherwise it won't work!
  129. self.is_cascaded = self.configuration_manager.previous_stage_name is not None
  130. self.folder_with_segs_from_previous_stage = \
  131. join(nnUNet_results, self.plans_manager.dataset_name,
  132. self.__class__.__name__ + '__' + self.plans_manager.plans_name + "__" +
  133. self.configuration_manager.previous_stage_name, 'predicted_next_stage', self.configuration_name) \
  134. if self.is_cascaded else None
  135. ### Some hyperparameters for you to fiddle with
  136. self.initial_lr = 1e-2
  137. self.weight_decay = 3e-5
  138. self.oversample_foreground_percent = 0.33
  139. self.num_iterations_per_epoch = 250
  140. self.num_val_iterations_per_epoch = 50
  141. self.current_epoch = 0
  142. self.enable_deep_supervision = True
  143. self.num_epochs = self.configuration_manager.configuration.get('num_epochs', 500)
  144. ### Dealing with labels/regions
  145. self.label_manager = self.plans_manager.get_label_manager(dataset_json)
  146. # labels can either be a list of int (regular training) or a list of tuples of int (region-based training)
  147. # needed for predictions. We do sigmoid in case of (overlapping) regions
  148. self.num_input_channels = None # -> self.initialize()
  149. self.network = None # -> self.build_network_architecture()
  150. self.optimizer = self.lr_scheduler = None # -> self.initialize
  151. self.grad_scaler = GradScaler() if self.device.type == 'cuda' else None
  152. self.loss = None # -> self.initialize
  153. ### Simple logging. Don't take that away from me!
  154. # initialize log file. This is just our log for the print statements etc. Not to be confused with lightning
  155. # logging
  156. timestamp = datetime.now()
  157. maybe_mkdir_p(self.output_folder)
  158. self.log_file = join(self.output_folder, "training_log_%d_%d_%d_%02.0d_%02.0d_%02.0d.txt" %
  159. (timestamp.year, timestamp.month, timestamp.day, timestamp.hour, timestamp.minute,
  160. timestamp.second))
  161. self.logger = nnUNetLogger()
  162. ### placeholders
  163. self.dataloader_train = self.dataloader_val = None # see on_train_start
  164. ### initializing stuff for remembering things and such
  165. self._best_ema = None
  166. ### inference things
  167. self.inference_allowed_mirroring_axes = None # this variable is set in
  168. # self.configure_rotation_dummyDA_mirroring_and_inital_patch_size and will be saved in checkpoints
  169. ### checkpoint saving stuff
  170. self.save_every = 50
  171. self.disable_checkpointing = False
  172. ## DDP batch size and oversampling can differ between workers and needs adaptation
  173. # we need to change the batch size in DDP because we don't use any of those distributed samplers
  174. self._set_batch_size_and_oversample()
  175. self.was_initialized = False
  176. self.print_to_log_file("\n#######################################################################\n"
  177. "Please cite the following paper when using nnU-Net:\n"
  178. "Isensee, F., Jaeger, P. F., Kohl, S. A., Petersen, J., & Maier-Hein, K. H. (2021). "
  179. "nnU-Net: a self-configuring method for deep learning-based biomedical image segmentation. "
  180. "Nature methods, 18(2), 203-211.\n"
  181. "#######################################################################\n",
  182. also_print_to_console=True, add_timestamp=False)
  183. def initialize(self):
  184. if not self.was_initialized:
  185. self.num_input_channels = determine_num_input_channels(self.plans_manager, self.configuration_manager,
  186. self.dataset_json)
  187. self.network = self.build_network_architecture(
  188. self.plans_manager,
  189. self.dataset_json,
  190. self.configuration_manager,
  191. self.num_input_channels,
  192. self.enable_deep_supervision,
  193. ).to(self.device)
  194. # compile network for free speedup
  195. if self._do_i_compile():
  196. self.print_to_log_file('Using torch.compile...')
  197. self.network = torch.compile(self.network)
  198. self.optimizer, self.lr_scheduler = self.configure_optimizers()
  199. # if ddp, wrap in DDP wrapper
  200. if self.is_ddp:
  201. self.network = torch.nn.SyncBatchNorm.convert_sync_batchnorm(self.network)
  202. self.network = DDP(self.network, device_ids=[self.local_rank])
  203. self.loss = self._build_loss()
  204. self.was_initialized = True
  205. else:
  206. raise RuntimeError("You have called self.initialize even though the trainer was already initialized. "
  207. "That should not happen.")
  208. def _do_i_compile(self):
  209. return ('nnUNet_compile' in os.environ.keys()) and (os.environ['nnUNet_compile'].lower() in ('true', '1', 't'))
  210. def _save_debug_information(self):
  211. # saving some debug information
  212. if self.local_rank == 0:
  213. dct = {}
  214. for k in self.__dir__():
  215. if not k.startswith("__"):
  216. if not callable(getattr(self, k)) or k in ['loss', ]:
  217. dct[k] = str(getattr(self, k))
  218. elif k in ['network', ]:
  219. dct[k] = str(getattr(self, k).__class__.__name__)
  220. else:
  221. # print(k)
  222. pass
  223. if k in ['dataloader_train', 'dataloader_val']:
  224. if hasattr(getattr(self, k), 'generator'):
  225. dct[k + '.generator'] = str(getattr(self, k).generator)
  226. if hasattr(getattr(self, k), 'num_processes'):
  227. dct[k + '.num_processes'] = str(getattr(self, k).num_processes)
  228. if hasattr(getattr(self, k), 'transform'):
  229. dct[k + '.transform'] = str(getattr(self, k).transform)
  230. import subprocess
  231. hostname = subprocess.getoutput(['hostname'])
  232. dct['hostname'] = hostname
  233. torch_version = torch.__version__
  234. if self.device.type == 'cuda':
  235. gpu_name = torch.cuda.get_device_name()
  236. dct['gpu_name'] = gpu_name
  237. cudnn_version = torch.backends.cudnn.version()
  238. else:
  239. cudnn_version = 'None'
  240. dct['device'] = str(self.device)
  241. dct['torch_version'] = torch_version
  242. dct['cudnn_version'] = cudnn_version
  243. save_json(dct, join(self.output_folder, "debug.json"))
  244. @staticmethod
  245. def build_network_architecture(plans_manager: PlansManager,
  246. dataset_json,
  247. configuration_manager: ConfigurationManager,
  248. num_input_channels,
  249. enable_deep_supervision: bool = True) -> nn.Module:
  250. """
  251. This is where you build the architecture according to the plans. There is no obligation to use
  252. get_network_from_plans, this is just a utility we use for the nnU-Net default architectures. You can do what
  253. you want. Even ignore the plans and just return something static (as long as it can process the requested
  254. patch size)
  255. but don't bug us with your bugs arising from fiddling with this :-P
  256. This is the function that is called in inference as well! This is needed so that all network architecture
  257. variants can be loaded at inference time (inference will use the same nnUNetTrainer that was used for
  258. training, so if you change the network architecture during training by deriving a new trainer class then
  259. inference will know about it).
  260. If you need to know how many segmentation outputs your custom architecture needs to have, use the following snippet:
  261. > label_manager = plans_manager.get_label_manager(dataset_json)
  262. > label_manager.num_segmentation_heads
  263. (why so complicated? -> We can have either classical training (classes) or regions. If we have regions,
  264. the number of outputs is != the number of classes. Also there is the ignore label for which no output
  265. should be generated. label_manager takes care of all that for you.)
  266. """
  267. return get_network_from_plans(plans_manager, dataset_json, configuration_manager,
  268. num_input_channels, deep_supervision=enable_deep_supervision)
  269. def _get_deep_supervision_scales(self):
  270. if self.enable_deep_supervision:
  271. deep_supervision_scales = list(list(i) for i in 1 / np.cumprod(np.vstack(
  272. self.configuration_manager.pool_op_kernel_sizes), axis=0))[:-1]
  273. else:
  274. deep_supervision_scales = None # for train and val_transforms
  275. return deep_supervision_scales
  276. def _set_batch_size_and_oversample(self):
  277. if not self.is_ddp:
  278. # set batch size to what the plan says, leave oversample untouched
  279. self.batch_size = self.configuration_manager.batch_size
  280. else:
  281. # batch size is distributed over DDP workers and we need to change oversample_percent for each worker
  282. world_size = dist.get_world_size()
  283. my_rank = dist.get_rank()
  284. global_batch_size = self.configuration_manager.batch_size
  285. assert global_batch_size >= world_size, 'Cannot run DDP if the batch size is smaller than the number of ' \
  286. 'GPUs... Duh.'
  287. batch_size_per_GPU = [global_batch_size // world_size] * world_size
  288. batch_size_per_GPU = [batch_size_per_GPU[i] + 1
  289. if (batch_size_per_GPU[i] * world_size + i) < global_batch_size
  290. else batch_size_per_GPU[i]
  291. for i in range(len(batch_size_per_GPU))]
  292. assert sum(batch_size_per_GPU) == global_batch_size
  293. sample_id_low = 0 if my_rank == 0 else np.sum(batch_size_per_GPU[:my_rank])
  294. sample_id_high = np.sum(batch_size_per_GPU[:my_rank + 1])
  295. # This is how oversampling is determined in DataLoader
  296. # round(self.batch_size * (1 - self.oversample_foreground_percent))
  297. # We need to use the same scheme here because an oversample of 0.33 with a batch size of 2 will be rounded
  298. # to an oversample of 0.5 (1 sample random, one oversampled). This may get lost if we just numerically
  299. # compute oversample
  300. oversample = [True if not i < round(global_batch_size * (1 - self.oversample_foreground_percent)) else False
  301. for i in range(global_batch_size)]
  302. if sample_id_high / global_batch_size < (1 - self.oversample_foreground_percent):
  303. oversample_percent = 0.0
  304. elif sample_id_low / global_batch_size > (1 - self.oversample_foreground_percent):
  305. oversample_percent = 1.0
  306. else:
  307. oversample_percent = sum(oversample[sample_id_low:sample_id_high]) / batch_size_per_GPU[my_rank]
  308. print("worker", my_rank, "oversample", oversample_percent)
  309. print("worker", my_rank, "batch_size", batch_size_per_GPU[my_rank])
  310. # self.print_to_log_file("worker", my_rank, "oversample", oversample_percents[my_rank])
  311. # self.print_to_log_file("worker", my_rank, "batch_size", batch_sizes[my_rank])
  312. self.batch_size = batch_size_per_GPU[my_rank]
  313. self.oversample_foreground_percent = oversample_percent
  314. def _build_loss(self):
  315. if self.label_manager.has_regions:
  316. loss = DC_and_BCE_loss({},
  317. {'batch_dice': self.configuration_manager.batch_dice,
  318. 'do_bg': True, 'smooth': 1e-5, 'ddp': self.is_ddp},
  319. use_ignore_label=self.label_manager.ignore_label is not None,
  320. dice_class=MemoryEfficientSoftDiceLoss)
  321. else:
  322. loss = DC_and_CE_loss({'batch_dice': self.configuration_manager.batch_dice,
  323. 'smooth': 1e-5, 'do_bg': False, 'ddp': self.is_ddp}, {}, weight_ce=1, weight_dice=1,
  324. ignore_label=self.label_manager.ignore_label, dice_class=MemoryEfficientSoftDiceLoss)
  325. # we give each output a weight which decreases exponentially (division by 2) as the resolution decreases
  326. # this gives higher resolution outputs more weight in the loss
  327. if self.enable_deep_supervision:
  328. deep_supervision_scales = self._get_deep_supervision_scales()
  329. weights = np.array([1 / (2**i) for i in range(len(deep_supervision_scales))])
  330. if self.is_ddp and not self._do_i_compile():
  331. # very strange and stupid interaction. DDP crashes and complains about unused parameters due to
  332. # weights[-1] = 0. Interestingly this crash doesn't happen with torch.compile enabled. Strange stuff.
  333. # Anywho, the simple fix is to set a very low weight to this.
  334. weights[-1] = 1e-6
  335. else:
  336. weights[-1] = 0
  337. # we don't use the lowest 2 outputs. Normalize weights so that they sum to 1
  338. weights = weights / weights.sum()
  339. # now wrap the loss
  340. loss = DeepSupervisionWrapper(loss, weights)
  341. return loss
  342. def configure_rotation_dummyDA_mirroring_and_inital_patch_size(self):
  343. """
  344. This function is stupid and certainly one of the weakest spots of this implementation. Not entirely sure how we can fix it.
  345. """
  346. patch_size = self.configuration_manager.patch_size
  347. dim = len(patch_size)
  348. # todo rotation should be defined dynamically based on patch size (more isotropic patch sizes = more rotation)
  349. if dim == 2:
  350. do_dummy_2d_data_aug = False
  351. # todo revisit this parametrization
  352. if max(patch_size) / min(patch_size) > 1.5:
  353. rotation_for_DA = {
  354. 'x': (-15. / 360 * 2. * np.pi, 15. / 360 * 2. * np.pi),
  355. 'y': (0, 0),
  356. 'z': (0, 0)
  357. }
  358. else:
  359. rotation_for_DA = {
  360. 'x': (-180. / 360 * 2. * np.pi, 180. / 360 * 2. * np.pi),
  361. 'y': (0, 0),
  362. 'z': (0, 0)
  363. }
  364. mirror_axes = (0, 1)
  365. elif dim == 3:
  366. # todo this is not ideal. We could also have patch_size (64, 16, 128) in which case a full 180deg 2d rot would be bad
  367. # order of the axes is determined by spacing, not image size
  368. do_dummy_2d_data_aug = (max(patch_size) / patch_size[0]) > ANISO_THRESHOLD
  369. if do_dummy_2d_data_aug:
  370. # why do we rotate 180 deg here all the time? We should also restrict it
  371. rotation_for_DA = {
  372. 'x': (-180. / 360 * 2. * np.pi, 180. / 360 * 2. * np.pi),
  373. 'y': (0, 0),
  374. 'z': (0, 0)
  375. }
  376. else:
  377. rotation_for_DA = {
  378. 'x': (-30. / 360 * 2. * np.pi, 30. / 360 * 2. * np.pi),
  379. 'y': (-30. / 360 * 2. * np.pi, 30. / 360 * 2. * np.pi),
  380. 'z': (-30. / 360 * 2. * np.pi, 30. / 360 * 2. * np.pi),
  381. }
  382. mirror_axes = (0, 1, 2)
  383. else:
  384. raise RuntimeError()
  385. # todo this function is stupid. It doesn't even use the correct scale range (we keep things as they were in the
  386. # old nnunet for now)
  387. initial_patch_size = get_patch_size(patch_size[-dim:],
  388. *rotation_for_DA.values(),
  389. (0.85, 1.25))
  390. if do_dummy_2d_data_aug:
  391. initial_patch_size[0] = patch_size[0]
  392. self.print_to_log_file(f'do_dummy_2d_data_aug: {do_dummy_2d_data_aug}')
  393. self.inference_allowed_mirroring_axes = mirror_axes
  394. return rotation_for_DA, do_dummy_2d_data_aug, initial_patch_size, mirror_axes
  395. def print_to_log_file(self, *args, also_print_to_console=True, add_timestamp=True):
  396. if self.local_rank == 0:
  397. timestamp = time()
  398. dt_object = datetime.fromtimestamp(timestamp)
  399. if add_timestamp:
  400. args = (f"{dt_object}:", *args)
  401. successful = False
  402. max_attempts = 5
  403. ctr = 0
  404. while not successful and ctr < max_attempts:
  405. try:
  406. with open(self.log_file, 'a+') as f:
  407. for a in args:
  408. f.write(str(a))
  409. f.write(" ")
  410. f.write("\n")
  411. successful = True
  412. except IOError:
  413. print(f"{datetime.fromtimestamp(timestamp)}: failed to log: ", sys.exc_info())
  414. sleep(0.5)
  415. ctr += 1
  416. if also_print_to_console:
  417. print(*args)
  418. elif also_print_to_console:
  419. print(*args)
  420. def print_plans(self):
  421. if self.local_rank == 0:
  422. dct = deepcopy(self.plans_manager.plans)
  423. del dct['configurations']
  424. self.print_to_log_file(f"\nThis is the configuration used by this "
  425. f"training:\nConfiguration name: {self.configuration_name}\n",
  426. self.configuration_manager, '\n', add_timestamp=False)
  427. self.print_to_log_file('These are the global plan.json settings:\n', dct, '\n', add_timestamp=False)
  428. def configure_optimizers(self):
  429. optimizer = torch.optim.SGD(self.network.parameters(), self.initial_lr, weight_decay=self.weight_decay,
  430. momentum=0.99, nesterov=True)
  431. lr_scheduler = PolyLRScheduler(
  432. optimizer,
  433. self.initial_lr,
  434. 1000,
  435. # self.num_epochs
  436. )
  437. return optimizer, lr_scheduler
  438. def plot_network_architecture(self):
  439. if self._do_i_compile():
  440. self.print_to_log_file("Unable to plot network architecture: nnUNet_compile is enabled!")
  441. return
  442. if self.local_rank == 0:
  443. try:
  444. # raise NotImplementedError('hiddenlayer no longer works and we do not have a viable alternative :-(')
  445. # pip install git+https://github.com/saugatkandel/hiddenlayer.git
  446. # from torchviz import make_dot
  447. # # not viable.
  448. # make_dot(tuple(self.network(torch.rand((1, self.num_input_channels,
  449. # *self.configuration_manager.patch_size),
  450. # device=self.device)))).render(
  451. # join(self.output_folder, "network_architecture.pdf"), format='pdf')
  452. # self.optimizer.zero_grad()
  453. # broken.
  454. import hiddenlayer as hl
  455. g = hl.build_graph(self.network,
  456. torch.rand((1, self.num_input_channels,
  457. *self.configuration_manager.patch_size),
  458. device=self.device),
  459. transforms=None)
  460. g.save(join(self.output_folder, "network_architecture.pdf"))
  461. del g
  462. except Exception as e:
  463. self.print_to_log_file("Unable to plot network architecture:")
  464. self.print_to_log_file(e)
  465. # self.print_to_log_file("\nprinting the network instead:\n")
  466. # self.print_to_log_file(self.network)
  467. # self.print_to_log_file("\n")
  468. finally:
  469. empty_cache(self.device)
  470. def do_split(self):
  471. """
  472. The default split is a 5 fold CV on all available training cases. nnU-Net will create a split (it is seeded,
  473. so always the same) and save it as splits_final.json file in the preprocessed data directory.
  474. Sometimes you may want to create your own split for various reasons. For this you will need to create your own
  475. splits_final.json file. If this file is present, nnU-Net is going to use it and whatever splits are defined in
  476. it. You can create as many splits in this file as you want. Note that if you define only 4 splits (fold 0-3)
  477. and then set fold=4 when training (that would be the fifth split), nnU-Net will print a warning and proceed to
  478. use a random 80:20 data split.
  479. :return:
  480. """
  481. if self.fold == "all":
  482. # if fold==all then we use all images for training and validation
  483. case_identifiers = get_case_identifiers(self.preprocessed_dataset_folder)
  484. tr_keys = case_identifiers
  485. val_keys = tr_keys
  486. else:
  487. splits_file = join(self.preprocessed_dataset_folder_base, "splits_final.json")
  488. dataset = nnUNetDataset(self.preprocessed_dataset_folder, case_identifiers=None,
  489. num_images_properties_loading_threshold=0,
  490. folder_with_segs_from_previous_stage=self.folder_with_segs_from_previous_stage)
  491. # if the split file does not exist we need to create it
  492. if not isfile(splits_file):
  493. self.print_to_log_file("Creating new 5-fold cross-validation split...")
  494. all_keys_sorted = list(np.sort(list(dataset.keys())))
  495. splits = generate_crossval_split(all_keys_sorted, seed=12345, n_splits=5)
  496. save_json(splits, splits_file)
  497. else:
  498. self.print_to_log_file("Using splits from existing split file:", splits_file)
  499. splits = load_json(splits_file)
  500. self.print_to_log_file(f"The split file contains {len(splits)} splits.")
  501. self.print_to_log_file("Desired fold for training: %d" % self.fold)
  502. if self.fold < len(splits):
  503. tr_keys = splits[self.fold]['train']
  504. val_keys = splits[self.fold]['val']
  505. self.print_to_log_file("This split has %d training and %d validation cases."
  506. % (len(tr_keys), len(val_keys)))
  507. else:
  508. self.print_to_log_file("INFO: You requested fold %d for training but splits "
  509. "contain only %d folds. I am now creating a "
  510. "random (but seeded) 80:20 split!" % (self.fold, len(splits)))
  511. # if we request a fold that is not in the split file, create a random 80:20 split
  512. rnd = np.random.RandomState(seed=12345 + self.fold)
  513. keys = np.sort(list(dataset.keys()))
  514. idx_tr = rnd.choice(len(keys), int(len(keys) * 0.8), replace=False)
  515. idx_val = [i for i in range(len(keys)) if i not in idx_tr]
  516. tr_keys = [keys[i] for i in idx_tr]
  517. val_keys = [keys[i] for i in idx_val]
  518. self.print_to_log_file("This random 80:20 split has %d training and %d validation cases."
  519. % (len(tr_keys), len(val_keys)))
  520. if any([i in val_keys for i in tr_keys]):
  521. self.print_to_log_file('WARNING: Some validation cases are also in the training set. Please check the '
  522. 'splits.json or ignore if this is intentional.')
  523. return tr_keys, val_keys
  524. def get_tr_and_val_datasets(self):
  525. # create dataset split
  526. tr_keys, val_keys = self.do_split()
  527. # load the datasets for training and validation. Note that we always draw random samples so we really don't
  528. # care about distributing training cases across GPUs.
  529. dataset_tr = nnUNetDataset(self.preprocessed_dataset_folder, tr_keys,
  530. folder_with_segs_from_previous_stage=self.folder_with_segs_from_previous_stage,
  531. num_images_properties_loading_threshold=0)
  532. dataset_val = nnUNetDataset(self.preprocessed_dataset_folder, val_keys,
  533. folder_with_segs_from_previous_stage=self.folder_with_segs_from_previous_stage,
  534. num_images_properties_loading_threshold=0)
  535. return dataset_tr, dataset_val
  536. def get_dataloaders(self):
  537. # we use the patch size to determine whether we need 2D or 3D dataloaders. We also use it to determine whether
  538. # we need to use dummy 2D augmentation (in case of 3D training) and what our initial patch size should be
  539. patch_size = self.configuration_manager.patch_size
  540. dim = len(patch_size)
  541. # needed for deep supervision: how much do we need to downscale the segmentation targets for the different
  542. # outputs?
  543. deep_supervision_scales = self._get_deep_supervision_scales()
  544. (
  545. rotation_for_DA,
  546. do_dummy_2d_data_aug,
  547. initial_patch_size,
  548. mirror_axes,
  549. ) = self.configure_rotation_dummyDA_mirroring_and_inital_patch_size()
  550. # training pipeline
  551. tr_transforms = self.get_training_transforms(
  552. patch_size, rotation_for_DA, deep_supervision_scales, mirror_axes, do_dummy_2d_data_aug,
  553. order_resampling_data=3, order_resampling_seg=1,
  554. use_mask_for_norm=self.configuration_manager.use_mask_for_norm,
  555. is_cascaded=self.is_cascaded, foreground_labels=self.label_manager.foreground_labels,
  556. regions=self.label_manager.foreground_regions if self.label_manager.has_regions else None,
  557. ignore_label=self.label_manager.ignore_label)
  558. # validation pipeline
  559. val_transforms = self.get_validation_transforms(deep_supervision_scales,
  560. is_cascaded=self.is_cascaded,
  561. foreground_labels=self.label_manager.foreground_labels,
  562. regions=self.label_manager.foreground_regions if
  563. self.label_manager.has_regions else None,
  564. ignore_label=self.label_manager.ignore_label)
  565. dl_tr, dl_val = self.get_plain_dataloaders(initial_patch_size, dim)
  566. allowed_num_processes = get_allowed_n_proc_DA()
  567. if allowed_num_processes == 0:
  568. mt_gen_train = SingleThreadedAugmenter(dl_tr, tr_transforms)
  569. mt_gen_val = SingleThreadedAugmenter(dl_val, val_transforms)
  570. else:
  571. mt_gen_train = LimitedLenWrapper(self.num_iterations_per_epoch, data_loader=dl_tr, transform=tr_transforms,
  572. num_processes=allowed_num_processes, num_cached=6, seeds=None,
  573. pin_memory=self.device.type == 'cuda', wait_time=0.02)
  574. mt_gen_val = LimitedLenWrapper(self.num_val_iterations_per_epoch, data_loader=dl_val,
  575. transform=val_transforms, num_processes=max(1, allowed_num_processes // 2),
  576. num_cached=3, seeds=None, pin_memory=self.device.type == 'cuda',
  577. wait_time=0.02)
  578. return mt_gen_train, mt_gen_val
  579. def get_plain_dataloaders(self, initial_patch_size: Tuple[int, ...], dim: int):
  580. dataset_tr, dataset_val = self.get_tr_and_val_datasets()
  581. if dim == 2:
  582. dl_tr = nnUNetDataLoader2D(dataset_tr, self.batch_size,
  583. initial_patch_size,
  584. self.configuration_manager.patch_size,
  585. self.label_manager,
  586. oversample_foreground_percent=self.oversample_foreground_percent,
  587. sampling_probabilities=None, pad_sides=None)
  588. dl_val = nnUNetDataLoader2D(dataset_val, self.batch_size,
  589. self.configuration_manager.patch_size,
  590. self.configuration_manager.patch_size,
  591. self.label_manager,
  592. oversample_foreground_percent=self.oversample_foreground_percent,
  593. sampling_probabilities=None, pad_sides=None)
  594. else:
  595. dl_tr = nnUNetDataLoader3D(dataset_tr, self.batch_size,
  596. initial_patch_size,
  597. self.configuration_manager.patch_size,
  598. self.label_manager,
  599. oversample_foreground_percent=self.oversample_foreground_percent,
  600. sampling_probabilities=None, pad_sides=None)
  601. dl_val = nnUNetDataLoader3D(dataset_val, self.batch_size,
  602. self.configuration_manager.patch_size,
  603. self.configuration_manager.patch_size,
  604. self.label_manager,
  605. oversample_foreground_percent=self.oversample_foreground_percent,
  606. sampling_probabilities=None, pad_sides=None)
  607. return dl_tr, dl_val
  608. @staticmethod
  609. def get_training_transforms(
  610. patch_size: Union[np.ndarray, Tuple[int]],
  611. rotation_for_DA: dict,
  612. deep_supervision_scales: Union[List, Tuple, None],
  613. mirror_axes: Tuple[int, ...],
  614. do_dummy_2d_data_aug: bool,
  615. order_resampling_data: int = 3,
  616. order_resampling_seg: int = 1,
  617. border_val_seg: int = -1,
  618. use_mask_for_norm: List[bool] = None,
  619. is_cascaded: bool = False,
  620. foreground_labels: Union[Tuple[int, ...], List[int]] = None,
  621. regions: List[Union[List[int], Tuple[int, ...], int]] = None,
  622. ignore_label: int = None,
  623. ) -> AbstractTransform:
  624. tr_transforms = []
  625. if do_dummy_2d_data_aug:
  626. ignore_axes = (0,)
  627. tr_transforms.append(Convert3DTo2DTransform())
  628. patch_size_spatial = patch_size[1:]
  629. else:
  630. patch_size_spatial = patch_size
  631. ignore_axes = None
  632. tr_transforms.append(SpatialTransform(
  633. patch_size_spatial, patch_center_dist_from_border=None,
  634. do_elastic_deform=False, alpha=(0, 0), sigma=(0, 0),
  635. do_rotation=True, angle_x=rotation_for_DA['x'], angle_y=rotation_for_DA['y'], angle_z=rotation_for_DA['z'],
  636. p_rot_per_axis=1, # todo experiment with this
  637. do_scale=True, scale=(0.7, 1.4),
  638. border_mode_data="constant", border_cval_data=0, order_data=order_resampling_data,
  639. border_mode_seg="constant", border_cval_seg=border_val_seg, order_seg=order_resampling_seg,
  640. random_crop=False, # random cropping is part of our dataloaders
  641. p_el_per_sample=0, p_scale_per_sample=0.2, p_rot_per_sample=0.2,
  642. independent_scale_for_each_axis=False # todo experiment with this
  643. ))
  644. if do_dummy_2d_data_aug:
  645. tr_transforms.append(Convert2DTo3DTransform())
  646. tr_transforms.append(GaussianNoiseTransform(p_per_sample=0.1))
  647. tr_transforms.append(GaussianBlurTransform((0.5, 1.), different_sigma_per_channel=True, p_per_sample=0.2,
  648. p_per_channel=0.5))
  649. tr_transforms.append(BrightnessMultiplicativeTransform(multiplier_range=(0.75, 1.25), p_per_sample=0.15))
  650. tr_transforms.append(ContrastAugmentationTransform(p_per_sample=0.15))
  651. tr_transforms.append(SimulateLowResolutionTransform(zoom_range=(0.5, 1), per_channel=True,
  652. p_per_channel=0.5,
  653. order_downsample=0, order_upsample=3, p_per_sample=0.25,
  654. ignore_axes=ignore_axes))
  655. tr_transforms.append(GammaTransform((0.7, 1.5), True, True, retain_stats=True, p_per_sample=0.1))
  656. tr_transforms.append(GammaTransform((0.7, 1.5), False, True, retain_stats=True, p_per_sample=0.3))
  657. if mirror_axes is not None and len(mirror_axes) > 0:
  658. tr_transforms.append(MirrorTransform(mirror_axes))
  659. if use_mask_for_norm is not None and any(use_mask_for_norm):
  660. tr_transforms.append(MaskTransform([i for i in range(len(use_mask_for_norm)) if use_mask_for_norm[i]],
  661. mask_idx_in_seg=0, set_outside_to=0))
  662. tr_transforms.append(RemoveLabelTransform(-1, 0))
  663. if is_cascaded:
  664. assert foreground_labels is not None, 'We need foreground_labels for cascade augmentations'
  665. tr_transforms.append(MoveSegAsOneHotToData(1, foreground_labels, 'seg', 'data'))
  666. tr_transforms.append(ApplyRandomBinaryOperatorTransform(
  667. channel_idx=list(range(-len(foreground_labels), 0)),
  668. p_per_sample=0.4,
  669. key="data",
  670. strel_size=(1, 8),
  671. p_per_label=1))
  672. tr_transforms.append(
  673. RemoveRandomConnectedComponentFromOneHotEncodingTransform(
  674. channel_idx=list(range(-len(foreground_labels), 0)),
  675. key="data",
  676. p_per_sample=0.2,
  677. fill_with_other_class_p=0,
  678. dont_do_if_covers_more_than_x_percent=0.15))
  679. tr_transforms.append(RenameTransform('seg', 'target', True))
  680. if regions is not None:
  681. # the ignore label must also be converted
  682. tr_transforms.append(ConvertSegmentationToRegionsTransform(list(regions) + [ignore_label]
  683. if ignore_label is not None else regions,
  684. 'target', 'target'))
  685. if deep_supervision_scales is not None:
  686. tr_transforms.append(DownsampleSegForDSTransform2(deep_supervision_scales, 0, input_key='target',
  687. output_key='target'))
  688. tr_transforms.append(NumpyToTensor(['data', 'target'], 'float'))
  689. tr_transforms = Compose(tr_transforms)
  690. return tr_transforms
  691. @staticmethod
  692. def get_validation_transforms(
  693. deep_supervision_scales: Union[List, Tuple, None],
  694. is_cascaded: bool = False,
  695. foreground_labels: Union[Tuple[int, ...], List[int]] = None,
  696. regions: List[Union[List[int], Tuple[int, ...], int]] = None,
  697. ignore_label: int = None,
  698. ) -> AbstractTransform:
  699. val_transforms = []
  700. val_transforms.append(RemoveLabelTransform(-1, 0))
  701. if is_cascaded:
  702. val_transforms.append(MoveSegAsOneHotToData(1, foreground_labels, 'seg', 'data'))
  703. val_transforms.append(RenameTransform('seg', 'target', True))
  704. if regions is not None:
  705. # the ignore label must also be converted
  706. val_transforms.append(ConvertSegmentationToRegionsTransform(list(regions) + [ignore_label]
  707. if ignore_label is not None else regions,
  708. 'target', 'target'))
  709. if deep_supervision_scales is not None:
  710. val_transforms.append(DownsampleSegForDSTransform2(deep_supervision_scales, 0, input_key='target',
  711. output_key='target'))
  712. val_transforms.append(NumpyToTensor(['data', 'target'], 'float'))
  713. val_transforms = Compose(val_transforms)
  714. return val_transforms
  715. def set_deep_supervision_enabled(self, enabled: bool):
  716. """
  717. This function is specific for the default architecture in nnU-Net. If you change the architecture, there are
  718. chances you need to change this as well!
  719. """
  720. if self.is_ddp:
  721. mod = self.network.module
  722. else:
  723. mod = self.network
  724. if isinstance(mod, OptimizedModule):
  725. mod = mod._orig_mod
  726. mod.decoder.deep_supervision = enabled
  727. def on_train_start(self):
  728. if not self.was_initialized:
  729. self.initialize()
  730. maybe_mkdir_p(self.output_folder)
  731. # make sure deep supervision is on in the network
  732. self.set_deep_supervision_enabled(self.enable_deep_supervision)
  733. self.print_plans()
  734. empty_cache(self.device)
  735. # maybe unpack
  736. if self.unpack_dataset and self.local_rank == 0:
  737. self.print_to_log_file('unpacking dataset...')
  738. unpack_dataset(self.preprocessed_dataset_folder, unpack_segmentation=True, overwrite_existing=False,
  739. num_processes=max(1, round(get_allowed_n_proc_DA() // 2)))
  740. self.print_to_log_file('unpacking done...')
  741. if self.is_ddp:
  742. dist.barrier()
  743. # dataloaders must be instantiated here because they need access to the training data which may not be present
  744. # when doing inference
  745. self.dataloader_train, self.dataloader_val = self.get_dataloaders()
  746. # copy plans and dataset.json so that they can be used for restoring everything we need for inference
  747. save_json(self.plans_manager.plans, join(self.output_folder_base, 'plans.json'), sort_keys=False)
  748. save_json(self.dataset_json, join(self.output_folder_base, 'dataset.json'), sort_keys=False)
  749. # we don't really need the fingerprint but its still handy to have it with the others
  750. shutil.copy(join(self.preprocessed_dataset_folder_base, 'dataset_fingerprint.json'),
  751. join(self.output_folder_base, 'dataset_fingerprint.json'))
  752. # produces a pdf in output folder
  753. self.plot_network_architecture()
  754. self._save_debug_information()
  755. # print(f"batch size: {self.batch_size}")
  756. # print(f"oversample: {self.oversample_foreground_percent}")
  757. def on_train_end(self):
  758. # dirty hack because on_epoch_end increments the epoch counter and this is executed afterwards.
  759. # This will lead to the wrong current epoch to be stored
  760. self.current_epoch -= 1
  761. self.save_checkpoint(join(self.output_folder, "checkpoint_final.pth"))
  762. self.current_epoch += 1
  763. # now we can delete latest
  764. if self.local_rank == 0 and isfile(join(self.output_folder, "checkpoint_latest.pth")):
  765. os.remove(join(self.output_folder, "checkpoint_latest.pth"))
  766. # shut down dataloaders
  767. old_stdout = sys.stdout
  768. with open(os.devnull, 'w') as f:
  769. sys.stdout = f
  770. if self.dataloader_train is not None and \
  771. isinstance(self.dataloader_train, (NonDetMultiThreadedAugmenter, MultiThreadedAugmenter)):
  772. self.dataloader_train._finish()
  773. if self.dataloader_val is not None and \
  774. isinstance(self.dataloader_train, (NonDetMultiThreadedAugmenter, MultiThreadedAugmenter)):
  775. self.dataloader_val._finish()
  776. sys.stdout = old_stdout
  777. empty_cache(self.device)
  778. self.print_to_log_file("Training done.")
  779. def on_train_epoch_start(self):
  780. self.network.train()
  781. self.lr_scheduler.step(self.current_epoch)
  782. self.print_to_log_file('')
  783. self.print_to_log_file(f'Epoch {self.current_epoch}')
  784. self.print_to_log_file(
  785. f"Current learning rate: {np.round(self.optimizer.param_groups[0]['lr'], decimals=5)}")
  786. # lrs are the same for all workers so we don't need to gather them in case of DDP training
  787. self.logger.log('lrs', self.optimizer.param_groups[0]['lr'], self.current_epoch)
  788. def train_step(self, batch: dict) -> dict:
  789. data = batch['data']
  790. target = batch['target']
  791. data = data.to(self.device, non_blocking=True)
  792. if isinstance(target, list):
  793. target = [i.to(self.device, non_blocking=True) for i in target]
  794. else:
  795. target = target.to(self.device, non_blocking=True)
  796. self.optimizer.zero_grad(set_to_none=True)
  797. # Autocast can be annoying
  798. # If the device_type is 'cpu' then it's slow as heck and needs to be disabled.
  799. # If the device_type is 'mps' then it will complain that mps is not implemented, even if enabled=False is set. Whyyyyyyy. (this is why we don't make use of enabled=False)
  800. # So autocast will only be active if we have a cuda device.
  801. with autocast(self.device.type, enabled=True) if self.device.type == 'cuda' else dummy_context():
  802. output = self.network(data)
  803. # del data
  804. l = self.loss(output, target)
  805. if self.grad_scaler is not None:
  806. self.grad_scaler.scale(l).backward()
  807. self.grad_scaler.unscale_(self.optimizer)
  808. torch.nn.utils.clip_grad_norm_(self.network.parameters(), 12)
  809. self.grad_scaler.step(self.optimizer)
  810. self.grad_scaler.update()
  811. else:
  812. l.backward()
  813. torch.nn.utils.clip_grad_norm_(self.network.parameters(), 12)
  814. self.optimizer.step()
  815. return {'loss': l.detach().cpu().numpy()}
  816. def on_train_epoch_end(self, train_outputs: List[dict]):
  817. outputs = collate_outputs(train_outputs)
  818. if self.is_ddp:
  819. losses_tr = [None for _ in range(dist.get_world_size())]
  820. dist.all_gather_object(losses_tr, outputs['loss'])
  821. loss_here = np.vstack(losses_tr).mean()
  822. else:
  823. loss_here = np.mean(outputs['loss'])
  824. self.logger.log('train_losses', loss_here, self.current_epoch)
  825. def on_validation_epoch_start(self):
  826. self.network.eval()
  827. def validation_step(self, batch: dict) -> dict:
  828. data = batch['data']
  829. target = batch['target']
  830. data = data.to(self.device, non_blocking=True)
  831. if isinstance(target, list):
  832. target = [i.to(self.device, non_blocking=True) for i in target]
  833. else:
  834. target = target.to(self.device, non_blocking=True)
  835. # Autocast can be annoying
  836. # If the device_type is 'cpu' then it's slow as heck and needs to be disabled.
  837. # If the device_type is 'mps' then it will complain that mps is not implemented, even if enabled=False is set. Whyyyyyyy. (this is why we don't make use of enabled=False)
  838. # So autocast will only be active if we have a cuda device.
  839. with autocast(self.device.type, enabled=True) if self.device.type == 'cuda' else dummy_context():
  840. output = self.network(data)
  841. del data
  842. l = self.loss(output, target)
  843. # we only need the output with the highest output resolution (if DS enabled)
  844. if self.enable_deep_supervision:
  845. output = output[0]
  846. target = target[0]
  847. # the following is needed for online evaluation. Fake dice (green line)
  848. axes = [0] + list(range(2, output.ndim))
  849. if self.label_manager.has_regions:
  850. predicted_segmentation_onehot = (torch.sigmoid(output) > 0.5).long()
  851. else:
  852. # no need for softmax
  853. output_seg = output.argmax(1)[:, None]
  854. predicted_segmentation_onehot = torch.zeros(output.shape, device=output.device, dtype=torch.float32)
  855. predicted_segmentation_onehot.scatter_(1, output_seg, 1)
  856. del output_seg
  857. if self.label_manager.has_ignore_label:
  858. if not self.label_manager.has_regions:
  859. mask = (target != self.label_manager.ignore_label).float()
  860. # CAREFUL that you don't rely on target after this line!
  861. target[target == self.label_manager.ignore_label] = 0
  862. else:
  863. mask = 1 - target[:, -1:]
  864. # CAREFUL that you don't rely on target after this line!
  865. target = target[:, :-1]
  866. else:
  867. mask = None
  868. tp, fp, fn, _ = get_tp_fp_fn_tn(predicted_segmentation_onehot, target, axes=axes, mask=mask)
  869. tp_hard = tp.detach().cpu().numpy()
  870. fp_hard = fp.detach().cpu().numpy()
  871. fn_hard = fn.detach().cpu().numpy()
  872. if not self.label_manager.has_regions:
  873. # if we train with regions all segmentation heads predict some kind of foreground. In conventional
  874. # (softmax training) there needs tobe one output for the background. We are not interested in the
  875. # background Dice
  876. # [1:] in order to remove background
  877. tp_hard = tp_hard[1:]
  878. fp_hard = fp_hard[1:]
  879. fn_hard = fn_hard[1:]
  880. return {'loss': l.detach().cpu().numpy(), 'tp_hard': tp_hard, 'fp_hard': fp_hard, 'fn_hard': fn_hard}
  881. def on_validation_epoch_end(self, val_outputs: List[dict]):
  882. outputs_collated = collate_outputs(val_outputs)
  883. tp = np.sum(outputs_collated['tp_hard'], 0)
  884. fp = np.sum(outputs_collated['fp_hard'], 0)
  885. fn = np.sum(outputs_collated['fn_hard'], 0)
  886. if self.is_ddp:
  887. world_size = dist.get_world_size()
  888. tps = [None for _ in range(world_size)]
  889. dist.all_gather_object(tps, tp)
  890. tp = np.vstack([i[None] for i in tps]).sum(0)
  891. fps = [None for _ in range(world_size)]
  892. dist.all_gather_object(fps, fp)
  893. fp = np.vstack([i[None] for i in fps]).sum(0)
  894. fns = [None for _ in range(world_size)]
  895. dist.all_gather_object(fns, fn)
  896. fn = np.vstack([i[None] for i in fns]).sum(0)
  897. losses_val = [None for _ in range(world_size)]
  898. dist.all_gather_object(losses_val, outputs_collated['loss'])
  899. loss_here = np.vstack(losses_val).mean()
  900. else:
  901. loss_here = np.mean(outputs_collated['loss'])
  902. global_dc_per_class = [i for i in [2 * i / (2 * i + j + k) for i, j, k in zip(tp, fp, fn)]]
  903. mean_fg_dice = np.nanmean(global_dc_per_class)
  904. self.logger.log('mean_fg_dice', mean_fg_dice, self.current_epoch)
  905. self.logger.log('dice_per_class_or_region', global_dc_per_class, self.current_epoch)
  906. self.logger.log('val_losses', loss_here, self.current_epoch)
  907. def on_epoch_start(self):
  908. self.logger.log('epoch_start_timestamps', time(), self.current_epoch)
  909. def on_epoch_end(self):
  910. self.logger.log('epoch_end_timestamps', time(), self.current_epoch)
  911. self.print_to_log_file('train_loss', np.round(self.logger.my_fantastic_logging['train_losses'][-1], decimals=4))
  912. self.print_to_log_file('val_loss', np.round(self.logger.my_fantastic_logging['val_losses'][-1], decimals=4))
  913. self.print_to_log_file('Pseudo dice', [np.round(i, decimals=4) for i in
  914. self.logger.my_fantastic_logging['dice_per_class_or_region'][-1]])
  915. self.print_to_log_file(
  916. f"Epoch time: {np.round(self.logger.my_fantastic_logging['epoch_end_timestamps'][-1] - self.logger.my_fantastic_logging['epoch_start_timestamps'][-1], decimals=2)} s")
  917. # handling periodic checkpointing
  918. current_epoch = self.current_epoch
  919. if (current_epoch + 1) % self.save_every == 0 and current_epoch != (self.num_epochs - 1):
  920. self.save_checkpoint(join(self.output_folder, 'checkpoint_latest.pth'))
  921. # handle 'best' checkpointing. ema_fg_dice is computed by the logger and can be accessed like this
  922. if self._best_ema is None or self.logger.my_fantastic_logging['ema_fg_dice'][-1] > self._best_ema:
  923. self._best_ema = self.logger.my_fantastic_logging['ema_fg_dice'][-1]
  924. self.print_to_log_file(f"Yayy! New best EMA pseudo Dice: {np.round(self._best_ema, decimals=4)}")
  925. self.save_checkpoint(join(self.output_folder, 'checkpoint_best.pth'))
  926. if self.local_rank == 0:
  927. self.logger.plot_progress_png(self.output_folder)
  928. self.current_epoch += 1
  929. def save_checkpoint(self, filename: str) -> None:
  930. if self.local_rank == 0:
  931. if not self.disable_checkpointing:
  932. if self.is_ddp:
  933. mod = self.network.module
  934. else:
  935. mod = self.network
  936. if isinstance(mod, OptimizedModule):
  937. mod = mod._orig_mod
  938. checkpoint = {
  939. 'network_weights': mod.state_dict(),
  940. 'optimizer_state': self.optimizer.state_dict(),
  941. 'grad_scaler_state': self.grad_scaler.state_dict() if self.grad_scaler is not None else None,
  942. 'logging': self.logger.get_checkpoint(),
  943. '_best_ema': self._best_ema,
  944. 'current_epoch': self.current_epoch + 1,
  945. 'init_args': self.my_init_kwargs,
  946. 'trainer_name': self.__class__.__name__,
  947. 'inference_allowed_mirroring_axes': self.inference_allowed_mirroring_axes,
  948. }
  949. torch.save(checkpoint, filename)
  950. else:
  951. self.print_to_log_file('No checkpoint written, checkpointing is disabled')
  952. def load_checkpoint(self, filename_or_checkpoint: Union[dict, str]) -> None:
  953. if not self.was_initialized:
  954. self.initialize()
  955. if isinstance(filename_or_checkpoint, str):
  956. checkpoint = torch.load(filename_or_checkpoint, map_location=self.device)
  957. # if state dict comes from nn.DataParallel but we use non-parallel model here then the state dict keys do not
  958. # match. Use heuristic to make it match
  959. new_state_dict = {}
  960. for k, value in checkpoint['network_weights'].items():
  961. key = k
  962. if key not in self.network.state_dict().keys() and key.startswith('module.'):
  963. key = key[7:]
  964. new_state_dict[key] = value
  965. self.my_init_kwargs = checkpoint['init_args']
  966. self.current_epoch = checkpoint['current_epoch']
  967. self.logger.load_checkpoint(checkpoint['logging'])
  968. self._best_ema = checkpoint['_best_ema']
  969. self.inference_allowed_mirroring_axes = checkpoint[
  970. 'inference_allowed_mirroring_axes'] if 'inference_allowed_mirroring_axes' in checkpoint.keys() else self.inference_allowed_mirroring_axes
  971. # messing with state dict naming schemes. Facepalm.
  972. if self.is_ddp:
  973. if isinstance(self.network.module, OptimizedModule):
  974. self.network.module._orig_mod.load_state_dict(new_state_dict)
  975. else:
  976. self.network.module.load_state_dict(new_state_dict)
  977. else:
  978. if isinstance(self.network, OptimizedModule):
  979. self.network._orig_mod.load_state_dict(new_state_dict)
  980. else:
  981. self.network.load_state_dict(new_state_dict)
  982. self.optimizer.load_state_dict(checkpoint['optimizer_state'])
  983. if self.grad_scaler is not None:
  984. if checkpoint['grad_scaler_state'] is not None:
  985. self.grad_scaler.load_state_dict(checkpoint['grad_scaler_state'])
  986. def perform_actual_validation(self, save_probabilities: bool = False):
  987. self.set_deep_supervision_enabled(False)
  988. self.network.eval()
  989. if self.is_ddp and self.batch_size == 1 and self.enable_deep_supervision and self._do_i_compile():
  990. self.print_to_log_file("WARNING! batch size is 1 during training and torch.compile is enabled. If you "
  991. "encounter crashes in validation then this is because torch.compile forgets "
  992. "to trigger a recompilation of the model with deep supervision disabled. "
  993. "This causes torch.flip to complain about getting a tuple as input. Just rerun the "
  994. "validation with --val (exactly the same as before) and then it will work. "
  995. "Why? Because --val triggers nnU-Net to ONLY run validation meaning that the first "
  996. "forward pass (where compile is triggered) already has deep supervision disabled. "
  997. "This is exactly what we need in perform_actual_validation")
  998. predictor = nnUNetPredictor(tile_step_size=0.5, use_gaussian=True, use_mirroring=True,
  999. perform_everything_on_device=True, device=self.device, verbose=False,
  1000. verbose_preprocessing=False, allow_tqdm=False)
  1001. predictor.manual_initialization(self.network, self.plans_manager, self.configuration_manager, None,
  1002. self.dataset_json, self.__class__.__name__,
  1003. self.inference_allowed_mirroring_axes)
  1004. with multiprocessing.get_context("spawn").Pool(default_num_processes) as segmentation_export_pool:
  1005. worker_list = [i for i in segmentation_export_pool._pool]
  1006. validation_output_folder = join(self.output_folder, 'validation')
  1007. maybe_mkdir_p(validation_output_folder)
  1008. # we cannot use self.get_tr_and_val_datasets() here because we might be DDP and then we have to distribute
  1009. # the validation keys across the workers.
  1010. _, val_keys = self.do_split()
  1011. if self.is_ddp:
  1012. last_barrier_at_idx = len(val_keys) // dist.get_world_size() - 1
  1013. val_keys = val_keys[self.local_rank:: dist.get_world_size()]
  1014. # we cannot just have barriers all over the place because the number of keys each GPU receives can be
  1015. # different
  1016. dataset_val = nnUNetDataset(self.preprocessed_dataset_folder, val_keys,
  1017. folder_with_segs_from_previous_stage=self.folder_with_segs_from_previous_stage,
  1018. num_images_properties_loading_threshold=0)
  1019. next_stages = self.configuration_manager.next_stage_names
  1020. if next_stages is not None:
  1021. _ = [maybe_mkdir_p(join(self.output_folder_base, 'predicted_next_stage', n)) for n in next_stages]
  1022. results = []
  1023. for i, k in enumerate(dataset_val.keys()):
  1024. proceed = not check_workers_alive_and_busy(segmentation_export_pool, worker_list, results,
  1025. allowed_num_queued=2)
  1026. while not proceed:
  1027. sleep(0.1)
  1028. proceed = not check_workers_alive_and_busy(segmentation_export_pool, worker_list, results,
  1029. allowed_num_queued=2)
  1030. self.print_to_log_file(f"predicting {k}")
  1031. data, seg, properties = dataset_val.load_case(k)
  1032. if self.is_cascaded:
  1033. data = np.vstack((data, convert_labelmap_to_one_hot(seg[-1], self.label_manager.foreground_labels,
  1034. output_dtype=data.dtype)))
  1035. with warnings.catch_warnings():
  1036. # ignore 'The given NumPy array is not writable' warning
  1037. warnings.simplefilter("ignore")
  1038. data = torch.from_numpy(data)
  1039. self.print_to_log_file(f'{k}, shape {data.shape}, rank {self.local_rank}')
  1040. output_filename_truncated = join(validation_output_folder, k)
  1041. prediction = predictor.predict_sliding_window_return_logits(data)
  1042. prediction = prediction.cpu()
  1043. # this needs to go into background processes
  1044. results.append(
  1045. segmentation_export_pool.starmap_async(
  1046. export_prediction_from_logits, (
  1047. (prediction, properties, self.configuration_manager, self.plans_manager,
  1048. self.dataset_json, output_filename_truncated, save_probabilities),
  1049. )
  1050. )
  1051. )
  1052. # for debug purposes
  1053. # export_prediction(prediction_for_export, properties, self.configuration, self.plans, self.dataset_json,
  1054. # output_filename_truncated, save_probabilities)
  1055. # if needed, export the softmax prediction for the next stage
  1056. if next_stages is not None:
  1057. for n in next_stages:
  1058. next_stage_config_manager = self.plans_manager.get_configuration(n)
  1059. expected_preprocessed_folder = join(nnUNet_preprocessed, self.plans_manager.dataset_name,
  1060. next_stage_config_manager.data_identifier)
  1061. try:
  1062. # we do this so that we can use load_case and do not have to hard code how loading training cases is implemented
  1063. tmp = nnUNetDataset(expected_preprocessed_folder, [k],
  1064. num_images_properties_loading_threshold=0)
  1065. d, s, p = tmp.load_case(k)
  1066. except FileNotFoundError:
  1067. self.print_to_log_file(
  1068. f"Predicting next stage {n} failed for case {k} because the preprocessed file is missing! "
  1069. f"Run the preprocessing for this configuration first!")
  1070. continue
  1071. target_shape = d.shape[1:]
  1072. output_folder = join(self.output_folder_base, 'predicted_next_stage', n)
  1073. output_file = join(output_folder, k + '.npz')
  1074. # resample_and_save(prediction, target_shape, output_file, self.plans_manager, self.configuration_manager, properties,
  1075. # self.dataset_json)
  1076. results.append(segmentation_export_pool.starmap_async(
  1077. resample_and_save, (
  1078. (prediction, target_shape, output_file, self.plans_manager,
  1079. self.configuration_manager,
  1080. properties,
  1081. self.dataset_json),
  1082. )
  1083. ))
  1084. # if we don't barrier from time to time we will get nccl timeouts for large datasets. Yuck.
  1085. if self.is_ddp and i < last_barrier_at_idx and (i + 1) % 20 == 0:
  1086. dist.barrier()
  1087. _ = [r.get() for r in results]
  1088. if self.is_ddp:
  1089. dist.barrier()
  1090. if self.local_rank == 0:
  1091. metrics = compute_metrics_on_folder(join(self.preprocessed_dataset_folder_base, 'gt_segmentations'),
  1092. validation_output_folder,
  1093. join(validation_output_folder, 'summary.json'),
  1094. self.plans_manager.image_reader_writer_class(),
  1095. self.dataset_json["file_ending"],
  1096. self.label_manager.foreground_regions if self.label_manager.has_regions else
  1097. self.label_manager.foreground_labels,
  1098. self.label_manager.ignore_label, chill=True,
  1099. num_processes=default_num_processes * dist.get_world_size() if
  1100. self.is_ddp else default_num_processes)
  1101. self.print_to_log_file("Validation complete", also_print_to_console=True)
  1102. self.print_to_log_file("Mean Validation Dice: ", (metrics['foreground_mean']["Dice"]), also_print_to_console=True)
  1103. self.set_deep_supervision_enabled(True)
  1104. compute_gaussian.cache_clear()
  1105. def run_training(self):
  1106. self.on_train_start()
  1107. for epoch in range(self.current_epoch, self.num_epochs):
  1108. self.on_epoch_start()
  1109. self.on_train_epoch_start()
  1110. train_outputs = []
  1111. for batch_id in range(self.num_iterations_per_epoch):
  1112. train_outputs.append(self.train_step(next(self.dataloader_train)))
  1113. self.on_train_epoch_end(train_outputs)
  1114. with torch.no_grad():
  1115. self.on_validation_epoch_start()
  1116. val_outputs = []
  1117. for batch_id in range(self.num_val_iterations_per_epoch):
  1118. val_outputs.append(self.validation_step(next(self.dataloader_val)))
  1119. self.on_validation_epoch_end(val_outputs)
  1120. self.on_epoch_end()
  1121. self.on_train_end()

nnUNetTrainer.py at commit 25225a8, under Apache-2.0 · at the source

Overview

Authors: Domen Preložnik1, Žiga Špiclin1
  1. Faculty of Electrical Engineering, University of Ljubljana, Ljubljana, Slovenia
Institutions: University of Ljubljana (Slovenia)
Journal: Frontiers in medicine, volume 13, article 1875760
Dates: received 8 May 2026; accepted 16 July 2026; published online 31 July 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.3389/fmed.2026.1875760 · PMID 42602271 · PMCID PMC13474016 · OpenAlex W7172015730
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: structural MRI / diffusion (modality), methods / tools (subfield)
Methods: Connectivity, Machine learning
Keywords: adaptive domain unlearning, comparative evaluation, domain generalization, image segmentation, open-source and reproducible, scanner variability
Topic: Domain Adaptation and Few-Shot Learning (Artificial Intelligence, Computer Science), according to OpenAlex
Citations: not cited yet (Europe PMC); 61 references in the paper

Abstract

Introduction: Inter-scanner variability in magnetic resonance imaging (MRI) adversely affects the diagnostic and prognostic quality of scans and necessitates the development of models that are robust to domain shift arising from the unseen scanner data. A review of recent advances in domain adaptation and domain generalization showed that the efficacy of strategies involving modifications or constraints on the latent space appears to be contingent upon the level and/or depth of supervision during model training.

Methods: We propose an adaptive multi-stage domain unlearning (ADMU) technique to improve robustness to unseen scanner domains. Building on the state-of-the-art segmentation framework nnU-Net, we employ deep supervision at deep encoder stages by applying domain classifier unlearning, sequentially across these stages to reduce domain-discriminative latent features. Following the self-configurable approach of nnU-Net, the auxiliary feedback loop implements an adaptive backpropagation schedule for unlearning. Experiments were conducted on four public datasets (one for training, three for testing) to benchmark white-matter lesion segmentation methods. Five benchmark models and/or strategies, spanning passive to active domain-robust training strategies, were tested, and five state-of-the-art methods were compared.

Results: AMDU demonstrated consistent, robust and improved cross-dataset segmentation performance on three test sets versus baseline nn-Unet variants. The advantage of AMDU was in enhanced lesion sensitivity with balanced false detections, resulting in good overall segmentation quality, as measured by segmentation overlap and relative lesion volume error. Compared to continuous domain unlearning, the adaptive scheduling balanced the adverse impact of unlearning onto the main segmentation task. Intensity-based preprocessing was found to be detrimental to segmentation performance.

Discussion: The proposed AMDU strategy was shown to be complementary to data augmentation. Demonstrated for white-matter lesion segmentation it relied only the FLAIR modality, simplifying preprocessing to spatial normalization to brain atlas, with no intensity harmonization, for best cross-dataset segmentation performance. The source code is available at https://github.com/Pubec/nnunetv2-unlearning.

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

Repositories

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

MIC-DKFZ/nnUNet

License: Apache-2.0
State: the link answers, verified on 26 September 2026
Evidence: files inventoried
Commit: 202f6baa0adc2ef5f7b615df19cc4da970412cc0, 25 September 2026
Languages: Python (216), Shell (7)
Size: 303 files, 223 scripts
Software Heritage: archived
Found in: the end of the paper
Holds: README, license file, environment (pyproject.toml, setup.py), tests, continuous integration, documentation
Not found: CITATION.cff
Tools: nnU-Net (125 files), NumPy (73 files), PyTorch (56 files), SimpleITK (10 files), scikit-image (6 files), SciPy (5 files), NiBabel (4 files), pandas (4 files), tifffile (3 files), Matplotlib (2 files), scikit-learn (1 file), seaborn (1 file)
Availability: 1 check, the latest on 26 September 2026: the link answers
  • 26 September 2026: the link answers
225 files

Pubec/nnunetv2-unlearning

License: Apache-2.0
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 25225a808b6a363c48be705df0c6a06996bf035e, 16 May 2024
Languages: Python (189), Shell (4)
Size: 221 files, 193 scripts
Software Heritage: not archived
Found in: the end of the paper
Holds: README, license file, environment (pyproject.toml, setup.py), tests, continuous integration, documentation
Not found: CITATION.cff
Tools: nnU-Net (108 files), NumPy (58 files), PyTorch (44 files), scikit-image (4 files), SciPy (4 files), SimpleITK (4 files), Matplotlib (3 files), pandas (3 files), NiBabel (2 files), seaborn (2 files), tifffile (2 files), scikit-learn (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
195 files

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:

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

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

Data

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

Data availability statement

The original contributions presented in the study are included in the article/supplementary material, further inquiries can be directed to the corresponding author.

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

Versions

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

Version 1, 27 September 2026: the first record

Recorded: type, language, journal, volume, pages, dates, 2 authors, 6 keywords, 1 funder, 52 references.

Cite

This paper

Preložnik, D., & Špiclin, Ž. (2026). Adaptive multi-stage domain unlearning for white-matter lesion segmentation. Frontiers in medicine, 13, 1875760. https://doi.org/10.3389/fmed.2026.1875760

BibTeX

@article{preloznik2026adaptive,
author = {Preložnik, Domen and Špiclin, Žiga},
title = {{Adaptive multi-stage domain unlearning for white-matter lesion segmentation}},
journal = {Frontiers in medicine},
year = {2026},
month = jul,
volume = {13},
pages = {1875760},
publisher = {Frontiers Media SA},
issn = {2296-858X},
doi = {10.3389/fmed.2026.1875760},
url = {https://doi.org/10.3389/fmed.2026.1875760},
pmid = {42602271},
pmcid = {PMC13474016}
}

RIS

TY - JOUR
AU - Preložnik, Domen
AU - Špiclin, Žiga
TI - Adaptive multi-stage domain unlearning for white-matter lesion segmentation
T2 - Frontiers in medicine
J2 - Front Med (Lausanne)
PY - 2026
DA - 2026/07/31
VL - 13
SP - 1875760
SN - 2296-858X
PB - Frontiers Media SA
DO - 10.3389/fmed.2026.1875760
UR - https://doi.org/10.3389/fmed.2026.1875760
LA - en
ER -

CSL-JSON

{
"id": "10.3389/fmed.2026.1875760",
"type": "article-journal",
"title": "Adaptive multi-stage domain unlearning for white-matter lesion segmentation",
"container-title": "Frontiers in medicine",
"author": [
{
"family": "Preložnik",
"given": "Domen"
},
{
"family": "Špiclin",
"given": "Žiga"
}
],
"container-title-short": "Front Med (Lausanne)",
"volume": "13",
"page": "1875760",
"DOI": "10.3389/fmed.2026.1875760",
"PMID": "42602271",
"PMCID": "PMC13474016",
"ISSN": "2296-858X",
"publisher": "Frontiers Media SA",
"URL": "https://doi.org/10.3389/fmed.2026.1875760",
"language": "en",
"issued": {
"date-parts": [
[
2026,
7,
31
]
]
}
}

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.1002/hipo.70124 [code]
Association Between Anterior Hippocampal Gyrification and Episodic Memory Performance in Neurotypical Young Adults.
Journal: Hippocampus
In common: nnU-Net, SimpleITK, tifffile, 9 other tools, structural MRI / diffusion, 2 references
[2] doi:10.1186/s12880-026-02335-x [code]
Automatic lateral ventricle and choroid plexus segmentation method in infant brain MR images.
Journal: BMC medical imaging
In common: nnU-Net, SimpleITK, tifffile, 9 other tools, structural MRI / diffusion, 2 references
[3] doi:10.64898/2026.07.15.26357954 [code]
Portable Ultra-Low Field MRI Deep-Learning Algorithms for White Matter Lesion Segmentation Improve Accuracy and Reflect Clinical Disability in Multiple Sclerosis
Journal: medRxiv (preprint)
In common: nnU-Net, SimpleITK, tifffile, 9 other tools, structural MRI / diffusion, 1 reference
[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: nnU-Net, SimpleITK, tifffile, 9 other tools, methods / tools, structural MRI / diffusion, 1 reference
[5] 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: nnU-Net, SimpleITK, tifffile, 9 other tools, structural MRI / diffusion, 1 reference
[6] doi:10.1007/s12021-026-09817-x [code]
Circle of Willis-Guided Localization for Simultaneous Detection and Classification of Large Vessel Occlusions in Brain CTA.
Journal: Neuroinformatics
In common: nnU-Net, SimpleITK, tifffile, 9 other tools, 1 reference
[7] doi:10.1016/j.adro.2026.102092 [code]
Effect of Anatomic Contextual Information on the Performance of a Convolutional Neural Network Tasked With Brain Metastasis Detection.
Journal: Advances in radiation oncology
In common: nnU-Net, SimpleITK, tifffile, 9 other tools, 1 reference
[8] 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: scikit-image, NiBabel, PyTorch, 4 other tools, methods / tools, structural MRI / diffusion, 5 references
[9] doi:10.1136/jnnp-2025-335884 [code]
Diffusivity anisotropy signature of slowly expanding lesions predicts progression independent of relapse activity in multiple sclerosis.
Journal: Journal of neurology, neurosurgery, and psychiatry
In common: nnU-Net, SimpleITK, tifffile, 9 other tools, structural MRI / diffusion
[10] doi:10.3390/jimaging12070276 [code]
Hyperelastic Regularization for Near-Diffeomorphic Transformer-Based Brain MRI Registration.
Journal: Journal of imaging
In common: nnU-Net, SimpleITK, scikit-image, 7 other tools, methods / tools, structural MRI / diffusion, 1 reference

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.