{"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":14855360,"datasetId":9484229,"databundleVersionId":15716020},{"sourceType":"modelInstanceVersion","sourceId":769548,"databundleVersionId":15877044,"modelInstanceId":587868,"modelId":599446,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":297302543,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":297674933,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":298023302,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":300093425,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":300258358,"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-03-04T09:37:37.018307Z","iopub.execute_input":"2026-03-04T09:37:37.018553Z","iopub.status.idle":"2026-03-04T09:37:38.207543Z","shell.execute_reply.started":"2026-03-04T09:37:37.018531Z","shell.execute_reply":"2026-03-04T09:37:38.206831Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm import tqdm\nimport cv2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-04T09:37:38.209027Z","iopub.execute_input":"2026-03-04T09:37:38.209343Z","iopub.status.idle":"2026-03-04T09:37:38.558610Z","shell.execute_reply.started":"2026-03-04T09:37:38.209320Z","shell.execute_reply":"2026-03-04T09:37:38.557921Z"}},"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-03-04T09:37:38.559525Z","iopub.execute_input":"2026-03-04T09:37:38.560209Z","iopub.status.idle":"2026-03-04T09:37:44.793487Z","shell.execute_reply.started":"2026-03-04T09:37:38.560176Z","shell.execute_reply":"2026-03-04T09:37:44.792840Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\nprint(sys.version)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-04T09:37:44.794504Z","iopub.execute_input":"2026-03-04T09:37:44.794757Z","iopub.status.idle":"2026-03-04T09:37:44.799553Z","shell.execute_reply.started":"2026-03-04T09:37:44.794723Z","shell.execute_reply":"2026-03-04T09:37:44.798619Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import jax\nprint(f\"Current Device: {jax.devices()[0]}\") ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-04T09:37:44.800336Z","iopub.execute_input":"2026-03-04T09:37:44.800585Z","iopub.status.idle":"2026-03-04T09:37:49.566989Z","shell.execute_reply.started":"2026-03-04T09:37:44.800564Z","shell.execute_reply":"2026-03-04T09:37:49.566170Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # ==========================================\n# # 0. INSTALLATION & IMPORTS\n# # ==========================================\n# import sys\n# #!pip install -q imagecodecs flax optax scikit-image\n\n# import os\n# import zipfile\n# import numpy as np\n# import jax\n# import jax.numpy as jnp\n# import flax\n# import flax.linen as nn\n# import matplotlib.pyplot as plt\n\n# from PIL import Image, ImageSequence\n# from tqdm import tqdm\n# from flax.traverse_util import unflatten_dict\n# from scipy import ndimage\n# from skimage.measure import label, regionprops\n# print('ok')\n# from pathlib import Path","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-04T09:37:49.568043Z","iopub.execute_input":"2026-03-04T09:37:49.568313Z","iopub.status.idle":"2026-03-04T09:37:49.572020Z","shell.execute_reply.started":"2026-03-04T09:37:49.568292Z","shell.execute_reply":"2026-03-04T09:37:49.571307Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Post processing","metadata":{}},{"cell_type":"markdown","source":"* T_low\tThe Bridge Builder. Lowering this allows the mask to \"flow\" into fainter areas, provided they are connected to a strong start. <br>📉 Lower it (0.1 - 0.3) to fix broken lines.\n* T_high\tThe Anchor. Only pixels above this start a new segment. <br>📈 Keep it high (0.8 - 0.95) to prevent noise from creating new fake sheets.\n* xy_radius\tThe Glue. Determines how far to reach to connect two nearby pixels on the same image. <br>📈 Increase (2-4) to close gaps in lines.\n* z_radius\tThe 3D Glue. Determines how far to reach across slices. <br>📈 Increase (1-2) if your 3D volume feels disjointed vertically.\n* dust_min_size\tThe Cleaner. Removes isolated objects smaller than X voxels. <br>📈 Increase (500-2000) to remove that dotted line on the left edge.\n","metadata":{}},{"cell_type":"code","source":"import os\nimport sys\nimport zipfile\nimport time\nfrom concurrent.futures import ThreadPoolExecutor\nfrom pathlib import Path\nfrom functools import partial\nfrom tqdm import tqdm\nfrom pathlib import Path\n\nimport numpy as np\nimport jax\nimport jax.numpy as jnp\nfrom jax import jit\nfrom functools import partial\nfrom flax import linen as nn \nimport flax.linen as nn\nfrom flax.traverse_util import unflatten_dict\nfrom flax.training import train_state\nimport tifffile as tiff\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 = True  # 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=False\nENSEMBLE_MODE = \"hybrid\"   # '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_10.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_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        \"/kaggle/input/notebooks/crischir/tpu-training-jax-pix2pix/checkpoints/pix2pix_ep_80.npz\",\n        \"/kaggle/input/notebooks/crischir/tpu-training-jax-pix2pix/checkpoints/pix2pix_ep_90.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_ep_60.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_60.npz\",\n        \"/kaggle/input/notebooks/crischir/tpu-training-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    \"Cl-dice\": [\n        \"/kaggle/input/models/crischir/vesuvius-2-5-mask-cldice/jax/300epochs-9samplefof-file-cldice-loss-best-220-230/1/checkpoints/pix2pix_best.npz\",\n        \"/kaggle/input/models/crischir/vesuvius-2-5-mask-cldice/jax/300epochs-9samplefof-file-cldice-loss-best-220-230/1/checkpoints/pix2pix_ep_170.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    # \"Cl-dice\",\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.26,\n    \"T_high\": 0.55,\n    \"z_radius\": 2,\n    \"xy_radius\": 0,\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\ndef make_fused_predict_fn(model, pp_model, mode=\"hybrid\"):\n    \"\"\"\n    stacked_ensemble: {\n        \"params\":  PyTree with leading dim = num_models\n        \"batch_stats\": PyTree with leading dim = num_models\n    }\n    pp_vars: {\"params\": ..., \"batch_stats\": ...}\n    x: (B, H, W, C)\n    \"\"\"\n\n    @partial(jax.jit, static_argnums=(3,))\n    def _predict(stacked_ensemble, pp_vars, x, num_models):\n        # carry: (acc_logits, acc_max_prob, votes)\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.int32),\n        )\n\n        def scan_body(carry, vars_i):\n            acc_l, acc_m_p, vts = carry\n            params_i, stats_i = vars_i\n\n            vars_dict = {\n                \"params\": params_i,\n                \"batch_stats\": stats_i,\n            }\n\n            logits = model.apply(vars_dict, x, train=False)\n            p = nn.sigmoid(logits)\n\n            acc_l = acc_l + logits\n            acc_m_p = jnp.maximum(acc_m_p, p)\n            vts = vts + (p > 0.1).astype(jnp.int32)\n\n            return (acc_l, acc_m_p, vts), None\n\n        # stacked_ensemble[\"params\"] and [\"batch_stats\"] each have leading dim = num_models\n        (acc_l, acc_m_p, vts), _ = jax.lax.scan(\n            scan_body,\n            init_carry,\n            (stacked_ensemble[\"params\"], stacked_ensemble[\"batch_stats\"])\n        )\n\n        # --- aggregation ---\n        p_mean = nn.sigmoid(acc_l / num_models)\n\n        if mode == \"max\":\n            raw_p = acc_m_p\n        elif mode == \"hybrid\":\n            raw_p = p_mean + ALPHA * (acc_m_p - p_mean)\n        elif mode == \"agreement\":\n            raw_p = jnp.where(vts >= AGREEMENT_THRESHOLD, 1.0, p_mean)\n        else:\n            raw_p = p_mean\n\n        # --- 2D post‑processing (same kernel, still on device) ---\n        pp_logits = pp_model.apply(pp_vars, raw_p, train=False)\n        final_p = nn.sigmoid(pp_logits)\n\n        return raw_p, final_p\n\n    return _predict\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]\ndef make_unified_predict_fn(model, pp_model, mode=\"hybrid\"):\n    \"\"\"Creates a single JIT-able function for Ensemble + Post-Processing.\"\"\"\n    \n    @partial(jax.jit, static_argnums=(3,))\n    def _predict(ensemble_vars, pp_vars, x, num_models):\n        # Initial carry for ensemble aggregation\n        init_carry = (\n            jnp.zeros((*x.shape[:-1], 1), dtype=jnp.float32), # accum_logits\n            jnp.zeros((*x.shape[:-1], 1), dtype=jnp.float32), # accum_max_p\n            jnp.zeros((*x.shape[:-1], 1), dtype=jnp.int32)    # votes\n        )\n\n        def scan_body(carry, v):\n            acc_l, acc_m_p, vts = carry\n            logits = model.apply(v, x, train=False)\n            p = nn.sigmoid(logits)\n            return (acc_l + logits, jnp.maximum(acc_m_p, p), vts + (p > 0.1).astype(jnp.int32)), None\n\n        # 1. ENSEMBLE STEP (Inside TPU)\n        (acc_l, acc_m_p, vts), _ = jax.lax.scan(scan_body, init_carry, ensemble_vars)\n        \n        # Calculate raw probability based on mode\n        p_mean = nn.sigmoid(acc_l / num_models)\n        if mode == \"max\": raw_p = acc_m_p\n        elif mode == \"hybrid\": raw_p = p_mean + ALPHA * (acc_m_p - p_mean)\n        elif mode == \"agreement\": raw_p = jnp.where(vts >= AGREEMENT_THRESHOLD, 1.0, p_mean)\n        else: raw_p = p_mean\n\n        # 2. 2D POST-PROC STEP (Also inside TPU, no data transfer)\n        pp_logits = pp_model.apply(pp_vars, raw_p, train=False)\n        final_p = nn.sigmoid(pp_logits)\n        \n        return raw_p, final_p\n\n    return _predict\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# ==============================================================================\ndef run_inference_pipeline(model, ensemble_vars):\n    test_dir = DATA_PATH / \"test_images\"\n    if not test_dir.exists():\n        return [], {}\n\n    # ------------------------------------------------------------------\n    # 1. Build stacked ensemble (leading dim = num_models)\n    # ------------------------------------------------------------------\n    num_models = len(ensemble_vars)\n\n    stacked_params = jax.tree_util.tree_map(\n        lambda *xs: jnp.stack(xs, axis=0),\n        *[v[\"params\"] for v in ensemble_vars]\n    )\n    stacked_stats = jax.tree_util.tree_map(\n        lambda *xs: jnp.stack(xs, axis=0),\n        *[v[\"batch_stats\"] for v in ensemble_vars]\n    )\n\n    stacked_ensemble = {\n        \"params\": stacked_params,\n        \"batch_stats\": stacked_stats,\n    }\n\n    # pp_vars already defined globally: {\"params\": pp_params, \"batch_stats\": pp_stats}\n    fused_fn = make_fused_predict_fn(model, postproc_model, ENSEMBLE_MODE)\n\n    # ------------------------------------------------------------------\n    # 2. Multi‑device wrapper (only x is sharded)\n    # ------------------------------------------------------------------\n    @partial(jax.pmap, in_axes=(None, None, 0), out_axes=0)\n    def multi_dev_predict(stacked_ensemble, pp_vars, x):\n        # x: (n_devices, per_device_B, H, W, C)\n        return fused_fn(stacked_ensemble, pp_vars, x, num_models)\n\n    submission_files = []\n    debug_volumes = {}\n\n    test_ids = [f.stem for f in test_dir.glob(\"*.tif\")]\n    devices = jax.devices()\n    n_dev = len(devices)\n\n    for vid in test_ids:\n        vol_start_time = time.time()\n        print(f\"\\n🏁 Processing Volume: {vid} | Devices: {n_dev}\")\n\n        # --------------------------------------------------------------\n        # 3. Load volume\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)\n                             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\n                patch = vol_input[\n                    z_idx - Z_CONTEXT : z_idx + Z_CONTEXT + 1,\n                    y0:y1,\n                    x0:x1,\n                ].transpose(1, 2, 0)  # (H, W, C)\n\n                batch_patches.append(normalize_patch(patch))\n\n            return np.stack(batch_patches), coords_subset\n\n        # --------------------------------------------------------------\n        # 4. Slice‑wise prediction loop\n        # --------------------------------------------------------------\n        with ThreadPoolExecutor(max_workers=4) as executor:\n            for z in tqdm(range(Z_CONTEXT, z_max - Z_CONTEXT), desc=\"Predicting Slices\"):\n                coord_chunks = [\n                    all_coords[i : i + BATCH_SIZE]\n                    for i in range(0, len(all_coords), BATCH_SIZE)\n                ]\n\n                if not coord_chunks:\n                    continue\n\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                    if i + 1 < len(coord_chunks):\n                        future_batch = executor.submit(\n                            get_batch_worker, z, coord_chunks[i + 1]\n                        )\n\n                    b_size = current_batch_np.shape[0]\n\n                    # pad so batch is divisible by n_dev\n                    pad_len = (n_dev - (b_size % n_dev)) % n_dev\n                    if pad_len > 0:\n                        current_batch_np = np.pad(\n                            current_batch_np,\n                            ((0, pad_len), (0, 0), (0, 0), (0, 0)),\n                            mode=\"constant\",\n                        )\n\n                    # reshape to (n_dev, per_dev_B, H, W, C)\n                    per_dev_B = current_batch_np.shape[0] // n_dev\n                    gpu_in = jnp.array(current_batch_np).reshape(\n                        n_dev, per_dev_B, PATCH_SIZE, PATCH_SIZE, 2 * Z_CONTEXT + 1\n                    )\n\n                    # fused ensemble + post‑proc\n                    raw_out, pp_out = multi_dev_predict(stacked_ensemble, pp_vars, gpu_in)\n\n                    # back to CPU, flatten batch\n                    raw_np = np.array(raw_out).reshape(-1, PATCH_SIZE, PATCH_SIZE)\n                    pp_np = np.array(pp_out).reshape(-1, PATCH_SIZE, PATCH_SIZE)\n\n                    # remove padding\n                    raw_np = raw_np[:b_size]\n                    pp_np = pp_np[:b_size]\n\n                    # write patches back into volume\n                    for (y, x), raw_patch, pp_patch in zip(batch_coords, raw_np, pp_np):\n                        y1, x1 = min(y + PATCH_SIZE, h), min(x + PATCH_SIZE, w)\n                        y0, x0 = y1 - PATCH_SIZE, x1 - PATCH_SIZE\n\n                        vol_raw[z, y0:y1, x0:x1] = raw_patch[: (y1 - y0), : (x1 - x0)]\n                        vol_prob[z, y0:y1, x0:x1] = pp_patch[: (y1 - y0), : (x1 - x0)]\n\n        # --------------------------------------------------------------\n        # 5. 3D topology post‑processing (CPU or GPU version)\n        # --------------------------------------------------------------\n        if DO_3D_POSTPROCESSING:\n            try:\n                final_mask = topo_postprocess_3d_gpu(vol_prob, TOPO_PARAMS)\n            except Exception:\n                final_mask = topo_postprocess_3d(vol_prob, TOPO_PARAMS)\n        else:\n            final_mask = (vol_prob >= SimpleThreshold).astype(np.uint8) * 255\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        # optional debug storage\n        if ENABLE_VIZ:\n            debug_volumes[vid] = {\n                \"input\": vol_input,\n                \"raw_prob\": vol_raw,\n                \"postproc_prob\": vol_prob,\n                \"final_mask\": final_mask,\n            }\n\n        # save TIF, RLE, etc. (keep your existing saving logic here)\n        # submission_files.append(...)\n\n        print(f\"✅ Volume {vid} done in {time.time() - vol_start_time:.1f}s\")\n        \n    return submission_files, debug_volumes\n\n\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\ndef run_inference_pipeline(model, ensemble_vars):\n    test_dir = DATA_PATH / \"test_images\"\n    if not test_dir.exists():\n        print(\"⚠️ test_images nu există, întorc doar numele implicit de zip.\")\n        return [], {}, \"submission.zip\"\n\n    # 1. Build stacked ensemble\n    num_models = len(ensemble_vars)\n\n    stacked_params = jax.tree_util.tree_map(\n        lambda *xs: jnp.stack(xs, axis=0),\n        *[v[\"params\"] for v in ensemble_vars]\n    )\n    stacked_stats = jax.tree_util.tree_map(\n        lambda *xs: jnp.stack(xs, axis=0),\n        *[v[\"batch_stats\"] for v in ensemble_vars]\n    )\n\n    stacked_ensemble = {\n        \"params\": stacked_params,\n        \"batch_stats\": stacked_stats,\n    }\n\n    fused_fn = make_fused_predict_fn(model, postproc_model, ENSEMBLE_MODE)\n\n    @partial(jax.pmap, in_axes=(None, None, 0), out_axes=0)\n    def multi_dev_predict(stacked_ensemble, pp_vars, x):\n        return fused_fn(stacked_ensemble, pp_vars, x, num_models)\n\n    submission_files = []\n    debug_volumes = {}\n\n    test_ids = [f.stem for f in test_dir.glob(\"*.tif\")]\n    devices = jax.devices()\n    n_dev = len(devices)\n\n    out_dir = Path(\"predictions\")\n    out_dir.mkdir(exist_ok=True, parents=True)\n\n    for vid in test_ids:\n        vol_start_time = time.time()\n        print(f\"\\n🏁 Processing Volume: {vid} | Devices: {n_dev}\")\n\n        # 3. Load volume\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)\n                             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\n                patch = vol_input[\n                    z_idx - Z_CONTEXT : z_idx + Z_CONTEXT + 1,\n                    y0:y1,\n                    x0:x1,\n                ].transpose(1, 2, 0)  # (H, W, C)\n\n                batch_patches.append(normalize_patch(patch))\n\n            return np.stack(batch_patches), coords_subset\n\n        # 4. Slice-wise prediction loop\n        with ThreadPoolExecutor(max_workers=4) as executor:\n            for z in tqdm(range(Z_CONTEXT, z_max - Z_CONTEXT), desc=\"Predicting Slices\"):\n                coord_chunks = [\n                    all_coords[i : i + BATCH_SIZE]\n                    for i in range(0, len(all_coords), BATCH_SIZE)\n                ]\n\n                if not coord_chunks:\n                    continue\n\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                    if i + 1 < len(coord_chunks):\n                        future_batch = executor.submit(\n                            get_batch_worker, z, coord_chunks[i + 1]\n                        )\n\n                    b_size = current_batch_np.shape[0]\n\n                    pad_len = (n_dev - (b_size % n_dev)) % n_dev\n                    if pad_len > 0:\n                        current_batch_np = np.pad(\n                            current_batch_np,\n                            ((0, pad_len), (0, 0), (0, 0), (0, 0)),\n                            mode=\"constant\",\n                        )\n\n                    per_dev_B = current_batch_np.shape[0] // n_dev\n                    gpu_in = jnp.array(current_batch_np).reshape(\n                        n_dev, per_dev_B, PATCH_SIZE, PATCH_SIZE, 2 * Z_CONTEXT + 1\n                    )\n\n                    raw_out, pp_out = multi_dev_predict(stacked_ensemble, pp_vars, gpu_in)\n\n                    raw_np = np.array(raw_out).reshape(-1, PATCH_SIZE, PATCH_SIZE)\n                    pp_np = np.array(pp_out).reshape(-1, PATCH_SIZE, PATCH_SIZE)\n\n                    raw_np = raw_np[:b_size]\n                    pp_np = pp_np[:b_size]\n\n                    for (y, x), raw_patch, pp_patch in zip(batch_coords, raw_np, pp_np):\n                        y1, x1 = min(y + PATCH_SIZE, h), min(x + PATCH_SIZE, w)\n                        y0, x0 = y1 - PATCH_SIZE, x1 - PATCH_SIZE\n\n                        vol_raw[z, y0:y1, x0:x1] = raw_patch[: (y1 - y0), : (x1 - x0)]\n                        vol_prob[z, y0:y1, x0:x1] = pp_patch[: (y1 - y0), : (x1 - x0)]\n\n        # 5. 3D topology post-processing\n        if DO_3D_POSTPROCESSING:\n            try:\n                final_mask = topo_postprocess_3d_gpu(vol_prob, TOPO_PARAMS)\n            except Exception:\n                final_mask = topo_postprocess_3d(vol_prob, TOPO_PARAMS)\n        else:\n            final_mask = (vol_prob >= SimpleThreshold).astype(np.uint8) * 255\n\n        final_mask_uint8 = final_mask.astype(np.uint8)\n\n        # 6. Save TIF\n        out_path = out_dir / f\"{vid}.tif\"\n        tiff.imwrite(str(out_path), final_mask_uint8)\n        submission_files.append(str(out_path))\n\n        if ENABLE_VIZ:\n            debug_volumes[vid] = {\n                \"input\": vol_input,\n                \"raw_prob\": vol_raw,\n                \"postproc_prob\": vol_prob,\n                \"final_mask\": final_mask_uint8,\n            }\n\n        print(f\"✅ Volume {vid} done in {time.time() - vol_start_time:.1f}s\")\n\n    zip_path = \"submission.zip\"\n    return submission_files, debug_volumes, zip_path\n\n# ==============================================================================\n# 7. MAIN EXECUTION\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 found. Please check your SELECTED keys and WEIGHT_PATHS.\")\n    else:\n        output_tifs, debug_data, zip_path = run_inference_pipeline(netG, ensemble_vars)\n\n        if output_tifs:\n            print(f\"📦 Packaging {len(output_tifs)} files into {zip_path}...\")\n            with zipfile.ZipFile(zip_path, \"w\", zipfile.ZIP_DEFLATED) as z:\n                for f in output_tifs:\n                    if os.path.exists(f):\n                        z.write(f, arcname=os.path.basename(f))\n            print(f\"✅ {zip_path} created successfully.\")\n\n            if ENABLE_VIZ and debug_data:\n                try:\n                    visualize_prediction_samples_jax(debug_data)\n                except Exception as e:\n                    print(f\"⚠️ Visualization failed: {e}\")\n\n            for f in output_tifs:\n                if os.path.exists(f):\n                    os.remove(f)\n            print(\"🧹 Temporary .tif files removed.\")\n        else:\n            print(\"⚠️ No output files were generated. Check if test images exist in DATA_PATH.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-04T09:45:37.405111Z","iopub.execute_input":"2026-03-04T09:45:37.406190Z","iopub.status.idle":"2026-03-04T09:46:26.964284Z","shell.execute_reply.started":"2026-03-04T09:45:37.406162Z","shell.execute_reply":"2026-03-04T09:46:26.963676Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}