{"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":"gpu","dataSources":[{"sourceType":"competition","sourceId":117682,"databundleVersionId":15062069},{"sourceType":"datasetVersion","sourceId":14245247,"datasetId":9088503,"databundleVersionId":15043865},{"sourceType":"datasetVersion","sourceId":14871252,"datasetId":9513632,"databundleVersionId":15733419},{"sourceType":"datasetVersion","sourceId":14295835,"datasetId":9125518,"databundleVersionId":15099120},{"sourceType":"modelInstanceVersion","sourceId":768516,"databundleVersionId":15864953,"modelInstanceId":587120,"modelId":599446},{"sourceType":"kernelVersion","sourceId":288572598},{"sourceType":"kernelVersion","sourceId":294535277},{"sourceType":"kernelVersion","sourceId":297732241},{"sourceType":"kernelVersion","sourceId":297827765},{"sourceType":"kernelVersion","sourceId":298023302},{"sourceType":"kernelVersion","sourceId":298441506},{"sourceType":"kernelVersion","sourceId":299093011},{"sourceType":"kernelVersion","sourceId":299102216},{"sourceType":"kernelVersion","sourceId":300093425},{"sourceType":"kernelVersion","sourceId":300641787}],"dockerImageVersionId":31260,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"###### Copied from 02.02.2026-vesuvius-v3\n#### Models\n#### 3DModels\n* https://www.kaggle.com/datasets/hideyukizushi/colab-a-162v4-gpu-transunet-seresnext101-x160\n* https://www.kaggle.com/code/choudharymanas/train-transunet-baseline-lb-0-537\n* https://www.kaggle.com/code/ipythonx/train-vesuvius-surface-3d-detection-on-tpu\n#### 2.5 DModels \n* https://www.kaggle.com/code/crischir/test-code-sobel-jax-pix2pix\n* https://www.kaggle.com/code/crischir/2-5d-filter-training-pix2pix\n* https://www.kaggle.com/code/crischir/2-5d-filter-training-pix2pix-pred-mask\n#### Un-learning Model :)\n* https://www.kaggle.com/datasets/crischir/postprocessvesuviuscomplete\n\n#### Credits\n* Ibrat Usmonov ·02.02.2026-vesuvius-v3\n* yukiZ :\n    * COLAB-A-162v4-GPU-TransUNet-seresnext101-x160 \n    * Vesuvius25-packages-offline-installer-v20251226\n* Manas Choudhary [Train] TransUNet Baseline [LB 0.537]\n* Innat ·[Train] Vesuvius Surface 3D Detection on TPU\n\n## Gemini\n","metadata":{}},{"cell_type":"markdown","source":"Hybrid postprocessing = 2D learned refinement + 3D structural refinement.\n\n1. 2.5D model produces fine‑detail probabilities.\n2. (Optional) 2D PostProcessModel sharpens and denoises these. Use with mean.\n3. 3D model produces deep structural probabilities.\n4. Ensemble blends 3D structure with 2.5D detail.\n5. (Optional) 2D PostProcessModel refines the ensemble.\n6. 3D topology enforces continuity and removes noise.\n\n2D PP improves local detail.\n3D topology enforces global structure.\nTogether they produce the cleanest, most stable masks.","metadata":{}},{"cell_type":"markdown","source":"## 🔍 Visual Comparison Grid\n\nThis grid helps evaluate the effect of each processing stage.\n\n| Stage | Description | What to Look For |\n|-------|-------------|------------------|\n| **Raw CT** | Original slice | Papyrus texture, folds, cracks |\n| **3D Prob** | 3D TransUNet output | Smooth structure, continuity |\n| **2.5D Before 2D PP** | Raw Pix2Pix ensemble | Fine detail, noise |\n| **2.5D After 2D PP** | Refined 2.5D | Sharper edges, less noise |\n| **Ensemble Before 2D PP** | Weighted blend | Balance of detail + structure |\n| **Ensemble After 2D PP** | Refined blend | Sharpened ensemble (if enabled) |\n| **Final Mask** | After 3D topology | Clean, continuous ink surface |\n| **Confidence Difference** | |3D - 2.5D| | Where models disagree |\n| **2D PP Difference** | |After - Before| | Effect of 2D PP |\n| **3D MIP Overview** | Max projections | Global structure check |\n| **GIF** | Slice animation | Temporal consistency |","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"### HYBRID PIPELINE CONFIGURATION & VISUALIZATION GUIDE\n\nThis section explains the full |hybrid 3D + 2.5D inference pipeline, including configuration options, their impact, and how to interpret the visual diagnostics.\n\n#### HYBRID PIPELINE FLOW DIAGRAM\n\n* Raw CT Volume → 3D TransUNet → 3D Probabilities → Ensemble (W3D * P3D + W25D * P25D)→ (Optional) 2D PostProcessModel (after)→ 3D Topology Postprocessing → Final Mask\n\n* Raw CT Volume→ 2.5D Pix2Pix Ensemble→ (Optional) 2D PostProcessModel (before)→ 2.5D Refined Probabilities→ Ensemble (W3D * P3D + W25D * P25D)\n\n#### CONFIGURATION OPTIONS & THEIR IMPACT\n\n##### POSTPROC_2D_LOCATION Controls where the 2D PostProcessModel is applied.Options: none, before, after, both\n\n1. \"none\": 2D postprocessing disabled Fastest, most stable No denoising or sharpening\n\n2. \"before\": (recommended default) 2D PostProcessing applied only to 2.5D branch Sharpens 2.5D details Preserves 3D structure No probability collapse Most stable option\n\n3. \"after\": 2D PP applied only to ensemble Can sharpen final output High risk of probability collapse Often produces black masks Experimental only\n\n4. \"both\": 2D PP applied to 2.5D and ensemble Maximum sharpening Very heavy GPU load Highest collapse risk Not recommended for production\n\n#####  POSTPROC_VERSION Controls which 2D postprocessing model is used.Options: none, v1, v2, hybrid\n\n1. \"none\": disables 2D PP\n2. \"v1\": stable, gentle sharpening\"\n3. v2\": stronger sharpening, more aggressive\n4. \"hybrid\": average of v1 and v2, balanced but heavier\n\n#####  POSTPROC_MODE Controls 3D topology postprocessing.Options: topo, none\n\n1. \"topo\": (recommended) 3D hysteresis Anisotropic closing Dust removal Clean, continuous surfaces Can remove predictions if thresholds too high\n\n2. \"none\":Raw ensemble Useful for debugging Very noisy\n\n######  Thresholds (T_LOW, T_HIGH) just in theory bacause of normalization our pred can be relevant at 0.1\n\n* High thresholds (0.45 / 0.85): Clean mask Risk missing faint ink\n\n* Medium thresholds (0.25 / 0.55): Balanced Low thresholds (0.05 / 0.15): Recovers faint ink More noise\n\nNeeded if 2D PP collapses probabilities\n\n###### Ensemble Weights (W_3D, W_25D)\n\n* Higher W_3D: More structural consistency Fewer false positives Less fine detail\n\n* Higher W_25D: More ink detail More noise and artifacts Recommended:W_3D = 0.3W_25D = 0.7\n\n#### Debug Options\n\nDEBUG_ENABLE = TrueDEBUG_MAX_VOLUMES = 1\n\nSaves:\n\n\n\n> Raw CT 3D probabilities Final mask All plots GIF 3D MIP overview\n\n\n\n\n######  POSTPROC_2D_LOCATION affects: detail vs stability \n###### POSTPROC_VERSION affects: sharpening strength \n###### POSTPROC_MODE affects: mask cleanliness \n###### Thresholds affect: sensitivity vs noise\n###### Ensemble weights affect: structure vs detail\n\n##### RECOMMENDED STABLE CONFIGURATION\n\n\n* POSTPROC_2D_LOCATION = before\n* POSTPROC_VERSION = v2\n* POSTPROC_MODE = topo\n* W_3D = 0.3\n* W_25D = 0.7","metadata":{}},{"cell_type":"code","source":"\"\"\"\n================================================================================\n   VESUVIUS HYBRID ENSEMBLE: 3D TransUNet + 2.5D Pix2Pix (JAX/Flax)\n   - 3D prediction flow preserved\n   - 2.5D ensemble + PostProcessModel v1/v2/hybrid\n   - 2D postproc location: before ensemble / after ensemble / both / none\n   - Hybrid postproc: (optional) 2D ML postproc on ensemble → 3D topology\n   - Debug stored only for FIRST test volume\n   - Confidence difference map, GIF, per-slice viewer, 3D overview\n================================================================================\n\"\"\"\n\nimport os\n\n# ============================================================================\n# 1. CRITICAL: MEMORY MANAGEMENT (TF + JAX COEXISTENCE)\n# ============================================================================\nos.environ[\"KERAS_BACKEND\"] = \"tensorflow\"\nos.environ[\"TF_CPP_MIN_LOG_LEVEL\"] = \"2\"\nos.environ[\"TF_FORCE_GPU_ALLOW_GROWTH\"] = \"true\"\nos.environ[\"XLA_PYTHON_CLIENT_PREALLOCATE\"] = \"false\"\nos.environ[\"XLA_PYTHON_CLIENT_ALLOCATOR\"] = \"platform\"\n\n# ============================================================================\n# 2. OPTIONAL OFFLINE INSTALL\n# ============================================================================\nvar = \"/kaggle/input/vesuvius25-packages-offline-installer-v20251226/whls\"\nif os.path.exists(var):\n    print(f\"Installing packages from: {var}\")\n    import subprocess\n    subprocess.run([\n        \"pip\", \"install\", \"--quiet\",\n        f\"{var}/keras_nightly-3.12.0.dev2025100703-py3-none-any.whl\",\n        f\"{var}/tifffile-2025.10.16-py3-none-any.whl\",\n        f\"{var}/imagecodecs-2025.11.11-cp311-abi3-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl\",\n        f\"{var}/medicai-0.0.3-py3-none-any.whl\",\n        \"--no-index\",\n        \"--find-links\", var\n    ], check=False, capture_output=True)\n\n# ============================================================================\n# 3. IMPORTS\n# ============================================================================\nimport numpy as np\nimport pandas as pd\nimport zipfile\nimport tifffile\nimport gc\nimport scipy.ndimage as ndi\nfrom skimage.morphology import remove_small_objects\nfrom tqdm.auto import tqdm\nfrom pathlib import Path\nfrom PIL import Image, ImageSequence\nimport matplotlib.pyplot as plt\nimport random\n\n# TensorFlow / Keras (3D)\nimport keras\nfrom keras import ops\nimport tensorflow as tf\nfrom medicai.transforms import Compose, ScaleIntensityRange\nfrom medicai.models import TransUNet\nfrom medicai.utils.inference import SlidingWindowInference\n\n# JAX / Flax (2.5D)\nimport jax\nimport jax.numpy as jnp\nimport flax.linen as nn\nfrom flax.core.frozen_dict import freeze\nfrom flax.traverse_util import unflatten_dict\n\nprint(\"=\" * 60)\nprint(\"VESUVIUS HYBRID PRODUCTION: 3D + 2.5D\")\nprint(f\"TF Devices: {tf.config.list_physical_devices('GPU')}\")\nprint(f\"JAX Devices: {jax.devices()}\")\nprint(\"=\" * 60)\n\n# ============================================================================\n# 0.0 2.5D models\n# ============================================================================\n\nMODEL_CHECKPOINTS = {\n        \"2-5d-filter-training-pix2pix\": [\n            \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix/checkpoints/pix2pix_best.npz\",\n            \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix/checkpoints/pix2pix_ep_10.npz\",\n            \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix/checkpoints/pix2pix_ep_20.npz\",\n            \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix/checkpoints/pix2pix_ep_30.npz\",\n        ],\n        \"2-5d-filter-training-pix2pix-pred-mask\": [\n            \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix-pred-mask/checkpoints/pix2pix_best.npz\",\n            \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix-pred-mask/checkpoints/pix2pix_ep_10.npz\",\n            \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix-pred-mask/checkpoints/pix2pix_ep_20.npz\",\n            \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix-pred-mask/checkpoints/pix2pix_ep_70.npz\",\n            \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix-pred-mask/checkpoints/pix2pix_ep_90.npz\",\n            \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix-pred-mask/checkpoints/pix2pix_ep_120.npz\",\n            \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix-pred-mask/checkpoints/pix2pix_ep_170.npz\",\n            \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix-pred-mask/checkpoints/pix2pix_ep_180.npz\",\n            \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix-pred-mask/checkpoints/pix2pix_ep_190.npz\",\n            \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix-pred-mask/checkpoints/pix2pix_ep_200.npz\",\n        ],\n        \"test-code-for-jax-pix2pix\": [\n            \"/kaggle/input/notebooks/crischir/test-code-for-jax-pix2pix/checkpoints/pix2pix_best.npz\",\n            \"/kaggle/input/notebooks/crischir/test-code-for-jax-pix2pix/checkpoints/pix2pix_ep_30.npz\",\n            \"/kaggle/input/notebooks/crischir/test-code-for-jax-pix2pix/checkpoints/pix2pix_ep_40.npz\",\n            \"/kaggle/input/notebooks/crischir/test-code-for-jax-pix2pix/checkpoints/pix2pix_ep_50.npz\",\n        ],\n        \"test-code-sobel-jax-pix2pix\": [\n            \"/kaggle/input/notebooks/crischir/test-code-sobel-jax-pix2pix/checkpoints/pix2pix_ep_100.npz\",\n            \"/kaggle/input/notebooks/crischir/test-code-sobel-jax-pix2pix/checkpoints/pix2pix_ep_130.npz\",\n        ],\n        \"tpu-training-jax-pix2pix\": [\n            \"/kaggle/input/notebooks/crischir/tpu-training-jax-pix2pix/checkpoints/pix2pix_best.npz\",\n            \"/kaggle/input/notebooks/crischir/tpu-training-jax-pix2pix/checkpoints/pix2pix_ep_30.npz\",\n            \"/kaggle/input/notebooks/crischir/tpu-training-jax-pix2pix/checkpoints/pix2pix_ep_60.npz\",\n            \"/kaggle/input/notebooks/crischir/tpu-training-jax-pix2pix/checkpoints/pix2pix_ep_70.npz\",\n        ],\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         \"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        \"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    \n# Choose your experiments here:\nSELECTED = [\n    \"tpu-training-jax-pix2pix\",\n    # \"retromasking-vesuvius\",\n    # \"2-5d-filter-training-pix2pix-pred-mask\",\n    # \"test-code-for-jax-pix2pix\",\n    \"clDice-1\",\n    \"clDice-2\",   \n    \n]\n\nWEIGHT_PATHS = []\nfor key in SELECTED:\n    WEIGHT_PATHS.extend(MODEL_CHECKPOINTS.get(key, []))\n# ============================================================================\n# 4.0 Gate Gamma\n# ============================================================================\nGATE_GAMMA = 4  # > 1 favors 3D/Masking; < 1 favors 2.5D details\n# ============================================================================\n# 4. CONFIGURATION\n# ============================================================================\n\n\n\nROOT_DIR = \"/kaggle/input/vesuvius-challenge-surface-detection\"\nTEST_DIR = f\"{ROOT_DIR}/test_images\"\nOUTPUT_DIR = \"/kaggle/working/submission_masks\"\nZIP_PATH = \"/kaggle/working/submission.zip\"\nos.makedirs(OUTPUT_DIR, exist_ok=True)\n\n# Ensemble weights\nW_3D = 0.4\nW_25D = 0.6\n\n# 3D CONFIG\nNUM_CLASSES_3D = 3\nPATCH_SIZE_3D = (160, 160, 160)\nOVERLAP_3D = 0.25\n\n# 2.5D CONFIG\nPATCH_SIZE_25D = 320\nZ_CONTEXT_25D = 3\nBATCH_SIZE_25D = 128\nNORMALIZE_OUTPUT_25D = False\n# ENSEMBLE_25D_MODE = \"balanced_mean\"   # \"mean\", \"max\", \"agreement\", \"hybrid\", \"first\",\"smoothed_mean\",\"balanced_mean\" # do not use 2D post preocesssing for hybrid mode \nENSEMBLE_25D_MODE = \"hybrid\" \nAGREEMENT_THRESHOLD = 5   # tune as needed\nALPHA = 0.7               # hybrid mixing strength\nBETA = 0.8 #balanced smooth mode \n\n\n# Post-processing (3D topology)\nT_LOW = 0.37\nT_HIGH = 0.75\nZ_RADIUS = 1\nXY_RADIUS = 0\nDUST_MIN_SIZE = 8\n\n# Paths\nPATH_3D_1 = \"/kaggle/input/train-transunet-baseline-lb-0-537/fine_tuning_epoch_20.weights.h5\"\nPATH_3D_2 = \"/kaggle/input/train-vesuvius-surface-3d-detection-on-tpu/model.weights.h5\"\n\n# PATHS_25D = [\n#    # \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix/checkpoints/pix2pix_best.npz\",\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_120.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_180.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#    # \"/kaggle/input/notebooks/crischir/test-code-for-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_50.npz\",\n#    \"/kaggle/input/notebooks/crischir/tpu-training-jax-pix2pix/checkpoints/pix2pix_ep_60.npz\",\n# ]\nPATHS_25D = WEIGHT_PATHS \n\n# Post-processing 2D versions\n# POSTPROC_VERSION = \"v2\"  # \"none\", \"v1\", \"v2\", \"hybrid\"\nPOSTPROC_VERSION = \"none\"  # \"none\", \"v1\", \"v2\", \"hybrid\"\n\nPATH_POSTPROC_V1 = \"/kaggle/input/notebooks/crischir/postprocessing-2tta-vesuvius-training/post_proc_checkpoints/post_process_ep30.npz\"\nPATH_POSTPROC_V2 = \"/kaggle/input/notebooks/crischir/jaxvesuviuspostprocessingtrainingfull/checkpoints/postproc_epoch_15.npz\"\n\n# Where to apply 2D postprocessing:\n# \"none\"  -> no 2D postproc\n# \"before\" -> only on 2.5D branch (current default behavior)\n# \"after\"  -> only on ensemble probabilities\n# \"both\"   -> on 2.5D branch AND on ensemble probabilities\n# POSTPROC_2D_LOCATION = \"before\"  # \"none\", \"before\", \"after\", \"both\"\nPOSTPROC_2D_LOCATION = \"none\"  # \"none\", \"before\", \"after\", \"both\n# POSTPROC_2D_LOCATION = \"both\"  # \"none\", \"before\", \"after\", \"both\"\n# Debug / visualization config\nDEBUG_ENABLE = True\nDEBUG_MAX_VOLUMES = 1   # store debug only for first volume\nPOSTPROC_MODE = \"topo\"  # \"topo\" | \"none\" (3D topology on final probs)\n\n# ============================================================================\n# 5. 2.5D MODEL DEFINITIONS (FLAX)\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\n\nclass PostProcessModel2D(nn.Module):\n    @nn.compact\n    def __call__(self, x, train: bool = False):\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        u1 = 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        u1 = jnp.concatenate([u1, s2], axis=-1)\n        u2 = 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        u2 = jnp.concatenate([u2, s1], axis=-1)\n        return nn.ConvTranspose(1, (4, 4), strides=(2, 2), padding=\"SAME\")(u2)\n\n\ndef load_flax_weights_v1(path):\n    if not os.path.exists(path):\n        return None\n    with np.load(path, allow_pickle=False) as data:\n        flat = {k: v for k, v in data.items()}\n    p = {k.replace(\"params/\", \"\"): v for k, v in flat.items() if k.startswith(\"params/\")}\n    s = {k.replace(\"stats/\", \"\"): v for k, v in flat.items() if k.startswith(\"stats/\")}\n    return freeze({\n        \"params\": unflatten_dict({tuple(k.split(\"/\")): v for k, v in p.items()}),\n        \"batch_stats\": unflatten_dict({tuple(k.split(\"/\")): v for k, v in s.items()}),\n    })\n\n\ndef load_flax_weights_v2(path):\n    if not os.path.exists(path):\n        return None\n    with np.load(path, allow_pickle=True) as data:\n        keys = list(data.keys())\n\n        if \"params\" in keys and \"batch_stats\" in keys:\n            return freeze({\n                \"params\": data[\"params\"].item(),\n                \"batch_stats\": data[\"batch_stats\"].item(),\n            })\n\n        if \"params\" in keys:\n            return freeze({\n                \"params\": data[\"params\"].item(),\n                \"batch_stats\": {},\n            })\n\n        flat = {k: v for k, v in data.items()}\n        return freeze({\n            \"params\": flat,\n            \"batch_stats\": {},\n        })\n\n\ndef load_postproc_weights(path, version):\n    if version == \"v1\":\n        return load_flax_weights_v1(path)\n    if version == \"v2\":\n        return load_flax_weights_v2(path)\n    return None\n\n\n# def make_dual_ensemble_fn(ensemble_vars):\n#     model = Pix2PixGenerator()\n\n#     def _forward(x):\n#         all_logits = jnp.stack([model.apply(v, x, train=False) for v in ensemble_vars], axis=0)\n#         logit_avg_probs = nn.sigmoid(jnp.mean(all_logits, axis=0))\n#         prob_avg_probs = jnp.mean(nn.sigmoid(all_logits), axis=0)\n#         return logit_avg_probs * prob_avg_probs\n\n#     return jax.jit(_forward)\ndef make_dual_ensemble_fn(ensemble_vars, mode=\"hybrid\"):\n    model = Pix2PixGenerator()\n\n    def _forward(batch_np):\n        # batch_np: (B, H, W, C)\n        batch_jax = jnp.array(batch_np)\n\n        def process_single(patch):\n            return get_ensemble_prob_single(model, ensemble_vars, patch, mode=mode)\n\n        # vmap over batch dimension\n        return jax.vmap(process_single)(batch_jax)\n\n    return jax.jit(_forward)\n\n\ndef make_postproc_2d_fn(post_vars):\n    \"\"\"\n    Build a JIT-ed postprocessing function robust to:\n    - weights with params + batch_stats (v1)\n    - weights with params only (v2)\n    \"\"\"\n    model = PostProcessModel2D()\n\n    if \"batch_stats\" not in post_vars or not post_vars[\"batch_stats\"]:\n        dummy = jnp.zeros((1, 256, 256, 1), jnp.float32)\n        init_vars = model.init(jax.random.PRNGKey(0), dummy, train=False)\n        post_vars = freeze({\n            \"params\": post_vars[\"params\"],\n            \"batch_stats\": init_vars[\"batch_stats\"],\n        })\n\n    def _forward(x):\n        return nn.sigmoid(model.apply(post_vars, x, train=False))\n\n    return jax.jit(_forward)\n\n\ndef get_postproc_fn(version):\n    if version == \"none\":\n        return None\n\n    if version == \"v1\":\n        vars1 = load_postproc_weights(PATH_POSTPROC_V1, \"v1\")\n        if vars1 is None:\n            print(\"WARNING: Postproc v1 weights not found.\")\n            return None\n        return make_postproc_2d_fn(vars1)\n\n    if version == \"v2\":\n        vars2 = load_postproc_weights(PATH_POSTPROC_V2, \"v2\")\n        if vars2 is None:\n            print(\"WARNING: Postproc v2 weights not found.\")\n            return None\n        return make_postproc_2d_fn(vars2)\n\n    if version == \"hybrid\":\n        vars1 = load_postproc_weights(PATH_POSTPROC_V1, \"v1\")\n        vars2 = load_postproc_weights(PATH_POSTPROC_V2, \"v2\")\n        if vars1 is None or vars2 is None:\n            print(\"WARNING: One of hybrid postproc weights missing.\")\n            return None\n        fn1 = make_postproc_2d_fn(vars1)\n        fn2 = make_postproc_2d_fn(vars2)\n\n        def hybrid_fn(x):\n            return 0.5 * fn1(x) + 0.5 * fn2(x)\n\n        return jax.jit(hybrid_fn)\n\n    return None\n# def get_ensemble_prob_single(model, ensemble_vars, patch_jax, mode=\"hybrid\"):\n#     \"\"\"\n#     Compute ensemble probability for a single patch using different strategies.\n#     patch_jax: (H, W, C) or (1, H, W, C)\n#     Returns: (H, W, 1)\n#     \"\"\"\n\n#     # Ensure batch dimension\n#     if patch_jax.ndim == 3:\n#         patch_jax = patch_jax[None, ...]\n\n#     # Apply function for each model\n#     apply_fn = make_apply_fn(model)\n#     logits_list = [apply_fn(v, patch_jax) for v in ensemble_vars]\n\n#     # Stack logits: (N_models, 1, H, W, 1)\n#     stacked_logits = jnp.stack(logits_list, axis=0)\n#     stacked_probs = nn.sigmoid(stacked_logits)\n\n#     # --- Ensemble modes ---\n#     if mode == \"max\":\n#         final_prob = jnp.max(stacked_probs, axis=0)\n\n#     elif mode == \"mean\":\n#         avg_logits = jnp.mean(stacked_logits, axis=0)\n#         final_prob = nn.sigmoid(avg_logits)\n\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\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\n#     else:  # fallback: first model\n#         final_prob = stacked_probs[0]\n\n#     return final_prob[0]  # remove batch dimension\ndef get_ensemble_prob_single(model, ensemble_vars, patch_jax, mode=\"hybrid\"):\n    \"\"\"\n    Compute ensemble probability for a single patch using different strategies.\n    patch_jax: (H, W, C) or (1, H, W, C)\n    Returns: (H, W, 1)\n    \"\"\"\n\n    # Ensure batch dimension\n    if patch_jax.ndim == 3:\n        patch_jax = patch_jax[None, ...]\n\n    # Apply model to each set of weights\n    logits_list = [\n        model.apply(v, patch_jax, train=False)\n        for v in ensemble_vars\n    ]\n\n    stacked_logits = jnp.stack(logits_list, axis=0)\n    stacked_probs = nn.sigmoid(stacked_logits)\n\n    # Ensemble modes\n    if mode == \"max\":\n        final_prob = jnp.max(stacked_probs, axis=0)\n\n    elif mode == \"mean\":\n        avg_logits = jnp.mean(stacked_logits, axis=0)\n        final_prob = nn.sigmoid(avg_logits)\n\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\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    elif mode == \"smoothed_mean\":\n        # Old ensemble behavior\n        # 1) sigmoid(mean(logits))\n        logit_avg_probs = nn.sigmoid(jnp.mean(stacked_logits, axis=0))\n        # 2) mean(sigmoid(logits))\n        prob_avg_probs = jnp.mean(stacked_probs, axis=0)\n        # 3) multiply them\n        final_prob = logit_avg_probs * prob_avg_probs\n    elif mode == \"balanced_mean\":\n        p_mean = nn.sigmoid(jnp.mean(stacked_logits, axis=0))\n        p_smooth = p_mean * jnp.mean(stacked_probs, axis=0)\n        final_prob = (1 - BETA) * p_mean + BETA * p_smooth\n\n    else:  # fallback\n        final_prob = stacked_probs[0]\n\n    return final_prob[0]  # remove batch dimension\n\n# ============================================================================\n# 6. MODEL LOADERS\n# ============================================================================\ndef get_3d_models():\n    models = []\n    for path in [PATH_3D_1, PATH_3D_2]:\n        if os.path.exists(path):\n            print(f\"Loading 3D Model: {os.path.basename(path)}\")\n            m = TransUNet(\n                input_shape=(160, 160, 160, 1),\n                encoder_name=\"seresnext50\",\n                classifier_activation=None,\n                num_classes=NUM_CLASSES_3D,\n            )\n            m.load_weights(path)\n            models.append(m)\n    return models\n\n\ndef get_25d_funcs():\n    print(\"Loading 2.5D Flax Weights...\")\n    weights = []\n    for p in PATHS_25D:\n        if os.path.exists(p):\n            w = load_flax_weights_v1(p)\n            if w is not None:\n                weights.append(w)\n\n    if not weights:\n        print(\"WARNING: No 2.5D weights found.\")\n        return None, None\n\n    # ensemble_fn = make_dual_ensemble_fn(weights)\n    ensemble_fn = make_dual_ensemble_fn(weights, mode=ENSEMBLE_25D_MODE)\n\n    postproc_fn = get_postproc_fn(POSTPROC_VERSION)\n    if postproc_fn is not None and POSTPROC_2D_LOCATION in [\"before\", \"after\", \"both\"]:\n        print(f\"Using 2D Post-Processor: {POSTPROC_VERSION} at {POSTPROC_2D_LOCATION}\")\n    else:\n        print(\"No 2D Post-Processor active.\")\n    return ensemble_fn, postproc_fn\n\n# ============================================================================\n# 7. INFERENCE FUNCTIONS\n# ============================================================================\ndef predict_3d_volume(vol_input, models):\n    if not models:\n        return None\n\n    print(f\"  -> 3D Inference on {vol_input.shape}...\")\n    data = {\"image\": vol_input}\n    pipeline = Compose([\n        ScaleIntensityRange(keys=[\"image\"], a_min=0, a_max=255, b_min=0, b_max=1, clip=True)\n    ])\n    vol_norm = pipeline(data)[\"image\"]\n\n    ensemble_logits = []\n    for i, model in enumerate(models):\n        swi = SlidingWindowInference(\n            model,\n            num_classes=NUM_CLASSES_3D,\n            roi_size=PATCH_SIZE_3D,\n            sw_batch_size=1,\n            mode=\"gaussian\",\n            overlap=OVERLAP_3D,\n        )\n        logits = []\n        logits.append(swi(vol_norm))\n        for k in [1, 2, 3]:\n            img_r = np.rot90(vol_norm, k=k, axes=(2, 3))\n            p = swi(img_r)\n            p = np.rot90(p, k=-k, axes=(2, 3))\n            logits.append(p)\n        ensemble_logits.append(np.mean(logits, axis=0))\n\n    total_logits = np.mean(ensemble_logits, axis=0)\n    probs = ops.softmax(total_logits, axis=-1)\n    return np.squeeze(probs[..., 1])\n\n\ndef predict_25d_volume(vol_raw, ensemble_fn, postproc_fn):\n    if ensemble_fn is None:\n        return None, None\n\n    vol = np.squeeze(vol_raw)  # (D, H, W)\n    z_max, h, w = vol.shape\n    vol_probs_before_pp = np.zeros((z_max, h, w), dtype=np.float32)\n    vol_probs_after_pp = np.zeros((z_max, h, w), dtype=np.float32)\n\n    print(f\"  -> 2.5D Inference (Context={Z_CONTEXT_25D}, Batch={BATCH_SIZE_25D})...\")\n\n    def normalize_patch(patch):\n        return (patch.astype(np.float32) - patch.mean()) / (patch.std() + 1e-6)\n\n    coords = [(y, x) for y in range(0, h, PATCH_SIZE_25D) for x in range(0, w, PATCH_SIZE_25D)]\n    z_range = range(Z_CONTEXT_25D, z_max - Z_CONTEXT_25D)\n\n    apply_before = POSTPROC_2D_LOCATION in [\"before\", \"both\"]\n\n    for z in tqdm(z_range, desc=\"2.5D Slices\", leave=False):\n        patches, valid_coords = [], []\n        for y, x in coords:\n            y1, x1 = min(y + PATCH_SIZE_25D, h), min(x + PATCH_SIZE_25D, w)\n            y0, x0 = y1 - PATCH_SIZE_25D, x1 - PATCH_SIZE_25D\n            p = vol[z - Z_CONTEXT_25D : z + Z_CONTEXT_25D + 1, y0:y1, x0:x1]\n            p = p.transpose(1, 2, 0)\n            patches.append(normalize_patch(p))\n            valid_coords.append((y0, y1, x0, x1))\n\n        for i in range(0, len(patches), BATCH_SIZE_25D):\n            batch_np = np.stack(patches[i : i + BATCH_SIZE_25D])\n            preds = ensemble_fn(batch_np)  # (B, H, W, 1)\n\n            preds_before_pp = np.array(preds)[..., 0]\n\n            if postproc_fn is not None and apply_before:\n                p_min, p_max = jnp.min(preds), jnp.max(preds)\n                preds = postproc_fn((preds - p_min) / (p_max - p_min + 1e-6))\n\n            preds_after_pp = np.array(preds)[..., 0]\n\n            for k in range(len(preds_after_pp)):\n                y0, y1, x0, x1 = valid_coords[i + k]\n                vol_probs_before_pp[z, y0:y1, x0:x1] = preds_before_pp[k]\n                vol_probs_after_pp[z, y0:y1, x0:x1] = preds_after_pp[k]\n    \n    print(\"  -> Scaling 2.5D probabilities by 10 and clipping...\")\n    # vol_probs_after_pp = np.clip(vol_probs_after_pp * 10.0, 0.0, 1.0)\n\n    if NORMALIZE_OUTPUT_25D:\n        v_min, v_max = vol_probs_after_pp.min(), vol_probs_after_pp.max()\n        vol_probs_after_pp = (vol_probs_after_pp - v_min) / (v_max - v_min + 1e-6)\n\n    return vol_probs_after_pp, vol_probs_before_pp\n\n# ============================================================================\n# 8. POST-PROCESSING (3D TOPOLOGY)\n# ============================================================================\ndef build_anisotropic_struct(z_radius, xy_radius):\n    z, r = z_radius, xy_radius\n    if z == 0 and r == 0:\n        return None\n    depth = 2 * z + 1\n    size = 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:\n                    struct[cz + dz, cy + dy, cx + dx] = True\n    return struct\n\n\ndef topo_postprocess(probs, t_low, t_high, z_rad, xy_rad, dust):\n    print(f\"  -> Post-processing (T_High={t_high}, T_Low={t_low}, Z_Rad={z_rad})...\")\n    strong = probs >= t_high\n    weak = probs >= t_low\n\n    if not strong.any():\n        return np.zeros_like(probs, 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(probs, dtype=np.uint8)\n\n    struct = build_anisotropic_struct(z_rad, xy_rad)\n    if struct is not None:\n        mask = ndi.binary_closing(mask, structure=struct)\n\n    if dust > 0:\n        mask = remove_small_objects(mask.astype(bool), min_size=dust)\n\n    return mask.astype(np.uint8)\n\n# ============================================================================\n# 9. DEBUG + VISUALIZATION HELPERS\n# ============================================================================\n#To make the model more aggressive, we introduce a Power Factor (γ). This allows you to control how quickly the model switches between 3D and 2.5D based on disagreement.\n\n#To favor 3D (Safety First): Increase γ. This makes the gate stay closer to 0 unless agreement is nearly perfect, forcing the model to rely on the 3D mask.\n\n#To favor 2.5D (Ink Recovery): Decrease γ (towards 0.5). This allows more 2.5D signal through even if there is slight disagreement.\n\n\n\ndef visualize_slice_comparison(debug_data, vid=None, z=None):\n    if not debug_data:\n        return\n    if vid is None:\n        vid = list(debug_data.keys())[0]\n    d = debug_data[vid]\n    vol_raw = d[\"vol_raw\"]\n    p3d = d[\"probs_3d\"]\n    p25_before = d[\"probs_25d_before_pp\"]\n    p25_after = d[\"probs_25d_after_pp\"]\n    p_final_before2d = d[\"final_probs_before_2d\"]\n    p_final_after2d = d[\"final_probs_after_2d\"]\n    mask = d[\"final_mask\"]\n\n    if z is None:\n        z = vol_raw.shape[1] // 2\n\n    fig, axes = plt.subplots(1, 7, figsize=(32, 4))\n\n    axes[0].imshow(np.squeeze(vol_raw[0, z]), cmap=\"gray\")\n    axes[0].set_title(\"Raw CT\")\n\n    axes[1].imshow(p3d[z], cmap=\"magma\", vmin=0, vmax=1)\n    axes[1].set_title(\"3D Prob\")\n\n    axes[2].imshow(p25_before[z], cmap=\"magma\", vmin=0, vmax=1)\n    axes[2].set_title(\"2.5D Before 2D-PP\")\n\n    axes[3].imshow(p25_after[z], cmap=\"magma\", vmin=0, vmax=1)\n    axes[3].set_title(\"2.5D After 2D-PP\")\n\n    axes[4].imshow(p_final_before2d[z], cmap=\"magma\", vmin=0, vmax=1)\n    axes[4].set_title(\"Ensemble Before 2D-PP\")\n\n    axes[5].imshow(p_final_after2d[z], cmap=\"magma\", vmin=0, vmax=1)\n    axes[5].set_title(\"Ensemble After 2D-PP\")\n\n    axes[6].imshow(mask[z], cmap=\"gray\")\n    axes[6].set_title(\"Final Mask\")\n\n    for ax in axes:\n        ax.axis(\"off\")\n\n    plt.tight_layout()\n    plt.savefig(\"debug_slice_comparison.png\", dpi=150)\n    plt.show()\n\ndef overlay_prediction(img, mask, alpha=0.4):\n    \"\"\"\n    img: (H, W) or (H, W, 3) original grayscale slice\n    mask: (H, W) predicted mask (0..1)\n    alpha: transparency of overlay\n    \"\"\"\n    if img.ndim == 2:\n        img_rgb = np.stack([img, img, img], axis=-1)\n    else:\n        img_rgb = img.copy()\n\n    # Normalize image for display\n    img_rgb = img_rgb.astype(np.float32)\n    img_rgb /= (img_rgb.max() + 1e-6)\n\n    # Create red overlay\n    red = np.zeros_like(img_rgb)\n    red[..., 0] = mask  # red channel = mask\n\n    # Blend\n    overlay = (1 - alpha) * img_rgb + alpha * red\n    overlay = np.clip(overlay, 0, 1)\n\n    return overlay\ndef preview_slice_with_overlay(vol_raw, probs_25d, slice_idx):\n    \"\"\"\n    vol_raw: (D, H, W) original CT volume\n    probs_25d: (D, H, W) predicted mask\n    slice_idx: which slice to preview\n    \"\"\"\n    img = vol_raw[slice_idx]\n    mask = probs_25d[slice_idx]\n\n    overlay = overlay_prediction(img, mask)\n\n    plt.figure(figsize=(14, 6))\n\n    plt.subplot(1, 3, 1)\n    plt.title(\"Original Slice\")\n    plt.imshow(img, cmap='gray')\n    plt.axis('off')\n\n    plt.subplot(1, 3, 2)\n    plt.title(\"Prediction (Mask)\")\n    plt.imshow(mask, cmap='Reds')\n    plt.axis('off')\n\n    plt.subplot(1, 3, 3)\n    plt.title(\"Overlay (Red Mask)\")\n    plt.imshow(overlay)\n    plt.axis('off')\n\n    plt.show()\n\ndef visualize_confidence_diff(debug_data, vid=None, z=None):\n    if not debug_data:\n        return\n    if vid is None:\n        vid = list(debug_data.keys())[0]\n    d = debug_data[vid]\n    p3d = d[\"probs_3d\"]\n    p25_after = d[\"probs_25d_after_pp\"]\n    diff = np.abs(p3d - p25_after)\n\n    if z is None:\n        z = diff.shape[0] // 2\n\n    plt.figure(figsize=(6, 5))\n    plt.imshow(diff[z], cmap=\"inferno\")\n    plt.title(f\"Confidence Difference |3D - 2.5D| (z={z})\")\n    plt.colorbar()\n    plt.axis(\"off\")\n    plt.tight_layout()\n    plt.savefig(\"debug_confidence_diff.png\", dpi=150)\n    plt.show()\n\n\ndef visualize_2d_postproc_diff(debug_data, vid=None, z=None):\n    if not debug_data:\n        return\n    if vid is None:\n        vid = list(debug_data.keys())[0]\n    d = debug_data[vid]\n    p25_before = d[\"probs_25d_before_pp\"]\n    p25_after = d[\"probs_25d_after_pp\"]\n    diff = np.abs(p25_after - p25_before)\n\n    if z is None:\n        z = diff.shape[0] // 2\n\n    plt.figure(figsize=(6, 5))\n    plt.imshow(diff[z], cmap=\"magma\")\n    plt.title(f\"2D Postproc Effect |After - Before| (z={z})\")\n    plt.colorbar()\n    plt.axis(\"off\")\n    plt.tight_layout()\n    plt.savefig(\"debug_2d_postproc_diff.png\", dpi=150)\n    plt.show()\n\n\ndef make_side_by_side_gif(debug_data, vid=None, out_path=\"debug_comparison.gif\"):\n    if not debug_data:\n        return\n    if vid is None:\n        vid = list(debug_data.keys())[0]\n    d = debug_data[vid]\n    vol_raw = d[\"vol_raw\"]\n    p25_before = d[\"probs_25d_before_pp\"]\n    p25_after = d[\"probs_25d_after_pp\"]\n    final_mask = d[\"final_mask\"]\n\n    z_max = vol_raw.shape[1]\n    frames = []\n\n    for z in range(0, z_max, max(1, z_max // 30)):\n        fig, axes = plt.subplots(1, 4, figsize=(10, 3))\n\n        axes[0].imshow(np.squeeze(vol_raw[0, z]), cmap=\"gray\")\n        axes[0].set_title(f\"CT z={z}\")\n        axes[0].axis(\"off\")\n\n        axes[1].imshow(p25_before[z], cmap=\"magma\", vmin=0, vmax=1)\n        axes[1].set_title(\"2.5D Before 2D-PP\")\n        axes[1].axis(\"off\")\n\n        axes[2].imshow(p25_after[z], cmap=\"magma\", vmin=0, vmax=1)\n        axes[2].set_title(\"2.5D After 2D-PP\")\n        axes[2].axis(\"off\")\n\n        axes[3].imshow(final_mask[z], cmap=\"gray\")\n        axes[3].set_title(\"Final Mask\")\n        axes[3].axis(\"off\")\n\n        plt.tight_layout()\n        fig.canvas.draw()\n        buf = np.asarray(fig.canvas.buffer_rgba())\n        frames.append(Image.fromarray(buf))\n        plt.close(fig)\n\n    if frames:\n        frames[0].save(\n            out_path,\n            save_all=True,\n            append_images=frames[1:],\n            duration=100,\n            loop=0,\n        )\n        print(f\"GIF saved to {out_path}\")\n\n\ndef visualize_3d_overview(debug_data, vid=None):\n\n\n    if not debug_data: return\n    d = debug_data[vid]\n    # Use the probability map instead of the binary mask for better debugging\n    data_to_viz = d[\"probs_3d\"] \n    \n    proj_axial = data_to_viz.max(axis=0)\n    proj_coronal = data_to_viz.max(axis=1)\n    proj_sagittal = data_to_viz.max(axis=2)\n\n    fig, axes = plt.subplots(1, 3, figsize=(12, 4))\n    # Using 'magma' or 'inferno' helps see intensity variations better than 'gray'\n    axes[0].imshow(proj_axial, cmap=\"magma\", vmin=0, vmax=1)\n    axes[1].imshow(proj_coronal, cmap=\"magma\", vmin=0, vmax=1)\n    axes[2].imshow(proj_sagittal, cmap=\"magma\", vmin=0, vmax=1)\n    # if not debug_data:\n    #     return\n    # if vid is None:\n    #     vid = list(debug_data.keys())[0]\n    # d = debug_data[vid]\n    # final_mask = d[\"final_mask\"]\n\n    # proj_axial = final_mask.max(axis=0)\n    # proj_coronal = final_mask.max(axis=1)\n    # proj_sagittal = final_mask.max(axis=2)\n\n    # fig, axes = plt.subplots(1, 3, figsize=(12, 4))\n    # axes[0].imshow(proj_axial, cmap=\"gray\")\n    # axes[0].set_title(\"Axial MIP\")\n    # axes[1].imshow(proj_coronal, cmap=\"gray\")\n    # axes[1].set_title(\"Coronal MIP\")\n    # axes[2].imshow(proj_sagittal, cmap=\"gray\")\n    # axes[2].set_title(\"Sagittal MIP\")\n    for ax in axes:\n        ax.axis(\"off\")\n    plt.tight_layout()\n    plt.savefig(\"debug_3d_overview.png\", dpi=150)\n    plt.show()\n\n# ============================================================================\n# 10. MAIN PIPELINE\n# ============================================================================\nmodels_3d = get_3d_models()\nens_fn_25d, pp_fn_25d = get_25d_funcs()\ntest_df = pd.read_csv(f\"{ROOT_DIR}/test.csv\")\n\ndebug_data = {}\n\nprint(\"\\nStarting Hybrid Inference...\")\nwith zipfile.ZipFile(ZIP_PATH, \"w\", compression=zipfile.ZIP_DEFLATED) as zf:\n    for idx, row in test_df.iterrows():\n        image_id = row[\"id\"]\n        print(f\"\\n[{idx + 1}/{len(test_df)}] Processing {image_id}...\")\n\n        vol_path = f\"{TEST_DIR}/{image_id}.tif\"\n        vol_raw = tifffile.imread(vol_path).astype(np.float32)\n        vol_raw = vol_raw[None, ..., None]  # (1, D, H, W, 1)\n\n        probs_3d = predict_3d_volume(vol_raw, models_3d)\n        probs_25d_after, probs_25d_before = predict_25d_volume(vol_raw, ens_fn_25d, pp_fn_25d)\n\n        # if probs_3d is not None and probs_25d_after is not None:\n        #     print(f\"  -> Ensembling 3D ({W_3D}) + 2.5D ({W_25D})\")\n        #     final_probs_before_2d = probs_3d * W_3D + probs_25d_after * W_25D\n        # elif probs_3d is not None:\n        #     print(\"  -> Using 3D Only\")\n        #     final_probs_before_2d = probs_3d\n        # elif probs_25d_after is not None:\n        #     print(\"  -> Using 2.5D Only\")\n        #     final_probs_before_2d = probs_25d_after\n        # else:\n        #     final_probs_before_2d = np.zeros(vol_raw.shape[1:4], dtype=np.float32)\n        ########\n        if probs_3d is not None and probs_25d_after is not None:\n            # 1. Create a \"Hard Mask\" from 3D\n            # Anything below this threshold in 3D is considered absolute background\n            hard_mask = (probs_3d > 0.1).astype(np.float32) \n            \n            # 2. Calculate Agreement Gate\n            # High difference = lower gate = trust 3D more\n            conf_diff = np.abs(probs_3d - probs_25d_after)\n            gate = np.power(np.clip(1.0 - conf_diff, 0, 1), GATE_GAMMA)\n            \n            # 3. Apply the Hybrid Mix\n            # Result is 2.5D where they agree, 3D where they don't\n            hybrid_mix = (probs_25d_after * gate) + (probs_3d * (1.0 - gate))\n            \n            # 4. Strict Suppression\n            # Force the 3D \"Black Mask\" to be final\n            final_probs_before_2d = hybrid_mix * hard_mask\n        ########\n        # if probs_3d is not None and probs_25d_after is not None:\n        #     print(f\"  -> Applying 3D Masking and Confidence Gating...\")\n            \n        #     # 1. Create the \"3D Black Mask\" \n        #     # This identifies where the 3D model is confident there is NO surface\n        #     mask_3d_binary = (probs_3d > T_LOW).astype(np.float32)\n            \n        #     # 2. Predict 2.5D using masked 3D (Element-wise multiplication)\n        #     # This forces 2.5D noise in the \"black\" areas of 3D to zero\n        #     masked_25d = probs_25d_after * mask_3d_binary\n            \n        #     # 3. Confidence Map Calculation |3D - 2.5D|\n        #     # High diff means the models disagree on the ink/surface\n        #     conf_diff = np.abs(probs_3d - probs_25d_after)\n            \n        #     # 4. Gated Mixing\n        #     # Where confidence diff is high, we lean on 3D (W_3D). \n        #     # Where they agree, we allow the 2.5D details to shine.\n        #     gate = np.clip(1.0 - conf_diff, 0, 1) # High agreement = 1.0\n            \n        #     # New Hybrid Formula: \n        #     # Base is the masked 2.5D, but we blend in 3D based on disagreement\n        #     final_probs_before_2d = (masked_25d * gate) + (probs_3d * (1 - gate))\n            \n        #     # Optional: Hard suppress based on 3D \"Black Mask\"\n        #     # This ensures \"all the 0 resulted from 3D should be mostly 0\"\n        #     final_probs_before_2d = final_probs_before_2d * mask_3d_binary\n        \n        elif probs_3d is not None:\n            final_probs_before_2d = probs_3d\n        else:\n            final_probs_before_2d = probs_25d_after\n                \n        final_probs_after_2d = final_probs_before_2d.copy()\n        apply_after = (pp_fn_25d is not None) and (POSTPROC_2D_LOCATION in [\"after\", \"both\"])\n\n        # if apply_after:\n        #     print(\"  -> Applying 2D Post-Processor on ensemble probabilities...\")\n        #     fp = final_probs_before_2d  # (D, H, W)\n        #     D, H, W = fp.shape\n        #     final_probs_after_2d = np.zeros_like(fp, dtype=np.float32)\n            \n        #     for z in range(D):\n        #         slice_np = fp[z][None, ..., None]  # (1, H, W, 1)\n        #         slice_jax = jnp.array(slice_np)\n        #         p_min, p_max = jnp.min(slice_jax), jnp.max(slice_jax)\n        #         slice_norm = (slice_jax - p_min) / (p_max - p_min + 1e-6)\n        #         slice_pp = pp_fn_25d(slice_norm)\n        #         final_probs_after_2d[z] = np.array(slice_pp)[0, ..., 0]\n        if apply_after:\n            print(\"  -> Applying 2D Post-Processor on ensemble probabilities (slice-wise)...\")\n            fp = final_probs_before_2d  # (D, H, W)\n            D, H, W = fp.shape\n            final_probs_after_2d = np.zeros_like(fp, dtype=np.float32)\n        \n            for z in range(D):\n                slice_np = fp[z][None, ..., None]  # (1, H, W, 1)\n                slice_jax = jnp.array(slice_np)\n        \n                # ❌ DO NOT NORMALIZE HERE\n                slice_pp = pp_fn_25d(slice_jax)\n        \n                final_probs_after_2d[z] = np.array(slice_pp)[0, ..., 0]\n\n        if POSTPROC_MODE == \"topo\":\n            final_mask = topo_postprocess(\n                final_probs_after_2d,\n                t_low=T_LOW,\n                t_high=T_HIGH,\n                z_rad=Z_RADIUS,\n                xy_rad=XY_RADIUS,\n                dust=DUST_MIN_SIZE,\n            )\n        else:\n            final_mask = (final_probs_after_2d > T_HIGH).astype(np.uint8)\n\n        print(f\"    Foreground voxels: {final_mask.sum():,}\")\n        out_path = f\"{OUTPUT_DIR}/{image_id}.tif\"\n        tifffile.imwrite(out_path, final_mask)\n        zf.write(out_path, arcname=f\"{image_id}.tif\")\n        os.remove(out_path)\n\n        if DEBUG_ENABLE and len(debug_data) < DEBUG_MAX_VOLUMES:\n            debug_data[image_id] = {\n                \"vol_raw\": vol_raw,\n                \"probs_3d\": probs_3d if probs_3d is not None else np.zeros_like(final_probs_before_2d),\n                \"probs_25d_before_pp\": probs_25d_before if probs_25d_before is not None else np.zeros_like(final_probs_before_2d),\n                \"probs_25d_after_pp\": probs_25d_after if probs_25d_after is not None else np.zeros_like(final_probs_before_2d),\n                \"final_probs_before_2d\": final_probs_before_2d,\n                \"final_probs_after_2d\": final_probs_after_2d,\n                \"final_mask\": final_mask,\n            }\n\n        del vol_raw, probs_3d, probs_25d_before, probs_25d_after, final_probs_before_2d, final_probs_after_2d, final_mask\n        gc.collect()\n\nprint(\"\\n\" + \"=\" * 60)\nprint(\"HYBRID PIPELINE COMPLETE\")\nprint(f\"Submission: {ZIP_PATH}\")\nprint(\"=\" * 60)\n\n# ============================================================================\n# 11. SIMPLE SUBMISSION VISUALIZATION + DEBUG VIEWERS\n# ============================================================================\ndef visualize_submission(zip_path=\"submission.zip\"):\n    if not os.path.exists(zip_path):\n        return\n    with zipfile.ZipFile(zip_path, \"r\") as z:\n        tif_files = [f for f in z.namelist() if f.endswith(\".tif\")]\n        if not tif_files:\n            return\n        chosen_file = random.choice(tif_files)\n        z.extract(chosen_file, path=\"temp_viz\")\n        temp_path = Path(f\"temp_viz/{chosen_file}\")\n    try:\n        with Image.open(str(temp_path)) as img:\n            vol = np.stack([np.array(f) for f in ImageSequence.Iterator(img)], axis=0)\n        z_dim = vol.shape[0]\n        indices = [int(z_dim * p) for p in [0.25, 0.50, 0.75]]\n        fig, axes = plt.subplots(1, 3, figsize=(12, 4))\n        for i, idx in enumerate(indices):\n            if idx < z_dim:\n                axes[i].imshow(vol[idx], cmap=\"gray\")\n                axes[i].set_title(f\"Slice {idx}\")\n                axes[i].axis(\"off\")\n        plt.suptitle(f\"Preview: {chosen_file}\")\n        plt.tight_layout()\n        plt.savefig(\"submission_preview.png\", dpi=150)\n        plt.show()\n    finally:\n        if os.path.exists(\"temp_viz\"):\n            import shutil\n            shutil.rmtree(\"temp_viz\")\ndef visualize_hybrid_logic(debug_data, vid=None, z=160):\n    d = debug_data[vid]\n    p3d = d[\"probs_3d\"][z]\n    p25 = d[\"probs_25d_after_pp\"][z]\n    \n    diff = np.abs(p3d - p25)\n    mask = (p3d > T_LOW).astype(np.float32)\n    final = d[\"final_probs_after_2d\"][z]\n\n    fig, axes = plt.subplots(1, 4, figsize=(20, 5))\n    axes[0].imshow(p3d, cmap='magma'); axes[0].set_title(\"3D Prob\")\n    axes[1].imshow(p25, cmap='magma'); axes[1].set_title(\"2.5D Prob\")\n    axes[2].imshow(diff, cmap='inferno'); axes[2].set_title(\"Confidence Diff\")\n    axes[3].imshow(final, cmap='magma'); axes[3].set_title(\"Result (Masked & Gated)\")\n    plt.show()\n\nif __name__ == \"__main__\":\n    visualize_submission(ZIP_PATH)\n    if DEBUG_ENABLE and debug_data:\n        vid = list(debug_data.keys())[0]\n        visualize_slice_comparison(debug_data, vid=vid)\n        visualize_confidence_diff(debug_data, vid=vid)\n        visualize_2d_postproc_diff(debug_data, vid=vid)\n        make_side_by_side_gif(debug_data, vid=vid, out_path=\"debug_comparison.gif\")\n        visualize_3d_overview(debug_data, vid=vid)\n        visualize_hybrid_logic(debug_data, vid=vid)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T18:04:29.835717Z","iopub.execute_input":"2026-02-28T18:04:29.836314Z","iopub.status.idle":"2026-02-28T18:06:59.963008Z","shell.execute_reply.started":"2026-02-28T18:04:29.836279Z","shell.execute_reply":"2026-02-28T18:06:59.962102Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"* Color intensity = level of disagreement\n* - Dark / near‑black → models agree (difference close to 0)\n* - Bright yellow / white → models strongly disagree (difference close to 1)\n* We use a perceptually strong colormap (e.g., inferno) so disagreements stand out clearly.\n\n🧠 Why this plot is useful\n1. Detect weak or uncertain regions\nIf both models agree, the difference is low → high reliability.\nIf they disagree, the region is ambiguous → needs attention.\n2. Identify model‑specific biases\n- 3D model may detect long continuous structures\n- 2.5D model may detect fine surface texture\nThe difference map highlights where each model “sees” something the other doesn’t.\n3. Tune ensemble weights\nIf you see:\n- 3D is consistently more confident → increase W_3D\n- 2.5D is consistently sharper → increase W_25D\n4. Debug post‑processing\nIf disagreement is high but final mask is clean, your hybrid post‑processing is doing its job.\nIf disagreement is high and final mask is noisy, thresholds or topology parameters may need tuning.\n","metadata":{}},{"cell_type":"markdown","source":"### Overlay map","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport matplotlib.pyplot as plt\nfrom tifffile import imread\nimport zipfile\nimport os\n\n# ---------------------------------------\n# Overlay helper\n# ---------------------------------------\ndef overlay_prediction(img, mask, alpha=0.4):\n    if img.ndim == 2:\n        img_rgb = np.stack([img, img, img], axis=-1)\n    else:\n        img_rgb = img.copy()\n\n    img_rgb = img_rgb.astype(np.float32)\n    img_rgb /= (img_rgb.max() + 1e-6)\n\n    red = np.zeros_like(img_rgb)\n    red[..., 0] = mask  # red channel = mask\n\n    overlay = (1 - alpha) * img_rgb + alpha * red\n    return np.clip(overlay, 0, 1)\n\n\n# ---------------------------------------\n# Preview 4 slices\n# ---------------------------------------\ndef preview_4_slices_from_submission(\n    submission_zip,\n    tif_folder,\n    tif_name,\n    slice_indices=[50, 100, 150, 200],\n    alpha=0.4\n):\n    # 1. Load original TIFF\n    tif_path = os.path.join(tif_folder, tif_name)\n    print(f\"Loading original volume: {tif_path}\")\n    vol_raw = imread(tif_path)  # (D, H, W)\n\n    # 2. Extract prediction from submission.zip\n    print(f\"Extracting prediction from: {submission_zip}\")\n    with zipfile.ZipFile(submission_zip, \"r\") as z:\n        pred_filename = tif_name  # prediction has same name\n        z.extract(pred_filename, \"/kaggle/working/\")\n        pred_path = f\"/kaggle/working/{pred_filename}\"\n\n    print(f\"Loading prediction: {pred_path}\")\n    pred = imread(pred_path)  # (D, H, W)\n\n    # 3. Plot 4 slices\n    plt.figure(figsize=(18, 12))\n\n    for i, idx in enumerate(slice_indices):\n        img = vol_raw[idx]\n        mask = pred[idx]\n        overlay = overlay_prediction(img, mask, alpha=alpha)\n\n        plt.subplot(2, 2, i + 1)\n        plt.title(f\"Slice {idx}\")\n        plt.imshow(overlay)\n        plt.axis(\"off\")\n\n    plt.show()\n\n\n# ---------------------------------------\n# Example usage\n# ---------------------------------------\npreview_4_slices_from_submission(\n    submission_zip=\"/kaggle/working/submission.zip\",\n    tif_folder=\"/kaggle/input/vesuvius-challenge-surface-detection/test_images\",\n    tif_name=\"1407735.tif\",\n    slice_indices=[50, 120, 240, 300],\n    alpha=0.45\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T18:06:59.964739Z","iopub.execute_input":"2026-02-28T18:06:59.965218Z","iopub.status.idle":"2026-02-28T18:07:00.747506Z","shell.execute_reply.started":"2026-02-28T18:06:59.965192Z","shell.execute_reply":"2026-02-28T18:07:00.746598Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}