{"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},{"sourceType":"datasetVersion","sourceId":14913554,"datasetId":9439627,"databundleVersionId":15779479},{"sourceType":"datasetVersion","sourceId":13444908,"datasetId":8534135,"databundleVersionId":14160985},{"sourceType":"datasetVersion","sourceId":14766118,"datasetId":9381724,"databundleVersionId":15617777},{"sourceType":"datasetVersion","sourceId":14689677,"datasetId":9159789,"databundleVersionId":15533364},{"sourceType":"datasetVersion","sourceId":14924723,"datasetId":8897381,"databundleVersionId":15791716},{"sourceType":"datasetVersion","sourceId":13866980,"datasetId":8835017,"databundleVersionId":14630030},{"sourceType":"modelInstanceVersion","sourceId":766754,"databundleVersionId":15844636,"modelInstanceId":560051,"modelId":572638},{"sourceType":"modelInstanceVersion","sourceId":752197,"databundleVersionId":15707024,"modelInstanceId":560051,"modelId":572638},{"sourceType":"modelInstanceVersion","sourceId":755796,"databundleVersionId":15752385,"modelInstanceId":560051,"modelId":572638},{"sourceType":"modelInstanceVersion","sourceId":760513,"databundleVersionId":15799501,"modelInstanceId":560051,"modelId":572638},{"sourceType":"modelInstanceVersion","sourceId":758869,"databundleVersionId":15792096,"modelInstanceId":560051,"modelId":572638},{"sourceType":"modelInstanceVersion","sourceId":757999,"databundleVersionId":15782074,"modelInstanceId":560051,"modelId":572638},{"sourceType":"modelInstanceVersion","sourceId":743345,"databundleVersionId":15603642,"modelInstanceId":558962,"modelId":571547},{"sourceType":"modelInstanceVersion","sourceId":751724,"databundleVersionId":15702935,"modelInstanceId":558512,"modelId":571079},{"sourceType":"modelInstanceVersion","sourceId":742034,"databundleVersionId":15588971,"modelInstanceId":558512,"modelId":571079},{"sourceType":"modelInstanceVersion","sourceId":745216,"databundleVersionId":15625646,"modelInstanceId":558512,"modelId":571079},{"sourceType":"modelInstanceVersion","sourceId":740545,"databundleVersionId":15571225,"modelInstanceId":558512,"modelId":571079},{"sourceType":"kernelVersion","sourceId":280319414},{"sourceType":"kernelVersion","sourceId":280491629},{"sourceType":"kernelVersion","sourceId":296450067},{"sourceType":"kernelVersion","sourceId":296464035},{"sourceType":"kernelVersion","sourceId":297764448},{"sourceType":"kernelVersion","sourceId":298080599},{"sourceType":"kernelVersion","sourceId":298229995},{"sourceType":"kernelVersion","sourceId":298307380},{"sourceType":"kernelVersion","sourceId":298386559},{"sourceType":"kernelVersion","sourceId":299552984},{"sourceType":"kernelVersion","sourceId":300087604}],"dockerImageVersionId":31193,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\n\npkg_dir = \"/kaggle/input/rsna-2025-7th-place-packages\"\n\nsafe = [\n    \"acvl_utils-0.2.5-py3-none-any.whl\",\n    \"batchgenerators-0.25.1-py3-none-any.whl\",\n    \"batchgeneratorsv2-0.3.0-py3-none-any.whl\",\n    \"dynamic_network_architectures-0.4.2-py3-none-any.whl\",\n    \"connected_components_3d-3.26.0-cp311-cp311-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl\",\n    \"simpleitk-2.5.2-cp311-abi3-manylinux2014_x86_64.manylinux_2_17_x86_64.whl\",\n    \"scikit_image-0.25.2-cp311-cp311-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\",\n    \"einops-0.8.1-py3-none-any.whl\",\n    #\"pillow-12.0.0-cp311-cp311-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl\",\n    \"tqdm-4.67.1-py3-none-any.whl\",\n    'fft_conv_pytorch-1.2.0-py3-none-any.whl',\n    'imagecodecs-2025.8.2-cp311-cp311-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl'\n    \n]\n# !pip install --no-deps /kaggle/input/rsna-2025-7th-place-packages/nnunetv2-2.6.2-py3-none-any.whl\n\nfor w in safe:\n    print(\"Installing:\", w)\n    !pip install --no-deps {os.path.join(pkg_dir, w)}","metadata":{"_uuid":"53ded4e2-318d-4b83-8a12-7869d290dce3","_cell_guid":"e7541922-e25c-4c71-bfdd-9aeb34ce6fc1","trusted":true,"collapsed":false,"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2026-02-24T22:14:42.765639Z","iopub.execute_input":"2026-02-24T22:14:42.765843Z","iopub.status.idle":"2026-02-24T22:15:12.743583Z","shell.execute_reply.started":"2026-02-24T22:14:42.765824Z","shell.execute_reply":"2026-02-24T22:15:12.74259Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip uninstall -y tensorflow protobuf\n!pip install --no-deps /kaggle/input/wheels-for-vesuvius/monai-1.5.1-py3-none-any.whl\n!pip install --no-deps /kaggle/input/wheels-for-vesuvius/imagecodecs-2025.11.11-cp311-abi3-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl\n!pip install /kaggle/input/pip-install-hydra/hydra_core-1.3.2-py3-none-any.whl","metadata":{"_uuid":"d0ce61f6-60cd-400a-ac88-df4a6e21159e","_cell_guid":"104bb1c7-254f-418e-a5ba-2f88575ca7cc","trusted":true,"collapsed":false,"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2026-02-24T22:15:12.744796Z","iopub.execute_input":"2026-02-24T22:15:12.745116Z","iopub.status.idle":"2026-02-24T22:15:48.061769Z","shell.execute_reply.started":"2026-02-24T22:15:12.745084Z","shell.execute_reply":"2026-02-24T22:15:48.060749Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install tensorrt==10.12.0.36 \\\n    --no-index \\\n    --find-links /kaggle/input/tensorrt-install/trt_wheels","metadata":{"_uuid":"0131803f-480a-4403-850f-2fa8f5abbf1a","_cell_guid":"165b3816-6b50-4cfe-a400-6b2f94bf1395","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-02-24T22:15:48.063707Z","iopub.execute_input":"2026-02-24T22:15:48.064039Z","iopub.status.idle":"2026-02-24T22:22:06.027554Z","shell.execute_reply.started":"2026-02-24T22:15:48.063984Z","shell.execute_reply":"2026-02-24T22:22:06.02671Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport sys\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport tifffile as tiff\nfrom pathlib import Path\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom tqdm import tqdm\nfrom monai.networks.nets import UNet\nfrom monai.inferers import SlidingWindowInferer\nimport pytorch_lightning as pl\nfrom dataclasses import dataclass\nfrom typing import List, Tuple, Union\nimport warnings\nfrom monai.networks.blocks import UnetResBlock\nimport torch.nn as nn\nfrom dynamic_network_architectures.architectures.unet import ResidualEncoderUNet\nfrom dynamic_network_architectures.architectures.primus import PrimusB\nimport timm\nimport shutil\nfrom collections import OrderedDict, defaultdict\nimport inspect\nfrom copy import deepcopy\nimport multiprocessing\nfrom monai import transforms\nimport hydra\nfrom hydra import initialize, compose\nfrom omegaconf import OmegaConf\nfrom PIL import Image, ImageSequence\nfrom skimage.morphology import skeletonize\nfrom skimage.morphology import binary_dilation, square\nfrom scipy.ndimage import convolve\nimport yaml\nimport gc\nfrom concurrent.futures import ThreadPoolExecutor\nwarnings.filterwarnings('ignore')","metadata":{"_uuid":"6490fbb9-be2e-491b-93cb-ddfda91349de","_cell_guid":"6bf8a338-6a9c-44bd-aff7-66aa92b4693f","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-02-24T22:22:06.028564Z","iopub.execute_input":"2026-02-24T22:22:06.028769Z","iopub.status.idle":"2026-02-24T22:22:37.024149Z","shell.execute_reply.started":"2026-02-24T22:22:06.028745Z","shell.execute_reply":"2026-02-24T22:22:37.023292Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Deformnet Utilities","metadata":{"_uuid":"03b4bbcd-f9a2-44e5-8c73-cc2be0b54704","_cell_guid":"868bf866-6a7a-4a22-9ae8-68dd045daeb4","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"sys.path.append('/kaggle/input/vesuvius-challenge-resource-new')\nsys.path.append('/kaggle/input/vesuvius-challenge-resource-new/deformnetv2')\nfrom src.models.deformNet3d import *\nfrom monai.inferers.inferer import SlidingWindowInfererAdapt","metadata":{"_uuid":"046b0e65-cdaf-48fa-9d3b-fb5bab7578ac","_cell_guid":"79d04e9c-b788-43cf-bd47-3cb987664852","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-02-24T22:22:37.025058Z","iopub.execute_input":"2026-02-24T22:22:37.026229Z","iopub.status.idle":"2026-02-24T22:22:37.125474Z","shell.execute_reply.started":"2026-02-24T22:22:37.026205Z","shell.execute_reply":"2026-02-24T22:22:37.124642Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"deformnet_ckpt_paths_1 = [\n    ('/kaggle/input/vesuvius-challenge-resource-new/deformnetv2-refined-label-ckpt/deform-dynunet-v2-k3-s5-customFalse-fold0-epoch=19-val_bias_comp_metric=0.7012.ckpt', \"cuda:0\"),\n    ('/kaggle/input/vesuvius-challenge-resource-new/deformnetv2-refined-label-ckpt/deform-dynunet-v2-k3-s5-customFalse-fold1-epoch=34-val_bias_comp_metric=0.7526.ckpt', \"cuda:1\"),\n    ('/kaggle/input/vesuvius-challenge-resource-new/deformnetv2/models/deform-dynunet-v2-k3-s5-fold0-epoch=29-val_bias_comp_metric=0.6885.ckpt', \"cuda:0\"),\n    ('/kaggle/input/vesuvius-challenge-resource-new/deformnetv2/models/deform-dynunet-v2-k3-s5-fold1-epoch=54-val_bias_comp_metric=0.7348.ckpt', \"cuda:1\")\n]\n\n\nkernel_size0 = 3\nkernel_size1 = 3\nkernel_size2 = 3\nsigma0 = 10\nsigma1 = 5\nsigma2 = 10\ndeformnet_threshold = 0.5\noverlap = 0.5\nnum_iters = 1\nmin_obj_size=200\nclosing_radius = 1\nalpha_loop = 0.6\ndevice=\"cuda:0\"\n\n\ndef load_model_from_checkpoint(model, ckpt_path):\n    ckpt = torch.load(ckpt_path, weights_only=False)\n    state_dict = ckpt['state_dict']\n    new_state_dict = OrderedDict()\n    \n    for k, v in state_dict.items():\n        new_key = k.replace(\"model.\", \"\") if k.startswith(\"model.\") else k\n        new_state_dict[new_key] = v\n    model.load_state_dict(new_state_dict)\n\n\n\ndef load_average_model_soup(model, ckpt_paths):\n    avg_state = None\n    n = len(ckpt_paths)\n\n    for ckpt_path in ckpt_paths:\n        ckpt = torch.load(ckpt_path, weights_only=False)\n        state_dict = ckpt[\"state_dict\"]\n\n        clean_state = OrderedDict()\n        for k, v in state_dict.items():\n            k = k.replace(\"model.\", \"\") if k.startswith(\"model.\") else k\n            clean_state[k] = v.float() \n\n        if avg_state is None:\n            avg_state = OrderedDict({k: v.clone() for k, v in clean_state.items()})\n        else:\n            for k in avg_state:\n                avg_state[k] += clean_state[k]\n\n    # average\n    for k in avg_state:\n        avg_state[k] /= n\n\n    model.load_state_dict(avg_state, strict=True)\n\n\ndef load_volume(path: Path) -> np.ndarray:\n    \"\"\"\n    Load a multi-page TIFF into a 3D NumPy array: (slices, H, W)\n    \"\"\"\n    try:\n        with Image.open(path) as img:\n            frames = [np.array(frame) for frame in ImageSequence.Iterator(img)]\n        volume = np.stack(frames)\n        return volume\n    except Exception as e:\n        raise RuntimeError(f\"Error loading TIFF {path}: {e}\")\n\n\ndef generate_transforms(\n    transforms_config: list[dict],\n) -> list[transforms.Transform]:\n    transform_list = []\n    #logger.debug(f\"Generating {len(transforms_config)} transforms\")\n\n    for transform_config in transforms_config:\n        transform_name = next(iter(transform_config))\n        transform_kwargs = transform_config[transform_name]\n        # logger.debug(\n        #     f\"Generating transform {transform_name} with kwargs {transform_kwargs}\"\n        # )\n        transform: transforms.Transform = getattr(transforms, transform_name)(\n            **transform_kwargs\n        )  # type: ignore\n        transform_list.append(transform)\n    return transforms.Compose(transform_list)\n\n\ndef gaussian_kernel_3d(kernel_size=5, sigma=1.0, device=\"cuda\"):\n    \"\"\"Returns a normalized 3D Gaussian kernel (1,1,K,K,K).\"\"\"\n    ax = torch.arange(kernel_size, device=device) - kernel_size // 2\n    xx, yy, zz = torch.meshgrid(ax, ax, ax, indexing='ij')\n    kernel = torch.exp(-(xx**2 + yy**2 + zz**2) / (2 * sigma**2))\n    kernel = kernel / kernel.sum()\n    return kernel\n\n\ndef gaussian_blur_3d(x, kernel_size=7, sigma=10.0):\n    \"\"\"\n    x: (B, C, D, H, W)\n    \"\"\"\n    B, C, D, H, W = x.shape\n    kernel = gaussian_kernel_3d(kernel_size, sigma, device=CFG.DEVICE)\n\n    # shape: (C, 1, K, K, K)\n    kernel = kernel.expand(C, 1, kernel_size, kernel_size, kernel_size)\n\n    # depthwise convolution\n    return F.conv3d(x, kernel, padding=kernel_size // 2, groups=C)\n\n\ndef build_config(config_path, config_name, new_config_path, new_config_name):\n    SRC = Path(config_path)\n    DST = Path(new_config_path)\n    DST.mkdir(parents=True, exist_ok=True)\n    shutil.copytree(SRC, DST, dirs_exist_ok=True)\n    \n    with initialize(version_base=None, config_path=new_config_name):\n        cfg = compose(config_name=config_name)\n    return cfg\n\n\ncfg_0 = build_config(\"/kaggle/input/vesuvius-challenge-resource-new/deformnetv2/configs\", 'config_deform',\n             \"/kaggle/working/configs_0\", 'configs_0'\n            )\n\ncfg_1 = build_config(\"/kaggle/input/vesuvius-challenge-resource-new/deformnetv2/configs\", 'config_deform_v2',\n             \"/kaggle/working/configs_1\", 'configs_1'\n            )","metadata":{"_uuid":"eed4bbb3-7438-43bb-af13-e12ee730c391","_cell_guid":"1a60a532-ffd9-4479-9ed7-906d404981eb","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-02-24T22:22:37.126371Z","iopub.execute_input":"2026-02-24T22:22:37.126628Z","iopub.status.idle":"2026-02-24T22:22:40.243954Z","shell.execute_reply.started":"2026-02-24T22:22:37.126602Z","shell.execute_reply":"2026-02-24T22:22:40.243349Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.inference_mode()\ndef deformnet_step_worker(model, x, device, inferer):\n\n    x_device = x.to(device)\n    pred_warped = inferer(x_device, model)\n    pred_cpu = pred_warped.detach().cpu()\n\n    # Cleanup\n    del x_device\n    del pred_warped\n\n    return pred_cpu\n\n\n@torch.inference_mode()\ndef one_step_inference(prob_mask, vol, kernel_size, sigma,\n                       deformnet_list, is_prob, apply_gaussian, threshold,\n                       devices=[\"cuda:0\", \"cuda:1\"]):\n\n    if not is_prob:\n        prob_mask = (prob_mask > threshold).float()\n\n    if apply_gaussian:\n        prev_mask_pred = gaussian_blur_3d(prob_mask, kernel_size, sigma)\n    else:\n        prev_mask_pred = prob_mask\n\n    x = torch.cat([vol, prev_mask_pred], dim=1).cpu()\n\n    # Cleanup early\n    del prob_mask\n    del prev_mask_pred\n    del vol\n\n    torch.cuda.synchronize()\n    torch.cuda.empty_cache()\n\n    future_predictions = []\n\n    with ThreadPoolExecutor(max_workers=len(devices)) as executor:\n        for i, model in enumerate(deformnet_list):\n            target_device = devices[i % len(devices)]\n\n            future = executor.submit(\n                deformnet_step_worker,\n                model,\n                x,\n                target_device,\n                sliding_window_inferer_1\n            )\n            future_predictions.append(future)\n\n    predictions = [f.result() for f in future_predictions]\n\n    del future_predictions\n\n    result = torch.cat(predictions, dim=0).mean(dim=0)\n\n    del predictions\n\n    torch.cuda.synchronize()\n    torch.cuda.empty_cache()\n\n    return result","metadata":{"_uuid":"d64f64d8-6c76-4aa0-af66-9fc89dca0ee3","_cell_guid":"cc1a3e21-1d9e-4b65-93f2-fab3af9002bd","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-02-24T22:22:40.244592Z","iopub.execute_input":"2026-02-24T22:22:40.244805Z","iopub.status.idle":"2026-02-24T22:22:40.25183Z","shell.execute_reply.started":"2026-02-24T22:22:40.244786Z","shell.execute_reply":"2026-02-24T22:22:40.25092Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"alpha = 0.5\n@torch.no_grad()\ndef inference_deformnet(data_path, prob_mask):\n    print('start inference with deformnet')\n    vol = load_volume(data_path)\n    raw = {\"Image\": vol, \"Mask_OOF\": prob_mask}\n    _data = deformnet_transforms(raw)\n    vol, prob_mask = _data['Image'][None, ], _data['Mask_OOF'][None, ]\n    vol = vol.to(device)\n    prob_mask = prob_mask.to(device)\n    prob_mask = prob_mask.float()\n\n    #diffeomorphic\n    prev_mask = prob_mask\n\n    prediction_ensemble_v2 = one_step_inference(prev_mask, vol, kernel_size1,\n                                                  sigma1, deformnet_list_1, False, True, CFG.THRESHOLD\n                                                 ) #(c, d, h, w)\n    \n    #final\n    # segmentation_final = prediction_ensemble_v2 > deformnet_threshold\n    segmentation_final = prediction_ensemble_v2[0].cpu().numpy()\n    return segmentation_final #(D, H, W)","metadata":{"_uuid":"3fcbd8b7-50e4-4255-b917-3f0ce9183502","_cell_guid":"cd140d0c-08fa-4960-b038-3c89a336f2ff","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-02-24T22:22:40.252644Z","iopub.execute_input":"2026-02-24T22:22:40.252869Z","iopub.status.idle":"2026-02-24T22:22:40.268038Z","shell.execute_reply.started":"2026-02-24T22:22:40.252852Z","shell.execute_reply":"2026-02-24T22:22:40.267503Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## nnUNet","metadata":{"_uuid":"bdfadc1f-2d54-4e89-a659-bcdf871293ef","_cell_guid":"5f786d5f-18e1-4bd3-ad95-ac7ace9ac55c","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"\nimport torch.nn as nn\nfrom dynamic_network_architectures.architectures.unet import ResidualEncoderUNet\n\n\ndef create_residual_unet(\n    in_channels=1,\n    out_channels=2,\n    channels=(32, 64, 128, 256, 320, 320),\n    strides=(1, 2, 2, 2, 2, 2),\n    n_blocks_per_stage=(1, 3, 4, 6, 6, 6),\n    deep_supervision=False,\n):\n    \"\"\"Create ResidualEncoderUNet matching nnUNet 3d_fullres configuration.\n\n    Args:\n        in_channels: Input channels (default: 1)\n        out_channels: Output channels (default: 2)\n        channels: Feature channels at each level (tuple of ints)\n        strides: Strides for downsampling at each level (tuple of ints)\n        n_blocks_per_stage: Number of residual blocks per stage (tuple of ints)\n        deep_supervision: Whether to use deep supervision\n\n    Returns:\n        ResidualEncoderUNet model\n    \"\"\"\n    # Number of stages in decoder is len(channels) - 1\n    n_conv_per_stage_decoder = [1] * (len(channels) - 1)\n\n    model = ResidualEncoderUNet(\n        input_channels=in_channels,\n        n_stages=len(channels),\n        features_per_stage=channels,\n        conv_op=nn.Conv3d,\n        kernel_sizes=3,\n        strides=strides,\n        n_blocks_per_stage=n_blocks_per_stage,\n        num_classes=out_channels,\n        n_conv_per_stage_decoder=n_conv_per_stage_decoder,\n        conv_bias=True,\n        norm_op=nn.InstanceNorm3d,\n        norm_op_kwargs={},\n        dropout_op=None,\n        nonlin=nn.LeakyReLU,\n        nonlin_kwargs={'inplace': True},\n        deep_supervision=deep_supervision,\n    )\n    return model\n\n\ndef create_primus(\n        in_channels=1,\n        out_channels=2,\n        input_shape = 160,\n        patch_embed_size = 8\n):\n    model = PrimusB(in_channels,\n                    out_channels,\n                    (patch_embed_size, patch_embed_size, patch_embed_size),\n                    (input_shape, input_shape, input_shape))\n    return model","metadata":{"_uuid":"b931eefb-4bbd-4943-bc68-5f6abb1a6915","_cell_guid":"b92e6eb1-6c40-4564-840b-c87de3d4aa09","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-02-24T22:22:40.270414Z","iopub.execute_input":"2026-02-24T22:22:40.270641Z","iopub.status.idle":"2026-02-24T22:22:40.284631Z","shell.execute_reply.started":"2026-02-24T22:22:40.270624Z","shell.execute_reply":"2026-02-24T22:22:40.283856Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class SegmentationModule(pl.LightningModule):\n    def __init__(\n        self,\n        in_channels=2,\n        out_channels=2,\n        channels=(32, 64, 128, 256, 320, 320),\n        strides=(1, 2, 2, 2, 2, 2),\n        **kwargs  # Accept extra hparams from checkpoint\n    ):\n        super().__init__()\n        self.save_hyperparameters()\n\n\n        n_blocks = kwargs.get('n_blocks_per_stage', (1, 3, 4, 6, 6, 6))\n        use_ds = kwargs.get('use_deep_supervision', False)\n            \n        self.model = create_residual_unet(\n            in_channels=in_channels,\n            out_channels=out_channels,\n            channels=channels,\n            strides=strides,\n            n_blocks_per_stage=n_blocks,\n            deep_supervision=use_ds\n        )\n\n    def forward(self, x):\n        return self.model(x)\n\n\nclass SegmentationModulePrimus(pl.LightningModule):\n    def __init__(\n        self,\n        model_type='segformer',\n        in_channels=2,\n        out_channels=2,\n        channels=(32, 64, 128, 256, 320, 320),\n        strides=(1, 2, 2, 2, 2, 2),\n        **kwargs  # Accept extra hparams from checkpoint\n    ):\n        super().__init__()\n        self.save_hyperparameters()\n\n\n        n_blocks = kwargs.get('n_blocks_per_stage', (1, 3, 4, 6, 6, 6))\n        use_ds = kwargs.get('use_deep_supervision', False)\n            \n        self.model = create_primus()\n\n    def forward(self, x):\n        return self.model(x)","metadata":{"_uuid":"a36bfaad-0555-44e2-b54a-c8106e490923","_cell_guid":"91688e59-878c-40a9-8a52-7333ed6d2be9","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-02-24T22:22:40.285931Z","iopub.execute_input":"2026-02-24T22:22:40.286189Z","iopub.status.idle":"2026-02-24T22:22:40.299761Z","shell.execute_reply.started":"2026-02-24T22:22:40.286173Z","shell.execute_reply":"2026-02-24T22:22:40.299191Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class SegmentationModule2ndStage(pl.LightningModule):\n    def __init__(\n        self,\n        in_channels=2,\n        out_channels=2,\n        channels=(32, 64, 128, 256, 320, 320),\n        strides=(1, 2, 2, 2, 2, 2),\n        **kwargs  # Accept extra hparams from checkpoint\n    ):\n        super().__init__()\n        self.save_hyperparameters()\n\n\n        n_blocks = kwargs.get('n_blocks_per_stage', (1, 3, 4, 6, 6, 6))\n        use_ds = kwargs.get('use_deep_supervision', False)\n            \n        self.model = create_residual_unet(\n            in_channels=in_channels,\n            out_channels=out_channels,\n            channels=channels,\n            strides=strides,\n            n_blocks_per_stage=n_blocks,\n            deep_supervision=use_ds\n        )\n\n    def forward(self, x):\n        return self.model(x)\n\n\nclass SegmentationModule2ndStageSmall(pl.LightningModule):\n    def __init__(\n        self,\n        in_channels=2,\n        out_channels=2,\n        channels=(32, 64, 128, 256),\n        strides=(1, 2, 2, 2),\n        n_blocks_per_stage = (1, 3, 4, 6),\n        use_deep_supervision = False,\n        **kwargs  # Accept extra hparams from checkpoint\n    ):\n        super().__init__()\n        self.save_hyperparameters()\n\n\n        n_blocks = n_blocks_per_stage\n        use_ds = use_deep_supervision\n            \n        self.model = create_residual_unet(\n            in_channels=in_channels,\n            out_channels=out_channels,\n            channels=channels,\n            strides=strides,\n            n_blocks_per_stage=n_blocks,\n            deep_supervision=use_ds\n        )\n\n    def forward(self, x):\n        return self.model(x)\n\n\n\nclass SegmentationModule2ndStageV2(pl.LightningModule):\n    def __init__(\n        self,\n        in_channels=3,\n        out_channels=2,\n        channels=(32, 64, 128, 256, 320, 320),\n        strides=(1, 2, 2, 2, 2, 2),\n        **kwargs  # Accept extra hparams from checkpoint\n    ):\n        super().__init__()\n        self.save_hyperparameters()\n\n\n        n_blocks = kwargs.get('n_blocks_per_stage', (1, 3, 4, 6, 6, 6))\n        use_ds = kwargs.get('use_deep_supervision', False)\n            \n        self.model = create_residual_unet(\n            in_channels=in_channels,\n            out_channels=out_channels,\n            channels=channels,\n            strides=strides,\n            n_blocks_per_stage=n_blocks,\n            deep_supervision=use_ds\n        )\n\n    def forward(self, x):\n        return self.model(x)","metadata":{"_uuid":"f30c3b10-e728-4ed1-bc80-d72a0a19c754","_cell_guid":"e17d71c7-08b9-49b8-b7f6-d85f6d22eb38","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-02-24T22:22:40.300406Z","iopub.execute_input":"2026-02-24T22:22:40.300637Z","iopub.status.idle":"2026-02-24T22:22:40.315537Z","shell.execute_reply.started":"2026-02-24T22:22:40.30062Z","shell.execute_reply":"2026-02-24T22:22:40.31485Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Configuration","metadata":{"_uuid":"4c5e096e-98c3-487f-a51d-2e2abb528b88","_cell_guid":"eed3baee-9ea1-4b44-b6bb-f9cbdd66dbd3","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"class CFG:\n    # Data directories\n    TEST_IMG_DIR = Path(\"/kaggle/input/vesuvius-challenge-surface-detection/test_images\")\n    MODEL_DIR = Path(\"\")\n    \n    # Inference settings - sliding window\n    ROI_SIZE = (160, 160, 160)       # Sliding window ROI size\n    SW_BATCH_SIZE = 1                 # Batch size for sliding window\n    OVERLAP = 0.5                    # Overlap ratio for sliding window\n    SW_MODE = \"gaussian\"              # Mode: \"constant\", \"gaussian\"\n    PADDING_MODE = \"reflect\"          # Padding mode\n    REPEAT = 2 #repeat time of refinement\n    REPEAT_V2 = 1\n    \n    # TTA settings\n    USE_TTA = True                  # Enable Test-Time Augmentation\n    TTA_FLIPS = True                  # Use flip augmentations\n    TTA_ROTATIONS = True             # Use 90-degree rotations (warning: slower and memory intensive)\n    \n    # Model ensemble settings\n    THRESHOLD = 0.3\n    THRESHOLD_1ST_STAGE = 0.3\n    THRESHOLD_2ND_STAGE = 0.3\n    THRESHOLD_2ND_V2_STAGE = 0.3\n    \n    # Output settings\n    OUTPUT_DIR = Path(\"./predictions\")\n    SAVE_VISUALIZATIONS = True  # Set to True to save visualization images\n    \n    # Post-processing settings\n    USE_POST_PROCESSING = True  \n    POST_PROCESS_MIN_CC_VOLUME = 3000  \n    \n    # Device\n    DEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n# Create output directories\nCFG.OUTPUT_DIR.mkdir(exist_ok=True, parents=True)\n(CFG.OUTPUT_DIR / \"submission_tifs\").mkdir(exist_ok=True, parents=True)\n\nprint(f\"Device: {CFG.DEVICE}\")\nprint(f\"Threshold: {CFG.THRESHOLD}\")\nprint(f\"ROI Size: {CFG.ROI_SIZE}\")\nprint(f\"Overlap: {CFG.OVERLAP}\")\nprint(f\"Mode: {CFG.SW_MODE}\")\nprint(f\"TTA Enabled: {CFG.USE_TTA}\")\nif CFG.USE_TTA:\n    print(f\"  - Flips: {CFG.TTA_FLIPS}\")\n    print(f\"  - Rotations: {CFG.TTA_ROTATIONS}\")\nprint(f\"Post-processing: {CFG.USE_POST_PROCESSING}\")\nif CFG.USE_POST_PROCESSING:\n    print(f\"  - Min CC Volume: {CFG.POST_PROCESS_MIN_CC_VOLUME}\")\nprint(f\"Visualizations: {CFG.SAVE_VISUALIZATIONS}\")","metadata":{"_uuid":"3ea5acc9-25b8-4a1e-8675-41dccf37e4c1","_cell_guid":"7cabefda-87f9-4544-ae5b-4fa70e026057","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-02-24T22:22:40.316253Z","iopub.execute_input":"2026-02-24T22:22:40.316432Z","iopub.status.idle":"2026-02-24T22:22:40.329247Z","shell.execute_reply.started":"2026-02-24T22:22:40.316418Z","shell.execute_reply":"2026-02-24T22:22:40.328488Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Utility Functions","metadata":{"_uuid":"39c37e32-9b36-46ff-bdf9-cf46cdcb8d8d","_cell_guid":"bf7aff01-3565-4079-b1c2-724f3a02c7e5","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"def load_array(path, fmt):\n    \"\"\"Load array from various formats\"\"\"\n    if fmt == \"tiff\" or fmt == \"tif\":\n        return tiff.imread(path)\n    elif fmt == \"npy\":\n        return np.load(path)\n    elif fmt == \"npz\":\n        return np.load(path)[\"arr_0\"]\n    elif fmt == \"rle\":\n        rle = np.load(path)\n        shape = tuple(rle[\"shape\"])\n        vals = rle[\"vals\"]\n        runs = rle[\"runs\"]\n        flat = np.repeat(vals, runs)\n        return flat.reshape(shape)\n    else:\n        raise ValueError(f\"Unsupported format: {fmt}\")\n\n\ndef normalize_volume(volume):\n    volume = volume.astype(np.float32)\n\n    volume = volume / 255.0\n    \n    return volume\n\n\ndef rle_encode(mask):\n    \"\"\"Run-length encoding for submission\"\"\"\n    pixels = mask.flatten()\n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)","metadata":{"_uuid":"ba017169-2d6b-4791-b742-8966cd0d4e7c","_cell_guid":"46233bb9-1aeb-4079-b161-84e496238747","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-02-24T22:22:40.329893Z","iopub.execute_input":"2026-02-24T22:22:40.330128Z","iopub.status.idle":"2026-02-24T22:22:40.344945Z","shell.execute_reply.started":"2026-02-24T22:22:40.330102Z","shell.execute_reply":"2026-02-24T22:22:40.344321Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Test-Time Augmentation Functions","metadata":{"_uuid":"2b685715-e2e0-4440-a191-c295bcc8fcd7","_cell_guid":"0c7cbad9-ae43-4e4e-839b-6b416b32e330","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"def apply_tta_transform(volume, flip_dims=None, rotation_k=0):\n    \"\"\"\n    Apply TTA transformation to volume.\n    \n    Args:\n        volume: Input volume tensor (1, 1, D, H, W)\n        flip_dims: List of dimensions to flip (e.g., [2, 3, 4] for D, H, W)\n        rotation_k: Number of 90-degree rotations (0-3) on H-W plane\n        \n    Returns:\n        Transformed volume\n    \"\"\"\n    transformed = volume.clone()\n    \n    # Apply flips\n    if flip_dims:\n        for dim in flip_dims:\n            transformed = torch.flip(transformed, dims=[dim])\n    \n    # Apply rotation on H-W plane (dims 3 and 4)\n    if rotation_k > 0:\n        transformed = torch.rot90(transformed, k=rotation_k, dims=[3, 4])\n    \n    return transformed\n\n\ndef reverse_tta_transform(prediction, flip_dims=None, rotation_k=0):\n    \"\"\"\n    Reverse TTA transformation on prediction.\n    \n    Args:\n        prediction: Prediction tensor (D, H, W) or (1, D, H, W)\n        flip_dims: List of dimensions to flip (adjusted for prediction dims)\n        rotation_k: Number of 90-degree rotations to reverse\n        \n    Returns:\n        Reversed prediction\n    \"\"\"\n    # Handle both (D, H, W) and (1, D, H, W) shapes\n    if prediction.ndim == 3:\n        pred = torch.from_numpy(prediction).unsqueeze(0)  # (1, D, H, W)\n        squeeze_output = True\n    else:\n        pred = torch.from_numpy(prediction)\n        squeeze_output = False\n    \n    # Reverse rotation first (opposite direction)\n    if rotation_k > 0:\n        pred = torch.rot90(pred, k=(4 - rotation_k), dims=[2, 3])\n    \n    # Reverse flips\n    if flip_dims:\n        # Adjust flip dims for (1, D, H, W) shape\n        adjusted_dims = [d - 1 for d in flip_dims] if flip_dims else None\n        if adjusted_dims:\n            for dim in adjusted_dims:\n                pred = torch.flip(pred, dims=[dim])\n    \n    result = pred.squeeze(0).numpy() if squeeze_output else pred.numpy()\n    return result\n\n\ndef get_tta_transforms():\n    \"\"\"\n    Get list of TTA transformations to apply.\n    \n    Returns:\n        List of (flip_dims, rotation_k) tuples\n    \"\"\"\n    transforms = [\n        (None, 0),  # Original (no transform)\n    ]\n    \n    if CFG.TTA_FLIPS:\n        transforms.extend([\n            ([2], 0),      # Flip D\n            ([3], 0),      # Flip H\n            ([4], 0),      # Flip W\n            #([2, 3], 0),   # Flip D+H\n            #([2, 4], 0),   # Flip D+W\n            #([3, 4], 0),   # Flip H+W\n            #([2, 3, 4], 0) # Flip all\n        ])\n    \n    if CFG.TTA_ROTATIONS:\n        # Add rotation augmentations (90, 180, 270 degrees)\n        transforms.extend([\n            (None, 1),  # 90° rotation\n            #(None, 2),  # 180° rotation\n            #(None, 3),  # 270° rotation\n        ])\n    \n    return transforms\n\n\ndef get_tta_transforms_primus():\n    \"\"\"\n    Get list of TTA transformations to apply.\n    \n    Returns:\n        List of (flip_dims, rotation_k) tuples\n    \"\"\"\n    transforms = [\n        (None, 0),  # Original (no transform)\n    ]\n    \n    if CFG.TTA_FLIPS:\n        transforms.extend([\n            # ([2], 0),      # Flip D\n            # ([3], 0),      # Flip H\n            # ([4], 0),      # Flip W\n            # ([2, 3], 0),   # Flip D+H\n            # ([2, 4], 0),   # Flip D+W\n            #([3, 4], 0),   # Flip H+W\n            #([2, 3, 4], 0) # Flip all\n        ])\n    \n    # if CFG.TTA_ROTATIONS:\n    #     # Add rotation augmentations (90, 180, 270 degrees)\n    #     transforms.extend([\n    #         (None, 1),  # 90° rotation\n    #         #(None, 2),  # 180° rotation\n    #         #(None, 3),  # 270° rotation\n    #     ])\n    \n    return transforms\n\n\ndef get_tta_transforms_2nd():\n    \"\"\"\n    Get list of TTA transformations to apply.\n    \n    Returns:\n        List of (flip_dims, rotation_k) tuples\n    \"\"\"\n    transforms = [\n        (None, 0),  # Original (no transform)\n    ]\n    \n    # if CFG.TTA_FLIPS:\n    #     transforms.extend([\n    #         ([2], 0),      # Flip D\n    #         ([3], 0),      # Flip H\n    #         ([4], 0),      # Flip W\n    #         ([2, 3], 0),   # Flip D+H\n    #         ([2, 4], 0),   # Flip D+W\n    #         #([3, 4], 0),   # Flip H+W\n    #         #([2, 3, 4], 0) # Flip all\n    #     ])\n    \n    # if CFG.TTA_ROTATIONS:\n    #     # Add rotation augmentations (90, 180, 270 degrees)\n    #     transforms.extend([\n    #         (None, 1),  # 90° rotation\n    #         #(None, 2),  # 180° rotation\n    #         #(None, 3),  # 270° rotation\n    #     ])\n    \n    return transforms\n\n\ndef get_tta_transforms_2nd_v2():\n    \"\"\"\n    Get list of TTA transformations to apply.\n    \n    Returns:\n        List of (flip_dims, rotation_k) tuples\n    \"\"\"\n    transforms = [\n        (None, 0),  # Original (no transform)\n    ]\n    \n    # if CFG.TTA_FLIPS:\n    #     transforms.extend([\n    #         ([2], 0),      # Flip D\n    #         ([3], 0),      # Flip H\n    #         ([4], 0),      # Flip W\n    #         ([2, 3], 0),   # Flip D+H\n    #         ([2, 4], 0),   # Flip D+W\n    #         ([3, 4], 0),   # Flip H+W\n    #         ([2, 3, 4], 0) # Flip all\n    #     ])\n    \n    # if CFG.TTA_ROTATIONS:\n    #     # Add rotation augmentations (90, 180, 270 degrees)\n    #     transforms.extend([\n    #         (None, 1),  # 90° rotation\n    #         # (None, 2),  # 180° rotation\n    #         # (None, 3),  # 270° rotation\n    #     ])\n    \n    return transforms\n\n\ndef get_tta_transforms_3nd():\n    \"\"\"\n    Get list of TTA transformations to apply.\n    \n    Returns:\n        List of (flip_dims, rotation_k) tuples\n    \"\"\"\n    transforms = [\n        (None, 0),  # Original (no transform)\n    ]\n    \n    # if CFG.TTA_FLIPS:\n    #     transforms.extend([\n    #         ([2], 0),      # Flip D\n    #         ([3], 0),      # Flip H\n    #         ([4], 0),      # Flip W\n    #         ([2, 3], 0),   # Flip D+H\n    #         ([2, 4], 0),   # Flip D+W\n    #         #([3, 4], 0),   # Flip H+W\n    #         #([2, 3, 4], 0) # Flip all\n    #     ])\n    \n    # if CFG.TTA_ROTATIONS:\n    #     # Add rotation augmentations (90, 180, 270 degrees)\n    #     transforms.extend([\n    #         (None, 1),  # 90° rotation\n    #         #(None, 2),  # 180° rotation\n    #         #(None, 3),  # 270° rotation\n    #     ])\n    \n    return transforms\n\n\ndef get_tta_transforms_3nd_refine():\n    \"\"\"\n    Get list of TTA transformations to apply.\n    \n    Returns:\n        List of (flip_dims, rotation_k) tuples\n    \"\"\"\n    transforms = [\n        (None, 0),  # Original (no transform)\n    ]\n    \n    # if CFG.TTA_FLIPS:\n    #     transforms.extend([\n    #         ([2], 0),      # Flip D\n    #         ([3], 0),      # Flip H\n    #         ([4], 0),      # Flip W\n    #         ([2, 3], 0),   # Flip D+H\n    #         ([2, 4], 0),   # Flip D+W\n    #         #([3, 4], 0),   # Flip H+W\n    #         #([2, 3, 4], 0) # Flip all\n         # ])\n    \n    #if CFG.TTA_ROTATIONS:\n    #    # Add rotation augmentations (90, 180, 270 degrees)\n    #    transforms.extend([\n    #        (None, 1),  # 90° rotation\n    #        #(None, 2),  # 180° rotation\n    #        #(None, 3),  # 270° rotation\n    #    ])\n    \n    return transforms","metadata":{"_uuid":"95e53930-5f3a-4ed8-987d-c0382b9793bf","_cell_guid":"a8271bb7-f92b-4173-b84c-e7658e3aaadf","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-02-24T22:22:40.34575Z","iopub.execute_input":"2026-02-24T22:22:40.346075Z","iopub.status.idle":"2026-02-24T22:22:40.362591Z","shell.execute_reply.started":"2026-02-24T22:22:40.346059Z","shell.execute_reply":"2026-02-24T22:22:40.361852Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import tensorrt as trt","metadata":{"_uuid":"a9273dfd-7f6c-4cf0-ab36-9051ad744b0c","_cell_guid":"89ae66c2-8b8b-4f08-a0d8-6e9afa5fa524","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-02-24T22:22:40.363355Z","iopub.execute_input":"2026-02-24T22:22:40.363551Z","iopub.status.idle":"2026-02-24T22:22:40.377827Z","shell.execute_reply.started":"2026-02-24T22:22:40.363529Z","shell.execute_reply":"2026-02-24T22:22:40.377238Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class TRTWrapper:\n    def __init__(self, engine_path, device_id=0):\n        self.device_id = device_id\n        self.device = torch.device(f\"cuda:{self.device_id}\")\n        \n        with torch.cuda.device(self.device):\n            # Initialize CUDA context explicitly before TRT\n            torch.cuda.init()\n            _ = torch.zeros(1, device=self.device)  # Force CUDA context creation\n            \n            self.logger = trt.Logger(trt.Logger.WARNING)\n            self.runtime = trt.Runtime(self.logger)\n            \n            print(f\"Loading engine from {engine_path} on GPU {self.device_id}...\")\n            with open(engine_path, \"rb\") as f:\n                self.engine = self.runtime.deserialize_cuda_engine(f.read())\n            \n            if self.engine is None:\n                raise RuntimeError(f\"Failed to deserialize engine: {engine_path}\")\n            \n            self.context = self.engine.create_execution_context()\n            \n            if self.context is None:\n                raise RuntimeError(f\"Failed to create execution context for: {engine_path}\")\n            \n            \n            # Inspect IO (Same as before)\n            self.inputs = []\n            self.outputs = []\n            for i in range(self.engine.num_io_tensors):\n                name = self.engine.get_tensor_name(i)\n                mode = self.engine.get_tensor_mode(name)\n                shape = self.engine.get_tensor_shape(name)\n                dtype = self.engine.get_tensor_dtype(name)\n                info = {'name': name, 'shape': shape, 'dtype': dtype}\n                \n                if mode == trt.TensorIOMode.INPUT:\n                    self.inputs.append(info)\n                elif mode == trt.TensorIOMode.OUTPUT:\n                    self.outputs.append(info)\n\n    def eval(self):\n        return\n    def __call__(self, input_tensor):\n        with torch.cuda.device(self.device):\n    \n            # 3. MOVE DATA TO CORRECT GPU\n            if input_tensor.device.type != 'cuda' or input_tensor.device.index != self.device_id:\n                input_tensor = input_tensor.to(self.device)\n            \n            input_tensor = input_tensor.contiguous()\n            \n            # Bind Input\n            self.context.set_tensor_address(self.inputs[0]['name'], input_tensor.data_ptr())\n    \n            # Allocate Outputs on CORRECT GPU\n            output_tensors = []\n            for out_info in self.outputs:\n                name = out_info['name']\n                shape = tuple(out_info['shape'])\n                \n                # Map TRT dtype to Torch dtype\n                dtype_map = {\n                    trt.DataType.FLOAT: torch.float32,\n                    trt.DataType.HALF: torch.float16,\n                    trt.DataType.INT32: torch.int32,\n                    trt.DataType.INT8: torch.int8,\n                    trt.DataType.BOOL: torch.bool,\n                }\n                torch_dtype = dtype_map.get(out_info['dtype'], torch.float32)\n    \n                # Create tensor on specific device\n                out_tensor = torch.empty(shape, device=self.device, dtype=torch_dtype)\n                output_tensors.append(out_tensor)\n                \n                self.context.set_tensor_address(name, out_tensor.data_ptr())\n    \n            # Execute\n            # Use the stream for the specific device\n            stream = torch.cuda.current_stream(device=self.device).cuda_stream\n            self.context.execute_async_v3(stream_handle=stream)\n            \n            # Sync just this device\n            torch.cuda.current_stream(device=self.device).synchronize()\n            \n            if len(output_tensors) == 1:\n                return output_tensors[0]\n            return output_tensors","metadata":{"_uuid":"fedb7ddc-0397-415b-9aee-6ab2ed4d12fd","_cell_guid":"6cb077cc-ca12-41d5-a03e-d71a0f1a1b56","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-02-24T22:22:40.378504Z","iopub.execute_input":"2026-02-24T22:22:40.37912Z","iopub.status.idle":"2026-02-24T22:22:40.394483Z","shell.execute_reply.started":"2026-02-24T22:22:40.379102Z","shell.execute_reply":"2026-02-24T22:22:40.393874Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model Loading","metadata":{"_uuid":"512f64c8-3720-45c0-8048-9d5afbd06b6e","_cell_guid":"ca9a8772-9d7f-4e02-afa3-739d837cbb26","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"# ==============================================================================\n# 1st Stage models\n# ==============================================================================\n\n\nMODEL_PATHS = [\n    (\"/kaggle/input/unet3d-1st-stage-vesuvius-challenge/pytorch/default/14/trt-models-1st-stage-resenc/best-epoch369-val_loss0.3720-val_dice0.5789.engine\", 1, TRTWrapper, \"resenc\", get_tta_transforms),\n    (\"/kaggle/input/unet3d-1st-stage-vesuvius-challenge/pytorch/default/14/trt-models-1st-stage-resenc/best-epoch364-val_loss0.3767-val_dice0.5813.engine\", 0, TRTWrapper, \"resenc\", get_tta_transforms),\n    (\"/kaggle/input/notebooks/iamparadox/tensorrt-1st-stage-all-models/1st_stage_models/fold-2-best-epoch384-val_loss0.3647-val_dice0.5865.engine\", 0, TRTWrapper, \"resenc\", get_tta_transforms),\n    (\"/kaggle/input/notebooks/iamparadox/tensorrt-1st-stage-all-models/1st_stage_models/fold-3-best-epoch349-val_loss0.3471-val_dice0.6062.engine\", 1, TRTWrapper, \"resenc\", get_tta_transforms),\n##########\n    (\"/kaggle/input/tensorrt-models-primus/primus-best-epoch=264-val_loss=0.3657-val_dice=0.5852.engine\", 1, TRTWrapper, \"primus\", get_tta_transforms_primus),\n    (\"/kaggle/input/tensorrt-models-primus/primus-best-epoch=294-val_loss=0.3653-val_dice=0.5917.engine\", 0, TRTWrapper, \"primus\", get_tta_transforms_primus),\n    (\"/kaggle/input/notebooks/iamparadox/tensorrt-1st-stage-all-models/primusV1/fold2_best-epoch359-val_loss0.3567-val_dice0.5960.engine\", 1, TRTWrapper, \"primus\", get_tta_transforms_primus),\n    (\"/kaggle/input/notebooks/iamparadox/tensorrt-1st-stage-all-models/primusV1/fold3_best-epoch344-val_loss0.3405-val_dice0.6139.engine\", 0, TRTWrapper, \"primus\", get_tta_transforms_primus),\n##############\n    (\"/kaggle/input/notebooks/tom99763/tensorrt-models-primus-v2/primus-primus-v2-best-epoch=404-val_loss=0.3655-val_dice=0.5881.engine\", 1, TRTWrapper, \"primus_v2\", get_tta_transforms_primus),\n    (\"/kaggle/input/notebooks/tom99763/tensorrt-models-primus-v2/primus-primus-v2-best-epoch=379-val_loss=0.3646-val_dice=0.5936.engine\", 0, TRTWrapper, \"primus_v2\", get_tta_transforms_primus),\n    (\"/kaggle/input/notebooks/iamparadox/tensorrt-1st-stage-all-models/primus_v2/fold2_best-epoch389-val_loss0.3553-val_dice0.5971.engine\", 1, TRTWrapper, \"primus_v2\", get_tta_transforms_primus),\n    (\"/kaggle/input/notebooks/iamparadox/tensorrt-1st-stage-all-models/primus_v2/fold3-best-epoch409-val_loss0.3412-val_dice0.6135.engine\", 0, TRTWrapper, \"primus_v2\", get_tta_transforms_primus),    \n]\n\ndef load_models_simple(model_paths):\n    \"\"\"Load models using Lightning's built-in checkpoint loading, handling EMA weights.\n\n    Accepts tuples of either:\n      (path, device, module_class)                       — legacy 3-tuple\n      (path, device, module_class, group_name, tta_fn)   — extended 5-tuple\n\n    Returns: list of dicts with keys: model, group_name, tta_fn\n    \"\"\"\n    entries = []\n    for item in model_paths:\n        if len(item) == 5:\n            path, device, module, group_name, tta_fn = item\n        else:\n            path, device, module = item[:3]\n            group_name, tta_fn = None, None\n        is_trt = (module == TRTWrapper)\n        \n        if isinstance(device, int) and not is_trt:\n            device = f\"cuda:{device}\"\n        print(f\"Loading: {path}\" + (f\"  [group={group_name}]\" if group_name else \"\"))\n\n        if module != TRTWrapper:\n\n            # First, load the checkpoint to check for EMA weights \n            checkpoint = torch.load(path, map_location=device)\n            state_dict = checkpoint.get('state_dict', checkpoint)\n\n            # Check if checkpoint has EMA weights\n            has_ema = any(k.startswith('ema.module.') for k in state_dict.keys())\n            use_ema_hparam = checkpoint.get('hyper_parameters', {}).get('use_ema', False)\n\n            # Create model from checkpoint\n            model = module.load_from_checkpoint(\n                path,\n                map_location=device,\n                pretrained_backbone=False,  # ensure no downloads\n                strict=False\n            )\n\n            # If EMA weights exist, extract and load them into model\n            if has_ema:\n                print(f\"  Found EMA weights in checkpoint, extracting...\")\n                ema_state_dict = {}\n                for k, v in state_dict.items():\n                    if k.startswith('ema.module.'):\n                        # Remove 'ema.module.' prefix to get 'model.*' key\n                        new_key = 'model.' + k[len('ema.module.'):]\n                        ema_state_dict[new_key] = v\n\n                # Load EMA weights into the model\n                missing, unexpected = model.load_state_dict(ema_state_dict, strict=False)\n                if missing:\n                    # Filter out non-model keys from missing (like loss functions, etc.)\n                    model_missing = [k for k in missing if k.startswith('model.')]\n                    if model_missing:\n                        print(f\"  Warning: Missing EMA keys: {model_missing[:5]}...\")\n                print(f\"  ✓ Loaded EMA weights into model\")\n\n            model.to(device)\n            model.eval()\n            model_type = getattr(model.hparams, 'model_type', 'unet')\n        else:\n            has_ema = False\n            use_ema_hparam = False\n            model = TRTWrapper(path, device)\n            model_type = \"tensorrt\"\n\n        # Print model info\n        print(f\"  Model type: {model_type}\")\n        if model_type == 'flex_unet':\n            backbone = getattr(model.hparams, 'backbone', 'resnet18')\n            print(f\"  Backbone: {backbone}\")\n        if has_ema or use_ema_hparam:\n            print(f\"  ✓ Using EMA weights for inference\")\n        else:\n            print(f\"  Standard weights (no EMA)\")\n        print(f\"  ✓ Loaded successfully\")\n\n        entries.append({\"model\": model, \"group_name\": group_name, \"tta_fn\": tta_fn})\n    print(f\"\\n✓ Total models loaded: {len(entries)}\")\n    return entries\n\n\ndef get_model_groups(entries):\n    \"\"\"Group loaded model entries by group_name, preserving insertion order.\n\n    Returns: OrderedDict of group_name -> {models: [...], tta_fns: [...], weights: [...]}\n    \"\"\"\n    groups = OrderedDict()\n    for e in entries:\n        name = e[\"group_name\"]\n        if name not in groups:\n            groups[name] = {\"models\": [], \"tta_fns\": [], \"weights\": []}\n        groups[name][\"models\"].append(e[\"model\"])\n        groups[name][\"tta_fns\"].append(e[\"tta_fn\"])\n    # Equal weights within each group\n    for g in groups.values():\n        n = len(g[\"models\"])\n        g[\"weights\"] = [1.0 / n] * n\n    return groups\n\n\n# Group paths without loading them into VRAM yet (avoids OOM from loading all 12 TRT models at once)\ngrouped_model_paths = defaultdict(list)\nfor item in MODEL_PATHS:\n    group_name = item[3]  # Extracting group_name from 5-tuple (path, device, module, group_name, tta_fn)\n    grouped_model_paths[group_name].append(item)\n\ngroup_names_1st = list(grouped_model_paths.keys())\nprint(f\"1st stage groups scheduled: {group_names_1st} (channels: image + {len(group_names_1st)} OOF)\")","metadata":{"_uuid":"f05c5ccd-7f8e-4cab-a3bc-bcf055fe4c1d","_cell_guid":"7b2a8ed2-44f3-4ac4-bd66-f1b33133d917","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-02-24T22:22:40.395174Z","iopub.execute_input":"2026-02-24T22:22:40.39558Z","iopub.status.idle":"2026-02-24T22:22:40.416612Z","shell.execute_reply.started":"2026-02-24T22:22:40.395555Z","shell.execute_reply":"2026-02-24T22:22:40.415905Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Inference Dataset","metadata":{"_uuid":"eacf93e5-7b30-4daf-b862-0a451123020b","_cell_guid":"3d3dfd79-7148-4635-ad49-d6c7d1fe7a46","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"class InferenceDataset(Dataset):\n    \"\"\"Dataset for inference on test volumes - no resizing\"\"\"\n    def __init__(self, image_paths):\n        self.image_paths = image_paths\n\n    def __len__(self):\n        return len(self.image_paths)\n\n    def __getitem__(self, idx):\n        img_path = self.image_paths[idx]\n        \n        # Load original volume\n        vol = load_array(img_path, \"tif\")\n        vol_shape = vol.shape  # Get actual input shape\n        \n        # Normalize volume (no resizing)\n        vol_normalized = normalize_volume(vol)\n        \n        # Convert to tensor (add channel dimension)\n        vol_tensor = torch.from_numpy(vol_normalized).float().unsqueeze(0)  # (1, D, H, W)\n        \n        return {\n            'volume': vol_tensor,\n            'shape': vol_shape,\n            'filename': img_path.name\n        }\n\n\ndef custom_collate_fn(batch):\n    \"\"\"Custom collate function\"\"\"\n    item = batch[0]\n    return {\n        'volume': item['volume'].unsqueeze(0),  # Add batch dimension\n        'shape': torch.tensor([item['shape']]),\n        'filename': [item['filename']]\n    }","metadata":{"_uuid":"59779fd3-d467-4f79-a15e-feb7d08429c7","_cell_guid":"92d183c8-01ad-4837-b2bc-772e7642abc9","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-02-24T22:22:40.417412Z","iopub.execute_input":"2026-02-24T22:22:40.417691Z","iopub.status.idle":"2026-02-24T22:22:40.434015Z","shell.execute_reply.started":"2026-02-24T22:22:40.417665Z","shell.execute_reply":"2026-02-24T22:22:40.433274Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Inference Function with Sliding Window and TTA","metadata":{"_uuid":"94efef83-8720-42fa-820b-d5c3284b6fcb","_cell_guid":"f820fa18-346d-45bb-9022-c8426d6d3b5d","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"import torch\nimport numpy as np\nfrom concurrent.futures import ThreadPoolExecutor\n\ndef _get_model_device(model):\n    \"\"\"Get the device a model lives on (TRTWrapper or Lightning/PyTorch).\"\"\"\n    if hasattr(model, 'device'):\n        return model.device\n    return next(model.parameters()).device\n\n\n@torch.no_grad()\ndef _predict_single_model_worker(model, volume_tensor, target_device, inferer, tta_transforms):\n\n    torch.cuda.set_device(target_device)\n    model.eval()\n\n    tta_predictions = []\n\n    for flip_dims, rotation_k in tta_transforms:\n\n        transformed_volume = apply_tta_transform(volume_tensor, flip_dims, rotation_k)\n        transformed_volume = transformed_volume.to(target_device, non_blocking=True)\n\n        with torch.amp.autocast(device_type='cuda'):\n            logits = inferer(transformed_volume, model)\n\n            \n        if isinstance(logits, (list, tuple)):\n            logits = logits[0]\n\n        if logits.shape[1]!=1:\n            probs = torch.softmax(logits, dim=1)[:, 1]\n        else:\n            probs = logits[:, 0]\n        probs = probs.squeeze(0).cpu().numpy() # Move result back to CPU immediately\n        \n        probs_reversed = reverse_tta_transform(probs, flip_dims, rotation_k)\n        tta_predictions.append(probs_reversed)\n        \n        del transformed_volume, logits, probs\n        # Optional: Empty cache only if strictly necessary to avoid synchronization overhead\n        # torch.cuda.empty_cache() \n\n    # Average TTA predictions for this specific model\n    return np.mean(tta_predictions, axis=0)\n\n@torch.no_grad()\ndef predict_volume_sliding_window(\n    models,\n    volume_tensor,\n    weights,\n    devices,\n    tta_fns,\n    use_tta=False\n):\n    \"\"\"\n    Args:\n        models:   List of PyTorch models\n        volume_tensor: torch.Tensor (C, D, H, W) or similar\n        weights:  List of floats, one per model\n        devices:  List of devices, e.g. ['cuda:0', 'cuda:1']\n        use_tta:  Whether to use test-time augmentation\n    Returns:\n        ensemble_pred: np.ndarray\n    \"\"\"\n\n    assert len(models) == len(weights), \"Number of models must match weights\"\n\n    # ------------------------------------------------------------------\n    # Sliding window inferer (stateless, can be shared)\n    # ------------------------------------------------------------------\n    inferer = SlidingWindowInferer(\n        roi_size=CFG.ROI_SIZE,\n        sw_batch_size=CFG.SW_BATCH_SIZE,\n        overlap=CFG.OVERLAP,\n        mode=CFG.SW_MODE,\n        padding_mode=CFG.PADDING_MODE,\n    )\n\n    # TTA setup\n    # tta_transforms = get_tta_transforms() if use_tta else [(None, 0)]\n\n    # Keep volume on CPU to avoid GPU-to-GPU transfer\n    volume_tensor = volume_tensor.cpu()\n\n    # ------------------------------------------------------------------\n    # Run inference in parallel (use each model's actual device - TRTWrapper is GPU-bound)\n    # ------------------------------------------------------------------\n    futures = []\n    with ThreadPoolExecutor(max_workers=len(devices)) as executor:\n        for i, (model, tta_fn) in enumerate(zip(models, tta_fns)):\n            device = _get_model_device(model)\n            futures.append(\n                executor.submit(\n                    _predict_single_model_worker,\n                    model,\n                    volume_tensor,\n                    device,\n                    inferer,\n                    tta_fn(),\n                )\n            )\n\n    # Collect predictions (each is np.ndarray with same shape)\n    all_predictions = [f.result() for f in futures]\n\n    # ------------------------------------------------------------------\n    # Weighted ensemble\n    # ------------------------------------------------------------------\n    if len(all_predictions) == 1:\n        return all_predictions[0]\n\n    # Normalize weights\n    weights_np = np.asarray(weights, dtype=np.float32)\n    weights_np /= weights_np.sum()\n\n    # Stack predictions: (N_models, ...)\n    stacked_preds = np.stack(all_predictions, axis=0)\n\n    # Expand weights for broadcasting\n    # (N_models,) -> (N_models, 1, 1, 1, ...)\n    expand_shape = (len(weights_np),) + (1,) * (stacked_preds.ndim - 1)\n    weights_np = weights_np.reshape(expand_shape)\n\n    # Weighted sum\n    ensemble_pred = np.sum(stacked_preds * weights_np, axis=0)\n\n    return ensemble_pred","metadata":{"_uuid":"71531732-5387-4dcd-980f-6ae4d1987b13","_cell_guid":"db34f649-0166-4361-ae54-67b5910dc486","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-02-24T22:22:40.434856Z","iopub.execute_input":"2026-02-24T22:22:40.435186Z","iopub.status.idle":"2026-02-24T22:22:40.44947Z","shell.execute_reply.started":"2026-02-24T22:22:40.435158Z","shell.execute_reply":"2026-02-24T22:22:40.448701Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Post process\n\n\nimport scipy.ndimage as ndimage\nfrom skimage import morphology\nimport numpy as np","metadata":{"_uuid":"2c9ba066-996d-4f82-9da6-aa68a9bc4a17","_cell_guid":"ea286340-eab0-4113-aede-176308dc9cd5","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-02-24T22:22:40.450275Z","iopub.execute_input":"2026-02-24T22:22:40.450544Z","iopub.status.idle":"2026-02-24T22:22:40.463354Z","shell.execute_reply.started":"2026-02-24T22:22:40.450521Z","shell.execute_reply":"2026-02-24T22:22:40.462692Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def postprocess_mask_voxel(mask, min_cc_volume=3000, verbose=True):\n    \"\"\"\n    Post-process binary mask using voxel-based operations:\n    1. Remove small connected components\n    2. Fill small holes/gaps using binary closing\n    \n    Returns:\n        processed_mask, stats_dict (if verbose) or just processed_mask\n    \"\"\"\n    if mask.sum() == 0:\n        if verbose:\n            return mask, {\"original_cc\": 0, \"removed_cc\": 0, \"kept_cc\": 0}\n        return mask\n\n    try:\n        mask_bool = mask.astype(bool)\n        \n        # Count original connected components\n        labeled_original, num_original = ndimage.label(mask_bool)\n        \n        # Get sizes of each component\n        component_sizes = ndimage.sum(mask_bool, labeled_original, range(1, num_original + 1))\n        \n        # Count how many are below threshold\n        num_small = np.sum(component_sizes < min_cc_volume)\n        num_kept = num_original - num_small\n        \n        # Remove small connected components\n        cleaned = morphology.remove_small_objects(mask_bool, min_size=min_cc_volume)\n        \n        # Fill holes / close gaps\n        #closed = ndimage.binary_closing(cleaned, structure=np.ones((3,3,3)))\n        \n        if verbose:\n            stats = {\n                \"original_cc\": num_original,\n                \"removed_cc\": int(num_small),\n                \"kept_cc\": int(num_kept),\n                \"component_sizes\": component_sizes\n            }\n            return cleaned.astype(np.uint8), stats\n        \n        return cleaned.astype(np.uint8)\n        \n    except Exception as e:\n        print(f\"    Error in voxel post-processing: {e}\")\n        if verbose:\n            return mask, {\"error\": str(e)}\n        return mask","metadata":{"_uuid":"e373404b-2361-4341-b937-111a891b63a6","_cell_guid":"42a2eee9-feda-49e5-9b5e-f49e0c1afbf0","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-02-24T22:22:40.464117Z","iopub.execute_input":"2026-02-24T22:22:40.464426Z","iopub.status.idle":"2026-02-24T22:22:40.476628Z","shell.execute_reply.started":"2026-02-24T22:22:40.464362Z","shell.execute_reply":"2026-02-24T22:22:40.476042Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Run Inference on Test Data","metadata":{"_uuid":"31c84d4b-22cf-4ee6-928f-e89cc7828976","_cell_guid":"34d51a30-c241-4d34-b5dd-c24eed5a5e43","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"if CFG.TEST_IMG_DIR.exists():\n    test_files = sorted([f for f in CFG.TEST_IMG_DIR.glob(\"*.tif\")])\n    print(f\"Found {len(test_files)} test .tif files\")\n    \n    if len(test_files) > 0:\n        print(\"\\nTest files:\")\n        for f in test_files[:5]:\n            print(f\"  - {f.name}\")\n        if len(test_files) > 5:\n            print(f\"  ... and {len(test_files) - 5} more\")\nelse:\n    print(f\"Warning: Test directory not found at {CFG.TEST_IMG_DIR}\")\n    print(\"Please update CFG.TEST_IMG_DIR to point to the correct directory\")\n    test_files = []","metadata":{"_uuid":"2f8f507d-5233-4d26-acd2-37af6c99d396","_cell_guid":"de0936c3-5ec4-4f58-a4cd-e36367f894c0","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-02-24T22:22:40.477358Z","iopub.execute_input":"2026-02-24T22:22:40.477642Z","iopub.status.idle":"2026-02-24T22:22:40.499286Z","shell.execute_reply.started":"2026-02-24T22:22:40.47762Z","shell.execute_reply":"2026-02-24T22:22:40.498571Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!mkdir first_stage_masks","metadata":{"_uuid":"1d133c03-879a-45a9-aa48-b4474081fef8","_cell_guid":"1010c76f-459e-419e-9a6f-3a09c9e8bd89","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-02-24T22:22:40.500071Z","iopub.execute_input":"2026-02-24T22:22:40.500376Z","iopub.status.idle":"2026-02-24T22:22:40.644479Z","shell.execute_reply.started":"2026-02-24T22:22:40.500358Z","shell.execute_reply":"2026-02-24T22:22:40.643356Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create dataset and dataloader\ntest_dataset = InferenceDataset(test_files)\ntest_loader = DataLoader(test_dataset, batch_size=1, shuffle=False, num_workers=0, collate_fn=custom_collate_fn)\n\ntta_status = \"WITH TTA\" if CFG.USE_TTA else \"(no TTA)\"\nprint(f\"\\nRunning 2-stage inference with sliding window {tta_status}...\\n\")\n\nif CFG.USE_TTA:\n    tta_count = len(get_tta_transforms())\n    print(f\"Using {tta_count} TTA transforms per model\")\n    total_models = sum(len(paths) for paths in grouped_model_paths.values())\n    print(f\"Total predictions per volume (1st stage): {total_models} models × {tta_count} TTA = {total_models * tta_count}\\n\")\n\ntif_dir = CFG.OUTPUT_DIR / \"submission_tifs\"","metadata":{"_uuid":"de2d00de-813c-49c8-a7e4-93a1d90642b9","_cell_guid":"7d191ea7-1e1c-42fc-950b-365f1a472014","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-02-24T22:22:40.645737Z","iopub.execute_input":"2026-02-24T22:22:40.646037Z","iopub.status.idle":"2026-02-24T22:22:40.652986Z","shell.execute_reply.started":"2026-02-24T22:22:40.645986Z","shell.execute_reply":"2026-02-24T22:22:40.652431Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"first_stage_mask_dir = Path(\"./first_stage_masks\")","metadata":{"_uuid":"3b610cc7-e1c2-4ba6-974b-78c5c06de7e2","_cell_guid":"f20b0c7f-69a7-4008-8980-ee6de37d08e4","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-02-24T22:22:40.653964Z","iopub.execute_input":"2026-02-24T22:22:40.65423Z","iopub.status.idle":"2026-02-24T22:22:40.66635Z","shell.execute_reply.started":"2026-02-24T22:22:40.654214Z","shell.execute_reply":"2026-02-24T22:22:40.665552Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Stage1 Inference","metadata":{"_uuid":"6c14ca44-20f3-4bf3-bbe8-d53585c11e1f","_cell_guid":"1a5cd3a1-6d35-4794-8da0-f00e536a8754","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"# Run inference and save predictions directly to .npz files\n# LOOP OVER GROUPS FIRST: load one group -> process all volumes -> unload -> next group (avoids OOM)\nif len(test_files) > 0 and len(grouped_model_paths) > 0:\n    processed_count = 0\n    print(f\"Doing 1st stage inference sequentially — groups: {group_names_1st}\")\n\n    for group_name, group_paths in grouped_model_paths.items():\n        print(f\"\\n--- Loading and processing group: {group_name} ---\")\n\n        # 1. Load models strictly for THIS group\n        group_entries = load_models_simple(group_paths)\n        group_models = [e[\"model\"] for e in group_entries]\n        group_tta_fns = [e[\"tta_fn\"] for e in group_entries]\n        group_weights = [1.0 / len(group_models)] * len(group_models)\n\n        # 2. Process all volumes through this group\n        for batch in tqdm(test_loader, desc=f\"Processing volumes for {group_name}\"):\n            volume = batch['volume']  # (1, 1, D, H, W)\n            vol_shape = tuple(batch['shape'][0].numpy())\n            filename = batch['filename'][0]\n            scroll_id = filename.replace('.tif', '')\n\n            probs = predict_volume_sliding_window(\n                group_models,\n                volume,\n                group_weights,\n                devices=[\"cuda:0\", \"cuda:1\"],\n                tta_fns=group_tta_fns,\n                use_tta=CFG.USE_TTA\n            )\n            binary = (probs > CFG.THRESHOLD_1ST_STAGE)\n            np.savez_compressed(first_stage_mask_dir / f'{group_name}_{scroll_id}.npz', mask=binary)\n\n            del probs, binary\n\n        # 3. UNLOAD models to free VRAM for the next group\n        print(f\"Unloading group {group_name} from VRAM...\")\n        for e in group_entries:\n            if hasattr(e[\"model\"], \"context\"):\n                del e[\"model\"].context\n            if hasattr(e[\"model\"], \"engine\"):\n                del e[\"model\"].engine\n            del e[\"model\"]\n        del group_entries, group_models, group_tta_fns\n        torch.cuda.empty_cache()\n        gc.collect()\n\n    processed_count = len(test_files)\n    print(f\"\\n✓ 1st stage Inference complete! Saved {processed_count * len(group_names_1st)} .npz files to {first_stage_mask_dir}\")\nelse:\n    print(\"No test files or models available.\")\n    processed_count = 0","metadata":{"_uuid":"09232c4d-ba4b-49ea-ad81-e09c9299e680","_cell_guid":"2be50fba-5769-43a7-8ac2-15653b3ae5e3","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-02-24T22:22:40.667149Z","iopub.execute_input":"2026-02-24T22:22:40.667412Z","iopub.status.idle":"2026-02-24T22:30:05.77378Z","shell.execute_reply.started":"2026-02-24T22:22:40.667392Z","shell.execute_reply":"2026-02-24T22:30:05.772987Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==============================================================================\n# Model Loading - 2nd stage (ensemble model)\n# ==============================================================================\n\n\nMODEL_PATHS_2ND_STAGE = [\n    (\"/kaggle/input/notebooks/tom99763/tensorrt-models-refine-v2-4ch-more-epoch/resenc-refine-v2-4ch-resenc-refinev2-4ch-best-epoch119-val_dice0.6102-val_loss0.3507.engine\", 0, TRTWrapper),\n    (\"/kaggle/input/notebooks/tom99763/tensorrt-models-refine-v2-4ch-more-epoch/resenc-refine-v2-4ch-resenc-refinev2-4ch-best-epoch134-val_dice0.6015-val_loss0.3528.engine\", 1, TRTWrapper),\n    #(\"/kaggle/input/notebooks/tom99763/tensorrt-models-refine-v2-4ch/resenc-refine-v2-4ch-resenc-refinev2-4ch-best-epoch69-val_dice0.6100-val_loss0.3517.engine\", 0, TRTWrapper),\n    #(\"/kaggle/input/notebooks/tom99763/tensorrt-models-refine-v2-4ch/resenc-refine-v2-4ch-resenc-refinev2-4ch-best-epoch74-val_dice0.6005-val_loss0.3544.engine\", 1, TRTWrapper),\n]\n\n_entries_2nd_stage = load_models_simple(MODEL_PATHS_2ND_STAGE)\nmodels_2nd_stage = [e[\"model\"] for e in _entries_2nd_stage]\ntta_2nd_stage = [get_tta_transforms_2nd_v2] * len(models_2nd_stage)\nsecond_stage_w = [1.0 / len(models_2nd_stage)] * len(models_2nd_stage)","metadata":{"_uuid":"c1ee3d9f-b561-4820-a3e4-d1cc5bab2740","_cell_guid":"02dd2cf3-8451-4ea9-8363-d68db9157b75","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==============================================================================\n# Model Loading - 3rd stage Refinement\n# ==============================================================================\n\n\nMODEL_PATHS_3RD_STAGE = [\n    (\"/kaggle/input/notebooks/iamparadox/newrefinev1-tensorrt/refinenewV1/refine-v1-fold0-best-epoch69-val_dice0.6158-val_loss0.3396.engine\", 0,TRTWrapper),\n    (\"/kaggle/input/notebooks/iamparadox/newrefinev1-tensorrt/refinenewV1/refine-v1-fold1-best-epoch74-val_dice0.6267-val_loss0.3363.engine\", 1, TRTWrapper),\n    #(\"/kaggle/input/tensorrt-models-0683bc/trt_models/2nd_stage_best-epoch104-val_dice0.5820-val_loss0.3709_2nd_v1_fp16.engine\", 0, TRTWrapper),\n    #(\"/kaggle/input/tensorrt-models-0683bc/trt_models/best-epoch=49-val_dice=0.5794-val_loss=0.3862_2nd_v1_fp16.engine\", 1,TRTWrapper)\n]\n\n\n_entries_3rd_stage = load_models_simple(MODEL_PATHS_3RD_STAGE) if MODEL_PATHS_3RD_STAGE else []\nmodels_3rd_stage = [e[\"model\"] for e in _entries_3rd_stage]\ntta_3rd_stage = [get_tta_transforms_2nd] * len(models_3rd_stage)\nthird_stage_w = [1.0 / len(models_3rd_stage)] * len(models_3rd_stage) if models_3rd_stage else []","metadata":{"_uuid":"44469cdf-1f0d-4158-9551-97f779f83bc5","_cell_guid":"de7fc8ad-5b26-4e76-8ed5-ac0fed87fa30","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-02-24T22:30:05.782099Z","iopub.execute_input":"2026-02-24T22:30:05.782309Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==============================================================================\n# Model Loading - 4th stage Refinement\n# ==============================================================================\n\n\nMODEL_PATHS_4TH_STAGE = [\n    (\"/kaggle/input/models/tom99763/diffeomophic-3stage-vesuvius-challenge/pytorch/default/24/refine-v1-3nd-stage/fold-0-best-epoch=99-val_dice=0.5995-val_loss=0.3575.ckpt\", 0, SegmentationModule2ndStage),\n    (\"/kaggle/input/models/tom99763/diffeomophic-3stage-vesuvius-challenge/pytorch/default/34/refine-v1-3nd-stage/fold-1-best-epoch=64-val_dice=0.6097-val_loss=0.3562.ckpt\", 1, SegmentationModule2ndStage)\n   \n]\n\n\n_entries_4th_stage = load_models_simple(MODEL_PATHS_4TH_STAGE)\nmodels_4th_stage = [e[\"model\"] for e in _entries_4th_stage]\ntta_4th_stage = [get_tta_transforms_3nd_refine] * len(models_4th_stage)\nfourth_stage_w = [1.0 / len(models_4th_stage)] * len(models_4th_stage)","metadata":{"_uuid":"4c923e59-e5f4-4567-99dc-ffd8509ff96e","_cell_guid":"183b290e-64e6-4199-a34e-a0210be9ef45","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Diffeomorphic Network","metadata":{"_uuid":"1f72161f-529b-494f-9da2-b8af610f69ff","_cell_guid":"d2ff98ee-a3df-4e6e-9f71-b07f96bf07b5","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"# Diffeo exponentiation and warper (same as before)\ndef make_base_grid(B, D, H, W, device):\n    zz = torch.linspace(0, D-1, D, device=device)\n    yy = torch.linspace(0, H-1, H, device=device)\n    xx = torch.linspace(0, W-1, W, device=device)\n    zz, yy, xx = torch.meshgrid(zz, yy, xx, indexing='ij')  # D,H,W\n    grid = torch.stack((xx, yy, zz), dim=3)  # D,H,W,3 (x,y,z)\n    grid = grid.unsqueeze(0).repeat(B,1,1,1,1)  # B,D,H,W,3\n    return grid\n\ndef disp_to_grid_for_sampling(disp_voxel: torch.Tensor):\n    B, C, D, H, W = disp_voxel.shape\n    device = disp_voxel.device\n    grid = make_base_grid(B, D, H, W, device)  # B,D,H,W,3 (x,y,z)\n    disp = disp_voxel.permute(0,2,3,4,1)  # B,D,H,W,3 (dx,dy,dz)\n    pos = grid + disp\n    pos_norm = torch.empty_like(pos)\n    pos_norm[...,0] = 2.0 * pos[...,0] / max(W-1,1) - 1.0  # x\n    pos_norm[...,1] = 2.0 * pos[...,1] / max(H-1,1) - 1.0  # y\n    pos_norm[...,2] = 2.0 * pos[...,2] / max(D-1,1) - 1.0  # z\n    return pos_norm\n\ndef warp_vol_using_disp(vol: torch.Tensor, disp_voxel: torch.Tensor, mode='bilinear'):\n    pos_norm = disp_to_grid_for_sampling(disp_voxel)\n    warped = F.grid_sample(vol, pos_norm, mode=mode, padding_mode='border', align_corners=True)\n    return warped\n\ndef warp_displacement(disp_voxel: torch.Tensor, by_disp_voxel: torch.Tensor):\n    warped = warp_vol_using_disp(by_disp_voxel, disp_voxel, mode='bilinear')\n    return warped\n\ndef scaling_and_squaring(v, n_steps=6) -> torch.Tensor:\n    flow = v / (2.0 ** n_steps)\n    for _ in range(n_steps):\n        flowed = warp_displacement(flow, flow)\n        flow = flow + flowed\n    return flow\n\ndef soft_sdf(x, eps=1e-4):\n    # x in [0,1]\n    return torch.log(x + eps) - torch.log(1 - x + eps)\n\nclass TopoFix(nn.Module):\n    def __init__(self, max_offset=2.0):\n        super().__init__()\n        self.max_offset = max_offset\n\n    def forward(self, warped_mask, raw_t):\n        \"\"\"\n        warped_mask: (B,1,D,H,W) in [0,1]\n        raw_t: network raw 4th channel (B,1,D,H,W), can be positive or negative\n        \"\"\"\n        sdf = soft_sdf(warped_mask)          # convert to SDF\n        t = torch.sigmoid(raw_t)             # gate: where to apply\n        delta = self.max_offset * torch.tanh(raw_t)  # signed magnitude\n        sdf_corr = sdf + delta * t            # apply offset only where t>0\n        corrected = torch.sigmoid(sdf_corr)  # back to probability\n        return corrected, t, delta\n\nclass DiffeomorphicNetwork(nn.Module):\n    def __init__(self, in_channels, out_channels, n_steps, max_v, max_topo_offset):\n        super().__init__()\n        '''\n        max_v: 1.5\n        n_steps: 6\n        lambda_jac: 0.3\n        lambda_smooth: 0.05\n        lambda_ce: 0.5\n        lambda_dice: 1.5\n        lambda_sparse: 0.1\n        lambda_tv: 0.02\n        lambda_boundary: 0.1\n        max_topo_offset: 1.0\n        '''\n        self.predictor = create_residual_unet(\n            in_channels=in_channels,\n            out_channels=out_channels\n        )\n        self.max_v = max_v\n        self.n_steps = n_steps\n        self.topofix = TopoFix(max_offset=max_topo_offset)\n\n    def forward(self, x, return_params=False):\n        raw = self.predictor(x)\n        raw_v = raw[:, :3]\n        raw_t = raw[:, 3:4]\n\n        # SVF\n        v = torch.tanh(raw_v) * self.max_v\n        phi = scaling_and_squaring(v, n_steps=self.n_steps)\n\n        # warp\n        warped = warp_vol_using_disp(x[:, 1:2], phi)\n\n        # topo fix\n        corrected, t, delta = self.topofix(warped, raw_t)\n\n        if return_params:\n            return corrected, v, phi, t\n        return corrected\n\n\ndef gaussian_kernel_3d(kernel_size=5, sigma=1.0, device=\"cuda\"):\n    \"\"\"Returns a normalized 3D Gaussian kernel (1,1,K,K,K).\"\"\"\n    ax = torch.arange(kernel_size, device=device) - kernel_size // 2\n    xx, yy, zz = torch.meshgrid(ax, ax, ax, indexing='ij')\n    kernel = torch.exp(-(xx**2 + yy**2 + zz**2) / (2 * sigma**2))\n    kernel = kernel / kernel.sum()\n    return kernel\n\n\ndef gaussian_blur_3d(x, kernel_size=3, sigma=5.0):\n    \"\"\"\n    x: (B, C, D, H, W)\n    \"\"\"\n    B, C, D, H, W = x.shape\n    kernel = gaussian_kernel_3d(kernel_size, sigma, device=x.device)\n\n    # shape: (C, 1, K, K, K)\n    kernel = kernel.expand(C, 1, kernel_size, kernel_size, kernel_size)\n\n    # depthwise convolution\n    return F.conv3d(x, kernel, padding=kernel_size // 2, groups=C)\n\n\nclass SegmentatioModule(pl.LightningModule):\n    def __init__(self, in_channels=2, out_channels=4,\n                 n_steps=6, max_v=1.5, max_topo_offset=1.0\n                 ):\n        super().__init__()\n        self.model = DiffeomorphicNetwork(\n            in_channels, out_channels, n_steps, max_v, max_topo_offset)\n\n    def forward(self, x, return_params=False):\n        return self.model(x, return_params = return_params)","metadata":{"_uuid":"85e58c31-9e7a-4c3f-b915-dde5ecf33338","_cell_guid":"0338b2cd-d289-4e78-8e7f-9867f8f95c61","trusted":true,"collapsed":false,"_kg_hide-input":true,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Deformnets from Sergio codebase\n# MODEL_PATHS_3ND_STAGE = [\n#     (\"/kaggle/input/models/tom99763/diffeomophic-3stage-vesuvius-challenge/pytorch/default/17/deformnet-refine-v2-4ch-oof/best-epoch=34-val_dice=0.6282-val_loss=0.4139.ckpt\", \"cuda:0\", SegmentatioModule),\n#     (\"/kaggle/input/models/tom99763/diffeomophic-3stage-vesuvius-challenge/pytorch/default/17/deformnet-refine-v2-4ch-oof/best-epoch=44-val_dice=0.6166-val_loss=0.4182.ckpt\", \"cuda:0\", SegmentatioModule),\n# ]\n\n# _entries_3nd_stage = load_models_simple(MODEL_PATHS_3ND_STAGE)\n# models_3nd_stage = [e[\"model\"] for e in _entries_3nd_stage]\n# tta_3nd_stage = [get_tta_transforms_3nd] * len(models_3nd_stage)\n# third_stage_w = [1.0 / len(models_3nd_stage)] * len(models_3nd_stage)\n# _kernel_size = 3\n# _sigma = 5","metadata":{"_uuid":"c6e5c585-6366-471f-9343-515a089a2870","_cell_guid":"01407683-bd25-4d36-87ef-ce37644c3753","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load Deformnet\ndef load_model_list(ckpts, cfg, model_arch):\n    models = []\n    for (ckpt_path, device) in tqdm(ckpts):\n        m = model_arch(cfg).to(device)\n        m.eval()\n        load_model_from_checkpoint(m, ckpt_path)\n        models.append(m)\n    return models\n\ndeformnet_list_1 = load_model_list(deformnet_ckpt_paths_1, cfg_1, DeformDynUnetV2)\n\ndeformnet_transforms = generate_transforms(cfg_0.data.transforms.test)\n\nsliding_window_inferer_1 = SlidingWindowInfererAdapt(\n    roi_size=cfg_1.input_size, sw_batch_size=1, overlap=0.5, mode=\"gaussian\",\n    progress = True\n)","metadata":{"_uuid":"b7abefd3-a580-4181-8728-906b8901aaf0","_cell_guid":"99bbcc2c-ad70-4089-bb0d-ca215474e8f4","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Stage 2, 3, 4 Inference","metadata":{"_uuid":"ed6991fe-52bb-4f91-bbb7-01814ddbe080","_cell_guid":"9c438892-6c48-4e6b-9ade-575abf1b9dc8","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"print(f\"2nd stage models: {len(models_2nd_stage)}, 3rd stage: {len(models_3rd_stage)}, 4th stage: {len(models_4th_stage)}\\n\")","metadata":{"_uuid":"a6ea9e59-9950-4c94-8d83-4a461539dcb0","_cell_guid":"5d023366-6e15-4b11-a249-fdfb7e54149d","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Run inference and save predictions directly to .tif files\nif len(test_files) > 0 and len(models_2nd_stage) > 0:\n    processed_count = 0\n    \n    print(\"Doing 2nd, 3rd, 4th stage inference\")\n    for batch in tqdm(test_loader, desc=\"Processing volumes\"):\n\n        torch.cuda.empty_cache()\n        gc.collect()\n        volume = batch['volume'].cpu()  # FORCE CPU\n        vol_shape = tuple(batch['shape'][0].numpy())\n        filename = batch['filename'][0]\n        scroll_id = filename.replace('.tif', '')\n\n        # Load 1st stage masks (CPU only)\n        first_stage_tensors = []\n        for group_name in group_names_1st:\n            mask = np.load(first_stage_mask_dir / f'{group_name}_{scroll_id}.npz')[\"mask\"].astype(np.float32)\n            first_stage_tensors.append(torch.from_numpy(mask).unsqueeze(0).unsqueeze(0))\n\n        # ------------------\n        # Ensemble model\n        # ------------------\n        input_2nd_stage_4ch = torch.cat([volume] + first_stage_tensors, dim=1)\n        p_ref_next_4ch = predict_volume_sliding_window(\n                models_2nd_stage,\n                input_2nd_stage_4ch,\n                second_stage_w,\n                devices=[\"cuda:0\",\"cuda:1\"],\n                tta_fns = tta_2nd_stage,\n                use_tta=True\n            )\n        p_ref_next = torch.from_numpy(p_ref_next_4ch)\n        b_ref_next = (p_ref_next > CFG.THRESHOLD_2ND_V2_STAGE).unsqueeze(0).unsqueeze(0).float()\n        del input_2nd_stage_4ch\n        torch.cuda.synchronize()\n        torch.cuda.empty_cache()\n        gc.collect()\n        print(\"done Ensemble model\")\n\n        # ------------------\n        # 3rd stage inference\n        # ------------------\n        to_third_stage_pred = b_ref_next\n        for _ in range(CFG.REPEAT):\n            volume_2ch = torch.cat([volume, to_third_stage_pred], dim=1)\n            third_stage_probs = predict_volume_sliding_window(\n                models_3rd_stage,\n                volume_2ch,\n                third_stage_w,\n                devices=[\"cuda:0\",\"cuda:1\"],\n                tta_fns = tta_3rd_stage,\n                use_tta=False\n            )\n            third_stage_binary = (third_stage_probs > CFG.THRESHOLD).astype(np.float32)\n            to_third_stage_pred = torch.from_numpy(third_stage_binary).unsqueeze(0).unsqueeze(0)\n            del volume_2ch\n            torch.cuda.synchronize()\n            torch.cuda.empty_cache()\n            gc.collect()\n        \n        print(\"done 3rd stage inference\")\n\n        # ------------------\n        # 4th stage inference\n        # ------------------\n        to_fourth_stage_pred = to_third_stage_pred\n        for _ in range(1):\n            volume_2ch = torch.cat([volume, to_fourth_stage_pred], dim=1)\n            fourth_stage_probs = predict_volume_sliding_window(\n                models_4th_stage,\n                volume_2ch,\n                fourth_stage_w,\n                devices=[\"cuda:0\",\"cuda:1\"],\n                tta_fns = tta_4th_stage,\n                use_tta=False\n            )\n            fourth_stage_binary = (fourth_stage_probs > CFG.THRESHOLD).astype(np.float32)\n            to_fourth_stage_pred = torch.from_numpy(fourth_stage_binary).unsqueeze(0).unsqueeze(0)\n            del volume_2ch\n            torch.cuda.synchronize()\n            torch.cuda.empty_cache()\n            gc.collect()\n        \n        print(\"done 4th stage inference\")\n\n        \n        # ------------------\n        # Deform\n        # ------------------\n        # to_third_stage_pred = gaussian_blur_3d(\n        #         to_third_stage_pred, _kernel_size, _sigma\n        #     )\n        # _input_to_3nd_stage = torch.cat([volume, to_third_stage_pred], dim=1)\n        # third_stage_probs = predict_volume_sliding_window(\n        #     models_3rd_stage,\n        #     _input_to_3rd_stage,\n        #     third_stage_w,\n        #     devices=[\"cuda:0\",\"cuda:1\"],\n        #     tta_fns = tta_3rd_stage,\n        #     use_tta=CFG.USE_TTA\n        # )\n        fourth_stage_probs = inference_deformnet(CFG.TEST_IMG_DIR/filename, fourth_stage_probs)\n        fourth_stage_binary = fourth_stage_probs > 0.5\n        binary_mask = fourth_stage_binary.astype(np.uint8)\n        torch.cuda.synchronize()\n        torch.cuda.empty_cache()\n        gc.collect()\n        print(\"done Deform + 4th stage\")\n        \n        # Apply post-processing if enabled\n        if CFG.USE_POST_PROCESSING:\n            try:\n                tqdm.write(f\"  Post-processing {scroll_id}...\", end=\"\")\n                binary_mask, stats = postprocess_mask_voxel(\n                    binary_mask, \n                    min_cc_volume=CFG.POST_PROCESS_MIN_CC_VOLUME,\n                    verbose=True\n                )\n                tqdm.write(f\" ✓ (removed {stats['removed_cc']}/{stats['original_cc']} components)\")\n            except Exception as e:\n                tqdm.write(f\"  ⚠ Warning: Post-processing failed for {scroll_id}: {e}\")\n                tqdm.write(\"  Saving original mask without post-processing...\")\n        \n        # Save to .tif\n        tif_path = tif_dir / f\"{scroll_id}.tif\"\n        tiff.imwrite(tif_path, binary_mask)\n        processed_count += 1\n        \n        # Free memory\n        #del first_stage_probs, first_stage_binary, first_stage_binary_tensor, volume_2ch, second_stage_probs, binary_mask\n        del p_ref_next, binary_mask\n        torch.cuda.empty_cache() if torch.cuda.is_available() else None\n        gc.collect()\n    \n    print(f\"\\n✓ Inference complete! Saved {processed_count} .tif files to {tif_dir}\")\nelse:\n    print(\"No test files or models available (need both 1st and 2nd stage models).\")\n    processed_count = 0","metadata":{"_uuid":"d64b0e27-ab03-45de-a9f7-d04ba1d1e77e","_cell_guid":"82123ebf-0fe2-4833-bece-2068256722b3","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Create Submission","metadata":{"_uuid":"3b22a2d9-2eee-432e-8cc4-bf6f37f1b544","_cell_guid":"dacf4137-aa49-4246-9fb9-8adc2791cd39","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"import zipfile\n\n# Zip all .tif files for submission\nif processed_count > 0:\n    tif_dir = CFG.OUTPUT_DIR / \"submission_tifs\"\n    tif_files = sorted(tif_dir.glob(\"*.tif\"))\n    zip_path = \"submission.zip\"\n    \n    with zipfile.ZipFile(zip_path, 'w', zipfile.ZIP_DEFLATED) as zipf:\n        for tif_file in tif_files:\n            zipf.write(tif_file, tif_file.name)\n    \n    print(\"\\n=== Submission Created ===\")\n    print(f\"Files: {len(tif_files)} .tif files\")\n    print(f\"Zip: {zip_path}\")\n    print(f\"Threshold: {CFG.THRESHOLD}\")\n    # print(f\"Models used: {len(models)}\")\n    print(f\"TTA: {'Enabled' if CFG.USE_TTA else 'Disabled'}\")\n    if CFG.USE_TTA:\n        print(f\"  - Total augmentations: {len(get_tta_transforms())}\")\n    print(f\"Post-processing: {'Enabled' if CFG.USE_POST_PROCESSING else 'Disabled'}\")\nelse:\n    print(\"No predictions to zip.\")","metadata":{"_uuid":"d66f79a5-b091-49bd-adb5-76336805c7a5","_cell_guid":"5eedadf3-69bf-4a48-8a7a-45144c6ca6e0","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Preview Visualizations","metadata":{"_uuid":"775d1eb1-192f-48dd-bcc7-fc37b4551bd0","_cell_guid":"46e129c4-68cd-4264-9d7d-947bcc432a03","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"# Display binary mask visualizations (if enabled)\nif CFG.SAVE_VISUALIZATIONS and processed_count > 0:\n    tif_dir = CFG.OUTPUT_DIR / \"submission_tifs\"\n    tif_files = sorted(tif_dir.glob(\"*.tif\"))[:3]  # Show first 3 files\n    \n    if tif_files:\n        print(\"\\n=== Binary Mask Previews ===\\n\")\n        \n        for tif_file in tif_files:\n            # Load the binary mask\n            binary_mask = tiff.imread(tif_file)\n            scroll_id = tif_file.stem\n            \n            # Show 2 slices\n            mid_slice = binary_mask.shape[0] // 2\n            fig, axes = plt.subplots(1, 2, figsize=(12, 6))\n            \n            axes[0].imshow(binary_mask[mid_slice], cmap='gray')\n            axes[0].set_title(f'{scroll_id} - Slice {mid_slice}')\n            axes[0].axis('off')\n            \n            axes[1].imshow(binary_mask[mid_slice + binary_mask.shape[0]//4], cmap='gray')\n            axes[1].set_title(f'{scroll_id} - Slice {mid_slice + binary_mask.shape[0]//4}')\n            axes[1].axis('off')\n            \n            plt.tight_layout()\n            plt.show()\n            \n        print(f\"\\n✓ Displayed previews for {len(tif_files)} files\")\n    else:\n        print(\"No .tif files found for visualization.\")\nelse:\n    if not CFG.SAVE_VISUALIZATIONS:\n        print(\"Visualizations disabled (set CFG.SAVE_VISUALIZATIONS = True to enable)\")\n    else:\n        print(\"No predictions to visualize.\")","metadata":{"_uuid":"527e43ac-4b68-4468-a793-e16c0286400c","_cell_guid":"b04e5df0-b706-4aec-9416-aaafbcdbf2fa","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"_uuid":"6c209c14-2905-4cb7-a1c7-13ac368c525e","_cell_guid":"5bffaeec-c083-4d57-ba65-323c9656dd73","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null}]}