{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":117682,"databundleVersionId":15062069,"isSourceIdPinned":false},{"sourceType":"datasetVersion","sourceId":14810989,"datasetId":9471057,"databundleVersionId":15666974},{"sourceType":"datasetVersion","sourceId":14903537,"datasetId":9536090,"databundleVersionId":15768535},{"sourceType":"datasetVersion","sourceId":14861012,"datasetId":9506451,"databundleVersionId":15722169},{"sourceType":"datasetVersion","sourceId":14829308,"datasetId":9484229,"databundleVersionId":15687042,"isSourceIdPinned":true},{"sourceType":"modelInstanceVersion","sourceId":768516,"databundleVersionId":15864953,"modelInstanceId":587120,"modelId":599446,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":297302543,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":297674933,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":297732241,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":297771298,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":297827765,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":298023302,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":299093011,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":299102216,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":300093425,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":300258358,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":300641787,"isSourceIdPinned":false}],"dockerImageVersionId":31260,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\n# import os\n# for dirname, _, filenames in os.walk('/kaggle/input'):\n#     for filename in filenames:\n#         print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-28T08:50:46.557512Z","iopub.execute_input":"2026-02-28T08:50:46.557668Z","iopub.status.idle":"2026-02-28T08:50:50.140743Z","shell.execute_reply.started":"2026-02-28T08:50:46.557651Z","shell.execute_reply":"2026-02-28T08:50:50.140000Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 2.5 D filtermask","metadata":{}},{"cell_type":"markdown","source":"","metadata":{},"attachments":{}},{"cell_type":"code","source":"# var=\"/kaggle/input/notebooks/crischir/vsdetection-packages-offline-installer-only/whls\"\n# # clear_output()\n# !pip install \\\n#   \"$var\"/tifffile-*.whl \\\n#   \"$var\"/imagecodecs-*.whl \\\n#   \"$var\"/scikit_image-*.whl \\\n#   --no-index \\\n#   --find-links \"$var\"\n\n# clear_output()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T08:50:50.141221Z","iopub.execute_input":"2026-02-28T08:50:50.141464Z","iopub.status.idle":"2026-02-28T08:50:50.144284Z","shell.execute_reply.started":"2026-02-28T08:50:50.141446Z","shell.execute_reply":"2026-02-28T08:50:50.143665Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !pip install fast_simplification","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T08:52:33.999798Z","iopub.execute_input":"2026-02-28T08:52:34.000059Z","iopub.status.idle":"2026-02-28T08:54:01.405078Z","shell.execute_reply.started":"2026-02-28T08:52:34.000037Z","shell.execute_reply":"2026-02-28T08:54:01.404358Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import fast_simplification","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T08:54:01.405889Z","iopub.execute_input":"2026-02-28T08:54:01.406061Z","iopub.status.idle":"2026-02-28T08:54:01.432728Z","shell.execute_reply.started":"2026-02-28T08:54:01.406043Z","shell.execute_reply":"2026-02-28T08:54:01.432028Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"var=\"/kaggle/input/notebooks/crischir/vsdetection-packages-offline-installer-only/whls\"\n\n!pip install --no-index --find-links=\"$var\" \\\n    tifffile \\\n    imagecodecs \\\n    scikit-image \\\n    connected-components-3d \\\n    trimesh\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T08:50:50.145001Z","iopub.execute_input":"2026-02-28T08:50:50.145154Z","iopub.status.idle":"2026-02-28T08:50:55.438221Z","shell.execute_reply.started":"2026-02-28T08:50:50.145139Z","shell.execute_reply":"2026-02-28T08:50:55.437400Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import trimesh","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T08:50:55.438819Z","iopub.execute_input":"2026-02-28T08:50:55.438982Z","iopub.status.idle":"2026-02-28T08:50:56.044575Z","shell.execute_reply.started":"2026-02-28T08:50:55.438965Z","shell.execute_reply":"2026-02-28T08:50:56.043730Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\nprint(sys.version)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T08:50:56.045043Z","iopub.execute_input":"2026-02-28T08:50:56.045304Z","iopub.status.idle":"2026-02-28T08:50:56.048058Z","shell.execute_reply.started":"2026-02-28T08:50:56.045288Z","shell.execute_reply":"2026-02-28T08:50:56.047483Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm import tqdm\nimport cv2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T08:50:56.048396Z","iopub.execute_input":"2026-02-28T08:50:56.048547Z","iopub.status.idle":"2026-02-28T08:50:58.467451Z","shell.execute_reply.started":"2026-02-28T08:50:56.048532Z","shell.execute_reply":"2026-02-28T08:50:58.466726Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport sys\nimport zipfile\nimport time\nfrom tqdm import tqdm\nfrom concurrent.futures import ThreadPoolExecutor\nfrom pathlib import Path\nimport numpy as np\nimport jax\nimport jax.numpy as jnp\nfrom jax import jit\nfrom functools import partial\nimport flax.linen as nn\nfrom flax.traverse_util import unflatten_dict\nfrom flax.training import train_state\nfrom flax.jax_utils import replicate  # <-- TPU specific addition\nimport matplotlib.pyplot as plt\nfrom PIL import Image, ImageSequence\nfrom scipy import ndimage as ndi\nfrom skimage.morphology import remove_small_objects, skeletonize\nfrom skimage.segmentation import watershed\nfrom skimage.measure import regionprops\nfrom skimage.measure import marching_cubes\nfrom skimage.feature import peak_local_max     \nimport cc3d\n\n# ==============================================================================\n# 0. QUICK RUNTIME INFO\n# ==============================================================================\nprint(\"Python:\", sys.version)\nprint(\"JAX devices:\", jax.devices())\n\n# ==============================================================================\n# 1. VISUALIZATION HELPER (5-VIEW)\n# ==============================================================================\ndef visualize_prediction_samples_jax(debug_data, vid=None):\n    if vid is None:\n        vid = list(debug_data.keys())[0]\n\n    data = debug_data[vid]\n    vol_in = data['input']\n    vol_raw = data['raw_prob']      \n    vol_pp = data['postproc_prob']  \n    vol_final = data['final_mask']  \n\n    z = vol_in.shape[0] // 2\n    \n    p2, p98 = np.percentile(vol_in[z], (2, 98))\n    orig_norm = np.clip((vol_in[z] - p2) / (p98 - p2 + 1e-7), 0, 1)\n\n    fig, axes = plt.subplots(1, 5, figsize=(30, 6))\n    fig.suptitle(f\"Pipeline Analysis: Weight-based Post-Processing (Slice: {z})\", fontsize=16)\n\n    axes[0].imshow(orig_norm, cmap='gray')\n    axes[0].set_title(\"1. Original CT\")\n\n    im1 = axes[1].imshow(vol_raw[z], cmap='magma', vmin=0, vmax=1)\n    axes[1].set_title(\"2. Ensemble Output\\n(Pre-Weights)\")\n    plt.colorbar(im1, ax=axes[1], fraction=0.046, pad=0.04)\n\n    im2 = axes[2].imshow(vol_pp[z], cmap='magma', vmin=0, vmax=1)\n    axes[2].set_title(\"3. 2D Post-Proc Model\\n(With Loaded Weights)\")\n    plt.colorbar(im2, ax=axes[2], fraction=0.046, pad=0.04)\n\n    diff = vol_pp[z] - vol_raw[z]\n    im3 = axes[3].imshow(diff, cmap='coolwarm', vmin=-0.5, vmax=0.5)\n    axes[3].set_title(\"4. Weight Impact\\n(PP - Raw)\")\n    plt.colorbar(im3, ax=axes[3], fraction=0.046, pad=0.04)\n\n    axes[4].imshow(vol_final[z], cmap='gray')\n    axes[4].set_title(\"5. Final 3D Mask\")\n\n    for ax in axes:\n        ax.axis('off')\n\n    plt.tight_layout()\n    plt.show()\n\n# ==============================================================================\n# 0. Debug mode \n# ==============================================================================\nPRODUCTION_MODE = True \nENABLE_VIZ = not PRODUCTION_MODE\n\n# ==============================================================================\n# 2. CONFIGURATION\n# ==============================================================================\nTTA_INFERENCE = True  \n\nENABLE_2D_POSTPROC = False\nENSEMBLE_MODE = \"agreement\"   # 'mean', 'max', 'agreement', 'hybrid', 'single'\nAGREEMENT_THRESHOLD = 2\nALPHA = 0.8\n\nMODEL_CHECKPOINTS = {\n    \"2-5d-filter-training-pix2pix\":[\n        \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix/checkpoints/pix2pix_ep_20.npz\",\n        \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix/checkpoints/pix2pix_ep_30.npz\",\n    ],\n    \"2-5d-filter-training-pix2pix-pred-mask\":[\n        \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix-pred-mask/checkpoints/pix2pix_best.npz\",\n        \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix-pred-mask/checkpoints/pix2pix_ep_70.npz\",\n        \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix-pred-mask/checkpoints/pix2pix_ep_90.npz\",\n        \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix-pred-mask/checkpoints/pix2pix_ep_120.npz\",\n        \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix-pred-mask/checkpoints/pix2pix_ep_170.npz\",\n    ],\n    \"test-code-for-jax-pix2pix\":[\n        \"/kaggle/input/notebooks/crischir/test-code-for-jax-pix2pix/checkpoints/pix2pix_best.npz\",\n        \"/kaggle/input/notebooks/crischir/test-code-for-jax-pix2pix/checkpoints/pix2pix_ep_20.npz\",\n       \"/kaggle/input/notebooks/crischir/test-code-for-jax-pix2pix/checkpoints/pix2pix_ep_30.npz\",\n        \"/kaggle/input/notebooks/crischir/test-code-for-jax-pix2pix/checkpoints/pix2pix_ep_40.npz\",\n        \"/kaggle/input/notebooks/crischir/test-code-sobel-jax-pix2pix/checkpoints/pix2pix_ep_110.npz\",\n        \"/kaggle/input/notebooks/crischir/test-code-sobel-jax-pix2pix/checkpoints/pix2pix_ep_120.npz\",\n       \"/kaggle/input/notebooks/crischir/test-code-sobel-jax-pix2pix/checkpoints/pix2pix_ep_130.npz\",\n        \"/kaggle/input/notebooks/crischir/test-code-sobel-jax-pix2pix/checkpoints/pix2pix_ep_140.npz\",\n        \n    ],\n    \"test-code-sobel-jax-pix2pix\":[\n        \"/kaggle/input/notebooks/crischir/test-code-sobel-jax-pix2pix/checkpoints/pix2pix_ep_100.npz\",\n        \"/kaggle/input/notebooks/crischir/test-code-sobel-jax-pix2pix/checkpoints/pix2pix_ep_130.npz\",\n    ],\n    \"tpu-training-jax-pix2pix\":[\n        \"/kaggle/input/notebooks/crischir/tpu-training-jax-pix2pix/checkpoints/pix2pix_ep_30.npz\",\n        \"/kaggle/input/notebooks/crischir/tpu-training-jax-pix2pix/checkpoints/pix2pix_ep_40.npz\",\n        \"//kaggle/input/notebooks/crischir/tpu-training-jax-pix2pix/checkpoints/pix2pix_ep_50.npz\",\n        \"/kaggle/input/notebooks/crischir/tpu-training-jax-pix2pix/checkpoints/pix2pix_ep_60.npz\",\n        \"/kaggle/input/notebooks/crischir/tpu-training-jax-pix2pix/checkpoints/pix2pix_ep_70.npz\",\n    ],\n    \"prewitt-loss-model-jax-pix2pix\":[\n        \"/kaggle/input/notebooks/crischir/prewitt-loss-model-jax-pix2pix/checkpoints/pix2pix_best.npz\",\n        \"/kaggle/input/notebooks/crischir/prewitt-loss-model-jax-pix2pix/checkpoints/pix2pix_ep_40.npz\",\n        \"/kaggle/input/notebooks/crischir/prewitt-loss-model-jax-pix2pix/checkpoints/pix2pix_ep_50.npz\",\n        \"/kaggle/input/notebooks/crischir/prewitt-loss-model-jax-pix2pix/checkpoints/pix2pix_ep_70.npz\",\n    ],\n    \"jax-pix2pix\":[\n        \"/kaggle/input/notebooks/crischir/jax-pix2pix/checkpoints/pix2pix_ep_10.npz\",\n        \"/kaggle/input/notebooks/crischir/jax-pix2pix/checkpoints/pix2pix_ep_30.npz\",\n        \"/kaggle/input/notebooks/crischir/jax-pix2pix/checkpoints/pix2pix_ep_30.npz\",\n        \"/kaggle/input/notebooks/crischir/jax-pix2pix/checkpoints/pix2pix_ep_50.npz\",\n        \"/kaggle/input/notebooks/crischir/jax-pix2pix/checkpoints/pix2pix_ep_70.npz\",\n        \"/kaggle/input/notebooks/crischir/jax-pix2pix/checkpoints/pix2pix_ep_120.npz\",\n        \"/kaggle/input/notebooks/crischir/jax-pix2pix/checkpoints/pix2pix_ep_140.npz\",\n        \"/kaggle/input/notebooks/crischir/jax-pix2pix/checkpoints/pix2pix_ep_160.npz\",\n    ],\n    \"retromasking-vesuvius\":[\n        \"/kaggle/input/datasets/crischir/retromasking-vesuvius-2-5d-filter-checkpoints/checkpoints/pix2pix_ep_30.npz\",\n        \"/kaggle/input/datasets/crischir/retromasking-vesuvius-2-5d-filter-checkpoints/checkpoints/pix2pix_ep_60.npz\",\n        \"/kaggle/input/datasets/crischir/retromasking-vesuvius-2-5d-filter-checkpoints/checkpoints/pix2pix_ep_90.npz\",\n        \"/kaggle/input/datasets/crischir/retromasking-vesuvius-2-5d-filter-checkpoints/checkpoints/pix2pix_ep_120.npz\",\n        \"/kaggle/input/datasets/crischir/retromasking-vesuvius-2-5d-filter-checkpoints/checkpoints/pix2pix_ep_120.npz\",\n        \"/kaggle/input/datasets/crischir/retromasking-vesuvius-2-5d-filter-checkpoints/checkpoints/pix2pix_ep_150.npz\",\n        \n    ],\n        \"vesuvius-2-5-mask-cldice/\":[\n        \"/kaggle/input/models/crischir/vesuvius-2-5-mask-cldice/flax/2.5d-cldice-24-samples-40-epochs/1/checkpoints/pix2pix_ep_20.npz\",\n        \"/kaggle/input/models/crischir/vesuvius-2-5-mask-cldice/flax/2.5d-cldice-24-samples-40-epochs/1/checkpoints/pix2pix_ep_30.npz\",\n        \"/kaggle/input/models/crischir/vesuvius-2-5-mask-cldice/flax/2.5d-cldice-24-samples-40-epochs/1/checkpoints/pix2pix_ep_40.npz\",\n        \"/kaggle/input/models/crischir/vesuvius-2-5-mask-cldice/flax/2.5d-cldice-24-samples-40-epochs/1/checkpoints/pix2pix_ep_50.npz\",\n        \"/kaggle/input/models/crischir/vesuvius-2-5-mask-cldice/flax/2.5d-cldice-24-samples-40-epochs/1/checkpoints/pix2pix_ep_60.npz\",\n    ],\n    \"clDice-1\": [\n        \"/kaggle/input/models/crischir/vesuvius-2-5-mask-cldice/flax/2.5d-cldice-24-samples-40-epochs/1/checkpoints/pix2pix_ep_10.npz\",\n        \"/kaggle/input/models/crischir/vesuvius-2-5-mask-cldice/flax/2.5d-cldice-24-samples-40-epochs/1/checkpoints/pix2pix_ep_20.npz\",\n       \"/kaggle/input/models/crischir/vesuvius-2-5-mask-cldice/flax/2.5d-cldice-24-samples-40-epochs/1/checkpoints/pix2pix_ep_30.npz\",\n        \"/kaggle/input/models/crischir/vesuvius-2-5-mask-cldice/flax/2.5d-cldice-24-samples-40-epochs/1/checkpoints/pix2pix_ep_40.npz\",\n        \"/kaggle/input/models/crischir/vesuvius-2-5-mask-cldice/flax/2.5d-cldice-24-samples-40-epochs/1/checkpoints/pix2pix_ep_50.npz\",\n       \"/kaggle/input/models/crischir/vesuvius-2-5-mask-cldice/flax/2.5d-cldice-24-samples-40-epochs/1/checkpoints/pix2pix_ep_60.npz\",\n    ],\n    \"clDice-2\": [\n        \"/kaggle/input/notebooks/crischir/cldice-tpu-training-pix2pix/checkpoints/pix2pix_ep_10.npz\",\n        \"/kaggle/input/notebooks/crischir/cldice-tpu-training-pix2pix/checkpoints/pix2pix_ep_20.npz\",\n        \"/kaggle/input/notebooks/crischir/cldice-tpu-training-pix2pix/checkpoints/pix2pix_ep_30.npz\",\n        \"/kaggle/input/notebooks/crischir/cldice-tpu-training-pix2pix/checkpoints/pix2pix_ep_40.npz\",\n        # \"/kaggle/input/notebooks/crischir/cldice-tpu-training-pix2pix/checkpoints/pix2pix_ep_50.npz\",\n        # \"/kaggle/input/notebooks/crischir/cldice-tpu-training-pix2pix/checkpoints/pix2pix_ep_60.npz\",\n        \n    ],\n\n    \"vesuvius-checkpoint-dataset\":[\n        \"/kaggle/input/datasets/crischir/checkpoint-vesuvius-1/checkpoints/pix2pix_ep_60.npz\",\n        \"/kaggle/input/datasets/crischir/vesuvius-checkpoint-dataset/checkpoints/pix2pix_ep_290.npz\",\n    ],\n}\n\nSELECTED =[\n    # \"tpu-training-jax-pix2pix\",\n    \"clDice-1\",\n    \"clDice-2\",\n    # \"retromasking-vesuvius\",\n    # \"retromasking-vesuvius\",\n    # \"test-code-for-jax-pix2pix\",\n    # \"2-5d-filter-training-pix2pix\",\n    # \"2-5d-filter-training-pix2pix-pred-mask\",\n]\n\nWEIGHT_PATHS =[]\nfor key in SELECTED:\n    WEIGHT_PATHS.extend(MODEL_CHECKPOINTS.get(key,[]))\n\n\nDO_3D_POSTPROCESSING = False\n \n\n# TOPO_PARAMS = {\n#     \"T_low\": 0.01,\n#     \"T_high\": 0.35, \n#     \"snake_iter\": 3,\n#     \"balloon_threshold\": 0.01,  \n#     \"smooth_threshold\": 12.0,  \n# }\n\n# TOPO_PARAMS = {\n#     \"T_low\": 0.15,               # Increased from 0.01. Acts as a strict wall so lines can't bleed into dark areas.\n#     \"T_high\": 0.45,              # Increased from 0.35. Seeds only from confident predictions.\n#     \"snake_iter\": 2,\n#     \"balloon_threshold\": 4.0,    # Increased from 0.01. Now requires at least 4 (out of 27) neighbors to expand, stopping explosive growth.\n#     \"smooth_threshold\": 16.0,    # Increased from 12.0. Prunes the edges more aggressively to keep the sheet thin.\n# }\n# TOPO_PARAMS = {\n#     \"T_low\": 0.10,\n#     \"T_high\": 0.45, \n#     \"snake_iter\": 2,\n#     \"balloon_threshold\": 4.0,  \n#     \"smooth_threshold\": 16.0,  \n#     \"dust_min_size\": 1000,     # NEW: Keeps all lines larger than 1000 voxels (adjust if some lines disappear)\n# }\n##For retromasking inly\n# TOPO_PARAMS = {\n#     \"T_low\": 0.5,              # Dropped from 0.15: Lets the mask expand into the very faint purple edges\n#     \"T_high\": 0.25,             # Dropped from 0.45: Allows weaker predicted lines to spawn their own masks\n#     \"snake_iter\": 2,\n#     \"balloon_threshold\": 3.0,   # Dropped from 4.0: Makes it easier to bridge tiny gaps in faint lines\n#     \"smooth_threshold\": 10.0,   # Dropped from 16.0: Crucial! Allows thin, 1-2 voxel thick sheets to survive without being erased\n#     \"dust_min_size\": 300,       # Dropped from 1000: Prevents small but valid faint line fragments from being deleted\n# }\nSkeleton3D = False\n\nTOPO_PARAMS = {\n    \"T_low\": 0.001,              # Keep at 0.15: Allows expansion into faint areas\n    \"T_high\": 0.25,             # Bumped slightly to 0.35: Stops absolute noise from spawning blocks, but still low enough to catch faint lines\n    \"snake_iter\": 2,\n    \"balloon_threshold\": 3.0,   # Increased from 3.0: Requires 5 neighbors to expand. This stops the masks from bridging the gap between two separate sheets.\n    \"smooth_threshold\": 20.0,   # Increased from 10.0: The most important change! This forcefully \"shaves off\" the edges of the mask, keeping the lines strictly thin.\n    \"dust_min_size\": 250,       # Keep at 300: Protects the small, faint, fragmented lines from being deleted.\n    \"skeleton_radius\": 3,\n    \n}\n\n# --- NEW: Foil Separation & Meshing Options ---\nENABLE_FOIL_SEPARATION = False   # Set to True to run 3D Watershed & keep only the largest sheets\nMAX_FOILS_TO_KEEP = 10           # Drops tiny exfoliations and keeps the top N largest foils\n\nENABLE_MESH_EXPORT = False       # Set to True to export a smoothed 3D .obj file for Blender/MeshLab\nMESH_SMOOTHING_ITERATIONS = 10   # How aggressively to smooth the 3D mesh (removes voxel stair-\n\nDATA_PATH = Path(\"/kaggle/input/competitions/vesuvius-challenge-surface-detection/\")\nPATCH_SIZE = 256\nZ_CONTEXT = 3\nSimpleThreshold = 0.1\nBATCH_SIZE = 256\ndevices = jax.devices()\nN_DEV = len(devices)\nprint(\"Num devices:\", N_DEV)\n\n# ==============================================================================\n# 3. HELPERS\n# ==============================================================================\ndef to_skeleton(vol_mask):\n    skel = skeletonize(vol_mask > 0)\n    return skel.astype(np.uint8) * 255\n\ndef build_anisotropic_struct(z_radius: int, xy_radius: int):\n    z, r = int(z_radius), int(xy_radius)\n    if z == 0 and r == 0: return None\n    if z == 0 and r > 0:\n        size = 2 * r + 1\n        struct = np.zeros((1, size, size), dtype=bool)\n        cy = cx = r\n        for dy in range(-r, r + 1):\n            for dx in range(-r, r + 1):\n                if dy * dy + dx * dx <= r * r: struct[0, cy + dy, cx + dx] = True\n        return struct\n    if z > 0 and r == 0:\n        struct = np.zeros((2 * z + 1, 1, 1), dtype=bool)\n        struct[:, 0, 0] = True\n        return struct\n    depth, size = 2 * z + 1, 2 * r + 1\n    struct = np.zeros((depth, size, size), dtype=bool)\n    cz, cy, cx = z, r, r\n    for dz in range(-z, z + 1):\n        for dy in range(-r, r + 1):\n            for dx in range(-r, r + 1):\n                if dy * dy + dx * dx <= r * r: struct[cz + dz, cy + dy, cx + dx] = True\n    return struct\n\n@partial(jax.jit, static_argnums=(1,))\ndef jax_morphological_smooth(mask, iterations, params):\n    kernel = jnp.ones((3, 3, 3))\n    kernel = kernel[None, None, ...] \n\n    dn = jax.lax.ConvDimensionNumbers(\n        lhs_spec=(0, 1, 2, 3, 4), rhs_spec=(0, 1, 2, 3, 4), out_spec=(0, 1, 2, 3, 4)\n    )\n\n    def body_fn(i, val):\n        x = val[None, None, ...]\n        \n        dil_logits = jax.lax.conv_general_dilated(x, kernel, (1,1,1), 'SAME', dimension_numbers=dn)\n        dilated = (dil_logits > params.get(\"balloon_threshold\", 4.0)).astype(jnp.float32)\n        \n        ero_logits = jax.lax.conv_general_dilated(dilated, kernel, (1,1,1), 'SAME', dimension_numbers=dn)\n        eroded = (ero_logits > params.get(\"smooth_threshold\", 16.0)).astype(jnp.float32)\n        \n        return eroded[0, 0, ...]\n\n    return jax.lax.fori_loop(0, iterations, body_fn, mask.astype(jnp.float32))\n\n@partial(jax.jit, static_argnums=(1,))\ndef jax_morphological_smooth_box(mask, iterations, params):\n    kernel = jnp.ones((3, 3, 3))\n    kernel = kernel[None, None, ...] \n\n    dn = jax.lax.ConvDimensionNumbers(\n        lhs_spec=(0, 1, 2, 3, 4), rhs_spec=(0, 1, 2, 3, 4), out_spec=(0, 1, 2, 3, 4)\n    )\n\n    def body_fn(i, val):\n        x = val[None, None, ...]\n        \n        dil_logits = jax.lax.conv_general_dilated(x, kernel, (1,1,1), 'SAME', dimension_numbers=dn)\n        dilated = (dil_logits > (params.get(\"balloon_threshold\", 0.3) * 27)).astype(jnp.float32)\n        \n        ero_logits = jax.lax.conv_general_dilated(dilated, kernel, (1,1,1), 'SAME', dimension_numbers=dn)\n        eroded = (ero_logits > (params.get(\"smooth_threshold\", 0.5) * 27)).astype(jnp.float32)\n        \n        return eroded[0, 0, ...]\n\n    return\n#Good results Before strict hysteresis \n# def topo_postprocess_3d(prob_vol, params):\n#     print(\"   ...Advanced 3D Post-processing (Snake + CCL + Skeleton)...\")\n    \n#     prob_jax = jnp.array(prob_vol)\n#     strong = prob_jax >= params[\"T_high\"]\n#     weak = prob_jax >= params[\"T_low\"]\n    \n#     mask_jax = (jax_morphological_smooth(strong.astype(jnp.float32), 2, params) > 0)\n#     mask_jax = jnp.logical_and(mask_jax, weak)\n\n#     refined_jax = jax_morphological_smooth(mask_jax, params.get(\"snake_iter\", 3), params)\n    \n#     refined_np = np.array(refined_jax) > 0.5\n    \n#     del prob_jax, strong, weak, mask_jax, refined_jax\n#     import gc; gc.collect()\n\n#     # CPU: Connected Components (Dust Removal)\n#     labels = cc3d.connected_components(refined_np.astype(np.uint8), connectivity=26)\n    \n#     if labels.max() == 0:\n#         return np.zeros_like(prob_vol, dtype=np.uint8)\n\n#     stats = cc3d.statistics(labels)\n#     voxel_counts = stats['voxel_counts']\n    \n#     min_size = params.get(\"dust_min_size\", 300)\n#     valid_labels = voxel_counts > min_size\n#     valid_labels[0] = False  \n    \n#     final_mask = valid_labels[labels]\n\n#     # Final Thinning & Controlled Expansion\n#     if Skeleton3D:\n#         # 1. Reduce to perfect 1-voxel centerlines\n#         final_mask = skeletonize(final_mask > 0)\n        \n#         # 2. Expand to the desired uniform thickness\n#         skel_radius = params.get(\"skeleton_radius\", 0)\n#         if skel_radius > 0:\n#             # Generate a 3D cross structure (prevents corners from becoming too blocky)\n#             struct = ndi.generate_binary_structure(3, 1)\n#             final_mask = ndi.binary_dilation(final_mask, structure=struct, iterations=skel_radius)\n\n#     return final_mask.astype(np.uint8) * 255\n\ndef topo_postprocess_3d(prob_vol, params):\n    print(\"   ...Strict Hysteresis 3D Post-processing...\")\n    \n    prob_jax = jnp.array(prob_vol)\n    strong = prob_jax >= params[\"T_high\"]\n    weak = prob_jax >= params[\"T_low\"]\n    \n    # 1. Grow strong seeds slightly to bridge gaps\n    mask_jax = (jax_morphological_smooth(strong.astype(jnp.float32), 1, params) > 0)\n    \n    # 2. Constrain to the weak map\n    mask_jax = jnp.logical_and(mask_jax, weak)\n\n    # 3. Final smoothing to connect lines, BUT strictly constrained to the probability map!\n    refined_jax = jax_morphological_smooth(mask_jax.astype(jnp.float32), params.get(\"snake_iter\", 1), params)\n    \n    # THIS IS THE MAGIC FIX: Slice off all the bloat\n    refined_jax = jnp.logical_and((refined_jax > 0.5), weak)\n    \n    refined_np = np.array(refined_jax)\n    \n    del prob_jax, strong, weak, mask_jax, refined_jax\n    import gc; gc.collect()\n\n    # 4. CPU: Connected Components (Dust Removal)\n    labels = cc3d.connected_components(refined_np.astype(np.uint8), connectivity=26)\n    \n    if labels.max() == 0:\n        return np.zeros_like(prob_vol, dtype=np.uint8)\n\n    stats = cc3d.statistics(labels)\n    voxel_counts = stats['voxel_counts']\n    \n    min_size = params.get(\"dust_min_size\", 300)\n    valid_labels = voxel_counts > min_size\n    valid_labels[0] = False  \n    \n    final_mask = valid_labels[labels]\n\n    return final_mask.astype(np.uint8) * 255\n# def topo_postprocess_3d(prob_vol, params):\n#     print(\"   ...Advanced 3D Post-processing (Snake + CCL)...\")\n    \n#     prob_jax = jnp.array(prob_vol)\n#     strong = prob_jax >= params[\"T_high\"]\n#     weak = prob_jax >= params[\"T_low\"]\n    \n#     mask_jax = (jax_morphological_smooth(strong.astype(jnp.float32), 2, params) > 0)\n#     mask_jax = jnp.logical_and(mask_jax, weak)\n\n#     refined_jax = jax_morphological_smooth(mask_jax, params.get(\"snake_iter\", 3), params)\n    \n#     refined_np = np.array(refined_jax) > 0.5\n    \n#     del prob_jax, strong, weak, mask_jax, refined_jax\n#     import gc; gc.collect()\n\n#     labels = cc3d.connected_components(refined_np.astype(np.uint8), connectivity=26)\n    \n#     if labels.max() == 0:\n#         return np.zeros_like(prob_vol, dtype=np.uint8)\n\n#     stats = cc3d.statistics(labels)\n#     largest_label = np.argmax(stats['voxel_counts'][1:]) + 1\n#     final_mask = (labels == largest_label)\n\n#     if Skeleton3D:\n#         final_mask = skeletonize(final_mask > 0)\n\n#     return final_mask.astype(np.uint8) * 255\n\n@partial(jax.jit, static_argnums=(1, 2))\ndef jax_morphology_closing(mask, z_radius, xy_radius):\n    zk, yk, xk = 2 * int(z_radius) + 1, 2 * int(xy_radius) + 1, 2 * int(xy_radius) + 1\n    kernel = jnp.ones((zk, yk, xk), dtype=jnp.float32)\n    kernel_norm = kernel / jnp.sum(kernel)\n    data = mask[None, ..., None].astype(jnp.float32)\n    kernel_jax = kernel_norm[..., None, None]\n    dn = jax.lax.ConvDimensionNumbers(\n        lhs_spec=(0, 1, 2, 3, 4), rhs_spec=(0, 1, 2, 3, 4), out_spec=(0, 1, 2, 3, 4)\n    )\n    dilated = jax.lax.conv_general_dilated(data, kernel_jax, (1, 1, 1), 'SAME', dimension_numbers=dn)\n    dilated_mask = (dilated > 0).astype(jnp.float32)\n    eroded = jax.lax.conv_general_dilated(dilated_mask, kernel_jax, (1, 1, 1), 'SAME', dimension_numbers=dn)\n    return (eroded >= 0.99).astype(jnp.uint8)[0, ..., 0]\n\ndef normalize_patch(patch: np.ndarray) -> np.ndarray:\n    patch = patch.astype(np.float32)\n    return (patch - patch.mean()) / (patch.std() + 1e-6)\n\ndef load_npz_weights(path):\n    with np.load(path, allow_pickle=False) as data:\n        flat_dict = {k: v for k, v in data.items()}\n    params_flat = {k.replace('params/', ''): v for k, v in flat_dict.items() if k.startswith('params/')}\n    stats_flat = {k.replace('stats/', ''): v for k, v in flat_dict.items() if k.startswith('stats/')}\n    return {\n        'params': unflatten_dict({tuple(k.split('/')): v for k, v in params_flat.items()}),\n        'batch_stats': unflatten_dict({tuple(k.split('/')): v for k, v in stats_flat.items()})\n    }\n\ndef load_ensemble_weights(paths):\n    ensemble_vars =[]\n    if ENSEMBLE_MODE == 'single': paths = [paths[0]]\n    for p in paths:\n        if Path(p).exists():\n            print(f\"📦 Loading weights: {p}\")\n            ensemble_vars.append(load_npz_weights(p))\n        else:\n            print(f\"⚠️ Warning: Path not found {p}\")\n    return ensemble_vars\n# def separate_and_identify_foils(binary_mask_3d, max_foils):\n#     \"\"\" Separates touching sheets using 3D Watershed and filters by mass. \"\"\"\n#     print(\"   ...Calculating 3D Distance Transform (This may take a moment)...\")\n#     distance = ndi.distance_transform_edt(binary_mask_3d)\n    \n#     print(\"   ...Finding foil centerlines (Seeds)...\")\n#     # footprint=(3,3,3) splits touching sheets aggressively\n#     local_max_coords = peak_local_max(distance, min_distance=3, footprint=np.ones((3, 3, 3)))\n    \n#     markers = np.zeros_like(binary_mask_3d, dtype=np.int32)\n#     for i, coord in enumerate(local_max_coords):\n#         markers[tuple(coord)] = i + 1  \n\n#     print(\"   ...Running 3D Watershed to slice foils...\")\n#     separated_labels = watershed(-distance, markers, mask=binary_mask_3d)\n\n#     print(\"   ...Sorting separated foils by Mass...\")\n#     props = regionprops(separated_labels)\n#     sorted_foils = sorted(props, key=lambda r: r.area, reverse=True)\n    \n#     kept_count = min(max_foils, len(sorted_foils))\n#     print(f\"   => Found {len(sorted_foils)} unique foil fragments. Keeping top {kept_count}.\")\n    \n#     # Create a clean mask keeping only the largest foils (removes exfoliations)\n#     clean_separated_mask = np.zeros_like(binary_mask_3d, dtype=np.uint8)\n#     for rank, foil in enumerate(sorted_foils):\n#         if rank < kept_count:\n#             # Assign them distinct greyscale values (e.g., 25, 50, 75...) to view them as distinct layers\n#             intensity = int(255 * ((rank + 1) / kept_count))\n#             clean_separated_mask[separated_labels == foil.label] = intensity\n            \n#     return clean_separated_mask\ndef separate_and_identify_foils(binary_mask_3d, max_foils):\n    \"\"\" Identifies separated sheets using fast 3D CCL and filters by mass. \"\"\"\n    print(\"   ...Identifying connected 3D sheets...\")\n    \n    # 1. Find all distinct, disconnected 3D objects\n    labels = cc3d.connected_components(binary_mask_3d.astype(np.uint8), connectivity=26)\n    \n    if labels.max() == 0:\n        print(\"   ⚠️ No sheets found!\")\n        return np.zeros_like(binary_mask_3d, dtype=np.uint8)\n\n    print(\"   ...Sorting foils by Mass (Volume)...\")\n    stats = cc3d.statistics(labels)\n    voxel_counts = stats['voxel_counts'][1:]  # [1:] skips the background (label 0)\n    \n    # Create a list of (label_id, volume) so we can sort them\n    foil_stats = [(i + 1, count) for i, count in enumerate(voxel_counts)]\n    \n    # Sort from largest to smallest volume\n    sorted_foils = sorted(foil_stats, key=lambda x: x[1], reverse=True)\n    \n    kept_count = min(max_foils, len(sorted_foils))\n    print(f\"   => Found {len(sorted_foils)} unique foil fragments. Keeping top {kept_count}.\")\n    \n    # 2. Create the clean, hierarchical mask\n    clean_separated_mask = np.zeros_like(binary_mask_3d, dtype=np.uint8)\n    \n    for rank in range(kept_count):\n        label_id = sorted_foils[rank][0]\n        \n        # Assign distinct greyscale intensities. \n        # Rank 1 (Largest) gets the brightest white. Smaller sheets get darker grey.\n        intensity = int(255 * (1.0 - (rank / kept_count))) \n        \n        clean_separated_mask[labels == label_id] = intensity\n            \n    return clean_separated_mask\n\ndef simplify_to_mesh(binary_mask_3d, out_name, smoothing_iters=20):\n    if 'trimesh' not in sys.modules:\n        return None\n    \n    # 1. Check if there is actually data to mesh\n    if binary_mask_3d.sum() == 0:\n        print(f\"⚠️ Skipping {out_name}: Mask is empty.\")\n        return None\n\n    # try:\n    #     # 2. Ensure data type is compatible\n    #     verts, faces, normals, values = marching_cubes(binary_mask_3d.astype(float), level=0.5)\n        \n    #     if len(verts) == 0:\n    #         print(f\"⚠️ No surface found for {out_name}.\")\n    #         return None\n\n    #     mesh = trimesh.Trimesh(vertices=verts, faces=faces, vertex_normals=normals)\n        \n    #     if smoothing_iters > 0:\n    #         trimesh.smoothing.filter_taubin(mesh, iterations=smoothing_iters)\n        \n    #     # 3. Final check before export\n    #     if not mesh.is_empty:\n    #         mesh.export(out_name)\n    #         print(f\"=> Successfully saved {out_name}!\")\n    #         return out_name\n    #     else:\n    #         print(\"⚠️ Resulting mesh was empty.\")\n    #         return None\n\n    # except Exception as e:\n    #     print(f\"⚠️ Failed: {e}\")\n    #     return None\n\n    try:\n        # Ensure there is a 0-value border\n        padded_mask = np.pad(binary_mask_3d, pad_width=1, mode='constant', constant_values=0)\n        \n        verts, faces, normals, values = marching_cubes(padded_mask, level=0.5)\n        \n        print(f\"    DEBUG: Found {len(verts)} vertices and {len(faces)} faces.\")\n        \n        if len(faces) == 0:\n            print(\"    ⚠️ Error: No faces were generated. Check if mask values are all identical.\")\n            return None\n\n        mesh = trimesh.Trimesh(vertices=verts, faces=faces, vertex_normals=normals)\n        \n        if smoothing_iters > 0:\n            # Check if mesh is valid before smoothing\n            if mesh.is_empty:\n                print(\"    ⚠️ Mesh is empty before smoothing!\")\n            mesh = trimesh.smoothing.filter_taubin(mesh, iterations=smoothing_iters)\n        mesh = mesh.simplify_quadric_decimation(0.1)\n        \n        # Explicitly check for content before saving\n        if len(mesh.faces) > 0:\n            mesh.export(out_name)\n            import os\n            print(f\"    => Saved! File size: {os.path.getsize(out_name)} bytes\")\n        else:\n            print(\"    ⚠️ Export aborted: Mesh has no geometry.\")\n            \n        return out_name\n    except Exception as e:\n        print(f\"⚠️ Failed: {e}\")\n        return None\n# ==============================================================================\n# 4. MODELS (PIX2PIX + POST-PROCESS)\n# ==============================================================================\nclass Pix2PixGenerator(nn.Module):\n    @nn.compact\n    def __call__(self, x, train: bool = False):\n        def conv_bn_leaky(feat, out_c):\n            feat = nn.Conv(out_c, (4, 4), strides=(2, 2), padding='SAME')(feat)\n            feat = nn.BatchNorm(use_running_average=not train)(feat)\n            return nn.leaky_relu(feat, 0.2)\n        \n        s1 = nn.leaky_relu(nn.Conv(64, (4, 4), strides=(2, 2), padding='SAME')(x), 0.2)\n        s2 = conv_bn_leaky(s1, 128)\n        s3 = conv_bn_leaky(s2, 256)\n        \n        def up_bn_relu(feat, skip, out_c):\n            feat = nn.ConvTranspose(out_c, (4, 4), strides=(2, 2), padding='SAME')(feat)\n            feat = nn.BatchNorm(use_running_average=not train)(feat)\n            feat = nn.relu(feat)\n            return jnp.concatenate([feat, skip], axis=-1)\n            \n        u1 = up_bn_relu(s3, s2, 128)\n        u2 = up_bn_relu(u1, s1, 64)\n        return nn.ConvTranspose(1, (4, 4), strides=(2, 2), padding='SAME')(u2)\n\nclass PostProcessModel(nn.Module):\n    @nn.compact\n    def __call__(self, x, train: bool = True):\n        x = nn.Conv(64, (4, 4), strides=(2, 2), padding='SAME')(x)\n        s1 = nn.leaky_relu(x, 0.2)\n        s2 = nn.leaky_relu(nn.BatchNorm(use_running_average=not train)(nn.Conv(128, (4, 4), strides=(2, 2), padding='SAME')(s1)), 0.2)\n        s3 = nn.leaky_relu(nn.BatchNorm(use_running_average=not train)(nn.Conv(256, (4, 4), strides=(2, 2), padding='SAME')(s2)), 0.2)\n        u1 = jnp.concatenate([nn.relu(nn.BatchNorm(use_running_average=not train)(nn.ConvTranspose(128, (4, 4), strides=(2, 2), padding='SAME')(s3))), s2], axis=-1)\n        u2 = jnp.concatenate([nn.relu(nn.BatchNorm(use_running_average=not train)(nn.ConvTranspose(64, (4, 4), strides=(2, 2), padding='SAME')(u1))), s1], axis=-1)\n        return nn.ConvTranspose(1, (4, 4), strides=(2, 2), padding='SAME')(u2)\n\ndef load_postproc_weights(path):\n    with np.load(path, allow_pickle=True) as data:\n        params, stats = data[\"params\"].item(), data[\"stats\"].item()\n    return params, stats\n\nPOSTPROC_PATH =\"/kaggle/input/notebooks/crischir/jaxvesuviuspostprocessingtraining/checkpoints/postproc_best.npz\"\npostproc_model = PostProcessModel()\n# Only load if path exists to prevent crash\nif Path(POSTPROC_PATH).exists():\n    pp_params, pp_stats = load_postproc_weights(POSTPROC_PATH)\n    pp_vars = {\"params\": pp_params, \"batch_stats\": pp_stats}\n\n@jax.jit\ndef apply_postproc(vars, x):\n    logits = postproc_model.apply(vars, x, train=False)\n    return nn.sigmoid(logits)\n\n# ==============================================================================\n# 5. ENSEMBLE + MULTI-GPU + TTA (TPU OPTIMIZED)\n# ==============================================================================\ndef stack_ensemble_vars(ensemble_vars):\n    # Stacks pytrees so the sequence is along axis 0 for jax.lax.scan\n    return jax.tree_util.tree_map(lambda *xs: jnp.stack(xs), *ensemble_vars)\n\ndef make_predict_ensemble(model, num_models, mode=\"hybrid\"):\n    @jax.pmap\n    def _predict(vars_stacked, x):\n        def scan_body(carry, v):\n            accum_logits, accum_max_probs, accum_probs, votes = carry\n            logits = model.apply(v, x, train=False)\n            p = nn.sigmoid(logits)\n            \n            accum_logits += logits\n            accum_max_probs = jnp.maximum(accum_max_probs, p)\n            accum_probs += p\n            votes += (p > 0.1).astype(jnp.int32)\n            \n            return (accum_logits, accum_max_probs, accum_probs, votes), None\n            \n        init_carry = (\n            jnp.zeros((*x.shape[:-1], 1), dtype=jnp.float32),\n            jnp.zeros((*x.shape[:-1], 1), dtype=jnp.float32),\n            jnp.zeros((*x.shape[:-1], 1), dtype=jnp.float32),\n            jnp.zeros((*x.shape[:-1], 1), dtype=jnp.int32)\n        )\n        \n        # Sequentially evaluate models heavily decreasing TPU HBM Requirements\n        (accum_logits, accum_max_probs, accum_probs, votes), _ = jax.lax.scan(scan_body, init_carry, vars_stacked)\n        \n        mean_logits = accum_logits / num_models\n        mean_probs = nn.sigmoid(mean_logits)\n        \n        if mode == \"mean\":\n            return mean_probs\n        elif mode == \"max\":\n            return accum_max_probs\n        elif mode == \"hybrid\":\n            return mean_probs + ALPHA * (accum_max_probs - mean_probs)\n        elif mode == \"agreement\":\n            mean_p = accum_probs / num_models\n            return jnp.where(votes >= AGREEMENT_THRESHOLD, 1.0, mean_p)\n        return mean_probs\n\n    return _predict\n\n_predict_fn_cache = {}\ndef get_cached_predict_fn(model, num_models, mode):\n    key = (id(model), num_models, mode)\n    if key not in _predict_fn_cache:\n        _predict_fn_cache[key] = make_predict_ensemble(model, num_models, mode)\n    return _predict_fn_cache[key]\n\ndef get_ensemble_prob_batch(model, vars_stacked_rep, num_models, patches_jax, mode=\"hybrid\"):\n    B = patches_jax.shape[0]\n    \n    pad = (N_DEV - (B % N_DEV)) % N_DEV\n    if pad > 0:\n        patches_jax = jnp.concatenate([\n            patches_jax, \n            jnp.zeros(((pad,) + patches_jax.shape[1:]), patches_jax.dtype)\n        ], axis=0)\n    \n    B_padded = patches_jax.shape[0]\n    per_dev = B_padded // N_DEV\n    \n    predict_fn = get_cached_predict_fn(model, num_models, mode)\n    \n    def run_tta(x):\n        x_sharded = x.reshape(N_DEV, per_dev, *x.shape[1:])\n        return predict_fn(vars_stacked_rep, x_sharded).reshape(B_padded, PATCH_SIZE, PATCH_SIZE, 1)\n\n    if not TTA_INFERENCE:\n        probs_flat = run_tta(patches_jax)\n    else:\n        p_orig = run_tta(patches_jax)\n        p_h    = jnp.flip(run_tta(jnp.flip(patches_jax, axis=2)), axis=2)\n        p_v    = jnp.flip(run_tta(jnp.flip(patches_jax, axis=1)), axis=1)\n        p_hv   = jnp.flip(run_tta(jnp.flip(patches_jax, axis=(1, 2))), axis=(1, 2))\n        \n        probs_flat = (p_orig + p_h + p_v + p_hv) / 4.0\n\n    return probs_flat[:B]\n\n# ==============================================================================\n# 6. INFERENCE PIPELINE\n# ==============================================================================\ndef run_inference_pipeline(model, vars_stacked_rep, num_models):\n    test_dir = DATA_PATH / 'test_images'\n    if not test_dir.exists(): return[], {}\n    \n    submission_files, debug_volumes = [], {}\n    test_ids =[f.stem for f in test_dir.glob('*.tif')]\n\n    for vid in test_ids:\n        vol_start_time = time.time()\n        print(f\"\\n🏁 Volume: {vid} | TTA: {TTA_INFERENCE} | Mode: {ENSEMBLE_MODE}\")\n        with Image.open(str(test_dir / f\"{vid}.tif\")) as img:\n            vol_input = np.stack([np.array(f) for f in ImageSequence.Iterator(img)], axis=0)\n        z_max, h, w = vol_input.shape\n        vol_raw, vol_prob = np.zeros((z_max, h, w), dtype=np.float32), np.zeros((z_max, h, w), dtype=np.float32)\n        all_coords =[(y, x) for y in range(0, h, PATCH_SIZE) for x in range(0, w, PATCH_SIZE)]\n\n        def get_batch_worker(z_idx, coords_subset):\n            batch_patches =[]\n            for y, x in coords_subset:\n                y1, x1 = min(y + PATCH_SIZE, h), min(x + PATCH_SIZE, w)\n                y0, x0 = y1 - PATCH_SIZE, x1 - PATCH_SIZE\n                patch = vol_input[z_idx - Z_CONTEXT: z_idx + Z_CONTEXT + 1, y0:y1, x0:x1].transpose(1, 2, 0)\n                batch_patches.append(normalize_patch(patch))\n            return np.stack(batch_patches), coords_subset\n\n        with ThreadPoolExecutor(max_workers=4) as executor:\n            for z in tqdm(range(Z_CONTEXT, z_max - Z_CONTEXT), desc=\"Slices\"):\n                coord_chunks = [all_coords[i:i + BATCH_SIZE] for i in range(0, len(all_coords), BATCH_SIZE)]\n                future_batch = executor.submit(get_batch_worker, z, coord_chunks[0])\n                for i in range(len(coord_chunks)):\n                    current_batch_np, batch_coords = future_batch.result()\n                    if i + 1 < len(coord_chunks): future_batch = executor.submit(get_batch_worker, z, coord_chunks[i + 1])\n                    \n                    prob_np = np.array(get_ensemble_prob_batch(model, vars_stacked_rep, num_models, current_batch_np, ENSEMBLE_MODE))[..., 0]\n                    \n                    for idx, (y, x) in enumerate(batch_coords):\n                        y1, x1 = min(y + PATCH_SIZE, h), min(x + PATCH_SIZE, w)\n                        y0, x0 = y1 - PATCH_SIZE, x1 - PATCH_SIZE\n                        raw_patch = prob_np[idx]\n                        vol_raw[z, y0:y1, x0:x1] = raw_patch\n                        \n                        if ENABLE_2D_POSTPROC:\n                            pp_out = apply_postproc(pp_vars, jnp.array(raw_patch[..., None].astype(np.float32)[None, ...]))\n                            vol_prob[z, y0:y1, x0:x1] = np.array(pp_out[0, ..., 0])\n                        else: vol_prob[z, y0:y1, x0:x1] = raw_patch\n\n        print(f\"🪄 GPU/TPU Post-processing...\")\n        vol_final = topo_postprocess_3d(vol_prob, TOPO_PARAMS) if DO_3D_POSTPROCESSING else (vol_prob > SimpleThreshold).astype(np.uint8) * 255\n        # 2. OPTIONAL: Foil Separation & Identification\n        if ENABLE_FOIL_SEPARATION:\n            print(f\"🧩 Separating touching foils and removing exfoliations...\")\n            # We pass vol_final > 0 to ensure it is boolean\n            vol_final = separate_and_identify_foils(vol_final > 0, MAX_FOILS_TO_KEEP)\n\n        # 3. OPTIONAL: 3D Mesh Export\n        if ENABLE_MESH_EXPORT:\n            mesh_name = f\"{vid}_3D_model.obj\"\n            # Pass vol_final > 0 so it meshes all the kept foils together\n            exported_mesh = simplify_to_mesh(vol_final > 0, mesh_name, MESH_SMOOTHING_ITERATIONS)\n            if exported_mesh:\n                submission_files.append(exported_mesh)\n        out_name = f\"{vid}.tif\"\n        import tifffile\n        tifffile.imwrite(out_name, vol_final, compression='deflate')\n        submission_files.append(out_name)\n        if not PRODUCTION_MODE: debug_volumes[vid] = {'input': vol_input, 'raw_prob': vol_raw, 'postproc_prob': vol_prob, 'final_mask': vol_final}\n        if PRODUCTION_MODE: \n            del vol_input, vol_raw, vol_prob, vol_final\n            import gc; gc.collect()\n        print(f\"✅ {vid} done in {time.time() - vol_start_time:.2f}s\")\n    return submission_files, debug_volumes\n\n# ==============================================================================\n# 7. VISUALIZATION\n# ==============================================================================\ndef visualize_results(debug_data):\n    for vid, data in debug_data.items():\n        vol_in, vol_raw, vol_final = data['input'], data['raw_prob'], data['final_mask']\n        z_max = vol_in.shape[0]\n        indices, labels =[Z_CONTEXT + 2, z_max // 2, z_max - Z_CONTEXT - 2], [\"Start\", \"Mid\", \"End\"]\n        fig, axes = plt.subplots(3, 4, figsize=(20, 15))\n        for i, z in enumerate(indices):\n            axes[i, 0].imshow(vol_in[z], cmap='gray'); axes[i, 0].set_title(f\"{labels[i]} (z={z}) Input\")\n            im2 = axes[i, 1].imshow(vol_raw[z], cmap='magma', vmin=0, vmax=1); plt.colorbar(im2, ax=axes[i, 1])\n            axes[i, 2].imshow(vol_final[z], cmap='gray'); axes[i, 2].set_title(\"3D Mask\")\n            overlay = np.zeros((*vol_in[z].shape, 3)); norm_in = vol_in[z] / 255.0\n            overlay[..., 0] = np.clip(norm_in + (vol_final[z] > 0) * 0.5, 0, 1)\n            overlay[..., 1], overlay[..., 2] = norm_in, norm_in\n            axes[i, 3].imshow(overlay); axes[i, 3].set_title(\"Overlay\")\n            for j in range(4): axes[i, j].axis('off')\n        plt.tight_layout()\n\ndef visualize_postproc_comparison(debug_data, vid=None):\n    if vid is None: vid = list(debug_data.keys())[0]\n    data = debug_data[vid]\n    vol_in, vol_raw, vol_pp, vol_final = data['input'], data['raw_prob'], data['postproc_prob'], data['final_mask']\n    z = vol_in.shape[0] // 2\n    fig, axes = plt.subplots(1, 4, figsize=(28, 6))\n    axes[0].imshow(vol_in[z], cmap='gray'); axes[0].axis('off')\n    im1 = axes[1].imshow(vol_raw[z], cmap='magma', vmin=0, vmax=1); plt.colorbar(im1, ax=axes[1])\n    im2 = axes[2].imshow(vol_pp[z], cmap='magma', vmin=0, vmax=1); plt.colorbar(im2, ax=axes[2])\n    axes[3].imshow(vol_final[z], cmap='gray'); axes[3].axis('off')\n    plt.tight_layout(); plt.show()\n\n# ==============================================================================\n# 8. MAIN\n# ==============================================================================\nif __name__ == \"__main__\":\n    netG = Pix2PixGenerator()\n    ensemble_vars = load_ensemble_weights(WEIGHT_PATHS)\n    \n    if not ensemble_vars:\n        print(\"❌ Error: No weights loaded.\")\n    else:\n        num_models = len(ensemble_vars)\n        print(f\"🔥 Stacking and replicating {num_models} model weights to {N_DEV} TPU cores...\")\n        vars_stacked = stack_ensemble_vars(ensemble_vars)\n        vars_stacked_rep = replicate(vars_stacked)\n\n        print(\"🔥 JIT Warmup...\")\n        dummy = jnp.zeros((BATCH_SIZE, PATCH_SIZE, PATCH_SIZE, 2 * Z_CONTEXT + 1), jnp.float32)\n        _ = get_ensemble_prob_batch(netG, vars_stacked_rep, num_models, dummy, ENSEMBLE_MODE)\n        \n        output_tifs, debug_data = run_inference_pipeline(netG, vars_stacked_rep, num_models)\n        if not PRODUCTION_MODE and debug_data:\n            visualize_postproc_comparison(debug_data)\n            if ENABLE_VIZ: visualize_prediction_samples_jax(debug_data); visualize_results(debug_data)\n        if output_tifs:\n            with zipfile.ZipFile('submission.zip', 'w') as z:\n                for f in output_tifs: z.write(f)\n            for f in output_tifs: \n                if os.path.exists(f): os.remove(f)\n                    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T08:54:13.417626Z","iopub.execute_input":"2026-02-28T08:54:13.418194Z","iopub.status.idle":"2026-02-28T08:54:54.812261Z","shell.execute_reply.started":"2026-02-28T08:54:13.418176Z","shell.execute_reply":"2026-02-28T08:54:54.811264Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Gemini","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}}]}