{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"tpuV5e8","dataSources":[{"sourceId":117682,"databundleVersionId":15062069,"sourceType":"competition"},{"sourceId":14266465,"sourceType":"datasetVersion","datasetId":8751895},{"sourceId":14266755,"sourceType":"datasetVersion","datasetId":8766236},{"sourceId":14312596,"sourceType":"datasetVersion","datasetId":9136869},{"sourceId":14314377,"sourceType":"datasetVersion","datasetId":9137964},{"sourceId":286271954,"sourceType":"kernelVersion"},{"sourceId":288572598,"sourceType":"kernelVersion"},{"sourceId":674747,"sourceType":"modelInstanceVersion","modelInstanceId":503784,"modelId":510647},{"sourceId":681152,"sourceType":"modelInstanceVersion","modelInstanceId":516822,"modelId":510647},{"sourceId":694469,"sourceType":"modelInstanceVersion","modelInstanceId":522802,"modelId":536801},{"sourceId":696478,"sourceType":"modelInstanceVersion","modelInstanceId":522802,"modelId":536801},{"sourceId":697465,"sourceType":"modelInstanceVersion","modelInstanceId":522802,"modelId":536801},{"sourceId":702889,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":533420,"modelId":547127}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"- **Inference**: [[Inference] Vesuvius Surface 3D Detection](https://www.kaggle.com/code/ipythonx/inference-vesuvius-surface-3d-detection)","metadata":{}},{"cell_type":"code","source":"from IPython.display import clear_output\n\n# This is required for TPU training at the moment in kaggel env.\n# Use the offline wheels that match your weights!\nvar=\"/kaggle/input/vsdetection-packages-offline-installer-only/whls\"\n!pip install \\\n    \"$var\"/keras_nightly-3.12.0.dev2025100703-py3-none-any.whl \\\n    \"$var\"/tifffile-2025.12.12-py3-none-any.whl \\\n    \"$var\"/imagecodecs-2025.11.11-cp311-abi3-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl \\\n    \"$var\"/medicai-0.0.3-py3-none-any.whl \\\n    --no-index \\\n    --find-links \"$var\"\nclear_output()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:10:58.220701Z","iopub.execute_input":"2025-12-28T07:10:58.221288Z","iopub.status.idle":"2025-12-28T07:11:05.513913Z","shell.execute_reply.started":"2025-12-28T07:10:58.221269Z","shell.execute_reply":"2025-12-28T07:11:05.513054Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# The `medicai` is medical-based 2D and 3D ML library. \n# We'll use it for segmentaiton model, 3D volume transformation, etc.\n!pip install git+https://github.com/innat/medic-ai.git -q\n\n# Installing is optional, we'll be using `npy` format instead of `tif`.\n# !pip install imagecodecs tifffile -q","metadata":{"trusted":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2025-12-28T07:11:05.515690Z","iopub.execute_input":"2025-12-28T07:11:05.515938Z","iopub.status.idle":"2025-12-28T07:11:16.878499Z","shell.execute_reply.started":"2025-12-28T07:11:05.515914Z","shell.execute_reply":"2025-12-28T07:11:16.877644Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, warnings\n\nos.environ[\"KERAS_BACKEND\"] = \"jax\"\nwarnings.filterwarnings('ignore')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:11:16.879565Z","iopub.execute_input":"2025-12-28T07:11:16.879839Z","iopub.status.idle":"2025-12-28T07:11:16.884009Z","shell.execute_reply.started":"2025-12-28T07:11:16.879814Z","shell.execute_reply":"2025-12-28T07:11:16.883273Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import glob\nimport numpy as np\nimport pandas as pd\n\nfrom PIL import Image\nimport matplotlib.pyplot as plt\n\n# mainly for training API\nimport keras\nfrom keras import ops\nfrom keras.optimizers.schedules import CosineDecay\n\n# only for tf.data API\nimport tensorflow as tf\n\n# mainly for 3D or 2D models, transformation, loss, metrics etc\nimport medicai\nfrom medicai.transforms import (\n    Compose,\n    ScaleIntensityRange,\n    Resize,\n    RandShiftIntensity,\n    RandRotate90,\n    RandFlip,\n    RandSpatialCrop\n)\nfrom medicai.models import (\n    UNet, SegFormer, TransUNet, SwinUNETR, UPerNet\n)\nfrom medicai.losses import (\n    SparseDiceCELoss, SparseTverskyLoss\n)\nfrom medicai.metrics import SparseDiceMetric\nfrom medicai.callbacks import SlidingWindowInferenceCallback\nfrom medicai.utils.inference import SlidingWindowInference","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:11:16.884766Z","iopub.execute_input":"2025-12-28T07:11:16.885113Z","iopub.status.idle":"2025-12-28T07:11:31.923068Z","shell.execute_reply.started":"2025-12-28T07:11:16.885096Z","shell.execute_reply":"2025-12-28T07:11:31.922485Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# due to distributed training only\nkeras.config.disable_flash_attention()\n\n# reproducibility\nkeras.utils.set_random_seed(101)\n\n# distributed config\ndevices = keras.distribution.list_devices()\ndata_parallel = keras.distribution.DataParallel(devices=devices)\nkeras.distribution.set_distribution(data_parallel)\ntotal_device = len(devices)\n\nprint(f'detected devices: {devices}')\nprint(f'total device: {total_device}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:11:31.924716Z","iopub.execute_input":"2025-12-28T07:11:31.925124Z","iopub.status.idle":"2025-12-28T07:11:32.513565Z","shell.execute_reply.started":"2025-12-28T07:11:31.925105Z","shell.execute_reply":"2025-12-28T07:11:32.512804Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"keras.version(), keras.config.backend(), medicai.version()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:11:32.515119Z","iopub.execute_input":"2025-12-28T07:11:32.515340Z","iopub.status.idle":"2025-12-28T07:11:32.520777Z","shell.execute_reply.started":"2025-12-28T07:11:32.515323Z","shell.execute_reply":"2025-12-28T07:11:32.520154Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:11:32.521378Z","iopub.execute_input":"2025-12-28T07:11:32.521587Z","iopub.status.idle":"2025-12-28T07:11:32.535110Z","shell.execute_reply.started":"2025-12-28T07:11:32.521568Z","shell.execute_reply":"2025-12-28T07:11:32.534505Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_tif_volume(path):\n    \"\"\"\n    Memory optimized loader. Pre-allocates array to avoid \n    memory spike during stacking.\n    \"\"\"\n    \n    img = Image.open(path)\n    W, H = img.size\n    Z = getattr(img, 'n_frames', 1)\n    \n    print(f\"Loading {path} ({Z}x{H}x{W})...\")\n    \n    # Pre-allocate buffer (Standard RAM)\n    vol = np.zeros((Z, H, W), dtype=np.uint8)\n    \n    for i in range(Z):\n        try:\n            img.seek(i)\n            # Write directly to buffer\n            vol[i] = np.array(img, dtype=np.uint8)\n        except EOFError:\n            break\n            \n    return vol\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:11:32.535798Z","iopub.execute_input":"2025-12-28T07:11:32.536015Z","iopub.status.idle":"2025-12-28T07:11:32.548749Z","shell.execute_reply.started":"2025-12-28T07:11:32.535995Z","shell.execute_reply":"2025-12-28T07:11:32.548005Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport matplotlib.pyplot as plt\nfrom PIL import Image\n\nimg = np.load(\"/kaggle/input/vesuvius-signed-distance-fiedl-npz/npy-dt-labels-20251227T160433Z-1-002/npy-dt-labels/1354910392.npy\")\norig = load_tif_volume(\"/kaggle/input/vesuvius-challenge-surface-detection/train_labels/1354910392.tif\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:11:32.549412Z","iopub.execute_input":"2025-12-28T07:11:32.549666Z","iopub.status.idle":"2025-12-28T07:11:33.605021Z","shell.execute_reply.started":"2025-12-28T07:11:32.549646Z","shell.execute_reply":"2025-12-28T07:11:33.604457Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.imshow(img[160])\nplt.show()\nplt.imshow(orig[160])\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:11:33.605757Z","iopub.execute_input":"2025-12-28T07:11:33.606018Z","iopub.status.idle":"2025-12-28T07:11:33.934052Z","shell.execute_reply.started":"2025-12-28T07:11:33.606000Z","shell.execute_reply":"2025-12-28T07:11:33.933330Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Data Loader","metadata":{}},{"cell_type":"code","source":"input_shape=(160, 160, 160)\nbatch_size=1 * total_device\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:11:33.934802Z","iopub.execute_input":"2025-12-28T07:11:33.935062Z","iopub.status.idle":"2025-12-28T07:11:33.938636Z","shell.execute_reply.started":"2025-12-28T07:11:33.935044Z","shell.execute_reply":"2025-12-28T07:11:33.937954Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install scikit-image","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:11:33.939282Z","iopub.execute_input":"2025-12-28T07:11:33.939508Z","iopub.status.idle":"2025-12-28T07:11:37.137296Z","shell.execute_reply.started":"2025-12-28T07:11:33.939492Z","shell.execute_reply":"2025-12-28T07:11:37.136560Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import glob\n\nbase = \"/kaggle/input/vesuvius-signed-distance-fiedl-npz\"\npattern = f\"{base}/npy-dt-labels-*/npy-dt-labels/*.npy\"\n\nlabels = sorted(glob.glob(pattern))\n\nprint(len(labels))\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:11:37.138429Z","iopub.execute_input":"2025-12-28T07:11:37.138704Z","iopub.status.idle":"2025-12-28T07:11:37.650757Z","shell.execute_reply.started":"2025-12-28T07:11:37.138682Z","shell.execute_reply":"2025-12-28T07:11:37.650067Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from pathlib import Path\n\nLABEL_BASE = Path(\"/kaggle/input/vesuvius-signed-distance-fiedl-npz\")\nIMAGE_BASE = Path(\"/kaggle/input/vesuvius-npy/train_images\")\nMASK_BASE = Path(\"/kaggle/input/vesuvius-npy/train_labels\")\n\nlabel_paths = sorted(\n    LABEL_BASE.glob(\"npy-dt-labels-*/npy-dt-labels/*.npy\")\n)\n\ndef label_to_image_path(label_path: Path) -> Path:\n    return IMAGE_BASE/label_path.name\ndef label_to_mask_path(label_path: Path) -> Path:\n    return MASK_BASE/label_path.name\n\n\nimage_paths = [label_to_image_path(p) for p in label_paths]\nmask_paths =  [label_to_mask_path(p) for p in label_paths]\n\n# Optional safety check\nfor img_p in image_paths[:10]:\n    assert img_p.exists(), f\"Missing image: {img_p}\"\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:11:37.654255Z","iopub.execute_input":"2025-12-28T07:11:37.654471Z","iopub.status.idle":"2025-12-28T07:11:37.737660Z","shell.execute_reply.started":"2025-12-28T07:11:37.654454Z","shell.execute_reply":"2025-12-28T07:11:37.736929Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"label_paths = [str(p) for p in label_paths]\nimage_paths = [str(p) for p in image_paths]\nmask_paths = [str(p) for p in mask_paths]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:11:37.738532Z","iopub.execute_input":"2025-12-28T07:11:37.739277Z","iopub.status.idle":"2025-12-28T07:11:37.745809Z","shell.execute_reply.started":"2025-12-28T07:11:37.739258Z","shell.execute_reply":"2025-12-28T07:11:37.745073Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Preprocessing and Augmentation**","metadata":{"execution":{"iopub.status.busy":"2025-12-02T18:19:07.08048Z","iopub.execute_input":"2025-12-02T18:19:07.080625Z","iopub.status.idle":"2025-12-02T18:19:07.092173Z","shell.execute_reply.started":"2025-12-02T18:19:07.080612Z","shell.execute_reply":"2025-12-02T18:19:07.09145Z"}}},{"cell_type":"code","source":"import tensorflow as tf\n\ndef random_occlusions_3d_tf(volume,\n                            occ_prob=1,\n                            max_blocks=6,\n                            min_size=2,\n                            max_size=8):\n    \"\"\"\n    volume: [D, H, W, C] float tensor\n    Returns volume with random cuboid occlusions (set to 0) with prob occ_prob.\n    \"\"\"\n\n    def no_aug():\n        return volume\n\n    def do_aug():\n        v = volume\n        shape = tf.shape(v)\n        D = shape[0]\n        H = shape[1]\n        W = shape[2]\n        C = shape[3]\n\n        # Start with all-ones occlusion mask [D, H, W, 1]\n        occlusion_mask = tf.ones([D, H, W, 1], dtype=v.dtype)\n\n        # Precompute ranges once\n        d_range = tf.range(D)\n        h_range = tf.range(H)\n        w_range = tf.range(W)\n\n        # Loop over blocks in TF graph\n        def body(i, mask):\n            # random block size\n            block_d = tf.random.uniform([], min_size, max_size, dtype=tf.int32)\n            block_h = tf.random.uniform([], min_size, max_size, dtype=tf.int32)\n            block_w = tf.random.uniform([], min_size, max_size, dtype=tf.int32)\n\n            # random start coords\n            d0 = tf.random.uniform([], 0, tf.maximum(D - block_d, 1), dtype=tf.int32)\n            h0 = tf.random.uniform([], 0, tf.maximum(H - block_h, 1), dtype=tf.int32)\n            w0 = tf.random.uniform([], 0, tf.maximum(W - block_w, 1), dtype=tf.int32)\n\n            d1 = tf.minimum(d0 + block_d, D)\n            h1 = tf.minimum(h0 + block_h, H)\n            w1 = tf.minimum(w0 + block_w, W)\n\n            # Boolean selection per dimension\n            d_sel = tf.logical_and(d_range >= d0, d_range < d1)   # [D]\n            h_sel = tf.logical_and(h_range >= h0, h_range < h1)   # [H]\n            w_sel = tf.logical_and(w_range >= w0, w_range < w1)   # [W]\n\n            # Broadcast to [D, H, W, 1]\n            d_sel = tf.reshape(d_sel, [D, 1, 1, 1])\n            h_sel = tf.reshape(h_sel, [1, H, 1, 1])\n            w_sel = tf.reshape(w_sel, [1, 1, W, 1])\n\n            block_mask = tf.cast(d_sel & h_sel & w_sel, v.dtype)  # 1 inside block\n\n            # Zero inside the block: mask *= (1 - block_mask)\n            new_mask = mask * (1.0 - block_mask)\n            return i + 1, new_mask\n\n        def cond(i, mask):\n            return i < max_blocks\n\n        _, occlusion_mask = tf.while_loop(\n            cond,\n            body,\n            loop_vars=[tf.constant(0, dtype=tf.int32), occlusion_mask],\n            shape_invariants=[\n                tf.TensorShape([]),\n                tf.TensorShape([None, None, None, 1])\n            ]\n        )\n\n        # Apply mask to all channels\n        v = v * occlusion_mask  # broadcasts over C\n        return v\n\n    return tf.cond(\n        tf.random.uniform([]) < occ_prob,\n        do_aug,\n        no_aug\n    )\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:11:37.746584Z","iopub.execute_input":"2025-12-28T07:11:37.747461Z","iopub.status.idle":"2025-12-28T07:11:37.761684Z","shell.execute_reply.started":"2025-12-28T07:11:37.747444Z","shell.execute_reply":"2025-12-28T07:11:37.760978Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_transformation(image, label):\n    data = {\"image\": image, \"label\": label}\n    pipeline = Compose([\n        RandSpatialCrop(\n            keys=[\"image\", \"label\"],\n            roi_size=input_shape,\n        ),\n        RandFlip(keys=[\"image\", \"label\"], spatial_axis=[0], prob=0.5),\n        RandFlip(keys=[\"image\", \"label\"], spatial_axis=[1], prob=0.5),\n        RandFlip(keys=[\"image\", \"label\"], spatial_axis=[2], prob=0.5),\n        RandRotate90(\n            keys=[\"image\", \"label\"], \n            prob=0.4, \n            max_k=3, \n            spatial_axes=(0, 1)\n        ),\n        RandShiftIntensity(keys=[\"image\"], offsets=0.15, prob=0.5),\n        ScaleIntensityRange(\n            keys=[\"image\"],\n            a_min = 0,\n            a_max = 255,\n            b_min = 0,\n            b_max = 1,\n            clip = True,\n        ),\n    ])\n    result = pipeline(data)\n\n    result[\"image\"] = random_occlusions_3d_tf(result[\"image\"])\n\n    \n\n    return result[\"image\"], result[\"label\"]\n\n\ndef val_transformation(image, label):\n    data = {\"image\": image, \"label\": label}\n    pipeline = Compose([\n        ScaleIntensityRange(\n            keys=[\"image\"],\n            a_min = 0,\n            a_max = 255,\n            b_min = 0,\n            b_max = 1,\n            clip = True,\n        ),\n    ])\n    result = pipeline(data)\n    return result[\"image\"], result[\"label\"]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:11:37.762553Z","iopub.execute_input":"2025-12-28T07:11:37.762892Z","iopub.status.idle":"2025-12-28T07:11:37.779195Z","shell.execute_reply.started":"2025-12-28T07:11:37.762867Z","shell.execute_reply":"2025-12-28T07:11:37.778285Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_npy_tf(path):\n    def _load(p):\n        arr = np.load(p.decode(\"utf-8\"))\n        return arr.astype(np.float32)\n    return tf.numpy_function(_load, [path], tf.float32)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:11:37.779932Z","iopub.execute_input":"2025-12-28T07:11:37.780157Z","iopub.status.idle":"2025-12-28T07:11:37.794118Z","shell.execute_reply.started":"2025-12-28T07:11:37.780141Z","shell.execute_reply":"2025-12-28T07:11:37.793365Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def parse_npy_pair(image_path, sdf_path, seg_path):\n    image = load_npy_tf(image_path)\n    sdf   = load_npy_tf(sdf_path)\n    seg   = load_npy_tf(seg_path)\n\n    image.set_shape([None, None, None])\n    sdf.set_shape([None, None, None])\n    seg.set_shape([None, None, None])\n\n    # valid region\n    valid = tf.cast(seg != 2, tf.float32)\n\n    # 1️⃣ mask input\n    image = image * valid\n\n    # 2️⃣ encode invalid region in sdf\n    INVALID = 69.0\n    sdf = tf.where(valid > 0, sdf, INVALID)\n\n    image = tf.expand_dims(image, -1)\n    sdf   = tf.expand_dims(sdf, -1)\n\n    return image, sdf\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:11:37.794856Z","iopub.execute_input":"2025-12-28T07:11:37.795159Z","iopub.status.idle":"2025-12-28T07:11:37.807459Z","shell.execute_reply.started":"2025-12-28T07:11:37.795137Z","shell.execute_reply":"2025-12-28T07:11:37.806746Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def npy_loader(image_paths, label_paths, mask_paths , batch_size=1, shuffle=True):\n    dataset = tf.data.Dataset.from_tensor_slices(\n        (image_paths, label_paths, mask_paths)\n    )\n\n    dataset = dataset.shuffle(buffer_size=100) if shuffle else dataset\n\n    dataset = dataset.map(\n        parse_npy_pair,\n        num_parallel_calls=tf.data.AUTOTUNE\n    )\n\n    if shuffle:\n        # TRAINING PATH\n        dataset = dataset.map(\n            train_transformation,\n            num_parallel_calls=tf.data.AUTOTUNE\n        )\n    else:\n        # VALIDATION PATH\n        dataset = dataset.map(\n            val_transformation,\n            num_parallel_calls=tf.data.AUTOTUNE\n        )\n\n    dataset = dataset.batch(\n        batch_size,\n        drop_remainder=True   # REQUIRED for TPU\n    ).prefetch(tf.data.AUTOTUNE)\n\n    return dataset\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:11:37.808214Z","iopub.execute_input":"2025-12-28T07:11:37.808542Z","iopub.status.idle":"2025-12-28T07:11:37.824060Z","shell.execute_reply.started":"2025-12-28T07:11:37.808520Z","shell.execute_reply":"2025-12-28T07:11:37.823170Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"val_image_paths = image_paths[780:]\nval_label_paths = label_paths[780:]\nval_mask_paths = mask_paths[780:]\n\n\nimage_paths = image_paths[:780]\nlabel_paths = label_paths[:780]\nmask_paths = mask_paths[:780]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:11:37.824656Z","iopub.execute_input":"2025-12-28T07:11:37.824863Z","iopub.status.idle":"2025-12-28T07:11:37.837201Z","shell.execute_reply.started":"2025-12-28T07:11:37.824849Z","shell.execute_reply":"2025-12-28T07:11:37.836606Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_loader = npy_loader(\n    image_paths,\n    label_paths,\n    mask_paths,\n    batch_size=batch_size,\n    shuffle=True\n)\n\nval_loader = npy_loader(\n    val_image_paths,\n    val_label_paths,\n    val_mask_paths,\n    batch_size=1,\n    shuffle=False\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:11:37.838388Z","iopub.execute_input":"2025-12-28T07:11:37.839035Z","iopub.status.idle":"2025-12-28T07:11:42.119248Z","shell.execute_reply.started":"2025-12-28T07:11:37.839019Z","shell.execute_reply":"2025-12-28T07:11:42.118629Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"model = TransUNet(\n        input_shape=(160, 160, 160, 1),\n        encoder_name=\"seresnext50\",\n        classifier_activation=None,\n        num_classes=1,\n    )\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:11:42.120048Z","iopub.execute_input":"2025-12-28T07:11:42.120379Z","iopub.status.idle":"2025-12-28T07:11:56.197503Z","shell.execute_reply.started":"2025-12-28T07:11:42.120354Z","shell.execute_reply":"2025-12-28T07:11:56.196901Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ALERT: This attributes only available in medicai (not in core keras)\nmodel.instance_describe()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:11:56.198228Z","iopub.execute_input":"2025-12-28T07:11:56.198434Z","iopub.status.idle":"2025-12-28T07:11:56.256868Z","shell.execute_reply.started":"2025-12-28T07:11:56.198419Z","shell.execute_reply":"2025-12-28T07:11:56.256242Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.load_weights(\"/kaggle/input/sdf-transunet/keras/default/1/fine_tuning_epoch_100.weights.h5\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:11:56.257470Z","iopub.execute_input":"2025-12-28T07:11:56.258254Z","iopub.status.idle":"2025-12-28T07:12:01.146485Z","shell.execute_reply.started":"2025-12-28T07:11:56.258236Z","shell.execute_reply":"2025-12-28T07:12:01.145732Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import keras\nfrom keras import ops\n\n\nclass SignedDistanceWithEikonalLoss(keras.losses.Loss):\n    def __init__(\n        self,\n        tau=20.0,\n        delta=2.0,\n        eikonal_weight=0.05,\n        eikonal_band=3.0,\n        sentinel=69.0,\n        name=\"sdf_eikonal_loss\",\n    ):\n        super().__init__(name=name)\n        self.tau = tau\n        self.delta = delta\n        self.eikonal_weight = eikonal_weight\n        self.eikonal_band = eikonal_band\n        self.sentinel = sentinel\n\n    def call(self, y_true, y_pred):\n        # ---------------------------\n        # 1. VALID MASK\n        # ---------------------------\n        valid_mask = ops.cast(y_true != self.sentinel, \"float32\")\n\n        # ---------------------------\n        # 2. TSDF CLIP\n        # ---------------------------\n        y_true = ops.clip(y_true, -self.tau, self.tau)\n        y_pred = ops.clip(y_pred, -self.tau, self.tau)\n\n        # ---------------------------\n        # 3. HUBER SDF LOSS (masked)\n        # ---------------------------\n        error = y_pred - y_true\n        abs_error = ops.abs(error)\n\n        huber = ops.where(\n            abs_error <= self.delta,\n            0.5 * ops.square(error),\n            self.delta * (abs_error - 0.5 * self.delta),\n        )\n\n        sdf_loss = ops.sum(huber * valid_mask) / (\n            ops.sum(valid_mask) + 1e-6\n        )\n\n        # ---------------------------\n        # 4. EIKONAL LOSS (FIXED)\n        # ---------------------------\n\n        # finite differences\n        dx = y_pred[:, 1:, :-1, :-1, :] - y_pred[:, :-1, :-1, :-1, :]\n        dy = y_pred[:, :-1, 1:, :-1, :] - y_pred[:, :-1, :-1, :-1, :]\n        dz = y_pred[:, :-1, :-1, 1:, :] - y_pred[:, :-1, :-1, :-1, :]\n\n        grad_norm = ops.sqrt(dx**2 + dy**2 + dz**2 + 1e-6)\n\n        # band + validity mask (aligned!)\n        band_mask = ops.cast(\n            (ops.abs(y_true[:, :-1, :-1, :-1, :]) <= self.eikonal_band)\n            & (y_true[:, :-1, :-1, :-1, :] != self.sentinel),\n            \"float32\",\n        )\n\n        eikonal_loss = ops.sum(\n            ops.square(grad_norm - 1.0) * band_mask\n        ) / (ops.sum(band_mask) + 1e-6)\n\n        # ---------------------------\n        # 5. FINAL LOSS\n        # ---------------------------\n        return sdf_loss + self.eikonal_weight * eikonal_loss\n\n\n\n\nclass SurfaceMAEMonitor(keras.metrics.Metric):\n    def __init__(\n        self,\n        tau=1.0,\n        invalid_value=69.0,\n        name=\"surface_mae\",\n        **kwargs\n    ):\n        super().__init__(name=name, **kwargs)\n        self.tau = tau\n        self.invalid = invalid_value\n\n        self.total = self.add_variable(shape=(), initializer=\"zeros\")\n        self.count = self.add_variable(shape=(), initializer=\"zeros\")\n\n    def update_state(self, y_true, y_pred, sample_weight=None):\n        y_true = y_true[..., 0]\n        y_pred = y_pred[..., 0]\n\n        valid = ops.cast(y_true != self.invalid, \"float32\")\n        surface = ops.cast(ops.abs(y_true) <= self.tau, \"float32\")\n\n        mask = valid * surface\n\n        err = ops.abs(y_pred - y_true) * mask\n        val = ops.sum(err) / (ops.sum(mask) + 1e-6)\n\n        self.total.assign_add(val)\n        self.count.assign_add(1.0)\n\n    def result(self):\n        return -self.total / (self.count + 1e-6)\n\n    def reset_state(self):\n        self.total.assign(0.0)\n        self.count.assign(0.0)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:22:17.262869Z","iopub.execute_input":"2025-12-28T07:22:17.263507Z","iopub.status.idle":"2025-12-28T07:22:17.277571Z","shell.execute_reply.started":"2025-12-28T07:22:17.263483Z","shell.execute_reply":"2025-12-28T07:22:17.276915Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_samples = 780\nepochs = 100\ntotal_steps = (num_samples // batch_size) * epochs\nwarmup_steps = (num_samples // batch_size) * 5\n\n\nlr_schedule = keras.optimizers.schedules.CosineDecay(\n    initial_learning_rate=1e-4,\n    decay_steps = total_steps,\n    alpha=0.05,\n)\n\n\noptim = keras.optimizers.AdamW(\n    learning_rate=lr_schedule,\n    weight_decay=1e-5,\n)\n\nloss_fn = SignedDistanceWithEikonalLoss(\n    tau=20.0,\n    delta=2.0,\n    eikonal_weight=0.05,\n    eikonal_band=3.0,\n    sentinel=69.0,\n)\n\nmetrics = [\n    SurfaceMAEMonitor(tau=1.0),\n]\n\nswi_callback_metric = SurfaceMAEMonitor(\n    tau=1.0,\n    name=\"val_surface_mae\",\n)\n\n\nswi_callback = SlidingWindowInferenceCallback(\n    model,\n    num_classes = 1,\n    dataset=val_loader,\n    metrics=swi_callback_metric,\n    interval=10,\n    overlap=0.5,\n    roi_size=input_shape,\n    sw_batch_size=1* total_device,\n    save_path=\"model.weights.h5\",\n)\n\n\n\n\nclass PeriodicWeightsSaver(keras.callbacks.Callback):\n    def __init__(self, interval=25, save_path_template=\"checkpoint_epoch_{epoch}.weights.h5\"):\n        super().__init__()\n        self.interval = interval\n        self.save_path_template = save_path_template\n\n    def on_epoch_end(self, epoch, logs=None):\n        # epoch is 0-indexed, so we check (epoch + 1)\n        if (epoch + 1) % self.interval == 0:\n            save_path = self.save_path_template.format(epoch=epoch + 1)\n            self.model.save_weights(save_path)\n            print(f\"\\n[Snapshot] Saved periodic weights to: {save_path}\")\n\n\n\nsnapshot_cb = PeriodicWeightsSaver(\n    interval=20,\n    save_path_template=\"fine_tuning_epoch_{epoch}.weights.h5\",\n)\n\nmodel.compile(\n    optimizer=optim,  # your AdamW + cosine schedule\n    loss=loss_fn,\n    metrics=[SurfaceMAEMonitor(tau=1.0)],\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:22:17.466785Z","iopub.execute_input":"2025-12-28T07:22:17.467027Z","iopub.status.idle":"2025-12-28T07:22:17.484228Z","shell.execute_reply.started":"2025-12-28T07:22:17.467010Z","shell.execute_reply":"2025-12-28T07:22:17.483645Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.fit(\n    train_loader,\n    epochs=epochs,\n    callbacks=[swi_callback, snapshot_cb],\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:22:20.661201Z","iopub.execute_input":"2025-12-28T07:22:20.661766Z","execution_failed":"2025-12-28T07:39:54.843Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Eval","metadata":{}},{"cell_type":"code","source":"# model.load_weights(\n#     \"/kaggle/input/train-vesuvius-surface-3d-detection-on-tpu/fine_tuning_epoch_200.weights.h5\"\n# )\n# swi = SlidingWindowInference(\n#     model,\n#     num_classes=num_classes,\n#     roi_size=input_shape,\n#     sw_batch_size=1 * total_device,\n#     overlap=0.5,\n# )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:12:08.304805Z","iopub.status.idle":"2025-12-28T07:12:08.305166Z","shell.execute_reply.started":"2025-12-28T07:12:08.304986Z","shell.execute_reply":"2025-12-28T07:12:08.305003Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# dice = SkeletonRecallPlusDiceLoss(\n#     num_classes=num_classes,\n# )\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:12:08.306435Z","iopub.status.idle":"2025-12-28T07:12:08.306882Z","shell.execute_reply.started":"2025-12-28T07:12:08.306681Z","shell.execute_reply":"2025-12-28T07:12:08.306700Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# loss = []","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:12:08.307548Z","iopub.status.idle":"2025-12-28T07:12:08.307923Z","shell.execute_reply.started":"2025-12-28T07:12:08.307737Z","shell.execute_reply":"2025-12-28T07:12:08.307754Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# i =0\n# for sample in val_loader:\n#     x, y = sample\n#     output = swi(x)\n#     y = ops.convert_to_tensor(y)\n#     output = ops.convert_to_tensor(output)\n#     loss.append(dice.call(y, output))\n    \n    \n    \n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:12:08.309324Z","iopub.status.idle":"2025-12-28T07:12:08.309667Z","shell.execute_reply.started":"2025-12-28T07:12:08.309488Z","shell.execute_reply":"2025-12-28T07:12:08.309502Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# dice_score = np.mean((ops.convert_to_numpy(loss)))\n# print(f\"Dice Score: {dice_score/6:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:12:08.310561Z","iopub.status.idle":"2025-12-28T07:12:08.310943Z","shell.execute_reply.started":"2025-12-28T07:12:08.310750Z","shell.execute_reply":"2025-12-28T07:12:08.310770Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# x, y = next(iter(val_loader))\n# x.shape, y.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:12:08.311680Z","iopub.status.idle":"2025-12-28T07:12:08.311952Z","shell.execute_reply.started":"2025-12-28T07:12:08.311813Z","shell.execute_reply":"2025-12-28T07:12:08.311827Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# y_pred = swi(x)\n# y_pred.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:12:08.312929Z","iopub.status.idle":"2025-12-28T07:12:08.313421Z","shell.execute_reply.started":"2025-12-28T07:12:08.313234Z","shell.execute_reply":"2025-12-28T07:12:08.313249Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# segment = y_pred.argmax(-1).astype(np.uint8)\n# segment.shape, np.unique(segment)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:12:08.315004Z","iopub.status.idle":"2025-12-28T07:12:08.315294Z","shell.execute_reply.started":"2025-12-28T07:12:08.315135Z","shell.execute_reply":"2025-12-28T07:12:08.315151Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# plot_sample(\n#     x, segment, sample_idx=0, max_slices=4\n# )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T07:12:08.319016Z","iopub.status.idle":"2025-12-28T07:12:08.319283Z","shell.execute_reply.started":"2025-12-28T07:12:08.319174Z","shell.execute_reply":"2025-12-28T07:12:08.319185Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Next Stop**\n\n- [Affinity Feature Strengthening](https://arxiv.org/pdf/2211.06578)\n- [Meta-Tubular-Net: A Robust Topology-Aware Re-Weighting Network](file:///C:/Users/ASUS/Pictures/Screenshots/ssrn-4132287.pdf)\n- [Landmark-Assisted Anatomy-Sensitive](file:///C:/Users/ASUS/Downloads/diagnostics-13-02260-v2.pdf)\n- [LEAD: Self-Supervised Landmark Estimation](https://arxiv.org/pdf/2204.02958)\n- [TopoSeg: Topology-Aware](https://openaccess.thecvf.com/content/ICCV2023/papers/He_TopoSeg_Topology-Aware_Nuclear_Instance_Segmentation_ICCV_2023_paper.pdf)\n- [Virtually Unrolling the Herculaneum Papyri](https://arxiv.org/pdf/2512.04927v1)","metadata":{}}]}