{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":117682,"databundleVersionId":15062069,"isSourceIdPinned":false},{"sourceType":"datasetVersion","sourceId":14704075,"datasetId":9393785,"databundleVersionId":15549135},{"sourceType":"datasetVersion","sourceId":14577821,"datasetId":9148843,"databundleVersionId":15411317},{"sourceType":"datasetVersion","sourceId":14748104,"datasetId":9303801,"databundleVersionId":15597878},{"sourceType":"datasetVersion","sourceId":14687597,"datasetId":9145348,"databundleVersionId":15531096},{"sourceType":"datasetVersion","sourceId":14066012,"datasetId":8953198,"databundleVersionId":14846463},{"sourceType":"datasetVersion","sourceId":14330461,"datasetId":8944423,"databundleVersionId":15137528},{"sourceType":"modelInstanceVersion","sourceId":760140,"databundleVersionId":15795592,"modelInstanceId":569748,"modelId":582040,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":289294474,"isSourceIdPinned":false}],"dockerImageVersionId":31193,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!mkdir -p /kaggle/temp\n!mkdir predictions_tiff\n!pip install nnunetv2 nibabel tifffile tqdm -q --no-index -f \"/kaggle/input/surface-packages-offline\"\n!pip install --no-index --find-links=\"/kaggle/input/surface-package-scraper\" -q monai albumentations imagecodecs --no-deps # \"numpy==1.26.4\" \"scipy==1.15.3\"\n!pip uninstall -q -y tensorflow\n","metadata":{"_uuid":"82a5abad-98d7-4e54-91bf-78afe62ddbb7","_cell_guid":"69b7c4ca-d393-41d0-ae77-e96a9e952a13","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-01-23T02:42:25.043591Z","iopub.execute_input":"2026-01-23T02:42:25.043786Z","iopub.status.idle":"2026-01-23T02:43:57.753066Z","shell.execute_reply.started":"2026-01-23T02:42:25.043768Z","shell.execute_reply":"2026-01-23T02:43:57.752294Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile /usr/local/lib/python3.11/dist-packages/nnunetv2/training/nnUNetTrainer/nnUNetTrainer_RotFlip_ClDice.py\nfrom nnunetv2.training.nnUNetTrainer.nnUNetTrainer import nnUNetTrainer\nfrom nnunetv2.utilities.helpers import softmax_helper_dim1\n\nimport inspect\nimport multiprocessing\nimport os\nimport shutil\nimport sys\nimport warnings\nfrom copy import deepcopy\nfrom datetime import datetime\nfrom time import time, sleep\nfrom typing import Tuple, Union, List\n\nimport numpy as np\nimport torch\nfrom batchgenerators.dataloading.multi_threaded_augmenter import MultiThreadedAugmenter\nfrom batchgenerators.dataloading.nondet_multi_threaded_augmenter import NonDetMultiThreadedAugmenter\nfrom batchgenerators.dataloading.single_threaded_augmenter import SingleThreadedAugmenter\nfrom batchgenerators.utilities.file_and_folder_operations import join, load_json, isfile, save_json, maybe_mkdir_p\nfrom batchgeneratorsv2.helpers.scalar_type import RandomScalar\nfrom batchgeneratorsv2.transforms.base.basic_transform import BasicTransform\nfrom batchgeneratorsv2.transforms.intensity.brightness import MultiplicativeBrightnessTransform\nfrom batchgeneratorsv2.transforms.intensity.contrast import ContrastTransform, BGContrast\nfrom batchgeneratorsv2.transforms.intensity.gamma import GammaTransform\nfrom batchgeneratorsv2.transforms.intensity.gaussian_noise import GaussianNoiseTransform\nfrom batchgeneratorsv2.transforms.nnunet.random_binary_operator import ApplyRandomBinaryOperatorTransform\nfrom batchgeneratorsv2.transforms.nnunet.remove_connected_components import \\\n    RemoveRandomConnectedComponentFromOneHotEncodingTransform\nfrom batchgeneratorsv2.transforms.nnunet.seg_to_onehot import MoveSegAsOneHotToDataTransform\nfrom batchgeneratorsv2.transforms.noise.gaussian_blur import GaussianBlurTransform\nfrom batchgeneratorsv2.transforms.spatial.low_resolution import SimulateLowResolutionTransform\nfrom batchgeneratorsv2.transforms.spatial.mirroring import MirrorTransform\nfrom batchgeneratorsv2.transforms.spatial.spatial import SpatialTransform\nfrom batchgeneratorsv2.transforms.utils.compose import ComposeTransforms\nfrom batchgeneratorsv2.transforms.utils.deep_supervision_downsampling import DownsampleSegForDSTransform\nfrom batchgeneratorsv2.transforms.utils.nnunet_masking import MaskImageTransform\nfrom batchgeneratorsv2.transforms.utils.pseudo2d import Convert3DTo2DTransform, Convert2DTo3DTransform\nfrom batchgeneratorsv2.transforms.utils.random import RandomTransform\nfrom batchgeneratorsv2.transforms.utils.remove_label import RemoveLabelTansform\nfrom batchgeneratorsv2.transforms.utils.seg_to_regions import ConvertSegmentationToRegionsTransform\nfrom torch import autocast, nn\nfrom torch import distributed as dist\nfrom torch._dynamo import OptimizedModule\nfrom torch.cuda import device_count\nfrom torch import GradScaler\nfrom torch.nn.parallel import DistributedDataParallel as DDP\n\nfrom nnunetv2.configuration import ANISO_THRESHOLD, default_num_processes\nfrom nnunetv2.evaluation.evaluate_predictions import compute_metrics_on_folder\nfrom nnunetv2.inference.export_prediction import export_prediction_from_logits, resample_and_save\nfrom nnunetv2.inference.predict_from_raw_data import nnUNetPredictor\nfrom nnunetv2.inference.sliding_window_prediction import compute_gaussian\nfrom nnunetv2.paths import nnUNet_preprocessed, nnUNet_results\nfrom nnunetv2.training.data_augmentation.compute_initial_patch_size import get_patch_size\nfrom nnunetv2.training.dataloading.nnunet_dataset import infer_dataset_class\nfrom nnunetv2.training.dataloading.data_loader import nnUNetDataLoader\nfrom nnunetv2.training.logging.nnunet_logger import nnUNetLogger\nfrom nnunetv2.training.loss.compound_losses import DC_and_CE_loss, DC_and_BCE_loss\nfrom nnunetv2.training.loss.deep_supervision import DeepSupervisionWrapper\nfrom nnunetv2.training.loss.dice import get_tp_fp_fn_tn, MemoryEfficientSoftDiceLoss\nfrom nnunetv2.training.lr_scheduler.polylr import PolyLRScheduler\nfrom nnunetv2.utilities.collate_outputs import collate_outputs\nfrom nnunetv2.utilities.crossval_split import generate_crossval_split\nfrom nnunetv2.utilities.default_n_proc_DA import get_allowed_n_proc_DA\nfrom nnunetv2.utilities.file_path_utilities import check_workers_alive_and_busy\nfrom nnunetv2.utilities.get_network_from_plans import get_network_from_plans\nfrom nnunetv2.utilities.helpers import empty_cache, dummy_context\nfrom nnunetv2.utilities.label_handling.label_handling import convert_labelmap_to_one_hot, determine_num_input_channels\nfrom nnunetv2.utilities.plans_handling.plans_handler import PlansManager\n\n\n# ============================================================================\nROT_PROB = 0.6 \nFORCE_MIRROR_AXES = (0, 1, 2)\n\nfrom monai.losses import SoftclDiceLoss\n\n\nclass AddMonaiClDice(nn.Module):\n    \"\"\"\n    total_loss = base_loss(net_output, target) + cldice_weight * SoftclDiceLoss(y_true, y_pred)\n    \"\"\"\n    def __init__(\n        self,\n        base_loss: nn.Module,\n        ignore_label: int = None,\n        cldice_weight: float = 0.2,\n        cldice_iters: int = 3,\n        smooth: float = 1.0,\n    ):\n        super().__init__()\n        self.base_loss = base_loss\n        self.ignore_label = ignore_label\n        self.w = float(cldice_weight)\n        self.cl = SoftclDiceLoss(iter_=cldice_iters, smooth=smooth)\n\n    def forward(self, net_output: torch.Tensor, target: torch.Tensor) -> torch.Tensor:\n        base = self.base_loss(net_output, target)\n\n        if self.w <= 0:\n            return base\n\n        if target.ndim == net_output.ndim:\n            if target.shape[1] == 1:\n                tgt = target[:, 0].long()\n            else:\n                tgt = torch.argmax(target, dim=1).long()\n        else:\n            tgt = target.long()\n\n        valid = None\n        if self.ignore_label is not None:\n            valid = (tgt != self.ignore_label).float()  # (B,...)\n\n        # ---- probs: (B,C,...)\n        probs = softmax_helper_dim1(net_output)\n        C = probs.shape[1]\n        if C < 2:\n            return base\n\n        cl_terms = []\n        for c in range(1, C):\n            y_true = (tgt == c).float()          # (B,...)\n            y_pred = probs[:, c]                # (B,...)\n\n            if valid is not None:\n                y_true = y_true * valid\n                y_pred = y_pred * valid\n\n            # MONAI SoftclDiceLoss: forward(y_true, y_pred)\n            cl_terms.append(self.cl(y_true.unsqueeze(1), y_pred.unsqueeze(1)))\n\n        if len(cl_terms) == 0:\n            return base\n\n        cldice = torch.stack(cl_terms).mean()\n        return base + self.w * cldice\n\n\n\nclass nnUNetTrainer_RotFlip_ClDice(nnUNetTrainer):\n    \"\"\"\n    - loss: DeepSupervisionWrapper( DC+CE + λ*clDice )\n    - aug: ROT_PROB MIRROR_AXES\n    \"\"\"\n\n    @staticmethod\n    def get_training_transforms(\n        patch_size,\n        rotation_for_DA,\n        deep_supervision_scales,\n        mirror_axes,\n        do_dummy_2d_data_aug,\n        use_mask_for_norm=None,\n        is_cascaded=False,\n        foreground_labels=None,\n        regions=None,\n        ignore_label=None,\n    ):\n        # --- xyz flip ---\n        mirror_axes = FORCE_MIRROR_AXES\n\n        transforms = []\n        if do_dummy_2d_data_aug:\n            ignore_axes = (0,)\n            transforms.append(Convert3DTo2DTransform())\n            patch_size_spatial = patch_size[1:]\n        else:\n            patch_size_spatial = patch_size\n            ignore_axes = None\n\n        transforms.append(\n            SpatialTransform(\n                patch_size_spatial, patch_center_dist_from_border=0, random_crop=False, p_elastic_deform=0,\n                p_rotation=ROT_PROB,         \n                rotation=rotation_for_DA,\n                p_scaling=0.2, scaling=(0.7, 1.4), p_synchronize_scaling_across_axes=1,\n                bg_style_seg_sampling=False\n            )\n        )\n\n        if do_dummy_2d_data_aug:\n            transforms.append(Convert2DTo3DTransform())\n\n        transforms.append(RandomTransform(\n            GaussianNoiseTransform(\n                noise_variance=(0, 0.1),\n                p_per_channel=1,\n                synchronize_channels=True\n            ), apply_probability=0.1\n        ))\n        transforms.append(RandomTransform(\n            GaussianBlurTransform(\n                blur_sigma=(0.5, 1.),\n                synchronize_channels=False,\n                synchronize_axes=False,\n                p_per_channel=0.5, benchmark=True\n            ), apply_probability=0.2\n        ))\n        transforms.append(RandomTransform(\n            MultiplicativeBrightnessTransform(\n                multiplier_range=BGContrast((0.75, 1.25)),\n                synchronize_channels=False,\n                p_per_channel=1\n            ), apply_probability=0.15\n        ))\n        transforms.append(RandomTransform(\n            ContrastTransform(\n                contrast_range=BGContrast((0.75, 1.25)),\n                preserve_range=True,\n                synchronize_channels=False,\n                p_per_channel=1\n            ), apply_probability=0.15\n        ))\n        transforms.append(RandomTransform(\n            SimulateLowResolutionTransform(\n                scale=(0.5, 1),\n                synchronize_channels=False,\n                synchronize_axes=True,\n                ignore_axes=ignore_axes,\n                allowed_channels=None,\n                p_per_channel=0.5\n            ), apply_probability=0.25\n        ))\n        transforms.append(RandomTransform(\n            GammaTransform(\n                gamma=BGContrast((0.7, 1.5)),\n                p_invert_image=1,\n                synchronize_channels=False,\n                p_per_channel=1,\n                p_retain_stats=1\n            ), apply_probability=0.1\n        ))\n        transforms.append(RandomTransform(\n            GammaTransform(\n                gamma=BGContrast((0.7, 1.5)),\n                p_invert_image=0,\n                synchronize_channels=False,\n                p_per_channel=1,\n                p_retain_stats=1\n            ), apply_probability=0.3\n        ))\n\n        if mirror_axes is not None and len(mirror_axes) > 0:\n            transforms.append(MirrorTransform(allowed_axes=mirror_axes))\n\n        if use_mask_for_norm is not None and any(use_mask_for_norm):\n            transforms.append(MaskImageTransform(\n                apply_to_channels=[i for i in range(len(use_mask_for_norm)) if use_mask_for_norm[i]],\n                channel_idx_in_seg=0,\n                set_outside_to=0,\n            ))\n\n        transforms.append(RemoveLabelTansform(-1, 0))\n\n        if is_cascaded:\n            assert foreground_labels is not None\n            transforms.append(\n                MoveSegAsOneHotToDataTransform(\n                    source_channel_idx=1,\n                    all_labels=foreground_labels,\n                    remove_channel_from_source=True\n                )\n            )\n            transforms.append(RandomTransform(\n                ApplyRandomBinaryOperatorTransform(\n                    channel_idx=list(range(-len(foreground_labels), 0)),\n                    strel_size=(1, 8),\n                    p_per_label=1\n                ), apply_probability=0.4\n            ))\n            transforms.append(RandomTransform(\n                RemoveRandomConnectedComponentFromOneHotEncodingTransform(\n                    channel_idx=list(range(-len(foreground_labels), 0)),\n                    fill_with_other_class_p=0,\n                    dont_do_if_covers_more_than_x_percent=0.15,\n                    p_per_label=1\n                ), apply_probability=0.2\n            ))\n\n        if regions is not None:\n            transforms.append(\n                ConvertSegmentationToRegionsTransform(\n                    regions=list(regions) + [ignore_label] if ignore_label is not None else regions,\n                    channel_in_seg=0\n                )\n            )\n\n        if deep_supervision_scales is not None:\n            transforms.append(DownsampleSegForDSTransform(ds_scales=deep_supervision_scales))\n\n        return ComposeTransforms(transforms)\n\n\n    def _build_loss(self):\n        if self.label_manager.has_regions:\n            loss = DC_and_BCE_loss({},\n                                {'batch_dice': self.configuration_manager.batch_dice,\n                                    'do_bg': True, 'smooth': 1e-5, 'ddp': self.is_ddp},\n                                use_ignore_label=self.label_manager.ignore_label is not None,\n                                dice_class=MemoryEfficientSoftDiceLoss)\n        else:\n            loss = DC_and_CE_loss({'batch_dice': self.configuration_manager.batch_dice,\n                                'smooth': 1e-5, 'do_bg': False, 'ddp': self.is_ddp},\n                                {},\n                                weight_ce=1, weight_dice=1,\n                                ignore_label=self.label_manager.ignore_label,\n                                dice_class=MemoryEfficientSoftDiceLoss)\n\n            # === clDice ==\n            loss = AddMonaiClDice(\n                base_loss=loss,\n                ignore_label=self.label_manager.ignore_label,  \n                cldice_weight=0.2,   \n                cldice_iters=3       \n            )\n\n        if self._do_i_compile():\n            base = getattr(loss, \"base_loss\", loss)\n            if hasattr(base, \"dc\"):\n                base.dc = torch.compile(base.dc)\n\n        if self.enable_deep_supervision:\n            deep_supervision_scales = self._get_deep_supervision_scales()\n            weights = np.array([1 / (2 ** i) for i in range(len(deep_supervision_scales))])\n\n            if self.is_ddp and not self._do_i_compile():\n                weights[-1] = 1e-6\n            else:\n                weights[-1] = 0\n\n            weights = weights / weights.sum()\n            loss = DeepSupervisionWrapper(loss, weights)\n\n        return loss\n\n\n    def configure_optimizers(self):\n        from schedulefree import RAdamScheduleFree\n        optimizer = RAdamScheduleFree(self.network.parameters(), 1e-3, weight_decay=1e-4)\n        return optimizer, None\n    \n    def on_train_epoch_start(self):\n        self.network.train()\n        # self.lr_scheduler.step(self.current_epoch)\n        self.print_to_log_file('')\n        self.print_to_log_file(f'Epoch {self.current_epoch}')\n        self.print_to_log_file(\n            f\"Current learning rate: {np.round(self.optimizer.param_groups[0]['lr'], decimals=5)}\")\n        # lrs are the same for all workers so we don't need to gather them in case of DDP training\n        self.logger.log('lrs', self.optimizer.param_groups[0]['lr'], self.current_epoch)\n        self.optimizer.train()\n    \n    def on_train_epoch_end(self, train_outputs: List[dict]):\n        outputs = collate_outputs(train_outputs)\n\n        if self.is_ddp:\n            losses_tr = [None for _ in range(dist.get_world_size())]\n            dist.all_gather_object(losses_tr, outputs['loss'])\n            loss_here = np.vstack(losses_tr).mean()\n        else:\n            loss_here = np.mean(outputs['loss'])\n\n        self.logger.log('train_losses', loss_here, self.current_epoch)\n        self.optimizer.eval()","metadata":{"trusted":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2026-01-23T02:43:57.754824Z","iopub.execute_input":"2026-01-23T02:43:57.755138Z","iopub.status.idle":"2026-01-23T02:43:57.767929Z","shell.execute_reply.started":"2026-01-23T02:43:57.755114Z","shell.execute_reply":"2026-01-23T02:43:57.767193Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile /usr/local/lib/python3.11/dist-packages/nnunetv2/inference/predict_from_raw_data.py\nimport inspect\nimport itertools\nimport multiprocessing\nimport os\nfrom copy import deepcopy\nfrom queue import Queue\nfrom threading import Thread\nfrom time import sleep\nfrom typing import Tuple, Union, List, Optional\n\nimport numpy as np\nimport torch\nfrom acvl_utils.cropping_and_padding.padding import pad_nd_image\nfrom batchgenerators.dataloading.multi_threaded_augmenter import MultiThreadedAugmenter\nfrom batchgenerators.utilities.file_and_folder_operations import load_json, join, isfile, maybe_mkdir_p, isdir, subdirs, \\\n    save_json\nfrom torch import nn\nfrom torch._dynamo import OptimizedModule\nfrom torch.nn.parallel import DistributedDataParallel\nfrom tqdm import tqdm\n\nimport nnunetv2\nfrom nnunetv2.configuration import default_num_processes\nfrom nnunetv2.inference.data_iterators import PreprocessAdapterFromNpy, preprocessing_iterator_fromfiles, \\\n    preprocessing_iterator_fromnpy\nfrom nnunetv2.inference.export_prediction import export_prediction_from_logits, \\\n    convert_predicted_logits_to_segmentation_with_correct_shape\nfrom nnunetv2.inference.sliding_window_prediction import compute_gaussian, \\\n    compute_steps_for_sliding_window\nfrom nnunetv2.utilities.file_path_utilities import get_output_folder, check_workers_alive_and_busy\nfrom nnunetv2.utilities.find_class_by_name import recursive_find_python_class\nfrom nnunetv2.utilities.helpers import empty_cache, dummy_context\nfrom nnunetv2.utilities.json_export import recursive_fix_for_json_export\nfrom nnunetv2.utilities.label_handling.label_handling import determine_num_input_channels\nfrom nnunetv2.utilities.plans_handling.plans_handler import PlansManager, ConfigurationManager\nfrom nnunetv2.utilities.utils import create_lists_from_splitted_dataset_folder\n\n\nclass nnUNetPredictor(object):\n    def __init__(self,\n                 tile_step_size: float = 0.5,\n                 use_gaussian: bool = True,\n                 use_mirroring: bool = True,\n                 perform_everything_on_device: bool = True,\n                 device: torch.device = torch.device('cuda'),\n                 verbose: bool = False,\n                 verbose_preprocessing: bool = False,\n                 allow_tqdm: bool = True):\n        self.verbose = verbose\n        self.verbose_preprocessing = verbose_preprocessing\n        self.allow_tqdm = allow_tqdm\n\n        self.plans_manager, self.configuration_manager, self.list_of_parameters, self.network, self.dataset_json, \\\n        self.trainer_name, self.allowed_mirroring_axes, self.label_manager = None, None, None, None, None, None, None, None\n\n        self.tile_step_size = tile_step_size\n        self.use_gaussian = use_gaussian\n        self.use_mirroring = use_mirroring\n        if device.type == 'cuda':\n            torch.backends.cudnn.benchmark = True\n        else:\n            print(f'perform_everything_on_device=True is only supported for cuda devices! Setting this to False')\n            perform_everything_on_device = False\n        self.device = device\n        self.perform_everything_on_device = perform_everything_on_device\n\n    def initialize_from_trained_model_folder(self, model_training_output_dir: str,\n                                             use_folds: Union[Tuple[Union[int, str]], None],\n                                             checkpoint_name: str = 'checkpoint_final.pth'):\n        \"\"\"\n        This is used when making predictions with a trained model\n        \"\"\"\n        if use_folds is None:\n            use_folds = nnUNetPredictor.auto_detect_available_folds(model_training_output_dir, checkpoint_name)\n\n        dataset_json = load_json(join(model_training_output_dir, 'dataset.json'))\n        plans = load_json(join(model_training_output_dir, 'plans.json'))\n        plans_manager = PlansManager(plans)\n\n        if isinstance(use_folds, str):\n            use_folds = [use_folds]\n\n        parameters = []\n        for i, f in enumerate(use_folds):\n            f = int(f) if f != 'all' else f\n            checkpoint = torch.load(join(model_training_output_dir, f'fold_{f}', checkpoint_name),\n                                    map_location=torch.device('cpu'), weights_only=False)\n            if i == 0:\n                trainer_name = checkpoint['trainer_name']\n                configuration_name = checkpoint['init_args']['configuration']\n                inference_allowed_mirroring_axes = checkpoint['inference_allowed_mirroring_axes'] if \\\n                    'inference_allowed_mirroring_axes' in checkpoint.keys() else None\n\n            parameters.append(checkpoint['network_weights'])\n\n        configuration_manager = plans_manager.get_configuration(configuration_name)\n        # restore network\n        num_input_channels = determine_num_input_channels(plans_manager, configuration_manager, dataset_json)\n        trainer_class = recursive_find_python_class(join(nnunetv2.__path__[0], \"training\", \"nnUNetTrainer\"),\n                                                    trainer_name, 'nnunetv2.training.nnUNetTrainer')\n        if trainer_class is None:\n            raise RuntimeError(f'Unable to locate trainer class {trainer_name} in nnunetv2.training.nnUNetTrainer. '\n                               f'Please place it there (in any .py file)!')\n        network = trainer_class.build_network_architecture(\n            configuration_manager.network_arch_class_name,\n            configuration_manager.network_arch_init_kwargs,\n            configuration_manager.network_arch_init_kwargs_req_import,\n            num_input_channels,\n            plans_manager.get_label_manager(dataset_json).num_segmentation_heads,\n            enable_deep_supervision=False\n        )\n\n        self.plans_manager = plans_manager\n        self.configuration_manager = configuration_manager\n        self.list_of_parameters = parameters\n\n        # initialize network with first set of parameters, also see https://github.com/MIC-DKFZ/nnUNet/issues/2520\n        network.load_state_dict(parameters[0])\n\n        self.network = network\n\n        self.dataset_json = dataset_json\n        self.trainer_name = trainer_name\n        self.allowed_mirroring_axes = inference_allowed_mirroring_axes\n        self.label_manager = plans_manager.get_label_manager(dataset_json)\n        if ('nnUNet_compile' in os.environ.keys()) and (os.environ['nnUNet_compile'].lower() in ('true', '1', 't')) \\\n                and not isinstance(self.network, OptimizedModule):\n            print('Using torch.compile')\n            self.network = torch.compile(self.network)\n\n    def manual_initialization(self, network: nn.Module, plans_manager: PlansManager,\n                              configuration_manager: ConfigurationManager, parameters: Optional[List[dict]],\n                              dataset_json: dict, trainer_name: str,\n                              inference_allowed_mirroring_axes: Optional[Tuple[int, ...]]):\n        \"\"\"\n        This is used by the nnUNetTrainer to initialize nnUNetPredictor for the final validation\n        \"\"\"\n        self.plans_manager = plans_manager\n        self.configuration_manager = configuration_manager\n        self.list_of_parameters = parameters\n        self.network = network\n        self.dataset_json = dataset_json\n        self.trainer_name = trainer_name\n        self.allowed_mirroring_axes = inference_allowed_mirroring_axes\n        self.label_manager = plans_manager.get_label_manager(dataset_json)\n        allow_compile = True\n        allow_compile = allow_compile and ('nnUNet_compile' in os.environ.keys()) and (\n                    os.environ['nnUNet_compile'].lower() in ('true', '1', 't'))\n        allow_compile = allow_compile and not isinstance(self.network, OptimizedModule)\n        if isinstance(self.network, DistributedDataParallel):\n            allow_compile = allow_compile and isinstance(self.network.module, OptimizedModule)\n        if allow_compile:\n            print('Using torch.compile')\n            self.network = torch.compile(self.network)\n\n    @staticmethod\n    def auto_detect_available_folds(model_training_output_dir, checkpoint_name):\n        print('use_folds is None, attempting to auto detect available folds')\n        fold_folders = subdirs(model_training_output_dir, prefix='fold_', join=False)\n        fold_folders = [i for i in fold_folders if i != 'fold_all']\n        fold_folders = [i for i in fold_folders if isfile(join(model_training_output_dir, i, checkpoint_name))]\n        use_folds = [int(i.split('_')[-1]) for i in fold_folders]\n        print(f'found the following folds: {use_folds}')\n        return use_folds\n\n    def _manage_input_and_output_lists(self, list_of_lists_or_source_folder: Union[str, List[List[str]]],\n                                       output_folder_or_list_of_truncated_output_files: Union[None, str, List[str]],\n                                       folder_with_segs_from_prev_stage: str = None,\n                                       overwrite: bool = True,\n                                       part_id: int = 0,\n                                       num_parts: int = 1,\n                                       save_probabilities: bool = False):\n        if isinstance(list_of_lists_or_source_folder, str):\n            list_of_lists_or_source_folder = create_lists_from_splitted_dataset_folder(list_of_lists_or_source_folder,\n                                                                                       self.dataset_json['file_ending'])\n        print(f'There are {len(list_of_lists_or_source_folder)} cases in the source folder')\n        list_of_lists_or_source_folder = list_of_lists_or_source_folder[part_id::num_parts]\n        caseids = [os.path.basename(i[0])[:-(len(self.dataset_json['file_ending']) + 5)] for i in\n                   list_of_lists_or_source_folder]\n        print(\n            f'I am processing {part_id} out of {num_parts} (max process ID is {num_parts - 1}, we start counting with 0!)')\n        print(f'There are {len(caseids)} cases that I would like to predict')\n\n        if isinstance(output_folder_or_list_of_truncated_output_files, str):\n            output_filename_truncated = [join(output_folder_or_list_of_truncated_output_files, i) for i in caseids]\n        elif isinstance(output_folder_or_list_of_truncated_output_files, list):\n            output_filename_truncated = output_folder_or_list_of_truncated_output_files[part_id::num_parts]\n        else:\n            output_filename_truncated = None\n        seg_from_prev_stage_files = [join(folder_with_segs_from_prev_stage, i + self.dataset_json['file_ending']) if\n                                     folder_with_segs_from_prev_stage is not None else None for i in caseids]\n        # remove already predicted files from the lists\n        if not overwrite and output_filename_truncated is not None:\n            tmp = [isfile(i + self.dataset_json['file_ending']) for i in output_filename_truncated]\n            if save_probabilities:\n                tmp2 = [isfile(i + '.npz') for i in output_filename_truncated]\n                tmp = [i and j for i, j in zip(tmp, tmp2)]\n            not_existing_indices = [i for i, j in enumerate(tmp) if not j]\n\n            output_filename_truncated = [output_filename_truncated[i] for i in not_existing_indices]\n            list_of_lists_or_source_folder = [list_of_lists_or_source_folder[i] for i in not_existing_indices]\n            seg_from_prev_stage_files = [seg_from_prev_stage_files[i] for i in not_existing_indices]\n            print(f'overwrite was set to {overwrite}, so I am only working on cases that haven\\'t been predicted yet. '\n                  f'That\\'s {len(not_existing_indices)} cases.')\n        return list_of_lists_or_source_folder, output_filename_truncated, seg_from_prev_stage_files\n\n    def predict_from_files(self,\n                           list_of_lists_or_source_folder: Union[str, List[List[str]]],\n                           output_folder_or_list_of_truncated_output_files: Union[str, None, List[str]],\n                           save_probabilities: bool = False,\n                           overwrite: bool = True,\n                           num_processes_preprocessing: int = default_num_processes,\n                           num_processes_segmentation_export: int = default_num_processes,\n                           folder_with_segs_from_prev_stage: str = None,\n                           num_parts: int = 1,\n                           part_id: int = 0):\n        \"\"\"\n        This is nnU-Net's default function for making predictions. It works best for batch predictions\n        (predicting many images at once).\n        \"\"\"\n        assert part_id <= num_parts, (\"Part ID must be smaller than num_parts. Remember that we start counting with 0. \"\n                                      \"So if there are 3 parts then valid part IDs are 0, 1, 2\")\n        if isinstance(output_folder_or_list_of_truncated_output_files, str):\n            output_folder = output_folder_or_list_of_truncated_output_files\n        elif isinstance(output_folder_or_list_of_truncated_output_files, list):\n            output_folder = os.path.dirname(output_folder_or_list_of_truncated_output_files[0])\n        else:\n            output_folder = None\n\n        ########################\n        # let's store the input arguments so that its clear what was used to generate the prediction\n        if output_folder is not None:\n            my_init_kwargs = {}\n            for k in inspect.signature(self.predict_from_files).parameters.keys():\n                my_init_kwargs[k] = locals()[k]\n            my_init_kwargs = deepcopy(\n                my_init_kwargs)  # let's not unintentionally change anything in-place. Take this as a\n            recursive_fix_for_json_export(my_init_kwargs)\n            maybe_mkdir_p(output_folder)\n            save_json(my_init_kwargs, join(output_folder, 'predict_from_raw_data_args.json'))\n\n            # we need these two if we want to do things with the predictions like for example apply postprocessing\n            save_json(self.dataset_json, join(output_folder, 'dataset.json'), sort_keys=False)\n            save_json(self.plans_manager.plans, join(output_folder, 'plans.json'), sort_keys=False)\n        #######################\n\n        # check if we need a prediction from the previous stage\n        if self.configuration_manager.previous_stage_name is not None:\n            assert folder_with_segs_from_prev_stage is not None, \\\n                f'The requested configuration is a cascaded network. It requires the segmentations of the previous ' \\\n                f'stage ({self.configuration_manager.previous_stage_name}) as input. Please provide the folder where' \\\n                f' they are located via folder_with_segs_from_prev_stage'\n\n        # sort out input and output filenames\n        list_of_lists_or_source_folder, output_filename_truncated, seg_from_prev_stage_files = \\\n            self._manage_input_and_output_lists(list_of_lists_or_source_folder,\n                                                output_folder_or_list_of_truncated_output_files,\n                                                folder_with_segs_from_prev_stage, overwrite, part_id, num_parts,\n                                                save_probabilities)\n        if len(list_of_lists_or_source_folder) == 0:\n            return\n\n        data_iterator = self._internal_get_data_iterator_from_lists_of_filenames(list_of_lists_or_source_folder,\n                                                                                 seg_from_prev_stage_files,\n                                                                                 output_filename_truncated,\n                                                                                 num_processes_preprocessing)\n\n        return self.predict_from_data_iterator(data_iterator, save_probabilities, num_processes_segmentation_export)\n\n    def _internal_get_data_iterator_from_lists_of_filenames(self,\n                                                            input_list_of_lists: List[List[str]],\n                                                            seg_from_prev_stage_files: Union[List[str], None],\n                                                            output_filenames_truncated: Union[List[str], None],\n                                                            num_processes: int):\n        return preprocessing_iterator_fromfiles(input_list_of_lists, seg_from_prev_stage_files,\n                                                output_filenames_truncated, self.plans_manager, self.dataset_json,\n                                                self.configuration_manager, num_processes, self.device.type == 'cuda',\n                                                self.verbose_preprocessing)\n        # preprocessor = self.configuration_manager.preprocessor_class(verbose=self.verbose_preprocessing)\n        # # hijack batchgenerators, yo\n        # # we use the multiprocessing of the batchgenerators dataloader to handle all the background worker stuff. This\n        # # way we don't have to reinvent the wheel here.\n        # num_processes = max(1, min(num_processes, len(input_list_of_lists)))\n        # ppa = PreprocessAdapter(input_list_of_lists, seg_from_prev_stage_files, preprocessor,\n        #                         output_filenames_truncated, self.plans_manager, self.dataset_json,\n        #                         self.configuration_manager, num_processes)\n        # if num_processes == 0:\n        #     mta = SingleThreadedAugmenter(ppa, None)\n        # else:\n        #     mta = MultiThreadedAugmenter(ppa, None, num_processes, 1, None, pin_memory=pin_memory)\n        # return mta\n\n    def get_data_iterator_from_raw_npy_data(self,\n                                            image_or_list_of_images: Union[np.ndarray, List[np.ndarray]],\n                                            segs_from_prev_stage_or_list_of_segs_from_prev_stage: Union[None,\n                                                                                                        np.ndarray,\n                                                                                                        List[\n                                                                                                            np.ndarray]],\n                                            properties_or_list_of_properties: Union[dict, List[dict]],\n                                            truncated_ofname: Union[str, List[str], None],\n                                            num_processes: int = 3):\n\n        list_of_images = [image_or_list_of_images] if not isinstance(image_or_list_of_images, list) else \\\n            image_or_list_of_images\n\n        if isinstance(segs_from_prev_stage_or_list_of_segs_from_prev_stage, np.ndarray):\n            segs_from_prev_stage_or_list_of_segs_from_prev_stage = [\n                segs_from_prev_stage_or_list_of_segs_from_prev_stage]\n\n        if isinstance(truncated_ofname, str):\n            truncated_ofname = [truncated_ofname]\n\n        if isinstance(properties_or_list_of_properties, dict):\n            properties_or_list_of_properties = [properties_or_list_of_properties]\n\n        num_processes = min(num_processes, len(list_of_images))\n        pp = preprocessing_iterator_fromnpy(\n            list_of_images,\n            segs_from_prev_stage_or_list_of_segs_from_prev_stage,\n            properties_or_list_of_properties,\n            truncated_ofname,\n            self.plans_manager,\n            self.dataset_json,\n            self.configuration_manager,\n            num_processes,\n            self.device.type == 'cuda',\n            self.verbose_preprocessing\n        )\n\n        return pp\n\n    def predict_from_list_of_npy_arrays(self,\n                                        image_or_list_of_images: Union[np.ndarray, List[np.ndarray]],\n                                        segs_from_prev_stage_or_list_of_segs_from_prev_stage: Union[None,\n                                                                                                    np.ndarray,\n                                                                                                    List[\n                                                                                                        np.ndarray]],\n                                        properties_or_list_of_properties: Union[dict, List[dict]],\n                                        truncated_ofname: Union[str, List[str], None],\n                                        num_processes: int = 3,\n                                        save_probabilities: bool = False,\n                                        num_processes_segmentation_export: int = default_num_processes):\n        iterator = self.get_data_iterator_from_raw_npy_data(image_or_list_of_images,\n                                                            segs_from_prev_stage_or_list_of_segs_from_prev_stage,\n                                                            properties_or_list_of_properties,\n                                                            truncated_ofname,\n                                                            num_processes)\n        return self.predict_from_data_iterator(iterator, save_probabilities, num_processes_segmentation_export)\n\n    def predict_from_data_iterator(self,\n                                   data_iterator,\n                                   save_probabilities: bool = False,\n                                   num_processes_segmentation_export: int = default_num_processes):\n        \"\"\"\n        each element returned by data_iterator must be a dict with 'data', 'ofile' and 'data_properties' keys!\n        If 'ofile' is None, the result will be returned instead of written to a file\n        \"\"\"\n        with multiprocessing.get_context(\"spawn\").Pool(num_processes_segmentation_export) as export_pool:\n            worker_list = [i for i in export_pool._pool]\n            r = []\n            for preprocessed in data_iterator:\n                data = preprocessed['data']\n                if isinstance(data, str):\n                    delfile = data\n                    data = torch.from_numpy(np.load(data))\n                    os.remove(delfile)\n\n                ofile = preprocessed['ofile']\n                if ofile is not None:\n                    print(f'\\nPredicting {os.path.basename(ofile)}:')\n                else:\n                    print(f'\\nPredicting image of shape {data.shape}:')\n\n                print(f'perform_everything_on_device: {self.perform_everything_on_device}')\n\n                properties = preprocessed['data_properties']\n\n                # let's not get into a runaway situation where the GPU predicts so fast that the disk has to be swamped with\n                # npy files\n                proceed = not check_workers_alive_and_busy(export_pool, worker_list, r, allowed_num_queued=2)\n                while not proceed:\n                    sleep(0.1)\n                    proceed = not check_workers_alive_and_busy(export_pool, worker_list, r, allowed_num_queued=2)\n\n                # convert to numpy to prevent uncatchable memory alignment errors from multiprocessing serialization of torch tensors\n                prediction = self.predict_logits_from_preprocessed_data(data).cpu().detach().numpy()\n\n                if ofile is not None:\n                    print('sending off prediction to background worker for resampling and export')\n                    r.append(\n                        export_pool.starmap_async(\n                            export_prediction_from_logits,\n                            ((prediction, properties, self.configuration_manager, self.plans_manager,\n                              self.dataset_json, ofile, save_probabilities),)\n                        )\n                    )\n                else:\n                    print('sending off prediction to background worker for resampling')\n                    r.append(\n                        export_pool.starmap_async(\n                            convert_predicted_logits_to_segmentation_with_correct_shape, (\n                                (prediction, self.plans_manager,\n                                 self.configuration_manager, self.label_manager,\n                                 properties,\n                                 save_probabilities),)\n                        )\n                    )\n                if ofile is not None:\n                    print(f'done with {os.path.basename(ofile)}')\n                else:\n                    print(f'\\nDone with image of shape {data.shape}:')\n            ret = [i.get()[0] for i in r]\n\n        if isinstance(data_iterator, MultiThreadedAugmenter):\n            data_iterator._finish()\n\n        # clear lru cache\n        compute_gaussian.cache_clear()\n        # clear device cache\n        empty_cache(self.device)\n        return ret\n\n    def predict_single_npy_array(self, input_image: np.ndarray, image_properties: dict,\n                                 segmentation_previous_stage: np.ndarray = None,\n                                 output_file_truncated: str = None,\n                                 save_or_return_probabilities: bool = False):\n        \"\"\"\n        WARNING: SLOW. ONLY USE THIS IF YOU CANNOT GIVE NNUNET MULTIPLE IMAGES AT ONCE FOR SOME REASON.\n\n\n        input_image: Make sure to load the image in the way nnU-Net expects! nnU-Net is trained on a certain axis\n                     ordering which cannot be disturbed in inference,\n                     otherwise you will get bad results. The easiest way to achieve that is to use the same I/O class\n                     for loading images as was used during nnU-Net preprocessing! You can find that class in your\n                     plans.json file under the key \"image_reader_writer\". If you decide to freestyle, know that the\n                     default axis ordering for medical images is the one from SimpleITK. If you load with nibabel,\n                     you need to transpose your axes AND your spacing from [x,y,z] to [z,y,x]!\n        image_properties must only have a 'spacing' key!\n        \"\"\"\n        ppa = PreprocessAdapterFromNpy([input_image], [segmentation_previous_stage], [image_properties],\n                                       [output_file_truncated],\n                                       self.plans_manager, self.dataset_json, self.configuration_manager,\n                                       num_threads_in_multithreaded=1, verbose=self.verbose)\n        if self.verbose:\n            print('preprocessing')\n        dct = next(ppa)\n\n        if self.verbose:\n            print('predicting')\n        predicted_logits = self.predict_logits_from_preprocessed_data(dct['data']).cpu()\n\n        if self.verbose:\n            print('resampling to original shape')\n        if output_file_truncated is not None:\n            export_prediction_from_logits(predicted_logits, dct['data_properties'], self.configuration_manager,\n                                          self.plans_manager, self.dataset_json, output_file_truncated,\n                                          save_or_return_probabilities)\n        else:\n            ret = convert_predicted_logits_to_segmentation_with_correct_shape(predicted_logits, self.plans_manager,\n                                                                              self.configuration_manager,\n                                                                              self.label_manager,\n                                                                              dct['data_properties'],\n                                                                              return_probabilities=\n                                                                              save_or_return_probabilities)\n            if save_or_return_probabilities:\n                return ret[0], ret[1]\n            else:\n                return ret\n\n    @torch.inference_mode()\n    def predict_logits_from_preprocessed_data(self, data: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        IMPORTANT! IF YOU ARE RUNNING THE CASCADE, THE SEGMENTATION FROM THE PREVIOUS STAGE MUST ALREADY BE STACKED ON\n        TOP OF THE IMAGE AS ONE-HOT REPRESENTATION! SEE PreprocessAdapter ON HOW THIS SHOULD BE DONE!\n\n        RETURNED LOGITS HAVE THE SHAPE OF THE INPUT. THEY MUST BE CONVERTED BACK TO THE ORIGINAL IMAGE SIZE.\n        SEE convert_predicted_logits_to_segmentation_with_correct_shape\n        \"\"\"\n        n_threads = torch.get_num_threads()\n        torch.set_num_threads(default_num_processes if default_num_processes < n_threads else n_threads)\n        prediction = None\n\n        for params in self.list_of_parameters:\n\n            # messing with state dict names...\n            if not isinstance(self.network, OptimizedModule):\n                self.network.load_state_dict(params)\n            else:\n                self.network._orig_mod.load_state_dict(params)\n\n            # why not leave prediction on device if perform_everything_on_device? Because this may cause the\n            # second iteration to crash due to OOM. Grabbing that with try except cause way more bloated code than\n            # this actually saves computation time\n            if prediction is None:\n                prediction = self.predict_sliding_window_return_logits(data).to('cpu')\n            else:\n                prediction += self.predict_sliding_window_return_logits(data).to('cpu')\n\n        if len(self.list_of_parameters) > 1:\n            prediction /= len(self.list_of_parameters)\n\n        if self.verbose: print('Prediction done')\n        torch.set_num_threads(n_threads)\n        return prediction\n\n    def _internal_get_sliding_window_slicers(self, image_size: Tuple[int, ...]):\n        slicers = []\n        if len(self.configuration_manager.patch_size) < len(image_size):\n            assert len(self.configuration_manager.patch_size) == len(\n                image_size) - 1, 'if tile_size has less entries than image_size, ' \\\n                                 'len(tile_size) ' \\\n                                 'must be one shorter than len(image_size) ' \\\n                                 '(only dimension ' \\\n                                 'discrepancy of 1 allowed).'\n            steps = compute_steps_for_sliding_window(image_size[1:], self.configuration_manager.patch_size,\n                                                     self.tile_step_size)\n            if self.verbose: print(f'n_steps {image_size[0] * len(steps[0]) * len(steps[1])}, image size is'\n                                   f' {image_size}, tile_size {self.configuration_manager.patch_size}, '\n                                   f'tile_step_size {self.tile_step_size}\\nsteps:\\n{steps}')\n            for d in range(image_size[0]):\n                for sx in steps[0]:\n                    for sy in steps[1]:\n                        slicers.append(\n                            tuple([slice(None), d, *[slice(si, si + ti) for si, ti in\n                                                     zip((sx, sy), self.configuration_manager.patch_size)]]))\n        else:\n            steps = compute_steps_for_sliding_window(image_size, self.configuration_manager.patch_size,\n                                                     self.tile_step_size)\n            if self.verbose: print(\n                f'n_steps {np.prod([len(i) for i in steps])}, image size is {image_size}, tile_size {self.configuration_manager.patch_size}, '\n                f'tile_step_size {self.tile_step_size}\\nsteps:\\n{steps}')\n            for sx in steps[0]:\n                for sy in steps[1]:\n                    for sz in steps[2]:\n                        slicers.append(\n                            tuple([slice(None), *[slice(si, si + ti) for si, ti in\n                                                  zip((sx, sy, sz), self.configuration_manager.patch_size)]]))\n        return slicers\n\n    @torch.inference_mode()\n    def _internal_maybe_mirror_and_predict(self, x: torch.Tensor) -> torch.Tensor:\n        mirror_axes = self.allowed_mirroring_axes if self.use_mirroring else None\n        prediction = self.network(x)\n\n        if mirror_axes is not None:\n            # check for invalid numbers in mirror_axes\n            # x should be 5d for 3d images and 4d for 2d. so the max value of mirror_axes cannot exceed len(x.shape) - 3\n            assert max(mirror_axes) <= x.ndim - 3, 'mirror_axes does not match the dimension of the input!'\n\n            mirror_axes = [m + 2 for m in mirror_axes]\n            axes_combinations = [\n                c for i in range(len(mirror_axes)) for c in itertools.combinations(mirror_axes, i + 1)\n            ]\n            for axes in axes_combinations:\n                prediction += torch.flip(self.network(torch.flip(x, axes)), axes)\n            dxy = (x.ndim - 2, x.ndim - 1)\n\n            def _rotate_inplane(t: torch.Tensor, angle_deg: float) -> torch.Tensor:\n                n = t.shape[0]\n                angle = torch.tensor(angle_deg, device=t.device, dtype=t.dtype) * (torch.pi / 180.0)\n                cos_a = torch.cos(angle)\n                sin_a = torch.sin(angle)\n\n                if t.ndim == 4:\n                    theta = torch.zeros((n, 2, 3), device=t.device, dtype=t.dtype)\n                    theta[:, 0, 0] = cos_a\n                    theta[:, 0, 1] = -sin_a\n                    theta[:, 1, 0] = sin_a\n                    theta[:, 1, 1] = cos_a\n                    grid = torch.nn.functional.affine_grid(theta, t.size(), align_corners=False)\n                    return torch.nn.functional.grid_sample(\n                        t, grid, mode='bilinear', padding_mode='zeros', align_corners=False\n                    )\n\n                if t.ndim == 5:\n                    theta = torch.zeros((n, 3, 4), device=t.device, dtype=t.dtype)\n                    theta[:, 0, 0] = cos_a\n                    theta[:, 0, 1] = -sin_a\n                    theta[:, 1, 0] = sin_a\n                    theta[:, 1, 1] = cos_a\n                    theta[:, 2, 2] = 1.0\n                    grid = torch.nn.functional.affine_grid(theta, t.size(), align_corners=False)\n                    return torch.nn.functional.grid_sample(\n                        t, grid, mode='bilinear', padding_mode='zeros', align_corners=False\n                    )\n\n                raise RuntimeError('Unsupported input dimensionality for rotation!')\n\n            prediction += _rotate_inplane(self.network(_rotate_inplane(x, 15.0)), -15.0)\n            prediction += _rotate_inplane(self.network(_rotate_inplane(x, -15.0)), 15.0)\n            prediction /= (len(axes_combinations) + 1 + 2)\n\n        return prediction\n\n    @torch.inference_mode()\n    def _internal_predict_sliding_window_return_logits(self,\n                                                       data: torch.Tensor,\n                                                       slicers,\n                                                       do_on_device: bool = True,\n                                                       ):\n        predicted_logits = n_predictions = prediction = gaussian = workon = None\n        results_device = self.device if do_on_device else torch.device('cpu')\n\n        def producer(d, slh, q):\n            for s in slh:\n                q.put((torch.clone(d[s][None], memory_format=torch.contiguous_format).to(self.device), s))\n            q.put('end')\n\n        try:\n            empty_cache(self.device)\n\n            # move data to device\n            if self.verbose:\n                print(f'move image to device {results_device}')\n            data = data.to(results_device)\n            queue = Queue(maxsize=2)\n            t = Thread(target=producer, args=(data, slicers, queue))\n            t.start()\n\n            # preallocate arrays\n            if self.verbose:\n                print(f'preallocating results arrays on device {results_device}')\n            predicted_logits = torch.zeros((self.label_manager.num_segmentation_heads, *data.shape[1:]),\n                                           dtype=torch.half,\n                                           device=results_device)\n            n_predictions = torch.zeros(data.shape[1:], dtype=torch.half, device=results_device)\n\n            if self.use_gaussian:\n                gaussian = compute_gaussian(tuple(self.configuration_manager.patch_size), sigma_scale=1. / 8,\n                                            value_scaling_factor=10,\n                                            device=results_device)\n            else:\n                gaussian = 1\n\n            if not self.allow_tqdm and self.verbose:\n                print(f'running prediction: {len(slicers)} steps')\n\n            with tqdm(desc=None, total=len(slicers), disable=not self.allow_tqdm) as pbar:\n                while True:\n                    item = queue.get()\n                    if item == 'end':\n                        queue.task_done()\n                        break\n                    workon, sl = item\n                    prediction = self._internal_maybe_mirror_and_predict(workon)[0].to(results_device)\n\n                    if self.use_gaussian:\n                        prediction *= gaussian\n                    predicted_logits[sl] += prediction\n                    n_predictions[sl[1:]] += gaussian\n                    queue.task_done()\n                    pbar.update()\n            queue.join()\n\n            # predicted_logits /= n_predictions\n            torch.div(predicted_logits, n_predictions, out=predicted_logits)\n            # check for infs\n            if torch.any(torch.isinf(predicted_logits)):\n                raise RuntimeError('Encountered inf in predicted array. Aborting... If this problem persists, '\n                                   'reduce value_scaling_factor in compute_gaussian or increase the dtype of '\n                                   'predicted_logits to fp32')\n        except Exception as e:\n            del predicted_logits, n_predictions, prediction, gaussian, workon\n            empty_cache(self.device)\n            empty_cache(results_device)\n            raise e\n        return predicted_logits\n\n    @torch.inference_mode()\n    def predict_sliding_window_return_logits(self, input_image: torch.Tensor) \\\n            -> Union[np.ndarray, torch.Tensor]:\n        assert isinstance(input_image, torch.Tensor)\n        self.network = self.network.to(self.device)\n        self.network.eval()\n\n        empty_cache(self.device)\n\n        # Autocast can be annoying\n        # If the device_type is 'cpu' then it's slow as heck on some CPUs (no auto bfloat16 support detection)\n        # and needs to be disabled.\n        # If the device_type is 'mps' then it will complain that mps is not implemented, even if enabled=False\n        # is set. Whyyyyyyy. (this is why we don't make use of enabled=False)\n        # So autocast will only be active if we have a cuda device.\n        with torch.autocast(self.device.type, enabled=True) if self.device.type == 'cuda' else dummy_context():\n            assert input_image.ndim == 4, 'input_image must be a 4D np.ndarray or torch.Tensor (c, x, y, z)'\n\n            if self.verbose:\n                print(f'Input shape: {input_image.shape}')\n                print(\"step_size:\", self.tile_step_size)\n                print(\"mirror_axes:\", self.allowed_mirroring_axes if self.use_mirroring else None)\n\n            # if input_image is smaller than tile_size we need to pad it to tile_size.\n            data, slicer_revert_padding = pad_nd_image(input_image, self.configuration_manager.patch_size,\n                                                       'constant', {'value': 0}, True,\n                                                       None)\n\n            slicers = self._internal_get_sliding_window_slicers(data.shape[1:])\n\n            if self.perform_everything_on_device and self.device != 'cpu':\n                # we need to try except here because we can run OOM in which case we need to fall back to CPU as a results device\n                try:\n                    predicted_logits = self._internal_predict_sliding_window_return_logits(data, slicers,\n                                                                                           self.perform_everything_on_device)\n                except RuntimeError:\n                    print(\n                        'Prediction on device was unsuccessful, probably due to a lack of memory. Moving results arrays to CPU')\n                    empty_cache(self.device)\n                    predicted_logits = self._internal_predict_sliding_window_return_logits(data, slicers, False)\n            else:\n                predicted_logits = self._internal_predict_sliding_window_return_logits(data, slicers,\n                                                                                       self.perform_everything_on_device)\n\n            empty_cache(self.device)\n            # revert padding\n            predicted_logits = predicted_logits[(slice(None), *slicer_revert_padding[1:])]\n        return predicted_logits\n\n    def predict_from_files_sequential(self,\n                           list_of_lists_or_source_folder: Union[str, List[List[str]]],\n                           output_folder_or_list_of_truncated_output_files: Union[str, None, List[str]],\n                           save_probabilities: bool = False,\n                           overwrite: bool = True,\n                           folder_with_segs_from_prev_stage: str = None):\n        \"\"\"\n        Just like predict_from_files but doesn't use any multiprocessing. Slow, but sometimes necessary\n        \"\"\"\n        if isinstance(output_folder_or_list_of_truncated_output_files, str):\n            output_folder = output_folder_or_list_of_truncated_output_files\n        elif isinstance(output_folder_or_list_of_truncated_output_files, list):\n            output_folder = os.path.dirname(output_folder_or_list_of_truncated_output_files[0])\n            if len(output_folder) == 0:  # just a file was given without a folder\n                output_folder = os.path.curdir\n        else:\n            output_folder = None\n\n        ########################\n        # let's store the input arguments so that its clear what was used to generate the prediction\n        if output_folder is not None:\n            my_init_kwargs = {}\n            for k in inspect.signature(self.predict_from_files_sequential).parameters.keys():\n                my_init_kwargs[k] = locals()[k]\n            my_init_kwargs = deepcopy(\n                my_init_kwargs)  # let's not unintentionally change anything in-place. Take this as a\n            recursive_fix_for_json_export(my_init_kwargs)\n            save_json(my_init_kwargs, join(output_folder, 'predict_from_raw_data_args.json'))\n\n            # we need these two if we want to do things with the predictions like for example apply postprocessing\n            save_json(self.dataset_json, join(output_folder, 'dataset.json'), sort_keys=False)\n            save_json(self.plans_manager.plans, join(output_folder, 'plans.json'), sort_keys=False)\n        #######################\n\n        # check if we need a prediction from the previous stage\n        if self.configuration_manager.previous_stage_name is not None:\n            assert folder_with_segs_from_prev_stage is not None, \\\n                f'The requested configuration is a cascaded network. It requires the segmentations of the previous ' \\\n                f'stage ({self.configuration_manager.previous_stage_name}) as input. Please provide the folder where' \\\n                f' they are located via folder_with_segs_from_prev_stage'\n\n        # sort out input and output filenames\n        list_of_lists_or_source_folder, output_filename_truncated, seg_from_prev_stage_files = \\\n            self._manage_input_and_output_lists(list_of_lists_or_source_folder,\n                                                output_folder_or_list_of_truncated_output_files,\n                                                folder_with_segs_from_prev_stage, overwrite, 0, 1,\n                                                save_probabilities)\n        if len(list_of_lists_or_source_folder) == 0:\n            return\n\n        label_manager = self.plans_manager.get_label_manager(self.dataset_json)\n        preprocessor = self.configuration_manager.preprocessor_class(verbose=self.verbose)\n\n        if output_filename_truncated is None:\n            output_filename_truncated = [None] * len(list_of_lists_or_source_folder)\n        if seg_from_prev_stage_files is None:\n            seg_from_prev_stage_files = [None] * len(seg_from_prev_stage_files)\n\n        ret = []\n        for li, of, sps in zip(list_of_lists_or_source_folder, output_filename_truncated, seg_from_prev_stage_files):\n            data, seg, data_properties = preprocessor.run_case(\n                li,\n                sps,\n                self.plans_manager,\n                self.configuration_manager,\n                self.dataset_json\n            )\n\n            print(f'perform_everything_on_device: {self.perform_everything_on_device}')\n\n            prediction = self.predict_logits_from_preprocessed_data(torch.from_numpy(data)).cpu()\n\n            if of is not None:\n                export_prediction_from_logits(prediction, data_properties, self.configuration_manager, self.plans_manager,\n                  self.dataset_json, of, save_probabilities)\n            else:\n                ret.append(convert_predicted_logits_to_segmentation_with_correct_shape(prediction, self.plans_manager,\n                     self.configuration_manager, self.label_manager,\n                     data_properties,\n                     save_probabilities))\n\n        # clear lru cache\n        compute_gaussian.cache_clear()\n        # clear device cache\n        empty_cache(self.device)\n        return ret\n\ndef _getDefaultValue(env: str, dtype: type, default: any,) -> any:\n    try:\n        val = dtype(os.environ.get(env) or default)\n    except:\n        val = default\n    return val\n\ndef predict_entry_point_modelfolder():\n    import argparse\n    parser = argparse.ArgumentParser(description='Use this to run inference with nnU-Net. This function is used when '\n                                                 'you want to manually specify a folder containing a trained nnU-Net '\n                                                 'model. This is useful when the nnunet environment variables '\n                                                 '(nnUNet_results) are not set.')\n    parser.add_argument('-i', type=str, required=True,\n                        help='input folder. Remember to use the correct channel numberings for your files (_0000 etc). '\n                             'File endings must be the same as the training dataset!')\n    parser.add_argument('-o', type=str, required=True,\n                        help='Output folder. If it does not exist it will be created. Predicted segmentations will '\n                             'have the same name as their source images.')\n    parser.add_argument('-m', type=str, required=True,\n                        help='Folder in which the trained model is. Must have subfolders fold_X for the different '\n                             'folds you trained')\n    parser.add_argument('-f', nargs='+', type=str, required=False, default=(0, 1, 2, 3, 4),\n                        help='Specify the folds of the trained model that should be used for prediction. '\n                             'Default: (0, 1, 2, 3, 4)')\n    parser.add_argument('-step_size', type=float, required=False, default=0.5,\n                        help='Step size for sliding window prediction. The larger it is the faster but less accurate '\n                             'the prediction. Default: 0.5. Cannot be larger than 1. We recommend the default.')\n    parser.add_argument('--disable_tta', action='store_true', required=False, default=False,\n                        help='Set this flag to disable test time data augmentation in the form of mirroring. Faster, '\n                             'but less accurate inference. Not recommended.')\n    parser.add_argument('--verbose', action='store_true', help=\"Set this if you like being talked to. You will have \"\n                                                               \"to be a good listener/reader.\")\n    parser.add_argument('--save_probabilities', action='store_true',\n                        help='Set this to export predicted class \"probabilities\". Required if you want to ensemble '\n                             'multiple configurations.')\n    parser.add_argument('--continue_prediction', '--c', action='store_true',\n                        help='Continue an aborted previous prediction (will not overwrite existing files)')\n    parser.add_argument('-chk', type=str, required=False, default='checkpoint_final.pth',\n                        help='Name of the checkpoint you want to use. Default: checkpoint_final.pth')\n    parser.add_argument('-npp', type=int, required=False, default=3,\n                        help='Number of processes used for preprocessing. More is not always better. Beware of '\n                             'out-of-RAM issues. Default: 3')\n    parser.add_argument('-nps', type=int, required=False, default=3,\n                        help='Number of processes used for segmentation export. More is not always better. Beware of '\n                             'out-of-RAM issues. Default: 3')\n    parser.add_argument('-prev_stage_predictions', type=str, required=False, default=None,\n                        help='Folder containing the predictions of the previous stage. Required for cascaded models.')\n    parser.add_argument('-device', type=str, default='cuda', required=False,\n                        help=\"Use this to set the device the inference should run with. Available options are 'cuda' \"\n                             \"(GPU), 'cpu' (CPU) and 'mps' (Apple M1/M2). Do NOT use this to set which GPU ID! \"\n                             \"Use CUDA_VISIBLE_DEVICES=X nnUNetv2_predict [...] instead!\")\n    parser.add_argument('--disable_progress_bar', action='store_true', required=False, default=False,\n                        help='Set this flag to disable progress bar. Recommended for HPC environments (non interactive '\n                             'jobs)')\n    parser.add_argument(\n        '-tile_size', type=int, required=False, default=None,\n        help='Override sliding window tile size at inference time. Example: -tile_size 128 160 160'\n    )\n\n\n    print(\n        \"\\n#######################################################################\\nPlease cite the following paper \"\n        \"when using nnU-Net:\\n\"\n        \"Isensee, F., Jaeger, P. F., Kohl, S. A., Petersen, J., & Maier-Hein, K. H. (2021). \"\n        \"nnU-Net: a self-configuring method for deep learning-based biomedical image segmentation. \"\n        \"Nature methods, 18(2), 203-211.\\n#######################################################################\\n\")\n\n    args = parser.parse_args()\n    args.f = [i if i == 'all' else int(i) for i in args.f]\n\n    if not isdir(args.o):\n        maybe_mkdir_p(args.o)\n\n    assert args.device in ['cpu', 'cuda',\n                           'mps'], f'-device must be either cpu, mps or cuda. Other devices are not tested/supported. Got: {args.device}.'\n    if args.device == 'cpu':\n        # let's allow torch to use hella threads\n        import multiprocessing\n        torch.set_num_threads(multiprocessing.cpu_count())\n        device = torch.device('cpu')\n    elif args.device == 'cuda':\n        # multithreading in torch doesn't help nnU-Net if run on GPU\n        torch.set_num_threads(1)\n        torch.set_num_interop_threads(1)\n        device = torch.device('cuda')\n    else:\n        device = torch.device('mps')\n\n    predictor = nnUNetPredictor(tile_step_size=args.step_size,\n                                use_gaussian=True,\n                                use_mirroring=not args.disable_tta,\n                                perform_everything_on_device=True,\n                                device=device,\n                                verbose=args.verbose,\n                                allow_tqdm=not args.disable_progress_bar,\n                                verbose_preprocessing=args.verbose)\n    predictor.initialize_from_trained_model_folder(args.m, args.f, args.chk)\n    if args.tile_size is not None:\n        print(\"old plan :\",predictor.configuration_manager.configuration['patch_size'])\n        predictor.configuration_manager.configuration['patch_size'] = (args.tile_size,args.tile_size,args.tile_size)\n        print(\"new plan :\",predictor.configuration_manager.configuration['patch_size'])\n\n    predictor.predict_from_files(args.i, args.o, save_probabilities=args.save_probabilities,\n                                 overwrite=not args.continue_prediction,\n                                 num_processes_preprocessing=args.npp,\n                                 num_processes_segmentation_export=args.nps,\n                                 folder_with_segs_from_prev_stage=args.prev_stage_predictions,\n                                 num_parts=1, part_id=0)\n\n\ndef predict_entry_point():\n    import argparse\n    parser = argparse.ArgumentParser(description='Use this to run inference with nnU-Net. This function is used when '\n                                                 'you want to manually specify a folder containing a trained nnU-Net '\n                                                 'model. This is useful when the nnunet environment variables '\n                                                 '(nnUNet_results) are not set.')\n    parser.add_argument('-i', type=str, required=True,\n                        help='input folder. Remember to use the correct channel numberings for your files (_0000 etc). '\n                             'File endings must be the same as the training dataset!')\n    parser.add_argument('-o', type=str, required=True,\n                        help='Output folder. If it does not exist it will be created. Predicted segmentations will '\n                             'have the same name as their source images.')\n    parser.add_argument('-d', type=str, required=True,\n                        help='Dataset with which you would like to predict. You can specify either dataset name or id')\n    parser.add_argument('-p', type=str, required=False, default='nnUNetPlans',\n                        help='Plans identifier. Specify the plans in which the desired configuration is located. '\n                             'Default: nnUNetPlans')\n    parser.add_argument('-tr', type=str, required=False, default='nnUNetTrainer',\n                        help='What nnU-Net trainer class was used for training? Default: nnUNetTrainer')\n    parser.add_argument('-c', type=str, required=True,\n                        help='nnU-Net configuration that should be used for prediction. Config must be located '\n                             'in the plans specified with -p')\n    parser.add_argument('-f', nargs='+', type=str, required=False, default=(0, 1, 2, 3, 4),\n                        help='Specify the folds of the trained model that should be used for prediction. '\n                             'Default: (0, 1, 2, 3, 4)')\n    parser.add_argument('-step_size', type=float, required=False, default=0.5,\n                        help='Step size for sliding window prediction. The larger it is the faster but less accurate '\n                             'the prediction. Default: 0.5. Cannot be larger than 1. We recommend the default.')\n    parser.add_argument('--disable_tta', action='store_true', required=False, default=False,\n                        help='Set this flag to disable test time data augmentation in the form of mirroring. Faster, '\n                             'but less accurate inference. Not recommended.')\n    parser.add_argument('--verbose', action='store_true', help=\"Set this if you like being talked to. You will have \"\n                                                               \"to be a good listener/reader.\")\n    parser.add_argument('--save_probabilities', action='store_true',\n                        help='Set this to export predicted class \"probabilities\". Required if you want to ensemble '\n                             'multiple configurations.')\n    parser.add_argument('--continue_prediction', action='store_true',\n                        help='Continue an aborted previous prediction (will not overwrite existing files)')\n    parser.add_argument('-chk', type=str, required=False, default='checkpoint_final.pth',\n                        help='Name of the checkpoint you want to use. Default: checkpoint_final.pth')\n    parser.add_argument('-npp', type=int, required=False, default=_getDefaultValue('nnUNet_npp', int, 3),\n                        help='Number of processes used for preprocessing. More is not always better. Beware of '\n                             'out-of-RAM issues. Default: 3')\n    parser.add_argument('-nps', type=int, required=False, default=_getDefaultValue('nnUNet_nps', int, 3),\n                        help='Number of processes used for segmentation export. More is not always better. Beware of '\n                             'out-of-RAM issues. Default: 3')\n    parser.add_argument('-prev_stage_predictions', type=str, required=False, default=None,\n                        help='Folder containing the predictions of the previous stage. Required for cascaded models.')\n    parser.add_argument('-num_parts', type=int, required=False, default=1,\n                        help='Number of separate nnUNetv2_predict call that you will be making. Default: 1 (= this one '\n                             'call predicts everything)')\n    parser.add_argument('-part_id', type=int, required=False, default=0,\n                        help='If multiple nnUNetv2_predict exist, which one is this? IDs start with 0 can end with '\n                             'num_parts - 1. So when you submit 5 nnUNetv2_predict calls you need to set -num_parts '\n                             '5 and use -part_id 0, 1, 2, 3 and 4. Simple, right? Note: You are yourself responsible '\n                             'to make these run on separate GPUs! Use CUDA_VISIBLE_DEVICES (google, yo!)')\n    parser.add_argument('-device', type=str, default='cuda', required=False,\n                        help=\"Use this to set the device the inference should run with. Available options are 'cuda' \"\n                             \"(GPU), 'cpu' (CPU) and 'mps' (Apple M1/M2). Do NOT use this to set which GPU ID! \"\n                             \"Use CUDA_VISIBLE_DEVICES=X nnUNetv2_predict [...] instead!\")\n    parser.add_argument('--disable_progress_bar', action='store_true', required=False, default=False,\n                        help='Set this flag to disable progress bar. Recommended for HPC environments (non interactive '\n                             'jobs)')\n    parser.add_argument(\n        '-tile_size', type=int, required=False, default=None,\n        help='Override sliding window tile size at inference time. Example: -tile_size 128 160 160'\n    )\n    print(\n        \"\\n#######################################################################\\nPlease cite the following paper \"\n        \"when using nnU-Net:\\n\"\n        \"Isensee, F., Jaeger, P. F., Kohl, S. A., Petersen, J., & Maier-Hein, K. H. (2021). \"\n        \"nnU-Net: a self-configuring method for deep learning-based biomedical image segmentation. \"\n        \"Nature methods, 18(2), 203-211.\\n#######################################################################\\n\")\n\n    args = parser.parse_args()\n    args.f = [i if i == 'all' else int(i) for i in args.f]\n\n    model_folder = get_output_folder(args.d, args.tr, args.p, args.c)\n\n    if not isdir(args.o):\n        maybe_mkdir_p(args.o)\n\n    # slightly passive aggressive haha\n    assert args.part_id < args.num_parts, 'Do you even read the documentation? See nnUNetv2_predict -h.'\n\n    assert args.device in ['cpu', 'cuda',\n                           'mps'], f'-device must be either cpu, mps or cuda. Other devices are not tested/supported. Got: {args.device}.'\n    if args.device == 'cpu':\n        # let's allow torch to use hella threads\n        import multiprocessing\n        torch.set_num_threads(multiprocessing.cpu_count())\n        device = torch.device('cpu')\n    elif args.device == 'cuda':\n        # multithreading in torch doesn't help nnU-Net if run on GPU\n        torch.set_num_threads(1)\n        torch.set_num_interop_threads(1)\n        device = torch.device('cuda')\n    else:\n        device = torch.device('mps')\n\n    predictor = nnUNetPredictor(tile_step_size=args.step_size,\n                                use_gaussian=True,\n                                use_mirroring=not args.disable_tta,\n                                perform_everything_on_device=True,\n                                device=device,\n                                verbose=args.verbose,\n                                verbose_preprocessing=args.verbose,\n                                allow_tqdm=not args.disable_progress_bar)\n    predictor.initialize_from_trained_model_folder(\n        model_folder,\n        args.f,\n        checkpoint_name=args.chk\n    )\n    if args.tile_size is not None:\n        print(\"old plan :\",predictor.configuration_manager.configuration['patch_size'])\n        predictor.configuration_manager.configuration['patch_size'] = (args.tile_size,args.tile_size,args.tile_size)\n        print(\"new plan :\",predictor.configuration_manager.configuration['patch_size'])\n\n    run_sequential = args.nps == 0 and args.npp == 0\n    \n    if run_sequential:\n        \n        print(\"Running in non-multiprocessing mode\")\n        predictor.predict_from_files_sequential(args.i, args.o, save_probabilities=args.save_probabilities,\n                                                overwrite=not args.continue_prediction,\n                                                folder_with_segs_from_prev_stage=args.prev_stage_predictions)\n    \n    else:\n        \n        predictor.predict_from_files(args.i, args.o, save_probabilities=args.save_probabilities,\n                                    overwrite=not args.continue_prediction,\n                                    num_processes_preprocessing=args.npp,\n                                    num_processes_segmentation_export=args.nps,\n                                    folder_with_segs_from_prev_stage=args.prev_stage_predictions,\n                                    num_parts=args.num_parts,\n                                    part_id=args.part_id)\n    \n    # r = predict_from_raw_data(args.i,\n    #                           args.o,\n    #                           model_folder,\n    #                           args.f,\n    #                           args.step_size,\n    #                           use_gaussian=True,\n    #                           use_mirroring=not args.disable_tta,\n    #                           perform_everything_on_device=True,\n    #                           verbose=args.verbose,\n    #                           save_probabilities=args.save_probabilities,\n    #                           overwrite=not args.continue_prediction,\n    #                           checkpoint_name=args.chk,\n    #                           num_processes_preprocessing=args.npp,\n    #                           num_processes_segmentation_export=args.nps,\n    #                           folder_with_segs_from_prev_stage=args.prev_stage_predictions,\n    #                           num_parts=args.num_parts,\n    #                           part_id=args.part_id,\n    #                           device=device)\n\n\nif __name__ == '__main__':\n    ########################## predict a bunch of files\n    from nnunetv2.paths import nnUNet_results, nnUNet_raw\n\n    predictor = nnUNetPredictor(\n        tile_step_size=0.5,\n        use_gaussian=True,\n        use_mirroring=True,\n        perform_everything_on_device=True,\n        device=torch.device('cuda', 0),\n        verbose=False,\n        verbose_preprocessing=False,\n        allow_tqdm=True\n    )\n    predictor.initialize_from_trained_model_folder(\n        join(nnUNet_results, 'Dataset004_Hippocampus/nnUNetTrainer_5epochs__nnUNetPlans__3d_fullres'),\n        use_folds=(0,),\n        checkpoint_name='checkpoint_final.pth',\n    )\n    # predictor.predict_from_files(join(nnUNet_raw, 'Dataset003_Liver/imagesTs'),\n    #                              join(nnUNet_raw, 'Dataset003_Liver/imagesTs_predlowres'),\n    #                              save_probabilities=False, overwrite=False,\n    #                              num_processes_preprocessing=2, num_processes_segmentation_export=2,\n    #                              folder_with_segs_from_prev_stage=None, num_parts=1, part_id=0)\n    #\n    # # predict a numpy array\n    # from nnunetv2.imageio.simpleitk_reader_writer import SimpleITKIO\n    #\n    # img, props = SimpleITKIO().read_images([join(nnUNet_raw, 'Dataset003_Liver/imagesTr/liver_63_0000.nii.gz')])\n    # ret = predictor.predict_single_npy_array(img, props, None, None, False)\n    #\n    # iterator = predictor.get_data_iterator_from_raw_npy_data([img], None, [props], None, 1)\n    # ret = predictor.predict_from_data_iterator(iterator, False, 1)\n\n    ret = predictor.predict_from_files_sequential(\n        [['/media/isensee/raw_data/nnUNet_raw/Dataset004_Hippocampus/imagesTs/hippocampus_002_0000.nii.gz'], ['/media/isensee/raw_data/nnUNet_raw/Dataset004_Hippocampus/imagesTs/hippocampus_005_0000.nii.gz']],\n        '/home/isensee/temp/tmp', False, True, None\n    )\n\n\n","metadata":{"trusted":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2026-01-23T02:43:57.769018Z","iopub.execute_input":"2026-01-23T02:43:57.769272Z","iopub.status.idle":"2026-01-23T02:43:58.046363Z","shell.execute_reply.started":"2026-01-23T02:43:57.769251Z","shell.execute_reply":"2026-01-23T02:43:58.045685Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile infer.py\n#!/usr/bin/env python3\n# -*- coding: utf-8 -*-\n\nimport os\nimport json\nimport shutil\nimport subprocess\nfrom pathlib import Path\nfrom typing import Optional, Union, Literal, List\n\nimport tifffile\nfrom tqdm.auto import tqdm\n\n# =============================================================================\n# TYPE DEFINITIONS\n# =============================================================================\nEpochs = Literal[1, 5, 10, 20, 50, 100, 250, 500, 750, 1000, 2000, 4000, 8000]\n\n# =============================================================================\n# DEFAULTS (override via CLI)\n# =============================================================================\nDEFAULT_INPUT_DIR = Path(\"/kaggle/input/vesuvius-challenge-surface-detection\")\nDEFAULT_PREPARED_PREPROCESSED_PATH = Path(\"/kaggle/input/vesuvius-surface-nnunet-preprocessed\")\n\nDEFAULT_WORKING_DIR = Path(\"/kaggle/temp\")\nDEFAULT_OUTPUT_ROOT = Path(\"/kaggle/working\")\n\nDEFAULT_DATASET_ID = 100\nDEFAULT_DATASET_NAME = f\"Dataset{DEFAULT_DATASET_ID:03d}_VesuviusSurface\"\n\nDEFAULT_CONFIGURATION = \"3d_fullres\"\nDEFAULT_PLANS_NAME = \"nnUNetResEncUNetMPlans\"\n\nDEFAULT_FOLD: Union[int, str] = \"all\"\nDEFAULT_EPOCHS: Optional[Epochs] = None  # None -> nnUNetTrainer (default 1000 epochs)\n\nDEFAULT_NUM_WORKERS = os.cpu_count() or 4\n\n\n# =============================================================================\n# HELPERS\n# =============================================================================\ndef _get_trainer_name(epochs: Optional[Epochs]) -> str:\n    if epochs is None or epochs == 1000:\n        return \"nnUNetTrainer\"\n    elif epochs == 1:\n        return \"nnUNetTrainer_1epoch\"\n    else:\n        return f\"nnUNetTrainer_{epochs}epochs\"\n\n\ndef create_spacing_json(output_path: Path, spacing=(1.0, 1.0, 1.0)) -> None:\n    output_path.parent.mkdir(parents=True, exist_ok=True)\n    with open(output_path, \"w\") as f:\n        json.dump({\"spacing\": list(spacing)}, f)\n\n\ndef _safe_symlink(src: Path, dst: Path) -> None:\n    dst.parent.mkdir(parents=True, exist_ok=True)\n    if dst.exists() or dst.is_symlink():\n        dst.unlink()\n    dst.symlink_to(src.resolve())\n\n\ndef setup_environment(\n    nnunet_raw: Path,\n    nnunet_preprocessed: Path,\n    nnunet_results: Path,\n    nnunet_compile: str = \"true\",\n    nnunet_use_blosc2: Optional[str] = \"1\",\n) -> None:\n    for d in [nnunet_raw, nnunet_preprocessed, nnunet_results]:\n        d.mkdir(parents=True, exist_ok=True)\n\n    os.environ[\"nnUNet_raw\"] = str(nnunet_raw)\n    os.environ[\"nnUNet_preprocessed\"] = str(nnunet_preprocessed)\n    os.environ[\"nnUNet_results\"] = str(nnunet_results)\n    os.environ[\"nnUNet_compile\"] = nnunet_compile\n\n    if nnunet_use_blosc2 is not None:\n        os.environ[\"nnUNet_USE_BLOSC2\"] = str(nnunet_use_blosc2)\n\n    print(f\"[Env] nnUNet_raw          = {nnunet_raw}\")\n    print(f\"[Env] nnUNet_preprocessed = {nnunet_preprocessed}\")\n    print(f\"[Env] nnUNet_results      = {nnunet_results}\")\n    print(f\"[Env] nnUNet_compile      = {os.environ.get('nnUNet_compile')}\")\n    print(f\"[Env] nnUNet_USE_BLOSC2   = {os.environ.get('nnUNet_USE_BLOSC2', 'not set')}\")\n    print(f\"[Env] NUM_WORKERS         = {DEFAULT_NUM_WORKERS}\")\n\n\ndef link_prepared_preprocessed(\n    prepared_preprocessed_path: Path,\n    nnunet_preprocessed: Path,\n    dataset_name: str,\n) -> bool:\n    \"\"\"\n    Link/copy prepared nnUNet_preprocessed dataset into nnunet_preprocessed/DatasetXXX_...\n    - Copies metadata (*.json/*.pkl/*.txt)\n    - Symlinks heavy data (*.npz/*.npy/*.b2nd)\n    \"\"\"\n    if not prepared_preprocessed_path.exists():\n        print(f\"[Preprocessed] Prepared path not found: {prepared_preprocessed_path}\")\n        return False\n\n    source_dir = prepared_preprocessed_path\n    if not (source_dir / \"dataset.json\").exists():\n        # try find DatasetXXX_* folder inside\n        candidates = list(prepared_preprocessed_path.glob(\"Dataset*\"))\n        if not candidates:\n            print(f\"[Preprocessed] No Dataset* folder found in {prepared_preprocessed_path}\")\n            return False\n        source_dir = candidates[0]\n\n    target_dir = nnunet_preprocessed / dataset_name\n    if target_dir.exists():\n        print(f\"[Preprocessed] Already exists: {target_dir}\")\n        return True\n\n    print(f\"[Preprocessed] Linking from: {source_dir}\")\n    target_dir.mkdir(parents=True, exist_ok=True)\n\n    symlink_suffixes = (\".npz\", \".npy\", \".b2nd\")\n    copied = 0\n    linked = 0\n\n    for src_path in source_dir.rglob(\"*\"):\n        if src_path.is_dir():\n            continue\n\n        rel = src_path.relative_to(source_dir)\n        dst_path = target_dir / rel\n        dst_path.parent.mkdir(parents=True, exist_ok=True)\n\n        if src_path.suffix.lower() in symlink_suffixes:\n            if not dst_path.exists():\n                dst_path.symlink_to(src_path.resolve())\n                linked += 1\n        else:\n            if not dst_path.exists():\n                shutil.copy2(src_path, dst_path)\n                copied += 1\n\n    print(f\"[Preprocessed] copied={copied}, symlinked={linked}\")\n    print(f\"[Preprocessed] target={target_dir}\")\n    return True\n\n\ndef ensure_results_metadata(\n    prepared_preprocessed_path: Path,\n    results_trainer_dir: Path,\n    plans_source_filename: str = \"nnUNetResEncUNetMPlans.json\",\n) -> None:\n    \"\"\"\n    Ensure results trainer dir has:\n      - dataset.json\n      - dataset_fingerprint.json\n      - plans.json\n    \"\"\"\n    results_trainer_dir.mkdir(parents=True, exist_ok=True)\n\n    src_dataset_json = prepared_preprocessed_path / \"dataset.json\"\n    src_fp_json = prepared_preprocessed_path / \"dataset_fingerprint.json\"\n    src_plans_json = prepared_preprocessed_path / plans_source_filename\n\n    # fallback: if prepared_preprocessed_path is a container, try auto-detect\n    if not src_dataset_json.exists():\n        candidates = list(prepared_preprocessed_path.glob(\"Dataset*/dataset.json\"))\n        if candidates:\n            src_dataset_json = candidates[0]\n    if not src_fp_json.exists():\n        candidates = list(prepared_preprocessed_path.glob(\"Dataset*/dataset_fingerprint.json\"))\n        if candidates:\n            src_fp_json = candidates[0]\n    if not src_plans_json.exists():\n        candidates = list(prepared_preprocessed_path.glob(\"**/*Plans*.json\"))\n        if candidates:\n            src_plans_json = candidates[0]\n\n    dst_dataset_json = results_trainer_dir / \"dataset.json\"\n    dst_fp_json = results_trainer_dir / \"dataset_fingerprint.json\"\n    dst_plans_json = results_trainer_dir / \"plans.json\"\n\n    if src_dataset_json.exists() and not dst_dataset_json.exists():\n        shutil.copy(src_dataset_json, dst_dataset_json)\n    if src_fp_json.exists() and not dst_fp_json.exists():\n        shutil.copy(src_fp_json, dst_fp_json)\n    if src_plans_json.exists() and not dst_plans_json.exists():\n        shutil.copy(src_plans_json, dst_plans_json)\n\n    print(f\"[ResultsMeta] ensured in: {results_trainer_dir}\")\n    print(f\"  - dataset.json              : {dst_dataset_json.exists()}\")\n    print(f\"  - dataset_fingerprint.json  : {dst_fp_json.exists()}\")\n    print(f\"  - plans.json                : {dst_plans_json.exists()}\")\n\n\ndef prepare_test_data(\n    input_dir: Path,\n    output_dir: Path,\n    use_symlinks: bool = True,\n    spacing=(1.0, 1.0, 1.0),\n) -> Path:\n    \"\"\"\n    Prepare test TIFF images for nnUNet inference.\n    input_dir/test_images/*.tif -> output_dir/<case>_0000.tif + <case>_0000.json\n    \"\"\"\n    test_images_dir = input_dir / \"test_images\"\n    output_dir.mkdir(parents=True, exist_ok=True)\n\n    if not test_images_dir.exists():\n        raise FileNotFoundError(f\"test_images not found: {test_images_dir}\")\n\n    test_files = sorted(test_images_dir.glob(\"*.tif\"))\n    if not test_files:\n        raise FileNotFoundError(f\"No .tif found in: {test_images_dir}\")\n\n    print(f\"[TestPrep] Found {len(test_files)} test cases. symlink={use_symlinks}\")\n    for img_path in tqdm(test_files, desc=\"Preparing test input\"):\n        case_id = img_path.stem\n        dst_tif = output_dir / f\"{case_id}_0000.tif\"\n        dst_json = output_dir / f\"{case_id}_0000.json\"\n\n        if use_symlinks:\n            if not dst_tif.exists():\n                dst_tif.symlink_to(img_path.resolve())\n        else:\n            if not dst_tif.exists():\n                shutil.copy2(img_path, dst_tif)\n\n        create_spacing_json(dst_json, spacing=spacing)\n\n    return output_dir\n\n\ndef _run_command(cmd: str, name: str, timeout: Optional[int] = None, tail_lines: int = 60) -> None:\n    print(f\"[Run] {name}: {cmd}\")\n    result = subprocess.run(\n        cmd,\n        shell=True,\n        capture_output=True,\n        text=True,\n        timeout=timeout,\n    )\n    if result.returncode != 0:\n        print(f\"[ERROR] {name} failed.\\nSTDERR:\\n{result.stderr[-4000:]}\")\n        raise RuntimeError(f\"{name} failed with code {result.returncode}\")\n    if result.stdout.strip():\n        lines = result.stdout.strip().splitlines()\n        print(\"\\n\".join(lines[-tail_lines:]))\n\n\ndef keep_only_npz_files(output_dir: Path) -> None:\n    \"\"\"\n    Delete all files in output_dir that are NOT .npz.\n    Also removes empty subdirectories afterwards.\n    \"\"\"\n    output_dir = output_dir.resolve()\n    if not output_dir.exists():\n        return\n\n    removed = 0\n    kept = 0\n\n    for p in output_dir.rglob(\"*\"):\n        if p.is_dir():\n            continue\n        if p.suffix.lower() == \".npz\":\n            kept += 1\n            continue\n        try:\n            p.unlink()\n            removed += 1\n        except Exception as e:\n            print(f\"[WARN] Failed to remove {p}: {e}\")\n\n    # remove empty dirs (bottom-up)\n    dirs = [d for d in output_dir.rglob(\"*\") if d.is_dir()]\n    for d in sorted(dirs, reverse=True):\n        try:\n            if not any(d.iterdir()):\n                d.rmdir()\n        except Exception:\n            pass\n\n    print(f\"[Cleanup] keep_only_npz: kept={kept}, removed={removed}, dir={output_dir}\")\n\n\ndef initialize_once(\n    *,\n    input_dir: Path,\n    working_dir: Path,\n    prepared_preprocessed_path: Path,\n    nnunet_preprocessed: Path,\n    nnunet_results: Path,\n    dataset_name: str,\n    trainer: str,\n    plans: str,\n    config: str,\n    reuse_test_cache: bool,\n    test_cache_dir: Path,\n) -> Path:\n    \"\"\"\n    One-time initialization:\n    - link prepared preprocessed dataset into nnUNet_preprocessed\n    - ensure results metadata (dataset.json / fingerprint / plans.json)\n    - prepare test_input cache once (symlinks + json)\n    Returns test_cache_dir.\n    \"\"\"\n    ok = link_prepared_preprocessed(prepared_preprocessed_path, nnunet_preprocessed, dataset_name)\n    if not ok:\n        raise RuntimeError(\n            \"Preprocessed data not available. Provide valid --prepared_preprocessed_path \"\n            \"or place dataset under nnUNet_preprocessed.\"\n        )\n\n    results_trainer_dir = nnunet_results / dataset_name / f\"{trainer}__{plans}__{config}\"\n    ensure_results_metadata(prepared_preprocessed_path, results_trainer_dir)\n\n    test_cache_dir.mkdir(parents=True, exist_ok=True)\n    if reuse_test_cache:\n        if any(test_cache_dir.glob(\"*_0000.tif\")) and any(test_cache_dir.glob(\"*_0000.json\")):\n            print(f\"[TestPrep] Reusing cached test input: {test_cache_dir}\")\n            return test_cache_dir\n\n    if test_cache_dir.exists():\n        shutil.rmtree(test_cache_dir, ignore_errors=True)\n    test_cache_dir.mkdir(parents=True, exist_ok=True)\n\n    prepare_test_data(input_dir, test_cache_dir, use_symlinks=True, spacing=(1.0, 1.0, 1.0))\n    print(f\"[TestPrep] Cached test input prepared at: {test_cache_dir}\")\n    return test_cache_dir\n\n\ndef run_inference_npz(\n    *,\n    ckpt_path: Path,\n    output_dir: Path,\n    test_input_dir: Path,\n    dataset_id: int,\n    dataset_name: str,\n    config: str,\n    fold: Union[int, str],\n    plans: str,\n    epochs: Optional[Epochs],\n    nnunet_results: Path,\n    # your custom predict args (keep your defaults)\n    step_size: float = 0.85,\n    tile_size: int = 192,\n    npp: int = 2,\n    nps: int = 2,\n    timeout: Optional[int] = None,\n    keep_only_npz: bool = False,\n) -> None:\n    \"\"\"\n    Per-model inference:\n    - map external ckpt into expected nnUNet results folder as checkpoint_final.pth\n    - call nnUNetv2_predict with --save_probabilities\n    - outputs (.npz probabilities) stored in output_dir\n    \"\"\"\n    ckpt_path = ckpt_path.expanduser().resolve()\n    output_dir = output_dir.expanduser().resolve()\n    output_dir.mkdir(parents=True, exist_ok=True)\n\n    if not ckpt_path.exists():\n        raise FileNotFoundError(f\"ckpt not found: {ckpt_path}\")\n\n    trainer = _get_trainer_name(epochs)\n\n    results_trainer_dir = nnunet_results / dataset_name / f\"{trainer}__{plans}__{config}\"\n    fold_dir = results_trainer_dir / f\"fold_{fold}\"\n    fold_dir.mkdir(parents=True, exist_ok=True)\n\n    checkpoint_final = fold_dir / \"checkpoint_final.pth\"\n    _safe_symlink(ckpt_path, checkpoint_final)\n    print(f\"[Model] mapped ckpt -> {checkpoint_final}\")\n\n    cmd = (\n        f\"nnUNetv2_predict -d {dataset_id:03d} -c {config} -f {fold} \"\n        f\"-i {test_input_dir} -o {output_dir} -p {plans} -tr {trainer} \"\n        f\"-npp {npp} -nps {nps} --verbose \"\n        f\"-step_size {step_size} -tile_size {tile_size} \"\n        f\"--save_probabilities\"\n    )\n    _run_command(cmd, \"Inference\", timeout=timeout)\n\n    npz_files = list(output_dir.glob(\"*.npz\"))\n    if not npz_files:\n        raise RuntimeError(f\"No .npz probabilities found in output_dir: {output_dir}\")\n    print(f\"[OK] Saved {len(npz_files)} probability maps (.npz) to: {output_dir}\")\n\n    if keep_only_npz:\n        keep_only_npz_files(output_dir)\n\n\ndef _read_ckpt_list_file(path: Path) -> List[Path]:\n    ckpts: List[Path] = []\n    if not path.exists():\n        raise FileNotFoundError(f\"ckpt_list not found: {path}\")\n    with open(path, \"r\") as f:\n        for line in f:\n            p = line.strip()\n            if p:\n                ckpts.append(Path(p))\n    return ckpts\n\n\n# =============================================================================\n# MAIN\n# =============================================================================\ndef main():\n    import argparse\n\n    parser = argparse.ArgumentParser(\n        description=\"nnUNet Vesuvius multi-ckpt inference script: one-time preprocessing/test prep, per-ckpt predict -> .npz probability maps.\"\n    )\n\n    # multiple ckpts\n    parser.add_argument(\"--ckpt\", type=str, action=\"append\", default=[],\n                        help=\"Path to checkpoint .pth. Can be used multiple times.\")\n    parser.add_argument(\"--ckpt_list\", type=str, default=\"\",\n                        help=\"Optional text file listing ckpt paths (one per line).\")\n\n    # output\n    parser.add_argument(\"--output_dir\", type=str, required=True,\n                        help=\"Output root directory. Each ckpt writes into a subfolder <output_dir>/<ckpt_stem>/\")\n\n    # optional overrides\n    parser.add_argument(\"--input_dir\", type=str, default=str(DEFAULT_INPUT_DIR),\n                        help=\"Dataset root containing test_images/\")\n    parser.add_argument(\"--prepared_preprocessed_path\", type=str, default=str(DEFAULT_PREPARED_PREPROCESSED_PATH),\n                        help=\"Prepared nnUNet_preprocessed dataset path (or container dir).\")\n\n    parser.add_argument(\"--working_dir\", type=str, default=str(DEFAULT_WORKING_DIR),\n                        help=\"Working dir for temp files\")\n    parser.add_argument(\"--output_root\", type=str, default=str(DEFAULT_OUTPUT_ROOT),\n                        help=\"Root for nnUNet_results default\")\n\n    parser.add_argument(\"--dataset_id\", type=int, default=DEFAULT_DATASET_ID)\n    parser.add_argument(\"--dataset_name\", type=str, default=DEFAULT_DATASET_NAME)\n\n    parser.add_argument(\"--config\", type=str, default=DEFAULT_CONFIGURATION)\n    parser.add_argument(\"--plans\", type=str, default=DEFAULT_PLANS_NAME)\n\n    parser.add_argument(\"--fold\", type=str, default=str(DEFAULT_FOLD))\n    parser.add_argument(\"--epochs\", type=int, default=(DEFAULT_EPOCHS if DEFAULT_EPOCHS is not None else 0),\n                        help=\"Training epochs for selecting trainer. 0 means default (nnUNetTrainer).\")\n\n    # test cache\n    parser.add_argument(\"--test_cache_dir\", type=str, default=\"\",\n                        help=\"Cache dir for prepared test inputs. Default: <working_dir>/test_input_cache\")\n    parser.add_argument(\"--reuse_test_cache\", action=\"store_true\",\n                        help=\"If set and test_cache_dir exists, skip preparing test data again.\")\n\n    # predict args (your modified nnUNetv2_predict supports these)\n    parser.add_argument(\"--step_size\", type=float, default=0.85)\n    parser.add_argument(\"--tile_size\", type=int, default=192)\n    parser.add_argument(\"--npp\", type=int, default=2)\n    parser.add_argument(\"--nps\", type=int, default=2)\n\n    parser.add_argument(\"--timeout\", type=int, default=0, help=\"0 means no timeout\")\n    parser.add_argument(\"--keep_only_npz\", action=\"store_true\",\n                        help=\"After inference, delete all non-.npz files inside each model output subdir.\")\n\n    args = parser.parse_args()\n\n    # collect ckpts\n    ckpts: List[Path] = [Path(x) for x in args.ckpt]\n    if args.ckpt_list:\n        ckpts.extend(_read_ckpt_list_file(Path(args.ckpt_list)))\n\n    if not ckpts:\n        raise ValueError(\"No checkpoints provided. Use --ckpt (multiple allowed) or --ckpt_list.\")\n\n    # parse fold\n    fold: Union[int, str]\n    if args.fold.isdigit():\n        fold = int(args.fold)\n    else:\n        fold = args.fold\n\n    # parse epochs\n    epochs: Optional[Epochs]\n    if args.epochs in (0, 1000):\n        epochs = None\n    else:\n        epochs = args.epochs  # type: ignore\n\n    timeout = None if args.timeout == 0 else int(args.timeout)\n\n    # paths\n    input_dir = Path(args.input_dir)\n    prepared_preprocessed_path = Path(args.prepared_preprocessed_path)\n    working_dir = Path(args.working_dir)\n    output_root = Path(args.output_root)\n    output_dir_root = Path(args.output_dir)\n\n    # nnUNet layout\n    nnunet_base = working_dir / \"nnUNet_data\"\n    nnunet_raw = nnunet_base / \"nnUNet_raw\"\n    nnunet_preprocessed = nnunet_base / \"nnUNet_preprocessed\"\n    nnunet_results = output_root / \"nnUNet_results\"\n\n    # env\n    setup_environment(nnunet_raw, nnunet_preprocessed, nnunet_results, nnunet_use_blosc2=\"1\")\n\n    # one-time init\n    trainer = _get_trainer_name(epochs)\n    test_cache_dir = Path(args.test_cache_dir) if args.test_cache_dir else (working_dir / \"test_input_cache\")\n\n    test_input_dir = initialize_once(\n        input_dir=input_dir,\n        working_dir=working_dir,\n        prepared_preprocessed_path=prepared_preprocessed_path,\n        nnunet_preprocessed=nnunet_preprocessed,\n        nnunet_results=nnunet_results,\n        dataset_name=args.dataset_name,\n        trainer=trainer,\n        plans=args.plans,\n        config=args.config,\n        reuse_test_cache=bool(args.reuse_test_cache),\n        test_cache_dir=test_cache_dir,\n    )\n\n    # loop models (serial)\n    output_dir_root.mkdir(parents=True, exist_ok=True)\n\n    for ckpt_path in ckpts:\n        tag = ckpt_path.stem\n        model_out = output_dir_root / tag\n\n        print(\"\\n\" + \"=\" * 80)\n        print(f\"[Predict] ckpt={ckpt_path}\")\n        print(f\"[Predict] out ={model_out}\")\n        print(\"=\" * 80)\n\n        run_inference_npz(\n            ckpt_path=ckpt_path,\n            output_dir=model_out,\n            test_input_dir=test_input_dir,\n            dataset_id=args.dataset_id,\n            dataset_name=args.dataset_name,\n            config=args.config,\n            fold=fold,\n            plans=args.plans,\n            epochs=epochs,\n            nnunet_results=nnunet_results,\n            step_size=float(args.step_size),\n            tile_size=int(args.tile_size),\n            npp=int(args.npp),\n            nps=int(args.nps),\n            timeout=timeout,\n            keep_only_npz=bool(args.keep_only_npz),\n        )\n\n    print(\"\\n[Done] All checkpoints finished.\")\n    print(f\"[Done] Outputs under: {output_dir_root}\")\n\n\nif __name__ == \"__main__\":\n    main()\n","metadata":{"_uuid":"463f2b49-d26e-45c9-9472-d0779a99322e","_cell_guid":"07aca3f2-b73c-47c9-916c-0683c19971d1","trusted":true,"execution":{"iopub.status.busy":"2026-01-23T02:43:58.047295Z","iopub.execute_input":"2026-01-23T02:43:58.047526Z","iopub.status.idle":"2026-01-23T02:43:58.06967Z","shell.execute_reply.started":"2026-01-23T02:43:58.047504Z","shell.execute_reply":"2026-01-23T02:43:58.069122Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!python infer.py \\\n  --ckpt /kaggle/input/nnunet-updateddata/4000epochs_iter6_checkpoint_best.pth \\\n  --ckpt /kaggle/input/nnunet-updateddata/4000epochs_org_prob0.8_checkpoint_best.pth \\\n  --output_dir /kaggle/temp/prob_npz_all \\\n  --reuse_test_cache \\\n  --keep_only_npz  \\\n  --step_size 0.85 \\\n  --tile_size 192 \\\n  --npp 2 --nps 2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T02:43:58.070414Z","iopub.execute_input":"2026-01-23T02:43:58.070649Z","iopub.status.idle":"2026-01-23T02:49:10.866707Z","shell.execute_reply.started":"2026-01-23T02:43:58.070627Z","shell.execute_reply":"2026-01-23T02:49:10.865803Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"weights = {\n    \"4000epochs_iter6_checkpoint_best\": 0.65,\n    \"4000epochs_org_prob0.8_checkpoint_best\": 0.35,\n    #\"4000_iter12_checkpoint_best\":0.3\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T02:49:10.869211Z","iopub.execute_input":"2026-01-23T02:49:10.869498Z","iopub.status.idle":"2026-01-23T02:49:10.873562Z","shell.execute_reply.started":"2026-01-23T02:49:10.869472Z","shell.execute_reply":"2026-01-23T02:49:10.872946Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nfrom skimage.morphology import (\n    ball,\n    disk,\n    binary_closing,\n    remove_small_objects,\n    remove_small_holes,\n)\nimport pandas as pd\nimport tifffile\nfrom pathlib import Path\nfrom typing import Optional, Union, Literal, List\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T02:49:10.87437Z","iopub.execute_input":"2026-01-23T02:49:10.874599Z","iopub.status.idle":"2026-01-23T02:49:11.362527Z","shell.execute_reply.started":"2026-01-23T02:49:10.874579Z","shell.execute_reply":"2026-01-23T02:49:11.361875Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile post.py\nimport os\nimport numpy as np\nfrom skimage.morphology import (\n    ball,\n    disk,\n    binary_closing,\n    remove_small_objects,\n    remove_small_holes,\n)\nimport pandas as pd\nimport tifffile\nfrom pathlib import Path\nfrom typing import Optional, Union, Literal, List\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\n\nimport numpy as np\nfrom skimage.morphology import (\n    remove_small_objects,\n    remove_small_holes,\n    binary_closing,\n    disk,\n)\n\nimport numpy as np\nfrom skimage.morphology import (\n    remove_small_objects,\n    remove_small_holes,\n    binary_closing,\n    disk,\n)\n\ndef prune_isolated_z(pred_bin: np.ndarray) -> np.ndarray:\n    Z = pred_bin.shape[0]\n    keep = np.zeros_like(pred_bin, dtype=bool)\n\n    for z in range(1, Z - 1):\n        keep[z] = pred_bin[z] & (pred_bin[z - 1] | pred_bin[z + 1])\n\n    return keep\n\n\ndef topo_postprocess_volume(\n    pred_volume: np.ndarray,\n    min_component_size: int = 2000,\n    closing_radius_xy: int = 1,\n    hole_area_threshold_2d: int = 64,\n    hole_area_threshold_3d: int = 64,\n) -> np.ndarray:\n\n    # --- 0. binarize ---\n    pred_bin = pred_volume > 0\n\n    # --- 1. remove tiny 3D blobs ---\n    pred_bin = remove_small_objects(pred_bin, min_size=min_component_size)\n\n    # --- 2. strong 2D morphology per slice ---\n    se2d = disk(closing_radius_xy)\n    for z in range(pred_bin.shape[0]):\n        pred_bin[z] = binary_closing(pred_bin[z], se2d)\n        # pred_bin[z] = remove_small_holes(\n        #     pred_bin[z], area_threshold=hole_area_threshold_2d\n        # )\n\n    # --- 3. Z-consistency pruning (CRITICAL) ---\n    pred_bin = prune_isolated_z(pred_bin)\n\n    # --- 4. anisotropic 3D closing (XY only) ---\n    se3d = np.zeros((1, 3, 3), dtype=bool)\n    se3d[0, :, :] = True\n    pred_bin = binary_closing(pred_bin, se3d)\n\n    # --- 5. conservative 3D hole filling ---\n    # pred_bin = remove_small_holes(\n    #     pred_bin, area_threshold=hole_area_threshold_3d\n    # )\n\n    return pred_bin.astype(np.uint8)\n\n    \n\n    \nfrom typing import Mapping, Tuple\n\nimport numpy as np\n\nDEFAULT_TARGET_DTYPE = np.float32\nEPS = 1e-8\n\n__all__ = [\n    \"normalize_minmax\",\n    \"normalize_zscore\",\n    \"normalize_ct\",\n    \"normalize_robust\",\n    \"DEFAULT_TARGET_DTYPE\",\n    \"NORMALIZATION_FUNCTIONS\",\n]\n\n\n# --------------------------------------------------------------------------- #\n# Helpers\n# --------------------------------------------------------------------------- #\n\ndef _prepare_image(image: np.ndarray, target_dtype: np.dtype | type = DEFAULT_TARGET_DTYPE) -> np.ndarray:\n    arr = np.asarray(image)\n    if target_dtype is not None:\n        arr = arr.astype(target_dtype, copy=False)\n    return arr\n\n\ndef _prepare_mask(mask: np.ndarray, image_shape: Tuple[int, ...]) -> np.ndarray:\n    \"\"\"\n    Broadcast the provided mask to match the image shape, ensuring boolean dtype.\n    Raises ValueError when broadcasting is impossible.\n    \"\"\"\n    mask_arr = np.asarray(mask)\n    if mask_arr.ndim > len(image_shape):\n        raise ValueError(\n            f\"Mask with shape {mask_arr.shape} has higher dimensionality than image shape {image_shape}\"\n        )\n\n    if mask_arr.ndim < len(image_shape):\n        expand_dims = (1,) * (len(image_shape) - mask_arr.ndim)\n        mask_arr = mask_arr.reshape(expand_dims + mask_arr.shape)\n\n    try:\n        broadcast = np.broadcast_to(mask_arr, image_shape, subok=True)\n    except ValueError as exc:\n        raise ValueError(\n            f\"Mask with shape {mask_arr.shape} cannot broadcast to image shape {image_shape}\"\n        ) from exc\n\n    return broadcast.astype(bool, copy=False)\n\n\ndef _select_valid_region(\n    image: np.ndarray,\n    mask: np.ndarray | None,\n    use_mask: bool,\n) -> Tuple[np.ndarray, np.ndarray | None]:\n    \"\"\"\n    Return the flattened view of valid pixels and the boolean mask used.\n\n    When no valid mask pixels exist, the full image is returned and mask_bool is None.\n    \"\"\"\n    if use_mask and mask is not None:\n        mask_bool = _prepare_mask(mask, image.shape)\n        if np.any(mask_bool):\n            return image[mask_bool], mask_bool\n    return image.reshape(-1), None\n\n\n# --------------------------------------------------------------------------- #\n# Normalisation functions\n# --------------------------------------------------------------------------- #\n\ndef normalize_minmax(\n    image: np.ndarray,\n    *,\n    target_dtype: np.dtype | type = DEFAULT_TARGET_DTYPE,\n) -> np.ndarray:\n    \"\"\"\n    Min-max normalisation that rescales intensities into the [0, 1] range.\n    \"\"\"\n    arr = _prepare_image(image, target_dtype)\n    min_val = float(arr.min())\n    max_val = float(arr.max())\n\n    if max_val > min_val:\n        arr -= min_val\n        arr /= max(max_val - min_val, EPS)\n    else:\n        arr.fill(0.0)\n\n    return arr\n\n\ndef normalize_zscore(\n    image: np.ndarray,\n    *,\n    mask: np.ndarray | None = None,\n    use_mask: bool = False,\n    target_dtype: np.dtype | type = DEFAULT_TARGET_DTYPE,\n) -> np.ndarray:\n    \"\"\"\n    Standard score normalisation (z-score).\n\n    Parameters\n    ----------\n    image : np.ndarray\n        Input image.\n    mask : np.ndarray, optional\n        Optional mask restricting the statistics to masked pixels.\n    use_mask : bool\n        When true and mask is provided, restrict mean/std computation to masked region.\n    \"\"\"\n    arr = _prepare_image(image, target_dtype)\n    valid, mask_bool = _select_valid_region(arr, mask, use_mask)\n\n    mean = float(valid.mean()) if valid.size else 0.0\n    std = float(valid.std()) if valid.size else 0.0\n    std = max(std, EPS)\n\n    if mask_bool is not None:\n        arr_masked = arr[mask_bool]\n        arr_masked -= mean\n        arr_masked /= std\n        arr[mask_bool] = arr_masked\n    else:\n        arr -= mean\n        arr /= std\n\n    return arr\n\n\ndef normalize_ct(\n    image: np.ndarray,\n    *,\n    intensity_properties: Mapping[str, float] | None,\n    target_dtype: np.dtype | type = DEFAULT_TARGET_DTYPE,\n) -> np.ndarray:\n    \"\"\"\n    CT-style normalisation using global percentiles, mean, and standard deviation.\n    \"\"\"\n    if intensity_properties is None:\n        raise ValueError(\"normalize_ct requires intensity_properties\")\n\n    required_keys = (\"mean\", \"std\", \"percentile_00_5\", \"percentile_99_5\")\n    missing = [key for key in required_keys if key not in intensity_properties]\n    if missing:\n        raise ValueError(\n            f\"normalize_ct missing intensity properties: {missing}. \"\n            f\"Provided keys: {list(intensity_properties.keys())}\"\n        )\n\n    arr = _prepare_image(image, target_dtype)\n    mean_intensity = float(intensity_properties[\"mean\"])\n    std_intensity = max(float(intensity_properties[\"std\"]), EPS)\n    lower_bound = float(intensity_properties[\"percentile_00_5\"])\n    upper_bound = float(intensity_properties[\"percentile_99_5\"])\n\n    np.clip(arr, lower_bound, upper_bound, out=arr)\n    arr -= mean_intensity\n    arr /= std_intensity\n\n    return arr\n\n\ndef normalize_robust(\n    image: np.ndarray,\n    *,\n    mask: np.ndarray | None = None,\n    use_mask: bool = False,\n    percentile_lower: float = 1.0,\n    percentile_upper: float = 99.0,\n    clip_values: bool = True,\n    target_dtype: np.dtype | type = DEFAULT_TARGET_DTYPE,\n) -> np.ndarray:\n    \"\"\"\n    Robust normalisation based on the median and MAD (Median Absolute Deviation).\n    \"\"\"\n    arr = _prepare_image(image, target_dtype)\n    valid, mask_bool = _select_valid_region(arr, mask, use_mask)\n\n    if valid.size == 0:\n        return arr\n\n    std_preclip = float(np.std(valid))\n    min_preclip = float(np.min(valid))\n    max_preclip = float(np.max(valid))\n\n    lower_val = upper_val = None\n    if clip_values and valid.size > 0:\n        lower_val = float(np.percentile(valid, percentile_lower))\n        upper_val = float(np.percentile(valid, percentile_upper))\n        np.clip(arr, lower_val, upper_val, out=arr)\n        valid, mask_bool = _select_valid_region(arr, mask, use_mask)\n        if valid.size == 0:\n            return arr\n\n    median = float(np.median(valid))\n    mad = float(np.median(np.abs(valid - median)))\n    scaled_mad = 1.4826 * mad\n\n    if not np.isfinite(scaled_mad) or scaled_mad < 1e-6:\n        percentile_span = None\n        if lower_val is not None and upper_val is not None:\n            percentile_span = upper_val - lower_val\n        else:\n            percentile_span = max_preclip - min_preclip\n\n        fallback_values = []\n        if np.isfinite(std_preclip):\n            fallback_values.append(abs(std_preclip))\n        if percentile_span is not None and np.isfinite(percentile_span):\n            fallback_values.append(abs(percentile_span) / 2.0)\n\n        scaled_mad = next(\n            (candidate for candidate in fallback_values if candidate >= 1e-6),\n            None,\n        )\n\n        if scaled_mad is None or not np.isfinite(scaled_mad):\n            scaled_mad = 1.0\n\n    if mask_bool is not None:\n        arr_masked = arr[mask_bool]\n        arr_masked -= median\n        arr_masked /= scaled_mad\n        arr[mask_bool] = arr_masked\n    else:\n        arr -= median\n        arr /= scaled_mad\n\n    np.nan_to_num(arr, copy=False, nan=0.0, posinf=0.0, neginf=0.0)\n    return arr\n\n\ndef _normalize_identity(\n    image: np.ndarray,\n    *,\n    target_dtype: np.dtype | type = DEFAULT_TARGET_DTYPE,\n    **_: object,\n) -> np.ndarray:\n    \"\"\"No-op normalisation helper.\"\"\"\n    arr = _prepare_image(image, target_dtype)\n    return arr\n\n\nNORMALIZATION_FUNCTIONS = {\n    \"zscore\": normalize_zscore,\n    \"ct\": normalize_ct,\n    \"rescale_to_01\": normalize_minmax,\n    \"minmax\": normalize_minmax,\n    \"robust\": normalize_robust,\n    \"none\": _normalize_identity,\n}\n\nimport numpy as np\nfrom numpy import linalg as LA\n\n\ndef divide_nonzero(array1, array2, eps=1e-10):\n    denominator = np.copy(array2)\n    denominator[denominator == 0] = eps\n    return np.divide(array1, denominator)\n\nimport numpy as np\nfrom scipy.ndimage.filters import gaussian_filter\n\ndef hessian_curvature_2d(image, gauss_sigma=2, sigma=6):\n    image_smoothed = gaussian_filter(image, sigma=gauss_sigma)\n    image_smoothed = normalize_minmax(image_smoothed)\n\n    joint_hessian = np.zeros((image.shape[0], image.shape[1], 2, 2), dtype=float)\n\n    Dy = np.gradient(image_smoothed, axis=0, edge_order=2)\n    joint_hessian[:, :, 1, 1] = np.gradient(Dy, axis=0, edge_order=2)\n    joint_hessian[:, :, 0, 1] = np.gradient(Dy, axis=1, edge_order=2)\n    del Dy\n\n    Dx = np.gradient(image_smoothed, axis=1, edge_order=2)\n    joint_hessian[:, :, 0, 0] = np.gradient(Dx, axis=1, edge_order=2)\n    joint_hessian[:, :, 1, 0] = joint_hessian[:, :, 0, 1]\n    del Dx\n\n    joint_hessian = joint_hessian * (sigma ** 2)\n    zero_mask = np.trace(joint_hessian, axis1=2, axis2=3) == 0\n    return joint_hessian, zero_mask\n\n\ndef hessian_curvature_3d(volume, gauss_sigma=2, sigma=6):\n    volume_smoothed = gaussian_filter(volume, sigma=gauss_sigma)\n    volume_smoothed = normalize_minmax(volume_smoothed)\n\n    joint_hessian = np.zeros((volume.shape[0], volume.shape[1], volume.shape[2], 3, 3), dtype=float)\n\n    Dz = np.gradient(volume_smoothed, axis=0, edge_order=2)\n    joint_hessian[:, :, :, 2, 2] = np.gradient(Dz, axis=0, edge_order=2)\n    del Dz\n\n    Dy = np.gradient(volume_smoothed, axis=1, edge_order=2)\n    joint_hessian[:, :, :, 1, 1] = np.gradient(Dy, axis=1, edge_order=2)\n    joint_hessian[:, :, :, 1, 2] = np.gradient(Dy, axis=0, edge_order=2)\n    del Dy\n\n    Dx = np.gradient(volume_smoothed, axis=2, edge_order=2)\n    joint_hessian[:, :, :, 0, 0] = np.gradient(Dx, axis=2, edge_order=2)\n    joint_hessian[:, :, :, 0, 1] = np.gradient(Dx, axis=1, edge_order=2)\n    joint_hessian[:, :, :, 0, 2] = np.gradient(Dx, axis=0, edge_order=2)\n    del Dx\n\n    joint_hessian = joint_hessian * (sigma ** 2)\n    zero_mask = np.trace(joint_hessian, axis1=3, axis2=4) == 0\n    return joint_hessian, zero_mask\n\ndef detect_ridges_2d(image, gamma=1.5, beta=0.5, gauss_sigma=2, sigma=6):\n    joint_hessian, zero_mask = hessian_curvature_2d(image, gauss_sigma, sigma)\n    eigvals = LA.eigvalsh(joint_hessian, 'U')\n    idxs = np.argsort(np.abs(eigvals), axis=-1)\n    eigvals = np.take_along_axis(eigvals, idxs, axis=-1)\n    eigvals[zero_mask, :] = 0\n\n    L1 = np.abs(eigvals[:, :, 0])\n    L2 = eigvals[:, :, 1]\n    L2abs = np.abs(L2)\n\n    S = np.sqrt(np.square(eigvals).sum(axis=-1))\n    background_term = 1 - np.exp(-0.5 * np.square(S / gamma))\n\n    Rb = divide_nonzero(L1, L2abs)\n    blob_term = np.exp(-0.5 * np.square(Rb / beta))\n\n    ridges = background_term * blob_term\n    ridges[L2 > 0] = 0\n    return ridges\n\ndef detect_ridges_3d(volume, gamma=1.5, beta1=0.5, beta2=0.5, gauss_sigma=2, sigma=6):\n    joint_hessian, zero_mask = hessian_curvature_3d(volume, gauss_sigma, sigma)\n    eigvals = LA.eigvalsh(joint_hessian, 'U')\n    idxs = np.argsort(np.abs(eigvals), axis=-1)\n    eigvals = np.take_along_axis(eigvals, idxs, axis=-1)\n    eigvals[zero_mask, :] = 0\n\n    L1 = np.abs(eigvals[:, :, :, 0])\n    L2 = np.abs(eigvals[:, :, :, 1])\n    L3 = eigvals[:, :, :, 2]\n    L3abs = np.abs(L3)\n\n    S = np.sqrt(np.square(eigvals).sum(axis=-1))\n    background_term = 1 - np.exp(-0.5 * np.square(S / gamma))\n\n    Ra = divide_nonzero(L2, L3abs)\n    planar_term = np.exp(-0.5 * np.square(Ra / beta1))\n\n    Rb = divide_nonzero(L1, np.sqrt(L2 * L3abs))\n    blob_term = np.exp(-0.5 * np.square(Rb / beta2))\n\n    ridges = background_term * planar_term * blob_term\n    ridges[L3 > 0] = 0\n    return ridges\n\n#!/usr/bin/env python\nimport os\nimport glob\nimport argparse\nimport multiprocessing\n\nimport numpy as np\nimport tifffile\nfrom tqdm import tqdm\n\nfrom scipy.ndimage import distance_transform_edt\nimport numpy as np\n\ndef dilate_by_inverse_edt(binary_volume, dilation_distance):\n    eps = 1e-6\n    edt = distance_transform_edt(1 - binary_volume)\n    inv_edt = 1.0 / (edt + eps)\n    threshold = 1.0 / dilation_distance\n    dilated = (inv_edt > threshold).astype(np.uint8)\n    return dilated\n    \n\n\n\ndef process_file(file_path, output_folder, dilation_distance=3, ridge_threshold=0.5):\n    try:\n        volume = tifffile.imread(file_path)\n        \n        # Check if input is 2D or 3D\n        if volume.ndim == 2:\n            # 2D processing\n            binary_volume = (volume > 0).astype(np.uint8)\n            dilated_volume = dilate_by_inverse_edt(binary_volume, dilation_distance)\n            dilated_float = dilated_volume.astype(np.float32)\n            ridges = detect_ridges_2d(dilated_float)\n            binary_ridges = (ridges > ridge_threshold).astype(np.uint8)\n        elif volume.ndim == 3:\n            # 3D processing\n            binary_volume = (volume > 0).astype(np.uint8)\n            dilated_volume = dilate_by_inverse_edt(binary_volume, dilation_distance)\n            dilated_float = dilated_volume.astype(np.float32)\n            ridges = detect_ridges_3d(dilated_float)\n            binary_ridges = (ridges > ridge_threshold).astype(np.uint8)\n        else:\n            raise ValueError(f\"Unsupported image dimensions: {volume.ndim}. Only 2D and 3D images are supported.\")\n        binary_ridges = topo_postprocess_volume(binary_ridges)\n        filename = os.path.basename(file_path)\n        output_path = os.path.join(output_folder, filename)\n        tifffile.imwrite(output_path, binary_ridges, compression='packbits')\n        return file_path, True\n    except Exception as e:\n        return file_path, False, str(e)\n\n\n# Define a top-level worker function to avoid lambda pickling issues\ndef worker(args):\n    return process_file(*args)\n\n\n# ---------------------------\n# Main Routine Using Multiprocessing and tqdm\n# ---------------------------\ndef main():\n    parser = argparse.ArgumentParser(\n        description=\"Process a folder of 2D or 3D TIFF files with inverse EDT dilation and custom ridge detection.\"\n    )\n    parser.add_argument(\"input_folder\", help=\"Folder containing input TIFF files (2D or 3D).\")\n    parser.add_argument(\"output_folder\", help=\"Folder where processed TIFF files will be saved.\")\n    parser.add_argument(\"--dilation_distance\", type=float, default=1, help=\"Dilation distance (in pixels/voxels).\")\n    parser.add_argument(\"--ridge_threshold\", type=float, default=0.75,\n                        help=\"Threshold for binarizing the ridge detection.\")\n    parser.add_argument(\"--num_workers\", type=int, default=multiprocessing.cpu_count(),\n                        help=\"Number of worker processes.\")\n    args = parser.parse_args()\n    \n    if not os.path.exists(args.output_folder):\n        os.makedirs(args.output_folder)\n\n    file_list = glob.glob(os.path.join(args.input_folder, \"*.tif\"))\n    if not file_list:\n        print(f\"No .tif files found in {args.input_folder}\")\n        return\n\n    tasks = [\n        (f, args.output_folder, args.dilation_distance, args.ridge_threshold)\n        for f in file_list\n    ]\n\n    with multiprocessing.Pool(args.num_workers) as pool:\n        results = list(\n            tqdm(\n                pool.imap_unordered(worker, tasks),\n                total=len(tasks),\n                desc=\"Processing Files\"\n            )\n        )\n\n    successes = [res for res in results if res[1] is True]\n    failures = [res for res in results if res[1] is not True]\n\n    print(f\"Processed {len(successes)} files successfully.\")\n    if failures:\n        print(f\"{len(failures)} files failed:\")\n        for fail in failures:\n            print(f\"File: {fail[0]}, Error: {fail[2]}\")\n\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T02:49:11.363315Z","iopub.execute_input":"2026-01-23T02:49:11.363629Z","iopub.status.idle":"2026-01-23T02:49:11.374275Z","shell.execute_reply.started":"2026-01-23T02:49:11.36361Z","shell.execute_reply":"2026-01-23T02:49:11.373574Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"output = '/kaggle/temp/prob_npz_all'\nensemble_out = '/kaggle/temp/predictions_tiff'\nos.mkdir(ensemble_out)\nOUTPUT_DIR = Path(\"/kaggle/working\")\ndir_list = os.listdir(output)\ndir_list = [os.path.join(output, i) for i in dir_list]\nids = list(pd.read_csv('/kaggle/input/vesuvius-challenge-surface-detection/test.csv')['id'])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T02:49:11.375087Z","iopub.execute_input":"2026-01-23T02:49:11.375428Z","iopub.status.idle":"2026-01-23T02:49:11.408103Z","shell.execute_reply.started":"2026-01-23T02:49:11.375403Z","shell.execute_reply":"2026-01-23T02:49:11.40739Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"weights_and_dir = []\nfor d in dir_list:\n    weight = weights[os.path.basename(d)]\n    weights_and_dir.append((weight, d))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T02:49:11.408777Z","iopub.execute_input":"2026-01-23T02:49:11.40899Z","iopub.status.idle":"2026-01-23T02:49:11.412839Z","shell.execute_reply.started":"2026-01-23T02:49:11.408968Z","shell.execute_reply":"2026-01-23T02:49:11.41232Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for id in ids:\n    res = []\n    for (weight, out) in weights_and_dir:\n        npz_path = os.path.join(out, f'{id}.npz')\n        probs = np.load(npz_path)['probabilities']\n        res.append(probs * weight)\n    res = np.stack(res, axis=0)          \n    avg_probs = res.sum(axis=0)\n    avg_probs = np.argmax(avg_probs, axis=0).astype(np.uint8)\n    tifffile.imwrite(os.path.join(ensemble_out,f'{id}.tif'), avg_probs)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T02:49:11.413543Z","iopub.execute_input":"2026-01-23T02:49:11.413806Z","iopub.status.idle":"2026-01-23T02:49:15.627972Z","shell.execute_reply.started":"2026-01-23T02:49:11.413783Z","shell.execute_reply":"2026-01-23T02:49:15.627352Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import shutil\nimport os\n\n# 1. Define where you found the files\npreprocessed_dir = \"/kaggle/input/vesuvius-surface-nnunet-preprocessed\"\n\n# 2. Define where the model results are\nresults_dir = \"/kaggle/working/nnUNet_results/Dataset100_VesuviusSurface/nnUNetTrainer__nnUNetResEncUNetMPlans__3d_cascade_fullres\"\nos.makedirs(results_dir, exist_ok=True)\n\n# 3. Copy and Rename as needed\n# nnU-Net prediction ALWAYS looks for \"dataset.json\" and \"plans.json\"\nshutil.copy(os.path.join(preprocessed_dir, \"dataset.json\"), \n            os.path.join(results_dir, \"dataset.json\"))\n\nshutil.copy(os.path.join(preprocessed_dir, \"dataset_fingerprint.json\"), \n            os.path.join(results_dir, \"dataset_fingerprint.json\"))\n\n# IMPORTANT: Rename your specific plans file to the generic 'plans.json'\nshutil.copy(os.path.join(preprocessed_dir, \"nnUNetResEncUNetMPlans.json\"), \n            os.path.join(results_dir, \"plans.json\"))\n\nprint(\"✅ Files moved and renamed. Your results folder is now complete!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T02:55:19.870684Z","iopub.execute_input":"2026-01-23T02:55:19.871328Z","iopub.status.idle":"2026-01-23T02:55:19.884218Z","shell.execute_reply.started":"2026-01-23T02:55:19.871304Z","shell.execute_reply":"2026-01-23T02:55:19.883384Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess\nfrom pathlib import Path\n\ncascade_ckpt = Path(\"/kaggle/input/models/arunodhayan/cascade-updated/pytorch/default/5/80003dlowres_80003dcascade_checkpoint_final.pth\")\n\ndataset_id = 100\nplans = \"nnUNetResEncUNetMPlans\"\ntrainer = \"nnUNetTrainer\"              \nconfig = \"3d_cascade_fullres\"\nfold = \"all\"\n\nos.environ[\"nnUNet_raw\"] = \"/kaggle/temp/nnUNet_data/nnUNet_raw\"\nos.environ[\"nnUNet_preprocessed\"] = \"/kaggle/temp/nnUNet_data/nnUNet_preprocessed\"\nos.environ[\"nnUNet_results\"] = \"/kaggle/working/nnUNet_results\"\nos.environ[\"nnUNet_compile\"] = \"true\"\nos.environ[\"nnUNet_USE_BLOSC2\"] = \"1\"\n\nnnunet_results = Path(\"/kaggle/working/nnUNet_results\")\ndataset_name = f\"Dataset{dataset_id:03d}_VesuviusSurface\"\n\ntarget_fold_dir = nnunet_results / dataset_name / f\"{trainer}__{plans}__{config}\" / f\"fold_{fold}\"\ntarget_fold_dir.mkdir(parents=True, exist_ok=True)\ntarget_ckpt = target_fold_dir / \"checkpoint_final.pth\"\nif target_ckpt.exists() or target_ckpt.is_symlink():\n    target_ckpt.unlink()\ntarget_ckpt.symlink_to(cascade_ckpt.resolve())\n\ntest_input_dir = Path(\"/kaggle/temp/test_input_cache\")              \nprev_stage_dir = Path(ensemble_out)         \nfinal_out = Path(\"/kaggle/temp/final_cascade_out\")\nfinal_out.mkdir(parents=True, exist_ok=True)\n\ncmd = (\n    f\"nnUNetv2_predict -d {dataset_id:03d} -c {config} -f {fold} \"\n    f\"-i {test_input_dir} -o {final_out} -p {plans} -tr {trainer} \"\n    f\"-npp 2 -nps 2 --verbose \"\n    f\"-step_size 0.85 -tile_size 192 \"\n    f\"-prev_stage_predictions {prev_stage_dir}\"\n)\nprint(cmd)\nsubprocess.run(cmd, shell=True, check=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T02:55:21.602218Z","iopub.execute_input":"2026-01-23T02:55:21.602457Z","iopub.status.idle":"2026-01-23T02:57:21.383801Z","shell.execute_reply.started":"2026-01-23T02:55:21.602439Z","shell.execute_reply":"2026-01-23T02:57:21.382985Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!python post.py /kaggle/temp/final_cascade_out /kaggle/working/predictions_tiff --dilation_distance 1 --ridge_threshold 0.20","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T02:57:45.624507Z","iopub.execute_input":"2026-01-23T02:57:45.624977Z","iopub.status.idle":"2026-01-23T02:58:36.67314Z","shell.execute_reply.started":"2026-01-23T02:57:45.624956Z","shell.execute_reply":"2026-01-23T02:58:36.672379Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_three_axis_cuts(image_vol_path, mask_vol_path):\n    \"\"\"Plots the middle slice of the XY, XZ, and YZ planes for both the image volume and the predicted mask.\"\"\"\n    print(f\"Visualizing cuts for: {os.path.basename(image_vol_path)}\")\n    # Load volumes\n    image_vol = tifffile.imread(image_vol_path)\n    mask_vol = tifffile.imread(mask_vol_path)\n    \n    # Get dimensions\n    d, h, w = image_vol.shape\n    z_mid, y_mid, x_mid = d // 2, h // 2, w // 2\n    \n    # Extract slices\n    slices = {\n        'XY Plane (Z-axis)': (image_vol[z_mid, :, :], mask_vol[z_mid, :, :]),\n        'XZ Plane (Y-axis)': (image_vol[:, y_mid, :], mask_vol[:, y_mid, :]),\n        'YZ Plane (X-axis)': (image_vol[:, :, x_mid], mask_vol[:, :, x_mid])\n    }\n    \n    fig, axes = plt.subplots(3, 2, figsize=(12, 15))\n    for i, (plane_name, (img_slice, mask_slice)) in enumerate(slices.items()):\n        # Image Volume\n        axes[i, 0].imshow(img_slice, cmap='gray')\n        axes[i, 0].set_title(f\"{plane_name} - Image Volume\")\n        axes[i, 0].axis('off')\n        \n        # Mask\n        axes[i, 1].imshow(mask_slice, cmap='gray')\n        axes[i, 1].set_title(f\"{plane_name} - Predicted Mask\")\n        axes[i, 1].axis('off')\n        \n    plt.tight_layout()\n    plt.show()\n    \nif not os.getenv(\"KAGGLE_IS_COMPETITION_RERUN\"):\n    mask_path = os.path.join('/kaggle/working/predictions_tiff', f'{id}.tif')\n    image_path = os.path.join('/kaggle/input/vesuvius-challenge-surface-detection/test_images', f'{id}.tif')\n    plot_three_axis_cuts(image_path, mask_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T02:58:36.674746Z","iopub.execute_input":"2026-01-23T02:58:36.67499Z","iopub.status.idle":"2026-01-23T02:58:37.828659Z","shell.execute_reply.started":"2026-01-23T02:58:36.674968Z","shell.execute_reply":"2026-01-23T02:58:37.827877Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def generate_submission(\n    predictions_tiff_dir: Path = OUTPUT_DIR / \"predictions_tiff\",\n    output_zip: Path = OUTPUT_DIR / \"submission.zip\",\n    delete_after_zip: bool = True  # Default True to save space on Kaggle\n) -> Optional[Path]:\n    \"\"\"\n    Create submission ZIP from TIFF predictions.\n    \n    Args:\n        predictions_tiff_dir: Directory containing predicted TIFF files\n        output_zip: Output ZIP file path\n        delete_after_zip: Delete TIFF files after adding to ZIP (saves space)\n    \n    Returns:\n        Path to submission ZIP if successful, None otherwise\n    \"\"\"\n    import zipfile\n    \n    if not predictions_tiff_dir.exists():\n        print(f\"ERROR: Predictions directory not found: {predictions_tiff_dir}\")\n        print(\"Run inference first!\")\n        return None\n    \n    tiff_files = sorted(predictions_tiff_dir.glob(\"*.tif\"))\n    \n    if not tiff_files:\n        print(f\"No TIFF files found in {predictions_tiff_dir}\")\n        return None\n    \n    print(f\"Creating submission ZIP with {len(tiff_files)} files...\")\n    \n    with zipfile.ZipFile(output_zip, 'w', zipfile.ZIP_DEFLATED) as zipf:\n        for tiff_path in tqdm(tiff_files, desc=\"Zipping predictions\"):\n            # Add file with just the filename (no directory structure)\n            zipf.write(tiff_path, tiff_path.name)\n            \n            if delete_after_zip:\n                tiff_path.unlink()\n    \n    zip_size_mb = output_zip.stat().st_size / (1024 * 1024)\n    print(f\"Submission saved: {output_zip} ({zip_size_mb:.1f} MB)\")\n    \n    return output_zip","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T02:58:37.829631Z","iopub.execute_input":"2026-01-23T02:58:37.829958Z","iopub.status.idle":"2026-01-23T02:58:37.83621Z","shell.execute_reply.started":"2026-01-23T02:58:37.829938Z","shell.execute_reply":"2026-01-23T02:58:37.835438Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"generate_submission()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T02:58:37.837363Z","iopub.execute_input":"2026-01-23T02:58:37.837897Z","iopub.status.idle":"2026-01-23T02:58:37.990783Z","shell.execute_reply.started":"2026-01-23T02:58:37.837875Z","shell.execute_reply":"2026-01-23T02:58:37.990103Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}