Adaptive multi-stage domain unlearning for white-matter lesion segmentation.
The 3 matches
- [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] § 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] § 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
- import inspect
- import multiprocessing
- import os
- import shutil
- import sys
- import warnings
- from copy import deepcopy
- from datetime import datetime
- from time import time, sleep
- from typing import Union, Tuple, List
- import numpy as np
- import torch
- from batchgenerators.dataloading.multi_threaded_augmenter import MultiThreadedAugmenter
- from batchgenerators.dataloading.nondet_multi_threaded_augmenter import NonDetMultiThreadedAugmenter
- from batchgenerators.dataloading.single_threaded_augmenter import SingleThreadedAugmenter
- from batchgenerators.transforms.abstract_transforms import AbstractTransform, Compose
- from batchgenerators.transforms.color_transforms import BrightnessMultiplicativeTransform, \
- ContrastAugmentationTransform, GammaTransform
- from batchgenerators.transforms.noise_transforms import GaussianNoiseTransform, GaussianBlurTransform
- from batchgenerators.transforms.resample_transforms import SimulateLowResolutionTransform
- from batchgenerators.transforms.spatial_transforms import SpatialTransform, MirrorTransform
- from batchgenerators.transforms.utility_transforms import RemoveLabelTransform, RenameTransform, NumpyToTensor
- from batchgenerators.utilities.file_and_folder_operations import join, load_json, isfile, save_json, maybe_mkdir_p, isdir
- from torch._dynamo import OptimizedModule
- from nnunetv2.configuration import ANISO_THRESHOLD, default_num_processes
- from nnunetv2.evaluation.evaluate_predictions import compute_metrics_on_folder
- from nnunetv2.inference.export_prediction import export_prediction_from_logits, resample_and_save
- from nnunetv2.inference.predict_from_raw_data import nnUNetPredictor
- from nnunetv2.inference.sliding_window_prediction import compute_gaussian
- from nnunetv2.paths import nnUNet_preprocessed, nnUNet_results
- from nnunetv2.training.data_augmentation.compute_initial_patch_size import get_patch_size
- from nnunetv2.training.data_augmentation.custom_transforms.cascade_transforms import MoveSegAsOneHotToData, \
- ApplyRandomBinaryOperatorTransform, RemoveRandomConnectedComponentFromOneHotEncodingTransform
- from nnunetv2.training.data_augmentation.custom_transforms.deep_supervision_donwsampling import \
- DownsampleSegForDSTransform2
- from nnunetv2.training.data_augmentation.custom_transforms.limited_length_multithreaded_augmenter import \
- LimitedLenWrapper
- from nnunetv2.training.data_augmentation.custom_transforms.masking import MaskTransform
- from nnunetv2.training.data_augmentation.custom_transforms.region_based_training import \
- ConvertSegmentationToRegionsTransform
- from nnunetv2.training.data_augmentation.custom_transforms.transforms_for_dummy_2d import Convert2DTo3DTransform, \
- Convert3DTo2DTransform
- from nnunetv2.training.dataloading.data_loader_2d import nnUNetDataLoader2D
- from nnunetv2.training.dataloading.data_loader_3d import nnUNetDataLoader3D
- from nnunetv2.training.dataloading.nnunet_dataset import nnUNetDataset
- from nnunetv2.training.dataloading.utils import get_case_identifiers, unpack_dataset
- from nnunetv2.training.logging.nnunet_logger import nnUNetLogger
- from nnunetv2.training.loss.compound_losses import DC_and_CE_loss, DC_and_BCE_loss
- from nnunetv2.training.loss.deep_supervision import DeepSupervisionWrapper
- from nnunetv2.training.loss.dice import get_tp_fp_fn_tn, MemoryEfficientSoftDiceLoss
- from nnunetv2.training.lr_scheduler.polylr import PolyLRScheduler
- from nnunetv2.utilities.collate_outputs import collate_outputs
- from nnunetv2.utilities.crossval_split import generate_crossval_split
- from nnunetv2.utilities.default_n_proc_DA import get_allowed_n_proc_DA
- from nnunetv2.utilities.file_path_utilities import check_workers_alive_and_busy
- from nnunetv2.utilities.get_network_from_plans import get_network_from_plans
- from nnunetv2.utilities.helpers import empty_cache, dummy_context
- from nnunetv2.utilities.label_handling.label_handling import convert_labelmap_to_one_hot, determine_num_input_channels
- from nnunetv2.utilities.plans_handling.plans_handler import PlansManager, ConfigurationManager
- from torch import autocast, nn
- from torch import distributed as dist
- from torch.cuda import device_count
- from torch.cuda.amp import GradScaler
- from torch.nn.parallel import DistributedDataParallel as DDP
- class nnUNetTrainer(object):
- def __init__(self, plans: dict, configuration: str, fold: int, dataset_json: dict, experiment_identifier: str = "", unpack_dataset: bool = True,
- device: torch.device = torch.device('cuda'), dir_can_exist: bool = False):
- # From https://grugbrain.dev/. Worth a read ya big brains ;-)
- # apex predator of grug is complexity
- # complexity bad
- # say again:
- # complexity very bad
- # you say now:
- # complexity very, very bad
- # given choice between complexity or one on one against t-rex, grug take t-rex: at least grug see t-rex
- # 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
- # one day code base understandable and grug can get work done, everything good!
- # next day impossible: complexity demon spirit has entered code and very dangerous situation!
- # OK OK I am guilty. But I tried.
- # https://www.osnews.com/images/comics/wtfm.jpg
- # https://i.pinimg.com/originals/26/b2/50/26b250a738ea4abc7a5af4d42ad93af0.jpg
- self.is_ddp = dist.is_available() and dist.is_initialized()
- self.local_rank = 0 if not self.is_ddp else dist.get_rank()
- self.device = device
- # print what device we are using
- if self.is_ddp: # implicitly it's clear that we use cuda in this case
- print(f"I am local rank {self.local_rank}. {device_count()} GPUs are available. The world size is "
- f"{dist.get_world_size()}."
- f"Setting device to {self.device}")
- self.device = torch.device(type='cuda', index=self.local_rank)
- else:
- if self.device.type == 'cuda':
- # we might want to let the user pick this but for now please pick the correct GPU with CUDA_VISIBLE_DEVICES=X
- self.device = torch.device(type='cuda', index=0)
- print(f"Using device: {self.device}")
- # loading and saving this class for continuing from checkpoint should not happen based on pickling. This
- # would also pickle the network etc. Bad, bad. Instead we just reinstantiate and then load the checkpoint we
- # need. So let's save the init args
- self.my_init_kwargs = {}
- for k in inspect.signature(self.__init__).parameters.keys():
- self.my_init_kwargs[k] = locals()[k]
- ### Saving all the init args into class variables for later access
- self.plans_manager = PlansManager(plans)
- self.configuration_manager = self.plans_manager.get_configuration(configuration)
- self.configuration_name = configuration
- self.dataset_json = dataset_json
- self.fold = fold
- self.unpack_dataset = unpack_dataset
- ### Setting all the folder names. We need to make sure things don't crash in case we are just running
- # inference and some of the folders may not be defined!
- self.preprocessed_dataset_folder_base = join(nnUNet_preprocessed, self.plans_manager.dataset_name) \
- if nnUNet_preprocessed is not None else None
- dir_name = f"{self.__class__.__name__}__{self.plans_manager.plans_name}__{configuration}"
- self.output_folder_base = join(nnUNet_results, self.plans_manager.dataset_name, dir_name) if nnUNet_results is not None else None
- # added experiment_identifier to make easier experiment numbering
- # experiment identifier is appended to the back of output_folder_name
- self.experiment_identifier = experiment_identifier
- if experiment_identifier:
- self.output_folder_base += f"__{experiment_identifier}"
- self.output_folder = join(self.output_folder_base, f'fold_{fold}')
- if isdir(self.output_folder) and not dir_can_exist:
- raise ValueError(f"Folder {dir_name} exist. Safety check to prevent overwriting of experiments")
- self.preprocessed_dataset_folder = join(self.preprocessed_dataset_folder_base,
- self.configuration_manager.data_identifier)
- # unlike the previous nnunet folder_with_segs_from_previous_stage is now part of the plans. For now it has to
- # be a different configuration in the same plans
- # IMPORTANT! the mapping must be bijective, so lowres must point to fullres and vice versa (using
- # "previous_stage" and "next_stage"). Otherwise it won't work!
- self.is_cascaded = self.configuration_manager.previous_stage_name is not None
- self.folder_with_segs_from_previous_stage = \
- join(nnUNet_results, self.plans_manager.dataset_name,
- self.__class__.__name__ + '__' + self.plans_manager.plans_name + "__" +
- self.configuration_manager.previous_stage_name, 'predicted_next_stage', self.configuration_name) \
- if self.is_cascaded else None
- ### Some hyperparameters for you to fiddle with
- self.initial_lr = 1e-2
- self.weight_decay = 3e-5
- self.oversample_foreground_percent = 0.33
- self.num_iterations_per_epoch = 250
- self.num_val_iterations_per_epoch = 50
- self.current_epoch = 0
- self.enable_deep_supervision = True
- self.num_epochs = self.configuration_manager.configuration.get('num_epochs', 500)
- ### Dealing with labels/regions
- self.label_manager = self.plans_manager.get_label_manager(dataset_json)
- # labels can either be a list of int (regular training) or a list of tuples of int (region-based training)
- # needed for predictions. We do sigmoid in case of (overlapping) regions
- self.num_input_channels = None # -> self.initialize()
- self.network = None # -> self.build_network_architecture()
- self.optimizer = self.lr_scheduler = None # -> self.initialize
- self.grad_scaler = GradScaler() if self.device.type == 'cuda' else None
- self.loss = None # -> self.initialize
- ### Simple logging. Don't take that away from me!
- # initialize log file. This is just our log for the print statements etc. Not to be confused with lightning
- # logging
- timestamp = datetime.now()
- maybe_mkdir_p(self.output_folder)
- self.log_file = join(self.output_folder, "training_log_%d_%d_%d_%02.0d_%02.0d_%02.0d.txt" %
- (timestamp.year, timestamp.month, timestamp.day, timestamp.hour, timestamp.minute,
- timestamp.second))
- self.logger = nnUNetLogger()
- ### placeholders
- self.dataloader_train = self.dataloader_val = None # see on_train_start
- ### initializing stuff for remembering things and such
- self._best_ema = None
- ### inference things
- self.inference_allowed_mirroring_axes = None # this variable is set in
- # self.configure_rotation_dummyDA_mirroring_and_inital_patch_size and will be saved in checkpoints
- ### checkpoint saving stuff
- self.save_every = 50
- self.disable_checkpointing = False
- ## DDP batch size and oversampling can differ between workers and needs adaptation
- # we need to change the batch size in DDP because we don't use any of those distributed samplers
- self._set_batch_size_and_oversample()
- self.was_initialized = False
- self.print_to_log_file("\n#######################################################################\n"
- "Please cite the following paper when using nnU-Net:\n"
- "Isensee, F., Jaeger, P. F., Kohl, S. A., Petersen, J., & Maier-Hein, K. H. (2021). "
- "nnU-Net: a self-configuring method for deep learning-based biomedical image segmentation. "
- "Nature methods, 18(2), 203-211.\n"
- "#######################################################################\n",
- also_print_to_console=True, add_timestamp=False)
- def initialize(self):
- if not self.was_initialized:
- self.num_input_channels = determine_num_input_channels(self.plans_manager, self.configuration_manager,
- self.dataset_json)
- self.network = self.build_network_architecture(
- self.plans_manager,
- self.dataset_json,
- self.configuration_manager,
- self.num_input_channels,
- self.enable_deep_supervision,
- ).to(self.device)
- # compile network for free speedup
- if self._do_i_compile():
- self.print_to_log_file('Using torch.compile...')
- self.network = torch.compile(self.network)
- self.optimizer, self.lr_scheduler = self.configure_optimizers()
- # if ddp, wrap in DDP wrapper
- if self.is_ddp:
- self.network = torch.nn.SyncBatchNorm.convert_sync_batchnorm(self.network)
- self.network = DDP(self.network, device_ids=[self.local_rank])
- self.loss = self._build_loss()
- self.was_initialized = True
- else:
- raise RuntimeError("You have called self.initialize even though the trainer was already initialized. "
- "That should not happen.")
- def _do_i_compile(self):
- return ('nnUNet_compile' in os.environ.keys()) and (os.environ['nnUNet_compile'].lower() in ('true', '1', 't'))
- def _save_debug_information(self):
- # saving some debug information
- if self.local_rank == 0:
- dct = {}
- for k in self.__dir__():
- if not k.startswith("__"):
- if not callable(getattr(self, k)) or k in ['loss', ]:
- dct[k] = str(getattr(self, k))
- elif k in ['network', ]:
- dct[k] = str(getattr(self, k).__class__.__name__)
- else:
- # print(k)
- pass
- if k in ['dataloader_train', 'dataloader_val']:
- if hasattr(getattr(self, k), 'generator'):
- dct[k + '.generator'] = str(getattr(self, k).generator)
- if hasattr(getattr(self, k), 'num_processes'):
- dct[k + '.num_processes'] = str(getattr(self, k).num_processes)
- if hasattr(getattr(self, k), 'transform'):
- dct[k + '.transform'] = str(getattr(self, k).transform)
- import subprocess
- hostname = subprocess.getoutput(['hostname'])
- dct['hostname'] = hostname
- torch_version = torch.__version__
- if self.device.type == 'cuda':
- gpu_name = torch.cuda.get_device_name()
- dct['gpu_name'] = gpu_name
- cudnn_version = torch.backends.cudnn.version()
- else:
- cudnn_version = 'None'
- dct['device'] = str(self.device)
- dct['torch_version'] = torch_version
- dct['cudnn_version'] = cudnn_version
- save_json(dct, join(self.output_folder, "debug.json"))
- @staticmethod
- def build_network_architecture(plans_manager: PlansManager,
- dataset_json,
- configuration_manager: ConfigurationManager,
- num_input_channels,
- enable_deep_supervision: bool = True) -> nn.Module:
- """
- This is where you build the architecture according to the plans. There is no obligation to use
- get_network_from_plans, this is just a utility we use for the nnU-Net default architectures. You can do what
- you want. Even ignore the plans and just return something static (as long as it can process the requested
- patch size)
- but don't bug us with your bugs arising from fiddling with this :-P
- This is the function that is called in inference as well! This is needed so that all network architecture
- variants can be loaded at inference time (inference will use the same nnUNetTrainer that was used for
- training, so if you change the network architecture during training by deriving a new trainer class then
- inference will know about it).
- If you need to know how many segmentation outputs your custom architecture needs to have, use the following snippet:
- > label_manager = plans_manager.get_label_manager(dataset_json)
- > label_manager.num_segmentation_heads
- (why so complicated? -> We can have either classical training (classes) or regions. If we have regions,
- the number of outputs is != the number of classes. Also there is the ignore label for which no output
- should be generated. label_manager takes care of all that for you.)
- """
- return get_network_from_plans(plans_manager, dataset_json, configuration_manager,
- num_input_channels, deep_supervision=enable_deep_supervision)
- def _get_deep_supervision_scales(self):
- if self.enable_deep_supervision:
- deep_supervision_scales = list(list(i) for i in 1 / np.cumprod(np.vstack(
- self.configuration_manager.pool_op_kernel_sizes), axis=0))[:-1]
- else:
- deep_supervision_scales = None # for train and val_transforms
- return deep_supervision_scales
- def _set_batch_size_and_oversample(self):
- if not self.is_ddp:
- # set batch size to what the plan says, leave oversample untouched
- self.batch_size = self.configuration_manager.batch_size
- else:
- # batch size is distributed over DDP workers and we need to change oversample_percent for each worker
- world_size = dist.get_world_size()
- my_rank = dist.get_rank()
- global_batch_size = self.configuration_manager.batch_size
- assert global_batch_size >= world_size, 'Cannot run DDP if the batch size is smaller than the number of ' \
- 'GPUs... Duh.'
- batch_size_per_GPU = [global_batch_size // world_size] * world_size
- batch_size_per_GPU = [batch_size_per_GPU[i] + 1
- if (batch_size_per_GPU[i] * world_size + i) < global_batch_size
- else batch_size_per_GPU[i]
- for i in range(len(batch_size_per_GPU))]
- assert sum(batch_size_per_GPU) == global_batch_size
- sample_id_low = 0 if my_rank == 0 else np.sum(batch_size_per_GPU[:my_rank])
- sample_id_high = np.sum(batch_size_per_GPU[:my_rank + 1])
- # This is how oversampling is determined in DataLoader
- # round(self.batch_size * (1 - self.oversample_foreground_percent))
- # We need to use the same scheme here because an oversample of 0.33 with a batch size of 2 will be rounded
- # to an oversample of 0.5 (1 sample random, one oversampled). This may get lost if we just numerically
- # compute oversample
- oversample = [True if not i < round(global_batch_size * (1 - self.oversample_foreground_percent)) else False
- for i in range(global_batch_size)]
- if sample_id_high / global_batch_size < (1 - self.oversample_foreground_percent):
- oversample_percent = 0.0
- elif sample_id_low / global_batch_size > (1 - self.oversample_foreground_percent):
- oversample_percent = 1.0
- else:
- oversample_percent = sum(oversample[sample_id_low:sample_id_high]) / batch_size_per_GPU[my_rank]
- print("worker", my_rank, "oversample", oversample_percent)
- print("worker", my_rank, "batch_size", batch_size_per_GPU[my_rank])
- # self.print_to_log_file("worker", my_rank, "oversample", oversample_percents[my_rank])
- # self.print_to_log_file("worker", my_rank, "batch_size", batch_sizes[my_rank])
- self.batch_size = batch_size_per_GPU[my_rank]
- self.oversample_foreground_percent = oversample_percent
- def _build_loss(self):
- if self.label_manager.has_regions:
- loss = DC_and_BCE_loss({},
- {'batch_dice': self.configuration_manager.batch_dice,
- 'do_bg': True, 'smooth': 1e-5, 'ddp': self.is_ddp},
- use_ignore_label=self.label_manager.ignore_label is not None,
- dice_class=MemoryEfficientSoftDiceLoss)
- else:
- loss = DC_and_CE_loss({'batch_dice': self.configuration_manager.batch_dice,
- 'smooth': 1e-5, 'do_bg': False, 'ddp': self.is_ddp}, {}, weight_ce=1, weight_dice=1,
- ignore_label=self.label_manager.ignore_label, dice_class=MemoryEfficientSoftDiceLoss)
- # we give each output a weight which decreases exponentially (division by 2) as the resolution decreases
- # this gives higher resolution outputs more weight in the loss
- if self.enable_deep_supervision:
- deep_supervision_scales = self._get_deep_supervision_scales()
- weights = np.array([1 / (2**i) for i in range(len(deep_supervision_scales))])
- if self.is_ddp and not self._do_i_compile():
- # very strange and stupid interaction. DDP crashes and complains about unused parameters due to
- # weights[-1] = 0. Interestingly this crash doesn't happen with torch.compile enabled. Strange stuff.
- # Anywho, the simple fix is to set a very low weight to this.
- weights[-1] = 1e-6
- else:
- weights[-1] = 0
- # we don't use the lowest 2 outputs. Normalize weights so that they sum to 1
- weights = weights / weights.sum()
- # now wrap the loss
- loss = DeepSupervisionWrapper(loss, weights)
- return loss
- def configure_rotation_dummyDA_mirroring_and_inital_patch_size(self):
- """
- This function is stupid and certainly one of the weakest spots of this implementation. Not entirely sure how we can fix it.
- """
- patch_size = self.configuration_manager.patch_size
- dim = len(patch_size)
- # todo rotation should be defined dynamically based on patch size (more isotropic patch sizes = more rotation)
- if dim == 2:
- do_dummy_2d_data_aug = False
- # todo revisit this parametrization
- if max(patch_size) / min(patch_size) > 1.5:
- rotation_for_DA = {
- 'x': (-15. / 360 * 2. * np.pi, 15. / 360 * 2. * np.pi),
- 'y': (0, 0),
- 'z': (0, 0)
- }
- else:
- rotation_for_DA = {
- 'x': (-180. / 360 * 2. * np.pi, 180. / 360 * 2. * np.pi),
- 'y': (0, 0),
- 'z': (0, 0)
- }
- mirror_axes = (0, 1)
- elif dim == 3:
- # 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
- # order of the axes is determined by spacing, not image size
- do_dummy_2d_data_aug = (max(patch_size) / patch_size[0]) > ANISO_THRESHOLD
- if do_dummy_2d_data_aug:
- # why do we rotate 180 deg here all the time? We should also restrict it
- rotation_for_DA = {
- 'x': (-180. / 360 * 2. * np.pi, 180. / 360 * 2. * np.pi),
- 'y': (0, 0),
- 'z': (0, 0)
- }
- else:
- rotation_for_DA = {
- 'x': (-30. / 360 * 2. * np.pi, 30. / 360 * 2. * np.pi),
- 'y': (-30. / 360 * 2. * np.pi, 30. / 360 * 2. * np.pi),
- 'z': (-30. / 360 * 2. * np.pi, 30. / 360 * 2. * np.pi),
- }
- mirror_axes = (0, 1, 2)
- else:
- raise RuntimeError()
- # todo this function is stupid. It doesn't even use the correct scale range (we keep things as they were in the
- # old nnunet for now)
- initial_patch_size = get_patch_size(patch_size[-dim:],
- *rotation_for_DA.values(),
- (0.85, 1.25))
- if do_dummy_2d_data_aug:
- initial_patch_size[0] = patch_size[0]
- self.print_to_log_file(f'do_dummy_2d_data_aug: {do_dummy_2d_data_aug}')
- self.inference_allowed_mirroring_axes = mirror_axes
- return rotation_for_DA, do_dummy_2d_data_aug, initial_patch_size, mirror_axes
- def print_to_log_file(self, *args, also_print_to_console=True, add_timestamp=True):
- if self.local_rank == 0:
- timestamp = time()
- dt_object = datetime.fromtimestamp(timestamp)
- if add_timestamp:
- args = (f"{dt_object}:", *args)
- successful = False
- max_attempts = 5
- ctr = 0
- while not successful and ctr < max_attempts:
- try:
- with open(self.log_file, 'a+') as f:
- for a in args:
- f.write(str(a))
- f.write(" ")
- f.write("\n")
- successful = True
- except IOError:
- print(f"{datetime.fromtimestamp(timestamp)}: failed to log: ", sys.exc_info())
- sleep(0.5)
- ctr += 1
- if also_print_to_console:
- print(*args)
- elif also_print_to_console:
- print(*args)
- def print_plans(self):
- if self.local_rank == 0:
- dct = deepcopy(self.plans_manager.plans)
- del dct['configurations']
- self.print_to_log_file(f"\nThis is the configuration used by this "
- f"training:\nConfiguration name: {self.configuration_name}\n",
- self.configuration_manager, '\n', add_timestamp=False)
- self.print_to_log_file('These are the global plan.json settings:\n', dct, '\n', add_timestamp=False)
- def configure_optimizers(self):
- optimizer = torch.optim.SGD(self.network.parameters(), self.initial_lr, weight_decay=self.weight_decay,
- momentum=0.99, nesterov=True)
- lr_scheduler = PolyLRScheduler(
- optimizer,
- self.initial_lr,
- 1000,
- # self.num_epochs
- )
- return optimizer, lr_scheduler
- def plot_network_architecture(self):
- if self._do_i_compile():
- self.print_to_log_file("Unable to plot network architecture: nnUNet_compile is enabled!")
- return
- if self.local_rank == 0:
- try:
- # raise NotImplementedError('hiddenlayer no longer works and we do not have a viable alternative :-(')
- # pip install git+https://github.com/saugatkandel/hiddenlayer.git
- # from torchviz import make_dot
- # # not viable.
- # make_dot(tuple(self.network(torch.rand((1, self.num_input_channels,
- # *self.configuration_manager.patch_size),
- # device=self.device)))).render(
- # join(self.output_folder, "network_architecture.pdf"), format='pdf')
- # self.optimizer.zero_grad()
- # broken.
- import hiddenlayer as hl
- g = hl.build_graph(self.network,
- torch.rand((1, self.num_input_channels,
- *self.configuration_manager.patch_size),
- device=self.device),
- transforms=None)
- g.save(join(self.output_folder, "network_architecture.pdf"))
- del g
- except Exception as e:
- self.print_to_log_file("Unable to plot network architecture:")
- self.print_to_log_file(e)
- # self.print_to_log_file("\nprinting the network instead:\n")
- # self.print_to_log_file(self.network)
- # self.print_to_log_file("\n")
- finally:
- empty_cache(self.device)
- def do_split(self):
- """
- The default split is a 5 fold CV on all available training cases. nnU-Net will create a split (it is seeded,
- so always the same) and save it as splits_final.json file in the preprocessed data directory.
- Sometimes you may want to create your own split for various reasons. For this you will need to create your own
- splits_final.json file. If this file is present, nnU-Net is going to use it and whatever splits are defined in
- it. You can create as many splits in this file as you want. Note that if you define only 4 splits (fold 0-3)
- and then set fold=4 when training (that would be the fifth split), nnU-Net will print a warning and proceed to
- use a random 80:20 data split.
- :return:
- """
- if self.fold == "all":
- # if fold==all then we use all images for training and validation
- case_identifiers = get_case_identifiers(self.preprocessed_dataset_folder)
- tr_keys = case_identifiers
- val_keys = tr_keys
- else:
- splits_file = join(self.preprocessed_dataset_folder_base, "splits_final.json")
- dataset = nnUNetDataset(self.preprocessed_dataset_folder, case_identifiers=None,
- num_images_properties_loading_threshold=0,
- folder_with_segs_from_previous_stage=self.folder_with_segs_from_previous_stage)
- # if the split file does not exist we need to create it
- if not isfile(splits_file):
- self.print_to_log_file("Creating new 5-fold cross-validation split...")
- all_keys_sorted = list(np.sort(list(dataset.keys())))
- splits = generate_crossval_split(all_keys_sorted, seed=12345, n_splits=5)
- save_json(splits, splits_file)
- else:
- self.print_to_log_file("Using splits from existing split file:", splits_file)
- splits = load_json(splits_file)
- self.print_to_log_file(f"The split file contains {len(splits)} splits.")
- self.print_to_log_file("Desired fold for training: %d" % self.fold)
- if self.fold < len(splits):
- tr_keys = splits[self.fold]['train']
- val_keys = splits[self.fold]['val']
- self.print_to_log_file("This split has %d training and %d validation cases."
- % (len(tr_keys), len(val_keys)))
- else:
- self.print_to_log_file("INFO: You requested fold %d for training but splits "
- "contain only %d folds. I am now creating a "
- "random (but seeded) 80:20 split!" % (self.fold, len(splits)))
- # if we request a fold that is not in the split file, create a random 80:20 split
- rnd = np.random.RandomState(seed=12345 + self.fold)
- keys = np.sort(list(dataset.keys()))
- idx_tr = rnd.choice(len(keys), int(len(keys) * 0.8), replace=False)
- idx_val = [i for i in range(len(keys)) if i not in idx_tr]
- tr_keys = [keys[i] for i in idx_tr]
- val_keys = [keys[i] for i in idx_val]
- self.print_to_log_file("This random 80:20 split has %d training and %d validation cases."
- % (len(tr_keys), len(val_keys)))
- if any([i in val_keys for i in tr_keys]):
- self.print_to_log_file('WARNING: Some validation cases are also in the training set. Please check the '
- 'splits.json or ignore if this is intentional.')
- return tr_keys, val_keys
- def get_tr_and_val_datasets(self):
- # create dataset split
- tr_keys, val_keys = self.do_split()
- # load the datasets for training and validation. Note that we always draw random samples so we really don't
- # care about distributing training cases across GPUs.
- dataset_tr = nnUNetDataset(self.preprocessed_dataset_folder, tr_keys,
- folder_with_segs_from_previous_stage=self.folder_with_segs_from_previous_stage,
- num_images_properties_loading_threshold=0)
- dataset_val = nnUNetDataset(self.preprocessed_dataset_folder, val_keys,
- folder_with_segs_from_previous_stage=self.folder_with_segs_from_previous_stage,
- num_images_properties_loading_threshold=0)
- return dataset_tr, dataset_val
- def get_dataloaders(self):
- # we use the patch size to determine whether we need 2D or 3D dataloaders. We also use it to determine whether
- # we need to use dummy 2D augmentation (in case of 3D training) and what our initial patch size should be
- patch_size = self.configuration_manager.patch_size
- dim = len(patch_size)
- # needed for deep supervision: how much do we need to downscale the segmentation targets for the different
- # outputs?
- deep_supervision_scales = self._get_deep_supervision_scales()
- (
- rotation_for_DA,
- do_dummy_2d_data_aug,
- initial_patch_size,
- mirror_axes,
- ) = self.configure_rotation_dummyDA_mirroring_and_inital_patch_size()
- # training pipeline
- tr_transforms = self.get_training_transforms(
- patch_size, rotation_for_DA, deep_supervision_scales, mirror_axes, do_dummy_2d_data_aug,
- order_resampling_data=3, order_resampling_seg=1,
- use_mask_for_norm=self.configuration_manager.use_mask_for_norm,
- is_cascaded=self.is_cascaded, foreground_labels=self.label_manager.foreground_labels,
- regions=self.label_manager.foreground_regions if self.label_manager.has_regions else None,
- ignore_label=self.label_manager.ignore_label)
- # validation pipeline
- val_transforms = self.get_validation_transforms(deep_supervision_scales,
- is_cascaded=self.is_cascaded,
- foreground_labels=self.label_manager.foreground_labels,
- regions=self.label_manager.foreground_regions if
- self.label_manager.has_regions else None,
- ignore_label=self.label_manager.ignore_label)
- dl_tr, dl_val = self.get_plain_dataloaders(initial_patch_size, dim)
- allowed_num_processes = get_allowed_n_proc_DA()
- if allowed_num_processes == 0:
- mt_gen_train = SingleThreadedAugmenter(dl_tr, tr_transforms)
- mt_gen_val = SingleThreadedAugmenter(dl_val, val_transforms)
- else:
- mt_gen_train = LimitedLenWrapper(self.num_iterations_per_epoch, data_loader=dl_tr, transform=tr_transforms,
- num_processes=allowed_num_processes, num_cached=6, seeds=None,
- pin_memory=self.device.type == 'cuda', wait_time=0.02)
- mt_gen_val = LimitedLenWrapper(self.num_val_iterations_per_epoch, data_loader=dl_val,
- transform=val_transforms, num_processes=max(1, allowed_num_processes // 2),
- num_cached=3, seeds=None, pin_memory=self.device.type == 'cuda',
- wait_time=0.02)
- return mt_gen_train, mt_gen_val
- def get_plain_dataloaders(self, initial_patch_size: Tuple[int, ...], dim: int):
- dataset_tr, dataset_val = self.get_tr_and_val_datasets()
- if dim == 2:
- dl_tr = nnUNetDataLoader2D(dataset_tr, self.batch_size,
- initial_patch_size,
- self.configuration_manager.patch_size,
- self.label_manager,
- oversample_foreground_percent=self.oversample_foreground_percent,
- sampling_probabilities=None, pad_sides=None)
- dl_val = nnUNetDataLoader2D(dataset_val, self.batch_size,
- self.configuration_manager.patch_size,
- self.configuration_manager.patch_size,
- self.label_manager,
- oversample_foreground_percent=self.oversample_foreground_percent,
- sampling_probabilities=None, pad_sides=None)
- else:
- dl_tr = nnUNetDataLoader3D(dataset_tr, self.batch_size,
- initial_patch_size,
- self.configuration_manager.patch_size,
- self.label_manager,
- oversample_foreground_percent=self.oversample_foreground_percent,
- sampling_probabilities=None, pad_sides=None)
- dl_val = nnUNetDataLoader3D(dataset_val, self.batch_size,
- self.configuration_manager.patch_size,
- self.configuration_manager.patch_size,
- self.label_manager,
- oversample_foreground_percent=self.oversample_foreground_percent,
- sampling_probabilities=None, pad_sides=None)
- return dl_tr, dl_val
- @staticmethod
- def get_training_transforms(
- patch_size: Union[np.ndarray, Tuple[int]],
- rotation_for_DA: dict,
- deep_supervision_scales: Union[List, Tuple, None],
- mirror_axes: Tuple[int, ...],
- do_dummy_2d_data_aug: bool,
- order_resampling_data: int = 3,
- order_resampling_seg: int = 1,
- border_val_seg: int = -1,
- use_mask_for_norm: List[bool] = None,
- is_cascaded: bool = False,
- foreground_labels: Union[Tuple[int, ...], List[int]] = None,
- regions: List[Union[List[int], Tuple[int, ...], int]] = None,
- ignore_label: int = None,
- ) -> AbstractTransform:
- tr_transforms = []
- if do_dummy_2d_data_aug:
- ignore_axes = (0,)
- tr_transforms.append(Convert3DTo2DTransform())
- patch_size_spatial = patch_size[1:]
- else:
- patch_size_spatial = patch_size
- ignore_axes = None
- tr_transforms.append(SpatialTransform(
- patch_size_spatial, patch_center_dist_from_border=None,
- do_elastic_deform=False, alpha=(0, 0), sigma=(0, 0),
- do_rotation=True, angle_x=rotation_for_DA['x'], angle_y=rotation_for_DA['y'], angle_z=rotation_for_DA['z'],
- p_rot_per_axis=1, # todo experiment with this
- do_scale=True, scale=(0.7, 1.4),
- border_mode_data="constant", border_cval_data=0, order_data=order_resampling_data,
- border_mode_seg="constant", border_cval_seg=border_val_seg, order_seg=order_resampling_seg,
- random_crop=False, # random cropping is part of our dataloaders
- p_el_per_sample=0, p_scale_per_sample=0.2, p_rot_per_sample=0.2,
- independent_scale_for_each_axis=False # todo experiment with this
- ))
- if do_dummy_2d_data_aug:
- tr_transforms.append(Convert2DTo3DTransform())
- tr_transforms.append(GaussianNoiseTransform(p_per_sample=0.1))
- tr_transforms.append(GaussianBlurTransform((0.5, 1.), different_sigma_per_channel=True, p_per_sample=0.2,
- p_per_channel=0.5))
- tr_transforms.append(BrightnessMultiplicativeTransform(multiplier_range=(0.75, 1.25), p_per_sample=0.15))
- tr_transforms.append(ContrastAugmentationTransform(p_per_sample=0.15))
- tr_transforms.append(SimulateLowResolutionTransform(zoom_range=(0.5, 1), per_channel=True,
- p_per_channel=0.5,
- order_downsample=0, order_upsample=3, p_per_sample=0.25,
- ignore_axes=ignore_axes))
- tr_transforms.append(GammaTransform((0.7, 1.5), True, True, retain_stats=True, p_per_sample=0.1))
- tr_transforms.append(GammaTransform((0.7, 1.5), False, True, retain_stats=True, p_per_sample=0.3))
- if mirror_axes is not None and len(mirror_axes) > 0:
- tr_transforms.append(MirrorTransform(mirror_axes))
- if use_mask_for_norm is not None and any(use_mask_for_norm):
- tr_transforms.append(MaskTransform([i for i in range(len(use_mask_for_norm)) if use_mask_for_norm[i]],
- mask_idx_in_seg=0, set_outside_to=0))
- tr_transforms.append(RemoveLabelTransform(-1, 0))
- if is_cascaded:
- assert foreground_labels is not None, 'We need foreground_labels for cascade augmentations'
- tr_transforms.append(MoveSegAsOneHotToData(1, foreground_labels, 'seg', 'data'))
- tr_transforms.append(ApplyRandomBinaryOperatorTransform(
- channel_idx=list(range(-len(foreground_labels), 0)),
- p_per_sample=0.4,
- key="data",
- strel_size=(1, 8),
- p_per_label=1))
- tr_transforms.append(
- RemoveRandomConnectedComponentFromOneHotEncodingTransform(
- channel_idx=list(range(-len(foreground_labels), 0)),
- key="data",
- p_per_sample=0.2,
- fill_with_other_class_p=0,
- dont_do_if_covers_more_than_x_percent=0.15))
- tr_transforms.append(RenameTransform('seg', 'target', True))
- if regions is not None:
- # the ignore label must also be converted
- tr_transforms.append(ConvertSegmentationToRegionsTransform(list(regions) + [ignore_label]
- if ignore_label is not None else regions,
- 'target', 'target'))
- if deep_supervision_scales is not None:
- tr_transforms.append(DownsampleSegForDSTransform2(deep_supervision_scales, 0, input_key='target',
- output_key='target'))
- tr_transforms.append(NumpyToTensor(['data', 'target'], 'float'))
- tr_transforms = Compose(tr_transforms)
- return tr_transforms
- @staticmethod
- def get_validation_transforms(
- deep_supervision_scales: Union[List, Tuple, None],
- is_cascaded: bool = False,
- foreground_labels: Union[Tuple[int, ...], List[int]] = None,
- regions: List[Union[List[int], Tuple[int, ...], int]] = None,
- ignore_label: int = None,
- ) -> AbstractTransform:
- val_transforms = []
- val_transforms.append(RemoveLabelTransform(-1, 0))
- if is_cascaded:
- val_transforms.append(MoveSegAsOneHotToData(1, foreground_labels, 'seg', 'data'))
- val_transforms.append(RenameTransform('seg', 'target', True))
- if regions is not None:
- # the ignore label must also be converted
- val_transforms.append(ConvertSegmentationToRegionsTransform(list(regions) + [ignore_label]
- if ignore_label is not None else regions,
- 'target', 'target'))
- if deep_supervision_scales is not None:
- val_transforms.append(DownsampleSegForDSTransform2(deep_supervision_scales, 0, input_key='target',
- output_key='target'))
- val_transforms.append(NumpyToTensor(['data', 'target'], 'float'))
- val_transforms = Compose(val_transforms)
- return val_transforms
- def set_deep_supervision_enabled(self, enabled: bool):
- """
- This function is specific for the default architecture in nnU-Net. If you change the architecture, there are
- chances you need to change this as well!
- """
- if self.is_ddp:
- mod = self.network.module
- else:
- mod = self.network
- if isinstance(mod, OptimizedModule):
- mod = mod._orig_mod
- mod.decoder.deep_supervision = enabled
- def on_train_start(self):
- if not self.was_initialized:
- self.initialize()
- maybe_mkdir_p(self.output_folder)
- # make sure deep supervision is on in the network
- self.set_deep_supervision_enabled(self.enable_deep_supervision)
- self.print_plans()
- empty_cache(self.device)
- # maybe unpack
- if self.unpack_dataset and self.local_rank == 0:
- self.print_to_log_file('unpacking dataset...')
- unpack_dataset(self.preprocessed_dataset_folder, unpack_segmentation=True, overwrite_existing=False,
- num_processes=max(1, round(get_allowed_n_proc_DA() // 2)))
- self.print_to_log_file('unpacking done...')
- if self.is_ddp:
- dist.barrier()
- # dataloaders must be instantiated here because they need access to the training data which may not be present
- # when doing inference
- self.dataloader_train, self.dataloader_val = self.get_dataloaders()
- # copy plans and dataset.json so that they can be used for restoring everything we need for inference
- save_json(self.plans_manager.plans, join(self.output_folder_base, 'plans.json'), sort_keys=False)
- save_json(self.dataset_json, join(self.output_folder_base, 'dataset.json'), sort_keys=False)
- # we don't really need the fingerprint but its still handy to have it with the others
- shutil.copy(join(self.preprocessed_dataset_folder_base, 'dataset_fingerprint.json'),
- join(self.output_folder_base, 'dataset_fingerprint.json'))
- # produces a pdf in output folder
- self.plot_network_architecture()
- self._save_debug_information()
- # print(f"batch size: {self.batch_size}")
- # print(f"oversample: {self.oversample_foreground_percent}")
- def on_train_end(self):
- # dirty hack because on_epoch_end increments the epoch counter and this is executed afterwards.
- # This will lead to the wrong current epoch to be stored
- self.current_epoch -= 1
- self.save_checkpoint(join(self.output_folder, "checkpoint_final.pth"))
- self.current_epoch += 1
- # now we can delete latest
- if self.local_rank == 0 and isfile(join(self.output_folder, "checkpoint_latest.pth")):
- os.remove(join(self.output_folder, "checkpoint_latest.pth"))
- # shut down dataloaders
- old_stdout = sys.stdout
- with open(os.devnull, 'w') as f:
- sys.stdout = f
- if self.dataloader_train is not None and \
- isinstance(self.dataloader_train, (NonDetMultiThreadedAugmenter, MultiThreadedAugmenter)):
- self.dataloader_train._finish()
- if self.dataloader_val is not None and \
- isinstance(self.dataloader_train, (NonDetMultiThreadedAugmenter, MultiThreadedAugmenter)):
- self.dataloader_val._finish()
- sys.stdout = old_stdout
- empty_cache(self.device)
- self.print_to_log_file("Training done.")
- def on_train_epoch_start(self):
- self.network.train()
- self.lr_scheduler.step(self.current_epoch)
- self.print_to_log_file('')
- self.print_to_log_file(f'Epoch {self.current_epoch}')
- self.print_to_log_file(
- f"Current learning rate: {np.round(self.optimizer.param_groups[0]['lr'], decimals=5)}")
- # lrs are the same for all workers so we don't need to gather them in case of DDP training
- self.logger.log('lrs', self.optimizer.param_groups[0]['lr'], self.current_epoch)
- def train_step(self, batch: dict) -> dict:
- data = batch['data']
- target = batch['target']
- data = data.to(self.device, non_blocking=True)
- if isinstance(target, list):
- target = [i.to(self.device, non_blocking=True) for i in target]
- else:
- target = target.to(self.device, non_blocking=True)
- self.optimizer.zero_grad(set_to_none=True)
- # Autocast can be annoying
- # If the device_type is 'cpu' then it's slow as heck and needs to be disabled.
- # 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)
- # So autocast will only be active if we have a cuda device.
- with autocast(self.device.type, enabled=True) if self.device.type == 'cuda' else dummy_context():
- output = self.network(data)
- # del data
- l = self.loss(output, target)
- if self.grad_scaler is not None:
- self.grad_scaler.scale(l).backward()
- self.grad_scaler.unscale_(self.optimizer)
- torch.nn.utils.clip_grad_norm_(self.network.parameters(), 12)
- self.grad_scaler.step(self.optimizer)
- self.grad_scaler.update()
- else:
- l.backward()
- torch.nn.utils.clip_grad_norm_(self.network.parameters(), 12)
- self.optimizer.step()
- return {'loss': l.detach().cpu().numpy()}
- def on_train_epoch_end(self, train_outputs: List[dict]):
- outputs = collate_outputs(train_outputs)
- if self.is_ddp:
- losses_tr = [None for _ in range(dist.get_world_size())]
- dist.all_gather_object(losses_tr, outputs['loss'])
- loss_here = np.vstack(losses_tr).mean()
- else:
- loss_here = np.mean(outputs['loss'])
- self.logger.log('train_losses', loss_here, self.current_epoch)
- def on_validation_epoch_start(self):
- self.network.eval()
- def validation_step(self, batch: dict) -> dict:
- data = batch['data']
- target = batch['target']
- data = data.to(self.device, non_blocking=True)
- if isinstance(target, list):
- target = [i.to(self.device, non_blocking=True) for i in target]
- else:
- target = target.to(self.device, non_blocking=True)
- # Autocast can be annoying
- # If the device_type is 'cpu' then it's slow as heck and needs to be disabled.
- # 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)
- # So autocast will only be active if we have a cuda device.
- with autocast(self.device.type, enabled=True) if self.device.type == 'cuda' else dummy_context():
- output = self.network(data)
- del data
- l = self.loss(output, target)
- # we only need the output with the highest output resolution (if DS enabled)
- if self.enable_deep_supervision:
- output = output[0]
- target = target[0]
- # the following is needed for online evaluation. Fake dice (green line)
- axes = [0] + list(range(2, output.ndim))
- if self.label_manager.has_regions:
- predicted_segmentation_onehot = (torch.sigmoid(output) > 0.5).long()
- else:
- # no need for softmax
- output_seg = output.argmax(1)[:, None]
- predicted_segmentation_onehot = torch.zeros(output.shape, device=output.device, dtype=torch.float32)
- predicted_segmentation_onehot.scatter_(1, output_seg, 1)
- del output_seg
- if self.label_manager.has_ignore_label:
- if not self.label_manager.has_regions:
- mask = (target != self.label_manager.ignore_label).float()
- # CAREFUL that you don't rely on target after this line!
- target[target == self.label_manager.ignore_label] = 0
- else:
- mask = 1 - target[:, -1:]
- # CAREFUL that you don't rely on target after this line!
- target = target[:, :-1]
- else:
- mask = None
- tp, fp, fn, _ = get_tp_fp_fn_tn(predicted_segmentation_onehot, target, axes=axes, mask=mask)
- tp_hard = tp.detach().cpu().numpy()
- fp_hard = fp.detach().cpu().numpy()
- fn_hard = fn.detach().cpu().numpy()
- if not self.label_manager.has_regions:
- # if we train with regions all segmentation heads predict some kind of foreground. In conventional
- # (softmax training) there needs tobe one output for the background. We are not interested in the
- # background Dice
- # [1:] in order to remove background
- tp_hard = tp_hard[1:]
- fp_hard = fp_hard[1:]
- fn_hard = fn_hard[1:]
- return {'loss': l.detach().cpu().numpy(), 'tp_hard': tp_hard, 'fp_hard': fp_hard, 'fn_hard': fn_hard}
- def on_validation_epoch_end(self, val_outputs: List[dict]):
- outputs_collated = collate_outputs(val_outputs)
- tp = np.sum(outputs_collated['tp_hard'], 0)
- fp = np.sum(outputs_collated['fp_hard'], 0)
- fn = np.sum(outputs_collated['fn_hard'], 0)
- if self.is_ddp:
- world_size = dist.get_world_size()
- tps = [None for _ in range(world_size)]
- dist.all_gather_object(tps, tp)
- tp = np.vstack([i[None] for i in tps]).sum(0)
- fps = [None for _ in range(world_size)]
- dist.all_gather_object(fps, fp)
- fp = np.vstack([i[None] for i in fps]).sum(0)
- fns = [None for _ in range(world_size)]
- dist.all_gather_object(fns, fn)
- fn = np.vstack([i[None] for i in fns]).sum(0)
- losses_val = [None for _ in range(world_size)]
- dist.all_gather_object(losses_val, outputs_collated['loss'])
- loss_here = np.vstack(losses_val).mean()
- else:
- loss_here = np.mean(outputs_collated['loss'])
- global_dc_per_class = [i for i in [2 * i / (2 * i + j + k) for i, j, k in zip(tp, fp, fn)]]
- mean_fg_dice = np.nanmean(global_dc_per_class)
- self.logger.log('mean_fg_dice', mean_fg_dice, self.current_epoch)
- self.logger.log('dice_per_class_or_region', global_dc_per_class, self.current_epoch)
- self.logger.log('val_losses', loss_here, self.current_epoch)
- def on_epoch_start(self):
- self.logger.log('epoch_start_timestamps', time(), self.current_epoch)
- def on_epoch_end(self):
- self.logger.log('epoch_end_timestamps', time(), self.current_epoch)
- self.print_to_log_file('train_loss', np.round(self.logger.my_fantastic_logging['train_losses'][-1], decimals=4))
- self.print_to_log_file('val_loss', np.round(self.logger.my_fantastic_logging['val_losses'][-1], decimals=4))
- self.print_to_log_file('Pseudo dice', [np.round(i, decimals=4) for i in
- self.logger.my_fantastic_logging['dice_per_class_or_region'][-1]])
- self.print_to_log_file(
- 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")
- # handling periodic checkpointing
- current_epoch = self.current_epoch
- if (current_epoch + 1) % self.save_every == 0 and current_epoch != (self.num_epochs - 1):
- self.save_checkpoint(join(self.output_folder, 'checkpoint_latest.pth'))
- # handle 'best' checkpointing. ema_fg_dice is computed by the logger and can be accessed like this
- if self._best_ema is None or self.logger.my_fantastic_logging['ema_fg_dice'][-1] > self._best_ema:
- self._best_ema = self.logger.my_fantastic_logging['ema_fg_dice'][-1]
- self.print_to_log_file(f"Yayy! New best EMA pseudo Dice: {np.round(self._best_ema, decimals=4)}")
- self.save_checkpoint(join(self.output_folder, 'checkpoint_best.pth'))
- if self.local_rank == 0:
- self.logger.plot_progress_png(self.output_folder)
- self.current_epoch += 1
- def save_checkpoint(self, filename: str) -> None:
- if self.local_rank == 0:
- if not self.disable_checkpointing:
- if self.is_ddp:
- mod = self.network.module
- else:
- mod = self.network
- if isinstance(mod, OptimizedModule):
- mod = mod._orig_mod
- checkpoint = {
- 'network_weights': mod.state_dict(),
- 'optimizer_state': self.optimizer.state_dict(),
- 'grad_scaler_state': self.grad_scaler.state_dict() if self.grad_scaler is not None else None,
- 'logging': self.logger.get_checkpoint(),
- '_best_ema': self._best_ema,
- 'current_epoch': self.current_epoch + 1,
- 'init_args': self.my_init_kwargs,
- 'trainer_name': self.__class__.__name__,
- 'inference_allowed_mirroring_axes': self.inference_allowed_mirroring_axes,
- }
- torch.save(checkpoint, filename)
- else:
- self.print_to_log_file('No checkpoint written, checkpointing is disabled')
- def load_checkpoint(self, filename_or_checkpoint: Union[dict, str]) -> None:
- if not self.was_initialized:
- self.initialize()
- if isinstance(filename_or_checkpoint, str):
- checkpoint = torch.load(filename_or_checkpoint, map_location=self.device)
- # if state dict comes from nn.DataParallel but we use non-parallel model here then the state dict keys do not
- # match. Use heuristic to make it match
- new_state_dict = {}
- for k, value in checkpoint['network_weights'].items():
- key = k
- if key not in self.network.state_dict().keys() and key.startswith('module.'):
- key = key[7:]
- new_state_dict[key] = value
- self.my_init_kwargs = checkpoint['init_args']
- self.current_epoch = checkpoint['current_epoch']
- self.logger.load_checkpoint(checkpoint['logging'])
- self._best_ema = checkpoint['_best_ema']
- self.inference_allowed_mirroring_axes = checkpoint[
- 'inference_allowed_mirroring_axes'] if 'inference_allowed_mirroring_axes' in checkpoint.keys() else self.inference_allowed_mirroring_axes
- # messing with state dict naming schemes. Facepalm.
- if self.is_ddp:
- if isinstance(self.network.module, OptimizedModule):
- self.network.module._orig_mod.load_state_dict(new_state_dict)
- else:
- self.network.module.load_state_dict(new_state_dict)
- else:
- if isinstance(self.network, OptimizedModule):
- self.network._orig_mod.load_state_dict(new_state_dict)
- else:
- self.network.load_state_dict(new_state_dict)
- self.optimizer.load_state_dict(checkpoint['optimizer_state'])
- if self.grad_scaler is not None:
- if checkpoint['grad_scaler_state'] is not None:
- self.grad_scaler.load_state_dict(checkpoint['grad_scaler_state'])
- def perform_actual_validation(self, save_probabilities: bool = False):
- self.set_deep_supervision_enabled(False)
- self.network.eval()
- if self.is_ddp and self.batch_size == 1 and self.enable_deep_supervision and self._do_i_compile():
- self.print_to_log_file("WARNING! batch size is 1 during training and torch.compile is enabled. If you "
- "encounter crashes in validation then this is because torch.compile forgets "
- "to trigger a recompilation of the model with deep supervision disabled. "
- "This causes torch.flip to complain about getting a tuple as input. Just rerun the "
- "validation with --val (exactly the same as before) and then it will work. "
- "Why? Because --val triggers nnU-Net to ONLY run validation meaning that the first "
- "forward pass (where compile is triggered) already has deep supervision disabled. "
- "This is exactly what we need in perform_actual_validation")
- predictor = nnUNetPredictor(tile_step_size=0.5, use_gaussian=True, use_mirroring=True,
- perform_everything_on_device=True, device=self.device, verbose=False,
- verbose_preprocessing=False, allow_tqdm=False)
- predictor.manual_initialization(self.network, self.plans_manager, self.configuration_manager, None,
- self.dataset_json, self.__class__.__name__,
- self.inference_allowed_mirroring_axes)
- with multiprocessing.get_context("spawn").Pool(default_num_processes) as segmentation_export_pool:
- worker_list = [i for i in segmentation_export_pool._pool]
- validation_output_folder = join(self.output_folder, 'validation')
- maybe_mkdir_p(validation_output_folder)
- # we cannot use self.get_tr_and_val_datasets() here because we might be DDP and then we have to distribute
- # the validation keys across the workers.
- _, val_keys = self.do_split()
- if self.is_ddp:
- last_barrier_at_idx = len(val_keys) // dist.get_world_size() - 1
- val_keys = val_keys[self.local_rank:: dist.get_world_size()]
- # we cannot just have barriers all over the place because the number of keys each GPU receives can be
- # different
- dataset_val = nnUNetDataset(self.preprocessed_dataset_folder, val_keys,
- folder_with_segs_from_previous_stage=self.folder_with_segs_from_previous_stage,
- num_images_properties_loading_threshold=0)
- next_stages = self.configuration_manager.next_stage_names
- if next_stages is not None:
- _ = [maybe_mkdir_p(join(self.output_folder_base, 'predicted_next_stage', n)) for n in next_stages]
- results = []
- for i, k in enumerate(dataset_val.keys()):
- proceed = not check_workers_alive_and_busy(segmentation_export_pool, worker_list, results,
- allowed_num_queued=2)
- while not proceed:
- sleep(0.1)
- proceed = not check_workers_alive_and_busy(segmentation_export_pool, worker_list, results,
- allowed_num_queued=2)
- self.print_to_log_file(f"predicting {k}")
- data, seg, properties = dataset_val.load_case(k)
- if self.is_cascaded:
- data = np.vstack((data, convert_labelmap_to_one_hot(seg[-1], self.label_manager.foreground_labels,
- output_dtype=data.dtype)))
- with warnings.catch_warnings():
- # ignore 'The given NumPy array is not writable' warning
- warnings.simplefilter("ignore")
- data = torch.from_numpy(data)
- self.print_to_log_file(f'{k}, shape {data.shape}, rank {self.local_rank}')
- output_filename_truncated = join(validation_output_folder, k)
- prediction = predictor.predict_sliding_window_return_logits(data)
- prediction = prediction.cpu()
- # this needs to go into background processes
- results.append(
- segmentation_export_pool.starmap_async(
- export_prediction_from_logits, (
- (prediction, properties, self.configuration_manager, self.plans_manager,
- self.dataset_json, output_filename_truncated, save_probabilities),
- )
- )
- )
- # for debug purposes
- # export_prediction(prediction_for_export, properties, self.configuration, self.plans, self.dataset_json,
- # output_filename_truncated, save_probabilities)
- # if needed, export the softmax prediction for the next stage
- if next_stages is not None:
- for n in next_stages:
- next_stage_config_manager = self.plans_manager.get_configuration(n)
- expected_preprocessed_folder = join(nnUNet_preprocessed, self.plans_manager.dataset_name,
- next_stage_config_manager.data_identifier)
- try:
- # we do this so that we can use load_case and do not have to hard code how loading training cases is implemented
- tmp = nnUNetDataset(expected_preprocessed_folder, [k],
- num_images_properties_loading_threshold=0)
- d, s, p = tmp.load_case(k)
- except FileNotFoundError:
- self.print_to_log_file(
- f"Predicting next stage {n} failed for case {k} because the preprocessed file is missing! "
- f"Run the preprocessing for this configuration first!")
- continue
- target_shape = d.shape[1:]
- output_folder = join(self.output_folder_base, 'predicted_next_stage', n)
- output_file = join(output_folder, k + '.npz')
- # resample_and_save(prediction, target_shape, output_file, self.plans_manager, self.configuration_manager, properties,
- # self.dataset_json)
- results.append(segmentation_export_pool.starmap_async(
- resample_and_save, (
- (prediction, target_shape, output_file, self.plans_manager,
- self.configuration_manager,
- properties,
- self.dataset_json),
- )
- ))
- # if we don't barrier from time to time we will get nccl timeouts for large datasets. Yuck.
- if self.is_ddp and i < last_barrier_at_idx and (i + 1) % 20 == 0:
- dist.barrier()
- _ = [r.get() for r in results]
- if self.is_ddp:
- dist.barrier()
- if self.local_rank == 0:
- metrics = compute_metrics_on_folder(join(self.preprocessed_dataset_folder_base, 'gt_segmentations'),
- validation_output_folder,
- join(validation_output_folder, 'summary.json'),
- self.plans_manager.image_reader_writer_class(),
- self.dataset_json["file_ending"],
- self.label_manager.foreground_regions if self.label_manager.has_regions else
- self.label_manager.foreground_labels,
- self.label_manager.ignore_label, chill=True,
- num_processes=default_num_processes * dist.get_world_size() if
- self.is_ddp else default_num_processes)
- self.print_to_log_file("Validation complete", also_print_to_console=True)
- self.print_to_log_file("Mean Validation Dice: ", (metrics['foreground_mean']["Dice"]), also_print_to_console=True)
- self.set_deep_supervision_enabled(True)
- compute_gaussian.cache_clear()
- def run_training(self):
- self.on_train_start()
- for epoch in range(self.current_epoch, self.num_epochs):
- self.on_epoch_start()
- self.on_train_epoch_start()
- train_outputs = []
- for batch_id in range(self.num_iterations_per_epoch):
- train_outputs.append(self.train_step(next(self.dataloader_train)))
- self.on_train_epoch_end(train_outputs)
- with torch.no_grad():
- self.on_validation_epoch_start()
- val_outputs = []
- for batch_id in range(self.num_val_iterations_per_epoch):
- val_outputs.append(self.validation_step(next(self.dataloader_val)))
- self.on_validation_epoch_end(val_outputs)
- self.on_epoch_end()
- self.on_train_end()
nnUNetTrainer.py at commit 25225a8, under Apache-2.0 · at the source
Overview
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/
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/
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://
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
202f6baa0adc2ef5f7b615df19cc4da970412cc0, 25 September 2026Availability: 1 check, the latest on 26 September 2026: the link answers
- 26 September 2026: the link answers
225 files
- .github/
scripts/ , Shell, 46 linessafe-label.sh - .github/
scripts/ , Shell, 29 linessafe-pr-review.sh - documentation/
__init__.py , Python, 1 line - documentation/
competitions/ , Python, 1 lineFLARE24/ Task_1/ __init__.py - documentation/
competitions/ , Python, 209 linesFLARE24/ Task_1/ inference_flare_task1.py - documentation/
competitions/ , Python, 1 lineFLARE24/ Task_2/ __init__.py - documentation/
competitions/ , Python, 467 linesFLARE24/ Task_2/ inference_flare_task2.py - documentation/
competitions/ , Python, 1 lineFLARE24/ __init__.py - documentation/
competitions/ , Python, 1 lineToothfairy2/ __init__.py - documentation/
competitions/ , Python, 359 linesToothfairy2/ inference_script_semseg_ only_customInf2.py - documentation/
competitions/ , Python, 1 line__init__.py - nnunetv2/
__init__.py , Python, 1 line - nnunetv2/
batch_running/ , Python, 1 line__init__.py - nnunetv2/
batch_running/ , Python, 1 linebenchmarking/ __init__.py - nnunetv2/
batch_running/ , Python, 41 linesbenchmarking/ generate_benchmarking_co mmands.py - nnunetv2/
batch_running/ , Python, 69 linesbenchmarking/ summarize_benchmark_resu lts.py - nnunetv2/
batch_running/ , Python, 121 linescollect_results_custom_D ecathlon.py - nnunetv2/
batch_running/ , Python, 17 linescollect_results_custom_D ecathlon_2d.py - nnunetv2/
batch_running/ , Python, 111 linesgenerate_lsf_runs_custom Decathlon.py - nnunetv2/
batch_running/ , Shell, 52 linesjobs.sh - nnunetv2/
batch_running/ , Python, 1 linerelease_trainings/ __init__.py - nnunetv2/
batch_running/ , Python, 1 linerelease_trainings/ nnunetv2_v1/ __init__.py - nnunetv2/
batch_running/ , Python, 112 linesrelease_trainings/ nnunetv2_v1/ collect_results.py - nnunetv2/
batch_running/ , Python, 91 linesrelease_trainings/ nnunetv2_v1/ generate_lsf_commands.py - nnunetv2/
configuration.py , Python, 10 lines - nnunetv2/
dataset_conversion/ , Python, 147 linesDataset015_018_RibFrac_R ibSeg.py - nnunetv2/
dataset_conversion/ , Python, 70 linesDataset021_CTAAorta.py - nnunetv2/
dataset_conversion/ , Python, 55 linesDataset023_AbdomenAtlas1 _1Mini.py - nnunetv2/
dataset_conversion/ , Python, 114 linesDataset027_ACDC.py - nnunetv2/
dataset_conversion/ , Python, 110 linesDataset042_BraTS18.py - nnunetv2/
dataset_conversion/ , Python, 110 linesDataset043_BraTS19.py - nnunetv2/
dataset_conversion/ , Python, 84 linesDataset073_Fluo_C3DH_A54 9_SIM.py - nnunetv2/
dataset_conversion/ , Python, 198 linesDataset114_MNMs.py - nnunetv2/
dataset_conversion/ , Python, 61 linesDataset115_EMIDEC.py - nnunetv2/
dataset_conversion/ , Python, 196 linesDataset119_ToothFairy2_A ll.py - nnunetv2/
dataset_conversion/ , Python, 86 linesDataset120_RoadSegmentat ion.py - nnunetv2/
dataset_conversion/ , Python, 97 linesDataset137_BraTS21.py - nnunetv2/
dataset_conversion/ , Python, 68 linesDataset218_Amos2022_task 1.py - nnunetv2/
dataset_conversion/ , Python, 64 linesDataset219_Amos2022_task 2.py - nnunetv2/
dataset_conversion/ , Python, 49 linesDataset220_KiTS2023.py - nnunetv2/
dataset_conversion/ , Python, 70 linesDataset221_AutoPETII_202 3.py - nnunetv2/
dataset_conversion/ , Python, 59 linesDataset223_AMOS2022postC hallenge.py - nnunetv2/
dataset_conversion/ , Python, 60 linesDataset224_AbdomenAtlas1 .0.py - nnunetv2/
dataset_conversion/ , Python, 55 linesDataset226_BraTS2024-Bra TS-GLI.py - nnunetv2/
dataset_conversion/ , Python, 54 linesDataset227_TotalSegmenta torMRI.py - nnunetv2/
dataset_conversion/ , Python, 32 linesDataset987_dummyDataset4 .py - nnunetv2/
dataset_conversion/ , Python, 32 linesDataset989_dummyDataset4 _2.py - nnunetv2/
dataset_conversion/ , Python, 1 line__init__.py - nnunetv2/
dataset_conversion/ , Python, 131 linesconvert_MSD_dataset.py - nnunetv2/
dataset_conversion/ , Python, 53 linesconvert_raw_dataset_from _old_nnunet_format.py - nnunetv2/
dataset_conversion/ , Python, 73 linesdatasets_for_integration _tests/ Dataset996_IntegrationTe st_Hippocampus_regions_i gnore.py - nnunetv2/
dataset_conversion/ , Python, 37 linesdatasets_for_integration _tests/ Dataset997_IntegrationTe st_Hippocampus_regions.p y - nnunetv2/
dataset_conversion/ , Python, 33 linesdatasets_for_integration _tests/ Dataset998_IntegrationTe st_Hippocampus_ignore.py - nnunetv2/
dataset_conversion/ , Python, 27 linesdatasets_for_integration _tests/ Dataset999_IntegrationTe st_Hippocampus.py - nnunetv2/
dataset_conversion/ , Python, 1 linedatasets_for_integration _tests/ __init__.py - nnunetv2/
dataset_conversion/ , Python, 111 linesgenerate_dataset_json.py - nnunetv2/
ensembling/ , Python, 1 line__init__.py - nnunetv2/
ensembling/ , Python, 207 linesensemble.py - nnunetv2/
evaluation/ , Python, 1 line__init__.py - nnunetv2/
evaluation/ , Python, 58 linesaccumulate_cv_results.py - nnunetv2/
evaluation/ , Python, 262 linesevaluate_predictions.py - nnunetv2/
evaluation/ , Python, 339 linesfind_best_configuration. py - nnunetv2/
experiment_planning/ , Python, 1 line__init__.py - nnunetv2/
experiment_planning/ , Python, 1 linedataset_fingerprint/ __init__.py - nnunetv2/
experiment_planning/ , Python, 211 linesdataset_fingerprint/ fingerprint_extractor.py - nnunetv2/
experiment_planning/ , Python, 1 lineexperiment_planners/ __init__.py - nnunetv2/
experiment_planning/ , Python, 603 linesexperiment_planners/ default_experiment_plann er.py - nnunetv2/
experiment_planning/ , Python, 108 linesexperiment_planners/ network_topology.py - nnunetv2/
experiment_planning/ , Python, 1 lineexperiment_planners/ resampling/ __init__.py - nnunetv2/
experiment_planning/ , Python, 54 linesexperiment_planners/ resampling/ planners_no_resampling.p y - nnunetv2/
experiment_planning/ , Python, 181 linesexperiment_planners/ resampling/ resample_with_torch.py - nnunetv2/
experiment_planning/ , Python, 26 linesexperiment_planners/ resencUNet_planner.py - nnunetv2/
experiment_planning/ , Python, 1 lineexperiment_planners/ residual_unets/ __init__.py - nnunetv2/
experiment_planning/ , Python, 291 linesexperiment_planners/ residual_unets/ residual_encoder_unet_pl anners.py - nnunetv2/
experiment_planning/ , Python, 512 lineslike_nnssl.py - nnunetv2/
experiment_planning/ , Python, 171 linesplan_and_preprocess_api. py - nnunetv2/
experiment_planning/ , Python, 235 linesplan_and_preprocess_entr ypoints.py - nnunetv2/
experiment_planning/ , Python, 1 lineplans_for_pretraining/ __init__.py - nnunetv2/
experiment_planning/ , Python, 82 linesplans_for_pretraining/ move_plans_between_datas ets.py - nnunetv2/
experiment_planning/ , Python, 239 linesverify_dataset_integrity .py - nnunetv2/
imageio/ , Python, 1 line__init__.py - nnunetv2/
imageio/ , Python, 107 linesbase_reader_writer.py - nnunetv2/
imageio/ , Python, 81 linesnatural_image_reader_wri ter.py - nnunetv2/
imageio/ , Python, 222 linesnibabel_reader_writer.py - nnunetv2/
imageio/ , Python, 88 linesreader_writer_registry.p y - nnunetv2/
imageio/ , Python, 234 linessimpleitk_reader_writer. py - nnunetv2/
imageio/ , Python, 100 linestif_reader_writer.py - nnunetv2/
inference/ , Python, 197 linesJHU_inference.py - nnunetv2/
inference/ , Python, 1 line__init__.py - nnunetv2/
inference/ , Python, 315 linesdata_iterators.py - nnunetv2/
inference/ , Python, 100 linesexamples.py - nnunetv2/
inference/ , Python, 160 linesexport_prediction.py - nnunetv2/
inference/ , Python, 1,099 linespredict_from_raw_data.py - nnunetv2/
inference/ , Python, 65 linessliding_window_predictio n.py - nnunetv2/
model_sharing/ , Python, 1 line__init__.py - nnunetv2/
model_sharing/ , Python, 61 linesentry_points.py - nnunetv2/
model_sharing/ , Python, 39 linesmodel_download.py - nnunetv2/
model_sharing/ , Python, 124 linesmodel_export.py - nnunetv2/
model_sharing/ , Python, 8 linesmodel_import.py - nnunetv2/
paths.py , Python, 79 lines - nnunetv2/
postprocessing/ , Python, 1 line__init__.py - nnunetv2/
postprocessing/ , Python, 361 linesremove_connected_compone nts.py - nnunetv2/
preprocessing/ , Python, 1 line__init__.py - nnunetv2/
preprocessing/ , Python, 1 linecropping/ __init__.py - nnunetv2/
preprocessing/ , Python, 39 linescropping/ cropping.py - nnunetv2/
preprocessing/ , Python, 1 linenormalization/ __init__.py - nnunetv2/
preprocessing/ , Python, 104 linesnormalization/ default_normalization_sc hemes.py - nnunetv2/
preprocessing/ , Python, 24 linesnormalization/ map_channel_name_to_norm alization.py - nnunetv2/
preprocessing/ , Python, 1 linepreprocessors/ __init__.py - nnunetv2/
preprocessing/ , Python, 492 linespreprocessors/ default_preprocessor.py - nnunetv2/
preprocessing/ , Python, 1 lineresampling/ __init__.py - nnunetv2/
preprocessing/ , Python, 231 linesresampling/ default_resampling.py - nnunetv2/
preprocessing/ , Python, 13 linesresampling/ no_resampling.py - nnunetv2/
preprocessing/ , Python, 174 linesresampling/ resample_torch.py - nnunetv2/
preprocessing/ , Python, 15 linesresampling/ utils.py - nnunetv2/
preprocessing/ , Python, 1 linesampling_locations/ __init__.py - nnunetv2/
preprocessing/ , Python, 171 linessampling_locations/ extract_sampling_locatio ns.py - nnunetv2/
run/ , Python, 1 line__init__.py - nnunetv2/
run/ , Python, 73 linesload_pretrained_weights. py - nnunetv2/
run/ , Python, 351 linesrun_training.py - nnunetv2/
run/ , Python, 252 linesrun_training_from_pretra ined.py - nnunetv2/
tests/ , Python, 1 line__init__.py - nnunetv2/
tests/ , Python, 1 lineintegration_tests/ __init__.py - nnunetv2/
tests/ , Python, 42 linesintegration_tests/ add_lowres_and_cascade.p y - nnunetv2/
tests/ , Python, 18 linesintegration_tests/ cleanup_integration_test .py - nnunetv2/
tests/ , Shell, 10 linesintegration_tests/ lsf_commands.sh - nnunetv2/
tests/ , Shell, 18 linesintegration_tests/ prepare_integration_test s.sh - nnunetv2/
tests/ , Shell, 27 linesintegration_tests/ run_integration_test.sh - nnunetv2/
tests/ , Python, 75 linesintegration_tests/ run_integration_test_bes tconfig_inference.py - nnunetv2/
tests/ , Shell, 1 lineintegration_tests/ run_integration_test_tra iningOnly_DDP.sh - nnunetv2/
tests/ , Python, 46 linesintegration_tests/ run_nnunet_inference.py - nnunetv2/
tests/ , Python, 53 linestest_copy_file_if_newer. py - nnunetv2/
tests/ , Python, 140 linestest_find_objects.py - nnunetv2/
tests/ , Python, 521 linestest_foreground_location s.py - nnunetv2/
tests/ , Python, 148 linestest_natural_image_tiff_ compression.py - nnunetv2/
tests/ , Python, 45 linestest_paths.py - nnunetv2/
tests/ , Python, 273 linestest_resampling.py - nnunetv2/
training/ , Python, 1 line__init__.py - nnunetv2/
training/ , Python, 1 linedata_augmentation/ __init__.py - nnunetv2/
training/ , Python, 24 linesdata_augmentation/ compute_initial_patch_si ze.py - nnunetv2/
training/ , Python, 1 linedata_augmentation/ custom_transforms/ __init__.py - nnunetv2/
training/ , Python, 136 linesdata_augmentation/ custom_transforms/ cascade_transforms.py - nnunetv2/
training/ , Python, 55 linesdata_augmentation/ custom_transforms/ deep_supervision_donwsam pling.py - nnunetv2/
training/ , Python, 22 linesdata_augmentation/ custom_transforms/ masking.py - nnunetv2/
training/ , Python, 31 linesdata_augmentation/ custom_transforms/ region_based_training.py - nnunetv2/
training/ , Python, 45 linesdata_augmentation/ custom_transforms/ transforms_for_dummy_2d. py - nnunetv2/
training/ , Python, 1 linedataloading/ __init__.py - nnunetv2/
training/ , Python, 217 linesdataloading/ data_loader.py - nnunetv2/
training/ , Python, 586 linesdataloading/ foreground_locations.py - nnunetv2/
training/ , Python, 530 linesdataloading/ nnunet_dataset.py - nnunetv2/
training/ , Python, 71 linesdataloading/ utils.py - nnunetv2/
training/ , Python, 1 linelogging/ __init__.py - nnunetv2/
training/ , Python, 303 lineslogging/ nnunet_logger.py - nnunetv2/
training/ , Python, 1 lineloss/ __init__.py - nnunetv2/
training/ , Python, 156 linesloss/ compound_losses.py - nnunetv2/
training/ , Python, 29 linesloss/ deep_supervision.py - nnunetv2/
training/ , Python, 200 linesloss/ dice.py - nnunetv2/
training/ , Python, 32 linesloss/ robust_ce_loss.py - nnunetv2/
training/ , Python, 1 linelr_scheduler/ __init__.py - nnunetv2/
training/ , Python, 26 lineslr_scheduler/ polylr.py - nnunetv2/
training/ , Python, 143 lineslr_scheduler/ warmup.py - nnunetv2/
training/ , Python, 1 linennUNetTrainer/ __init__.py - nnunetv2/
training/ , Python, 1,505 linesnnUNetTrainer/ nnUNetTrainer.py - nnunetv2/
training/ , Python, 1 linennUNetTrainer/ pretraining/ __init__.py - nnunetv2/
training/ , Python, 455 linesnnUNetTrainer/ pretraining/ dynamicPretrainedTrainer .py - nnunetv2/
training/ , Python, 760 linesnnUNetTrainer/ pretraining/ pretrainedTrainer.py - nnunetv2/
training/ , Python, 287 linesnnUNetTrainer/ pretraining/ thrp_primusx_finetuning. py - nnunetv2/
training/ , Python, 1 linennUNetTrainer/ primus/ __init__.py - nnunetv2/
training/ , Python, 511 linesnnUNetTrainer/ primus/ primus_trainers.py - nnunetv2/
training/ , Python, 1 linennUNetTrainer/ variants/ __init__.py - nnunetv2/
training/ , Python, 1 linennUNetTrainer/ variants/ benchmarking/ __init__.py - nnunetv2/
training/ , Python, 67 linesnnUNetTrainer/ variants/ benchmarking/ nnUNetTrainerBenchmark_5 epochs.py - nnunetv2/
training/ , Python, 68 linesnnUNetTrainer/ variants/ benchmarking/ nnUNetTrainerBenchmark_5 epochs_noDataLoading.py - nnunetv2/
training/ , Python, 1 linennUNetTrainer/ variants/ competitions/ __init__.py - nnunetv2/
training/ , Python, 5 linesnnUNetTrainer/ variants/ competitions/ aortaseg24.py - nnunetv2/
training/ , Python, 1 linennUNetTrainer/ variants/ data_augmentation/ __init__.py - nnunetv2/
training/ , Python, 390 linesnnUNetTrainer/ variants/ data_augmentation/ nnUNetTrainerDA5.py - nnunetv2/
training/ , Python, 191 linesnnUNetTrainer/ variants/ data_augmentation/ nnUNetTrainerDAOrd0.py - nnunetv2/
training/ , Python, 34 linesnnUNetTrainer/ variants/ data_augmentation/ nnUNetTrainerNoDA.py - nnunetv2/
training/ , Python, 215 lines, 1 matchnnUNetTrainer/ variants/ data_augmentation/ nnUNetTrainerNoMirroring .py - nnunetv2/
training/ , Python, 35 linesnnUNetTrainer/ variants/ data_augmentation/ nnUNetTrainer_noDummy2DD A.py - nnunetv2/
training/ , Python, 1 linennUNetTrainer/ variants/ loss/ __init__.py - nnunetv2/
training/ , Python, 40 linesnnUNetTrainer/ variants/ loss/ nnUNetTrainerCELoss.py - nnunetv2/
training/ , Python, 59 linesnnUNetTrainer/ variants/ loss/ nnUNetTrainerDiceLoss.py - nnunetv2/
training/ , Python, 76 linesnnUNetTrainer/ variants/ loss/ nnUNetTrainerTopkLoss.py - nnunetv2/
training/ , Python, 1 linennUNetTrainer/ variants/ lr_schedule/ __init__.py - nnunetv2/
training/ , Python, 12 linesnnUNetTrainer/ variants/ lr_schedule/ nnUNetTrainerCosAnneal.p y - nnunetv2/
training/ , Python, 128 linesnnUNetTrainer/ variants/ lr_schedule/ nnUNetTrainer_warmup.py - nnunetv2/
training/ , Python, 1 linennUNetTrainer/ variants/ network_architecture/ __init__.py - nnunetv2/
training/ , Python, 29 linesnnUNetTrainer/ variants/ network_architecture/ nnUNetTrainerBN.py - nnunetv2/
training/ , Python, 15 linesnnUNetTrainer/ variants/ network_architecture/ nnUNetTrainerNoDeepSuper vision.py - nnunetv2/
training/ , Python, 1 linennUNetTrainer/ variants/ optimizer/ __init__.py - nnunetv2/
training/ , Python, 58 linesnnUNetTrainer/ variants/ optimizer/ nnUNetTrainerAdam.py - nnunetv2/
training/ , Python, 65 linesnnUNetTrainer/ variants/ optimizer/ nnUNetTrainerAdan.py - nnunetv2/
training/ , Python, 1 linennUNetTrainer/ variants/ sampling/ __init__.py - nnunetv2/
training/ , Python, 57 linesnnUNetTrainer/ variants/ sampling/ nnUNetTrainer_probabilis ticOversampling.py - nnunetv2/
training/ , Python, 1 linennUNetTrainer/ variants/ training_length/ __init__.py - nnunetv2/
training/ , Python, 98 linesnnUNetTrainer/ variants/ training_length/ nnUNetTrainer_Xepochs.py - nnunetv2/
training/ , Python, 59 linesnnUNetTrainer/ variants/ training_length/ nnUNetTrainer_Xepochs_No Mirroring.py - nnunetv2/
utilities/ , Python, 1 line__init__.py - nnunetv2/
utilities/ , Python, 24 linescollate_outputs.py - nnunetv2/
utilities/ , Python, 16 linescrossval_split.py - nnunetv2/
utilities/ , Python, 75 linesdataset_name_id_conversi on.py - nnunetv2/
utilities/ , Python, 105 linesddp.py - nnunetv2/
utilities/ , Python, 53 linesddp_allgather.py - nnunetv2/
utilities/ , Python, 50 linesdefault_n_proc_DA.py - nnunetv2/
utilities/ , Python, 132 linesfile_path_utilities.py - nnunetv2/
utilities/ , Python, 174 linesfind_class_by_name.py - nnunetv2/
utilities/ , Python, 56 linesfind_objects.py - nnunetv2/
utilities/ , Python, 91 linesget_network_from_plans.p y - nnunetv2/
utilities/ , Python, 89 linesget_network_via_name.py - nnunetv2/
utilities/ , Python, 27 lineshelpers.py - nnunetv2/
utilities/ , Python, 60 linesjson_export.py - nnunetv2/
utilities/ , Python, 1 linelabel_handling/ __init__.py - nnunetv2/
utilities/ , Python, 351 lineslabel_handling/ label_handling.py - nnunetv2/
utilities/ , Python, 252 linesload_weights_utils.py - nnunetv2/
utilities/ , Python, 12 linesnetwork_initialization.p y - nnunetv2/
utilities/ , Python, 279 linesoverlay_plots.py - nnunetv2/
utilities/ , Python, 1 lineplans_handling/ __init__.py - nnunetv2/
utilities/ , Python, 341 linesplans_handling/ plans_handler.py - nnunetv2/
utilities/ , Python, 51 linespool_utils.py - nnunetv2/
utilities/ , Python, 76 linesutils.py - setup.py, Python, 4 lines
- LICENSE, License, 201 lines
- readme.md, Text, 78 lines
Pubec/nnunetv2-unlearning
25225a808b6a363c48be705df0c6a06996bf035e, 16 May 2024Availability: 1 check, the latest on 27 September 2026: the link answers
- 27 September 2026: the link answers
195 files
- documentation/
__init__.py , Python, 1 line - nnunetv2/
__init__.py , Python, 1 line - nnunetv2/
batch_running/ , Python, 1 line__init__.py - nnunetv2/
batch_running/ , Python, 1 linebenchmarking/ __init__.py - nnunetv2/
batch_running/ , Python, 41 linesbenchmarking/ generate_benchmarking_co mmands.py - nnunetv2/
batch_running/ , Python, 70 linesbenchmarking/ summarize_benchmark_resu lts.py - nnunetv2/
batch_running/ , Python, 114 linescollect_results_custom_D ecathlon.py - nnunetv2/
batch_running/ , Python, 18 linescollect_results_custom_D ecathlon_2d.py - nnunetv2/
batch_running/ , Python, 86 linesgenerate_lsf_runs_custom Decathlon.py - nnunetv2/
batch_running/ , Python, 1 linerelease_trainings/ __init__.py - nnunetv2/
batch_running/ , Python, 1 linerelease_trainings/ nnunetv2_v1/ __init__.py - nnunetv2/
batch_running/ , Python, 113 linesrelease_trainings/ nnunetv2_v1/ collect_results.py - nnunetv2/
batch_running/ , Python, 93 linesrelease_trainings/ nnunetv2_v1/ generate_lsf_commands.py - nnunetv2/
configuration.py , Python, 10 lines - nnunetv2/
dataset_conversion/ , Python, 87 linesDataset027_ACDC.py - nnunetv2/
dataset_conversion/ , Python, 85 linesDataset073_Fluo_C3DH_A54 9_SIM.py - nnunetv2/
dataset_conversion/ , Python, 198 linesDataset114_MNMs.py - nnunetv2/
dataset_conversion/ , Python, 61 linesDataset115_EMIDEC.py - nnunetv2/
dataset_conversion/ , Python, 87 linesDataset120_RoadSegmentat ion.py - nnunetv2/
dataset_conversion/ , Python, 98 linesDataset137_BraTS21.py - nnunetv2/
dataset_conversion/ , Python, 70 linesDataset218_Amos2022_task 1.py - nnunetv2/
dataset_conversion/ , Python, 65 linesDataset219_Amos2022_task 2.py - nnunetv2/
dataset_conversion/ , Python, 50 linesDataset220_KiTS2023.py - nnunetv2/
dataset_conversion/ , Python, 70 linesDataset221_AutoPETII_202 3.py - nnunetv2/
dataset_conversion/ , Python, 59 linesDataset223_AMOS2022postC hallenge.py - nnunetv2/
dataset_conversion/ , Python, 32 linesDataset988_dummyDataset4 .py - nnunetv2/
dataset_conversion/ , Python, 1 line__init__.py - nnunetv2/
dataset_conversion/ , Python, 132 linesconvert_MSD_dataset.py - nnunetv2/
dataset_conversion/ , Python, 53 linesconvert_raw_dataset_from _old_nnunet_format.py - nnunetv2/
dataset_conversion/ , Python, 75 linesdatasets_for_integration _tests/ Dataset996_IntegrationTe st_Hippocampus_regions_i gnore.py - nnunetv2/
dataset_conversion/ , Python, 37 linesdatasets_for_integration _tests/ Dataset997_IntegrationTe st_Hippocampus_regions.p y - nnunetv2/
dataset_conversion/ , Python, 33 linesdatasets_for_integration _tests/ Dataset998_IntegrationTe st_Hippocampus_ignore.py - nnunetv2/
dataset_conversion/ , Python, 27 linesdatasets_for_integration _tests/ Dataset999_IntegrationTe st_Hippocampus.py - nnunetv2/
dataset_conversion/ , Python, 1 linedatasets_for_integration _tests/ __init__.py - nnunetv2/
dataset_conversion/ , Python, 103 linesgenerate_dataset_json.py - nnunetv2/
ensembling/ , Python, 1 line__init__.py - nnunetv2/
ensembling/ , Python, 206 linesensemble.py - nnunetv2/
evaluation/ , Python, 1 line__init__.py - nnunetv2/
evaluation/ , Python, 58 linesaccumulate_cv_results.py - nnunetv2/
evaluation/ , Python, 264 linesevaluate_predictions.py - nnunetv2/
evaluation/ , Python, 357 linesfind_best_configuration. py - nnunetv2/
experiment_planning/ , Python, 1 line__init__.py - nnunetv2/
experiment_planning/ , Python, 1 linedataset_fingerprint/ __init__.py - nnunetv2/
experiment_planning/ , Python, 199 linesdataset_fingerprint/ fingerprint_extractor.py - nnunetv2/
experiment_planning/ , Python, 1 lineexperiment_planners/ __init__.py - nnunetv2/
experiment_planning/ , Python, 556 linesexperiment_planners/ default_experiment_plann er.py - nnunetv2/
experiment_planning/ , Python, 105 linesexperiment_planners/ network_topology.py - nnunetv2/
experiment_planning/ , Python, 54 linesexperiment_planners/ resencUNet_planner.py - nnunetv2/
experiment_planning/ , Python, 137 linesplan_and_preprocess_api. py - nnunetv2/
experiment_planning/ , Python, 201 linesplan_and_preprocess_entr ypoints.py - nnunetv2/
experiment_planning/ , Python, 1 lineplans_for_pretraining/ __init__.py - nnunetv2/
experiment_planning/ , Python, 83 linesplans_for_pretraining/ move_plans_between_datas ets.py - nnunetv2/
experiment_planning/ , Python, 234 linesverify_dataset_integrity .py - nnunetv2/
imageio/ , Python, 1 line__init__.py - nnunetv2/
imageio/ , Python, 107 linesbase_reader_writer.py - nnunetv2/
imageio/ , Python, 73 linesnatural_image_reader_wri ter.py - nnunetv2/
imageio/ , Python, 204 linesnibabel_reader_writer.py - nnunetv2/
imageio/ , Python, 79 linesreader_writer_registry.p y - nnunetv2/
imageio/ , Python, 129 linessimpleitk_reader_writer. py - nnunetv2/
imageio/ , Python, 100 linestif_reader_writer.py - nnunetv2/
inference/ , Python, 1 line__init__.py - nnunetv2/
inference/ , Python, 318 linesdata_iterators.py - nnunetv2/
inference/ , Python, 102 linesexamples.py - nnunetv2/
inference/ , Python, 651 linesexport_bottleneck_featur es.py - nnunetv2/
inference/ , Python, 143 linesexport_prediction.py - nnunetv2/
inference/ , Python, 928 linespredict_from_raw_data.py - nnunetv2/
inference/ , Python, 67 linessliding_window_predictio n.py - nnunetv2/
model_sharing/ , Python, 1 line__init__.py - nnunetv2/
model_sharing/ , Python, 61 linesentry_points.py - nnunetv2/
model_sharing/ , Python, 47 linesmodel_download.py - nnunetv2/
model_sharing/ , Python, 124 linesmodel_export.py - nnunetv2/
model_sharing/ , Python, 8 linesmodel_import.py - nnunetv2/
models/ , Python, 1 line__init__.py - nnunetv2/
models/ , Python, 58 linesmodels.py - nnunetv2/
paths.py , Python, 39 lines - nnunetv2/
postprocessing/ , Python, 1 line__init__.py - nnunetv2/
postprocessing/ , Python, 364 linesremove_connected_compone nts.py - nnunetv2/
preprocessing/ , Python, 1 line__init__.py - nnunetv2/
preprocessing/ , Python, 1 linecropping/ __init__.py - nnunetv2/
preprocessing/ , Python, 51 linescropping/ cropping.py - nnunetv2/
preprocessing/ , Python, 1 linenormalization/ __init__.py - nnunetv2/
preprocessing/ , Python, 98 linesnormalization/ default_normalization_sc hemes.py - nnunetv2/
preprocessing/ , Python, 24 linesnormalization/ map_channel_name_to_norm alization.py - nnunetv2/
preprocessing/ , Python, 1 linepreprocessors/ __init__.py - nnunetv2/
preprocessing/ , Python, 296 linespreprocessors/ default_preprocessor.py - nnunetv2/
preprocessing/ , Python, 1 lineresampling/ __init__.py - nnunetv2/
preprocessing/ , Python, 216 linesresampling/ default_resampling.py - nnunetv2/
preprocessing/ , Python, 15 linesresampling/ utils.py - nnunetv2/
run/ , Python, 1 line__init__.py - nnunetv2/
run/ , Python, 66 linesload_pretrained_weights. py - nnunetv2/
run/ , Python, 282 linesrun_training.py - nnunetv2/
tests/ , Python, 1 line__init__.py - nnunetv2/
tests/ , Python, 1 lineintegration_tests/ __init__.py - nnunetv2/
tests/ , Python, 33 linesintegration_tests/ add_lowres_and_cascade.p y - nnunetv2/
tests/ , Python, 19 linesintegration_tests/ cleanup_integration_test .py - nnunetv2/
tests/ , Shell, 10 linesintegration_tests/ lsf_commands.sh - nnunetv2/
tests/ , Shell, 18 linesintegration_tests/ prepare_integration_test s.sh - nnunetv2/
tests/ , Shell, 27 linesintegration_tests/ run_integration_test.sh - nnunetv2/
tests/ , Python, 75 linesintegration_tests/ run_integration_test_bes tconfig_inference.py - nnunetv2/
tests/ , Shell, 1 lineintegration_tests/ run_integration_test_tra iningOnly_DDP.sh - nnunetv2/
training/ , Python, 1 line__init__.py - nnunetv2/
training/ , Python, 1 linedata_augmentation/ __init__.py - nnunetv2/
training/ , Python, 24 linesdata_augmentation/ compute_initial_patch_si ze.py - nnunetv2/
training/ , Python, 1 linedata_augmentation/ custom_transforms/ __init__.py - nnunetv2/
training/ , Python, 136 linesdata_augmentation/ custom_transforms/ cascade_transforms.py - nnunetv2/
training/ , Python, 55 linesdata_augmentation/ custom_transforms/ deep_supervision_donwsam pling.py - nnunetv2/
training/ , Python, 10 linesdata_augmentation/ custom_transforms/ limited_length_multithre aded_augmenter.py - nnunetv2/
training/ , Python, 10 linesdata_augmentation/ custom_transforms/ manipulating_data_dict.p y - nnunetv2/
training/ , Python, 22 linesdata_augmentation/ custom_transforms/ masking.py - nnunetv2/
training/ , Python, 38 linesdata_augmentation/ custom_transforms/ region_based_training.py - nnunetv2/
training/ , Python, 45 linesdata_augmentation/ custom_transforms/ transforms_for_dummy_2d. py - nnunetv2/
training/ , Python, 1 linedataloading/ __init__.py - nnunetv2/
training/ , Python, 139 linesdataloading/ base_data_loader.py - nnunetv2/
training/ , Python, 94 linesdataloading/ data_loader_2d.py - nnunetv2/
training/ , Python, 66 linesdataloading/ data_loader_3d.py - nnunetv2/
training/ , Python, 105 linesdataloading/ data_loader_unlearn.py - nnunetv2/
training/ , Python, 146 linesdataloading/ nnunet_dataset.py - nnunetv2/
training/ , Python, 128 linesdataloading/ utils.py - nnunetv2/
training/ , Python, 1 linelogging/ __init__.py - nnunetv2/
training/ , Python, 103 lineslogging/ nnunet_logger.py - nnunetv2/
training/ , Python, 132 lineslogging/ unlearn_logger.py - nnunetv2/
training/ , Python, 1 lineloss/ __init__.py - nnunetv2/
training/ , Python, 150 linesloss/ compound_losses.py - nnunetv2/
training/ , Python, 30 linesloss/ deep_supervision.py - nnunetv2/
training/ , Python, 192 linesloss/ dice.py - nnunetv2/
training/ , Python, 32 linesloss/ robust_ce_loss.py - nnunetv2/
training/ , Python, 1 linelr_scheduler/ __init__.py - nnunetv2/
training/ , Python, 20 lineslr_scheduler/ polylr.py - nnunetv2/
training/ , Python, 1 linennUNetTrainer/ __init__.py - nnunetv2/
training/ , Python, 1,306 lines, 1 matchnnUNetTrainer/ nnUNetTrainer.py - nnunetv2/
training/ , Python, 1 linennUNetTrainer/ variants/ __init__.py - nnunetv2/
training/ , Python, 1 linennUNetTrainer/ variants/ benchmarking/ __init__.py - nnunetv2/
training/ , Python, 65 linesnnUNetTrainer/ variants/ benchmarking/ nnUNetTrainerBenchmark_5 epochs.py - nnunetv2/
training/ , Python, 65 linesnnUNetTrainer/ variants/ benchmarking/ nnUNetTrainerBenchmark_5 epochs_noDataLoading.py - nnunetv2/
training/ , Python, 1 linennUNetTrainer/ variants/ data_augmentation/ __init__.py - nnunetv2/
training/ , Python, 419 linesnnUNetTrainer/ variants/ data_augmentation/ nnUNetTrainerDA5.py - nnunetv2/
training/ , Python, 157 linesnnUNetTrainer/ variants/ data_augmentation/ nnUNetTrainerDAOrd0.py - nnunetv2/
training/ , Python, 40 linesnnUNetTrainer/ variants/ data_augmentation/ nnUNetTrainerNoDA.py - nnunetv2/
training/ , Python, 28 linesnnUNetTrainer/ variants/ data_augmentation/ nnUNetTrainerNoMirroring .py - nnunetv2/
training/ , Python, 1 linennUNetTrainer/ variants/ loss/ __init__.py - nnunetv2/
training/ , Python, 41 linesnnUNetTrainer/ variants/ loss/ nnUNetTrainerCELoss.py - nnunetv2/
training/ , Python, 60 linesnnUNetTrainer/ variants/ loss/ nnUNetTrainerDiceLoss.py - nnunetv2/
training/ , Python, 76 linesnnUNetTrainer/ variants/ loss/ nnUNetTrainerTopkLoss.py - nnunetv2/
training/ , Python, 1 linennUNetTrainer/ variants/ lr_schedule/ __init__.py - nnunetv2/
training/ , Python, 13 linesnnUNetTrainer/ variants/ lr_schedule/ nnUNetTrainerCosAnneal.p y - nnunetv2/
training/ , Python, 1 linennUNetTrainer/ variants/ network_architecture/ __init__.py - nnunetv2/
training/ , Python, 74 linesnnUNetTrainer/ variants/ network_architecture/ nnUNetTrainerBN.py - nnunetv2/
training/ , Python, 16 linesnnUNetTrainer/ variants/ network_architecture/ nnUNetTrainerNoDeepSuper vision.py - nnunetv2/
training/ , Python, 1 linennUNetTrainer/ variants/ optimizer/ __init__.py - nnunetv2/
training/ , Python, 58 linesnnUNetTrainer/ variants/ optimizer/ nnUNetTrainerAdam.py - nnunetv2/
training/ , Python, 66 linesnnUNetTrainer/ variants/ optimizer/ nnUNetTrainerAdan.py - nnunetv2/
training/ , Python, 1 linennUNetTrainer/ variants/ sampling/ __init__.py - nnunetv2/
training/ , Python, 84 linesnnUNetTrainer/ variants/ sampling/ nnUNetTrainer_probabilis ticOversampling.py - nnunetv2/
training/ , Python, 1 linennUNetTrainer/ variants/ training_length/ __init__.py - nnunetv2/
training/ , Python, 76 linesnnUNetTrainer/ variants/ training_length/ nnUNetTrainer_Xepochs.py - nnunetv2/
training/ , Python, 60 linesnnUNetTrainer/ variants/ training_length/ nnUNetTrainer_Xepochs_No Mirroring.py - nnunetv2/
training/ , Python, 1 linennUNetTrainer/ variants/ unlearning/ Unlearn/ __init__.py - nnunetv2/
training/ , Python, 29 linesnnUNetTrainer/ variants/ unlearning/ Unlearn/ nnUNetTrainer_Unlearn.py - nnunetv2/
training/ , Python, 89 linesnnUNetTrainer/ variants/ unlearning/ Unlearn/ nnUNetTrainer_UnlearnPer iodicPerBatch.py - nnunetv2/
training/ , Python, 8 linesnnUNetTrainer/ variants/ unlearning/ Unlearn/ nnUNetTrainer_UnlearnPer iodicPerBatchKLDivConfus ionLoss.py - nnunetv2/
training/ , Python, 8 linesnnUNetTrainer/ variants/ unlearning/ Unlearn/ nnUNetTrainer_UnlearnPer iodicPerBatchRefConfusio nLoss.py - nnunetv2/
training/ , Python, 8 linesnnUNetTrainer/ variants/ unlearning/ Unlearn/ nnUNetTrainer_UnlearnPer iodicPerBatchRefConfusio nLossNoNeg.py - nnunetv2/
training/ , Python, 90 linesnnUNetTrainer/ variants/ unlearning/ Unlearn/ nnUNetTrainer_UnlearnPer iodicPerEpoch.py - nnunetv2/
training/ , Python, 8 linesnnUNetTrainer/ variants/ unlearning/ Unlearn/ nnUNetTrainer_UnlearnRef ConfusionLoss.py - nnunetv2/
training/ , Python, 8 linesnnUNetTrainer/ variants/ unlearning/ Unlearn/ nnUNetTrainer_UnlearnRef ConfusionLossEpsilon.py - nnunetv2/
training/ , Python, 8 linesnnUNetTrainer/ variants/ unlearning/ Unlearn/ nnUNetTrainer_UnlearnRef ConfusionLossNoNeg.py - nnunetv2/
training/ , Python, 32 linesnnUNetTrainer/ variants/ unlearning/ Unlearn/ nnUNetTrainer_UnlearnSeq uential.py - nnunetv2/
training/ , Python, 1 linennUNetTrainer/ variants/ unlearning/ UnlearnNoDA/ __init__.py - nnunetv2/
training/ , Python, 39 linesnnUNetTrainer/ variants/ unlearning/ UnlearnNoDA/ nnUNetTrainer_UnlearnNoD A.py - nnunetv2/
training/ , Python, 41 linesnnUNetTrainer/ variants/ unlearning/ UnlearnNoDA/ nnUNetTrainer_UnlearnNoD APeriodicPerBatch.py - nnunetv2/
training/ , Python, 8 linesnnUNetTrainer/ variants/ unlearning/ UnlearnNoDA/ nnUNetTrainer_UnlearnNoD APeriodicPerBatchKLDivCo nfusionLoss.py - nnunetv2/
training/ , Python, 1 linennUNetTrainer/ variants/ unlearning/ __init__.py - nnunetv2/
training/ , Python, 184 lines, 1 matchnnUNetTrainer/ variants/ unlearning/ models.py - nnunetv2/
training/ , Python, 1,451 linesnnUNetTrainer/ variants/ unlearning/ nnUNetTrainer_UnlearnBas e.py - nnunetv2/
utilities/ , Python, 1 line__init__.py - nnunetv2/
utilities/ , Python, 24 linescollate_outputs.py - nnunetv2/
utilities/ , Python, 16 linescrossval_split.py - nnunetv2/
utilities/ , Python, 74 linesdataset_name_id_conversi on.py - nnunetv2/
utilities/ , Python, 49 linesddp_allgather.py - nnunetv2/
utilities/ , Python, 44 linesdefault_n_proc_DA.py - nnunetv2/
utilities/ , Python, 127 linesfile_path_utilities.py - nnunetv2/
utilities/ , Python, 24 linesfind_class_by_name.py - nnunetv2/
utilities/ , Python, 79 linesget_network_from_plans.p y - nnunetv2/
utilities/ , Python, 27 lineshelpers.py - nnunetv2/
utilities/ , Python, 60 linesjson_export.py - nnunetv2/
utilities/ , Python, 1 linelabel_handling/ __init__.py - nnunetv2/
utilities/ , Python, 322 lineslabel_handling/ label_handling.py - nnunetv2/
utilities/ , Python, 12 linesnetwork_initialization.p y - nnunetv2/
utilities/ , Python, 275 linesoverlay_plots.py - nnunetv2/
utilities/ , Python, 1 lineplans_handling/ __init__.py - nnunetv2/
utilities/ , Python, 307 linesplans_handling/ plans_handler.py - nnunetv2/
utilities/ , Python, 69 linesutils.py - setup.py, Python, 4 lines
- LICENSE, License, 201 lines
- readme.md, Text, 16 lines
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/
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://
BibTeX
@article{preloznik2026ad
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/
url = {https://
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/
VL - 13
SP - 1875760
SN - 2296-858X
PB - Frontiers Media SA
DO - 10.3389/
UR - https://
LA - en
ER -
CSL-JSON
{
"id": "10.3389/
"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":
"volume": "13",
"page": "1875760",
"DOI": "10.3389/
"PMID": "42602271",
"PMCID": "PMC13474016",
"ISSN": "2296-858X",
"publisher": "Frontiers Media SA",
"URL": "https://
"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: HippocampusIn 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 imagingIn 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 SclerosisJournal: 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 surgeryIn 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 neuroscienceIn 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: NeuroinformaticsIn 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 oncologyIn 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 healthIn 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 psychiatryIn 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 imagingIn 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.
Claim this paper
Correct its record
Say what each link of this record is, remove the ones that are not the paper's, add the ones that are missing. The correction becomes a new version of the record, in its Versions section.
Validate its tracing map
You validate the map as this page shows it: 2 repositories of the authors' code, each at its verified commit and with its license, 416 scripts, and 3 matches between paragraphs and code (see the Code and Map sections). It then receives a DOI on Zenodo, with you (your ORCID iD) and OSCR as its creators; the code itself is not deposited.
The map's fingerprint: sha256:13b4090c8ae338d6…
Add the badge to its README
The badge links the code to this page. Copy one of these into the README of the paper's code: only you decide where it goes, and nothing is changed for you.
Markdown
[, paste the snippet at the top, then “Commit changes…” and, to review it first, “Create a new branch and start a pull request”. You open the pull request; OSCR asks for no permission.
Request its removal
To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).
Discussion, reproductions, activity
Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.
Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.
Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.
