{"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":[{"sourceType":"competition","sourceId":117682,"databundleVersionId":15062069},{"sourceType":"modelInstanceVersion","sourceId":773712,"databundleVersionId":15926923,"modelInstanceId":590869,"modelId":603164},{"sourceType":"modelInstanceVersion","sourceId":773709,"databundleVersionId":15926897,"modelInstanceId":590866,"modelId":603164},{"sourceType":"modelInstanceVersion","sourceId":773706,"databundleVersionId":15926873,"modelInstanceId":590863,"modelId":603164},{"sourceType":"modelInstanceVersion","sourceId":773701,"databundleVersionId":15926815,"modelInstanceId":590858,"modelId":603164},{"sourceType":"modelInstanceVersion","sourceId":773685,"databundleVersionId":15926598,"modelInstanceId":590851,"modelId":603164},{"sourceType":"modelInstanceVersion","sourceId":773697,"databundleVersionId":15926781,"modelInstanceId":590855,"modelId":603164},{"sourceType":"modelInstanceVersion","sourceId":773703,"databundleVersionId":15926846,"modelInstanceId":590860,"modelId":603164},{"sourceType":"modelInstanceVersion","sourceId":773704,"databundleVersionId":15926855,"modelInstanceId":590861,"modelId":603164},{"sourceType":"modelInstanceVersion","sourceId":773699,"databundleVersionId":15926800,"modelInstanceId":590856,"modelId":603164},{"sourceType":"modelInstanceVersion","sourceId":773696,"databundleVersionId":15926766,"modelInstanceId":590854,"modelId":603164},{"sourceType":"modelInstanceVersion","sourceId":773689,"databundleVersionId":15926626,"modelInstanceId":590853,"modelId":603164},{"sourceType":"modelInstanceVersion","sourceId":773657,"databundleVersionId":15926322,"modelInstanceId":590829,"modelId":603164}],"dockerImageVersionId":31286,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\n# for dirname, _, filenames in os.walk('/kaggle/input'):\n#     for filename in filenames:\n#         print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-03-04T21:23:46.568344Z","iopub.execute_input":"2026-03-04T21:23:46.568763Z","iopub.status.idle":"2026-03-04T21:23:46.574500Z","shell.execute_reply.started":"2026-03-04T21:23:46.568721Z","shell.execute_reply":"2026-03-04T21:23:46.573386Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from pathlib import Path\nfrom typing import Dict, List, Tuple, Set\n# ── Ensemble checkpoints ──────────────────────────────────────────────────────\nMODEL_CHECKPOINTS: Dict[str, List[Path]] = {\n    \"unu\": [\n        Path(\"/kaggle/input/models/crischir/2-5d-filter-vesuvius-surface-models-jax-collection/jax/cldice-tpu-training-pix2pix-v2/1/checkpoints/pix2pix_best.npz\"),\n        Path(\"/kaggle/input/models/crischir/2-5d-filter-vesuvius-surface-models-jax-collection/jax/cldice-tpu-training-pix2pix-v2/1/checkpoints/pix2pix_ep_30.npz\"),\n    ],\n    \"doi\": [\n        Path(\"/kaggle/input/models/crischir/2-5d-filter-vesuvius-surface-models-jax-collection/jax/tpu-training-70-epochs-v1/1/checkpoints/pix2pix_best.npz\"),\n        Path(\"/kaggle/input/models/crischir/2-5d-filter-vesuvius-surface-models-jax-collection/jax/tpu-training-70-epochs-v1/1/checkpoints/pix2pix_ep_30.npz\"),\n        Path(\"/kaggle/input/models/crischir/2-5d-filter-vesuvius-surface-models-jax-collection/jax/tpu-training-70-epochs-v1/1/checkpoints/pix2pix_ep_40.npz\"),\n        Path(\"/kaggle/input/models/crischir/2-5d-filter-vesuvius-surface-models-jax-collection/jax/tpu-training-70-epochs-v1/1/checkpoints/pix2pix_ep_40.npz\"),\n        Path(\"/kaggle/input/models/crischir/2-5d-filter-vesuvius-surface-models-jax-collection/jax/tpu-training-70-epochs-v1/1/checkpoints/pix2pix_ep_70.npz\"),\n    ],\n    \"trei\": [\n        Path(\"/kaggle/input/models/crischir/2-5d-filter-vesuvius-surface-models-jax-collection/jax/prewitt-loss-100-epochs/1/checkpoints/pix2pix_best.npz\"),\n        Path(\"/kaggle/input/models/crischir/2-5d-filter-vesuvius-surface-models-jax-collection/jax/prewitt-loss-100-epochs/1/checkpoints/pix2pix_ep_100.npz\"),\n    ],\n    \"patru\": [\n        Path(\"/kaggle/input/models/crischir/2-5d-filter-vesuvius-surface-models-jax-collection/jax/simple-model-no-tta/1/checkpoints/pix2pix_best.npz\"),\n        Path(\"/kaggle/input/models/crischir/2-5d-filter-vesuvius-surface-models-jax-collection/jax/simple-model-no-tta/1/checkpoints/pix2pix_best.npz\"),\n    ],\n    \"cinci\": [\n        Path(\"/kaggle/input/models/crischir/2-5d-filter-vesuvius-surface-models-jax-collection/jax/cldice-tpu-training-pix2pix-final-version/1/checkpoints/pix2pix_best.npz\"),\n        Path(\"/kaggle/input/models/crischir/2-5d-filter-vesuvius-surface-models-jax-collection/jax/cldice-tpu-training-pix2pix-final-version/1/checkpoints/pix2pix_ep_160.npz\"),\n        Path(\"/kaggle/input/models/crischir/2-5d-filter-vesuvius-surface-models-jax-collection/jax/cldice-tpu-training-pix2pix-final-version/1/checkpoints/pix2pix_ep_270.npz\"),\n    ],\n    \"sase\": [\n        Path(\"/kaggle/input/models/crischir/2-5d-filter-vesuvius-surface-models-jax-collection/jax/retro-masking-200-epochs-mask-preds-bigger-then07/1/checkpoints/pix2pix_best.npz\"),\n        Path(\"/kaggle/input/models/crischir/2-5d-filter-vesuvius-surface-models-jax-collection/jax/retro-masking-200-epochs-mask-preds-bigger-then07/1/checkpoints/pix2pix_ep_100.npz\"),\n        Path(\"/kaggle/input/models/crischir/2-5d-filter-vesuvius-surface-models-jax-collection/jax/retro-masking-200-epochs-mask-preds-bigger-then07/1/checkpoints/pix2pix_ep_100.npz\"),\n        Path(\"/kaggle/input/models/crischir/2-5d-filter-vesuvius-surface-models-jax-collection/jax/retro-masking-200-epochs-mask-preds-bigger-then07/1/checkpoints/pix2pix_ep_200.npz\"),\n        Path(\"/kaggle/input/models/crischir/2-5d-filter-vesuvius-surface-models-jax-collection/jax/retro-masking-200-epochs-mask-preds-bigger-then07/1/checkpoints/pix2pix_ep_200.npz\"),\n    ],\n    \"sapte\": [\n        Path(\"/kaggle/input/models/crischir/2-5d-filter-vesuvius-surface-models-jax-collection/jax/tta-training-30epoch-retrain/1/checkpoints/pix2pix_best.npz\"),\n        Path(\"/kaggle/input/models/crischir/2-5d-filter-vesuvius-surface-models-jax-collection/jax/tta-training-30epoch-retrain/1/checkpoints/pix2pix_ep_30.npz\"),\n    ],\n    \"opt\": [\n        Path(\"/kaggle/input/models/crischir/2-5d-filter-vesuvius-surface-models-jax-collection/jax/50-more-epochs-400simplemodel/1/checkpoints/pix2pix_best.npz\"),\n    ],\n    \"noua\": [\n        Path(\"/kaggle/input/models/crischir/2-5d-filter-vesuvius-surface-models-jax-collection/jax/cldice-tpu-training-pix2pix-v1/1/checkpoints/pix2pix_best.npz\"),\n        Path(\"/kaggle/input/models/crischir/2-5d-filter-vesuvius-surface-models-jax-collection/jax/cldice-tpu-training-pix2pix-v1/1/checkpoints/pix2pix_ep_30.npz\"),\n    ],\n    \"unsprezece\": [\n        Path(\"/kaggle/input/models/crischir/2-5d-filter-vesuvius-surface-models-jax-collection/jax/sobel-loss-170-epochs/1/checkpoints/pix2pix_best.npz\"),\n        Path(\"/kaggle/input/models/crischir/2-5d-filter-vesuvius-surface-models-jax-collection/jax/sobel-loss-170-epochs/1/checkpoints/pix2pix_ep_130.npz\"),\n    ],\n    \"doisprezece\": [\n        Path(\"/kaggle/input/models/crischir/2-5d-filter-vesuvius-surface-models-jax-collection/jax/300-epochs-imag2image/1/checkpoints/pix2pix_best.npz\"),\n        Path(\"/kaggle/input/models/crischir/2-5d-filter-vesuvius-surface-models-jax-collection/jax/300-epochs-imag2image/1/checkpoints/pix2pix_ep_260.npz\"),\n    ],\n    \"zero\": [\n        Path(\"/kaggle/input/models/crischir/2-5d-filter-vesuvius-surface-models-jax-collection/jax/tpu-training-100epochs-v2/1/checkpoints/pix2pix_best.npz\"),\n        Path(\"/kaggle/input/models/crischir/2-5d-filter-vesuvius-surface-models-jax-collection/jax/tpu-training-100epochs-v2/1/checkpoints/pix2pix_ep_90.npz\"),\n    ],\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-04T21:23:46.576436Z","iopub.execute_input":"2026-03-04T21:23:46.576790Z","iopub.status.idle":"2026-03-04T21:23:46.602609Z","shell.execute_reply.started":"2026-03-04T21:23:46.576764Z","shell.execute_reply":"2026-03-04T21:23:46.601421Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"TEST_IMG_DIR=\"/kaggle/input/competitions/vesuvius-challenge-surface-detection/test_images\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-04T21:23:46.604181Z","iopub.execute_input":"2026-03-04T21:23:46.604593Z","iopub.status.idle":"2026-03-04T21:23:46.624086Z","shell.execute_reply.started":"2026-03-04T21:23:46.604563Z","shell.execute_reply":"2026-03-04T21:23:46.623046Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport jax\nimport jax.numpy as jnp\nimport flax.linen as nn\nimport matplotlib.pyplot as plt\nimport numpy as np\n\nfrom typing import Dict, List, Tuple, Set\nfrom PIL import Image, ImageSequence\nfrom flax.traverse_util import unflatten_dict\n\n# --- 1. SETTINGS & PATHS ---\nos.environ[\"XLA_PYTHON_CLIENT_MEM_FRACTION\"] = \"0.95\"\nN_DEVICES = jax.device_count()\nZ_CONTEXT = 3  # 7-slice sandwich (3 up + 1 mid + 3 down)\nTIF_PATH = Path(\"/kaggle/input/competitions/vesuvius-challenge-surface-detection/test_images/1407735.tif\")\nSAVE_DIR = Path(\"/kaggle/working/showcase_results\")\nSAVE_DIR.mkdir(exist_ok=True)\n_DARK = \"#0e0e0e\"\nSHOWCASE_THRESHOLDS = [0.05, 0.1, 0.2, 0.3, 0.4]\n\n# --- 2. MODEL ARCHITECTURE ---\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        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        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        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_model_singleton = Pix2PixGenerator()\n\n# --- 3. JAX DISPATCH & TTA ---\ndef _pmap_fn(variables, x):\n    logit = _model_singleton.apply(variables, x[None], train=False)\n    return jax.nn.sigmoid(logit)[0, :, :, 0]\n\n_pmap_forward = jax.pmap(_pmap_fn, in_axes=(0, 0))\n\n_AUG_FWDS = [\n    lambda x: x,\n    lambda x: jnp.flip(x, axis=0),\n    lambda x: jnp.flip(x, axis=1),\n    lambda x: jnp.flip(jnp.flip(x, axis=0), axis=1),\n    lambda x: jnp.rot90(x, k=1, axes=(0, 1)),\n    lambda x: jnp.rot90(x, k=2, axes=(0, 1)),\n    lambda x: jnp.rot90(x, k=3, axes=(0, 1)),\n    lambda x: jnp.flip(jnp.rot90(x, k=1, axes=(0, 1)), axis=0),\n]\n_AUG_INVS = [\n    lambda p: p,\n    lambda p: np.flip(p, axis=0),\n    lambda p: np.flip(p, axis=1),\n    lambda p: np.flip(np.flip(p, axis=0), axis=1),\n    lambda p: np.rot90(p, k=-1),\n    lambda p: np.rot90(p, k=-2),\n    lambda p: np.rot90(p, k=-3),\n    lambda p: np.rot90(np.flip(p, axis=0), k=-1),\n]\n\ndef _tta_predict_pmap(variables_rep, stack_hwc):\n    x = jnp.asarray(stack_hwc)\n    aug_inputs = [fn(x) for fn in _AUG_FWDS]\n    padded_len = ((8 + N_DEVICES - 1) // N_DEVICES) * N_DEVICES\n    aug_padded = aug_inputs + [aug_inputs[0]] * (padded_len - 8)\n    preds = []\n    for i in range(0, padded_len, N_DEVICES):\n        shard = jnp.stack(aug_padded[i: i + N_DEVICES], axis=0)\n        out = _pmap_forward(variables_rep, shard)\n        out_np = np.asarray(out)\n        for j in range(N_DEVICES):\n            idx = i + j\n            if idx < 8: preds.append(_AUG_INVS[idx](out_np[j]))\n    return np.mean(preds, axis=0)\n\n# --- 4. WEIGHT LOADING (PREFIX AWARE) ---\ndef load_npz_params_local(path: Path) -> dict:\n    with np.load(str(path), allow_pickle=False) as f:\n        flat = dict(f)\n    p_flat = {k.replace(\"params/\", \"\"): v for k, v in flat.items() if k.startswith(\"params/\")}\n    s_flat = {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_flat.items()}),\n        \"batch_stats\": unflatten_dict({tuple(k.split(\"/\")): v for k, v in s_flat.items()}),\n    }\n\ndef load_ensemble_local(checkpoint_groups: Dict[str, List[Path]]):\n    loaded = []\n    seen = set()\n    for group, paths in checkpoint_groups.items():\n        for p in paths:\n            rp = p.resolve()\n            if rp in seen or not p.exists(): continue\n            seen.add(rp)\n            v = load_npz_params_local(p)\n            v_rep = jax.device_put_replicated(v, jax.devices())\n            loaded.append((f\"{group}/{p.stem}\", v_rep))\n    return loaded\n\n# --- 5. VISUALIZATION HELPERS ---\ndef _to_u8(arr):\n    mn, mx = arr.min(), arr.max()\n    if mx > mn: return ((arr - mn) / (mx - mn) * 255).astype(np.uint8)\n    return np.zeros_like(arr, dtype=np.uint8)\n\ndef read_volume_local(path: Path):\n    with Image.open(str(path)) as img:\n        frames = [np.array(f.copy()) for f in ImageSequence.Iterator(img)]\n    return np.stack(frames, axis=0)\n\n# --- 6. MAIN SHOWCASE ENGINE ---\ndef run_full_showcase(checkpoint_groups: Dict[str, List[Path]]):\n    print(f\"Initializing Showcase on {N_DEVICES} GPUs...\")\n    vol = read_volume_local(TIF_PATH)\n    z_max = vol.shape[0]\n    indices = np.linspace(0, z_max - 1, 5, dtype=int)\n    padded = np.pad(vol.astype(np.float32), ((Z_CONTEXT, Z_CONTEXT), (0, 0), (0, 0)), mode=\"reflect\")\n\n    all_models = load_ensemble_local(checkpoint_groups)\n    weight_map = {lbl: w for lbl, w in all_models}\n\n    def get_prob(w_rep, z):\n        stack = padded[z: z + 2 * Z_CONTEXT + 1]\n        stack = (stack - stack.mean()) / (stack.std() + 1e-6)\n        return _tta_predict_pmap(w_rep, stack.transpose(1, 2, 0))\n\n    for key, paths in checkpoint_groups.items():\n        # Handle Single\n        first_lbl = f\"{key}/{paths[0].stem}\"\n        if first_lbl in weight_map:\n            print(f\"Processing Group: {key}\")\n            probs = [get_prob(weight_map[first_lbl], z) for z in indices]\n            \n            fig, axes = plt.subplots(5, 6, figsize=(20, 15), facecolor=_DARK)\n            for i, z_idx in enumerate(indices):\n                axes[i, 0].imshow(_to_u8(vol[z_idx]), cmap='gray')\n                axes[i, 0].set_title(f\"Z={z_idx}\", color='white', fontsize=9)\n                axes[i, 0].axis('off')\n                for j, thr in enumerate(SHOWCASE_THRESHOLDS):\n                    bin_img = (probs[i] > thr)\n                    axes[i, j+1].imshow(bin_img, cmap='viridis')\n                    axes[i, j+1].set_title(f\"T={thr}\\n{bin_img.sum():,}px\", color='#00ff88', fontsize=8)\n                    axes[i, j+1].axis('off')\n            \n            plt.suptitle(f\"SHOWCASE: {key} (Single Model)\", color='white', fontsize=18)\n            plt.tight_layout()\n            plt.savefig(SAVE_DIR / f\"{key.replace('/','_')}_showcase.png\", facecolor=_DARK)\n            plt.show()\n\n# --- EXECUTION ---\n# Ensure MODEL_CHECKPOINTS is defined before running\nrun_full_showcase(MODEL_CHECKPOINTS)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-04T21:23:46.665190Z","iopub.execute_input":"2026-03-04T21:23:46.665870Z","iopub.status.idle":"2026-03-04T21:25:26.743450Z","shell.execute_reply.started":"2026-03-04T21:23:46.665828Z","shell.execute_reply":"2026-03-04T21:25:26.742426Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport jax\nimport jax.numpy as jnp\nimport flax.linen as nn\nimport matplotlib.pyplot as plt\nimport numpy as np\nfrom pathlib import Path\nfrom typing import Dict, List, Tuple, Set\nfrom PIL import Image, ImageSequence\nfrom flax.traverse_util import unflatten_dict\n\n# --- 1. SETTINGS & PATHS ---\nos.environ[\"XLA_PYTHON_CLIENT_MEM_FRACTION\"] = \"0.95\"\nN_DEVICES = jax.device_count()\nZ_CONTEXT = 3  # 7-slice sandwich (3 up + 1 mid + 3 down)\nTIF_PATH = Path(\"/kaggle/input/competitions/vesuvius-challenge-surface-detection/test_images/1407735.tif\")\nSAVE_DIR = Path(\"/kaggle/working/ensemble_showcase\")\nSAVE_DIR.mkdir(exist_ok=True)\n_DARK = \"#0e0e0e\"\nSHOWCASE_THRESHOLDS = [0.05, 0.1, 0.2, 0.3, 0.4]\n\n# --- 2. MODEL ARCHITECTURE ---\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        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        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        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_model_singleton = Pix2PixGenerator()\n\n# --- 3. JAX DISPATCH & TTA ---\ndef _pmap_fn(variables, x):\n    logit = _model_singleton.apply(variables, x[None], train=False)\n    return jax.nn.sigmoid(logit)[0, :, :, 0]\n\n_pmap_forward = jax.pmap(_pmap_fn, in_axes=(0, 0))\n\n_AUG_FWDS = [\n    lambda x: x,\n    lambda x: jnp.flip(x, axis=0),\n    lambda x: jnp.flip(x, axis=1),\n    lambda x: jnp.flip(jnp.flip(x, axis=0), axis=1),\n    lambda x: jnp.rot90(x, k=1, axes=(0, 1)),\n    lambda x: jnp.rot90(x, k=2, axes=(0, 1)),\n    lambda x: jnp.rot90(x, k=3, axes=(0, 1)),\n    lambda x: jnp.flip(jnp.rot90(x, k=1, axes=(0, 1)), axis=0),\n]\n_AUG_INVS = [\n    lambda p: p,\n    lambda p: np.flip(p, axis=0),\n    lambda p: np.flip(p, axis=1),\n    lambda p: np.flip(np.flip(p, axis=0), axis=1),\n    lambda p: np.rot90(p, k=-1),\n    lambda p: np.rot90(p, k=-2),\n    lambda p: np.rot90(p, k=-3),\n    lambda p: np.rot90(np.flip(p, axis=0), k=-1),\n]\n\ndef _tta_predict_pmap(variables_rep, stack_hwc):\n    x = jnp.asarray(stack_hwc)\n    aug_inputs = [fn(x) for fn in _AUG_FWDS]\n    padded_len = ((8 + N_DEVICES - 1) // N_DEVICES) * N_DEVICES\n    aug_padded = aug_inputs + [aug_inputs[0]] * (padded_len - 8)\n    preds = []\n    for i in range(0, padded_len, N_DEVICES):\n        shard = jnp.stack(aug_padded[i: i + N_DEVICES], axis=0)\n        out = _pmap_forward(variables_rep, shard)\n        out_np = np.asarray(out)\n        for j in range(N_DEVICES):\n            idx = i + j\n            if idx < 8: preds.append(_AUG_INVS[idx](out_np[j]))\n    return np.mean(preds, axis=0)\n\n# --- 4. WEIGHT LOADING ---\ndef load_npz_params_local(path: Path) -> dict:\n    with np.load(str(path), allow_pickle=False) as f:\n        flat = dict(f)\n    p_flat = {k.replace(\"params/\", \"\"): v for k, v in flat.items() if k.startswith(\"params/\")}\n    s_flat = {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_flat.items()}),\n        \"batch_stats\": unflatten_dict({tuple(k.split(\"/\")): v for k, v in s_flat.items()}),\n    }\n\ndef load_ensemble_local(checkpoint_groups: Dict[str, List[Path]]):\n    loaded = []\n    seen = set()\n    for group, paths in checkpoint_groups.items():\n        for p in paths:\n            rp = p.resolve()\n            if rp in seen or not p.exists(): continue\n            seen.add(rp)\n            v = load_npz_params_local(p)\n            v_rep = jax.device_put_replicated(v, jax.devices())\n            loaded.append((f\"{group}/{p.stem}\", v_rep))\n    return loaded\n\n# --- 5. VISUALIZATION HELPERS ---\ndef _to_u8(arr):\n    mn, mx = arr.min(), arr.max()\n    return ((arr - mn) / (mx - mn + 1e-8) * 255).astype(np.uint8) if mx > mn else np.zeros_like(arr, dtype=np.uint8)\n\ndef read_volume_local(path: Path):\n    with Image.open(str(path)) as img:\n        frames = [np.array(f.copy()) for f in ImageSequence.Iterator(img)]\n    return np.stack(frames, axis=0)\n\n# --- 6. ENSEMBLE SHOWCASE ENGINE ---\ndef run_ensemble_showcase(checkpoint_groups: Dict[str, List[Path]]):\n    print(f\"Reading Volume: {TIF_PATH.name} (320 pages expected)\")\n    vol = read_volume_local(TIF_PATH)\n    z_max = vol.shape[0]\n    indices = np.linspace(0, z_max - 1, 5, dtype=int)\n    padded = np.pad(vol.astype(np.float32), ((Z_CONTEXT, Z_CONTEXT), (0, 0), (0, 0)), mode=\"reflect\")\n\n    # Load all models once to avoid repeated disk I/O\n    all_models = load_ensemble_local(checkpoint_groups)\n    weight_map = {lbl: w for lbl, w in all_models}\n\n    def get_prob(w_rep, z):\n        stack = padded[z: z + 2 * Z_CONTEXT + 1]\n        stack = (stack - stack.mean()) / (stack.std() + 1e-6)\n        return _tta_predict_pmap(w_rep, stack.transpose(1, 2, 0))\n\n    for key, paths in checkpoint_groups.items():\n        # Identify all models belonging to this group\n        group_labels = [f\"{key}/{p.stem}\" for p in paths if f\"{key}/{p.stem}\" in weight_map]\n        if not group_labels: continue\n        \n        print(f\"\\n--- Processing Ensemble Group: {key} ({len(group_labels)} models) ---\")\n        \n        # Calculate ensemble probability for each of the 5 showcase slices\n        ensemble_probs = []\n        for z in indices:\n            # Average the TTA predictions of all models in the group\n            member_preds = [get_prob(weight_map[lbl], z) for lbl in group_labels]\n            ensemble_probs.append(np.mean(member_preds, axis=0))\n            \n        # Plot the 5x6 Grid (CT Slice | T=0.05 | T=0.1 | T=0.2 | T=0.4 | T=0.5)\n        fig, axes = plt.subplots(5, 6, figsize=(22, 16), facecolor=_DARK)\n        for i, z_idx in enumerate(indices):\n            # Col 0: Original CT\n            axes[i, 0].imshow(_to_u8(vol[z_idx]), cmap='gray')\n            axes[i, 0].set_title(f\"Slice Z={z_idx}\", color='white', fontsize=10)\n            axes[i, 0].axis('off')\n            \n            p_map = ensemble_probs[i]\n            for j, thr in enumerate(SHOWCASE_THRESHOLDS):\n                binary = (p_map > thr)\n                px_count = np.sum(binary)\n                \n                # Plot Binary Prediction\n                axes[i, j+1].imshow(binary, cmap='magma')\n                axes[i, j+1].set_title(f\"Ens Thr: {thr}\\n{px_count:,} px\", color='#00ff88', fontsize=9)\n                axes[i, j+1].axis('off')\n\n        plt.suptitle(f\"ENSEMBLE SHOWCASE: {key}\\n{len(group_labels)} models averaged | 8-fold TTA\", \n                     color='white', fontsize=20, y=0.98)\n        plt.tight_layout(rect=[0, 0.03, 1, 0.95])\n        \n        save_path = SAVE_DIR / f\"ensemble_{key.replace('/', '_')}.png\"\n        plt.savefig(save_path, facecolor=_DARK)\n        print(f\"  [SAVED] {save_path}\")\n        plt.show()\n\n# --- EXECUTION ---\nrun_ensemble_showcase(MODEL_CHECKPOINTS)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-04T21:25:26.745462Z","iopub.execute_input":"2026-03-04T21:25:26.745794Z","iopub.status.idle":"2026-03-04T21:28:16.239472Z","shell.execute_reply.started":"2026-03-04T21:25:26.745768Z","shell.execute_reply":"2026-03-04T21:28:16.238013Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport jax\nimport jax.numpy as jnp\nimport flax.linen as nn\nimport matplotlib.pyplot as plt\nimport numpy as np\nfrom pathlib import Path\nfrom typing import Dict, List, Tuple, Set\nfrom PIL import Image, ImageSequence\nfrom flax.traverse_util import unflatten_dict\n\n# --- 1. SETTINGS & PATHS ---\nos.environ[\"XLA_PYTHON_CLIENT_MEM_FRACTION\"] = \"0.95\"\nN_DEVICES = jax.device_count()\nZ_CONTEXT = 3  # 7-slice sandwich (3 up + 1 mid + 3 down)\nTIF_PATH = Path(\"/kaggle/input/competitions/vesuvius-challenge-surface-detection/test_images/1407735.tif\")\nSAVE_DIR = Path(\"/kaggle/working/showcase_overlay\")\nSAVE_DIR.mkdir(exist_ok=True)\n_DARK = \"#0e0e0e\"\nSHOWCASE_THRESHOLDS = [0.025, 0.15, 0.25, 0.4, 0.4]\nOVERLAY_ALPHA = 0.5  # Transparency of the green mask\n\n# --- 2. MODEL ARCHITECTURE ---\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        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        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        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_model_singleton = Pix2PixGenerator()\n\n# --- 3. JAX DISPATCH & TTA ---\ndef _pmap_fn(variables, x):\n    logit = _model_singleton.apply(variables, x[None], train=False)\n    return jax.nn.sigmoid(logit)[0, :, :, 0]\n\n_pmap_forward = jax.pmap(_pmap_fn, in_axes=(0, 0))\n\n_AUG_FWDS = [\n    lambda x: x,\n    lambda x: jnp.flip(x, axis=0),\n    lambda x: jnp.flip(x, axis=1),\n    lambda x: jnp.flip(jnp.flip(x, axis=0), axis=1),\n    lambda x: jnp.rot90(x, k=1, axes=(0, 1)),\n    lambda x: jnp.rot90(x, k=2, axes=(0, 1)),\n    lambda x: jnp.rot90(x, k=3, axes=(0, 1)),\n    lambda x: jnp.flip(jnp.rot90(x, k=1, axes=(0, 1)), axis=0),\n]\n_AUG_INVS = [\n    lambda p: p,\n    lambda p: np.flip(p, axis=0),\n    lambda p: np.flip(p, axis=1),\n    lambda p: np.flip(np.flip(p, axis=0), axis=1),\n    lambda p: np.rot90(p, k=-1),\n    lambda p: np.rot90(p, k=-2),\n    lambda p: np.rot90(p, k=-3),\n    lambda p: np.rot90(np.flip(p, axis=0), k=-1),\n]\n\ndef _tta_predict_pmap(variables_rep, stack_hwc):\n    x = jnp.asarray(stack_hwc)\n    aug_inputs = [fn(x) for fn in _AUG_FWDS]\n    padded_len = ((8 + N_DEVICES - 1) // N_DEVICES) * N_DEVICES\n    aug_padded = aug_inputs + [aug_inputs[0]] * (padded_len - 8)\n    preds = []\n    for i in range(0, padded_len, N_DEVICES):\n        shard = jnp.stack(aug_padded[i: i + N_DEVICES], axis=0)\n        out = _pmap_forward(variables_rep, shard)\n        out_np = np.asarray(out)\n        for j in range(N_DEVICES):\n            idx = i + j\n            if idx < 8: preds.append(_AUG_INVS[idx](out_np[j]))\n    return np.mean(preds, axis=0)\n\n# --- 4. WEIGHT LOADING ---\ndef load_npz_params_local(path: Path) -> dict:\n    with np.load(str(path), allow_pickle=False) as f:\n        flat = dict(f)\n    p_flat = {k.replace(\"params/\", \"\"): v for k, v in flat.items() if k.startswith(\"params/\")}\n    s_flat = {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_flat.items()}),\n        \"batch_stats\": unflatten_dict({tuple(k.split(\"/\")): v for k, v in s_flat.items()}),\n    }\n\ndef load_ensemble_local(checkpoint_groups: Dict[str, List[Path]]):\n    loaded = []\n    seen = set()\n    for group, paths in checkpoint_groups.items():\n        for p in paths:\n            rp = p.resolve()\n            if rp in seen or not p.exists(): continue\n            seen.add(rp)\n            v = load_npz_params_local(p)\n            v_rep = jax.device_put_replicated(v, jax.devices())\n            loaded.append((f\"{group}/{p.stem}\", v_rep))\n    return loaded\n\n# --- 5. VISUALIZATION HELPERS (UPDATED FOR OVERLAY) ---\ndef _to_u8(arr):\n    mn, mx = arr.min(), arr.max()\n    mnmax = max(mx - mn, 1e-8)\n    return ((arr - mn) / mnmax * 255).astype(np.uint8) if mx > mn else np.zeros_like(arr, dtype=np.uint8)\n\ndef _get_overlay_image(gray_img, bin_mask, alpha=OVERLAY_ALPHA):\n    \"\"\"\n    Blends a grayscale image with a green binary mask.\n    gray_img: np.ndarray (H, W) \n    bin_mask: np.ndarray (H, W) Boolean or 0/1\n    \"\"\"\n    g_u8 = _to_u8(gray_img)\n    # Create RGB representation of grayscale\n    rgb = np.stack([g_u8, g_u8, g_u8], axis=-1).astype(np.float32)\n    \n    # Define active foreground\n    fg = bin_mask > 0\n    \n    # Shift colors in foreground: Reduce Red/Blue, Keep Green high\n    rgb[fg, 0] *= (1 - alpha)\n    # rgb[fg, 1] keeps original green intensity mixed with alpha\n    rgb[fg, 1] = (rgb[fg, 1] * (1 - alpha)) + (255 * alpha)\n    rgb[fg, 2] *= (1 - alpha)\n    \n    return rgb.clip(0, 255).astype(np.uint8)\n\ndef read_volume_local(path: Path):\n    with Image.open(str(path)) as img:\n        frames = [np.array(f.copy()) for f in ImageSequence.Iterator(img)]\n    return np.stack(frames, axis=0)\n\n# --- 6. SINGLE MODEL OVERLAY ENGINE ---\ndef run_single_model_overlay(checkpoint_groups: Dict[str, List[Path]]):\n    print(f\"Reading Volume: {TIF_PATH.name} (320 pages expected)\")\n    vol = read_volume_local(TIF_PATH)\n    z_max = vol.shape[0]\n    indices = np.linspace(0, z_max - 1, 5, dtype=int)\n    padded = np.pad(vol.astype(np.float32), ((Z_CONTEXT, Z_CONTEXT), (0, 0), (0, 0)), mode=\"reflect\")\n\n    # Load only the first unique model found to save memory\n    all_models = load_ensemble_local(checkpoint_groups)\n    if not all_models:\n        print(\"Error: No models loaded.\")\n        return\n    label, weights = all_models[0]\n    print(f\"Processing Single Model for Overlay: {label}\")\n\n    def get_prob(w_rep, z):\n        stack = padded[z: z + 2 * Z_CONTEXT + 1]\n        stack = (stack - stack.mean()) / (stack.std() + 1e-6)\n        return _tta_predict_pmap(w_rep, stack.transpose(1, 2, 0))\n\n    # Run predictions for the 5 slices\n    probs = [get_prob(weights, z) for z in indices]\n            \n    # Plot the 5x6 Grid (CT Slice | T=0.05 MIP | T=0.1 MIP ... )\n    fig, axes = plt.subplots(5, 6, figsize=(22, 16), facecolor=_DARK)\n    \n    for i, z_idx in enumerate(indices):\n        original_u8 = _to_u8(vol[z_idx])\n        # Col 0: Original CT (Axial view)\n        axes[i, 0].imshow(original_u8, cmap='gray')\n        axes[i, 0].set_title(f\"Slice Z={z_idx}\", color='white', fontsize=10)\n        axes[i, 0].axis('off')\n        \n        p_map = probs[i]\n        for j, thr in enumerate(SHOWCASE_THRESHOLDS):\n            binary = (p_map > thr)\n            px_count = np.sum(binary)\n            \n            # Create Green Overlay\n            overlay_img = _get_overlay_image(vol[z_idx], binary)\n            \n            # Plot Overlay\n            axes[i, j+1].imshow(overlay_img)\n            axes[i, j+1].set_title(f\"Thr: {thr}\\n{px_count:,} px\", color='#00ff88', fontsize=9)\n            axes[i, j+1].axis('off')\n\n    plt.suptitle(f\"SINGLE MODEL OVERLAY: {label}\\n7-Slice Sandwich | 8-fold TTA (Alpha={OVERLAY_ALPHA})\", \n                 color='white', fontsize=20, y=0.98)\n    plt.tight_layout(rect=[0, 0.03, 1, 0.95])\n    \n    # Clean filename for saving\n    clean_label = label.replace(\"/\", \"_\").replace(\" \", \"_\")\n    save_path = SAVE_DIR / f\"{clean_label}_overlay.png\"\n    plt.savefig(save_path, facecolor=_DARK)\n    print(f\"  [SAVED] {save_path}\")\n    plt.show()\n\n# --- EXECUTION ---\nrun_single_model_overlay(MODEL_CHECKPOINTS)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-04T21:32:05.644690Z","iopub.execute_input":"2026-03-04T21:32:05.645730Z","iopub.status.idle":"2026-03-04T21:32:16.540845Z","shell.execute_reply.started":"2026-03-04T21:32:05.645695Z","shell.execute_reply":"2026-03-04T21:32:16.539500Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport jax\nimport jax.numpy as jnp\nimport flax.linen as nn\nimport matplotlib.pyplot as plt\nimport numpy as np\nfrom pathlib import Path\nfrom typing import Dict, List, Tuple\nfrom PIL import Image, ImageSequence\nfrom flax.traverse_util import unflatten_dict\n\n# --- 1. SETTINGS & PATHS ---\nos.environ[\"XLA_PYTHON_CLIENT_MEM_FRACTION\"] = \"0.90\" # Slightly lower to allow room for more models\nN_DEVICES = jax.device_count()\nZ_CONTEXT = 3  \nTIF_PATH = Path(\"/kaggle/input/competitions/vesuvius-challenge-surface-detection/test_images/1407735.tif\")\nSAVE_DIR = Path(\"/kaggle/working/all_models_overlay\")\nSAVE_DIR.mkdir(exist_ok=True)\n_DARK = \"#0e0e0e\"\nSHOWCASE_THRESHOLDS = [0.05, 0.1, 0.2, 0.4, 0.5]\nOVERLAY_ALPHA = 0.5 \n\n# --- 2. MODEL ARCHITECTURE ---\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        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        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        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_model_singleton = Pix2PixGenerator()\n\n# --- 3. JAX DISPATCH & TTA ---\ndef _pmap_fn(variables, x):\n    logit = _model_singleton.apply(variables, x[None], train=False)\n    return jax.nn.sigmoid(logit)[0, :, :, 0]\n\n_pmap_forward = jax.pmap(_pmap_fn, in_axes=(0, 0))\n\n_AUG_FWDS = [\n    lambda x: x,\n    lambda x: jnp.flip(x, axis=0),\n    lambda x: jnp.flip(x, axis=1),\n    lambda x: jnp.flip(jnp.flip(x, axis=0), axis=1),\n    lambda x: jnp.rot90(x, k=1, axes=(0, 1)),\n    lambda x: jnp.rot90(x, k=2, axes=(0, 1)),\n    lambda x: jnp.rot90(x, k=3, axes=(0, 1)),\n    lambda x: jnp.flip(jnp.rot90(x, k=1, axes=(0, 1)), axis=0),\n]\n_AUG_INVS = [\n    lambda p: p,\n    lambda p: np.flip(p, axis=0),\n    lambda p: np.flip(p, axis=1),\n    lambda p: np.flip(np.flip(p, axis=0), axis=1),\n    lambda p: np.rot90(p, k=-1),\n    lambda p: np.rot90(p, k=-2),\n    lambda p: np.rot90(p, k=-3),\n    lambda p: np.rot90(np.flip(p, axis=0), k=-1),\n]\n\ndef _tta_predict_pmap(variables_rep, stack_hwc):\n    x = jnp.asarray(stack_hwc)\n    aug_inputs = [fn(x) for fn in _AUG_FWDS]\n    padded_len = ((8 + N_DEVICES - 1) // N_DEVICES) * N_DEVICES\n    aug_padded = aug_inputs + [aug_inputs[0]] * (padded_len - 8)\n    preds = []\n    for i in range(0, padded_len, N_DEVICES):\n        shard = jnp.stack(aug_padded[i: i + N_DEVICES], axis=0)\n        out = _pmap_forward(variables_rep, shard)\n        out_np = np.asarray(out)\n        for j in range(N_DEVICES):\n            idx = i + j\n            if idx < 8: preds.append(_AUG_INVS[idx](out_np[j]))\n    return np.mean(preds, axis=0)\n\n# --- 4. PREFIX-AWARE LOADING ---\ndef load_npz_params_local(path: Path) -> dict:\n    with np.load(str(path), allow_pickle=False) as f:\n        flat = dict(f)\n    p_flat = {k.replace(\"params/\", \"\"): v for k, v in flat.items() if k.startswith(\"params/\")}\n    s_flat = {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_flat.items()}),\n        \"batch_stats\": unflatten_dict({tuple(k.split(\"/\")): v for k, v in s_flat.items()}),\n    }\n\n# --- 5. OVERLAY HELPERS ---\ndef _to_u8(arr):\n    mn, mx = arr.min(), arr.max()\n    return ((arr - mn) / (max(mx - mn, 1e-8)) * 255).astype(np.uint8) if mx > mn else np.zeros_like(arr, dtype=np.uint8)\n\ndef _get_overlay_image(gray_img, bin_mask, alpha=OVERLAY_ALPHA):\n    g_u8 = _to_u8(gray_img)\n    rgb = np.stack([g_u8, g_u8, g_u8], axis=-1).astype(np.float32)\n    fg = bin_mask > 0\n    rgb[fg, 0] *= (1 - alpha)\n    rgb[fg, 1] = (rgb[fg, 1] * (1 - alpha)) + (255 * alpha)\n    rgb[fg, 2] *= (1 - alpha)\n    return rgb.clip(0, 255).astype(np.uint8)\n\ndef read_volume_local(path: Path):\n    with Image.open(str(path)) as img:\n        frames = [np.array(f.copy()) for f in ImageSequence.Iterator(img)]\n    return np.stack(frames, axis=0)\n\n# --- 6. FULL MULTI-MODEL OVERLAY ENGINE ---\ndef run_all_models_overlay(checkpoint_groups: Dict[str, List[Path]]):\n    print(f\"Reading Volume: {TIF_PATH.name}\")\n    vol = read_volume_local(TIF_PATH)\n    z_max = vol.shape[0]\n    indices = np.linspace(0, z_max - 1, 5, dtype=int)\n    padded = np.pad(vol.astype(np.float32), ((Z_CONTEXT, Z_CONTEXT), (0, 0), (0, 0)), mode=\"reflect\")\n\n    for key, paths in checkpoint_groups.items():\n        for p in paths:\n            if not p.exists():\n                print(f\"  [SKIP] {p.name} not found.\")\n                continue\n            \n            label = f\"{key}/{p.stem}\"\n            print(f\"\\n---> Starting Overlay for: {label}\")\n            \n            # Load and replicate weights\n            variables = load_npz_params_local(p)\n            weights_rep = jax.device_put_replicated(variables, jax.devices())\n\n            def get_prob(z):\n                stack = padded[z: z + 2 * Z_CONTEXT + 1]\n                stack = (stack - stack.mean()) / (stack.std() + 1e-6)\n                return _tta_predict_pmap(weights_rep, stack.transpose(1, 2, 0))\n\n            # Run predictions for the 5 slices\n            probs = [get_prob(z) for z in indices]\n                    \n            fig, axes = plt.subplots(5, 6, figsize=(22, 16), facecolor=_DARK)\n            for i, z_idx in enumerate(indices):\n                axes[i, 0].imshow(_to_u8(vol[z_idx]), cmap='gray')\n                axes[i, 0].set_title(f\"Z={z_idx}\", color='white', fontsize=10)\n                axes[i, 0].axis('off')\n                \n                for j, thr in enumerate(SHOWCASE_THRESHOLDS):\n                    binary = (probs[i] > thr)\n                    overlay_img = _get_overlay_image(vol[z_idx], binary)\n                    axes[i, j+1].imshow(overlay_img)\n                    axes[i, j+1].set_title(f\"Thr: {thr}\\n{binary.sum():,} px\", color='#00ff88', fontsize=9)\n                    axes[i, j+1].axis('off')\n\n            plt.suptitle(f\"OVERLAY SHOWCASE: {label}\\n7-Slice Sandwich | 8-fold TTA\", color='white', fontsize=20, y=0.98)\n            plt.tight_layout(rect=[0, 0.03, 1, 0.95])\n            \n            save_name = f\"{key}_{p.stem}_overlay.png\"\n            plt.savefig(SAVE_DIR / save_name, facecolor=_DARK)\n            print(f\"  [SAVED] {save_name}\")\n            plt.show()\n            plt.close(fig) # Explicitly close to prevent memory leaks\n\n# --- EXECUTION ---\nrun_all_models_overlay(MODEL_CHECKPOINTS)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-04T21:34:08.941547Z","iopub.execute_input":"2026-03-04T21:34:08.942479Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}