OSCR

An nnU-Net-based framework with adaptive feature representation for 3D brain tumor segmentation.

Code ↔ Paper

8 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 8 matches · 2 of them tie a paragraph to a whole file, not to given lines: weak matches, whose lines are not tinted
  1. [1] § Methods › Training and inference protocol › Training settings ↔ nnunetv2/training/nnUNetTrainer/variants/pathology/nnUNetTrainer_WSD_undefined_dataloader.py, lines 60–194 · score 0.78 · oversample_foreground_percent, weight decay, scheduled, disabled, iterations, epoch
  2. [2] § Methods › Training and inference protocol › Training settings ↔ nnunetv2/training/nnUNetTrainer/nnUNetTrainer.py, lines 68–197 · score 0.77 · oversample_foreground_percent, weight decay, scheduled, disabled, iterations, epoch
  3. [3] § Methods ↔ nnunetv2/dataset_conversion/Dataset137_BraTS21.py, lines 59–98 · score 0.66 · T1ce, tumor core, BraTS, enhancing tumor, FLAIR, segmentation
  4. [4] § Results › Qualitative results and case-level failure analysis › Qualitative comparison ↔ nnunetv2/dataset_conversion/Dataset137_BraTS21.py, lines 59–98 · score 0.61 · T1ce, tumor core, BraTS, enhancing tumor, FLAIR, segmentation
  5. [5] § Results › Implementation details ↔ nnunetv2/batch_running/benchmarking/generate_benchmarking_commands.py, the whole file · a weak match · score 0.56 · Tesla V100 SXM2, NVIDIA, GPU, configuration, training
  6. [6] § Methods › Boundary-aware objective function ↔ nnunetv2/training/loss/compound_losses.py, lines 60–100 · score 0.52 · dice loss, aggregation, logits, network, class
  7. [7] § Methods › Boundary-aware objective function ↔ nnunetv2/training/nnUNetTrainer/variants/network_architecture/nnUNetTrainerNoDeepSupervision.py, the whole file · a weak match · score 0.51 · dice loss, softmax, batch, network, class, segmentation
  8. [8] § Methods › Improved network architecture › Backbone: nnU-Net v2 setup ↔ nnunetv2/training/nnUNetTrainer/nnUNetTrainer.py, lines 1104–1231 · score 0.51 · sliding window, deep supervision, foreground, prediction, inference, network

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,254 lines · 66 KB · Apache-2.0 · 2 matches

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

nnUNetTrainer.py at commit 23b4088, under Apache-2.0 · at the source

Overview

Authors: Chenghong Zhang1, Qiang Wei1
ORCID iDs: Chenghong Zhang
  1. Department of Electronics, School of Electronic Information Engineering, Guiyang University, Guiyang, China
Institutions: Guiyang University (China)
Journal: Quantitative imaging in medicine and surgery, volume 16, issue 9, article 712
Dates: received 1 April 2026; accepted 21 July 2026; published online 10 August 2026; in print 1 September 2026
Type: Research article · Language: English
License: CC BY-NC-ND
Identifiers: DOI 10.21037/qims-2026-0792 · PMID 42701467 · PMCID PMC13545556 · OpenAlex W7204256151
Open access: diamond, a free copy (OpenAlex)
Status: code verified
Categories: structural MRI / diffusion (modality), other condition (population), methods / tools (subfield)
Methods: Connectivity, Machine learning
Keywords: Glioma, multimodal magnetic resonance imaging (multimodal MRI), nnU-Net, medical image segmentation, boundary-aware learning
Topic: Brain Tumor Detection and Classification (Neurology, Neuroscience), according to OpenAlex
Citations: not cited yet (Europe PMC); 23 references in the paper

Abstract

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

Repository

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

DIAGNijmegen/nnUNet_v2

