{"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":"none","dataSources":[{"sourceId":117682,"databundleVersionId":15062069,"sourceType":"competition"},{"sourceId":14245247,"sourceType":"datasetVersion","datasetId":9088503},{"sourceId":14295835,"sourceType":"datasetVersion","datasetId":9125518},{"sourceId":14871252,"sourceType":"datasetVersion","datasetId":9513632},{"sourceId":288572598,"sourceType":"kernelVersion"},{"sourceId":294535277,"sourceType":"kernelVersion"},{"sourceId":297302543,"sourceType":"kernelVersion"},{"sourceId":297674933,"sourceType":"kernelVersion"},{"sourceId":297771298,"sourceType":"kernelVersion"},{"sourceId":297827765,"sourceType":"kernelVersion"},{"sourceId":298023302,"sourceType":"kernelVersion"}],"dockerImageVersionId":31260,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-18T08:49:39.594151Z","iopub.execute_input":"2026-02-18T08:49:39.594554Z","iopub.status.idle":"2026-02-18T08:49:39.606025Z","shell.execute_reply.started":"2026-02-18T08:49:39.594524Z","shell.execute_reply":"2026-02-18T08:49:39.604659Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# SETUP PACKAGES\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)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-18T08:49:39.616582Z","iopub.execute_input":"2026-02-18T08:49:39.617366Z","iopub.status.idle":"2026-02-18T08:49:43.894738Z","shell.execute_reply.started":"2026-02-18T08:49:39.617329Z","shell.execute_reply":"2026-02-18T08:49:43.893610Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"PRED_AVERGE=False\nif PRED_AVERGE:\n    import os\n    import gc\n    import zipfile\n    from pathlib import Path\n    \n    import numpy as np\n    import jax\n    import jax.numpy as jnp\n    import flax.linen as nn\n    from flax.traverse_util import unflatten_dict\n    from PIL import Image, ImageSequence\n    from tqdm import tqdm\n    import scipy.ndimage as ndi\n    from skimage.morphology import remove_small_objects\n    \n    # =========================================================\n    # CONFIG\n    # =========================================================\n    DATA_PATH = Path(\"/kaggle/input/vesuvius-challenge-surface-detection\")\n    \n    WEIGHT_PATHS_25D = [\n        # \"/kaggle/input/notebooks/crischir/test-code-for-jax-pix2pix/checkpoints/pix2pix_best.npz\",\n        # \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix/checkpoints/pix2pix_ep_30.npz\",\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_30.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-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/jax-pix2pix/checkpoints/pix2pix_ep_130.npz\",\n        # \"/kaggle/input/notebooks/crischir/jax-pix2pix/checkpoints/pix2pix_ep_80.npz\",\n    ]\n    \n    POST_PROC_WEIGHTS_2D = \"/kaggle/input/datasets/crischir/postprocessvesuviuscomplete/post_proc_checkpoints/post_proc_ep30.npz\"\n    \n    PATCH_SIZE = 320\n    Z_CONTEXT = 3\n    BATCH_SIZE = 128\n    \n    USE_ML_POST_PROC_2D = True\n    USE_ANISOTROPIC_HYST = True\n    NORMALIZE_OUTPUT = True\n    \n    T_LOW = 0.037\n    T_HIGH = 0.95\n    Z_RADIUS = 5\n    XY_RADIUS = 2\n    DUST_MIN_SIZE = 3\n    \n    # =========================================================\n    # MODELS (FLAX)\n    # =========================================================\n    class Pix2PixGenerator(nn.Module):\n        @nn.compact\n        def __call__(self, x, train: bool = False):\n    \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            # Encoder\n            s1 = nn.leaky_relu(\n                nn.Conv(64, (4, 4), strides=(2, 2), padding='SAME')(x), 0.2\n            )\n            s2 = conv_bn_leaky(s1, 128)\n            s3 = conv_bn_leaky(s2, 256)\n    \n            # Decoder\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    \n            # Output\n            return nn.ConvTranspose(1, (4, 4), strides=(2, 2), padding='SAME')(u2)\n    \n    \n    class 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            out = nn.ConvTranspose(1, (4, 4), strides=(2, 2), padding=\"SAME\")(u2)\n            return out\n    \n    # =========================================================\n    # UTILITIES\n    # =========================================================\n    def load_flax_weights(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 {\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    \n    def normalize_patch(patch):\n        patch = patch.astype(np.float32)\n        return (patch - patch.mean()) / (patch.std() + 1e-6)\n    \n    \n    def 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    \n    def apply_anisotropic_hysteresis(vol_probs):\n        print(f\"    ...Running 3D Hysteresis (High={T_HIGH}, Low={T_LOW})...\")\n        strong = vol_probs >= T_HIGH\n        weak = vol_probs >= T_LOW\n    \n        if not strong.any():\n            return np.zeros_like(vol_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(vol_probs, dtype=np.uint8)\n    \n        struct_close = build_anisotropic_struct(Z_RADIUS, XY_RADIUS)\n        if struct_close is not None:\n            mask = ndi.binary_closing(mask, structure=struct_close)\n    \n        if DUST_MIN_SIZE > 0:\n            mask = remove_small_objects(mask.astype(bool), min_size=DUST_MIN_SIZE)\n    \n        return mask.astype(np.uint8) * 255\n    \n    \n    def save_tif_stack(vol, out_name):\n        pages = [Image.fromarray(vol[i]) for i in range(vol.shape[0])]\n        pages[0].save(out_name, save_all=True, append_images=pages[1:], compression=\"tiff_deflate\")\n    \n    # =========================================================\n    # 2.5D INFERENCE FUNCTIONS\n    # =========================================================\n    def make_25d_ensemble_fn(ensemble_vars):\n        model = Pix2PixGenerator()\n    \n        def _single_forward(vars, x):\n            return nn.sigmoid(model.apply(vars, x, train=False))\n    \n        def _ensemble_forward(x):\n            all_probs = jnp.stack([_single_forward(v, x) for v in ensemble_vars], axis=0)\n            return jnp.mean(all_probs, axis=0)\n    \n        return jax.jit(_ensemble_forward)\n    \n    \n    def make_postproc_2d_fn(post_vars):\n        model = PostProcessModel2D()\n    \n        def _forward(x):\n            return nn.sigmoid(model.apply(post_vars, x, train=False))\n    \n        return jax.jit(_forward)\n    \n    # =========================================================\n    # MAIN 2.5D PIPELINE\n    # =========================================================\n    def run_inference_25d():\n        print(\"🔧 2.5D CONFIG: ML_PostProc_2D=\"\n              f\"{USE_ML_POST_PROC_2D} | Anisotropic_Hyst={USE_ANISOTROPIC_HYST}\")\n    \n        # 1. Load 2.5D weights\n        ensemble_vars_25d = [load_flax_weights(p) for p in WEIGHT_PATHS_25D]\n        ensemble_vars_25d = [v for v in ensemble_vars_25d if v is not None]\n        if not ensemble_vars_25d:\n            print(\"❌ No 2.5D weights found.\")\n            return\n    \n        # Init model once (shape check)\n        dummy = jnp.zeros((1, PATCH_SIZE, PATCH_SIZE, 2 * Z_CONTEXT + 1))\n        _ = Pix2PixGenerator().init(jax.random.PRNGKey(0), dummy, train=False)\n    \n        ensemble_25d_fn = make_25d_ensemble_fn(ensemble_vars_25d)\n    \n        # Optional 2D post-proc\n        postproc_2d_fn = None\n        if USE_ML_POST_PROC_2D and os.path.exists(POST_PROC_WEIGHTS_2D):\n            pv = load_flax_weights(POST_PROC_WEIGHTS_2D)\n            if pv:\n                dummy_post = jnp.zeros((1, PATCH_SIZE, PATCH_SIZE, 1))\n                _ = PostProcessModel2D().init(jax.random.PRNGKey(0), dummy_post, train=False)\n                postproc_2d_fn = make_postproc_2d_fn(pv)\n            else:\n                print(\"⚠️ 2D post-proc weights missing. Disabling ML post-proc.\")\n                postproc_2d_fn = None\n    \n        # 2. Process test volumes\n        test_dir = DATA_PATH / \"test_images\"\n        if not test_dir.exists():\n            print(f\"❌ Test dir not found: {test_dir}\")\n            return\n    \n        test_ids = [f.stem for f in test_dir.glob(\"*.tif\")]\n        submission_files = []\n    \n        for vid in test_ids:\n            print(f\"\\nProcessing volume {vid}...\")\n    \n            try:\n                with Image.open(str(test_dir / f\"{vid}.tif\")) as img:\n                    vol = np.stack([np.array(f) for f in ImageSequence.Iterator(img)], axis=0)\n            except Exception as e:\n                print(f\"Error loading {vid}: {e}\")\n                continue\n    \n            z_max, h, w = vol.shape\n            vol_probs = np.zeros((z_max, h, w), dtype=np.float32)\n    \n            coords = [(y, x) for y in range(0, h, PATCH_SIZE) for x in range(0, w, PATCH_SIZE)]\n    \n            # --- 2.5D INFERENCE LOOP ---\n            for z in tqdm(range(Z_CONTEXT, z_max - Z_CONTEXT), desc=f\"2.5D {vid}\"):\n                patches_list = []\n                valid_coords = []\n    \n                for y, x in coords:\n                    y1, x1 = min(y + PATCH_SIZE, h), min(x + PATCH_SIZE, w)\n                    y0, x0 = y1 - PATCH_SIZE, x1 - PATCH_SIZE\n                    p = vol[z - Z_CONTEXT : z + Z_CONTEXT + 1, y0:y1, x0:x1].transpose(1, 2, 0)\n                    patches_list.append(normalize_patch(p))\n                    valid_coords.append((y0, y1, x0, x1))\n    \n                num_patches = len(patches_list)\n                for i in range(0, num_patches, BATCH_SIZE):\n                    batch_in = patches_list[i : i + BATCH_SIZE]\n                    current_bs = len(batch_in)\n                    batch_np = np.stack(batch_in)  # (B, H, W, D)\n    \n                    preds = ensemble_25d_fn(batch_np)  # (B, H, W, 1)\n    \n                    if postproc_2d_fn is not None:\n                        p_min, p_max = jnp.min(preds), jnp.max(preds)\n                        norm_preds = (preds - p_min) / (p_max - p_min + 1e-6)\n                        preds = postproc_2d_fn(norm_preds)\n    \n                    preds_flat = np.array(preds)[..., 0]  # (B, H, W)\n    \n                    for k in range(current_bs):\n                        global_idx = i + k\n                        y0, y1, x0, x1 = valid_coords[global_idx]\n                        vol_probs[z, y0:y1, x0:x1] = preds_flat[k]\n    \n            # --- OPTIONAL NORMALIZATION ---\n            if NORMALIZE_OUTPUT:\n                v_min, v_max = vol_probs.min(), vol_probs.max()\n                if v_max > v_min:\n                    vol_probs = (vol_probs - v_min) / (v_max - v_min)\n    \n            # --- FINAL POST-PROCESSING ---\n            if USE_ANISOTROPIC_HYST:\n                final_mask = apply_anisotropic_hysteresis(vol_probs)\n            else:\n                final_mask = (vol_probs > T_LOW).astype(np.uint8) * 255\n    \n            out_name = f\"{vid}.tif\"\n            save_tif_stack(final_mask, out_name)\n            submission_files.append(out_name)\n    \n            del vol, vol_probs, final_mask\n            gc.collect()\n    \n        # Zip results\n        if submission_files:\n            with zipfile.ZipFile(\"submission.zip\", \"w\") as z:\n                for f in submission_files:\n                    z.write(f)\n                    os.remove(f)\n            print(\"✅ submission.zip created\")\n    \n    \n    if __name__ == \"__main__\":\n        run_inference_25d()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-18T08:49:43.896274Z","iopub.execute_input":"2026-02-18T08:49:43.896690Z","iopub.status.idle":"2026-02-18T08:49:43.944626Z","shell.execute_reply.started":"2026-02-18T08:49:43.896656Z","shell.execute_reply":"2026-02-18T08:49:43.943361Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nimport os\nimport zipfile\nimport random\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom PIL import Image, ImageSequence\nfrom pathlib import Path\n\ndef visualize_submission(zip_path=\"submission.zip\"):\n    \"\"\"\n    Extracts a random TIF from submission.zip, analyzes it, \n    and plots 3 sample slices (25%, 50%, 75% depth).\n    \"\"\"\n    if not os.path.exists(zip_path):\n        print(f\"❌ Error: {zip_path} not found.\")\n        return\n\n    # 1. Open Zip and Pick Random File\n    with zipfile.ZipFile(zip_path, 'r') as z:\n        tif_files = [f for f in z.namelist() if f.endswith('.tif')]\n        \n        if not tif_files:\n            print(\"❌ No .tif files found inside the zip.\")\n            return\n            \n        chosen_file = random.choice(tif_files)\n        print(f\"🎲 Selected random file: {chosen_file}\")\n        \n        # Extract to temporary path\n        z.extract(chosen_file, path=\"temp_viz\")\n        temp_path = Path(f\"temp_viz/{chosen_file}\")\n\n    # 2. Load Volume\n    try:\n        with Image.open(str(temp_path)) as img:\n            # Load all pages into a numpy stack\n            vol = np.stack([np.array(f) for f in ImageSequence.Iterator(img)], axis=0)\n    except Exception as e:\n        print(f\"❌ Error reading TIFF: {e}\")\n        return\n\n    z_dim, h, w = vol.shape\n    print(f\"📊 Dimensions: {z_dim} slices | {h}x{w} resolution\")\n    \n    # 3. Check Statistics\n    unique_vals = np.unique(vol)\n    print(f\"🔢 Unique values in volume: {unique_vals}\")\n    \n    if len(unique_vals) == 1 and unique_vals[0] == 0:\n        print(\"⚠️ WARNING: This volume is completely empty (All Black/Zero).\")\n    else:\n        print(\"✅ Predictions detected (Volume contains non-zero data).\")\n\n    # 4. Select Slices (25%, 50%, 75% mark to avoid empty starts/ends)\n    indices = [\n        int(z_dim * 0.25),\n        int(z_dim * 0.50),\n        int(z_dim * 0.75)\n    ]\n    \n    # 5. Plot\n    fig, axes = plt.subplots(1, 3, figsize=(15, 6))\n    fig.suptitle(f\"Visualization: {chosen_file}\", fontsize=16)\n\n    for i, idx in enumerate(indices):\n        if idx < z_dim:\n            slice_img = vol[idx]\n            axes[i].imshow(slice_img, cmap='gray', interpolation='nearest')\n            axes[i].set_title(f\"Slice Index: {idx}\\n(Depth: {idx/z_dim:.0%})\")\n            axes[i].axis('off')\n            \n            # Annotate if empty\n            if np.sum(slice_img) == 0:\n                axes[i].text(w//2, h//2, \"EMPTY SLICE\", color='red', \n                             ha='center', va='center', fontweight='bold')\n\n    plt.tight_layout()\n    plt.show()\n\n    # 6. Cleanup\n    try:\n        os.remove(temp_path)\n        os.rmdir(\"temp_viz\")\n    except:\n        pass\n\nif __name__ == \"__main__\":\n    visualize_submission()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-18T08:49:43.945947Z","iopub.execute_input":"2026-02-18T08:49:43.946612Z","iopub.status.idle":"2026-02-18T08:49:45.501451Z","shell.execute_reply.started":"2026-02-18T08:49:43.946548Z","shell.execute_reply":"2026-02-18T08:49:45.500228Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Logit average","metadata":{}},{"cell_type":"code","source":"LOGIT_AVERGE=False\nif LOGIT_AVERGE:\n    import os\n    import gc\n    import zipfile\n    from pathlib import Path\n    import numpy as np\n    import jax\n    import jax.numpy as jnp\n    import flax.linen as nn\n    from flax.traverse_util import unflatten_dict\n    from PIL import Image, ImageSequence\n    from tqdm import tqdm\n    import scipy.ndimage as ndi\n    from skimage.morphology import remove_small_objects\n    \n    # =========================================================\n    # CONFIG\n    # =========================================================\n    DATA_PATH = Path(\"/kaggle/input/vesuvius-challenge-surface-detection\")\n    WEIGHT_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/checkpoints/pix2pix_ep_30.npz\",\n        \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix/checkpoints/pix2pix_ep_20.npz\",\n    ]\n    WEIGHT_PATHS_25D = [\n        # \"/kaggle/input/notebooks/crischir/test-code-for-jax-pix2pix/checkpoints/pix2pix_best.npz\",\n        # \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix/checkpoints/pix2pix_ep_30.npz\",\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_30.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-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/jax-pix2pix/checkpoints/pix2pix_ep_130.npz\",\n        # \"/kaggle/input/notebooks/crischir/jax-pix2pix/checkpoints/pix2pix_ep_80.npz\",\n    ]\n    # --- NEW TOGGLE ---\n    USE_LOGIT_AVERAGING = True  # Set to False for Probability Averaging\n    # ------------------\n    \n    POST_PROC_WEIGHTS_2D = \"/kaggle/input/datasets/crischir/postprocessvesuviuscomplete/post_proc_checkpoints/post_proc_ep30.npz\"\n    PATCH_SIZE = 320\n    Z_CONTEXT = 3\n    BATCH_SIZE = 128\n    USE_ML_POST_PROC_2D = True\n    USE_ANISOTROPIC_HYST = False\n    NORMALIZE_OUTPUT = True\n    T_LOW = 0.017\n    T_HIGH = 0.95\n    Z_RADIUS = 5\n    XY_RADIUS = 2\n    DUST_MIN_SIZE = 3\n    \n    # =========================================================\n    # MODELS (FLAX)\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            \n            # Return raw logits (no sigmoid here anymore)\n            return nn.ConvTranspose(1, (4, 4), strides=(2, 2), padding='SAME')(u2)\n    \n    class 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(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 = nn.relu(nn.BatchNorm(use_running_average=not train)(nn.ConvTranspose(128, (4, 4), strides=(2, 2), padding=\"SAME\")(s3)))\n            u1 = jnp.concatenate([u1, s2], axis=-1)\n            u2 = nn.relu(nn.BatchNorm(use_running_average=not train)(nn.ConvTranspose(64, (4, 4), strides=(2, 2), padding=\"SAME\")(u1)))\n            u2 = jnp.concatenate([u2, s1], axis=-1)\n            out = nn.ConvTranspose(1, (4, 4), strides=(2, 2), padding=\"SAME\")(u2)\n            return out\n    \n    # =========================================================\n    # UTILITIES\n    # =========================================================\n    def load_flax_weights(path):\n        if not os.path.exists(path): 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 {\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    def normalize_patch(patch):\n        patch = patch.astype(np.float32)\n        return (patch - patch.mean()) / (patch.std() + 1e-6)\n    \n    def apply_anisotropic_hysteresis(vol_probs):\n        print(f\"    ...Running 3D Hysteresis (High={T_HIGH}, Low={T_LOW})...\")\n        strong = vol_probs >= T_HIGH\n        weak = vol_probs >= T_LOW\n        if not strong.any(): return np.zeros_like(vol_probs, dtype=np.uint8)\n        struct_hyst = ndi.generate_binary_structure(3, 3)\n        mask = ndi.binary_propagation(strong, mask=weak, structure=struct_hyst)\n        return mask.astype(np.uint8) * 255\n    \n    def save_tif_stack(vol, out_name):\n        pages = [Image.fromarray(vol[i]) for i in range(vol.shape[0])]\n        pages[0].save(out_name, save_all=True, append_images=pages[1:], compression=\"tiff_deflate\")\n    \n    # =========================================================\n    # ENSEMBLE LOGIC (The Switch)\n    # =========================================================\n    \n    def make_25d_ensemble_fn(ensemble_vars, use_logit_avg: bool):\n        model = Pix2PixGenerator()\n    \n        def _ensemble_forward(x):\n            # 1. Get raw logits from all models\n            # Shape: (NumModels, Batch, H, W, 1)\n            all_logits = jnp.stack([model.apply(v, x, train=False) for v in ensemble_vars], axis=0)\n            \n            if use_logit_avg:\n                # OPTION A: Average logits first, then sigmoid\n                # This results in sharper, more confident boundaries\n                avg_logits = jnp.mean(all_logits, axis=0)\n                return nn.sigmoid(avg_logits)\n            else:\n                # OPTION B: Sigmoid first, then average probabilities\n                # This is more conservative and reduces outlier influence\n                all_probs = nn.sigmoid(all_logits)\n                return jnp.mean(all_probs, axis=0)\n    \n        return jax.jit(_ensemble_forward)\n    \n    def make_postproc_2d_fn(post_vars):\n        model = PostProcessModel2D()\n        def _forward(x): return nn.sigmoid(model.apply(post_vars, x, train=False))\n        return jax.jit(_forward)\n    \n    # =========================================================\n    # MAIN PIPELINE\n    # =========================================================\n    def run_inference_25d():\n        mode_str = \"LOGIT\" if USE_LOGIT_AVERAGING else \"PROBABILITY\"\n        print(f\"🔧 2.5D CONFIG: Mode={mode_str} Averaging | ML_PostProc={USE_ML_POST_PROC_2D}\")\n    \n        ensemble_vars_25d = [load_flax_weights(p) for p in WEIGHT_PATHS_25D if load_flax_weights(p) is not None]\n        if not ensemble_vars_25d: return\n    \n        ensemble_25d_fn = make_25d_ensemble_fn(ensemble_vars_25d, USE_LOGIT_AVERAGING)\n    \n        postproc_2d_fn = None\n        if USE_ML_POST_PROC_2D and os.path.exists(POST_PROC_WEIGHTS_2D):\n            pv = load_flax_weights(POST_PROC_WEIGHTS_2D)\n            postproc_2d_fn = make_postproc_2d_fn(pv)\n    \n        test_dir = DATA_PATH / \"test_images\"\n        if not test_dir.exists(): return\n        test_ids = [f.stem for f in test_dir.glob(\"*.tif\")]\n        submission_files = []\n    \n        for vid in test_ids:\n            print(f\"\\nProcessing volume {vid}...\")\n            with Image.open(str(test_dir / f\"{vid}.tif\")) as img:\n                vol = np.stack([np.array(f) for f in ImageSequence.Iterator(img)], axis=0)\n    \n            z_max, h, w = vol.shape\n            vol_probs = np.zeros((z_max, h, w), dtype=np.float32)\n            coords = [(y, x) for y in range(0, h, PATCH_SIZE) for x in range(0, w, PATCH_SIZE)]\n    \n            for z in tqdm(range(Z_CONTEXT, z_max - Z_CONTEXT), desc=f\"2.5D {vid}\"):\n                patches_list, valid_coords = [], []\n                for y, x in coords:\n                    y1, x1 = min(y + PATCH_SIZE, h), min(x + PATCH_SIZE, w)\n                    y0, x0 = y1 - PATCH_SIZE, x1 - PATCH_SIZE\n                    p = vol[z - Z_CONTEXT : z + Z_CONTEXT + 1, y0:y1, x0:x1].transpose(1, 2, 0)\n                    patches_list.append(normalize_patch(p))\n                    valid_coords.append((y0, y1, x0, x1))\n    \n                for i in range(0, len(patches_list), BATCH_SIZE):\n                    batch_np = np.stack(patches_list[i : i + BATCH_SIZE])\n                    preds = ensemble_25d_fn(batch_np)\n    \n                    if postproc_2d_fn is not None:\n                        # ML Post-proc usually expects normalized 0-1 input\n                        p_min, p_max = jnp.min(preds), jnp.max(preds)\n                        preds = postproc_2d_fn((preds - p_min) / (p_max - p_min + 1e-6))\n    \n                    preds_flat = np.array(preds)[..., 0]\n                    for k in range(len(preds_flat)):\n                        y0, y1, x0, x1 = valid_coords[i + k]\n                        vol_probs[z, y0:y1, x0:x1] = preds_flat[k]\n    \n            if NORMALIZE_OUTPUT:\n                vol_probs = (vol_probs - vol_probs.min()) / (vol_probs.max() - vol_probs.min() + 1e-6)\n    \n            final_mask = apply_anisotropic_hysteresis(vol_probs) if USE_ANISOTROPIC_HYST else (vol_probs > T_LOW).astype(np.uint8) * 255\n            save_tif_stack(final_mask, f\"{vid}.tif\")\n            submission_files.append(f\"{vid}.tif\")\n            del vol, vol_probs; gc.collect()\n    \n        if submission_files:\n            with zipfile.ZipFile(\"submission.zip\", \"w\") as z:\n                for f in submission_files: z.write(f); os.remove(f)\n    \n    if __name__ == \"__main__\":\n        run_inference_25d()\n        visualize_submission()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-18T08:49:45.502995Z","iopub.execute_input":"2026-02-18T08:49:45.503493Z","iopub.status.idle":"2026-02-18T08:49:45.541937Z","shell.execute_reply.started":"2026-02-18T08:49:45.503463Z","shell.execute_reply":"2026-02-18T08:49:45.540682Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Logit and pred average","metadata":{}},{"cell_type":"markdown","source":"DUAL Inference: Logit AND Probability Intersection","metadata":{}},{"cell_type":"code","source":"#2.5D DUAL CONSENSUS PIPELINE\nimport os\nimport gc\nimport zipfile\nfrom pathlib import Path\nimport numpy as np\nimport jax\nimport jax.numpy as jnp\nimport flax.linen as nn\nfrom flax.traverse_util import unflatten_dict\nfrom PIL import Image, ImageSequence\nfrom tqdm import tqdm\nimport scipy.ndimage as ndi\nfrom skimage.morphology import remove_small_objects\n\n# =========================================================\n# CONFIG\n# =========================================================\nDATA_PATH = Path(\"/kaggle/input/vesuvius-challenge-surface-detection\")\nWEIGHT_PATHS_25D = [\n    #better without retro\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_30.npz\",\n    \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix/checkpoints/pix2pix_ep_20.npz\",\n    # worst only with retro\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_200.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_120.npz\",\n    # \"/kaggle/input/notebooks/crischir/2-5d-filter-training-pix2pix-pred-mask/checkpoints/pix2pix_ep_70.npz\",\n    #\n    \"/kaggle/input/notebooks/crischir/test-code-for-jax-pix2pix/checkpoints/pix2pix_ep_40.npz\",\n    \"/kaggle/input/notebooks/crischir/prewitt-loss-model-jax-pix2pix/checkpoints/pix2pix_best.npz\",\n]\n\nPOST_PROC_WEIGHTS_2D = \"/kaggle/input/datasets/crischir/postprocessvesuviuscomplete/post_proc_checkpoints/post_proc_ep30.npz\"\nPATCH_SIZE = 320\nZ_CONTEXT = 3\nBATCH_SIZE = 128\nUSE_ML_POST_PROC_2D = True\nUSE_ANISOTROPIC_HYST = True\nNORMALIZE_OUTPUT = True\n\n# Adjusted for the \"AND\" (multiplication) effect\nT_LOW = 0.023  # Slightly lowered because (prob * logit_prob) is smaller\nT_HIGH = 0.85  # Slightly lowered to ensure we still get \"seeds\" for hysteresis\nZ_RADIUS = 5\nXY_RADIUS = 2\nDUST_MIN_SIZE = 3\n\n# =========================================================\n# MODELS (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        s1 = nn.leaky_relu(nn.Conv(64, (4, 4), strides=(2, 2), padding='SAME')(x), 0.2)\n        s2 = conv_bn_leaky(s1, 128); s3 = conv_bn_leaky(s2, 256)\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        u1 = up_bn_relu(s3, s2, 128); u2 = up_bn_relu(u1, s1, 64)\n        return nn.ConvTranspose(1, (4, 4), strides=(2, 2), padding='SAME')(u2)\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(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 = nn.relu(nn.BatchNorm(use_running_average=not train)(nn.ConvTranspose(128, (4, 4), strides=(2, 2), padding=\"SAME\")(s3)))\n        u1 = jnp.concatenate([u1, s2], axis=-1)\n        u2 = nn.relu(nn.BatchNorm(use_running_average=not train)(nn.ConvTranspose(64, (4, 4), strides=(2, 2), padding=\"SAME\")(u1)))\n        u2 = jnp.concatenate([u2, s1], axis=-1)\n        return nn.ConvTranspose(1, (4, 4), strides=(2, 2), padding=\"SAME\")(u2)\n\n# =========================================================\n# UTILITIES\n# =========================================================\ndef load_flax_weights(path):\n    if not os.path.exists(path): 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 {\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\ndef normalize_patch(patch):\n    return (patch.astype(np.float32) - patch.mean()) / (patch.std() + 1e-6)\n\ndef apply_anisotropic_hysteresis(vol_probs):\n    print(f\"    ...Running 3D Hysteresis (High={T_HIGH}, Low={T_LOW})...\")\n    strong = vol_probs >= T_HIGH\n    weak = vol_probs >= T_LOW\n    if not strong.any(): return np.zeros_like(vol_probs, 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 DUST_MIN_SIZE > 0:\n        mask = remove_small_objects(mask.astype(bool), min_size=DUST_MIN_SIZE)\n    return mask.astype(np.uint8) * 255\n\ndef save_tif_stack(vol, out_name):\n    pages = [Image.fromarray(vol[i]) for i in range(vol.shape[0])]\n    pages[0].save(out_name, save_all=True, append_images=pages[1:], compression=\"tiff_deflate\")\n\n# =========================================================\n# ENSEMBLE LOGIC\n# =========================================================\ndef make_dual_ensemble_fn(ensemble_vars):\n    model = Pix2PixGenerator()\n    def _forward(x):\n        all_logits = jnp.stack([model.apply(v, x, train=False) for v in ensemble_vars], axis=0)\n        # Logit Averaging\n        logit_avg_probs = nn.sigmoid(jnp.mean(all_logits, axis=0))\n        # Prob Averaging\n        prob_avg_probs = jnp.mean(nn.sigmoid(all_logits), axis=0)\n        # Intersection (Soft AND)\n        return logit_avg_probs * prob_avg_probs\n    return jax.jit(_forward)\n\ndef make_postproc_2d_fn(post_vars):\n    model = PostProcessModel2D()\n    def _forward(x): return nn.sigmoid(model.apply(post_vars, x, train=False))\n    return jax.jit(_forward)\n\n# =========================================================\n# MAIN PIPELINE\n# =========================================================\ndef run_inference_25d():\n    print(\"🔧 2.5D DUAL CONSENSUS PIPELINE\")\n    \n    weights = [load_flax_weights(p) for p in WEIGHT_PATHS_25D if os.path.exists(p)]\n    if not weights: return\n    \n    ensemble_fn = make_dual_ensemble_fn(weights)\n    \n    postproc_fn = None\n    if USE_ML_POST_PROC_2D and os.path.exists(POST_PROC_WEIGHTS_2D):\n        postproc_fn = make_postproc_2d_fn(load_flax_weights(POST_PROC_WEIGHTS_2D))\n\n    test_dir = DATA_PATH / \"test_images\"\n    test_ids = [f.stem for f in test_dir.glob(\"*.tif\")]\n    submission_files = []\n\n    for vid in test_ids:\n        print(f\"\\nProcessing volume {vid}...\")\n        with Image.open(str(test_dir / f\"{vid}.tif\")) as img:\n            vol = np.stack([np.array(f) for f in ImageSequence.Iterator(img)], axis=0)\n\n        z_max, h, w = vol.shape\n        vol_probs = np.zeros((z_max, h, w), dtype=np.float32)\n        coords = [(y, x) for y in range(0, h, PATCH_SIZE) for x in range(0, w, PATCH_SIZE)]\n\n        for z in tqdm(range(Z_CONTEXT, z_max - Z_CONTEXT), desc=f\"Dual Consensus {vid}\"):\n            patches, valid_coords = [], []\n            for y, x in coords:\n                y1, x1 = min(y + PATCH_SIZE, h), min(x + PATCH_SIZE, w)\n                y0, x0 = y1 - PATCH_SIZE, x1 - PATCH_SIZE\n                p = vol[z - Z_CONTEXT : z + Z_CONTEXT + 1, y0:y1, x0:x1].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):\n                batch_np = np.stack(patches[i : i + BATCH_SIZE])\n                preds = ensemble_fn(batch_np)\n\n                if postproc_fn is not None:\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_flat = np.array(preds)[..., 0]\n                for k in range(len(preds_flat)):\n                    y0, y1, x0, x1 = valid_coords[i + k]\n                    vol_probs[z, y0:y1, x0:x1] = preds_flat[k]\n\n        if NORMALIZE_OUTPUT:\n            v_min, v_max = vol_probs.min(), vol_probs.max()\n            vol_probs = (vol_probs - v_min) / (v_max - v_min + 1e-6)\n\n        final_mask = apply_anisotropic_hysteresis(vol_probs) if USE_ANISOTROPIC_HYST else (vol_probs > T_LOW).astype(np.uint8) * 255\n        \n        out_name = f\"{vid}.tif\"\n        save_tif_stack(final_mask, out_name)\n        submission_files.append(out_name)\n        del vol, vol_probs; gc.collect()\n\n    if submission_files:\n        with zipfile.ZipFile(\"submission.zip\", \"w\") as z:\n            for f in submission_files: z.write(f); os.remove(f)\n        print(\"✅ submission.zip created\")\n\nif __name__ == \"__main__\":\n    run_inference_25d()\n    visualize_submission()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-18T09:27:27.322073Z","iopub.execute_input":"2026-02-18T09:27:27.322407Z","iopub.status.idle":"2026-02-18T09:31:25.092535Z","shell.execute_reply.started":"2026-02-18T09:27:27.322379Z","shell.execute_reply":"2026-02-18T09:31:25.091332Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"  visualize_submission()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-18T08:53:20.387985Z","iopub.execute_input":"2026-02-18T08:53:20.388442Z","iopub.status.idle":"2026-02-18T08:53:21.809590Z","shell.execute_reply.started":"2026-02-18T08:53:20.388411Z","shell.execute_reply":"2026-02-18T08:53:21.808269Z"}},"outputs":[],"execution_count":null}]}