{"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,"isSourceIdPinned":true},{"sourceType":"modelInstanceVersion","sourceId":769548,"databundleVersionId":15877044,"modelInstanceId":587868,"modelId":599446},{"sourceType":"modelInstanceVersion","sourceId":769070,"databundleVersionId":15871980,"modelInstanceId":587120,"modelId":599446},{"sourceType":"modelInstanceVersion","sourceId":768516,"databundleVersionId":15864953,"modelInstanceId":587120,"modelId":599446},{"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},{"sourceType":"kernelVersion","sourceId":300710378}],"dockerImageVersionId":31259,"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-01T09:16:32.184646Z","iopub.execute_input":"2026-03-01T09:16:32.184908Z","iopub.status.idle":"2026-03-01T09:16:33.388834Z","shell.execute_reply.started":"2026-03-01T09:16:32.184873Z","shell.execute_reply":"2026-03-01T09:16:33.388043Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"","metadata":{}},{"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-01T09:16:33.389722Z","iopub.execute_input":"2026-03-01T09:16:33.390044Z","iopub.status.idle":"2026-03-01T09:16:39.422411Z","shell.execute_reply.started":"2026-03-01T09:16:33.390021Z","shell.execute_reply":"2026-03-01T09:16:39.421767Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\nprint(sys.version)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T09:16:39.424312Z","iopub.execute_input":"2026-03-01T09:16:39.424584Z","iopub.status.idle":"2026-03-01T09:16:39.428847Z","shell.execute_reply.started":"2026-03-01T09:16:39.424548Z","shell.execute_reply":"2026-03-01T09:16:39.428050Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm import tqdm\nimport cv2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T09:16:39.429767Z","iopub.execute_input":"2026-03-01T09:16:39.430003Z","iopub.status.idle":"2026-03-01T09:16:39.764308Z","shell.execute_reply.started":"2026-03-01T09:16:39.429982Z","shell.execute_reply":"2026-03-01T09:16:39.763788Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport sys\nimport zipfile\nimport time\nfrom tqdm import tqdm\nfrom concurrent.futures import ThreadPoolExecutor\nfrom pathlib import Path\nimport numpy as np\nimport jax\nimport jax.numpy as jnp\nfrom jax import jit\nfrom functools import partial\nimport flax.linen as nn\nfrom flax.traverse_util import unflatten_dict\nfrom flax.training import train_state\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\n\n# ==============================================================================\n# 0. QUICK RUNTIME INFO\n# ==============================================================================\nprint(\"Python:\", sys.version)\nprint(\"JAX devices:\", jax.devices())\n\n# ==============================================================================\n# 1. VISUALIZATION HELPER (5-VIEW)\n# ==============================================================================\ndef visualize_prediction_samples_jax(debug_data, vid=None):\n    if vid is None:\n        vid = list(debug_data.keys())[0]\n\n    data = debug_data[vid]\n    vol_in = data['input']\n    vol_raw = data['raw_prob']      \n    vol_pp = data['postproc_prob']  \n    vol_final = data['final_mask']  \n\n    z = vol_in.shape[0] // 2\n    \n    p2, p98 = np.percentile(vol_in[z], (2, 98))\n    orig_norm = np.clip((vol_in[z] - p2) / (p98 - p2 + 1e-7), 0, 1)\n\n    fig, axes = plt.subplots(1, 5, figsize=(30, 6))\n    fig.suptitle(f\"Pipeline Analysis: Weight-based Post-Processing (Slice: {z})\", fontsize=16)\n\n    axes[0].imshow(orig_norm, cmap='gray')\n    axes[0].set_title(\"1. Original CT\")\n\n    im1 = axes[1].imshow(vol_raw[z], cmap='magma', vmin=0, vmax=1)\n    axes[1].set_title(\"2. Ensemble Output\\n(Pre-Weights)\")\n    plt.colorbar(im1, ax=axes[1], fraction=0.046, pad=0.04)\n\n    im2 = axes[2].imshow(vol_pp[z], cmap='magma', vmin=0, vmax=1)\n    axes[2].set_title(\"3. 2D Post-Proc Model\\n(With Loaded Weights)\")\n    plt.colorbar(im2, ax=axes[2], fraction=0.046, pad=0.04)\n\n    diff = vol_pp[z] - vol_raw[z]\n    im3 = axes[3].imshow(diff, cmap='coolwarm', vmin=-0.5, vmax=0.5)\n    axes[3].set_title(\"4. Weight Impact\\n(PP - Raw)\")\n    plt.colorbar(im3, ax=axes[3], fraction=0.046, pad=0.04)\n\n    axes[4].imshow(vol_final[z], cmap='gray')\n    axes[4].set_title(\"5. Final 3D Mask\")\n\n    for ax in axes:\n        ax.axis('off')\n\n    plt.tight_layout()\n    plt.show()\n# ==============================================================================\n# 0. Debug mode \n# ==============================================================================\nPRODUCTION_MODE = False  # Set to False when debugging/visualizing locally Without poduction mode it will crash with real test files \nENABLE_VIZ = not PRODUCTION_MODE\n# ==============================================================================\n# 2. CONFIGURATION\n# ==============================================================================\nTTA_INFERENCE = True  # TTA option: True or False\n\nENABLE_2D_POSTPROC = False\nENSEMBLE_MODE = \"hybrid\"   # 'mean', 'max', 'agreement', 'hybrid', 'single'\nAGREEMENT_THRESHOLD = 11\nALPHA = 0.8\n\nMODEL_CHECKPOINTS = {\n    \"2-5d-filter-training-pix2pix\": [\n        # \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix/checkpoints/pix2pix_ep_20.npz\",\n        \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix/checkpoints/pix2pix_ep_30.npz\",\n    ],\n    \"2-5d-filter-training-pix2pix-pred-mask\": [\n        \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix-pred-mask/checkpoints/pix2pix_best.npz\",\n        # \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix-pred-mask/checkpoints/pix2pix_ep_70.npz\",\n        \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix-pred-mask/checkpoints/pix2pix_ep_90.npz\",\n        \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix-pred-mask/checkpoints/pix2pix_ep_120.npz\",\n        # \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix-pred-mask/checkpoints/pix2pix_ep_170.npz\",\n    ],\n    \"test-code-for-jax-pix2pix\": [\n        \"/kaggle/input/notebooks/crischir/test-code-for-jax-pix2pix/checkpoints/pix2pix_best.npz\",\n       #  \"/kaggle/input/notebooks/crischir/test-code-for-jax-pix2pix/checkpoints/pix2pix_ep_20.npz\",\n       # \"/kaggle/input/notebooks/crischir/test-code-for-jax-pix2pix/checkpoints/pix2pix_ep_30.npz\",\n       #  \"/kaggle/input/notebooks/crischir/test-code-for-jax-pix2pix/checkpoints/pix2pix_ep_40.npz\",\n        \"/kaggle/input/notebooks/crischir/test-code-sobel-jax-pix2pix/checkpoints/pix2pix_ep_110.npz\",\n        \"/kaggle/input/notebooks/crischir/test-code-sobel-jax-pix2pix/checkpoints/pix2pix_ep_120.npz\",\n       # \"/kaggle/input/notebooks/crischir/test-code-sobel-jax-pix2pix/checkpoints/pix2pix_ep_130.npz\",\n       #  \"/kaggle/input/notebooks/crischir/test-code-sobel-jax-pix2pix/checkpoints/pix2pix_ep_140.npz\",\n        \n    ],\n    \"test-code-sobel-jax-pix2pix\": [\n        \"/kaggle/input/notebooks/crischir/test-code-sobel-jax-pix2pix/checkpoints/pix2pix_ep_100.npz\",\n        # \"/kaggle/input/notebooks/crischir/test-code-sobel-jax-pix2pix/checkpoints/pix2pix_ep_130.npz\",\n    ],\n    \"tpu-training-jax-pix2pix\": [\n        \"/kaggle/input/notebooks/crischir/tpu-training-jax-pix2pix/checkpoints/pix2pix_ep_30.npz\",\n        \"/kaggle/input/notebooks/crischir/tpu-training-jax-pix2pix/checkpoints/pix2pix_ep_40.npz\",\n        \"//kaggle/input/notebooks/crischir/tpu-training-jax-pix2pix/checkpoints/pix2pix_ep_50.npz\",\n        \"/kaggle/input/notebooks/crischir/tpu-training-jax-pix2pix/checkpoints/pix2pix_ep_60.npz\",\n        \"/kaggle/input/notebooks/crischir/tpu-training-jax-pix2pix/checkpoints/pix2pix_ep_70.npz\",\n    ],\n    \"prewitt-loss-model-jax-pix2pix\": [\n        \"/kaggle/input/notebooks/crischir/prewitt-loss-model-jax-pix2pix/checkpoints/pix2pix_best.npz\",\n        # \"/kaggle/input/notebooks/crischir/prewitt-loss-model-jax-pix2pix/checkpoints/pix2pix_ep_40.npz\",\n        # \"/kaggle/input/notebooks/crischir/prewitt-loss-model-jax-pix2pix/checkpoints/pix2pix_ep_50.npz\",\n        # \"/kaggle/input/notebooks/crischir/prewitt-loss-model-jax-pix2pix/checkpoints/pix2pix_ep_70.npz\",\n    ],\n    \"jax-pix2pix\": [\n        # \"/kaggle/input/notebooks/crischir/jax-pix2pix/checkpoints/pix2pix_ep_10.npz\",\n        \"/kaggle/input/notebooks/crischir/jax-pix2pix/checkpoints/pix2pix_ep_30.npz\",\n        # \"/kaggle/input/notebooks/crischir/jax-pix2pix/checkpoints/pix2pix_ep_30.npz\",\n        # \"/kaggle/input/notebooks/crischir/jax-pix2pix/checkpoints/pix2pix_ep_50.npz\",\n        # \"/kaggle/input/notebooks/crischir/jax-pix2pix/checkpoints/pix2pix_ep_70.npz\",\n        # \"/kaggle/input/notebooks/crischir/jax-pix2pix/checkpoints/pix2pix_ep_120.npz\",\n        # \"/kaggle/input/notebooks/crischir/jax-pix2pix/checkpoints/pix2pix_ep_140.npz\",\n        # \"/kaggle/input/notebooks/crischir/jax-pix2pix/checkpoints/pix2pix_ep_160.npz\",\n    ],\n    \"clDice-1\": [\n       #  \"/kaggle/input/models/crischir/vesuvius-2-5-mask-cldice/flax/2.5d-cldice-24-samples-40-epochs/1/checkpoints/pix2pix_ep_10.npz\",\n       #  \"/kaggle/input/models/crischir/vesuvius-2-5-mask-cldice/flax/2.5d-cldice-24-samples-40-epochs/1/checkpoints/pix2pix_ep_20.npz\",\n       # \"/kaggle/input/models/crischir/vesuvius-2-5-mask-cldice/flax/2.5d-cldice-24-samples-40-epochs/1/checkpoints/pix2pix_ep_30.npz\",\n        \"/kaggle/input/models/crischir/vesuvius-2-5-mask-cldice/flax/2.5d-cldice-24-samples-40-epochs/1/checkpoints/pix2pix_ep_40.npz\",\n       #  \"/kaggle/input/models/crischir/vesuvius-2-5-mask-cldice/flax/2.5d-cldice-24-samples-40-epochs/1/checkpoints/pix2pix_ep_50.npz\",\n       # \"/kaggle/input/models/crischir/vesuvius-2-5-mask-cldice/flax/2.5d-cldice-24-samples-40-epochs/1/checkpoints/pix2pix_ep_60.npz\",\n    ],\n    \"clDice-2\": [\n        # \"/kaggle/input/notebooks/crischir/cldice-tpu-training-pix2pix/checkpoints/pix2pix_ep_10.npz\",\n        # \"/kaggle/input/notebooks/crischir/cldice-tpu-training-pix2pix/checkpoints/pix2pix_ep_20.npz\",\n        # \"/kaggle/input/notebooks/crischir/cldice-tpu-training-pix2pix/checkpoints/pix2pix_ep_30.npz\",\n        \"/kaggle/input/notebooks/crischir/cldice-tpu-training-pix2pix/checkpoints/pix2pix_ep_40.npz\",\n        # \"/kaggle/input/notebooks/crischir/cldice-tpu-training-pix2pix/checkpoints/pix2pix_ep_50.npz\",\n        # \"/kaggle/input/notebooks/crischir/cldice-tpu-training-pix2pix/checkpoints/pix2pix_ep_60.npz\",\n        \n    ],\n        \"clDice-3\": [\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_30.npz\",\n        # \"/kaggle/input/models/crischir/vesuvius-2-5-mask-cldice/jax/300epochs-9samplefof-file-cldice-loss-best-220-230/1/checkpoints/pix2pix_ep_120.npz\",\n        \"/kaggle/input/models/crischir/vesuvius-2-5-mask-cldice/jax/300epochs-9samplefof-file-cldice-loss-best-220-230/1/checkpoints/pix2pix_ep_220.npz\",\n        # \"/kaggle/input/models/crischir/vesuvius-2-5-mask-cldice/jax/300epochs-9samplefof-file-cldice-loss-best-220-230/1/checkpoints/pix2pix_ep_280.npz\",\n        \"/kaggle/input/models/crischir/vesuvius-2-5-mask-cldice/jax/300epochs-9samplefof-file-cldice-loss-best-220-230/1/checkpoints/pix2pix_ep_300.npz\",\n        \n    ],\n        \"retromasking-vesuvius\": [\n        \"/kaggle/input/datasets/crischir/retromasking-vesuvius-2-5d-filter-checkpoints/checkpoints/pix2pix_ep_30.npz\",\n        \"/kaggle/input/datasets/crischir/retromasking-vesuvius-2-5d-filter-checkpoints/checkpoints/pix2pix_ep_60.npz\",\n        # \"/kaggle/input/datasets/crischir/retromasking-vesuvius-2-5d-filter-checkpoints/checkpoints/pix2pix_ep_90.npz\",\n        # \"/kaggle/input/datasets/crischir/retromasking-vesuvius-2-5d-filter-checkpoints/checkpoints/pix2pix_ep_120.npz\",\n        # \"/kaggle/input/datasets/crischir/retromasking-vesuvius-2-5d-filter-checkpoints/checkpoints/pix2pix_ep_120.npz\",\n        \"/kaggle/input/datasets/crischir/retromasking-vesuvius-2-5d-filter-checkpoints/checkpoints/pix2pix_ep_150.npz\",\n        \n    ],\n    \"vesuvius-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    # \"2-5d-filter-training-pix2pix\",\n    # \"test-code-sobel-jax-pix2pix\",\n    # \"prewitt-loss-model-jax-pix2pix\",\n    # \"jax-pix2pix\",\n    # \"vesuvius-checkpoint-dataset\",\n    \"clDice-1\",\n    # \"clDice-2\",\n    \"clDice-3\",\n    \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.10,\n    \"T_high\": 0.55,\n    \"z_radius\": 1,\n    \"xy_radius\": 1,\n    \"dust_min_size\": 50,\n}\n\nDATA_PATH = Path(\"/kaggle/input/competitions/vesuvius-challenge-surface-detection/\")\nPATCH_SIZE = 256\nZ_CONTEXT = 3\nSimpleThreshold = 0.1\nBATCH_SIZE = 256\ndevices = jax.devices()\nN_DEV = len(devices)\nprint(\"Num devices:\", N_DEV)\n\n# ==============================================================================\n# 3. HELPERS\n# ==============================================================================\ndef to_skeleton(vol_mask):\n    skel = skeletonize(vol_mask > 0)\n    return skel.astype(np.uint8) * 255\n\ndef build_anisotropic_struct(z_radius: int, xy_radius: int):\n    z, r = int(z_radius), int(xy_radius)\n    if z == 0 and r == 0: return None\n    if z == 0 and r > 0:\n        size = 2 * r + 1\n        struct = np.zeros((1, size, size), dtype=bool)\n        cy = cx = r\n        for dy in range(-r, r + 1):\n            for dx in range(-r, r + 1):\n                if dy * dy + dx * dx <= r * r: struct[0, cy + dy, cx + dx] = True\n        return struct\n    if z > 0 and r == 0:\n        struct = np.zeros((2 * z + 1, 1, 1), dtype=bool)\n        struct[:, 0, 0] = True\n        return struct\n    depth, size = 2 * z + 1, 2 * r + 1\n    struct = np.zeros((depth, size, size), dtype=bool)\n    cz, cy, cx = z, r, r\n    for dz in range(-z, z + 1):\n        for dy in range(-r, r + 1):\n            for dx in range(-r, r + 1):\n                if dy * dy + dx * dx <= r * r: struct[cz + dz, cy + dy, cx + dx] = True\n    return struct\n\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    if not strong.any(): return np.zeros_like(prob_vol, dtype=np.uint8)\n    struct_hyst = ndi.generate_binary_structure(3, 3)\n    mask = ndi.binary_propagation(strong, mask=weak, structure=struct_hyst)\n    if not mask.any(): return np.zeros_like(prob_vol, dtype=np.uint8)\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    if params[\"dust_min_size\"] > 0:\n        mask = remove_small_objects(mask.astype(bool), min_size=params[\"dust_min_size\"])\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    data = mask[None, ..., None].astype(jnp.float32)\n    kernel_jax = kernel_norm[..., None, None]\n    dn = jax.lax.ConvDimensionNumbers(\n        lhs_spec=(0, 1, 2, 3, 4), rhs_spec=(0, 1, 2, 3, 4), out_spec=(0, 1, 2, 3, 4)\n    )\n    dilated = jax.lax.conv_general_dilated(data, kernel_jax, (1, 1, 1), 'SAME', dimension_numbers=dn)\n    dilated_mask = (dilated > 0).astype(jnp.float32)\n    eroded = jax.lax.conv_general_dilated(dilated_mask, kernel_jax, (1, 1, 1), 'SAME', dimension_numbers=dn)\n    return (eroded >= 0.99).astype(jnp.uint8)[0, ..., 0]\n\ndef normalize_patch(patch: np.ndarray) -> np.ndarray:\n    patch = patch.astype(np.float32)\n    return (patch - patch.mean()) / (patch.std() + 1e-6)\n\ndef load_npz_weights(path):\n    with np.load(path, allow_pickle=False) as data:\n        flat_dict = {k: v for k, v in data.items()}\n    params_flat = {k.replace('params/', ''): v for k, v in flat_dict.items() if k.startswith('params/')}\n    stats_flat = {k.replace('stats/', ''): v for k, v in flat_dict.items() if k.startswith('stats/')}\n    return {\n        'params': unflatten_dict({tuple(k.split('/')): v for k, v in params_flat.items()}),\n        'batch_stats': unflatten_dict({tuple(k.split('/')): v for k, v in stats_flat.items()})\n    }\n\ndef load_ensemble_weights(paths):\n    ensemble_vars = []\n    if ENSEMBLE_MODE == 'single': paths = [paths[0]]\n    for p in paths:\n        if Path(p).exists():\n            print(f\"📦 Loading weights: {p}\")\n            ensemble_vars.append(load_npz_weights(p))\n        else:\n            print(f\"⚠️ Warning: Path not found {p}\")\n    return ensemble_vars\n\n# ==============================================================================\n# 4. MODELS (PIX2PIX + POST-PROCESS)\n# ==============================================================================\nclass Pix2PixGenerator(nn.Module):\n    @nn.compact\n    def __call__(self, x, train: bool = False):\n        def conv_bn_leaky(feat, out_c):\n            feat = nn.Conv(out_c, (4, 4), strides=(2, 2), padding='SAME')(feat)\n            feat = nn.BatchNorm(use_running_average=not train)(feat)\n            return nn.leaky_relu(feat, 0.2)\n        \n        # FIXED: assignments split to avoid UnboundLocalError\n        s1 = nn.leaky_relu(nn.Conv(64, (4, 4), strides=(2, 2), padding='SAME')(x), 0.2)\n        s2 = conv_bn_leaky(s1, 128)\n        s3 = conv_bn_leaky(s2, 256)\n        \n        def up_bn_relu(feat, skip, out_c):\n            feat = nn.ConvTranspose(out_c, (4, 4), strides=(2, 2), padding='SAME')(feat)\n            feat = nn.BatchNorm(use_running_average=not train)(feat)\n            feat = nn.relu(feat)\n            return jnp.concatenate([feat, skip], axis=-1)\n            \n        u1 = up_bn_relu(s3, s2, 128)\n        u2 = up_bn_relu(u1, s1, 64)\n        return nn.ConvTranspose(1, (4, 4), strides=(2, 2), padding='SAME')(u2)\n\nclass PostProcessModel(nn.Module):\n    @nn.compact\n    def __call__(self, x, train: bool = True):\n        x = nn.Conv(64, (4, 4), strides=(2, 2), padding='SAME')(x)\n        s1 = nn.leaky_relu(x, 0.2)\n        s2 = nn.leaky_relu(nn.BatchNorm(use_running_average=not train)(nn.Conv(128, (4, 4), strides=(2, 2), padding='SAME')(s1)), 0.2)\n        s3 = nn.leaky_relu(nn.BatchNorm(use_running_average=not train)(nn.Conv(256, (4, 4), strides=(2, 2), padding='SAME')(s2)), 0.2)\n        u1 = jnp.concatenate([nn.relu(nn.BatchNorm(use_running_average=not train)(nn.ConvTranspose(128, (4, 4), strides=(2, 2), padding='SAME')(s3))), s2], axis=-1)\n        u2 = jnp.concatenate([nn.relu(nn.BatchNorm(use_running_average=not train)(nn.ConvTranspose(64, (4, 4), strides=(2, 2), padding='SAME')(u1))), s1], axis=-1)\n        return nn.ConvTranspose(1, (4, 4), strides=(2, 2), padding='SAME')(u2)\n\ndef load_postproc_weights(path):\n    with np.load(path, allow_pickle=True) as data:\n        params, stats = data[\"params\"].item(), data[\"stats\"].item()\n    return params, stats\n\nPOSTPROC_PATH =\"/kaggle/input/notebooks/crischir/jaxvesuviuspostprocessingtraining/checkpoints/postproc_best.npz\"\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 + TTA\n# ==============================================================================\ndef make_apply_fn(model):\n    @jax.jit\n    def _apply(vars, x):\n        return model.apply(vars, x, train=False)\n    return _apply\n\ndef get_ensemble_prob_single(model, ensemble_vars, patch_jax, mode=\"hybrid\"):\n    if patch_jax.ndim == 3: patch_jax = patch_jax[None, ...]\n    apply_fn = make_apply_fn(model)\n\n    def predict_internal(x):\n        all_logits = [apply_fn(v, x) for v in ensemble_vars]\n        stacked_logits = jnp.stack(all_logits, axis=0)\n        stacked_probs = nn.sigmoid(stacked_logits)\n        if mode == \"max\": return jnp.max(stacked_probs, axis=0)\n        elif mode == \"mean\": return nn.sigmoid(jnp.mean(stacked_logits, axis=0))\n        elif mode == \"agreement\":\n            votes, mean_p = jnp.sum(stacked_probs > 0.1, axis=0), jnp.mean(stacked_probs, axis=0)\n            return jnp.where(votes >= AGREEMENT_THRESHOLD, 1.0, mean_p)\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            return p_mean + ALPHA * (p_max - p_mean)\n        return stacked_probs[0]\n\n    if not TTA_INFERENCE:\n        return predict_internal(patch_jax)[0]\n    else:\n        # 4-way TTA\n        x_orig = patch_jax\n        p_orig = predict_internal(x_orig)\n        \n        x_h = jnp.flip(patch_jax, axis=2)\n        p_h = jnp.flip(predict_internal(x_h), axis=2)\n        \n        x_v = jnp.flip(patch_jax, axis=1)\n        p_v = jnp.flip(predict_internal(x_v), axis=1)\n        \n        x_hv = jnp.flip(patch_jax, axis=(1, 2))\n        p_hv = jnp.flip(predict_internal(x_hv), axis=(1, 2))\n\n        return (p_orig + p_h + p_v + p_hv)[0] / 4.0\n\ndef get_ensemble_prob_batch_single_device(model, ensemble_vars, patches_jax, mode=\"hybrid\"):\n    def _single(p): 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    return jax.pmap(local_batch_fn, axis_name=\"devices\")\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    pad = (N_DEV - (B % N_DEV)) % N_DEV\n    if pad > 0:\n        patches_jax = jnp.concatenate([patches_jax, jnp.zeros(((pad,) + patches_jax.shape[1:]), patches_jax.dtype)], axis=0)\n    B_padded = patches_jax.shape[0]\n    per_dev = B_padded // N_DEV\n    patches_sharded = patches_jax.reshape(N_DEV, per_dev, *patches_jax.shape[1:])\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    return probs_flat[:B] if pad > 0 else probs_flat\n\n# ==============================================================================\n# 6. INFERENCE PIPELINE\n# ==============================================================================\ndef run_inference_pipeline(model, ensemble_vars):\n    test_dir = DATA_PATH / 'test_images'\n    if not test_dir.exists(): return [], {}\n    \n    def local_batch_fn(patches_local):\n        return get_ensemble_prob_batch_single_device(model, ensemble_vars, patches_local, ENSEMBLE_MODE)\n    multi_gpu_fn = jax.pmap(local_batch_fn, axis_name=\"devices\")\n\n    submission_files, debug_volumes = [], {}\n    test_ids = [f.stem for f in test_dir.glob('*.tif')]\n\n    for vid in test_ids:\n        vol_start_time = time.time()\n        print(f\"\\n🏁 Volume: {vid} | TTA: {TTA_INFERENCE} | Mode: {ENSEMBLE_MODE}\")\n        with Image.open(str(test_dir / f\"{vid}.tif\")) as img:\n            vol_input = np.stack([np.array(f) for f in ImageSequence.Iterator(img)], axis=0)\n        z_max, h, w = vol_input.shape\n        vol_raw, vol_prob = np.zeros((z_max, h, w), dtype=np.float32), np.zeros((z_max, h, w), dtype=np.float32)\n        all_coords = [(y, x) for y in range(0, h, PATCH_SIZE) for x in range(0, w, PATCH_SIZE)]\n\n        def get_batch_worker(z_idx, coords_subset):\n            batch_patches = []\n            for y, x in coords_subset:\n                y1, x1 = min(y + PATCH_SIZE, h), min(x + PATCH_SIZE, w)\n                y0, x0 = y1 - PATCH_SIZE, x1 - PATCH_SIZE\n                patch = vol_input[z_idx - Z_CONTEXT: z_idx + Z_CONTEXT + 1, y0:y1, x0:x1].transpose(1, 2, 0)\n                batch_patches.append(normalize_patch(patch))\n            return np.stack(batch_patches), coords_subset\n\n        with ThreadPoolExecutor(max_workers=4) as executor:\n            for z in tqdm(range(Z_CONTEXT, z_max - Z_CONTEXT), desc=\"Slices\"):\n                coord_chunks = [all_coords[i:i + BATCH_SIZE] for i in range(0, len(all_coords), BATCH_SIZE)]\n                future_batch = executor.submit(get_batch_worker, z, coord_chunks[0])\n                for i in range(len(coord_chunks)):\n                    current_batch_np, batch_coords = future_batch.result()\n                    if i + 1 < len(coord_chunks): future_batch = executor.submit(get_batch_worker, z, coord_chunks[i + 1])\n                    b_size = current_batch_np.shape[0]\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_np = np.array(multi_gpu_fn(gpu_input)).reshape(b_size, PATCH_SIZE, PATCH_SIZE)\n                    else:\n                        prob_np = np.array(get_ensemble_prob_batch_single_device(model, ensemble_vars, current_batch_np, ENSEMBLE_MODE))[..., 0]\n                    for idx, (y, x) in enumerate(batch_coords):\n                        y1, x1 = min(y + PATCH_SIZE, h), min(x + PATCH_SIZE, w)\n                        y0, x0 = y1 - PATCH_SIZE, x1 - PATCH_SIZE\n                        raw_patch = prob_np[idx]\n                        vol_raw[z, y0:y1, x0:x1] = raw_patch\n                        if ENABLE_2D_POSTPROC:\n                            pp_out = apply_postproc(pp_vars, jnp.array(raw_patch[..., None].astype(np.float32)[None, ...]))\n                            vol_prob[z, y0:y1, x0:x1] = np.array(pp_out[0, ..., 0])\n                        else: vol_prob[z, y0:y1, x0:x1] = raw_patch\n\n        print(f\"🪄 GPU Post-processing...\")\n        vol_final = topo_postprocess_3d(vol_prob, TOPO_PARAMS) if DO_3D_POSTPROCESSING else (vol_prob > SimpleThreshold).astype(np.uint8) * 255\n        out_name = f\"{vid}.tif\"\n        import tifffile\n        tifffile.imwrite(out_name, vol_final, compression='deflate')\n        submission_files.append(out_name)\n        if not PRODUCTION_MODE: debug_volumes[vid] = {'input': vol_input, 'raw_prob': vol_raw, 'postproc_prob': vol_prob, 'final_mask': vol_final}\n        if PRODUCTION_MODE: \n            del vol_input, vol_raw, vol_prob, vol_final\n            import gc; gc.collect()\n        print(f\"✅ {vid} done in {time.time() - vol_start_time:.2f}s\")\n    return submission_files, debug_volumes\n\n# ==============================================================================\n# 7. VISUALIZATION\n# ==============================================================================\ndef visualize_results(debug_data):\n    for vid, data in debug_data.items():\n        vol_in, vol_raw, vol_final = data['input'], data['raw_prob'], data['final_mask']\n        z_max = vol_in.shape[0]\n        indices, labels = [Z_CONTEXT + 2, z_max // 2, z_max - Z_CONTEXT - 2], [\"Start\", \"Mid\", \"End\"]\n        fig, axes = plt.subplots(3, 4, figsize=(20, 15))\n        for i, z in enumerate(indices):\n            axes[i, 0].imshow(vol_in[z], cmap='gray'); axes[i, 0].set_title(f\"{labels[i]} (z={z}) Input\")\n            im2 = axes[i, 1].imshow(vol_raw[z], cmap='magma', vmin=0, vmax=1); plt.colorbar(im2, ax=axes[i, 1])\n            axes[i, 2].imshow(vol_final[z], cmap='gray'); axes[i, 2].set_title(\"3D Mask\")\n            overlay = np.zeros((*vol_in[z].shape, 3)); norm_in = vol_in[z] / 255.0\n            overlay[..., 0] = np.clip(norm_in + (vol_final[z] > 0) * 0.5, 0, 1)\n            overlay[..., 1], overlay[..., 2] = norm_in, norm_in\n            axes[i, 3].imshow(overlay); axes[i, 3].set_title(\"Overlay\")\n            for j in range(4): axes[i, j].axis('off')\n        plt.tight_layout()\n\ndef visualize_postproc_comparison(debug_data, vid=None):\n    if vid is None: vid = list(debug_data.keys())[0]\n    data = debug_data[vid]\n    vol_in, vol_raw, vol_pp, vol_final = data['input'], data['raw_prob'], data['postproc_prob'], data['final_mask']\n    z = vol_in.shape[0] // 2\n    fig, axes = plt.subplots(1, 4, figsize=(28, 6))\n    axes[0].imshow(vol_in[z], cmap='gray'); axes[0].axis('off')\n    im1 = axes[1].imshow(vol_raw[z], cmap='magma', vmin=0, vmax=1); plt.colorbar(im1, ax=axes[1])\n    im2 = axes[2].imshow(vol_pp[z], cmap='magma', vmin=0, vmax=1); plt.colorbar(im2, ax=axes[2])\n    axes[3].imshow(vol_final[z], cmap='gray'); axes[3].axis('off')\n    plt.tight_layout(); plt.show()\n\n# ==============================================================================\n# 8. MAIN\n# ==============================================================================\nif __name__ == \"__main__\":\n    netG = Pix2PixGenerator()\n    ensemble_vars = load_ensemble_weights(WEIGHT_PATHS)\n    if not ensemble_vars:\n        print(\"❌ Error: No weights loaded.\")\n    else:\n        print(\"🔥 JIT Warmup...\")\n        dummy = jnp.zeros((BATCH_SIZE, PATCH_SIZE, PATCH_SIZE, 2 * Z_CONTEXT + 1), jnp.float32)\n        _ = get_ensemble_prob_batch(netG, ensemble_vars, dummy, ENSEMBLE_MODE)\n        \n        output_tifs, debug_data = run_inference_pipeline(netG, ensemble_vars)\n        if not PRODUCTION_MODE and debug_data:\n            visualize_postproc_comparison(debug_data)\n            if ENABLE_VIZ: visualize_prediction_samples_jax(debug_data); visualize_results(debug_data)\n        if output_tifs:\n            with zipfile.ZipFile('submission.zip', 'w') as z:\n                for f in output_tifs: z.write(f)\n            for f in output_tifs: \n                if os.path.exists(f): os.remove(f)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T09:20:10.523219Z","iopub.execute_input":"2026-03-01T09:20:10.523586Z","iopub.status.idle":"2026-03-01T09:20:35.991892Z","shell.execute_reply.started":"2026-03-01T09:20:10.523552Z","shell.execute_reply":"2026-03-01T09:20:35.990861Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"NO TTA version","metadata":{}},{"cell_type":"code","source":"# import os\n# import sys\n# import zipfile\n# import time\n# from concurrent.futures import ThreadPoolExecutor\n# from pathlib import Path\n\n# import numpy as np\n# import jax\n# import jax.numpy as jnp\n# from jax import jit\n# from functools import partial\n\n# import flax.linen as nn\n# from flax.traverse_util import unflatten_dict\n# from flax.training import train_state\n\n# import matplotlib.pyplot as plt\n# from PIL import Image, ImageSequence\n# from scipy import ndimage as ndi\n# from skimage.morphology import remove_small_objects, skeletonize\n# # ==============================================================================\n# # 0. Debug mode \n# # =============================================================================\n\n# PRODUCTION_MODE = False  # Set to False when debugging/visualizing locally\n# ENABLE_VIZ = not PRODUCTION_MODE\n# # ==============================================================================\n# # 0. QUICK RUNTIME INFO\n# # ==============================================================================\n# print(\"Python:\", sys.version)\n# print(\"JAX devices:\", jax.devices())\n\n# # ==============================================================================\n# # 1. VISUALIZATION HELPER (5-VIEW)\n# # ==============================================================================\n\n# def 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\n# ENABLE_VIZ = True\n# ENABLE_2D_POSTPROC=True\n# ENSEMBLE_MODE = \"hybrid\"   # 'mean', 'max', 'agreement', 'hybrid', 'single'\n# AGREEMENT_THRESHOLD = 7\n# ALPHA = 0.5\n\n# MODEL_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#     ],\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# }\n# MODEL_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_50.npz\",\n#         \"/kaggle/input/notebooks/crischir/tpu-training-jax-pix2pix/checkpoints/pix2pix_ep_60.npz\",\n#         \"/kaggle/input/notebooks/crischir/tpu-training-jax-pix2pix/checkpoints/pix2pix_ep_70.npz\",\n#     ],\n#     \"prewitt-loss-model-jax-pix2pix\": [\n#         \"/kaggle/input/notebooks/crischir/prewitt-loss-model-jax-pix2pix/checkpoints/pix2pix_best.npz\",\n#         \"/kaggle/input/notebooks/crischir/prewitt-loss-model-jax-pix2pix/checkpoints/pix2pix_ep_40.npz\",\n#         \"/kaggle/input/notebooks/crischir/prewitt-loss-model-jax-pix2pix/checkpoints/pix2pix_ep_50.npz\",\n#         \"/kaggle/input/notebooks/crischir/prewitt-loss-model-jax-pix2pix/checkpoints/pix2pix_ep_70.npz\",\n#     ],\n#     \"jax-pix2pix\": [\n#         \"/kaggle/input/notebooks/crischir/jax-pix2pix/checkpoints/pix2pix_ep_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\n# SELECTED = [\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#     \"vesuvius-checkpoint-dataset\",\n#     \"prewitt-loss-model-jax-pix2pix\",\n# ]\n\n# WEIGHT_PATHS = []\n# for key in SELECTED:\n#     WEIGHT_PATHS.extend(MODEL_CHECKPOINTS.get(key, []))\n\n# DO_3D_POSTPROCESSING = True\n# TOPO_PARAMS = {\n#     \"T_low\": 0.1,\n#     \"T_high\": 0.65,\n#     \"z_radius\": 3,\n#     \"xy_radius\": 2,\n#     \"dust_min_size\": 5,\n# }\n\n# DATA_PATH = Path(\"/kaggle/input/competitions/vesuvius-challenge-surface-detection/\")\n# PATCH_SIZE = 256\n# Z_CONTEXT = 3\n# SimpleThreshold = 0.15\n\n# BATCH_SIZE = 256\n\n# devices = jax.devices()\n# N_DEV = len(devices)\n# print(\"Num devices:\", N_DEV)\n\n# # ==============================================================================\n# # 3. HELPERS\n# # ==============================================================================\n\n# def to_skeleton(vol_mask):\n#     skel = skeletonize(vol_mask > 0)\n#     return skel.astype(np.uint8) * 255\n\n# def 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\n# def 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))\n# def 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\n# def 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\n# def normalize_patch(patch: np.ndarray) -> np.ndarray:\n#     patch = patch.astype(np.float32)\n#     return (patch - patch.mean()) / (patch.std() + 1e-6)\n\n# def 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\n# def 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\n# class 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\n# class 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\n# class PostProcState(train_state.TrainState):\n#     batch_stats: dict\n\n# def 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\n# POSTPROC_PATH =\"/kaggle/input/datasets/crischir/postprocessing2vesuviuscheckpoints/checkpoints/postproc_e20.npz\" #nice results\n# # POSTPROC_PATH =\"/kaggle/input/notebooks/crischir/jaxvesuviuspostprocessingtrainingfull/checkpoints/postproc_epoch_15.npz\"\n\n# postproc_model = PostProcessModel()\n# pp_params, pp_stats = load_postproc_weights(POSTPROC_PATH)\n# pp_vars = {\"params\": pp_params, \"batch_stats\": pp_stats}\n\n# @jax.jit\n# def 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\n# def make_apply_fn(model):\n#     def _apply(vars, x):\n#         return model.apply(vars, x, train=False)\n#     return jax.jit(_apply)\n\n# def 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\n# def 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\n# def 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\n# def 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\n# def 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\n# def 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()\n# def 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\n# if __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-03-01T09:17:37.342609Z","iopub.execute_input":"2026-03-01T09:17:37.343265Z","iopub.status.idle":"2026-03-01T09:17:37.489675Z","shell.execute_reply.started":"2026-03-01T09:17:37.343221Z","shell.execute_reply":"2026-03-01T09:17:37.488828Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Gemini","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}}]}