License: Apache-2.0
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 23b408840fbad2684e122aad08847f2bfac5f5f8, 7 February 2024
Languages: Python (185), Shell (8), Jupyter (4)
Size: 243 files, 197 scripts
Software Heritage: not archived
Found in: the text, “Implementation details”
Holds: README, license file, environment (pyproject.toml, setup.py, docker/DIAG/sol1/Dockerfile, docker/DIAG/sol2/Dockerfile), tests, continuous integration, documentation, 4 notebooks
Not found: CITATION.cff
Tools: nnU-Net (112 files), NumPy (61 files), PyTorch (57 files), SciPy (6 files), Matplotlib (5 files), scikit-image (5 files), pandas (4 files), SimpleITK (4 files), scikit-learn (3 files), NiBabel (2 files), tifffile (2 files), OpenCV (1 file), seaborn (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
199 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:

  • 1 repository of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 197 scripts, each with its path and the digest of its content;
  • 8 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.

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, issue, pages, dates, 2 authors, 5 keywords, 20 references.

Cite

This paper

Zhang, C., & Wei, Q. (2026). An nnU-Net-based framework with adaptive feature representation for 3D brain tumor segmentation. Quantitative imaging in medicine and surgery, 16(9), 712. https://doi.org/10.21037/qims-2026-0792

BibTeX

@article{zhang2026nnu,
author = {Zhang, Chenghong and Wei, Qiang},
title = {{An nnU-Net-based framework with adaptive feature representation for 3D brain tumor segmentation}},
journal = {Quantitative imaging in medicine and surgery},
year = {2026},
month = aug,
volume = {16},
number = {9},
pages = {712},
publisher = {AME Publications},
issn = {2223-4292},
doi = {10.21037/qims-2026-0792},
url = {https://doi.org/10.21037/qims-2026-0792},
pmid = {42701467},
pmcid = {PMC13545556}
}

RIS

TY - JOUR
AU - Zhang, Chenghong
AU - Wei, Qiang
TI - An nnU-Net-based framework with adaptive feature representation for 3D brain tumor segmentation
T2 - Quantitative imaging in medicine and surgery
J2 - Quant Imaging Med Surg
PY - 2026
DA - 2026/08/10
VL - 16
IS - 9
SP - 712
SN - 2223-4292
PB - AME Publications
DO - 10.21037/qims-2026-0792
UR - https://doi.org/10.21037/qims-2026-0792
LA - en
ER -

CSL-JSON

{
"id": "10.21037/qims-2026-0792",
"type": "article-journal",
"title": "An nnU-Net-based framework with adaptive feature representation for 3D brain tumor segmentation",
"container-title": "Quantitative imaging in medicine and surgery",
"author": [
{
"family": "Zhang",
"given": "Chenghong"
},
{
"family": "Wei",
"given": "Qiang"
}
],
"container-title-short": "Quant Imaging Med Surg",
"volume": "16",
"issue": "9",
"page": "712",
"DOI": "10.21037/qims-2026-0792",
"PMID": "42701467",
"PMCID": "PMC13545556",
"ISSN": "2223-4292",
"publisher": "AME Publications",
"URL": "https://doi.org/10.21037/qims-2026-0792",
"language": "en",
"issued": {
"date-parts": [
[
2026,
8,
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.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, 10 other tools, structural MRI / diffusion, 1 reference
[2] doi:10.3389/fmed.2026.1875760 [code]
Adaptive multi-stage domain unlearning for white-matter lesion segmentation.
Journal: Frontiers in medicine
In common: nnU-Net, SimpleITK, tifffile, 9 other tools, methods / tools, structural MRI / diffusion, 1 reference
[3] 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, 1 reference
[4] 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, 1 reference
[5] 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, other condition, 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.1364/boe.600665 [code]
NeuroSeg-MF: robust neuron segmentation in two-photon Ca&lt;sup&gt;2+&lt;/sup&gt; imaging using multi-feature fusion and detection-guided SAM.
Journal: Biomedical optics express
In common: tifffile, OpenCV, PyTorch, 5 other tools, methods / tools, 5 references
[8] doi:10.3389/frai.2026.1771088 [code]
Few-shot deployment of pretrained MRI transformers in brain imaging tasks.
Journal: Frontiers in artificial intelligence
In common: nnU-Net, SimpleITK, OpenCV, 9 other tools, methods / tools, structural MRI / diffusion
[9] doi:10.1002/alz.71649 [code]
Postmortem brain MRI reveals differential associations of subcortical and limbic volumes with cortical thinning and neurodegenerative pathologies.
Journal: Alzheimer's & dementia : the journal of the Alzheimer's Association
In common: nnU-Net, SimpleITK, OpenCV, 9 other tools, structural MRI / diffusion, other condition
[10] doi:10.1038/s41598-026-48496-1 [code]
A unified FLAIR hyperintensity segmentation model for various CNS tumor types and acquisition time points.
Journal: Scientific reports
In common: SimpleITK, scikit-image, NiBabel, 6 other tools, methods / tools, structural MRI / diffusion, other condition, 4 references

Contribute

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

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

Request its removal

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

Discussion, reproductions, activity

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

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

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