{"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":758869,"databundleVersionId":15792096,"modelInstanceId":560051,"modelId":572638},{"sourceType":"modelInstanceVersion","sourceId":755796,"databundleVersionId":15752385,"modelInstanceId":560051,"modelId":572638},{"sourceType":"modelInstanceVersion","sourceId":766754,"databundleVersionId":15844636,"modelInstanceId":560051,"modelId":572638},{"sourceType":"modelInstanceVersion","sourceId":760513,"databundleVersionId":15799501,"modelInstanceId":560051,"modelId":572638},{"sourceType":"modelInstanceVersion","sourceId":757999,"databundleVersionId":15782074,"modelInstanceId":560051,"modelId":572638},{"sourceType":"modelInstanceVersion","sourceId":752197,"databundleVersionId":15707024,"modelInstanceId":560051,"modelId":572638},{"sourceType":"modelInstanceVersion","sourceId":743345,"databundleVersionId":15603642,"modelInstanceId":558962,"modelId":571547},{"sourceType":"modelInstanceVersion","sourceId":742034,"databundleVersionId":15588971,"modelInstanceId":558512,"modelId":571079},{"sourceType":"modelInstanceVersion","sourceId":751724,"databundleVersionId":15702935,"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-03-01T16:08:15.224376Z","iopub.execute_input":"2026-03-01T16:08:15.224801Z","iopub.status.idle":"2026-03-01T16:08:35.166515Z","shell.execute_reply.started":"2026-03-01T16:08:15.224783Z","shell.execute_reply":"2026-03-01T16:08:35.165815Z"},"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-03-01T16:08:35.168385Z","iopub.execute_input":"2026-03-01T16:08:35.168639Z","iopub.status.idle":"2026-03-01T16:09:03.81484Z","shell.execute_reply.started":"2026-03-01T16:08:35.168615Z","shell.execute_reply":"2026-03-01T16:09:03.814089Z"},"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-03-01T16:09:03.815872Z","iopub.execute_input":"2026-03-01T16:09:03.816188Z","iopub.status.idle":"2026-03-01T16:10:14.731848Z","shell.execute_reply.started":"2026-03-01T16:09:03.816136Z","shell.execute_reply":"2026-03-01T16:10:14.730956Z"},"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-03-01T16:10:16.172223Z","iopub.execute_input":"2026-03-01T16:10:16.172825Z","execution_failed":"2026-03-01T16:10:33.214Z"}},"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":{"execution_failed":"2026-03-01T16:10:33.215Z"},"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":{"execution_failed":"2026-03-01T16:10:33.215Z"},"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":{"execution_failed":"2026-03-01T16:10:33.215Z"},"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":{"execution_failed":"2026-03-01T16:10:33.216Z"},"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":{"execution_failed":"2026-03-01T16:10:33.216Z"}},"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":{"execution_failed":"2026-03-01T16:10:33.216Z"}},"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":{"execution_failed":"2026-03-01T16:10:33.217Z"},"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":{"execution_failed":"2026-03-01T16:10:33.217Z"}},"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":{"execution_failed":"2026-03-01T16:10:33.218Z"}},"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":{"execution_failed":"2026-03-01T16:10:33.218Z"}},"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":{"execution_failed":"2026-03-01T16:10:33.219Z"},"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":{"execution_failed":"2026-03-01T16:10:33.22Z"},"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":{"execution_failed":"2026-03-01T16:10:33.22Z"}},"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":{"execution_failed":"2026-03-01T16:10:33.22Z"}},"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":{"execution_failed":"2026-03-01T16:10:33.22Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport math\nimport multiprocessing\nfrom itertools import product\nfrom typing import List, Tuple\n\nimport cv2\nfrom numba import njit\nfrom skimage.morphology import skeletonize, remove_small_objects\nfrom skimage.measure import label\nfrom skimage.draw import line as skimage_line\nfrom scipy.ndimage import convolve, distance_transform_edt, binary_closing, binary_fill_holes, binary_dilation, \\\n    gaussian_filter\nimport numpy as np\nfrom scipy import ndimage as ndi, ndimage\nfrom scipy.ndimage import find_objects, binary_fill_holes\nfrom skimage.morphology import remove_small_objects, ball\nfrom skimage.draw import line\nfrom skimage.measure import label as measure_label\nfrom skimage.segmentation import watershed, find_boundaries\nfrom collections import Counter\nimport matplotlib.pyplot as plt\nfrom pathlib import Path\nimport tifffile as tiff\nfrom skimage import measure\nimport zipfile\nimport csv\n\n# from topometrics._bm_loader import load_betti_matching\n\n# 嘗試導入路徑計算模組\ntry:\n    from skimage.graph import route_through_array\nexcept ImportError:\n    from skimage.graph import MCP_Geometric\n\n\n    def route_through_array(cost, start, end):\n        mcp = MCP_Geometric(cost)\n        cumulative_costs, traceback = mcp.find_costs([start], [end])\n        return mcp.traceback(end), 0\n\n\n\n\n# =========================================================================\n#  Numba 加速核心 (Geodesic, PCA, Ray-Casting)\n# =========================================================================\n\n@njit(cache=True)\ndef compute_geodesic_distance_numba(mask, start_point):\n    \"\"\" Geodesic Distance Transform\"\"\"\n    d, h, w = mask.shape\n    dist_map = np.full((d, h, w), np.inf, dtype=np.float32)\n    sz, sy, sx = start_point\n\n    if mask[sz, sy, sx] == 0:\n        found = False\n        for dz in range(-1, 2):\n            for dy in range(-1, 2):\n                for dx in range(-1, 2):\n                    nz, ny, nx = sz + dz, sy + dy, sx + dx\n                    if 0 <= nz < d and 0 <= ny < h and 0 <= nx < w:\n                        if mask[nz, ny, nx] != 0:\n                            sz, sy, sx = nz, ny, nx\n                            found = True\n                            break\n                if found: break\n            if found: break\n        if not found: return np.zeros((d, h, w), dtype=np.float32)\n\n    max_queue_len = d * h * w\n    queue_z = np.empty(max_queue_len, dtype=np.int32)\n    queue_y = np.empty(max_queue_len, dtype=np.int32)\n    queue_x = np.empty(max_queue_len, dtype=np.int32)\n    head, tail = 0, 0\n\n    queue_z[tail], queue_y[tail], queue_x[tail] = sz, sy, sx\n    tail += 1\n    dist_map[sz, sy, sx] = 0.0\n\n    dz_off = np.array([0, 0, 0, 0, 1, -1], dtype=np.int8)\n    dy_off = np.array([0, 0, 1, -1, 0, 0], dtype=np.int8)\n    dx_off = np.array([1, -1, 0, 0, 0, 0], dtype=np.int8)\n\n    while head < tail:\n        cz, cy, cx = queue_z[head], queue_y[head], queue_x[head]\n        head += 1\n        next_dist = dist_map[cz, cy, cx] + 1.0\n\n        for i in range(6):\n            nz, ny, nx = cz + int(dz_off[i]), cy + int(dy_off[i]), cx + int(dx_off[i])\n            if 0 <= nz < d and 0 <= ny < h and 0 <= nx < w:\n                if mask[nz, ny, nx] != 0 and dist_map[nz, ny, nx] == np.inf:\n                    dist_map[nz, ny, nx] = next_dist\n                    queue_z[tail], queue_y[tail], queue_x[tail] = nz, ny, nx\n                    tail += 1\n    return dist_map\n\n\n@njit(fastmath=True)\ndef compute_pca_field_numba(labels, radius, step=1):\n    h, w = labels.shape\n    norm_y = np.zeros((h, w), dtype=np.float32)\n    norm_x = np.zeros((h, w), dtype=np.float32)\n\n    for r in range(0, h, step):\n        for c in range(0, w, step):\n            current_id = labels[r, c]\n            if current_id == 0: continue\n            sum_xx, sum_yy, sum_xy, count = 0.0, 0.0, 0.0, 0.0\n            r_min, r_max = max(0, r - radius), min(h, r + radius + 1)\n            c_min, c_max = max(0, c - radius), min(w, c + radius + 1)\n            for nr in range(r_min, r_max):\n                for nc in range(c_min, c_max):\n                    if labels[nr, nc] == current_id:\n                        dy, dx = nr - r, nc - c\n                        sum_xx += dx * dx\n                        sum_yy += dy * dy\n                        sum_xy += dx * dy\n                        count += 1.0\n            if count < 3: continue\n            A, B, C = sum_yy, sum_xy, sum_xx\n            tangent_angle = 0.5 * np.arctan2(2 * B, C - A)\n            normal_angle = tangent_angle + np.pi / 2.0\n            norm_y[r, c] = np.sin(normal_angle)\n            norm_x[r, c] = np.cos(normal_angle)\n    return norm_y, norm_x\n\n\n@njit(fastmath=True)\ndef find_first_collision_pair(y_coords, x_coords, pca_ny, pca_nx, mask, labels, ray_len):\n    h, w = mask.shape\n    num_points = len(y_coords)\n    for i in range(num_points):\n        r, c = y_coords[i], x_coords[i]\n        ny, nx = pca_ny[r, c], pca_nx[r, c]\n        if ny == 0.0 and nx == 0.0: continue\n        source_id = labels[r, c]\n        for sign in (1, -1):\n            has_left_self = False\n            for t in range(1, ray_len):\n                cy = int(r + sign * t * ny + 0.5)\n                cx = int(c + sign * t * nx + 0.5)\n                if cy < 0 or cy >= h or cx < 0 or cx >= w: break\n                pixel_val = mask[cy, cx]\n                target_id = labels[cy, cx]\n                if not has_left_self:\n                    if pixel_val == 0:\n                        has_left_self = True\n                    elif target_id != source_id:\n                        return True, r, c, cy, cx\n                else:\n                    if pixel_val > 0:\n                        if target_id != source_id:\n                            return True, r, c, cy, cx\n                        else:\n                            break\n    return False, 0, 0, 0, 0\n\n\n@njit(fastmath=True)\ndef find_all_collision_pairs(y_coords, x_coords, pca_ny, pca_nx, mask, labels, ray_len):\n    \"\"\"\n    修改版：不再找到第一個就停，而是收集所有碰撞對。\n    回傳: results list, 每個元素為 (r1, c1, r2, c2)\n    \"\"\"\n    h, w = mask.shape\n    num_points = len(y_coords)\n\n    # 建立一個 List 來儲存結果\n    # Numba 會自動推斷這是 List(Tuple(int64, int64, int64, int64))\n    results = []\n\n    for i in range(num_points):\n        r, c = y_coords[i], x_coords[i]\n        ny, nx = pca_ny[r, c], pca_nx[r, c]\n\n        # 如果 PCA 向量為 0 (無法計算方向)，跳過\n        if ny == 0.0 and nx == 0.0: continue\n\n        source_id = labels[r, c]\n\n        # 標記這個點是否已經找到對象，避免同一個點向左向右都加，造成重複 (視需求可保留雙向)\n        # 這裡設定為：如果正向找到就不找反向，確保每個起始點最多貢獻一條路徑\n        found_for_this_point = False\n\n        for sign in (1, -1):\n            has_left_self = False\n            for t in range(1, ray_len):\n                cy = int(r + sign * t * ny + 0.5)\n                cx = int(c + sign * t * nx + 0.5)\n\n                # 邊界檢查\n                if cy < 0 or cy >= h or cx < 0 or cx >= w: break\n\n                pixel_val = mask[cy, cx]\n                target_id = labels[cy, cx]\n\n                if not has_left_self:\n                    # 還沒離開自己\n                    if pixel_val == 0:\n                        has_left_self = True\n                    elif target_id != source_id:\n                        # 緊鄰就是不同 ID (沾黏嚴重)，視為碰撞\n                        results.append((r, c, cy, cx))\n                        found_for_this_point = True\n                        break\n                else:\n                    # 已經離開自己 (在背景中移動)\n                    if pixel_val > 0:\n                        # 撞到某個東西\n                        if target_id != source_id:\n                            # 撞到別人 -> 有效切割對\n                            results.append((r, c, cy, cx))\n                            found_for_this_point = True\n                            break\n                        else:\n                            # 撞回自己 (U型彎曲) -> 無效，停止這條射線\n                            break\n\n            if found_for_this_point:\n                break\n\n    return results\n\n\n# =========================================================================\n#  Python 輔助 (Seeds, Path)\n# =========================================================================\n\n\ndef get_split_seeds(mask_2d, ray_len=64, pca_radius=8, pca_step=3):\n    mask_uint8 = (mask_2d > 0).astype(np.uint8)\n    y_coords, x_coords = np.where(mask_uint8 > 0)\n\n    if len(y_coords) < 10: return []  # 改回傳空 list\n\n    # 降採樣取點 (Step sampling)\n    valid_indices = (y_coords % pca_step == 0) & (x_coords % pca_step == 0)\n    y_sampled = y_coords[valid_indices]\n    x_sampled = x_coords[valid_indices]\n\n    if len(y_sampled) == 0: return []\n\n    _, labels_map = cv2.connectedComponents(mask_uint8, connectivity=8)\n    labels_map = labels_map.astype(np.int32)\n\n    # 計算 PCA 場\n    pca_ny, pca_nx = compute_pca_field_numba(labels_map, pca_radius, step=pca_step)\n\n    # 使用新的 Numba 函數取得所有碰撞對\n    raw_results = find_all_collision_pairs(\n        y_sampled, x_sampled, pca_ny, pca_nx, mask_uint8, labels_map, ray_len\n    )\n\n    # 格式化輸出 [((y1,x1), (y2,x2)), ...]\n    seeds_list = []\n    for (r1, c1, r2, c2) in raw_results:\n        seeds_list.append(((r1, c1), (r2, c2)))\n\n    return seeds_list\n\n\n#\n# def analyze_shortest_path(mask_3d, start_pt, end_pt):\n#     \"\"\"\n#     優化版最短路徑分析：\n#     1. 僅對 Z, Y 軸進行 2 倍降採樣，X 軸保持原解析度 (Scale=1)。\n#     2. 加入 K-Neighbor Search 容錯。\n#     3. 回傳還原後的 Full-Resolution Path 供視覺化。\n#     \"\"\"\n#     # 1. 如果 Mask 很小，直接算原圖\n#     if mask_3d.shape[0] < 20 or mask_3d.shape[1] < 20:\n#         return _analyze_path_original(mask_3d, start_pt, end_pt)\n#\n#     # ==========================================\n#     # 設定降採樣步長： (Z=2, Y=2, X=1)\n#     # ==========================================\n#     step_z, step_y, step_x = 2, 2, 1\n#\n#     small_mask = mask_3d[::step_z, ::step_y, ::step_x]\n#\n#     # 2. 映射 Start/End 座標到小圖空間\n#     d_s, h_s, w_s = small_mask.shape\n#\n#     def to_small(pt):\n#         sz = min(pt[0] // step_z, d_s - 1)\n#         sy = min(pt[1] // step_y, h_s - 1)\n#         sx = min(pt[2] // step_x, w_s - 1)\n#         return (sz, sy, sx)\n#\n#     small_start = to_small(start_pt)\n#     small_end = to_small(end_pt)\n#\n#     # 3. 在小圖上計算距離變換 (EDT) 與路徑\n#     dist = distance_transform_edt(small_mask)\n#     cost_map = np.max(dist) - dist + 1\n#     cost_map[small_mask == 0] = np.inf\n#\n#     try:\n#         indices, _ = route_through_array(cost_map, small_start, small_end)\n#         path_indices = np.array(indices)\n#\n#         path_len = len(path_indices)\n#         if path_len == 0:\n#             return _analyze_path_original(mask_3d, start_pt, end_pt)\n#\n#         # =======================================================\n#         # K-Neighbor Search 尋找有效 Midpoint\n#         # =======================================================\n#         mid_idx = path_len // 2\n#         search_k = 5\n#         valid_midpoint = None\n#\n#         # 產生搜尋順序: 0, 1, -1, 2, -2 ...\n#         offsets = [0]\n#         for i in range(1, search_k + 1):\n#             offsets.append(i)\n#             offsets.append(-i)\n#\n#         for offset in offsets:\n#             current_idx = mid_idx + offset\n#             if 0 <= current_idx < path_len:\n#                 small_pt = path_indices[current_idx]\n#\n#                 # 映射回原始尺寸 (針對單點)\n#                 real_z = small_pt[0] * step_z + (step_z // 2)\n#                 real_y = small_pt[1] * step_y + (step_y // 2)\n#                 real_x = small_pt[2] * step_x\n#\n#                 # 邊界檢查\n#                 real_z = np.clip(real_z, 0, mask_3d.shape[0] - 1)\n#                 real_y = np.clip(real_y, 0, mask_3d.shape[1] - 1)\n#                 real_x = np.clip(real_x, 0, mask_3d.shape[2] - 1)\n#\n#                 candidate_pt = (int(real_z), int(real_y), int(real_x))\n#\n#                 if mask_3d[candidate_pt] > 0:\n#                     valid_midpoint = candidate_pt\n#                     break\n#\n#         if valid_midpoint is not None:\n#             # ===================================================\n#             # 新增：還原整條路徑座標 (Upscaling Path)\n#             # ===================================================\n#             full_res_path = path_indices.copy()\n#\n#             # 向量化計算：還原座標並加上中心偏移量\n#             # Z 軸\n#             full_res_path[:, 0] = full_res_path[:, 0] * step_z + (step_z // 2)\n#             # Y 軸\n#             full_res_path[:, 1] = full_res_path[:, 1] * step_y + (step_y // 2)\n#             # X 軸 (不變)\n#             full_res_path[:, 2] = full_res_path[:, 2] * step_x\n#\n#             # 統一進行邊界限制 (Clip)，防止偏移後超出原圖範圍\n#             max_coords = np.array(mask_3d.shape) - 1\n#             # 利用 numpy 的廣播機制將所有點限制在 [0, max]\n#             # 注意：np.clip 需要 array 為 float 或 int，這裡保持 int\n#             full_res_path[:, 0] = np.clip(full_res_path[:, 0], 0, max_coords[0])\n#             full_res_path[:, 1] = np.clip(full_res_path[:, 1], 0, max_coords[1])\n#             full_res_path[:, 2] = np.clip(full_res_path[:, 2], 0, max_coords[2])\n#\n#             return full_res_path, valid_midpoint\n#         else:\n#             return _analyze_path_original(mask_3d, start_pt, end_pt)\n#\n#     except Exception as e:\n#         print(e)\n#         print(f\"Analyze Path Error: {start_pt} -> {end_pt}\")\n#         return _analyze_path_original(mask_3d, start_pt, end_pt)\n\n\ndef _analyze_path_original(mask_3d, start_pt, end_pt):\n    \"\"\" 原始的全解析度路徑分析 (Fallback) \"\"\"\n    dist = distance_transform_edt(mask_3d)\n    cost_map = np.max(dist) - dist + 1\n    cost_map[mask_3d == 0] = np.inf\n    try:\n        indices, _ = route_through_array(cost_map, start_pt, end_pt)\n        path_indices = np.array(indices)\n        if len(path_indices) == 0:\n            print(f\"找不到路徑: {start_pt} -> {end_pt}\")\n            return None, None\n        midpoint = tuple(map(int, path_indices[len(path_indices) // 2]))\n        return path_indices, midpoint\n    except Exception as e:\n        print(f\"Analyze Path Error: {start_pt} -> {end_pt}\")\n        print(e)\n        return None, None\n\n\n# =========================================================================\n#  主函數: split_paper (使用 High Conf Region 作為 Watershed Markers)\n# =========================================================================\n\ndef get_split_paper(mask_3d, high_conf_labels, ray_len=64, max_iter=5, cleanup_iter=2):\n    \"\"\"\n    對 3D Binary Mask 執行迭代式幾何分割。\n\n    1. 遍歷 Slice 找尋候選切割點。\n    2. 投票找出最主要的 High Conf 連通域組合 (Pair A-B)。\n    3. 計算分割幾何邊界：透過 0, Mid, End 的種子路徑計算中點，建立 Geodesic Basin。\n    4. 【核心改動】建立 Markers：直接使用 Pair A 和 Pair B 在此 ROI 內的像素作為種子區域。\n    5. 執行 Watershed。\n    \"\"\"\n\n    if mask_3d.dtype == bool or np.max(mask_3d) == 1:\n        refined_labels = measure_label(mask_3d)\n    else:\n        refined_labels = measure_label(mask_3d > 0)\n\n    refined_labels = refined_labels.astype(np.int32)\n    next_new_label = refined_labels.max() + 1\n    debug_paths = []\n\n    print(f\">> 啟動 split_paper: Ray={ray_len}, Iter={max_iter}, Cleanup={cleanup_iter}\")\n\n    for iteration in range(max_iter):\n        unique_ids = np.unique(refined_labels)\n        unique_ids = unique_ids[unique_ids != 0]\n        splits_this_round = 0\n\n        for obj_id in unique_ids:\n            obj_mask = (refined_labels == obj_id)\n            slices = np.where(obj_mask)\n            if len(slices[0]) == 0: continue\n\n            # 取得目前的 Bounding Box\n            z_min, z_max = np.min(slices[0]), np.max(slices[0])\n            y_min, y_max = np.min(slices[1]), np.max(slices[1])\n            x_min, x_max = np.min(slices[2]), np.max(slices[2])\n\n            roi_mask = obj_mask[z_min:z_max + 1, y_min:y_max + 1, x_min:x_max + 1]\n            d, h, w = roi_mask.shape\n\n            # -----------------------------------------------------------\n            # 1. 遍歷所有 Slice 進行 Dense Sampling 並做投票\n            # -----------------------------------------------------------\n            candidate_seeds = []\n\n            for z_local in range(d):\n                slice_2d = roi_mask[z_local, :, :]\n\n                # 這裡現在會回傳一個 List，包含該層所有的切割候選\n                seeds_list = get_split_seeds(slice_2d, ray_len=ray_len)\n\n                # 如果這層有找到種子，遍歷它們\n                if seeds_list:\n                    for seeds in seeds_list:\n                        (y1, x1), (y2, x2) = seeds\n\n                        # 轉全域座標查表\n                        z_global = z_min + z_local\n                        y1_global, x1_global = y_min + y1, x_min + x1\n                        y2_global, x2_global = y_min + y2, x_min + x2\n\n                        # 邊界檢查\n                        if not (0 <= y1_global < high_conf_labels.shape[1] and 0 <= x1_global < high_conf_labels.shape[\n                            2]): continue\n                        if not (0 <= y2_global < high_conf_labels.shape[1] and 0 <= x2_global < high_conf_labels.shape[\n                            2]): continue\n\n                        id_a = high_conf_labels[z_global, y1_global, x1_global]\n                        id_b = high_conf_labels[z_global, y2_global, x2_global]\n\n                        # 核心邏輯：必須連接兩個不同的高信心物件\n                        if id_a != 0 and id_b != 0 and id_a != id_b:\n                            pair = tuple(sorted((id_a, id_b)))\n                            candidate_seeds.append({\n                                'z_local': z_local,\n                                'start_local': (z_local, y1, x1),\n                                'end_local': (z_local, y2, x2),\n                                'conf_pair': pair\n                            })\n\n            # -----------------------------------------------------------\n            # 2. 投票選出最佳組合 (A-B)\n            # -----------------------------------------------------------\n            if not candidate_seeds:\n                continue\n\n            pair_counts = Counter([c['conf_pair'] for c in candidate_seeds])\n            best_pair, count = pair_counts.most_common(1)[0]\n            id_comp_1, id_comp_2 = best_pair  # 取出這兩個高信心度連通域的 ID\n\n            # 僅保留最佳組合的種子\n            filtered_seeds = [c for c in candidate_seeds if c['conf_pair'] == best_pair]\n\n            # -----------------------------------------------------------\n            # 3. 排序 (按照 y, z, x) 並選取 3 個 (計算 Basin 幾何中心用)\n            # -----------------------------------------------------------\n            filtered_seeds.sort(key=lambda x: (x['start_local'][1], x['start_local'][0], x['start_local'][2]))\n\n            num_seeds = len(filtered_seeds)\n            if num_seeds == 0:\n                continue\n\n            indices_to_pick = sorted(list(set([0, num_seeds // 2, num_seeds - 1])))\n\n            valid_midpoints = []\n\n            for idx in indices_to_pick:\n                seed_info = filtered_seeds[idx]\n                p1 = seed_info['start_local']\n                p2 = seed_info['end_local']\n                # 計算路徑求 Midpoint (定義邊界位置)\n                path_idx, midpoint = _analyze_path_original(roi_mask, p1, p2)\n\n                if path_idx is not None:\n                    valid_midpoints.append(midpoint)\n\n                    # 存入 Debug Path 供顯示\n                    global_path = path_idx + np.array([z_min, y_min, x_min])\n                    debug_paths.append(global_path)\n                else:\n                    print(\"path_idx is None\")\n\n            if not valid_midpoints:\n                continue\n\n            # -----------------------------------------------------------\n            # 4. 計算 Geodesic Distance Field (基於 Midpoints)\n            # -----------------------------------------------------------\n            # 這定義了 \"哪裡是分界線\" (脊線 = 0)\n            fused_dist_map = np.full(roi_mask.shape, np.inf, dtype=np.float32)\n            roi_mask_int = roi_mask.astype(np.uint8)\n\n            for midpoint in valid_midpoints:\n                start_pt = (int(midpoint[0]), int(midpoint[1]), int(midpoint[2]))\n                d_map = compute_geodesic_distance_numba(roi_mask_int, start_pt)\n                fused_dist_map = np.minimum(fused_dist_map, d_map)\n\n            fused_dist_map[fused_dist_map == np.inf] = 0\n\n            # -----------------------------------------------------------\n            # 5. 【修改】使用 High Confidence Regions 作為 Markers\n            # -----------------------------------------------------------\n            basin_map = -1.0 * fused_dist_map\n\n            # 取出目前 ROI 範圍內的 High Conf Labels\n            high_conf_roi = high_conf_labels[z_min:z_max + 1, y_min:y_max + 1, x_min:x_max + 1]\n\n            markers = np.zeros_like(roi_mask, dtype=np.int32)\n\n            # 將 High Conf ID A 設為 Label 1\n            markers[(high_conf_roi == id_comp_1) & (roi_mask > 0)] = 1\n\n            # 將 High Conf ID B 設為 Label 2\n            markers[(high_conf_roi == id_comp_2) & (roi_mask > 0)] = 2\n\n            # 確保有兩個標記才執行\n            unique_markers = np.unique(markers)\n            if 1 in unique_markers and 2 in unique_markers:\n\n                # 執行分水嶺\n                labels_ws = watershed(basin_map, markers, mask=roi_mask)\n\n                if np.max(labels_ws) >= 2:\n                    # 選定要分離的部分 (這裡是 Label 2)\n                    split_mask_2 = (labels_ws == 2)\n\n                    roi_view = refined_labels[z_min:z_max + 1, y_min:y_max + 1, x_min:x_max + 1]\n                    roi_view[split_mask_2] = next_new_label\n\n                    # 邊界清理\n                    mask_old = (roi_view == obj_id)\n                    mask_new = (roi_view == next_new_label)\n                    struct_26 = np.ones((3, 3, 3), dtype=bool)\n\n                    neighbors_have_old = binary_dilation(mask_old, structure=struct_26, iterations=cleanup_iter)\n                    boundary_pixels_in_new = mask_new & neighbors_have_old\n\n                    neighbors_have_new = binary_dilation(mask_new, structure=struct_26, iterations=cleanup_iter)\n                    boundary_pixels_in_old = mask_old & neighbors_have_new\n\n                    roi_view[boundary_pixels_in_new] = 0\n                    roi_view[boundary_pixels_in_old] = 0\n\n                    next_new_label += 1\n                    splits_this_round += 1\n\n        print(f\"   Iter {iteration + 1}/{max_iter}: 分割了 {splits_this_round} 處。\")\n        if splits_this_round == 0:\n            break\n\n    final_binary = (refined_labels > 0).astype(np.uint8)\n    return final_binary, debug_paths\n\n\ndef get_split_paper_8slice(vol, high_conf_labels, ray_len=64, max_iter=5, cleanup_iter=2):\n    \"\"\"\n    將 3D 體積切分為 8 個區塊分別執行 get_split_paper。\n    \"\"\"\n    output_vol = np.zeros_like(vol, dtype=np.uint8)\n    all_debug_paths = []\n\n    # 取得 8 個區塊的 slice (2, 2, 2)\n    slices = _octant_slices(vol.shape, (2, 2, 2))\n\n    for slc in slices:\n        vol_block = vol[slc]\n        # 如果區塊內沒有 mask，直接跳過\n        if not np.any(vol_block):\n            continue\n\n        labels_block = high_conf_labels[slc]\n\n        # 執行原本的分割邏輯\n        refined_block, block_paths = get_split_paper(\n            vol_block,\n            labels_block,\n            ray_len=ray_len,\n            max_iter=max_iter,\n            cleanup_iter=cleanup_iter\n        )\n\n        # 將結果填回\n        output_vol[slc] = refined_block\n\n        # 校正 Debug 路徑座標 (加上 slice 的起始偏移量)\n        offset = np.array([slc[0].start, slc[1].start, slc[2].start])\n        for path in block_paths:\n            all_debug_paths.append(path + offset)\n\n    return output_vol, all_debug_paths\n\n\ndef get_outward_vector(skel, ep_xy, pca_vec):\n    \"\"\"計算骨架端點的朝外向量\"\"\"\n    start_x, start_y = int(ep_xy[0]), int(ep_xy[1])\n    h, w = skel.shape\n    STEPS = 5\n    visited = set()\n    visited.add((start_y, start_x))\n    curr_x, curr_y = start_x, start_y\n    for _ in range(STEPS):\n        found_next = False\n        for dy in [-1, 0, 1]:\n            for dx in [-1, 0, 1]:\n                if dx == 0 and dy == 0: continue\n                ny, nx = curr_y + dy, curr_x + dx\n                if ny < 0 or ny >= h or nx < 0 or nx >= w: continue\n                if skel[ny, nx] > 0 and (ny, nx) not in visited:\n                    curr_x, curr_y = nx, ny\n                    visited.add((ny, nx))\n                    found_next = True\n                    break\n            if found_next: break\n        if not found_next: break\n    ref_vec_x, ref_vec_y = start_x - curr_x, start_y - curr_y\n    if ref_vec_x == 0 and ref_vec_y == 0: return pca_vec\n    if ref_vec_x * pca_vec[0] + ref_vec_y * pca_vec[1] < 0: return -pca_vec\n    return pca_vec\n\n\ndef check_path_blocked(p1, p2, mask_img):\n    \"\"\"檢查兩點之間的路徑是否被障礙物阻擋\"\"\"\n    x1, y1 = int(p1[0]), int(p1[1])\n    x2, y2 = int(p2[0]), int(p2[1])\n    rr, cc = skimage_line(y1, x1, y2, x2)\n    valid_mask = (rr >= 0) & (rr < mask_img.shape[0]) & (cc >= 0) & (cc < mask_img.shape[1])\n    rr, cc = rr[valid_mask], cc[valid_mask]\n    if len(rr) == 0: return False\n    safe_radius_sq = 4 ** 2\n\n    def is_obstacle(px, py):\n        if px < 0 or px >= mask_img.shape[1] or py < 0 or py >= mask_img.shape[0]: return False\n        if mask_img[py, px] == 0: return False\n        d1 = (px - x1) ** 2 + (py - y1) ** 2\n        d2 = (px - x2) ** 2 + (py - y2) ** 2\n        return (d1 > safe_radius_sq) and (d2 > safe_radius_sq)\n\n    for i in range(len(rr)):\n        curr_y, curr_x = rr[i], cc[i]\n        if is_obstacle(curr_x, curr_y): return True\n        if i > 0:\n            prev_y, prev_x = rr[i - 1], cc[i - 1]\n            if abs(curr_x - prev_x) == 1 and abs(curr_y - prev_y) == 1:\n                if is_obstacle(prev_x, curr_y) and is_obstacle(curr_x, prev_y): return True\n    return False\n\n\ndef get_pca_tangent(skeleton, center_yx, radius=16, line_len=16):\n    \"\"\"使用 PCA 計算端點的切線方向\"\"\"\n    #\n    cy, cx = center_yx\n    h, w = skeleton.shape\n    y_min, y_max = max(0, cy - radius), min(h, cy + radius + 1)\n    x_min, x_max = max(0, cx - radius), min(w, cx + radius + 1)\n    roi = skeleton[y_min:y_max, x_min:x_max].copy()\n    if not np.any(roi): return None\n    roi_labels = label(roi > 0, connectivity=2)\n    local_cy, local_cx = cy - y_min, cx - x_min\n    target_label = roi_labels[local_cy, local_cx]\n    if target_label == 0: return None\n    pts_y, pts_x = np.where(roi_labels == target_label)\n    pts_global_y, pts_global_x = pts_y + y_min, pts_x + x_min\n    dist_sq = (pts_global_y - cy) ** 2 + (pts_global_x - cx) ** 2\n    mask = dist_sq <= (radius ** 2)\n    valid_y, valid_x = pts_global_y[mask], pts_global_x[mask]\n    if len(valid_x) < 2: return None\n    data = np.vstack([valid_x, valid_y]).T\n    mean = np.mean(data, axis=0)\n    cov = np.cov((data - mean).T)\n    if np.isnan(cov).any() or np.isinf(cov).any(): return None\n    eig_vals, eig_vecs = np.linalg.eigh(cov)\n    principal_vec = eig_vecs[:, -1]\n    vec_len = np.linalg.norm(principal_vec)\n    if vec_len == 0: return None\n    vx, vy = principal_vec / vec_len\n    pt1 = (int(cx - vx * line_len), int(cy - vy * line_len))\n    pt2 = (int(cx + vx * line_len), int(cy + vy * line_len))\n    return pt1, pt2\n\n\ndef get_endpoints(mask_2d):\n    \"\"\"取得 2D mask 的骨架端點\"\"\"\n    #\n    if not np.any(mask_2d): return [], None\n    skeleton = skeletonize(mask_2d.astype(bool))\n    skeleton_uint8 = skeleton.astype(np.uint8) * 255\n    kernel = np.array([[1, 1, 1], [1, 0, 1], [1, 1, 1]], dtype=np.uint8)\n    neighbors = convolve(skeleton.astype(np.uint8), kernel, mode='constant', cval=0)\n    coords = np.argwhere(skeleton & (neighbors <= 1))\n    if len(coords) == 0: return [], skeleton_uint8\n    raw_eps = [[pt[1], pt[0]] for pt in coords]\n    raw_eps.sort(key=lambda p: (p[1], p[0]))\n    return raw_eps, skeleton_uint8\n\n\ndef compute_candidate_info(all_eps, skel_img, radius=16):\n    \"\"\"計算所有候選端點的詳細資訊 (位置、向量)\"\"\"\n    ep_info_list = []\n    for i, (x, y) in enumerate(all_eps):\n        t_pts = get_pca_tangent(skel_img, (y, x), radius=radius, line_len=radius)\n        if not t_pts: continue\n        pt1, pt2 = t_pts\n        vx, vy = pt2[0] - pt1[0], pt2[1] - pt1[1]\n        norm = np.sqrt(vx ** 2 + vy ** 2)\n        if norm == 0: continue\n        pca_vec = np.array([vx / norm, vy / norm])\n        out_vec = get_outward_vector(skel_img, (x, y), pca_vec)\n        ep_info_list.append({'id': i, 'xy': np.array([x, y]), 'vec': out_vec, 'tangent_pts': t_pts})\n    return ep_info_list\n\n\ndef dynamic_filter(ep_info_list, skel_img, use_angle=True, use_block=True, use_pair=True, max_dist=15):\n    \"\"\"\n    篩選並配對需要連接的端點\n\n    Args:\n        max_dist (float): 允許連線的最大距離 (預設 15)\n    \"\"\"\n    num_eps = len(ep_info_list)\n    if num_eps < 2: return []\n\n    candidate_links = []\n    ANGLE_WEIGHT = 10.0\n    COS_THRES = np.cos(np.deg2rad(60))\n\n    for i in range(num_eps):\n        for j in range(i + 1, num_eps):\n            p1 = ep_info_list[i]\n            p2 = ep_info_list[j]\n            vec_p1_to_p2 = p2['xy'] - p1['xy']\n            dist = np.linalg.norm(vec_p1_to_p2)\n\n            # --- [修改點 1] ---\n            # 距離檢查：如果距離為 0 或大於等於 max_dist (15)，則跳過\n            if dist == 0: continue\n            if dist >= max_dist: continue\n            # ------------------\n\n            dir_1_to_2 = vec_p1_to_p2 / dist\n            dir_2_to_1 = -dir_1_to_2\n            cos_p1 = np.dot(p1['vec'], dir_1_to_2)\n            cos_p2 = np.dot(p2['vec'], dir_2_to_1)\n\n            if use_angle and (cos_p1 < COS_THRES or cos_p2 < COS_THRES): continue\n            if use_block and check_path_blocked(tuple(p1['xy']), tuple(p2['xy']), skel_img): continue\n\n            avg_cos = (cos_p1 + cos_p2) / 2.0\n            score = dist * (1.0 + ANGLE_WEIGHT * (1.0 - min(avg_cos, 1.0)))\n            candidate_links.append((score, i, j))\n\n    candidate_links.sort(key=lambda x: x[0])\n    matched_indices = set()\n    final_connections = []\n\n    if use_pair:\n        for score, idx1, idx2 in candidate_links:\n            if idx1 in matched_indices or idx2 in matched_indices: continue\n            p1, p2 = ep_info_list[idx1], ep_info_list[idx2]\n            # 再次檢查遮擋\n            if use_block and check_path_blocked(tuple(p1['xy']), tuple(p2['xy']), skel_img): continue\n            matched_indices.add(idx1)\n            matched_indices.add(idx2)\n            final_connections.append((p1, p2))\n    else:\n        pass\n\n    return final_connections\n\n\ndef execute_line_repair(vol_3d, settings=None):\n    \"\"\"\n    執行 3D 體積的逐層補線操作\n    \"\"\"\n    # --- [修改點 2] ---\n    # 設定預設值，確保 max_dist 存在\n    default_settings = {\n        'use_angle': True,\n        'use_block': True,\n        'use_pair': True,\n        'trim_ends': True,\n        'curve_str': 50,\n        'max_dist': 8  # 預設距離限制\n    }\n\n    if settings is None:\n        settings = default_settings\n    else:\n        # 若使用者傳入部分 settings，補齊未傳入的預設值\n        for k, v in default_settings.items():\n            if k not in settings:\n                settings[k] = v\n    # ------------------\n\n    print(f\"正在執行斷線修復 (Line Endpoint Repair), Max Dist: {settings['max_dist']}...\")\n\n    if vol_3d.dtype == bool:\n        labeled_vol = label(vol_3d)\n    else:\n        labeled_vol = label(vol_3d > 0)\n\n    patched_vol = (vol_3d > 0).astype(np.uint8) * 255\n    d, h, w = labeled_vol.shape\n\n    for z in range(d):\n        slice_labels = labeled_vol[z, :, :]\n        if not np.any(slice_labels): continue\n\n        slice_img = patched_vol[z, :, :]\n        present_ids = np.unique(slice_labels)\n        present_ids = present_ids[present_ids > 0]\n\n        for pid in present_ids:\n            obj_mask = (slice_labels == pid)\n            raw_eps, skel_img = get_endpoints(obj_mask)\n            if len(raw_eps) < 2: continue\n\n            ep_info_all = compute_candidate_info(raw_eps, skel_img, radius=16)\n\n            candidates = list(ep_info_all)\n            # 這裡保留原有的 trim_ends 邏輯，如果需要的話\n            if settings.get('trim_ends', False) and len(candidates) > 2:\n                candidates.sort(key=lambda p: (p['xy'][1], p['xy'][0]))\n                candidates = candidates[1:-1]  # 去頭去尾\n\n            # --- [修改點 3] ---\n            # 將 max_dist 傳入 dynamic_filter\n            connections = dynamic_filter(\n                candidates, skel_img,\n                use_angle=settings['use_angle'],\n                use_block=settings['use_block'],\n                use_pair=settings['use_pair'],\n                max_dist=settings['max_dist']\n            )\n            # ------------------\n\n            for (p1_obj, p2_obj) in connections:\n                pt1 = tuple(map(int, p1_obj['xy']))\n                pt2 = tuple(map(int, p2_obj['xy']))\n                cv2.line(slice_img, pt1, pt2, 255, 1, cv2.LINE_AA)\n\n        patched_vol[z, :, :] = slice_img\n\n    return patched_vol\n\n\ndef execute_line_repair_8slice(vol_3d, settings=None, splits=(2, 2, 2), pad_size=32):\n    \"\"\"\n    針對 execute_line_repair 的分塊並行/依序處理版本。\n    自動處理 Padding 以避免在切割邊界產生錯誤的端點判定。\n\n    Args:\n        vol_3d (np.ndarray): 原始 3D 陣列 (Boolean 或 Label)。\n        settings (dict): 傳遞給 execute_line_repair 的參數設定。\n        splits (tuple): (z, y, x) 切分份數，預設 (2, 2, 2) 為 8 塊。\n        pad_size (int): 擴充邊界大小。\n                        注意：必須大於 execute_line_repair 內部的 radius (預設16)，\n                        建議設為 32 或更大以確保向量計算正確。\n\n    Returns:\n        np.ndarray: 修復後的完整 3D 陣列 (uint8, 0-255)。\n    \"\"\"\n    print(f\"啟動分塊修復 (Grid: {splits}, Padding: {pad_size})...\")\n\n    # 1. 初始化輸出容器 (確保是 uint8，因為 execute_line_repair 回傳 uint8)\n    patched_full = np.zeros(vol_3d.shape, dtype=np.uint8)\n\n    # 2. 取得切分邏輯 (沿用之前的切分函數)\n    base_slices = _octant_slices(vol_3d.shape, splits)\n\n    for i, (sl_z, sl_y, sl_x) in enumerate(base_slices):\n        # --- A. 計算原始座標 ---\n        z_start, z_end = sl_z.start, sl_z.stop\n        y_start, y_end = sl_y.start, sl_y.stop\n        x_start, x_end = sl_x.start, sl_x.stop\n\n        # --- B. 計算 Padding 後的座標 (限制在圖像範圍內) ---\n        p_z_start = max(0, z_start)\n        p_z_end = min(vol_3d.shape[0], z_end)\n        p_y_start = max(0, y_start)\n        p_y_end = min(vol_3d.shape[1], y_end)\n        p_x_start = max(0, x_start)\n        p_x_end = min(vol_3d.shape[2], x_end)\n\n        # --- C. 取出子區塊 (含 Ghost Cells) ---\n        sub_vol = vol_3d[p_z_start:p_z_end, p_y_start:p_y_end, p_x_start:p_x_end]\n\n        # 這裡加個簡單的檢查，如果該區塊全是空的，就跳過運算以節省時間\n        if not np.any(sub_vol):\n            continue\n\n        # --- D. 執行核心修復 ---\n        # 注意：傳入子區塊進行運算。\n        # 由於 execute_line_repair 內部會重新 label，這對於局部修復是正確的行為。\n        # 只要 Padding 足夠，跨邊界的物件就能被視為連續。\n        patched_sub = execute_line_repair(sub_vol, settings)\n\n        # --- E. 裁切 (Remove Padding) ---\n        # 計算相對於 sub_vol 的有效區域偏移量\n        offset_z = z_start - p_z_start\n        offset_y = y_start - p_y_start\n        offset_x = x_start - p_x_start\n\n        len_z = z_end - z_start\n        len_y = y_end - y_start\n        len_x = x_end - x_start\n\n        valid_sub = patched_sub[\n            offset_z: offset_z + len_z,\n            offset_y: offset_y + len_y,\n            offset_x: offset_x + len_x\n        ]\n\n        # --- F. 填入結果 ---\n        patched_full[sl_z, sl_y, sl_x] = valid_sub\n\n        # (可選) 顯示進度\n        # print(f\"  - Block {i+1}/{len(base_slices)} done.\")\n\n    return patched_full\n\n\ndef line_repair_msk(mask3d: np.ndarray, min_area: int = 10, max_link_dist: float = 30.0,\n                    pass_iters: int = 1, line_thickness: int = 2):\n    \"\"\"\n    核心修補邏輯：檢查 3D 連通物件是否在 2D 切片上斷裂並修復。\n    \"\"\"\n\n    def label_3d_26(mask3d: np.ndarray):\n        \"\"\"26-connectivity 3D CC labeling.\"\"\"\n        structure = np.ones((3, 3, 3), dtype=bool)\n        return ndi.label(mask3d.astype(bool), structure=structure)\n\n    def label_2d_8(mask2d: np.ndarray):\n        \"\"\"8-connectivity 2D CC labeling.\"\"\"\n        structure = np.ones((3, 3), dtype=bool)\n        return ndi.label(mask2d.astype(bool), structure=structure)\n\n    def find_splits_in_slice(cc3d: np.ndarray, z: int, min_area: int = 10):\n        \"\"\"\n        找出在 slice z 上：同一個 3D label 出現 >=2 個 2D 連通塊的情況。\n        \"\"\"\n        lab2d = cc3d[z]\n        Ls = np.unique(lab2d)\n        Ls = Ls[Ls != 0]\n\n        splits = []\n        for L in Ls:\n            m = (lab2d == L)\n            if m.sum() < min_area: continue\n\n            cc2d, n2d = label_2d_8(m)\n            if n2d >= 2:\n                # 過濾掉太小的 component，避免雜訊干擾連接\n                areas = np.array([(cc2d == k).sum() for k in range(1, n2d + 1)])\n                keep = np.where(areas >= min_area)[0] + 1\n                if len(keep) >= 2:\n                    splits.append((L, cc2d, keep))\n        return splits\n\n    def connect_components_by_nearest_points(mask2d: np.ndarray, cc2d: np.ndarray, comp_ids, line_thickness: int = 1):\n        \"\"\"\n        在同一張 2D slice 中，針對指定的多塊 component 找最近點連線。\n        \"\"\"\n        comps = [np.argwhere(cc2d == cid) for cid in comp_ids]\n        if len(comps) < 2:\n            return mask2d, None\n\n        # 暴力找最近兩塊 (若點非常多可考慮 KDTree 優化，但在分割圖上通常還好)\n        best = None\n        best_pts = None\n\n        # 簡化：只連接最近的一對，避免過度連接\n        # 如果希望串聯所有斷開部分，需要改為 Minimum Spanning Tree 邏輯，但這裡先維持你原本的邏輯\n        for i in range(len(comps)):\n            for j in range(i + 1, len(comps)):\n                A, B = comps[i], comps[j]\n                # 計算兩兩距離矩陣\n                d2 = ((A[:, None, :] - B[None, :, :]) ** 2).sum(axis=2)\n                idx = np.unravel_index(np.argmin(d2), d2.shape)\n                dist = np.sqrt(d2[idx])\n\n                if best is None or dist < best:\n                    best = dist\n                    best_pts = (tuple(A[idx[0]]), tuple(B[idx[1]]))\n\n        if best_pts is None: return mask2d, None\n\n        (r0, c0), (r1, c1) = best_pts\n        rr, cc = line(r0, c0, r1, c1)\n\n        # 邊界檢查，防止 skimage.draw.line 超出範圍 (雖然理論上不應發生)\n        valid = (rr >= 0) & (rr < mask2d.shape[0]) & (cc >= 0) & (cc < mask2d.shape[1])\n        rr, cc = rr[valid], cc[valid]\n\n        out = mask2d.copy()\n        out[rr, cc] = True\n\n        if line_thickness >= 2:\n            line_only = np.zeros_like(mask2d, dtype=bool)\n            line_only[rr, cc] = True\n            structure = np.ones((3, 3), dtype=bool)\n            # 膨脹線條使其變粗\n            line_fat = ndi.binary_dilation(line_only, structure=structure, iterations=line_thickness - 1)\n            out |= line_fat\n\n        return out, (best, best_pts)\n\n    mask = mask3d.astype(bool).copy()\n\n    for _ in range(pass_iters):\n        cc3d, n = label_3d_26(mask)\n        changed = 0\n\n        for z in range(mask.shape[0]):\n            splits = find_splits_in_slice(cc3d, z, min_area=min_area)\n            for (L, cc2d, comp_ids) in splits:\n                repaired2d, info = connect_components_by_nearest_points(\n                    mask[z], cc2d, comp_ids, line_thickness=line_thickness\n                )\n                if info is None: continue\n\n                dist, _ = info\n                if dist <= max_link_dist:\n                    mask[z] = repaired2d\n                    changed += 1\n\n        if changed == 0:\n            break\n\n    return mask\n\n\n# ==========================================\n# 2. 現有後處理函數 (Auxiliary Functions)\n# ==========================================\n\n\ndef normalize_segments_3d(vol_bool: np.ndarray, radius: float = 1.5, border: int = 8, axis: int = 0, sigma: float = 1.0,\n                          repair: bool = False, max_dist: int = 30, max_angle_deg: int = 30,\n                          safe_distance: int = 4, repair_dijkstra: bool = False, iterations: int = 1) -> np.ndarray:\n    \"\"\"\n    將 3D 體積切成 8 塊後，對每塊分別進行線段正規化。\n    axis: 0 為 Z 軸, 1 為 Y 軸, 2 為 X 軸\n    \"\"\"\n\n    normalized_vol = vol_bool.copy()\n    skel_vol = np.zeros_like(vol_bool)\n    labeled_vol, num_features = ndi.label(normalized_vol, structure=ndi.generate_binary_structure(rank=3, connectivity=1))\n    slices_list = _octant_slices(normalized_vol.shape, splits=(2, 2, 2))\n\n    def get_2d_endpoints(skeleton):\n        # 建立一個中間為 0，周圍為 1 的 kernel\n        kernel = np.array([[1, 1, 1],\n                           [1, 0, 1],\n                           [1, 1, 1]])\n        neighbor_count = convolve(skeleton.astype(int), kernel, mode='constant', cval=0)\n\n        # 骨架點且鄰居只有 1 個\n        endpoints = np.argwhere((skeleton == 1) & (neighbor_count == 1))\n        return endpoints\n\n    def trace_back_and_get_vector_2d(skeleton, endpoint, steps=3):\n        \"\"\"\n        從端點沿著骨架往回走指定的步數，計算生長方向向量。\n        \"\"\"\n        current_pt = endpoint\n        visited = {tuple(endpoint)}\n\n        for _ in range(steps):\n            y, x = current_pt\n            y_min, y_max = max(0, y - 1), min(skeleton.shape[0], y + 2)\n            x_min, x_max = max(0, x - 1), min(skeleton.shape[1], x + 2)\n\n            neighborhood = skeleton[y_min:y_max, x_min:x_max]\n            neighbors = np.argwhere(neighborhood == 1)\n\n            # 轉換為全局座標\n            neighbors += np.array([y_min, x_min])\n\n            next_pt = None\n            for n in neighbors:\n                if tuple(n) not in visited:\n                    next_pt = n\n                    break\n\n            if next_pt is None:\n                break  # 骨架太短，提早走到盡頭\n\n            visited.add(tuple(next_pt))\n            current_pt = next_pt\n\n        # 向量方向：從回溯點指向端點 (即未來的生長趨勢)\n        vec = endpoint - current_pt\n        norm = np.linalg.norm(vec)\n        if norm == 0:\n            return np.array([0.0, 0.0])\n        return vec / norm\n\n    def match_endpoints_2d(endpoints, vectors, label_slice, max_dist=30, max_angle_deg=45):\n        \"\"\"\n        根據距離和方向夾角配對端點\n        \"\"\"\n        matched_pairs = []\n        used_indices = set()\n        max_angle_rad = math.radians(max_angle_deg)\n\n        for i in range(len(endpoints)):\n            if i in used_indices: continue\n\n            best_match = -1\n            min_dist = max_dist\n\n            for j in range(len(endpoints)):\n                if i == j or j in used_indices: continue\n\n                p_A, v_A = endpoints[i], vectors[i]\n                p_B, v_B = endpoints[j], vectors[j]\n\n                label_A = label_slice[p_A[0], p_A[1]]\n                label_B = label_slice[p_B[0], p_B[1]]\n\n                if label_A != label_B:\n                    continue\n\n                # 1. 檢查距離\n                dist = np.linalg.norm(p_A - p_B)\n                if dist > max_dist: continue\n\n                # 2. 檢查方向性\n                vec_AB = (p_B - p_A) / dist\n                vec_BA = -vec_AB\n\n                # 確保向量長度不為 0 才計算夾角\n                if np.linalg.norm(v_A) == 0 or np.linalg.norm(v_B) == 0: continue\n\n                dot_A_forward = np.dot(v_A, vec_AB)\n                dot_B_forward = np.dot(v_B, vec_BA)\n\n                # 2. 檢查兩個端點本身的生長向量是否「相向」(例如夾角大於 135 度)\n                # v_B 應該要和 v_A 大致反向，所以 v_A 和 -v_B 應該要大致同向\n                cos_dirs = np.clip(np.dot(v_A, -v_B), -1.0, 1.0)\n                dirs_aligned = math.acos(cos_dirs) < max_angle_rad  # 這裡的 max_angle_rad 可以設為 45 度 (即允許 45 度的方向誤差)\n\n                if dot_A_forward > 0 and dot_B_forward > 0 and dirs_aligned:\n                    if dist < min_dist:\n                        min_dist = dist\n                        best_match = j\n\n            if best_match != -1:\n                matched_pairs.append({\n                    'pA': endpoints[i], 'pB': endpoints[best_match],\n                    'vA': vectors[i], 'vB': vectors[best_match]\n                })\n                used_indices.add(i)\n                used_indices.add(best_match)\n\n        return matched_pairs\n\n    def connect_skeleton_holes_2d(skeleton, label_slice, max_dist=30, max_angle_deg=45, trace_steps=5):\n        \"\"\"\n        主函數：輸入 2D 骨架矩陣，輸出修補好的骨架矩陣\n        \"\"\"\n        # 複製一份準備畫線用\n        repaired_skeleton = skeleton.copy()\n\n        # 1. 抓取端點\n        raw_endpoints = get_2d_endpoints(skeleton)\n        if len(raw_endpoints) < 2:\n            return repaired_skeleton\n\n        # 2. 【修改處】：直接先拿掉太短的骨架，不參與配對\n        valid_endpoints = []\n        valid_vectors = []\n        for ep in raw_endpoints:\n            # 這裡的 trace_steps 可以依需求設為 3 或 5\n            vec = trace_back_and_get_vector_2d(skeleton, ep, steps=trace_steps)\n\n            # 如果 vec 不是 None，代表這條骨架夠長，才允許加入配對池\n            if vec is not None:\n                valid_endpoints.append(ep)\n                valid_vectors.append(vec)\n\n        if len(valid_endpoints) < 2:\n            return repaired_skeleton\n\n        # 3. 執行配對\n        matched_pairs = match_endpoints_2d(valid_endpoints, valid_vectors, label_slice, max_dist, max_angle_deg)\n\n        # 4. Bresenham 畫線補洞\n        for match in matched_pairs:\n            pA, pB = match['pA'], match['pB']\n            vA, vB = match['vA'], match['vB']\n            # skimage.draw.line 會回傳線上所有點的 (y, x) 座標\n            rr, cc = line(pA[0], pA[1], pB[0], pB[1])\n            repaired_skeleton[rr, cc] = 1\n\n        return repaired_skeleton\n\n    def connect_skeleton_holes_dijkstra(skeleton, slice_2d, label_slice, max_dist=30, max_angle_deg=45, trace_steps=5):\n        \"\"\"\n        主函數：利用 Dijkstra 與成本圖連接骨架\n        \"\"\"\n        repaired_skeleton = skeleton.copy()\n\n        # 1. 抓取端點\n        raw_endpoints = get_2d_endpoints(skeleton)\n        if len(raw_endpoints) < 2:\n            return repaired_skeleton\n\n        # 2. 直接先拿掉太短的骨架，不參與配對\n        valid_endpoints = []\n        valid_vectors = []\n        for ep in raw_endpoints:\n            vec = trace_back_and_get_vector_2d(skeleton, ep, steps=trace_steps)\n            if vec is not None:\n                valid_endpoints.append(ep)\n                valid_vectors.append(vec)\n\n        if len(valid_endpoints) < 2:\n            return repaired_skeleton\n\n        # 3. 執行配對\n        matched_pairs = match_endpoints_2d(valid_endpoints, valid_vectors, label_slice, max_dist, max_angle_deg)\n        if not matched_pairs:\n            return repaired_skeleton\n\n        # ==========================================\n        # 【新增處】：利用連通集區分真正的目標區域與無關孤島\n        # ==========================================\n        # 取出當前 label 的所有區域\n        current_mask_full = (slice_2d == label_slice)\n\n        # 標記所有連通的孤島\n        labeled_mask, num_features = ndi.label(current_mask_full)\n\n        # 找出包含 valid_endpoints 的所有孤島 ID\n        target_island_ids = set()\n        for ep in valid_endpoints:\n            island_id = labeled_mask[ep[0], ep[1]]\n            if island_id != 0:\n                target_island_ids.add(island_id)\n\n        # 定義區域\n        # is_target_mask: 只有包含端點的孤島才享有最低成本\n        is_target_mask = np.isin(labeled_mask, list(target_island_ids))\n\n        # is_disconnected_same_label: 同 label 但不含端點的孤島 (視同背景處理，避免被借道)\n        is_disconnected_same_label = current_mask_full & (~is_target_mask)\n\n        is_background = (slice_2d == 0) | is_disconnected_same_label\n        is_other_mask = (~is_background) & (~is_target_mask)\n\n        # ==========================================\n\n        # 4. 建立成本圖 (Cost Map)\n        cost_map = np.ones_like(slice_2d, dtype=np.float32)\n\n        # --- A. 建立避開「其他 Mask」的護城河 ---\n        if np.any(is_other_mask):\n            dist_to_others = distance_transform_edt(~is_other_mask)\n            max_warning_penalty = 15.0\n\n            warning_penalty = np.clip(max_warning_penalty * (1 - dist_to_others / safe_distance), 0,\n                                      max_warning_penalty)\n            cost_map += warning_penalty\n\n        # --- B. 設定各地形基礎成本 ---\n\n        # 1. 自己的目標 Mask 內部：保持最低成本 1.0\n        # (cost_map[is_target_mask] 維持原樣)\n\n        # 2. 黑色背景與「無關的同 label 孤島」：給予基礎懲罰\n        bg_penalty = 2.0  # 建議稍微提高背景懲罰，強迫走最短直線\n        cost_map[is_background] = bg_penalty\n\n        # 3. 其他 Mask：絕對高牆\n        cost_map[is_other_mask] = 1e6\n\n        # 5. 使用 Dijkstra 畫線補洞\n        for match in matched_pairs:\n            pA, pB = tuple(match['pA']), tuple(match['pB'])\n\n            try:\n                path, cost = route_through_array(cost_map, pA, pB, fully_connected=True)\n\n                if cost >= max_dist * bg_penalty * 2:  # 注意這裡的閾值可能需要配合 bg_penalty 調整\n                    continue\n\n                path_y = [p[0] for p in path]\n                path_x = [p[1] for p in path]\n\n                repaired_skeleton[path_y, path_x] = 1\n            except ValueError:\n                continue\n\n        return repaired_skeleton\n\n    def _process_block(block: np.ndarray, labeled_block: np.ndarray, radius: float, border: int,\n                       axis: int, sigma: float = 0.0) -> tuple[np.ndarray, np.ndarray]:\n\n        block_out = block.copy()\n        skel_out = block.copy()\n        shape = block.shape\n\n        # 1. 檢查非處理軸的其餘兩個維度是否夠大\n        other_dims = [shape[i] for i in range(3) if i != axis]\n        if border > 0 and any(d <= 2 * border for d in other_dims):\n            return block_out, skel_out  # 修正：確保提早結束時也回傳兩個陣列\n\n        # 2. 定義統一的內部區域切片邊界\n        inner_slice = slice(border, -border if border > 0 else None)\n\n        # 3. 沿著指定的 axis 進行迭代\n        for i in range(shape[axis]):\n            # 動態生成 3D 讀取切片 (例如 axis=1, i=5 時，等同於 [:, 5, :])\n            read_idx = [slice(None)] * 3\n            read_idx[axis] = i\n            read_idx = tuple(read_idx)\n\n            slice_2d = block[read_idx]\n            if not np.any(slice_2d):\n                continue\n\n            label_slice_2d = labeled_block[read_idx]\n            skel = skeletonize(slice_2d, method='lee')\n            if repair:\n                for it in range(iterations):\n                    if repair_dijkstra:\n                        skel = connect_skeleton_holes_dijkstra(skeleton=skel, slice_2d=slice_2d,\n                                                               label_slice=label_slice_2d,\n                                                               max_dist=max_dist, max_angle_deg=max_angle_deg)\n                    else:\n                        skel = connect_skeleton_holes_2d(skeleton=skel, label_slice=label_slice_2d, max_dist=max_dist,\n                                                         max_angle_deg=max_angle_deg)\n\n            if np.any(skel):\n                # 距離變換擴張\n                reconstructed = distance_transform_edt(~skel) <= radius\n\n                # 加入 2D 高斯模糊\n                if sigma > 0:\n                    reconstructed = gaussian_filter(reconstructed.astype(float), sigma=sigma)\n                    reconstructed = reconstructed > 0.3\n\n                # 動態生成 3D 寫回切片 (保留 border，並指定當前層 i)\n                write_idx = [inner_slice] * 3\n                write_idx[axis] = i\n                write_idx = tuple(write_idx)\n\n                # reconstructed 是 2D，我們只需要擷取它內部的部分來填入 3D 結構\n                recon_idx = tuple([inner_slice] * 2)\n\n                block_out[write_idx] = reconstructed[recon_idx]\n                skel_out[write_idx] = skel[recon_idx]\n\n        return block_out, skel_out\n\n    # 疊代處理 8 個子塊\n    for sl in slices_list:\n        sub_vol = vol_bool[sl]\n        labeled_sub_vol = labeled_vol[sl]\n        processed_sub_vol, skel_sub_vol = _process_block(sub_vol, labeled_sub_vol, radius, border, axis)\n        normalized_vol[sl] = processed_sub_vol\n        skel_vol[sl] = skel_sub_vol\n\n    return normalized_vol, skel_vol\n\n\n\ndef normalize_segments_3d_normal(vol_bool: np.ndarray, radius: float = 1.5, border: int = 8) -> np.ndarray:\n    \"\"\"\n        逐 Z 軸切片進行線段正規化：先骨架化再擴張。\n\n        參數:\n        - vol_bool: 輸入的 3D 布林陣列 (D, H, W)\n        - radius: 擴張半徑，用於統一線段粗細\n        \"\"\"\n    # 建立輸出的容器，預設為全 False (或根據需求複製原圖)\n    normalized_vol = np.zeros_like(vol_bool, dtype=bool)\n    depth = vol_bool.shape[0]\n\n    for z in range(depth):\n        slice_2d = vol_bool[z, :, :]\n\n        # 如果該層沒有任何像素，直接跳過\n        if not np.any(slice_2d):\n            continue\n\n        # 1. 骨架化：將線條縮減為 1 像素寬度的中心線\n        # 注意：輸入必須是布林值\n        skel = skeletonize(slice_2d)\n\n        # 2. 擴張 (使用距離變換實現精準半徑控制)\n        if np.any(skel):\n            # 計算每個背景點到最近骨架點的距離\n            dist_map = distance_transform_edt(~skel)\n            # 距離在半徑內的點即為新的線段區域\n            reconstructed_slice = dist_map <= radius\n\n            # 3. 存入結果\n            normalized_vol[z, :, :] = reconstructed_slice\n\n    return normalized_vol\n\n\ndef apply_y_axis_closing(vol, iterations=1):\n    \"\"\"Y 軸方向閉運算\"\"\"\n    processed_vol = np.copy(vol)\n    for y in range(vol.shape[1]):\n        slice_2d = vol[:, y, :]\n        if np.any(slice_2d):\n            closed_slice = binary_closing(slice_2d, iterations=iterations)\n            processed_vol[:, y, :] = closed_slice\n    return processed_vol\n\n\ndef apply_z_axis_fill_holes(vol):\n    \"\"\"Z 軸方向孔洞填充\"\"\"\n    processed_vol = np.copy(vol)\n    for z in range(vol.shape[0]):\n        slice_2d = vol[z, :, :]\n        if np.any(slice_2d):\n            filled_slice = binary_fill_holes(slice_2d)\n            processed_vol[z, :, :] = filled_slice\n    return processed_vol\n\n\ndef apply_y_axis_fill_holes(vol):\n    \"\"\"Y 軸方向孔洞填充\"\"\"\n    processed_vol = np.copy(vol)\n    for y in range(vol.shape[1]):\n        slice_2d = vol[:, y, :]\n        if np.any(slice_2d):\n            filled_slice = binary_fill_holes(slice_2d)\n            processed_vol[:, y, :] = filled_slice\n    return processed_vol\n\n\nimport numpy as np\n\n\ndef robust_mask_sandwich(mask, max_gap=2, axis_weights=[0, 1, 1], min_neighbors=2):\n    \"\"\"\n    改良版 Mask Sandwich (加入 has_support 檢查):\n\n    Args:\n        mask: 3D Binary Mask\n        max_gap: 最大填補間隙\n        axis_weights: [z, y, x] 開關。\n        min_neighbors: 一個像素周圍 (3x3範圍內) 至少要有幾個鄰居才算「結實」。\n                       建議設為 2 或 3。如果設為 0 則等於沒檢查。\n    \"\"\"\n    refined = mask.copy()\n    ndim = 3\n\n    # 定義檢查函數：輸入一個 2D slice，回傳「結實像素」的 Mask\n    for axis in range(ndim):\n        if axis_weights[axis] == 0:\n            continue\n\n        current_max_gap = max_gap\n\n        for gap in range(1, current_max_gap + 1):\n            stride = gap + 1\n\n            s_prev = [slice(None)] * ndim\n            s_next = [slice(None)] * ndim\n\n            s_prev[axis] = slice(0, -stride)\n            s_next[axis] = slice(stride, None)\n\n            # 1. 取出兩端切片\n            slice_prev = refined[tuple(s_prev)]\n            slice_next = refined[tuple(s_next)]\n\n            # 2. 【核心修改】檢查支撐性 (Support Check)\n            # 這裡我們只對「當前操作面」做 2D 檢查\n            # 注意：slice_prev 和 slice_next 可能是 3D 的 (一部分的 volume)\n            # 為了效能，我們可以批次處理或簡化處理\n\n            # 如果是處理 Z 軸，slice_prev 是 (D', H, W)，我們希望在 H,W 平面檢查\n            # 如果是處理 Y 軸，slice_prev 是 (D, H', W)，我們希望在 D,W 平面檢查 (有點怪)\n            # 但通常「噪點」定義在 3D 空間都是通用的。\n            # 為了簡化且通用，我們直接比較兩端是否重疊，\n            # 並利用 Logical AND 的特性：噪音通常不會剛好在隔壁層的同個位置\n\n            # --- 實作支撐性過濾 ---\n            # 為了避免複雜的軸向判斷，這裡做一個取捨：\n            # 我們假設輸入的 slice 已經是二值化，直接計算鄰居太慢？\n            # 不會，scipy convolve 在 CPU 上對 binary mask 很快。\n\n            # 針對不同的軸向，我們需要正確的 Kernel\n            # 為了通用性，我們在函數內動態構建 N-dim Kernel 比較慢\n            # 但考慮到 Vesuvius 的各向異性，我們主要關心 \"XY平面\" 的鄰居\n\n            if axis == 0:  # 正在修補 Z 軸間隙，檢查 XY 平面的鄰居\n                # 這裡 slice_prev 的形狀是 (D_subset, H, W)\n                # 我們可以對每一層做 2D 卷積，或者直接用 3D 卷積但 Kernel 只有 XY 有值\n\n                # 建立 3D Kernel 但只在 XY 平面擴展 (3x3x1 concept)\n                kernel_3d = np.zeros((3, 3, 3), dtype=np.uint8)\n                kernel_3d[1, :, :] = 1  # 中間層的 3x3\n                kernel_3d[1, 1, 1] = 0  # 扣掉自己\n\n                # 計算支撐\n                count_prev = convolve(slice_prev.astype(np.uint8), kernel_3d, mode='constant', cval=0)\n                count_next = convolve(slice_next.astype(np.uint8), kernel_3d, mode='constant', cval=0)\n\n                valid_prev = (slice_prev) & (count_prev >= min_neighbors)\n                valid_next = (slice_next) & (count_next >= min_neighbors)\n\n            else:\n                # 針對 Y 或 X 軸填補，通常比較少用，或者可以直接忽略支撐檢查(設寬鬆)\n                # 或者簡單一點：只要有值就算 (不做額外檢查)，因為 XY 斷裂修復通常比較安全\n                valid_prev = slice_prev\n                valid_next = slice_next\n                # 如果你想非常嚴格，也可以在這裡實作對應軸的 convolve，但代碼會變很長\n\n            # 3. 結合條件：兩端都有值 + 兩端都結實\n            bridge_candidates = valid_prev & valid_next\n\n            if not np.any(bridge_candidates):\n                continue\n\n            # 4. 執行填補\n            for i in range(1, stride):\n                s_fill = [slice(None)] * ndim\n                if -stride + i == 0:\n                    s_fill[axis] = slice(i, None)\n                else:\n                    s_fill[axis] = slice(i, -stride + i)\n\n                refined[tuple(s_fill)] |= bridge_candidates\n\n    return refined\n\n\ndef robust_mask_sandwich_v3(mask, max_gap=2, axis_weights=[1, 1, 1], min_neighbors=2, kernel_size=5):\n    \"\"\"\n    V3 優化強健版（支援擴大 Kernel）：\n    1. 預計算支撐圖 (Support Map)，大幅提升效能。\n    2. 全軸向支持：所有軸向均可進行 2D 支撐檢查。\n    3. 記憶體優化：最小化陣列複製次數。\n    4. 動態 Kernel：可透過 kernel_size 參數（需為奇數）擴大鄰居檢查範圍。\n    \"\"\"\n    if kernel_size % 2 == 0:\n        raise ValueError(\"kernel_size 必須為奇數（例如 3, 5, 7）\")\n\n    # 統一轉換為 bool 進行運算，減少記憶體壓力\n    refined = mask.astype(bool, copy=True)\n    ndim = 3\n    half_k = kernel_size // 2\n\n    # --- 預計算：每一層的支撐圖 (只算一次) ---\n    # 這樣在後續遍歷不同的 gap 時，不需要重複捲積\n    support_masks = []\n    for axis in range(ndim):\n        if axis_weights[axis] == 0:\n            support_masks.append(None)\n            continue\n\n        # 建立該軸向的 2D 鄰居統計 Kernel\n        # 若 kernel_size=5，則 k_shape 預設為 [5, 5, 5]\n        k_shape = [kernel_size] * ndim\n        k_shape[axis] = 1  # 讓 Kernel 在該軸向上是扁平的 (例如 [1, 5, 5])\n        kernel = np.ones(k_shape, dtype=np.uint8)\n\n        # 移除中心點 (不計算自己)\n        # 動態定位中心，例如 kernel_size=5，half_k=2，center=[2, 2, 2]\n        center = [half_k] * ndim\n        center[axis] = 0\n        kernel[tuple(center)] = 0\n\n        # 一次性計算整體的鄰居數量\n        neighbor_count = convolve(refined.astype(np.uint8), kernel, mode='constant', cval=0)\n\n        # 只有「自己是1」且「鄰居夠多」的才算有效錨點\n        support_masks.append((refined) & (neighbor_count >= min_neighbors))\n\n    # --- 執行填補 ---\n    for axis in range(ndim):\n        if axis_weights[axis] == 0:\n            continue\n\n        valid_map = support_masks[axis]\n        shape_at_axis = refined.shape[axis]\n\n        for gap in range(1, max_gap + 1):\n            stride = gap + 1\n\n            # 取得兩端的「強健錨點」\n            s_prev = [slice(None)] * ndim\n            s_next = [slice(None)] * ndim\n            s_prev[axis] = slice(0, -stride)\n            s_next[axis] = slice(stride, None)\n\n            # 只有兩端都是「結實像素」時，才建立橋樑\n            bridge = valid_map[tuple(s_prev)] & valid_map[tuple(s_next)]\n\n            if not np.any(bridge):\n                continue\n\n            # 填補中間空隙\n            for i in range(1, stride):\n                s_fill = [slice(None)] * ndim\n                s_fill[axis] = slice(i, i + (shape_at_axis - stride))\n                refined[tuple(s_fill)] |= bridge\n\n    return refined\n\n\ndef robust_sandwich_fill(mask, max_gap=2, fix_planes=('XZ', 'YZ'), kernel_size=3, min_neighbors=2):\n    \"\"\"\n    支援大空隙 (max_gap) 的對角線強健版填補（動態更新錨點模式）。\n    每次填補後會即時更新可用的「強健錨點」，允許產生連鎖填補效應。\n    \"\"\"\n    if kernel_size % 2 == 0:\n        raise ValueError(\"kernel_size 必須為奇數（例如 3, 5, 7）\")\n\n    # 統一轉換為 bool 進行運算，減少記憶體壓力\n    img = mask.astype(bool, copy=True)\n    refined = img.copy()\n    ndim = 3\n    shape = refined.shape\n\n    # 準備 3D 鄰居統計 Kernel (若需要判斷鄰居才會用到)\n    if min_neighbors > 0:\n        support_kernel = np.ones((kernel_size, kernel_size, kernel_size), dtype=np.uint8)\n        half_k = kernel_size // 2\n        support_kernel[half_k, half_k, half_k] = 0\n\n    # 定義要修復的平面對角線方向 (dz, dy, dx)\n    plane_map = {\n        'XZ': [(1, 0, 1), (1, 0, -1)],\n        'YZ': [(1, 1, 0), (1, -1, 0)],\n        'XY': [(0, 1, 1), (0, 1, -1)]\n    }\n\n    selected_dirs = []\n    for p in fix_planes:\n        if p in plane_map:\n            selected_dirs.extend(plane_map[p])\n\n    # 執行切片位移對角線填補\n    for dz, dy, dx in selected_dirs:\n        dir_vector = (dz, dy, dx)\n\n        for gap in range(1, max_gap + 1):\n\n            # --- 【核心修改】：動態更新錨點 ---\n            # 每次處理新方向或新間距前，根據「最新的 refined 狀態」重新計算錨點\n            if min_neighbors > 0:\n                # 重新執行卷積以獲取最新鄰居數量\n                neighbor_count = convolve(refined.astype(np.uint8), support_kernel, mode='constant', cval=0)\n                valid_anchors = refined & (neighbor_count >= min_neighbors)\n            else:\n                # 效能優化：如果 min_neighbors=0，最新的錨點就等於當前的 refined 狀態，免算卷積\n                valid_anchors = refined.copy()\n            # ----------------------------------\n\n            stride = gap + 1\n            stride_offset = tuple(d * stride for d in dir_vector)\n\n            s_base = []\n            s_target = []\n            lengths = []\n            starts = []\n\n            for a in range(ndim):\n                offset_a = stride_offset[a]\n                if offset_a > 0:\n                    s_base.append(slice(0, -offset_a))\n                    s_target.append(slice(offset_a, None))\n                    lengths.append(shape[a] - offset_a)\n                    starts.append(0)\n                elif offset_a < 0:\n                    s_base.append(slice(-offset_a, None))\n                    s_target.append(slice(0, offset_a))\n                    lengths.append(shape[a] - abs(offset_a))\n                    starts.append(-offset_a)\n                else:\n                    s_base.append(slice(None))\n                    s_target.append(slice(None))\n                    lengths.append(shape[a])\n                    starts.append(0)\n\n            # 找出兩端都是「最新強健錨點」的橋樑\n            bridge = valid_anchors[tuple(s_base)] & valid_anchors[tuple(s_target)]\n\n            if not np.any(bridge):\n                continue\n\n            # 將橋樑中間的「所有」空隙點補齊\n            for i in range(1, stride):\n                fill_offset = tuple(d * i for d in dir_vector)\n                s_fill = []\n\n                for a in range(ndim):\n                    start = starts[a] + fill_offset[a]\n                    end = start + lengths[a]\n                    s_fill.append(slice(start, end))\n\n                # 更新 refined，供下一個迴圈使用\n                refined[tuple(s_fill)] |= bridge\n\n    return refined\n\n\ndef xyz_multi_gap_fill_v2(mask, max_gap=2):\n    \"\"\"\n    優化版：使用向量化位移檢查三明治結構。\n    \"\"\"\n    refined = mask.astype(bool, copy=True)\n\n    for axis in range(3):\n        # 針對每一種可能的 gap 長度進行檢查\n        for gap in range(1, max_gap + 1):\n            stride = gap + 1\n\n            # 建立位移切片\n            # 這裡利用 numpy 的 slice 技巧，找出距離為 stride 的兩端\n            s_prev = [slice(None)] * 3\n            s_next = [slice(None)] * 3\n            s_fill_base = [slice(None)] * 3\n\n            s_prev[axis] = slice(0, -stride)\n            s_next[axis] = slice(stride, None)\n\n            # 找出兩端皆為 1 的橋樑錨點\n            bridge = refined[tuple(s_prev)] & refined[tuple(s_next)]\n\n            # 填補中間的所有像素\n            for i in range(1, stride):\n                s_fill = list(s_fill_base)\n                s_fill[axis] = slice(i, i + (refined.shape[axis] - stride))\n                refined[tuple(s_fill)] |= bridge\n\n    return refined\n\n\ndef diagonal_sandwich_fill_v2(mask, fix_planes):\n    \"\"\"\n        精確版：只針對特定平面進行對角線修復，避免過度填補。\n        \"\"\"\n    img = mask.astype(np.uint8)\n    refined = img.copy()\n\n    plane_map = {\n        'XZ': [(1, 0, 1), (1, 0, -1)],\n        'YZ': [(1, 1, 0), (1, -1, 0)],\n        'XY': [(0, 1, 1), (0, 1, -1)]\n    }\n\n    selected_directions = []\n    for p in fix_planes:\n        if p in plane_map:\n            print(f'Fixing {p} plane.')\n            selected_directions.extend(plane_map[p])\n\n    for dz, dy, dx in selected_directions:\n        kernel = np.zeros((3, 3, 3), dtype=np.uint8)\n        kernel[1 + dz, 1 + dy, 1 + dx] = 1\n        kernel[1 - dz, 1 - dy, 1 - dx] = 1\n\n        # 卷積檢查\n        counts = convolve(refined, kernel, mode='constant', cval=0)\n        refined[counts == 2] = 1\n\n    return refined > 0\n\n\ndef fill_3d_projection_holes_blocked(binary_mask):\n    \"\"\"\n    將 binary_mask 切成 8 塊 (2x2x2) 後分別進行投影補洞，最後合併。\n    \"\"\"\n\n    def process_single_block(sub_mask):\n        \"\"\"\n        處理單個 3D 子區塊的核心邏輯。\n        包含：移除小物件、骨架檢查、投影填補。\n        \"\"\"\n        # 1. 複製並過濾小物件\n        # 注意：這裡的 min_size 作用於「子區塊內的物件體積」。\n        # 若物件被切開，其局部體積可能小於 5000，會被暫時過濾掉 (但最後會通過原始 mask 合併回來)。\n        mask_filtered = remove_small_objects(sub_mask, min_size=5000, connectivity=1)\n\n        # 用於運算的 output，基於過濾後的 mask 進行修改\n        processing_mask = mask_filtered.copy()\n\n        # 2. 提取 3D 連通域\n        struct_3d = ndimage.generate_binary_structure(3, 1)\n        labeled_array, num_features = ndimage.label(mask_filtered, structure=struct_3d)\n\n        # 用於檢查 2D (ZY平面) 連接關係的結構 (4 連通)\n        struct_2d = ndimage.generate_binary_structure(2, 2)\n\n        # 若沒有大型物件，直接回傳全 false 的 processing_mask (其實就是空的)\n        if num_features == 0:\n            return processing_mask\n\n        # 3. 遍歷大型連通域\n        for i in range(1, num_features + 1):\n            slices = ndimage.find_objects(labeled_array == i)[0]\n            comp_mask = (labeled_array[slices] == i)\n\n            # ==========================================================\n            # 骨架化檢查\n            # ==========================================================\n            total_pixels = np.sum(comp_mask)\n            if total_pixels == 0: continue\n\n            width_x = comp_mask.shape[2]\n            skeleton_pixels = 0\n\n            # 針對 X 軸的每個 2D 切片做骨架化\n            for x in range(width_x):\n                current_slice_zy = comp_mask[:, :, x]\n                if np.any(current_slice_zy):\n                    skel_slice = skeletonize(current_slice_zy)\n                    skeleton_pixels += np.sum(skel_slice)\n\n            ratio = skeleton_pixels / total_pixels\n            if ratio > 0.1:  # 骨架佔比過高，跳過\n                continue\n            # ==========================================================\n\n            # 4. 投影與填補邏輯\n            projection_zy = np.any(comp_mask, axis=2)\n            filled_projection = ndimage.binary_fill_holes(projection_zy)\n            hole_mask_zy = filled_projection ^ projection_zy\n\n            if not np.any(hole_mask_zy):\n                continue\n\n            hole_dilated = ndimage.binary_dilation(hole_mask_zy, structure=struct_2d)\n\n            x_start = slices[2].start\n            for x in range(width_x):\n                current_slice_zy = comp_mask[:, :, x]\n                is_connected = np.any(hole_dilated & current_slice_zy)\n\n                if is_connected:\n                    # 修改該區塊的 processing_mask\n                    target_slice = processing_mask[slices[0], slices[1], x_start + x]\n                    target_slice[hole_dilated] = True\n\n        return processing_mask\n\n    # 確保輸入是布林值\n    mask = binary_mask.astype(bool)\n\n    # 1. 取得切片清單 (2x2x2 = 8塊)\n    splits = (2, 2, 2)\n    slices_list = _octant_slices(mask.shape, splits)\n\n    # 建立一個全域的空陣列，用來存放處理過(填補後)的大物件\n    processed_global = np.zeros_like(mask)\n\n    print(f\"開始分塊處理：將影像切分為 {len(slices_list)} 塊...\")\n\n    # 2. 迭代每個區塊\n    for idx, (sz, sy, sx) in enumerate(slices_list):\n        # 取出子區塊\n        sub_mask = mask[sz, sy, sx]\n\n        # 如果子區塊全空，跳過運算\n        if not np.any(sub_mask):\n            continue\n\n        # 執行核心補洞邏輯\n        sub_result = process_single_block(sub_mask)\n\n        # 將結果放回全域陣列\n        processed_global[sz, sy, sx] = sub_result\n\n    # 3. 合併結果\n    # final_output = 原始小物件 (從原始 mask 保留) | 處理後的大物件與填補 (從 processed_global 取得)\n    final_output = mask | processed_global\n\n    return final_output\n\n\ndef compute_separation_constraint(mask):\n    \"\"\"\n    計算連通域間的不可侵犯領域與邊界。\n\n    規則：\n    1. 定義物件：使用 6 連通 (Face-connected) 區分物件。\n    2. 劃分領土：使用 Voronoi (EDT) 擴張，填滿背景。\n    3. 定義邊界：若體素的 26 連通鄰域 (3x3x3) 內包含 \">=2 個不同的連通域 ID\"，則視為邊界。\n    \"\"\"\n\n    # -----------------------------------------------------------\n    # 1. 計算 3D 連通域 (Labeling) - 嚴格區分對角線接觸\n    # -----------------------------------------------------------\n    mask = mask.copy()\n    mask = clear_boundary_faces(mask, margin=3)\n\n    mask = remove_small_objects(mask > 0, 5000)\n    # 使用 6 連通 (connectivity=1)，確保僅在對角線接觸的物件被視為不同個體\n    structure_6conn = ndimage.generate_binary_structure(rank=3, connectivity=1)\n    labeled_array, num_features = ndimage.label(mask, structure=structure_6conn)\n\n    # 如果全場只有 0 或 1 個物件，不存在\"交界\"，直接回傳全 True\n    if num_features <= 1:\n        return np.ones_like(mask, dtype=bool)\n\n    # -----------------------------------------------------------\n    # 2. 擴張勢力範圍 (Voronoi / EDT)\n    # -----------------------------------------------------------\n    # 我們需要一個\"填滿\"的空間圖，每個點都知道自己最近的物件是誰 (Territory ID)\n    # 對\"背景 (0)\"做距離變換，找到最近的前景索引\n    # indices shape: (3, D, H, W)\n    _, indices = ndimage.distance_transform_edt(labeled_array == 0, return_indices=True, return_distances=True)\n\n    # 映射回 Label ID，得到全空間的領土圖\n    # 這裡 territory 不再有 0 (背景)，全都是 1~N 的 ID\n    territory = labeled_array[indices[0], indices[1], indices[2]]\n\n    # -----------------------------------------------------------\n    # 3. 標記交界處 (26-Connectivity Multi-Label Detection)\n    # -----------------------------------------------------------\n    # 規則：如果一個體素的 26 連通鄰域內，包含 2 個以上的不同 ID，它就是邊界。\n    # 數學實作：在 territory map 上，若 Max(鄰域) != Min(鄰域)，則必有至少兩個不同 ID。\n\n    footprint_26conn = np.ones((3, 3, 3), dtype=int)  # 26 連通核心\n\n    # 找出鄰域內的最大 ID 與最小 ID\n    max_labels = ndimage.maximum_filter(territory, footprint=footprint_26conn)\n    min_labels = ndimage.minimum_filter(territory, footprint=footprint_26conn)\n\n    # 邊界判定：只要最大值不等於最小值，代表鄰域內混雜了不同的領土\n    # 這會偵測到：\n    # 1. 兩個物件擴張後的接觸面 (Voronoi 邊界)\n    # 2. 兩個原始物件物理上靠得很近或接觸的地方 (Original Contact)\n    boundary_mask = (max_labels != min_labels)\n\n    # -----------------------------------------------------------\n    # 4. 取得最終 Valid Mask\n    # -----------------------------------------------------------\n    # 只要不是邊界，就是合法填充區\n    valid_mask = ~boundary_mask\n\n    return valid_mask\n\n\ndef robust_mask_refine(mask, max_gap=2, iterations=1, min_neighbors=2, kernel_size=5):\n    \"\"\"\n    包含分離保護的綜合修復流程。\n    \"\"\"\n    # --- 前處理：計算不可侵犯領域 ---\n    # print(\"Calculating separation constraints...\")\n    # valid_constraint = compute_separation_constraint(mask)\n\n    res = mask.copy()\n\n    # 開始修復迴圈\n    for i in range(iterations):\n        # 假設這是您的修復函數 (需確保您有定義這些函數)\n        # res = xyz_multi_gap_fill_v2(res, max_gap=max_gap)\n\n        # 1. 執行填充/修復\n        # 注意：這裡假設您的 sandwich 或 diagonal 函數會讓 mask 變大\n        res = robust_mask_sandwich_v3(res, max_gap=max_gap, min_neighbors=min_neighbors, kernel_size=kernel_size)\n        res = robust_sandwich_fill(res, max_gap=1, min_neighbors=0, kernel_size=3,\n                                   fix_planes=('XZ', 'YZ'))\n\n        # 2. 應用不可侵犯領域約束\n        # 強制切斷跨越連通域邊界的填充\n    # res = res & valid_constraint\n    return res\n\n\nimport numpy as np\n\n\ndef robust_mask_refine_8slice(mask, max_gap=2, iterations=1, min_neighbors=2, splits=(2, 2, 2)):\n    \"\"\"\n    將 3D mask 切分為多個區塊（預設 2x2x2=8 份）平行或依序處理，\n    並處理邊界重疊以保持修復的連續性。\n\n    Args:\n        mask (np.ndarray): 原始 3D Boolean/Int Mask。\n        max_gap (int): 最大填補間隙 (傳遞給 robust_mask_refine)。\n        iterations (int): 迭代次數 (傳遞給 robust_mask_refine)。\n        min_neighbors (int): 最小鄰居數 (傳遞給 robust_mask_refine)。\n        splits (tuple): (z_parts, y_parts, x_parts) 指定各軸切分數量，預設為 (2, 2, 2) 即 8 等份。\n\n    Returns:\n        np.ndarray: 修復完成的完整 Mask。\n    \"\"\"\n    # 1. 準備輸出的容器\n    refined_full = np.zeros_like(mask)\n\n    # 2. 取得基礎切片 (根據 splits 數量切分)\n    # 假設 _octant_slices 傳回 [(slice_z, slice_y, slice_x), ...]\n    base_slices = _octant_slices(mask.shape, splits)\n\n    for sl_z, sl_y, sl_x in base_slices:\n        # --- A. 直接取出子區塊 (無 Padding) ---\n        sub_mask = mask[sl_z, sl_y, sl_x]\n\n        # --- B. 執行核心修復演算法 ---\n        # 直接在子區塊上運算\n        refined_sub = robust_mask_refine(\n            sub_mask,\n            max_gap=max_gap,\n            iterations=iterations,\n            min_neighbors=min_neighbors\n        )\n\n        # --- C. 直接填回結果陣列 ---\n        # 因為沒有 Padding，不需要計算 offset 或裁切\n        refined_full[sl_z, sl_y, sl_x] = refined_sub\n\n    return refined_full\n\n\n# def robust_mask_refine(mask, max_gap=2, max_iterations=15, min_neighbors=2):\n#     \"\"\"\n#     綜合修復流程：結合軸向與對角線填充。\n#     機制：若 res 經過處理後沒有變化就停止，否則最大執行 max_iterations 次。\n#     輸出：會列印最終執行的次數。\n#     \"\"\"\n#     res = mask.copy()\n#\n#     for i in range(max_iterations):\n#         prev_res = res.copy()\n#\n#         # 執行修復操作\n#         # res = xyz_multi_gap_fill_v2(res, max_gap=max_gap)\n#         res = robust_mask_sandwich_v3(res, max_gap=max_gap, min_neighbors=min_neighbors)\n#         res = diagonal_sandwich_fill_v2(res, fix_planes=('XZ', 'YZ'))\n#\n#         # 檢查是否收斂：若處理前後結果相同，則提早停止\n#         if np.array_equal(res, prev_res):\n#             # i 是從 0 開始，所以實際次數是 i + 1\n#             print(f\"處理已收斂，共執行了 {i + 1} 次。\")\n#             break\n#     else:\n#         # Python 的 for-else 語法：只有當迴圈「沒有」被 break 中斷（即跑滿 max_iterations）時才會執行\n#         print(f\"已達到最大設定次數，共執行了 {max_iterations} 次。\")\n#\n#     return res\n\ndef _axis_cuts(n: int, parts: int) -> List[Tuple[int, int]]:\n    \"\"\"將維度 n 切分為 parts 份，處理餘數確保覆蓋全域。\"\"\"\n    parts = max(1, min(parts, n if n > 0 else 1))\n    base = n // parts\n    extra = n % parts\n    cuts = []\n    start = 0\n    for i in range(parts):\n        size = base + (1 if i < extra else 0)\n        end = start + size\n        cuts.append((start, end))\n        start = end\n    return cuts\n\n\ndef _octant_slices(shape: Tuple[int, int, int], splits: Tuple[int, int, int]) -> List[Tuple[slice, slice, slice]]:\n    \"\"\"生成 3D 切片物件清單。\"\"\"\n    zcuts = _axis_cuts(shape[0], splits[0])\n    ycuts = _axis_cuts(shape[1], splits[1])\n    xcuts = _axis_cuts(shape[2], splits[2])\n    out = []\n    for z0, z1 in zcuts:\n        for y0, y1 in ycuts:\n            for x0, x1 in xcuts:\n                if (z1 - z0) > 0 and (y1 - y0) > 0 and (x1 - x0) > 0:\n                    out.append((slice(z0, z1), slice(y0, y1), slice(x0, x1)))\n    return out\n\n\ndef remove_small_8slice(vol, min_size=100, connectivity=1):\n    \"\"\"使用 8-slice 邏輯平行化/分段移除小物件。\"\"\"\n    vol_bool = vol > 0\n    output = np.zeros_like(vol_bool)\n\n    # 取得 8 個區塊的 slice (2x2x2)\n    slices = _octant_slices(vol_bool.shape, (2, 2, 2))\n    for slc in slices:\n        block = vol_bool[slc]\n        if np.any(block):\n            # 處理區塊並塞回對應位置\n            output[slc] = remove_small_objects(\n                block,\n                min_size=min_size,\n                connectivity=connectivity\n            )\n\n    return output.astype(np.uint8)\n\n\nimport numpy as np\nimport torch\nimport scipy.ndimage as ndimage\n\n\ndef reconstruct_surface_from_mask(\n        mask_3d: np.ndarray,\n        degree: int = 6,\n        min_pixels: int = 20,\n        device: str = 'cpu'\n) -> np.ndarray:\n    \"\"\"\n    輸入一個 3D Binary Mask，對每個 6-連通區域進行多項式曲面擬合重建。\n\n    參數:\n    - mask_3d: 輸入的 3D numpy array (0 或 1)\n    - degree: 多項式的階數 (建議 4-8)\n    - min_pixels: 忽略小於此像素數量的連通域\n    - device: 'cpu' 或 'cuda'\n\n    回傳:\n    - output_mask: 重建後的 3D mask (僅包含擬合出的曲面，非實心體積)\n    \"\"\"\n\n    # 1. 定義 6-連通結構 (3x3x3)\n    # structure 設為 1 代表只有上下左右前後相連 (6-connectivity)\n    s = ndimage.generate_binary_structure(3, 1)\n\n    # 2. 標記連通域\n    labeled_array, num_features = ndimage.label(mask_3d, structure=s)\n\n    # 準備輸出容器\n    output_mask = np.zeros_like(mask_3d)\n\n    if num_features == 0:\n        return output_mask\n\n    print(f\"檢測到 {num_features} 個連通域，開始處理...\")\n\n    # 3. 逐一處理每個連通域\n    for label_id in range(1, num_features + 1):\n        # 提取當前組件的點雲\n        component_mask = (labeled_array == label_id)\n        points_zyx = np.argwhere(component_mask)  # 格式: [z, y, x]\n\n        # 忽略過小的噪點\n        if len(points_zyx) < min_pixels:\n            continue\n\n        # 呼叫擬合核心函數\n        reconstructed_points = _fit_poly_component(points_zyx, degree, device)\n\n        # 4. 將重建的點雲填回輸出 Mask\n        if reconstructed_points is not None:\n            # 過濾掉超出原始邊界的點\n            D, H, W = mask_3d.shape\n            z, y, x = reconstructed_points[:, 0], reconstructed_points[:, 1], reconstructed_points[:, 2]\n\n            valid_mask = (\n                    (z >= 0) & (z < D) &\n                    (y >= 0) & (y < H) &\n                    (x >= 0) & (x < W)\n            )\n\n            valid_z = z[valid_mask].astype(int)\n            valid_y = y[valid_mask].astype(int)\n            valid_x = x[valid_mask].astype(int)\n\n            # 填入輸出 (設為 1)\n            output_mask[valid_z, valid_y, valid_x] = 1\n\n    return output_mask\n\n\ndef _fit_poly_component(points_np, degree, device):\n    \"\"\"\n    單一組件的 PCA + 多項式擬合核心邏輯 (基於 hengck23 的思路)\n    \"\"\"\n    try:\n        points = torch.tensor(points_np, dtype=torch.float64, device=device)\n\n        # --- 1. PCA 座標轉換 (Alignment) ---\n        mean = points.mean(dim=0)\n        centered = points - mean\n\n        # 使用 PCA 找到主軸 (SVD 分解)\n        # V 的最後一列通常是法向量方向 (變異最小軸)，我們將其視為新的 Z 軸\n        U, S, V = torch.pca_lowrank(centered, q=3)\n\n        # 旋轉到 PCA 空間: [x', y', z']\n        # 注意: 這裡我們假設前兩個主成分是展延面 (x, y)，第三個是高度 (z)\n        pca_points = centered @ V\n\n        x_pca = pca_points[:, 0]\n        y_pca = pca_points[:, 1]\n        z_pca = pca_points[:, 2]  # 這是我們要擬合的目標高度\n\n        # 歸一化以避免數值不穩定\n        x_scale = x_pca.abs().max() + 1e-6\n        y_scale = y_pca.abs().max() + 1e-6\n\n        # --- 2. 建構多項式矩陣 (Vandermonde Matrix) ---\n        # 擬合目標: z_pca = Poly(x_pca, y_pca)\n        A_list = []\n        for i in range(degree + 1):\n            for j in range(degree + 1 - i):\n                term = ((x_pca / x_scale) ** i) * ((y_pca / y_scale) ** j)\n                A_list.append(term)\n\n        A = torch.stack(A_list, dim=1)\n\n        # --- 3. 求解係數 (Least Squares) ---\n        # Ridge Regression (加上 lambda 避免奇異矩陣)\n        lam = 1e-3\n        I = torch.eye(A.shape[1], device=device, dtype=torch.float64)\n        coeffs = torch.linalg.solve(A.T @ A + lam * I, A.T @ z_pca)\n\n        # --- 4. 生成重建網格 ---\n        # 我們在 PCA 的 XY 平面上生成網格來重建曲面\n        # grid_size 取決於該組件的投影大小，這裡動態計算\n        min_x, max_x = x_pca.min(), x_pca.max()\n        min_y, max_y = y_pca.min(), y_pca.max()\n\n        # 密度設為 1.0 (每個像素採樣一次)\n        step = 0.8  # 稍微密一點可以填補空隙\n        grid_x_range = torch.arange(min_x, max_x + step, step, device=device, dtype=torch.float64)\n        grid_y_range = torch.arange(min_y, max_y + step, step, device=device, dtype=torch.float64)\n\n        grid_x, grid_y = torch.meshgrid(grid_x_range, grid_y_range, indexing='ij')\n        grid_x_flat = grid_x.flatten()\n        grid_y_flat = grid_y.flatten()\n\n        # 計算擬合後的 Z 值\n        A_grid_list = []\n        for i in range(degree + 1):\n            for j in range(degree + 1 - i):\n                term = ((grid_x_flat / x_scale) ** i) * ((grid_y_flat / y_scale) ** j)\n                A_grid_list.append(term)\n\n        # 預測 Z\n        pred_z_flat = (torch.stack(A_grid_list, dim=1) @ coeffs)\n\n        # --- 5. 轉回原始空間 ---\n        # 組合: [x', y', z_pred]\n        reconstructed_pca = torch.stack([grid_x_flat, grid_y_flat, pred_z_flat], dim=-1)\n\n        # 逆旋轉 + 加回均值\n        reconstructed_global = (reconstructed_pca @ V.T) + mean\n\n        return reconstructed_global.cpu().numpy()\n\n    except Exception as e:\n        print(f\"擬合失敗，跳過此組件: {e}\")\n        return None\n\n\ndef fill_hole_8slice(vol):\n    \"\"\"使用 8-slice 邏輯分段執行 binary_fill_holes。\"\"\"\n    vol_bool = vol > 0\n    output = np.zeros_like(vol_bool)\n\n    slices = _octant_slices(vol_bool.shape, (2, 2, 2))\n    struct_26 = ndimage.generate_binary_structure(3, 3)\n\n    for slc in slices:\n        block = vol_bool[slc]\n        if np.any(block):\n            output[slc] = binary_fill_holes(block, structure=struct_26)\n\n    return output.astype(np.uint8)\n\n\nimport numpy as np\n\n\ndef clear_boundary_faces(vol: np.ndarray, margin: int = 3) -> np.ndarray:\n    \"\"\"\n    將 3D Volume 六個面的邊界區域 (距離邊緣 <= margin) 設為 0。\n\n    Args:\n        vol (np.ndarray): 輸入的 3D 陣列 (Z, Y, X)。\n        margin (int): 邊界寬度，預設為 2 (即 <= 2vx)。\n        inplace (bool): 是否直接修改原陣列以節省記憶體。預設 False。\n\n    Returns:\n        np.ndarray: 處理後的陣列。\n    \"\"\"\n    vol = vol.copy()\n    vol = vol > 0\n    # 確保 vol 是 3D\n    if vol.ndim != 3:\n        raise ValueError(f\"Input volume must be 3D, but got shape {vol.shape}\")\n\n    # 1. Z 軸邊界 (上下)\n    vol[:margin, :, :] = False\n    vol[-margin:, :, :] = False\n\n    # 2. Y 軸邊界 (前後)\n    vol[:, :margin, :] = False\n    vol[:, -margin:, :] = False\n\n    # 3. X 軸邊界 (左右)\n    vol[:, :, :margin] = False\n    vol[:, :, -margin:] = False\n    return vol.astype(np.uint8)\n\n\ndef fill_hole_8slice_keep_one_per_slice(vol):\n    \"\"\"\n        1. 複製 vol 並移除小於 min_size (5000) 的物件 -> clean_vol\n        2. 在 clean_vol 上執行 8-slice 分塊邏輯：\n           - 找出所有洞\n           - 計算洞的體積\n           - 決定保留「體積最小」的那個洞 (不補)\n           - 產生「被填補的洞」的遮罩 (fill_mask)\n        3. 將 fill_mask 與 原始 vol 做 OR 運算回傳\n        \"\"\"\n    # 0. 基礎設定\n    vol_bool = vol > 0\n    struct_26 = ndimage.generate_binary_structure(3, 3)\n\n    # 建立一個全域的遮罩，用來存放「決定要補起來的洞」\n    total_fill_mask = np.zeros_like(vol_bool)\n\n    # --- 第一階段：前處理 (移除小物件) ---\n    # 這裡確保雜訊不影響洞的判斷\n    print(\"正在移除小物件以進行判斷...\")\n    clean_vol = remove_small_objects(vol_bool, min_size=5000, connectivity=1)\n\n    # 取得切片\n    slices = _octant_slices(clean_vol.shape, (2, 2, 2))\n\n    # --- 第二階段：在乾淨的圖上計算「要補哪些洞」 ---\n    for slc in slices:\n        # 注意：這裡我們只看 clean_vol\n        block = clean_vol[slc]\n\n        # 如果區塊全空，跳過\n        if not np.any(block):\n            continue\n\n        # 1. 全補\n        filled_block = ndimage.binary_fill_holes(block, structure=struct_26)\n\n        # 2. 找出洞 (填補後 - 原圖)\n        holes_mask = filled_block & ~block\n\n        # 如果沒有洞，跳過\n        if not np.any(holes_mask):\n            continue\n\n        # 3. 標記所有的洞 (不再需要標記宿主物件)\n        labeled_holes, num_holes = ndimage.label(holes_mask, structure=struct_26)\n\n        # 預設：所有洞都要補 (稍後把要保留的那個洞挖掉)\n        holes_to_fill_in_this_slice = holes_mask.copy()\n\n        # 4. 邏輯判斷：保留體積最小的洞\n        if num_holes > 0:\n            # 計算每個洞的體積\n            # sizes[0] 是背景，sizes[1:] 是各個洞的體積\n            sizes = np.bincount(labeled_holes.ravel())\n\n            # 找出最小洞的 Label\n            # argmin 回傳的是索引，因為我們切片了 [1:]，所以索引要 +1 才是 Label\n            target_hole_label = np.argmin(sizes[1:]) + 1\n\n            # 從「要補的洞」清單中，移除這個「要保留的洞」\n            # 將該 Label 的位置設為 False (不補)\n            holes_to_fill_in_this_slice[labeled_holes == target_hole_label] = False\n\n        # 5. 將決定好要補的洞，寫入全域遮罩\n        total_fill_mask[slc] = holes_to_fill_in_this_slice\n\n    # --- 第三階段：合併回原圖 ---\n    # 結果 = 原始圖 OR 填補遮罩 (這樣原本 < 5000 的物件也會回來，且被選中的洞也被補上了)\n    final_result = vol_bool | total_fill_mask\n\n    return final_result.astype(np.uint8)\n\n\n# # 1. 將 bm 宣告為全域變數，但在各進程內延遲載入\n# _worker_bm = None\n#\n#\n# def _get_bm_model():\n#     \"\"\"確保每個子進程都有自己的 bm 實體\"\"\"\n#     global _worker_bm\n#     if _worker_bm is None:\n#         _worker_bm = load_betti_matching()\n#     return _worker_bm\n#\n#\n# # ==========================================\n# # 2. 高斯修補副程式 (維持回傳邏輯)\n# # ==========================================\n# def _apply_gaussian_patch_and_return(chunk, center_coord, sigma_val):\n#     z, y, x = center_coord\n#     d, h, w = chunk.shape\n#     r = 2\n#\n#     z_min, z_max = max(0, z - r), min(d, z + r + 1)\n#     y_min, y_max = max(0, y - r), min(h, y + r + 1)\n#     x_min, x_max = max(0, x - r), min(w, x + r + 1)\n#\n#     local_patch = chunk[z_min:z_max, y_min:y_max, x_min:x_max].astype(np.float32)\n#     smoothed_patch = gaussian_filter(local_patch, sigma=sigma_val)\n#\n#     # 修改局部\n#     chunk[z_min:z_max, y_min:y_max, x_min:x_max] |= (smoothed_patch > 0.5)\n#     return chunk\n#\n#\n# # ==========================================\n# # 3. Worker 任務 (不再接收 bm_model)\n# # ==========================================\n# def _process_chunk_worker(args):\n#     \"\"\"\n#     args 現在只包含 (slc_idx, chunk_data, sigma)\n#     \"\"\"\n#     slc_idx, chunk, sigma = args\n#\n#     # 獲取該進程專用的 bm 物件\n#     bm = _get_bm_model()\n#\n#     # 偵測破洞 (依據你提供的邏輯：自己跟自己 match 找特徵)\n#     topo_pr = (~chunk).astype(np.uint8)\n#     result = bm.compute_matching(topo_pr, topo_pr)\n#\n#     # 取得 Matched 座標 (依據你最新的程式碼需求)\n#     if len(result.input1_matched_birth_coordinates) > 1:\n#         local_coords = result.input1_matched_birth_coordinates[1]\n#         for l_coord in local_coords:\n#             chunk = _apply_gaussian_patch_and_return(chunk, l_coord, sigma)\n#     print(\"處理完成\")\n#     return slc_idx, chunk\n#\n#\n# # ==========================================\n# # 4. 主函數\n# # ==========================================\n# def fill_holes_multiprocess_return(vol, split_parts=(2, 2, 2), sigma=1.0, n_procs=8):\n#     \"\"\"\n#     注意：參數中移除了 bm，改由子進程自行載入\n#     \"\"\"\n#     output_vol = np.zeros_like(vol)\n#     slices = _octant_slices(vol.shape, split_parts)\n#\n#     # 準備參數包 (剔除 bm_model)\n#     task_args = [\n#         (i, vol[slc].copy(), sigma)\n#         for i, slc in enumerate(slices)\n#     ]\n#\n#     print(f\"[-] 開始多進程修補 (進程數: {n_procs})...\")\n#\n#     # 使用 'spawn' 或 'forkserver' 模式在某些系統更穩定，但 Linux 預設 fork 即可\n#     with multiprocessing.Pool(processes=n_procs) as pool:\n#         results = pool.map(_process_chunk_worker, task_args)\n#\n#     print(\"[-] 正在組裝最終體積...\")\n#     for idx, processed_chunk in results:\n#         slc = slices[idx]\n#         output_vol[slc] = processed_chunk\n#\n#     print(\"[-] 處理完成。\")\n#     return output_vol\n\n\ndef process_gaussian_octants_dual_threshold(\n        vol: np.ndarray,\n        sigma: float = 1.0,\n        inner_threshold: float = 0.5,\n        border_threshold: float = 0.1,\n        border: int = 8\n) -> np.ndarray:\n    \"\"\"\n    將 3D 體積切塊處理：\n    1. 內部區域 (Inner): 使用 inner_threshold (預設 0.5)\n    2. 邊界區域 (Border): 使用 border_threshold (預設 0.1)\n    \"\"\"\n    output_vol = vol.copy()\n\n    # 取得 8 塊區域的切片範圍\n    slices_list = _octant_slices(vol.shape, splits=(2, 2, 2))\n\n    for sl in slices_list:\n        # 取出該區塊\n        sub_vol = vol[sl]\n        depth, height, width = sub_vol.shape\n\n        # 防呆：確保區塊大於兩倍 border\n        if height <= 2 * border or width <= 2 * border or depth <= 2 * border:\n            # 若區塊太小無法區分 border，則統一使用 border_threshold 處理或跳過\n            blurred_small = ndi.gaussian_filter(sub_vol.astype(float), sigma=sigma)\n            output_vol[sl] = (blurred_small > border_threshold).astype(vol.dtype)\n            continue\n\n        # 核心：高斯模糊運算\n        blurred_sub = ndi.gaussian_filter(sub_vol.astype(float), sigma=sigma)\n\n        # 建立該區塊的結果容器，先統一用 border_threshold 處理整塊\n        processed_block = (blurred_sub > border_threshold).astype(vol.dtype)\n\n        # 定義中心區域 (Inner area) 的切片\n        inner_z = slice(border, -border if border > 0 else None)\n        inner_h = slice(border, -border if border > 0 else None)\n        inner_w = slice(border, -border if border > 0 else None)\n        inner_slice = (inner_z, inner_h, inner_w)\n\n        # 針對「中心區域」覆蓋使用 inner_threshold 的結果\n        # 這會讓中心區域與邊界區域有不同的二值化敏感度\n        processed_block[inner_slice] = (blurred_sub[inner_slice] > inner_threshold).astype(vol.dtype)\n\n        # 將處理完的塊填回大圖\n        output_vol[sl] = processed_block\n\n    return output_vol\n\n\ndef fill_2x2diag(\n        mask_3d: np.ndarray,\n        axis: int = 0,\n        gap: int = 0,\n) -> np.ndarray:\n    \"\"\"\n    掃描相鄰切片，修補 2x2 對角跳空 (diagonal gap) 問題。\n\n    Args:\n        mask_3d: 3D binary mask，shape = (z, y, x)\n        axis:    沿哪個軸做切片比較。0 = z 軸（預設），1 = y 軸，2 = x 軸\n        gap:     prev 與 curr 之間允許的中間層數。\n                 0 = 原始行為（相鄰兩層）。\n                 N = prev 與 curr 相距 N+1 層，中間 N 層需滿足\n                     連續相鄰對在該 2x2 的 AND 總和 > 1，\n                     確認有連續性後再判斷 prev/curr 的 3/3/4 對角條件，\n                     全部通過才將整段全設為 True。\n\n    Returns:\n        修補後的 3D binary mask（不修改原始輸入）\n    \"\"\"\n    if axis not in (0, 1, 2):\n        raise ValueError(f\"axis 必須是 0、1 或 2，收到 {axis}\")\n    if gap < 0:\n        raise ValueError(f\"gap 必須 >= 0，收到 {gap}\")\n\n    work = np.moveaxis(mask_3d, axis, 0).copy().astype(bool)\n    d, h, w = work.shape\n\n    step = gap + 1  # prev 與 curr 的 index 距離\n    corner_offsets = [(0, 0), (0, 1), (1, 0), (1, 1)]\n\n    for z in range(step, d):\n        prev_idx = z - step\n        curr_idx = z  # inclusive，共 step+1 層需要填補\n\n        prev = work[prev_idx].astype(np.uint8)\n        curr = work[curr_idx].astype(np.uint8)\n        orr = prev | curr\n\n        for dr, dc in corner_offsets:\n            r0, r1 = dr, h - 1 - dr  # 2x2 左上角的 row 範圍 [r0, r1)\n            c0, c1 = dc, w - 1 - dc\n\n            if r1 <= r0 or c1 <= c0:\n                continue\n\n            def box_sum(arr: np.ndarray) -> np.ndarray:\n                return (arr[r0: r1, c0: c1]\n                        + arr[r0: r1, c0 + 1: c1 + 1]\n                        + arr[r0 + 1: r1 + 1, c0: c1]\n                        + arr[r0 + 1: r1 + 1, c0 + 1: c1 + 1])\n\n            # ── 條件一：prev / curr 滿足 3/3/4 對角規則 ──────────────────\n            cond = (box_sum(prev) == 3) & (box_sum(curr) == 3) & (box_sum(orr) == 4)\n\n            if not np.any(cond):\n                continue\n\n            # ── 條件二：中間所有相鄰對的 AND 在該 2x2 內 > 1 ─────────────\n            # 遍歷 (prev_idx, prev_idx+1), ..., (curr_idx-1, curr_idx)\n            # gap=0 時此迴圈不執行\n            if gap > 0:\n                for i in range(prev_idx, curr_idx):\n                    a = work[i].astype(np.uint8)\n                    b = work[i + 1].astype(np.uint8)\n                    cond = cond & (box_sum(a & b) > 1)\n\n                    if not np.any(cond):\n                        break  # 提早結束，此 corner 已無命中點\n\n            # ── 填補：整段 [prev_idx, curr_idx] 的該 2x2 全設為 True ─────\n            hit_rows, hit_cols = np.where(cond)\n            if len(hit_rows) == 0:\n                continue\n\n            abs_rows = hit_rows + r0\n            abs_cols = hit_cols + c0\n\n            for rr, cc in zip(abs_rows, abs_cols):\n                work[prev_idx: curr_idx + 1, rr: rr + 2, cc: cc + 2] = True\n\n    return np.moveaxis(work, 0, axis)\n\n\ndef apply_post_processing(\n        pred: np.ndarray,\n        min_size: int,\n        split_paper: bool,\n        split_paper_iter: int = 1,\n        split_high_prob: float = 0.8,\n        prob: np.ndarray = None,  # 機率圖\n        fg_threshold: float = None,  # 門檻值 (必填以啟用取代功能)\n        skip: bool = False,\n        line_norm: bool = False,\n        y_axis_closing: bool = False,\n        y_axis_fill_holes: bool = False,\n        z_axis_fill_holes: bool = False,\n        fill_hole: bool = False,\n        line_endpoints_repair: bool = False,\n        line_mask_repair: bool = False,\n        repair_settings: dict = None,\n        sandwich: bool = False,\n        sandwich_iterations: int = 1,\n        gaussian: bool = False,\n        clear_boundary: bool = False,\n        dig2x2: bool = False\n) -> np.ndarray:\n    \"\"\"\n    整合後處理流程。\n    \"\"\"\n    # 複製並初始化\n    vol = pred.copy()\n    vol[vol == 2] = 0  # 清除 Class 2 (若有)\n\n    if skip:\n        return vol\n\n\n    if min_size > 0:\n        vol_bool = vol > 0\n        vol_bool = remove_small_objects(vol_bool, min_size=5000, connectivity=1)\n\n        vol = vol_bool.astype(np.uint8)\n    else:\n        vol = (vol > 0).astype(np.uint8)\n\n    # ==========================================\n    # 3. 斷線修復\n    # ==========================================\n    if line_endpoints_repair:\n        vol = execute_line_repair(vol, settings=repair_settings)\n        # vol = execute_line_repair_8slice(vol, settings=repair_settings)\n\n        vol = (vol > 0).astype(np.uint8)\n\n    if line_mask_repair:\n        vol = line_repair_msk(vol, min_area=10, max_link_dist=120.0, pass_iters=2, line_thickness=2)\n        vol = (vol > 0).astype(np.uint8)\n\n    # ==========================================\n    # 4. 形態學與幾何優化\n    # ==========================================\n    if line_norm:\n        Z_15, _ = normalize_segments_3d(vol, radius=1.5, border=8, axis=0, sigma=0, repair=True, repair_dijkstra=False,\n                                        max_dist=30, max_angle_deg=45, safe_distance=2, iterations=1)\n\n        Y_15, _ = normalize_segments_3d(vol, radius=1.5, border=8, axis=1, sigma=0, repair=True, repair_dijkstra=False,\n                                        max_dist=30, max_angle_deg=45, safe_distance=2, iterations=1)\n\n        vol = vol | Z_15 | Y_15\n        vol = vol.astype(np.uint8)\n\n    if sandwich:\n        # 建議 gap=2 以應對連續斷層\n        vol = robust_mask_refine(vol, max_gap=3, iterations=sandwich_iterations, min_neighbors=9, kernel_size=5)\n\n        vol = vol.astype(np.uint8)\n\n    if gaussian:\n        vol = process_gaussian_octants_dual_threshold(\n            vol,\n            sigma=1,\n            inner_threshold=0.5,\n            border_threshold=0.5,\n            border=8)\n\n    if split_paper:  # 註：這裡依照您提供的 code 原樣保留變數名\n        foreground_prob = prob[1]\n        high_conf_mask = (foreground_prob > split_high_prob)\n\n        # 建議這裡也考慮使用 8-slice 版本的 remove_small_objects\n        high_conf_mask = remove_small_objects(high_conf_mask, min_size=5000, connectivity=1)\n        high_conf_labels = measure_label(high_conf_mask).astype(np.int32)\n\n        # --- 改用 8-slice 版本 ---\n        vol, debug_paths = get_split_paper_8slice(\n            vol,\n            high_conf_labels,\n            ray_len=64,\n            max_iter=split_paper_iter,\n            cleanup_iter=2\n        )\n        # -----------------------\n\n        vol = vol > 0\n        vol = vol.astype(np.uint8)\n\n    if y_axis_closing:\n        vol = apply_y_axis_closing(vol, iterations=1)\n        vol = vol.astype(np.uint8)\n\n    if y_axis_fill_holes:\n        vol = apply_y_axis_fill_holes(vol)\n        vol = vol.astype(np.uint8)\n\n    if z_axis_fill_holes:\n        vol = apply_z_axis_fill_holes(vol)\n        vol = vol.astype(np.uint8)\n\n    if fill_hole:\n        vol = fill_hole_8slice_keep_one_per_slice(vol)\n        # vol = binary_fill_holes(vol)\n        # vol = vol.astype(np.uint8)\n\n    if min_size > 0:\n        vol_bool = vol > 0\n        vol_bool = remove_small_objects(vol_bool, min_size=min_size, connectivity=1)\n        # vol_bool = clear_boundary_faces(vol_bool,margin=3)\n        vol_bool = remove_small_8slice(vol_bool, min_size=10)\n\n        vol = vol_bool.astype(np.uint8)\n    else:\n        vol = (vol > 0).astype(np.uint8)\n    if dig2x2:\n        vol = fill_2x2diag(vol, axis=0, gap=0)\n        vol = fill_2x2diag(vol, axis=0, gap=1)\n        vol = fill_2x2diag(vol, axis=1, gap=0)\n        vol = fill_2x2diag(vol, axis=1, gap=1)\n        vol.astype(np.uint8)\n\n    return vol.astype(np.uint8)","metadata":{"trusted":true,"_kg_hide-input":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import cupy as cp\nimport cupyx.scipy.ndimage as cndi\nimport numpy as np\n\n\ndef apply_ninth_place_post_proc(binary_mask, MIN_SIZE=3000):\n    # B. 執行新的後處理函數\n    processed_data = apply_post_processing(\n        binary_mask, \n        skip=False,\n        min_size=MIN_SIZE,\n        split_paper=False,\n        split_paper_iter=1,\n        split_high_prob=0.9,\n        line_norm=True,\n        line_mask_repair=False,\n        line_endpoints_repair=False,\n        y_axis_closing=False,\n        z_axis_fill_holes=False,\n        y_axis_fill_holes=False,\n        sandwich=True,\n        sandwich_iterations=5,\n        fill_hole=False,\n        gaussian=True,\n        dig2x2=True\n    )\n    return processed_data\n\n\ndef postprocess_mask_voxel(mask, min_cc_volume=3000, median_iters=5, verbose=True):\n    \"\"\"\n    GPU version:\n    1. Remove CC < min_cc_volume\n    2. Apply iterative 3x3x3 median per surviving CC\n    3. Reassemble volume\n    Returns stats for compatibility.\n    \"\"\"\n\n    if mask.sum() == 0:\n        if verbose:\n            return mask, {\n                \"original_cc\": 0,\n                \"final_cc\": 0,\n                \"removed_cc\": 0\n            }\n        return mask\n\n    try:\n        # Move mask to GPU\n        mask_gpu = cp.asarray(mask).astype(bool)\n\n        # Count original CCs (GPU)\n        labeled_before, n_before = cndi.label(mask_gpu)\n\n        final_mask = cp.zeros_like(mask_gpu, dtype=bool)\n\n        # Process each CC\n        for cc_id in range(1, int(n_before) + 1):\n            component = (labeled_before == cc_id)\n            comp_size = int(component.sum())\n\n            if comp_size < min_cc_volume:\n                continue\n\n            comp_uint8 = component.astype(cp.uint8)\n\n            # Median smoothing\n            #for _ in range(median_iters):\n            #    comp_uint8 = (\n            #        cndi.median_filter(comp_uint8, size=3) > 0\n            #    ).astype(cp.uint8)\n\n            final_mask |= comp_uint8.astype(bool)\n\n        # Count final CCs\n        labeled_after, n_after = cndi.label(final_mask)\n\n        # Move result back to CPU\n        final_mask_cpu = cp.asnumpy(final_mask).astype(np.uint8)\n        final_mask_cpu = apply_ninth_place_post_proc(final_mask_cpu, MIN_SIZE = 5000)\n\n        if verbose:\n            return final_mask_cpu, {\n                \"original_cc\": int(n_before),\n                \"final_cc\": int(n_after),\n                \"removed_cc\": max(0, int(n_before) - int(n_after))\n            }\n\n\n        return final_mask_cpu\n\n    except Exception as e:\n        import traceback\n        traceback.print_exc()\n\n        if verbose:\n            return mask, {\n                \"original_cc\": 0,\n                \"final_cc\": 0,\n                \"removed_cc\": 0,\n                \"error\": str(e)\n            }\n\n        return mask","metadata":{"trusted":true,"execution":{"execution_failed":"2026-03-01T16:10:33.22Z"}},"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":{"execution_failed":"2026-03-01T16:10:33.22Z"},"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":{"execution_failed":"2026-03-01T16:10:33.22Z"},"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":{"execution_failed":"2026-03-01T16:10:33.22Z"},"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":{"execution_failed":"2026-03-01T16:10:33.22Z"},"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":{"execution_failed":"2026-03-01T16:10:33.221Z"},"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},"execution":{"execution_failed":"2026-03-01T16:10:33.221Z"}},"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":{"execution_failed":"2026-03-01T16:10:33.221Z"},"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},"execution":{"execution_failed":"2026-03-01T16:10:33.221Z"}},"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},"execution":{"execution_failed":"2026-03-01T16:10:33.221Z"}},"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},"execution":{"execution_failed":"2026-03-01T16:10:33.221Z"}},"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},"execution":{"execution_failed":"2026-03-01T16:10:33.221Z"}},"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},"execution":{"execution_failed":"2026-03-01T16:10:33.222Z"}},"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},"execution":{"execution_failed":"2026-03-01T16:10:33.222Z"}},"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},"execution":{"execution_failed":"2026-03-01T16:10:33.222Z"}},"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},"execution":{"execution_failed":"2026-03-01T16:10:33.222Z"}},"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}]}