{"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":286271954,"sourceType":"kernelVersion"},{"sourceId":288402244,"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}],"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-24T15:40:22.605687Z","iopub.execute_input":"2025-12-24T15:40:22.605881Z","iopub.status.idle":"2025-12-24T15:40:31.592905Z","shell.execute_reply.started":"2025-12-24T15:40:22.605865Z","shell.execute_reply":"2025-12-24T15:40:31.592240Z"}},"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-24T15:40:31.594804Z","iopub.execute_input":"2025-12-24T15:40:31.595121Z","iopub.status.idle":"2025-12-24T15:40:41.985134Z","shell.execute_reply.started":"2025-12-24T15:40:31.595096Z","shell.execute_reply":"2025-12-24T15:40:41.984430Z"}},"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-24T15:40:41.986141Z","iopub.execute_input":"2025-12-24T15:40:41.986673Z","iopub.status.idle":"2025-12-24T15:40:41.990585Z","shell.execute_reply.started":"2025-12-24T15:40:41.986648Z","shell.execute_reply":"2025-12-24T15:40:41.990024Z"}},"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-24T15:40:41.991399Z","iopub.execute_input":"2025-12-24T15:40:41.991632Z","iopub.status.idle":"2025-12-24T15:41:03.615710Z","shell.execute_reply.started":"2025-12-24T15:40:41.991609Z","shell.execute_reply":"2025-12-24T15:41:03.614895Z"}},"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-24T15:41:03.616558Z","iopub.execute_input":"2025-12-24T15:41:03.617056Z","iopub.status.idle":"2025-12-24T15:41:04.735442Z","shell.execute_reply.started":"2025-12-24T15:41:03.617036Z","shell.execute_reply":"2025-12-24T15:41:04.734824Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"keras.version(), keras.config.backend(), medicai.version()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T15:41:04.736060Z","iopub.execute_input":"2025-12-24T15:41:04.736292Z","iopub.status.idle":"2025-12-24T15:41:04.742532Z","shell.execute_reply.started":"2025-12-24T15:41:04.736264Z","shell.execute_reply":"2025-12-24T15:41:04.741808Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"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\nnum_classes=3","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T15:41:04.744814Z","iopub.execute_input":"2025-12-24T15:41:04.745411Z","iopub.status.idle":"2025-12-24T15:41:04.757743Z","shell.execute_reply.started":"2025-12-24T15:41:04.745389Z","shell.execute_reply":"2025-12-24T15:41:04.757163Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install scikit-image","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T15:41:04.758500Z","iopub.execute_input":"2025-12-24T15:41:04.758755Z","iopub.status.idle":"2025-12-24T15:41:07.990222Z","shell.execute_reply.started":"2025-12-24T15:41:04.758738Z","shell.execute_reply":"2025-12-24T15:41:07.989460Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from skimage.morphology import skeletonize\nfrom scipy.ndimage import binary_dilation\n\ndef generate_tubed_skeleton_numpy(label_vol):\n    # label_vol shape: (D, H, W, 1)\n    # Extract binary mask for the class of interest (assuming class 1 is ink)\n    mask = (label_vol[..., 0] == 1)\n    \n    # 1. Skeletonize\n    skel = skeletonize(mask)\n    \n    # 2. Tubular Dilation (Tubed Skeleton)\n    # Iterations=1 is usually sufficient for \"thin\" tubes; increase for thicker targets\n    tubed_skel = binary_dilation(skel, iterations=1)\n    \n    # Return as float32 for loss calculation, keeping shape (D, H, W, 1)\n    return tubed_skel.astype(np.float32)[..., None]\n\ndef add_skeleton_target(image, label):\n    # Wrapper to run numpy code inside tf.data graph\n    # Inputs: image (D,H,W,1), label (D,H,W,1)\n    \n    tubed_skel = tf.numpy_function(\n        func=generate_tubed_skeleton_numpy,\n        inp=[label],\n        Tout=tf.float32\n    )\n    \n    # Explicitly set shape because numpy_function loses it\n    tubed_skel.set_shape(label.shape)\n    \n    # Pack both targets into y_true: Channel 0 = Mask, Channel 1 = Skeleton\n    # New label shape: (D, H, W, 2)\n    combined_label = tf.concat([tf.cast(label, tf.float32), tubed_skel], axis=-1)\n    \n    return image, combined_label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T15:41:07.991250Z","iopub.execute_input":"2025-12-24T15:41:07.991527Z","iopub.status.idle":"2025-12-24T15:41:08.178303Z","shell.execute_reply.started":"2025-12-24T15:41:07.991503Z","shell.execute_reply":"2025-12-24T15:41:08.177730Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def parse_tfrecord_fn(example):\n    feature_description = {\n        \"image\": tf.io.FixedLenFeature([], tf.string),\n        \"label\": tf.io.FixedLenFeature([], tf.string),\n        \"image_shape\": tf.io.FixedLenFeature([3], tf.int64),\n        \"label_shape\": tf.io.FixedLenFeature([3], tf.int64),\n    }\n    parsed_example = tf.io.parse_single_example(example, feature_description)\n    image = tf.io.decode_raw(parsed_example[\"image\"], tf.uint8)\n    label = tf.io.decode_raw(parsed_example[\"label\"], tf.uint8)\n    image_shape = tf.cast(parsed_example[\"image_shape\"], tf.int64)\n    label_shape = tf.cast(parsed_example[\"label_shape\"], tf.int64)\n    image = tf.reshape(image, image_shape)\n    label = tf.reshape(label, label_shape)\n    return image, label","metadata":{"trusted":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2025-12-24T15:41:08.179110Z","iopub.execute_input":"2025-12-24T15:41:08.180220Z","iopub.status.idle":"2025-12-24T15:41:08.184778Z","shell.execute_reply.started":"2025-12-24T15:41:08.180201Z","shell.execute_reply":"2025-12-24T15:41:08.184172Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def prepare_inputs(image, label):\n    # Add channel dimension\n    image = image[..., None] # (D, H, W, 1)\n    label = label[..., None] # (D, H, W, 1)\n\n    # Convert to float32\n    image = tf.cast(image, tf.float32)\n    label = tf.cast(label, tf.float32)\n    return image, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T15:41:08.185573Z","iopub.execute_input":"2025-12-24T15:41:08.185819Z","iopub.status.idle":"2025-12-24T15:41:08.208956Z","shell.execute_reply.started":"2025-12-24T15:41:08.185795Z","shell.execute_reply":"2025-12-24T15:41:08.208457Z"}},"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-24T15:41:08.209686Z","iopub.execute_input":"2025-12-24T15:41:08.209941Z","iopub.status.idle":"2025-12-24T15:41:08.228486Z","shell.execute_reply.started":"2025-12-24T15:41:08.209917Z","shell.execute_reply":"2025-12-24T15:41:08.228004Z"}},"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-24T15:41:08.229117Z","iopub.execute_input":"2025-12-24T15:41:08.229334Z","iopub.status.idle":"2025-12-24T15:41:08.250466Z","shell.execute_reply.started":"2025-12-24T15:41:08.229319Z","shell.execute_reply":"2025-12-24T15:41:08.249967Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def tfrecord_loader(tfrecord_pattern, batch_size=1, shuffle=True):\n    dataset = tf.data.TFRecordDataset(\n        tf.io.gfile.glob(tfrecord_pattern)\n    )\n    dataset = dataset.shuffle(buffer_size=100) if shuffle else dataset \n    dataset = dataset.map(\n        parse_tfrecord_fn, num_parallel_calls=tf.data.AUTOTUNE\n    )\n    dataset = dataset.map(\n        prepare_inputs,\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        # ONLY add skeleton for training\n        dataset = dataset.map(\n            add_skeleton_target, \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        # Skeleton target skipped here!\n        \n    dataset = dataset.batch(batch_size, drop_remainder=True).prefetch(tf.data.AUTOTUNE)\n    return dataset","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T15:41:08.251096Z","iopub.execute_input":"2025-12-24T15:41:08.251356Z","iopub.status.idle":"2025-12-24T15:41:08.271959Z","shell.execute_reply.started":"2025-12-24T15:41:08.251341Z","shell.execute_reply":"2025-12-24T15:41:08.271306Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Generate the exact list of filenames for 0 to 129\nbase_path = \"/kaggle/input/vesuvius-tfrecords/training_shard_{}.tfrec\"\ntrain_files = [base_path.format(i) for i in range(130)]  # Generates 0, 1, ... 129\n\n# Validation file\nval_files = [base_path.format(130)]\n\n# Pass the LIST directly (glob handles lists of paths correctly)\ntrain_loader = tfrecord_loader(\n    train_files, batch_size=batch_size, shuffle=True\n)\n\nval_loader = tfrecord_loader(\n    val_files, batch_size=1, shuffle=False\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T15:41:08.272755Z","iopub.execute_input":"2025-12-24T15:41:08.272995Z","iopub.status.idle":"2025-12-24T15:41:14.070571Z","shell.execute_reply.started":"2025-12-24T15:41:08.272958Z","shell.execute_reply":"2025-12-24T15:41:14.070018Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"x, y = next(iter(train_loader))\nx.shape, y.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T15:41:14.071318Z","iopub.execute_input":"2025-12-24T15:41:14.071550Z","iopub.status.idle":"2025-12-24T15:42:27.144154Z","shell.execute_reply.started":"2025-12-24T15:41:14.071524Z","shell.execute_reply":"2025-12-24T15:42:27.143480Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Viz**","metadata":{}},{"cell_type":"code","source":"def plot_sample(x, y, sample_idx=0, max_slices=16):\n    img = np.squeeze(x[sample_idx])  # (D, H, W)\n    mask = np.squeeze(y[sample_idx])  # (D, H, W)\n    D = img.shape[0]\n\n    # Decide which slices to plot\n    step = max(1, D // max_slices)\n    slices = range(0, D, step)\n\n    n_slices = len(slices)\n    fig, axes = plt.subplots(2, n_slices, figsize=(3*n_slices, 6))\n\n    for i, s in enumerate(slices):\n        axes[0, i].imshow(img[s], cmap='gray')\n        axes[0, i].set_title(f\"Slice {s}\")\n        axes[0, i].axis('off')\n\n        axes[1, i].imshow(mask[s], cmap='gray')\n        axes[1, i].set_title(f\"Mask {s}\")\n        axes[1, i].axis('off')\n\n    plt.suptitle(f\"Sample {sample_idx}\")\n    plt.tight_layout()\n    plt.show()\n","metadata":{"trusted":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2025-12-24T15:42:27.145009Z","iopub.execute_input":"2025-12-24T15:42:27.145242Z","iopub.status.idle":"2025-12-24T15:42:27.151319Z","shell.execute_reply.started":"2025-12-24T15:42:27.145224Z","shell.execute_reply":"2025-12-24T15:42:27.150550Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_planes(image, mask, alpha=0.4):\n    # Central slices\n    d, h, w = image.shape\n    axial_img    = image[d // 2]\n    coronal_img  = image[:, h // 2, :]\n    sagittal_img = image[:, :, w // 2]\n\n    axial_msk    = mask[d // 2]\n    coronal_msk  = mask[:, h // 2, :]\n    sagittal_msk = mask[:, :, w // 2]\n\n    slices_img = [axial_img, coronal_img, sagittal_img]\n    slices_msk = [axial_msk, coronal_msk, sagittal_msk]\n    \n    titles = [\"Axial (XY plane)\", \"Coronal (XZ plane)\", \"Sagittal (YZ plane)\"]\n    fig, axes = plt.subplots(1, 3, figsize=(15, 5))\n\n    for i, ax in enumerate(axes):\n        ax.imshow(slices_img[i], cmap=\"gray\")\n\n        # overlay jet only where mask > 0\n        m = slices_msk[i]\n        if m.max() > 0:\n            ax.imshow(m, cmap=\"jet\", alpha=alpha)\n\n        ax.set_title(titles[i])\n        ax.axis(\"off\")\n\n    plt.tight_layout()\n    plt.show()\n","metadata":{"trusted":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2025-12-24T15:42:27.152110Z","iopub.execute_input":"2025-12-24T15:42:27.152392Z","iopub.status.idle":"2025-12-24T15:42:27.174222Z","shell.execute_reply.started":"2025-12-24T15:42:27.152376Z","shell.execute_reply":"2025-12-24T15:42:27.173525Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"x, y = next(iter(val_loader))\nx.shape, y.shape ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T15:42:27.174921Z","iopub.execute_input":"2025-12-24T15:42:27.175260Z","iopub.status.idle":"2025-12-24T15:42:30.138111Z","shell.execute_reply.started":"2025-12-24T15:42:27.175242Z","shell.execute_reply":"2025-12-24T15:42:30.137597Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_sample(\n    x, y[:,:,:,:,0], sample_idx=0, max_slices=4\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T15:42:30.138679Z","iopub.execute_input":"2025-12-24T15:42:30.138856Z","iopub.status.idle":"2025-12-24T15:42:31.090558Z","shell.execute_reply.started":"2025-12-24T15:42:30.138841Z","shell.execute_reply":"2025-12-24T15:42:31.089862Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_planes(\n    np.squeeze(x[0]), # picking one sample\n    np.squeeze(y[0,:,:,:,0])  # picking one sample\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T15:42:31.091539Z","iopub.execute_input":"2025-12-24T15:42:31.091723Z","iopub.status.idle":"2025-12-24T15:42:31.970138Z","shell.execute_reply.started":"2025-12-24T15:42:31.091703Z","shell.execute_reply":"2025-12-24T15:42:31.969273Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"## check available models (classification + segmentation)\nmedicai.models.list_models()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T15:42:31.975283Z","iopub.execute_input":"2025-12-24T15:42:31.975764Z","iopub.status.idle":"2025-12-24T15:42:31.992939Z","shell.execute_reply.started":"2025-12-24T15:42:31.975745Z","shell.execute_reply":"2025-12-24T15:42:31.992408Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = TransUNet(\n        input_shape=(160, 160, 160, 1),\n        encoder_name=\"seresnext50\",\n        classifier_activation=\"softmax\",\n        num_classes=3,\n    )\n\nmodel.load_weights(\"/kaggle/input/train-vesuvius-surface-3d-detection-on-tpu/model.weights.h5\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T15:42:31.993715Z","iopub.execute_input":"2025-12-24T15:42:31.993916Z","execution_failed":"2025-12-24T15:42:39.284Z"}},"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":{"execution_failed":"2025-12-24T15:42:39.284Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import keras\nfrom keras import ops\nfrom medicai.losses import SparseDiceCELoss\n\n# --- Monitor for Skeleton Loss ---\nclass SkeletonLossMonitor(keras.metrics.Metric):\n    def __init__(self, name=\"skel_loss\", **kwargs):\n        super().__init__(name=name, **kwargs)\n        self.total = self.add_variable(shape=(), initializer=\"zeros\", name=\"total\")\n        self.count = self.add_variable(shape=(), initializer=\"zeros\", name=\"count\")\n\n    def update_state(self, y_true, y_pred, sample_weight=None):\n        # Only calc if we have the skeleton channel (Training data)\n        if y_true.shape[-1] == 2:\n            y_true_skel = y_true[..., 1]\n            pred_prob = y_pred[..., 1]\n            \n            # Re-calculate the specific term we want to watch\n            intersection = keras.ops.sum(pred_prob * y_true_skel, axis=(1, 2, 3))\n            skeleton_sum = keras.ops.sum(y_true_skel, axis=(1, 2, 3))\n            \n            has_skeleton = keras.ops.cast(skeleton_sum > 0, dtype=\"float32\")\n            recall = (intersection + 1e-6) / (skeleton_sum + 1e-6)\n            \n            # We want to see the loss value\n            val = (1.0 - recall) * has_skeleton\n            \n            # Average over batch\n            self.total.assign_add(keras.ops.mean(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\n# --- Monitor for Base (Dice) Loss ---\nclass BaseLossMonitor(keras.metrics.Metric):\n    def __init__(self, num_classes, name=\"base_loss\", **kwargs):\n        super().__init__(name=name, **kwargs)\n        self.base_fn = SparseDiceCELoss(from_logits=False, num_classes=num_classes, ignore_class_ids=2)\n        self.total = self.add_variable(shape=(), initializer=\"zeros\", name=\"total\")\n        self.count = self.add_variable(shape=(), initializer=\"zeros\", name=\"count\")\n\n    def update_state(self, y_true, y_pred, sample_weight=None):\n        # Handle shape: if 2 channels, take only the first (mask)\n        if len(y_true.shape) == 5 and y_true.shape[-1] == 2:\n            y_true_mask = y_true[..., 0:1]\n        else:\n            y_true_mask = y_true \n            \n        val = self.base_fn(y_true_mask, y_pred)\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\nclass SkeletonRecallPlusDiceLoss(keras.losses.Loss):\n    def __init__(\n        self,\n        num_classes,\n        w_srec=0.75,\n        w_fp=0.5,          # <-- NEW\n        name=\"skel_recall_fp_loss\",\n    ):\n        super().__init__(name=name)\n        self.num_classes = num_classes\n        self.w_srec = w_srec\n        self.w_fp = w_fp  # <-- NEW\n\n        self.base_loss_fn = SparseDiceCELoss(\n            from_logits=False,\n            num_classes=num_classes,\n            ignore_class_ids=2,\n        )\n\n    def call(self, y_true, y_pred):\n        # --------------------\n        # GT unpacking\n        # --------------------\n        y_true_mask = y_true[..., 0]   # 0/1/2 (2 = ignore)\n        y_true_skel = y_true[..., 1]\n\n        pred_ink_prob = y_pred[..., 1]\n\n        # Valid (non-ignore) mask\n        valid_mask = keras.ops.cast(y_true_mask != 2, \"float32\")\n\n        # --------------------\n        # 1. Base Dice+CE Loss\n        # --------------------\n        base_loss = self.base_loss_fn(\n            y_true_mask[..., None],\n            y_pred,\n        )\n\n        # --------------------\n        # 2. Skeleton Recall Loss\n        # --------------------\n        intersection = keras.ops.sum(\n            pred_ink_prob * y_true_skel * valid_mask,\n            axis=(1, 2, 3),\n        )\n        skeleton_sum = keras.ops.sum(\n            y_true_skel * valid_mask,\n            axis=(1, 2, 3),\n        )\n\n        has_skeleton = keras.ops.cast(skeleton_sum > 0, \"float32\")\n        recall = (intersection + 1e-6) / (skeleton_sum + 1e-6)\n        skel_loss = keras.ops.mean((1.0 - recall) * has_skeleton)\n\n        # --------------------\n        # 3. FP Volume Loss (NEW)\n        # --------------------\n        gt_fg = keras.ops.cast(y_true_mask == 1, \"float32\")\n        gt_bg = keras.ops.cast(y_true_mask == 0, \"float32\")\n\n        fp_volume = (\n            pred_ink_prob * gt_bg * valid_mask\n        )\n\n        fp_loss = keras.ops.sum(fp_volume) / (\n            keras.ops.sum(gt_bg * valid_mask) + 1e-6\n        )\n\n        # --------------------\n        # Final Loss\n        # --------------------\n        return (\n            base_loss\n            + self.w_srec * skel_loss\n            + self.w_fp * fp_loss\n        )","metadata":{"trusted":true,"execution":{"execution_failed":"2025-12-24T15:42:39.284Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_samples = 780\nepochs = 25\ntotal_steps = (num_samples // batch_size) * epochs\nwarmup_steps = (num_samples // batch_size) * 5\n\nlr_schedule = keras.optimizers.schedules.CosineDecay(\n    initial_learning_rate=5e-5, \n    decay_steps=total_steps, \n    alpha=0.1 # Don't go to 0, stay at 5e-7\n)\n\n# define optomizer, loss, metrics\noptim = keras.optimizers.AdamW(\n    learning_rate=lr_schedule,\n    weight_decay=1e-5,\n)\n# older loss fn without skeletal loss\n# loss_fn = SparseDiceCELoss(\n#     from_logits=False, \n#     num_classes=num_classes,\n#     ignore_class_ids=2,\n# )\nloss_fn = SkeletonRecallPlusDiceLoss(\n    num_classes=num_classes,\n    \n)\n\nmetrics = [\n    SparseDiceMetric(\n        from_logits=False, \n        num_classes=num_classes, \n        ignore_class_ids=2,\n        name='dice'\n    ),\n]\n\nmodel.compile(\n    optimizer=optim,\n    loss=loss_fn,\n    metrics=metrics,\n)\n\nswi_callback_metric = SparseDiceMetric(\n    from_logits=False,\n    ignore_class_ids=2,\n    num_classes=num_classes,\n    name='val_dice',\n)\n\nswi_callback = SlidingWindowInferenceCallback(\n    model,\n    dataset=val_loader,\n    metrics=swi_callback_metric,\n    num_classes=num_classes,\n    interval=10,\n    overlap=0.5,\n    roi_size=input_shape,\n    sw_batch_size=2 * total_device,\n    save_path=\"model.weights.h5\"\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}\")\nsnapshot_cb = PeriodicWeightsSaver(\n    interval=20, \n    save_path_template=\"fine_tuning_epoch_{epoch}.weights.h5\"\n)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-12-24T15:42:39.285Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import keras\nfrom keras import ops\nfrom medicai.metrics import SparseDiceMetric\n\nclass MaskOnlySparseDiceMetric(SparseDiceMetric):\n    def update_state(self, y_true, y_pred, sample_weight=None):\n        # y_true comes from train_loader with shape: (B, D, H, W, 2) -> [Mask, Skeleton]\n        # We only need Channel 0 (the Mask) for the Dice calculation\n        # Check if we have the extra channel dimension\n        if len(y_true.shape) == 5 and y_true.shape[-1] == 2:\n            y_true_mask = y_true[..., 0]\n            # Restore channel dim to match expectations: (B, D, H, W, 1)\n            y_true_mask = ops.expand_dims(y_true_mask, axis=-1)\n        else:\n            y_true_mask = y_true\n            \n        return super().update_state(y_true_mask, y_pred, sample_weight=sample_weight)\n\n# Re-define metrics using the wrapper to handle the training data format\nmetrics = [\n    MaskOnlySparseDiceMetric(\n        from_logits=False, num_classes=num_classes, ignore_class_ids=2, name='dice'\n    ),\n    SkeletonLossMonitor(name='skel_L'), # Short name to fit in progress bar\n    BaseLossMonitor(num_classes=num_classes, name='base_L')\n]\n\n# Re-compile the model with the new metric\n# Note: optim and loss_fn are preserved from your previous cells\nmodel.compile(\n    optimizer=optim,\n    loss=loss_fn,\n    metrics=metrics\n)\n\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-12-24T15:42:39.285Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ALERT: Starting may take time.\nmodel.fit(\n    train_loader,\n    epochs=epochs,\n    callbacks=[\n        swi_callback,snapshot_cb\n    ]\n)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-12-24T15:42:39.285Z"}},"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":{"execution_failed":"2025-12-24T15:42:39.285Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# dice = SkeletonRecallPlusDiceLoss(\n#     num_classes=num_classes,\n# )\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-12-24T15:42:39.285Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# loss = []","metadata":{"trusted":true,"execution":{"execution_failed":"2025-12-24T15:42:39.285Z"}},"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":{"execution_failed":"2025-12-24T15:42:39.285Z"}},"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":{"execution_failed":"2025-12-24T15:42:39.285Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# x, y = next(iter(val_loader))\n# x.shape, y.shape","metadata":{"trusted":true,"execution":{"execution_failed":"2025-12-24T15:42:39.286Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# y_pred = swi(x)\n# y_pred.shape","metadata":{"trusted":true,"execution":{"execution_failed":"2025-12-24T15:42:39.286Z"}},"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":{"execution_failed":"2025-12-24T15:42:39.286Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# plot_sample(\n#     x, segment, sample_idx=0, max_slices=4\n# )","metadata":{"trusted":true,"execution":{"execution_failed":"2025-12-24T15:42:39.286Z"}},"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":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}