{"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":"tpuV5e8","dataSources":[{"sourceType":"competition","sourceId":117682,"databundleVersionId":15062069},{"sourceType":"datasetVersion","sourceId":14810989,"datasetId":9471057,"databundleVersionId":15666974},{"sourceType":"datasetVersion","sourceId":14961699,"datasetId":9576610,"databundleVersionId":15832991},{"sourceType":"datasetVersion","sourceId":14855360,"datasetId":9484229,"databundleVersionId":15716020},{"sourceType":"modelInstanceVersion","sourceId":769070,"databundleVersionId":15871980,"modelInstanceId":587120,"modelId":599446},{"sourceType":"modelInstanceVersion","sourceId":768516,"databundleVersionId":15864953,"modelInstanceId":587120,"modelId":599446},{"sourceType":"kernelVersion","sourceId":290917305},{"sourceType":"kernelVersion","sourceId":296912426},{"sourceType":"kernelVersion","sourceId":297302543},{"sourceType":"kernelVersion","sourceId":297674933},{"sourceType":"kernelVersion","sourceId":300258358}],"dockerImageVersionId":31261,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Dependencies & Setup","metadata":{}},{"cell_type":"code","source":"# from IPython.display import clear_output\n\n# var=\"/kaggle/input/notebooks/ipythonx/vsdetection-packages-offline-installer-only/whls\"\n\n# # !pip install \\\n# #   \"$var\"/keras_nightly-*.whl \\\n# #   \"$var\"/tifffile-*.whl \\\n# #   \"$var\"/imagecodecs-*.whl \\\n# #   \"$var\"/medicai-*.whl \\\n# #   --no-index \\\n# #   --find-links \"$var\"\n# var=\"/kaggle/input/notebooks/crischir/vsdetection-packages-offline-installer-only/whls\"\n# # clear_output()\n# !pip install \\\n#   \"$var\"/tifffile-*.whl \\\n#   \"$var\"/imagecodecs-*.whl \\\n#   \"$var\"/scikit_image-*.whl \\\n#   --no-index \\\n#   --find-links \"$var\n# # clear_output()\n","metadata":{"trusted":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2026-02-28T20:58:50.053381Z","iopub.execute_input":"2026-02-28T20:58:50.053520Z","iopub.status.idle":"2026-02-28T20:58:50.060551Z","shell.execute_reply.started":"2026-02-28T20:58:50.053504Z","shell.execute_reply":"2026-02-28T20:58:50.059810Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"clDice- a Novel Topology-Preserving Loss Function for Tubular Structure\nSegmentation\nSuprosanna Shit *1\nJohannes C. Paetzold ∗ 1 Anjany Sekuboyina1 Ivan Ezhov1\nAlexander Unger1 Andrey Zhylka2 Josien P. W. Pluim2 Ulrich Bauer1 Bjoern H. Menze1\n1Technical University of Munich 2 Eindhoven University of Technology","metadata":{}},{"cell_type":"markdown","source":"\n## Credits \nclDice -- A Novel Topology-Preserving Loss Function for Tubular Structure Segmentation\nSuprosanna Shit, Johannes C. Paetzold, Anjany Sekuboyina, Ivan Ezhov, Alexander Unger, Andrey Zhylka, Josien P. W. Pluim, Ulrich Bauer, Bjoern H. Menze\nAccurate segmentation of tubular, network-like structures, such as vessels, neurons, or roads, is relevant to many fields of research. For such structures, the topology is their most important characteristic; particularly preserving connectedness: in the case of vascular networks, missing a connected vessel entirely alters the blood-flow dynamics. We introduce a novel similarity measure termed centerlineDice (short clDice), which is calculated on the intersection of the segmentation masks and their (morphological) skeleta. We theoretically prove that clDice guarantees topology preservation up to homotopy equivalence for binary 2D and 3D segmentation. Extending this, we propose a computationally efficient, differentiable loss function (soft-clDice) for training arbitrary neural segmentation networks. We benchmark the soft-clDice loss on five public datasets, including vessels, roads and neurons (2D and 3D). Training on soft-clDice leads to segmentation with more accurate connectivity information, higher graph similarity, and better volumetric scores.\nComments:\t* The authors Suprosanna Shit and Johannes C. Paetzold contributed equally to the work\nSubjects:\tComputer Vision and Pattern Recognition (cs.CV); Machine Learning (cs.LG); Image and Video Processing (eess.IV)\nReport number:\tCVPR 2021\nCite as:\tarXiv:2003.07311 [cs.CV]\n \t(or arXiv:2003.07311v7 [cs.CV] for this version)\n \nhttps://doi.org/10.48550/arXiv.2003.07311\nFocus to learn more\nRelated DOI:\nhttps://doi.org/10.1109/CVPR46437.2021.01629\nFocus to learn more\nSubmission history\nFrom: Johannes C. Paetzold [view email]\n[v1] Mon, 16 Mar 2020 16:27:49 UTC (7,206 KB)\n[v2] Mon, 23 Mar 2020 20:45:16 UTC (7,206 KB)\n[v3] Sun, 29 Mar 2020 22:46:43 UTC (7,206 KB)\n[v4] Thu, 3 Dec 2020 19:53:43 UTC (7,241 KB)\n[v5] Mon, 29 Mar 2021 13:36:28 UTC (14,293 KB)\n[v6] Tue, 30 Mar 2021 11:51:21 UTC (14,293 KB)\n[v7] Fri, 15 Jul 2022 10:39:38 UTC (7,146 KB)","metadata":{}},{"cell_type":"markdown","source":"https://github.com/cpuimage/clDice/tree/master","metadata":{}},{"cell_type":"code","source":"from IPython.display import clear_output\nvar = \"/kaggle/input/notebooks/crischir/vsdetection-packages-offline-installer-only/whls\"\n\n!pip install \\\n  {var}/tifffile-*.whl \\\n  {var}/imagecodecs-*.whl \\\n  {var}/scikit_image-*.whl \\\n  --no-index \\\n  --find-links {var}\nclear_output()","metadata":{"trusted":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2026-02-28T20:58:50.060877Z","iopub.execute_input":"2026-02-28T20:58:50.061024Z","iopub.status.idle":"2026-02-28T20:58:54.352856Z","shell.execute_reply.started":"2026-02-28T20:58:50.061009Z","shell.execute_reply":"2026-02-28T20:58:54.351973Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from flax.training import train_state, checkpoints\nfrom flax.traverse_util import flatten_dict, unflatten_dict\nfrom scipy import ndimage\nfrom skimage.measure import label, regionprops\nimport warnings\nwarnings.filterwarnings(\"ignore\", category=FutureWarning, module=\"keras.src.export.tf2onnx_lib\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T20:58:54.353468Z","iopub.execute_input":"2026-02-28T20:58:54.353641Z","iopub.status.idle":"2026-02-28T20:59:31.640905Z","shell.execute_reply.started":"2026-02-28T20:58:54.353621Z","shell.execute_reply":"2026-02-28T20:59:31.640228Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import imagecodecs\nimport tifffile\nimport numpy as np","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T20:59:31.641409Z","iopub.execute_input":"2026-02-28T20:59:31.641736Z","iopub.status.idle":"2026-02-28T20:59:31.652956Z","shell.execute_reply.started":"2026-02-28T20:59:31.641719Z","shell.execute_reply":"2026-02-28T20:59:31.652321Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\nimport os\nimport time\nimport random\nimport zipfile\nimport numpy as np\n\nimport multiprocessing\n\nos.environ[\"JAX_PLATFORMS\"] = \"tpu\"\n\nimport jax\nimport jax.numpy as jnp\nimport flax\nimport flax.linen as nn\nimport optax\nimport matplotlib.pyplot as plt\nfrom pathlib import Path\nfrom PIL import Image, ImageSequence\nfrom tqdm import tqdm\nfrom torch.utils.data import DataLoader, Dataset\nfrom flax.training import train_state, checkpoints\nfrom flax.traverse_util import flatten_dict, unflatten_dict\nfrom scipy import ndimage\nfrom skimage.measure import label, regionprops\n\n# ==========================================\n# 1. CONFIGURATION\n# ==========================================\nDATA_PATH = Path(\"/kaggle/input/vesuvius-challenge-surface-detection\")\nOUT_DIR = Path(\"/kaggle/temp/vesuvius_npy\")\nOUT_DIR.mkdir(exist_ok=True, parents=True)\nCHECKPOINT_DIR = Path(\"/kaggle/working/checkpoints\")\nCHECKPOINT_DIR.mkdir(parents=True, exist_ok=True)\n\nPATCH_SIZE = 256\nZ_CONTEXT = 3\nBATCH_SIZE = 128\nLEARNING_RATE_G = 2e-4\nLEARNING_RATE_D = 1e-6\n\nSAMPLES_PER_ID = 9\nMAX_BATCHES_PER_EPOCH = 1800\nMAX_TRAIN_TIME_HOURS = 2.8\nMAX_TRAIN_SECONDS = MAX_TRAIN_TIME_HOURS * 3600\nMAX_EPOCHS = 300\nLOADWeights = True\n\n# clDice hyperparameters\nALPHA_CLDICE = 0.2   # weight between hybrid loss and clDice\nSKEL_ITERS = 14      # k in the paper (soft-skeleton iterations)\n\n# ==========================================\n# SCHEDULER HELPER\n# ==========================================\ndef get_linear_schedule(init_lr, total_epochs, steps_per_epoch, constant_epochs=20):\n    constant_steps = constant_epochs * steps_per_epoch\n    total_steps = total_epochs * steps_per_epoch\n    decay_steps = total_steps - constant_steps\n    \n    schedule = optax.join_schedules(\n        schedules=[\n            optax.constant_schedule(init_lr),\n            optax.linear_schedule(init_lr, 0.0, decay_steps)\n        ],\n        boundaries=[constant_steps]\n    )\n    return schedule\n\n# ==========================================\n# 1. Load existing weights\n# ==========================================\npath_to_weights = \"/kaggle/input/models/crischir/vesuvius-2-5-mask-cldice/flax/2.5d-cldice-24-samples-40-epochs/2/checkpoints/pix2pix_ep_40.npz\"\nif LOADWeights:\n    with np.load(path_to_weights) as data:\n        flat_dict = {k: v for k, v in data.items()}\n    p_flat = {k.replace('params/', ''): v for k, v in flat_dict.items() if k.startswith('params/')}\n    s_flat = {k.replace('stats/', ''): v for k, v in flat_dict.items() if k.startswith('stats/')}\n    loaded_params = unflatten_dict({tuple(k.split('/')): v for k, v in p_flat.items()})\n    loaded_stats = unflatten_dict({tuple(k.split('/')): v for k, v in s_flat.items()})\n\n# ==========================================\n# 2. PREPROCESSING (TIFF to NPY)\n# ==========================================\ndef tiff_to_npy(tif_path: Path, out_path: Path):\n    if out_path.exists():\n        return\n    with Image.open(str(tif_path)) as img:\n        frames = [np.array(frame) for frame in ImageSequence.Iterator(img)]\n        vol = np.stack(frames, axis=0)\n    np.save(out_path, vol)\n\ndef run_preprocessing():\n    IMG_SRC = DATA_PATH / \"train_images\"\n    LBL_SRC = DATA_PATH / \"train_labels\"\n    for tif in tqdm(sorted(IMG_SRC.glob(\"*.tif\")), desc=\"Preprocessing Images\"):\n        tiff_to_npy(tif, OUT_DIR / f\"{tif.stem}_img.npy\")\n    for tif in tqdm(sorted(LBL_SRC.glob(\"*.tif\")), desc=\"Preprocessing Labels\"):\n        tiff_to_npy(tif, OUT_DIR / f\"{tif.stem}_lbl.npy\")\n    # for tif in tqdm(sorted(IMG_SRC.glob(\"*.tif\"))[:20], desc=\"Preprocessing Images\"):\n    #     tiff_to_npy(tif, OUT_DIR / f\"{tif.stem}_img.npy\")\n    # for tif in tqdm(sorted(LBL_SRC.glob(\"*.tif\"))[:20], desc=\"Preprocessing Labels\"):\n    #     tiff_to_npy(tif, OUT_DIR / f\"{tif.stem}_lbl.npy\")\n\n# ==========================================\n# 3. DATASET & AUGMENTATION\n# ==========================================\ndef normalize_patch(patch: np.ndarray) -> np.ndarray:\n    patch = patch.astype(np.float32)\n    return (patch - patch.mean()) / (patch.std() + 1e-6)\n\nclass VesuviusSurfaceDataset(Dataset):\n    def __init__(self, ids, npy_dir, samples_per_id, train=True):\n        self.patch_size, self.z_context = PATCH_SIZE, Z_CONTEXT\n        self.rng = np.random.default_rng(42)\n        self.coords, self.volumes, self.labels = [], {}, {}\n        self.train = train\n\n        for vid in ids:\n            img_npy, lbl_npy = npy_dir / f\"{vid}_img.npy\", npy_dir / f\"{vid}_lbl.npy\"\n            if not img_npy.exists() or not lbl_npy.exists():\n                continue\n            img_vol = np.load(img_npy, mmap_mode=\"r\")\n            z_max, h, w = img_vol.shape\n            z_min, z_max_r = max(self.z_context, 5), min(z_max - self.z_context - 1, z_max - 5)\n            if z_max_r <= z_min:\n                continue\n            zs = self.rng.integers(z_min, z_max_r, size=samples_per_id)\n            ys = self.rng.integers(0, max(1, h - self.patch_size), size=samples_per_id)\n            xs = self.rng.integers(0, max(1, w - self.patch_size), size=samples_per_id)\n            for z, y, x in zip(zs, ys, xs):\n                self.coords.append({\n                    \"vid\": vid, \"z\": int(z), \"y\": int(y), \"x\": int(x),\n                    \"img_npy\": img_npy, \"lbl_npy\": lbl_npy\n                })\n\n    def _random_augment(self, img, target):\n        if self.rng.random() > 0.5:\n            img, target = np.flip(img, axis=1), np.flip(target, axis=1)\n        if self.rng.random() > 0.5:\n            img, target = np.flip(img, axis=0), np.flip(target, axis=0)\n        k = self.rng.integers(0, 4)\n        if k > 0:\n            img, target = np.rot90(img, k=k, axes=(0, 1)), np.rot90(target, k=k, axes=(0, 1))\n        return img.copy(), target.copy()\n\n    def __getitem__(self, idx):\n        item = self.coords[idx]\n        vid = item[\"vid\"]\n        if vid not in self.volumes:\n            self.volumes[vid] = np.load(item[\"img_npy\"], mmap_mode=\"r\")\n            self.labels[vid] = np.load(item[\"lbl_npy\"], mmap_mode=\"r\")\n        z, y, x = item[\"z\"], item[\"y\"], item[\"x\"]\n        img_patch = self.volumes[vid][z-Z_CONTEXT:z+Z_CONTEXT+1, y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n        img_patch = img_patch.transpose(1, 2, 0)\n        z_lbl = z if self.labels[vid].shape[0] > 1 else 0\n        target_patch = self.labels[vid][z_lbl, y:y+PATCH_SIZE, x:x+PATCH_SIZE][..., None]\n        if self.train:\n            img_patch, target_patch = self._random_augment(img_patch, target_patch)\n        return normalize_patch(img_patch), target_patch.astype(np.float32)\n\n    def __len__(self):\n        return len(self.coords)\n\n# ==========================================\n# 4. MODELS\n# ==========================================\nclass Pix2PixGenerator(nn.Module):\n    @nn.compact\n    def __call__(self, x, train: bool = True):\n        def conv_bn_leaky(feat, out_c):\n            feat = nn.Conv(out_c, (4,4), strides=(2,2), padding='SAME')(feat)\n            feat = nn.BatchNorm(use_running_average=not train)(feat)\n            return nn.leaky_relu(feat, 0.2)\n        \n        s1 = nn.leaky_relu(nn.Conv(64, (4,4), strides=(2,2), padding='SAME')(x), 0.2)\n        s2 = conv_bn_leaky(s1, 128)\n        s3 = conv_bn_leaky(s2, 256)\n        \n        def up_bn_relu(feat, skip, out_c):\n            feat = nn.ConvTranspose(out_c, (4,4), strides=(2,2), padding='SAME')(feat)\n            feat = nn.BatchNorm(use_running_average=not train)(feat)\n            feat = nn.relu(feat)\n            return jnp.concatenate([feat, skip], axis=-1)\n        \n        u1 = up_bn_relu(s3, s2, 128)\n        u2 = up_bn_relu(u1, s1, 64)\n        return nn.ConvTranspose(1, (4,4), strides=(2,2), padding='SAME')(u2)\n\nclass TopologyDiscriminator(nn.Module):\n    @nn.compact\n    def __call__(self, x, label, train: bool = True):\n        inputs = jnp.concatenate([x, label], axis=-1)\n        y = nn.leaky_relu(nn.BatchNorm(use_running_average=not train)(nn.Conv(64, (4,4), strides=(2,2), padding='SAME')(inputs)), 0.2)\n        y = nn.leaky_relu(nn.BatchNorm(use_running_average=not train)(nn.Conv(128, (4,4), strides=(2,2), padding='SAME')(y)), 0.2)\n        y = nn.leaky_relu(nn.BatchNorm(use_running_average=not train)(nn.Conv(256, (4,4), strides=(2,2), padding='SAME')(y)), 0.2)\n        return nn.Conv(1, (4,4), padding='SAME')(y)\n\nclass JaxTrainState(train_state.TrainState):\n    batch_stats: flax.core.FrozenDict\n\ndef save_portable_npy(state, filename):\n    flat_params = flatten_dict(jax.device_get(state.params), sep='/')\n    flat_stats = flatten_dict(jax.device_get(state.batch_stats), sep='/')\n    save_dict = {f\"params/{k}\": v for k, v in flat_params.items()}\n    save_dict.update({f\"stats/{k}\": v for k, v in flat_stats.items()})\n    np.savez(CHECKPOINT_DIR / filename, **save_dict)\n    print(f\"⭐ Portable .npz weights saved: {filename}\")\n\n# ==========================================\n# 5. LOSSES (hybrid + clDice)\n# ==========================================\ndef hybrid_loss(logits, target):\n    valid_mask = (target != 2).astype(jnp.float32)\n    target_binary = (target == 1).astype(jnp.float32)\n    bce = optax.sigmoid_binary_cross_entropy(logits, target_binary)\n    masked_bce = jnp.sum(bce * valid_mask) / (jnp.sum(valid_mask) + 1e-6)\n    pred = nn.sigmoid(logits)\n    dice = 1 - (2.0 * jnp.sum(pred * target_binary * valid_mask) + 1e-6) / \\\n           (jnp.sum(pred * valid_mask) + jnp.sum(target_binary * valid_mask) + 1e-6)\n    return 0.5 * masked_bce + 0.5 * dice\n\n# >>> Soft-skeletonization (JAX)\ndef maxpool_2d(x):\n    return jax.lax.reduce_window(\n        x,\n        -jnp.inf,\n        jax.lax.max,\n        window_dimensions=(1, 3, 3, 1),\n        window_strides=(1, 1, 1, 1),\n        padding='SAME'\n    )\n\ndef minpool_2d(x):\n    return jax.lax.reduce_window(\n        x,\n        jnp.inf,\n        jax.lax.min,\n        window_dimensions=(1, 3, 3, 1),\n        window_strides=(1, 1, 1, 1),\n        padding='SAME'\n    )\n\n# def soft_skeletonize(img, k=SKEL_ITERS):\n#     I = img\n#     I_eroded = minpool_2d(I)\n#     I_dilated = maxpool_2d(I_eroded)\n#     S = jax.nn.relu(I - I_dilated)\n\n#     def body_fn(i, val):\n#         I, S = val\n#         I = minpool_2d(I)\n#         I_eroded = minpool_2d(I)\n#         I_dilated = maxpool_2d(I_eroded)\n#         new = jax.nn.relu(I - I_dilated)\n#         S = S + (1.0 - S) * new\n#         return (I, S)\n\n#     _, S = jax.lax.fori_loop(0, k, body_fn, (I, S))\n#     return S\n# def soft_skeletonize(x, maximum_iterations=10, kernel_size=3):\n#     def body_fn(_, x_curr):\n#         min_pool_x = erosion2d(x_curr, kernel_size=kernel_size)\n#         contour = jax.nn.relu(dilation2d(min_pool_x, kernel_size=kernel_size) - min_pool_x)\n#         x_next = jax.nn.relu(x_curr - contour)\n#         return x_next\n\n#     return jax.lax.fori_loop(0, maximum_iterations, body_fn, x)\n\ndef soft_cldice(pred, target, k=SKEL_ITERS, eps=1e-6):\n    # pred, target: (B, H, W, 1) in [0,1]\n    Sp = soft_skeletonize(pred, k)\n    St = soft_skeletonize(target, k)\n\n    tprec = (jnp.sum(Sp * target) + eps) / (jnp.sum(Sp) + eps)\n    tsens = (jnp.sum(St * pred) + eps) / (jnp.sum(St) + eps)\n\n    return 2.0 * tprec * tsens / (tprec + tsens + eps)\n\n\n\n\n# -----------------------------\n# Fixed-iteration soft skeleton\n# -----------------------------\n# def fixed_soft_skeletonize(x, maximum_iterations=10, kernel_size=3, dilations=1):\n#     \"\"\"\n#     Fixed number of thinning iterations.\n#     x: (B, H, W, C) in [0,1]\n#     \"\"\"\n#     def body_fn(_, x_curr):\n#         min_pool_x = erosion2d(x_curr, kernel_size=kernel_size, dilations=dilations)\n#         contour = jax.nn.relu(\n#             dilation2d(min_pool_x, kernel_size=kernel_size, dilations=dilations) - min_pool_x\n#         )\n#         x_next = jax.nn.relu(x_curr - contour)\n#         return x_next\n\n#     x_out = jax.lax.fori_loop(0, maximum_iterations, body_fn, x)\n#     return x_out\n\n\n# -----------------------------\n# Adaptive soft skeleton (while)\n# -----------------------------\n# def soft_skeletonize(x, maximum_iterations=10, kernel_size=3, dilations=1, threshold=1.0):\n#     \"\"\"\n#     While-loop version: stops when eroded mass < threshold or max iters reached.\n#     x: (B, H, W, C)\n#     \"\"\"\n#     def body_fn(loop_vars):\n#         prev_eroded, skel = loop_vars\n#         eroded = erosion2d(skel, kernel_size=kernel_size, dilations=dilations)\n#         skel_next = jax.nn.relu(\n#             skel - jax.nn.relu(dilation2d(eroded, kernel_size=kernel_size, dilations=dilations) - eroded)\n#         )\n#         return (eroded, skel_next)\n\n#     def cond_fn(loop_vars):\n#         prev_eroded, _ = loop_vars\n#         # sum over spatial + channel dims, per batch\n#         mass = jnp.sum(prev_eroded, axis=(1, 2, 3))\n#         return jnp.any(mass > threshold)\n\n#     # We emulate TF's while_loop with a bounded loop + cond\n#     def bounded_while(loop_vars):\n#         def inner_body(i, lv):\n#             return jax.lax.cond(\n#                 cond_fn(lv),\n#                 lambda v: body_fn(v),\n#                 lambda v: v,\n#                 lv\n#             )\n#         return jax.lax.fori_loop(0, maximum_iterations, inner_body, loop_vars)\n\n#     _, x_out = bounded_while((x, x))\n#     return x_out\ndef soft_skeletonize(x, maximum_iterations=10):\n    def body_fn(_, x_curr):\n        min_pool_x = erosion2d(x_curr)\n        contour = jax.nn.relu(dilation2d(min_pool_x) - min_pool_x)\n        x_next = jax.nn.relu(x_curr - contour)\n        return x_next\n\n    return jax.lax.fori_loop(0, maximum_iterations, body_fn, x)\n\n# -----------------------------\n# Normalized intersection\n# -----------------------------\ndef norm_intersection(center_line, vessel, eps=1e-8):\n    \"\"\"\n    center_line, vessel: (B, H, W, C)\n    returns: (B, 1, 1, 1) normalized intersection\n    \"\"\"\n    intersection = jnp.sum(center_line * vessel, axis=(1, 2, 3), keepdims=True)\n    denom = jnp.sum(center_line, axis=(1, 2, 3), keepdims=True)\n    return intersection / (denom + eps)\n\n\n# -----------------------------\n# Soft clDice loss\n# -----------------------------\n# def soft_cldice_losses(\n#     y_true,\n#     y_pred,\n#     true_skeleton=None,\n#     maximum_iterations=10,\n#     fixed_iterations=True,\n#     kernel_size=3,\n#     dilations=1,\n#     threshold=1.0,\n# ):\n#     \"\"\"\n#     y_true, y_pred: (B, H, W, C) in [0,1]\n#     returns: (B, 1, 1, 1) clDice loss per batch element\n#     \"\"\"\n#     if fixed_iterations:\n#         soft_skel_fn = lambda x: fixed_soft_skeletonize(\n#             x,\n#             maximum_iterations=maximum_iterations,\n#             kernel_size=kernel_size,\n#             dilations=dilations,\n#         )\n#     else:\n#         soft_skel_fn = lambda x: soft_skeletonize(\n#             x,\n#             maximum_iterations=maximum_iterations,\n#             kernel_size=kernel_size,\n#             dilations=dilations,\n#             threshold=threshold,\n#         )\n\n#     pred_skeleton = soft_skel_fn(y_pred)\n\n#     if true_skeleton is None:\n#         true_skeleton = soft_skel_fn(y_true)\n\n#     iflat = norm_intersection(pred_skeleton, y_true)\n#     tflat = norm_intersection(true_skeleton, y_pred)\n\n#     num = 2.0 * iflat * tflat\n#     den = iflat + tflat + 1e-8\n#     cldice = num / den\n\n#     loss = 1.0 - cldice\n#     return loss  # shape (B, 1, 1, 1)\n\ndef soft_cldice_losses(y_true, y_pred, maximum_iterations=10):\n    pred_skeleton = soft_skeletonize(y_pred, maximum_iterations)\n    true_skeleton = soft_skeletonize(y_true, maximum_iterations)\n\n    iflat = norm_intersection(pred_skeleton, y_true)\n    tflat = norm_intersection(true_skeleton, y_pred)\n\n    cl = (2 * iflat * tflat) / (iflat + tflat + 1e-8)\n    return 1.0 - cl\n\n# -----------------------------\n# 2D dilation and erosion (NHWC)\n# -----------------------------\ndef erosion2d(x, kernel_size=3, strides=1):\n    return jax.lax.reduce_window(\n        x,\n        jnp.inf,\n        jax.lax.min,\n        window_dimensions=(1, kernel_size, kernel_size, 1),\n        window_strides=(1, strides, strides, 1),\n        padding=\"SAME\"\n    )\n\ndef dilation2d(x, kernel_size=3, strides=1):\n    return jax.lax.reduce_window(\n        x,\n        -jnp.inf,\n        jax.lax.max,\n        window_dimensions=(1, kernel_size, kernel_size, 1),\n        window_strides=(1, strides, strides, 1),\n        padding=\"SAME\"\n    )\n\n# ==========================================\n# 5b. TRAIN STEP WITH clDice\n# ==========================================\n@jax.jit\ndef train_step(state_g, state_d, imgs, targets, d_train_flag):\n    # imgs: (B, H, W, C), targets: (B, H, W, 1) with {0,1,2}\n    def loss_d_fn(params):\n        fake_logits = state_g.apply_fn(\n            {'params': state_g.params, 'batch_stats': state_g.batch_stats},\n            imgs,\n            train=False\n        )\n        fake_targets = nn.sigmoid(fake_logits)\n        combined_imgs = jnp.concatenate([imgs, imgs], axis=0)\n        combined_lbls = jnp.concatenate(\n            [(targets == 1).astype(jnp.float32), fake_targets],\n            axis=0\n        )\n        (d_out, updates) = state_d.apply_fn(\n            {'params': params, 'batch_stats': state_d.batch_stats},\n            combined_imgs,\n            combined_lbls,\n            train=True,\n            mutable=['batch_stats']\n        )\n        d_real, d_fake = jnp.split(d_out, 2, axis=0)\n        loss_d = jnp.mean((d_real - 1.0) ** 2) + jnp.mean(d_fake ** 2)\n        return loss_d, updates\n\n    def loss_g_fn(params):\n        (gen_logits, updates) = state_g.apply_fn(\n            {'params': params, 'batch_stats': state_g.batch_stats},\n            imgs,\n            train=True,\n            mutable=['batch_stats']\n        )\n        pred = nn.sigmoid(gen_logits)\n\n        # base hybrid loss (BCE + Dice with ignore=2)\n        base_loss = hybrid_loss(gen_logits, targets)\n\n        # clDice on foreground only (targets == 1)\n        fg_target = (targets == 1).astype(jnp.float32)\n        # cl = soft_cldice(pred, fg_target)\n        # cl = soft_cldice_losses(\n        #                         fg_target,          # y_true\n        #                         pred,               # y_pred\n        #                         maximum_iterations=SKEL_ITERS,\n        #                         fixed_iterations=True,\n        #                         kernel_size=3,\n        #                         dilations=1\n        #                     ).mean()\n        cl = soft_cldice_losses(\n                            fg_target,\n                            pred,\n                            maximum_iterations=SKEL_ITERS\n                        ).mean()\n\n        topo_loss = (1.0 - ALPHA_CLDICE) * base_loss + ALPHA_CLDICE * (1.0 - cl)\n\n        d_out = state_d.apply_fn(\n            {'params': state_d.params, 'batch_stats': state_d.batch_stats},\n            imgs,\n            pred,\n            train=False\n        )\n        adv_loss = 0.1 * jnp.mean((d_out - 1.0) ** 2)\n\n        total_loss = topo_loss + adv_loss\n        return total_loss, updates\n\n    # Step D\n    grad_d_fn = jax.value_and_grad(loss_d_fn, has_aux=True)\n    (loss_d, d_stats), grad_d = grad_d_fn(state_d.params)\n    state_d = jax.lax.cond(\n        d_train_flag,\n        lambda s, g, st: s.apply_gradients(grads=g).replace(batch_stats=st['batch_stats']),\n        lambda s, g, st: s,\n        state_d,\n        grad_d,\n        d_stats\n    )\n\n    # Step G\n    grad_g_fn = jax.value_and_grad(loss_g_fn, has_aux=True)\n    (loss_g, g_stats), grad_g = grad_g_fn(state_g.params)\n    state_g = state_g.apply_gradients(grads=grad_g).replace(batch_stats=g_stats['batch_stats'])\n\n    return state_g, state_d, loss_g, loss_d\n\n# ==========================================\n# 6. VISUALIZATION & POST-PROCESSING\n# ==========================================\ndef apply_cc_filter(mask: np.ndarray, min_size: int = 150):\n    labeled, num_features = ndimage.label(mask)\n    if num_features == 0:\n        return mask\n    component_sizes = np.bincount(labeled.ravel())\n    mask[component_sizes[labeled] < min_size] = 0\n    return mask\n\ndef visualize_results(state_g, dataset, epoch):\n    rows, cols = 5, 7\n    fig, axes = plt.subplots(rows, cols, figsize=(cols * 4, rows * 4))\n    for i in range(rows):\n        img_np, target_np = dataset[random.randint(0, len(dataset)-1)]\n        logits = state_g.apply_fn(\n            {'params': state_g.params, 'batch_stats': state_g.batch_stats},\n            jnp.array(img_np[None, ...]),\n            train=False\n        )\n        pred = np.array(nn.sigmoid(logits)).squeeze()\n        thr = (pred > 0.3).astype(np.uint8)\n        cleaned = apply_cc_filter(thr.copy(), 150)\n        target_viz = np.zeros((*target_np.shape[:2], 3))\n        target_viz[target_np.squeeze() == 1] = [1, 1, 1]\n        target_viz[target_np.squeeze() == 2] = [0.5, 0.5, 0.5]\n        overlay_in = np.zeros((*cleaned.shape, 3))\n        overlay_in[..., 2] = 0.3\n        overlay_in[..., 0] = cleaned\n        overlay_lbl = target_viz.copy()\n        overlay_lbl[..., 0] = np.maximum(overlay_lbl[..., 0], cleaned.astype(float))\n        axes[i, 0].imshow(img_np[..., Z_CONTEXT], cmap=\"gray\")\n        axes[i, 1].imshow(target_viz)\n        axes[i, 2].imshow(pred, cmap=\"magma\")\n        axes[i, 3].imshow(thr, cmap=\"gray\")\n        axes[i, 4].imshow(cleaned, cmap=\"gray\")\n        axes[i, 5].imshow(overlay_in)\n        axes[i, 6].imshow(overlay_lbl)\n        for j in range(cols):\n            axes[i, j].axis(\"off\")\n    plt.tight_layout()\n    plt.savefig(f\"viz_{epoch}.png\")\n    plt.show()\n    plt.close()\n\n# ==========================================\n# 7. MAIN EXECUTION\n# ==========================================\nif __name__ == \"__main__\":\n    run_preprocessing()\n    ids = sorted({p.stem.replace(\"_img\", \"\") for p in OUT_DIR.glob(\"*_img.npy\")})\n    ds = VesuviusSurfaceDataset(ids, OUT_DIR, SAMPLES_PER_ID, train=True)\n    loader = DataLoader(ds, batch_size=BATCH_SIZE, shuffle=True, drop_last=True, num_workers=0)\n    steps_per_epoch = len(loader)\n    rng = jax.random.PRNGKey(42)\n    netG, netD = Pix2PixGenerator(), TopologyDiscriminator()\n    dummy = jnp.ones((1, PATCH_SIZE, PATCH_SIZE, 2*Z_CONTEXT+1))\n\n    sched_G = get_linear_schedule(LEARNING_RATE_G, MAX_EPOCHS, steps_per_epoch)\n    sched_D = get_linear_schedule(LEARNING_RATE_D, MAX_EPOCHS, steps_per_epoch)\n\n    v_g = netG.init(jax.random.split(rng)[0], dummy)\n    if LOADWeights:\n        state_g = JaxTrainState.create(\n            apply_fn=netG.apply,\n            params=loaded_params,\n            batch_stats=loaded_stats,\n            tx=optax.adam(sched_G, 0.5, 0.999)\n        )\n    else:\n        state_g = JaxTrainState.create(\n            apply_fn=netG.apply,\n            params=v_g['params'],\n            batch_stats=v_g['batch_stats'],\n            tx=optax.adam(sched_G, 0.5, 0.999)\n        )\n\n    v_d = netD.init(jax.random.split(rng)[1], dummy, jnp.ones((1, PATCH_SIZE, PATCH_SIZE, 1)))\n    state_d = JaxTrainState.create(\n        apply_fn=netD.apply,\n        params=v_d['params'],\n        batch_stats=v_d['batch_stats'],\n        tx=optax.adam(sched_D, 0.5, 0.999)\n    )\n\n    best_loss, start_time = float('inf'), time.time()\n\n    for ep in range(1, MAX_EPOCHS + 1):\n        if (time.time() - start_time) > MAX_TRAIN_SECONDS:\n            print(\"⏳ Max training time reached.\")\n            break\n        \n        epoch_g_loss, epoch_d_loss = [], []\n        for step, (b_imgs, b_tgts) in enumerate(loader):\n            state_g, state_d, lg, ld = train_step(\n                state_g,\n                state_d,\n                jnp.array(b_imgs),\n                jnp.array(b_tgts),\n                step % 2 == 0\n            )\n            epoch_g_loss.append(lg)\n            epoch_d_loss.append(ld)\n\n            if step % 100 == 0:\n                print(f\"Ep {ep} | Step {step}/{steps_per_epoch} | G_loss: {lg:.4f} | D_loss: {ld:.4f}\")\n\n        avg_g = float(np.mean(jax.device_get(jnp.array(epoch_g_loss))))\n        avg_d = float(np.mean(jax.device_get(jnp.array(epoch_d_loss))))\n        elapsed = (time.time() - start_time) / 3600\n        print(f\"✅ Epoch {ep} Complete | Avg G: {avg_g:.4f} | Avg D: {avg_d:.4f} | Time: {elapsed:.2f}h\")\n\n        if ep % 10 == 0:\n            visualize_results(state_g, ds, ep)\n            save_portable_npy(state_g, f\"pix2pix_ep_{ep}.npz\")\n            \n        if avg_g < best_loss:\n            best_loss = avg_g\n            save_portable_npy(state_g, \"pix2pix_best.npz\")\n\n    # Final Submission Inference\n    test_dir = DATA_PATH / 'test_images'\n    if test_dir.exists():\n        for vid in [f.stem for f in test_dir.glob('*.tif')]:\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            z_m, h, w = vol.shape\n            pages = []\n            for z in tqdm(range(z_m), desc=f\"Infer {vid}\"):\n                if z-Z_CONTEXT < 0 or z+Z_CONTEXT+1 > z_m:\n                    pages.append(Image.fromarray(np.zeros((h, w), dtype=np.uint8)))\n                    continue\n                pp = np.zeros((h, w))\n                for y in range(0, h, PATCH_SIZE):\n                    for x in range(0, w, PATCH_SIZE):\n                        y1, x1 = min(y+PATCH_SIZE, h), min(x+PATCH_SIZE, w)\n                        y0, x0 = y1-PATCH_SIZE, x1-PATCH_SIZE\n                        p = normalize_patch(\n                            vol[z-Z_CONTEXT:z+Z_CONTEXT+1, y0:y1, x0:x1].transpose(1, 2, 0)\n                        )\n                        o = state_g.apply_fn(\n                            {'params': state_g.params, 'batch_stats': state_g.batch_stats},\n                            jnp.array(p[None, ...]),\n                            train=False\n                        )\n                        pp[y0:y1, x0:x1] = np.array(nn.sigmoid(o)).squeeze()\n                pages.append(Image.fromarray(apply_cc_filter((pp > 0.2).astype(np.uint8), 150)))\n            pages[0].save(f\"{vid}.tif\", save_all=True, append_images=pages[1:], compression=\"tiff_deflate\")\n        with zipfile.ZipFile('submission.zip', 'w') as z:\n            for f in Path('.').glob('*.tif'):\n                z.write(f)\n                os.remove(f)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T20:59:31.653498Z","iopub.execute_input":"2026-02-28T20:59:31.653655Z","execution_failed":"2026-02-28T21:03:04.970Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}