{"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},{"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},{"sourceType":"kernelVersion","sourceId":297302543},{"sourceType":"kernelVersion","sourceId":297674933},{"sourceType":"kernelVersion","sourceId":297732241},{"sourceType":"kernelVersion","sourceId":297771298},{"sourceType":"kernelVersion","sourceId":297827765},{"sourceType":"kernelVersion","sourceId":298023302},{"sourceType":"kernelVersion","sourceId":299093011},{"sourceType":"kernelVersion","sourceId":299102216},{"sourceType":"kernelVersion","sourceId":300093425},{"sourceType":"kernelVersion","sourceId":300258358}],"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-27T12:31:52.629661Z","iopub.execute_input":"2026-02-27T12:31:52.629942Z","iopub.status.idle":"2026-02-27T12:31:54.110068Z","shell.execute_reply.started":"2026-02-27T12:31:52.629882Z","shell.execute_reply":"2026-02-27T12:31:54.109140Z"}},"outputs":[],"execution_count":null},{"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-27T12:31:54.111120Z","iopub.execute_input":"2026-02-27T12:31:54.111552Z","iopub.status.idle":"2026-02-27T12:32:01.466302Z","shell.execute_reply.started":"2026-02-27T12:31:54.111520Z","shell.execute_reply":"2026-02-27T12:32:01.465195Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\nprint(sys.version)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-27T12:32:01.468522Z","iopub.execute_input":"2026-02-27T12:32:01.468832Z","iopub.status.idle":"2026-02-27T12:32:01.474989Z","shell.execute_reply.started":"2026-02-27T12:32:01.468797Z","shell.execute_reply":"2026-02-27T12:32:01.474008Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm import tqdm\nimport cv2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-27T12:32:01.476079Z","iopub.execute_input":"2026-02-27T12:32:01.476319Z","iopub.status.idle":"2026-02-27T12:32:01.924669Z","shell.execute_reply.started":"2026-02-27T12:32:01.476295Z","shell.execute_reply":"2026-02-27T12:32:01.923886Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport sys\nimport zipfile\nimport time\nfrom concurrent.futures import ThreadPoolExecutor\nfrom pathlib import Path\n\nimport numpy as np\nimport jax\nimport jax.numpy as jnp\nfrom jax import jit\nfrom functools import partial\n\nimport flax.linen as nn\nfrom flax.traverse_util import unflatten_dict\nfrom flax.training import train_state\n\nimport matplotlib.pyplot as plt\nfrom PIL import Image, ImageSequence\nfrom scipy import ndimage as ndi\nfrom skimage.morphology import remove_small_objects, skeletonize\n# ==============================================================================\n# 0. Debug mode \n# =============================================================================\n\nPRODUCTION_MODE = False  # Set to False when debugging/visualizing locally\nENABLE_VIZ = not PRODUCTION_MODE\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# ==============================================================================\n\ndef visualize_prediction_samples_jax(debug_data, vid=None):\n    \"\"\"\n    Visualizes the transition from Ensemble -> 2D Post-Process Model -> Final 3D.\n    \"\"\"\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']      # BEFORE Post-process weights\n    vol_pp = data['postproc_prob']  # AFTER Post-process weights\n    vol_final = data['final_mask']  # AFTER 3D Topology\n\n    z = vol_in.shape[0] // 2\n    \n    # Normalizing input for display\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    # 1. Input\n    axes[0].imshow(orig_norm, cmap='gray')\n    axes[0].set_title(\"1. Original CT\")\n\n    # 2. Raw Ensemble (Before PP Weights)\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    # 3. 2D Post-Processed (The \"After\" of your loaded weights)\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    # 4. Difference Map (Visualizing what the weights actually changed)\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    # 5. Final Mask\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# 2. CONFIGURATION\n# ==============================================================================\n\nENABLE_VIZ = True\nENABLE_2D_POSTPROC=True\nENSEMBLE_MODE = \"mean\"   # 'mean', 'max', 'agreement', 'hybrid', 'single'\nAGREEMENT_THRESHOLD = 5\nALPHA = 0.8\n\nMODEL_CHECKPOINTS = {\n    \"2-5d-filter-training-pix2pix\": [\n        \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix/checkpoints/pix2pix_best.npz\",\n        \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix/checkpoints/pix2pix_ep_10.npz\",\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_10.npz\",\n        \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix-pred-mask/checkpoints/pix2pix_ep_20.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        \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix-pred-mask/checkpoints/pix2pix_ep_180.npz\",\n        \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix-pred-mask/checkpoints/pix2pix_ep_190.npz\",\n        \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix-pred-mask/checkpoints/pix2pix_ep_200.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_30.npz\",\n        \"/kaggle/input/notebooks/crischir/test-code-for-jax-pix2pix/checkpoints/pix2pix_ep_40.npz\",\n        \"/kaggle/input/notebooks/crischir/test-code-for-jax-pix2pix/checkpoints/pix2pix_ep_50.npz\",\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_best.npz\",\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_60.npz\",\n        \"/kaggle/input/notebooks/crischir/tpu-training-jax-pix2pix/checkpoints/pix2pix_ep_70.npz\",\n        \"/kaggle/input/notebooks/crischir/tpu-training-jax-pix2pix/checkpoints/pix2pix_ep_20.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    ],\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_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_best.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_140.npz\",\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}\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_30.npz\",\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_60.npz\",\n          \"/kaggle/input/notebooks/crischir/tpu-training-jax-pix2pix/checkpoints/pix2pix_best.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_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_120.npz\",\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}\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_30.npz\",\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_60.npz\",\n          \"/kaggle/input/notebooks/crischir/tpu-training-jax-pix2pix/checkpoints/pix2pix_best.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_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_120.npz\",\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    \"retromasking-vesuvius\",\n   \"2-5d-filter-training-pix2pix-pred-mask\",\n   \"test-code-for-jax-pix2pix\",\n    \"jax-pix2pix\",\n]\n\nWEIGHT_PATHS = []\nfor key in SELECTED:\n    WEIGHT_PATHS.extend(MODEL_CHECKPOINTS.get(key, []))\n\nDO_3D_POSTPROCESSING = True\nTOPO_PARAMS = {\n    \"T_low\": 0.05,\n    \"T_high\": 0.85,\n    \"z_radius\": 1,\n    \"xy_radius\": 2,\n    \"dust_min_size\": 25,\n}\n\nDATA_PATH = Path(\"/kaggle/input/competitions/vesuvius-challenge-surface-detection/\")\nPATCH_SIZE = 256\nZ_CONTEXT = 3\nSimpleThreshold = 0.15\n\nBATCH_SIZE = 256\n\ndevices = jax.devices()\nN_DEV = len(devices)\nprint(\"Num devices:\", N_DEV)\n\n# ==============================================================================\n# 3. HELPERS\n# ==============================================================================\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:\n        return None\n\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:\n                    struct[0, cy + dy, cx + dx] = True\n        return struct\n\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\n    depth = 2 * z + 1\n    size = 2 * r + 1\n    struct = np.zeros((depth, size, size), dtype=bool)\n    cz = z\n    cy = cx = 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:\n                    struct[cz + dz, cy + dy, cx + dx] = True\n    return struct\n\ndef topo_postprocess_3d(prob_vol, params):\n    print(\"   ...3D Topology Post-processing...\")\n    strong = prob_vol >= params[\"T_high\"]\n    weak = prob_vol >= params[\"T_low\"]\n\n    if not strong.any():\n        return np.zeros_like(prob_vol, dtype=np.uint8)\n\n    struct_hyst = ndi.generate_binary_structure(3, 3)\n    mask = ndi.binary_propagation(strong, mask=weak, structure=struct_hyst)\n\n    if not mask.any():\n        return np.zeros_like(prob_vol, dtype=np.uint8)\n\n    struct_close = build_anisotropic_struct(params[\"z_radius\"], params[\"xy_radius\"])\n    if struct_close is not None:\n        mask = ndi.binary_closing(mask, structure=struct_close)\n\n    if params[\"dust_min_size\"] > 0:\n        mask = remove_small_objects(mask.astype(bool), min_size=params[\"dust_min_size\"])\n\n    return 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\n    data = mask[None, ..., None].astype(jnp.float32)\n    kernel_jax = kernel_norm[..., None, None]\n\n    dn = jax.lax.ConvDimensionNumbers(\n        lhs_spec=(0, 1, 2, 3, 4),\n        rhs_spec=(0, 1, 2, 3, 4),\n        out_spec=(0, 1, 2, 3, 4)\n    )\n\n    dilated = jax.lax.conv_general_dilated(\n        data, kernel_jax, (1, 1, 1), 'SAME',\n        dimension_numbers=dn\n    )\n    dilated_mask = (dilated > 0).astype(jnp.float32)\n\n    eroded = jax.lax.conv_general_dilated(\n        dilated_mask, kernel_jax, (1, 1, 1), 'SAME',\n        dimension_numbers=dn\n    )\n\n    return (eroded >= 0.99).astype(jnp.uint8)[0, ..., 0]\n\ndef topo_postprocess_3d_gpu(prob_vol, params):\n    d_prob = jnp.array(prob_vol)\n    strong = (d_prob >= params[\"T_high\"])\n    weak = (d_prob >= params[\"T_low\"])\n\n    mask = strong\n    for _ in range(2):\n        mask = jax_morphology_closing(mask, 1, 1)\n        mask = jnp.logical_and(mask, weak)\n\n    final_mask = jax_morphology_closing(\n        mask,\n        int(params[\"z_radius\"]),\n        int(params[\"xy_radius\"])\n    )\n\n    return np.array(final_mask) * 255\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':\n        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\n# ==============================================================================\n# 4. MODELS (PIX2PIX + POST-PROCESS)\n# ==============================================================================\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(\n            nn.BatchNorm(use_running_average=not train)(\n                nn.Conv(128, (4, 4), strides=(2, 2), padding='SAME')(s1)\n            ),\n            0.2\n        )\n        s3 = nn.leaky_relu(\n            nn.BatchNorm(use_running_average=not train)(\n                nn.Conv(256, (4, 4), strides=(2, 2), padding='SAME')(s2)\n            ),\n            0.2\n        )\n\n        u1 = jnp.concatenate(\n            [\n                nn.relu(\n                    nn.BatchNorm(use_running_average=not train)(\n                        nn.ConvTranspose(128, (4, 4), strides=(2, 2), padding='SAME')(s3)\n                    )\n                ),\n                s2\n            ],\n            axis=-1\n        )\n        u2 = jnp.concatenate(\n            [\n                nn.relu(\n                    nn.BatchNorm(use_running_average=not train)(\n                        nn.ConvTranspose(64, (4, 4), strides=(2, 2), padding='SAME')(u1)\n                    )\n                ),\n                s1\n            ],\n            axis=-1\n        )\n        return nn.ConvTranspose(1, (4, 4), strides=(2, 2), padding='SAME')(u2)\n\nclass PostProcState(train_state.TrainState):\n    batch_stats: dict\n\ndef load_postproc_weights(path):\n    \"\"\"\n    Loads the post-processing model weights saved with:\n    np.savez(..., params=cpu_state.params, stats=cpu_state.batch_stats)\n    \"\"\"\n    with np.load(path, allow_pickle=True) as data:\n        params = data[\"params\"].item()\n        stats = data[\"stats\"].item()\n    return params, stats\n\n# Path to your trained post-processing checkpoint\n# POSTPROC_PATH = \"/kaggle/input/notebooks/crischir/jaxvesuviuspostprocessingtraining/checkpoints/postproc_e20.npz\"  # adjust if needed\nPOSTPROC_PATH =\"/kaggle/input/datasets/crischir/postprocessing2vesuviuscheckpoints/checkpoints/postproc_e20.npz\" #nice results\n# POSTPROC_PATH =\"/kaggle/input/notebooks/crischir/jaxvesuviuspostprocessingtrainingfull/checkpoints/postproc_epoch_5.npz\"\n\"\"\npostproc_model = PostProcessModel()\npp_params, pp_stats = load_postproc_weights(POSTPROC_PATH)\npp_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\n# ==============================================================================\n\ndef make_apply_fn(model):\n    def _apply(vars, x):\n        return model.apply(vars, x, train=False)\n    return jax.jit(_apply)\n\ndef get_ensemble_prob_single(model, ensemble_vars, patch_jax, mode=\"hybrid\"):\n    if patch_jax.ndim == 3:\n        patch_jax = patch_jax[None, ...]\n\n    apply_fn = make_apply_fn(model)\n\n    all_logits = [apply_fn(v, patch_jax) for v in ensemble_vars]\n    stacked_logits = jnp.stack(all_logits, axis=0)\n    stacked_probs = nn.sigmoid(stacked_logits)\n\n    if mode == \"max\":\n        final_prob = jnp.max(stacked_probs, axis=0)\n    elif mode == \"mean\":\n        avg_logits = jnp.mean(stacked_logits, axis=0)\n        final_prob = nn.sigmoid(avg_logits)\n    elif mode == \"agreement\":\n        votes = jnp.sum(stacked_probs > 0.1, axis=0)\n        mean_prob = jnp.mean(stacked_probs, axis=0)\n        final_prob = jnp.where(votes >= AGREEMENT_THRESHOLD, 1.0, mean_prob)\n    elif mode == \"hybrid\":\n        p_mean = nn.sigmoid(jnp.mean(stacked_logits, axis=0))\n        p_max = jnp.max(stacked_probs, axis=0)\n        final_prob = p_mean + ALPHA * (p_max - p_mean)\n    else:\n        final_prob = stacked_probs[0]\n\n    return final_prob[0]\n\ndef get_ensemble_prob_batch_single_device(model, ensemble_vars, patches_jax, mode=\"hybrid\"):\n    def _single(p):\n        return get_ensemble_prob_single(model, ensemble_vars, p, mode)\n    return jax.vmap(_single)(patches_jax)\n\ndef make_multi_gpu_batch_fn(model, ensemble_vars, mode=\"hybrid\"):\n    def local_batch_fn(patches_local):\n        return get_ensemble_prob_batch_single_device(model, ensemble_vars, patches_local, mode)\n    pmapped = jax.pmap(local_batch_fn, axis_name=\"devices\")\n    return pmapped\n\ndef get_ensemble_prob_batch(model, ensemble_vars, patches_jax, mode=\"hybrid\"):\n    B = patches_jax.shape[0]\n    if N_DEV <= 1 or B < N_DEV:\n        return get_ensemble_prob_batch_single_device(model, ensemble_vars, patches_jax, mode)\n\n    pad = (N_DEV - (B % N_DEV)) % N_DEV\n    if pad > 0:\n        pad_shape = (pad,) + patches_jax.shape[1:]\n        patches_jax = jnp.concatenate([patches_jax, jnp.zeros(pad_shape, patches_jax.dtype)], axis=0)\n        B_padded = patches_jax.shape[0]\n    else:\n        B_padded = B\n\n    per_dev = B_padded // N_DEV\n    patches_sharded = patches_jax.reshape(N_DEV, per_dev, *patches_jax.shape[1:])\n\n    multi_gpu_fn = make_multi_gpu_batch_fn(model, ensemble_vars, mode)\n    probs_sharded = multi_gpu_fn(patches_sharded)\n    probs_flat = probs_sharded.reshape(B_padded, *probs_sharded.shape[2:])\n\n    if pad > 0:\n        probs_flat = probs_flat[:B]\n    return probs_flat\n\n# ==============================================================================\n# 6. INFERENCE PIPELINE (WITH POST-PROCESS MODEL)\n# ==============================================================================\n\ndef run_inference_pipeline(model, ensemble_vars):\n    test_dir = DATA_PATH / 'test_images'\n    if not test_dir.exists():\n        print(\"Test directory not found.\")\n        return [], {}\n\n    # Replicate weights once for the whole run\n    sharded_vars = jax.device_put_replicated(ensemble_vars, jax.devices())\n\n    def local_batch_fn(patches_local):\n        return get_ensemble_prob_batch_single_device(model, ensemble_vars, patches_local, ENSEMBLE_MODE)\n\n    multi_gpu_fn = jax.pmap(local_batch_fn, axis_name=\"devices\")\n\n    submission_files = []\n    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🏁 Processing Volume: {vid} | Mode: {ENSEMBLE_MODE}\")\n\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\n        z_max, h, w = vol_input.shape\n        vol_raw = np.zeros((z_max, h, w), dtype=np.float32)\n        vol_prob = np.zeros((z_max, h, w), dtype=np.float32)\n\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        # --- PREDICTION LOOP ---\n        with ThreadPoolExecutor(max_workers=4) as executor:\n            for z in tqdm(range(Z_CONTEXT, z_max - Z_CONTEXT), desc=f\"Predicting Slices\"):\n                coord_chunks = [all_coords[i:i + BATCH_SIZE] for i in range(0, len(all_coords), BATCH_SIZE)]\n                \n                # Initial pre-fetch\n                future_batch = executor.submit(get_batch_worker, z, coord_chunks[0])\n\n                for i in range(len(coord_chunks)):\n                    current_batch_np, batch_coords = future_batch.result()\n\n                    # Pre-fetch next chunk while GPU works\n                    if i + 1 < len(coord_chunks):\n                        future_batch = executor.submit(get_batch_worker, z, coord_chunks[i + 1])\n\n                    b_size = current_batch_np.shape[0]\n                    \n                    # Multi-GPU logic\n                    if b_size % N_DEV == 0 and N_DEV > 1:\n                        gpu_input = current_batch_np.reshape(N_DEV, b_size // N_DEV, PATCH_SIZE, PATCH_SIZE, -1)\n                        prob_out = multi_gpu_fn(gpu_input)\n                        prob_np = np.array(prob_out).reshape(b_size, PATCH_SIZE, PATCH_SIZE)\n                    else:\n                        prob_np = np.array(\n                            get_ensemble_prob_batch_single_device(\n                                model, ensemble_vars, current_batch_np, ENSEMBLE_MODE\n                            )\n                        )[..., 0]\n\n                    # Map results back to volume\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                        \n                        raw_patch = prob_np[idx]\n                        vol_raw[z, y0:y1, x0:x1] = raw_patch\n\n                        if ENABLE_2D_POSTPROC:\n                            # Run 2D Post-Process Model\n                            pp_in = raw_patch[..., None].astype(np.float32)\n                            pp_out = apply_postproc(pp_vars, jnp.array(pp_in[None, ...]))\n                            vol_prob[z, y0:y1, x0:x1] = np.array(pp_out[0, ..., 0])\n                        else:\n                            vol_prob[z, y0:y1, x0:x1] = raw_patch\n\n        # --- FINAL 3D POST-PROCESSING ---\n        print(f\"🪄 Running GPU Post-processing for {vid}...\")\n        t_post_start = time.time()\n\n        if DO_3D_POSTPROCESSING:\n            vol_final = topo_postprocess_3d(vol_prob, TOPO_PARAMS)\n        else:\n            vol_final = (vol_prob > SimpleThreshold).astype(np.uint8) * 255\n\n        # --- MEMORY-SAFE SAVING ---\n        out_name = f\"{vid}.tif\"\n        print(f\"💾 Saving to {out_name}...\")\n        import tifffile\n        tifffile.imwrite(out_name, vol_final, compression='deflate')\n        submission_files.append(out_name)\n\n        # --- DEBUG vs PRODUCTION LOGIC ---\n        if not PRODUCTION_MODE:\n            debug_volumes[vid] = {\n                'input': vol_input,\n                'raw_prob': vol_raw,\n                'postproc_prob': vol_prob,\n                'final_mask': vol_final\n            }\n        \n        # Explicitly release RAM for the next volume in the loop\n        if PRODUCTION_MODE:\n            del vol_input, vol_raw, vol_prob, vol_final\n            import gc\n            gc.collect()\n\n        vol_duration = time.time() - vol_start_time\n        print(f\"✅ Volume {vid} done. Total: {vol_duration:.2f}s\")\n\n    return submission_files, debug_volumes\n\n# ==============================================================================\n# 7. VISUALIZATION OF RESULTS\n# ==============================================================================\n\ndef visualize_results(debug_data):\n    for vid, data in debug_data.items():\n        print(f\"\\n📸 Visualizing Results for Volume: {vid}\")\n\n        vol_in = data['input']\n        vol_raw = data['raw_prob']\n        vol_final = data['final_mask']\n\n        z_max = vol_in.shape[0]\n        indices = [Z_CONTEXT + 2, z_max // 2, z_max - Z_CONTEXT - 2]\n        labels = [\"Beginning\", \"Middle\", \"End\"]\n\n        fig, axes = plt.subplots(3, 4, figsize=(20, 15))\n\n        for i, z in enumerate(indices):\n            axes[i, 0].imshow(vol_in[z], cmap='gray')\n            axes[i, 0].set_title(f\"{labels[i]} (z={z})\\nInput\")\n\n            im2 = axes[i, 1].imshow(vol_raw[z], cmap='magma', vmin=0, vmax=1)\n            axes[i, 1].set_title(\"Raw Ensemble Prob (2D)\")\n            plt.colorbar(im2, ax=axes[i, 1], fraction=0.046, pad=0.04)\n\n            axes[i, 2].imshow(vol_final[z], cmap='gray')\n            axes[i, 2].set_title(\"3D Post-Processed\\n(Hysteresis + Closing)\")\n\n            overlay = np.zeros((*vol_in[z].shape, 3))\n            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] = norm_in\n            overlay[..., 2] = norm_in\n\n            axes[i, 3].imshow(overlay)\n            axes[i, 3].set_title(\"Overlay\")\n\n            for j in range(4):\n                axes[i, j].axis('off')\n\n        plt.tight_layout()\ndef visualize_postproc_comparison(debug_data, vid=None):\n    \"\"\"\n    Shows side-by-side comparison:\n    Input | Raw Ensemble | 2D Postproc | Final 3D Mask\n    \"\"\"\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  # middle slice\n\n    fig, axes = plt.subplots(1, 4, figsize=(28, 6))\n    fig.suptitle(f\"2D Post-Processing Comparison (Volume: {vid}, Slice: {z})\", fontsize=18)\n\n    # 1. Input\n    axes[0].imshow(vol_in[z], cmap='gray')\n    axes[0].set_title(\"Input CT\")\n    axes[0].axis('off')\n\n    # 2. Raw ensemble\n    im1 = axes[1].imshow(vol_raw[z], cmap='magma', vmin=0, vmax=1)\n    axes[1].set_title(\"Raw Ensemble Prob\")\n    plt.colorbar(im1, ax=axes[1], fraction=0.046, pad=0.04)\n    axes[1].axis('off')\n\n    # 3. 2D Postproc UNet\n    im2 = axes[2].imshow(vol_pp[z], cmap='magma', vmin=0, vmax=1)\n    axes[2].set_title(\"2D Postproc Model Output\")\n    plt.colorbar(im2, ax=axes[2], fraction=0.046, pad=0.04)\n    axes[2].axis('off')\n\n    # 4. Final 3D mask\n    axes[3].imshow(vol_final[z], cmap='gray')\n    axes[3].set_title(\"Final 3D Post-Processed Mask\")\n    axes[3].axis('off')\n\n    plt.tight_layout()\n    plt.show()\n\n\n# ==============================================================================\n# 8. MAIN\n# ==============================================================================\n\nif __name__ == \"__main__\":\n    start_total = time.time()\n    netG = Pix2PixGenerator()\n    ensemble_vars = load_ensemble_weights(WEIGHT_PATHS)\n\n    def warmup(model, ensemble_vars):\n        dummy = jnp.zeros((BATCH_SIZE, PATCH_SIZE, PATCH_SIZE, 2 * Z_CONTEXT + 1), jnp.float32)\n        _ = get_ensemble_prob_batch(model, ensemble_vars, dummy, ENSEMBLE_MODE)\n\n    if not ensemble_vars:\n        print(\"❌ Error: No weights loaded.\")\n    else:\n        print(\"🔥 Starting JIT warmup...\")\n        warmup(netG, ensemble_vars)\n        print(\"✅ JIT warmup done.\")\n\n        output_tifs, debug_data = run_inference_pipeline(netG, ensemble_vars)\n\n        # Only run these if we aren't in production\n        if not PRODUCTION_MODE and debug_data:\n            print(\"🔍 Visualizing 2D post-processing comparison...\")\n            visualize_postproc_comparison(debug_data)\n            \n            if ENABLE_VIZ:\n                print(\"📊 Generating 5-view visualization...\")\n                visualize_prediction_samples_jax(debug_data)\n                visualize_results(debug_data)\n    \n        # Always create the zip for submission\n        if output_tifs:\n            with zipfile.ZipFile('submission.zip', 'w') as z:\n                for f in output_tifs:\n                    z.write(f)\n            \n            # Clean up disk space\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-27T14:12:42.282530Z","iopub.execute_input":"2026-02-27T14:12:42.282806Z","iopub.status.idle":"2026-02-27T14:12:42.339539Z","shell.execute_reply.started":"2026-02-27T14:12:42.282779Z","shell.execute_reply":"2026-02-27T14:12:42.338174Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Gemini","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}}]}