{"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":[{"sourceId":117682,"databundleVersionId":15062069,"sourceType":"competition"},{"sourceId":14266465,"sourceType":"datasetVersion","datasetId":8751895},{"sourceId":290917305,"sourceType":"kernelVersion"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"<div align=\"center\">\n    <a href=\"https://github.com/innat/medic-ai\">\n        <img src=\"https://i.imgur.com/nWOYfUO.png\" width=\"350\">\n    </a>\n</div>\n\n## About\n\n- It is starter designed for **Vesuvius Challenge - Surface Detection** in kaggle.\n- We will utilize [Medic-AI](https://github.com/innat/medic-ai), a high-performance library built on Keras 3 tailored for 2D and 3D medical image analysis. Thanks to its multi-backend architecture, it operates seamlessly across `tensorflow`, `torch`, and `jax`. This flexibility allows developers, including dedicated `torch` users to integrate state-of-the-art 2D or 3D segmentation models directly into their existing training workflows. One key architectural detail to note: Medic-AI adopts the (`depth, y, x, channel`) convention for input shapes. [This guide](https://www.kaggle.com/code/ipythonx/medicai-x-isic-2017-starter-x-binary-segmentation) demonstrates how to plug-and-play these capabilities into your own pure `torch` projects.\n- The official documentaiton of [`Medic-AI`](https://github.com/innat/medic-ai) is bit out-dated at the moment, and so to get up-to-date documentation, please refer to the GitHub readme page created for each model, i.e. [segormer](https://github.com/innat/medic-ai/blob/main/medicai/models/segformer/README.md), [trans-unet](https://github.com/innat/medic-ai/blob/main/medicai/models/transunet/README.md), [unetr++](https://github.com/innat/medic-ai/blob/main/medicai/models/unetr_plus_plus/README.md), [swin-unetr](https://github.com/innat/medic-ai/blob/main/medicai/models/swin/README.md#swin-unetr), [upernet](https://github.com/innat/medic-ai/blob/main/medicai/models/upernet/README.md), [convnext](https://github.com/innat/medic-ai/blob/main/medicai/models/convnext/README.md) etc.\n- Since [medicia](https://github.com/innat/medic-ai) is a new project, any feedback, suggestions, or contributions to help make it more reliable, robust, and improved are highly appreciated.\n\n## Code\n\n- **Training**: This notebook is the main training code, updated time to time. It uses `jax` backend. To use other backend (i.e., `torch`), see [this](https://www.kaggle.com/code/ipythonx/inference-vesuvius-surface-3d-detection).\n-  **Inference**: The inference code can be found here: [[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 with Jax backend.\n!pip install tensorflow -qU\n\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    --no-index \\\n    --find-links \"$var\"\n\nclear_output()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-25T17:39:39.922477Z","iopub.execute_input":"2026-01-25T17:39:39.922774Z","iopub.status.idle":"2026-01-25T17:39:42.111391Z","shell.execute_reply.started":"2026-01-25T17:39:39.922741Z","shell.execute_reply":"2026-01-25T17:39:42.110494Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# To get up-to-date feature, installing from scource is safe.\n!pip install git+https://github.com/innat/medic-ai.git -q\n\n# Installing is optional, we'll be using `tfrecord` format instead of `tif`.\n# !pip install imagecodecs tifffile -q","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-25T17:39:42.112010Z","iopub.execute_input":"2026-01-25T17:39:42.112175Z","iopub.status.idle":"2026-01-25T17:39:46.579880Z","shell.execute_reply.started":"2026-01-25T17:39:42.112156Z","shell.execute_reply":"2026-01-25T17:39:46.578805Z"},"_kg_hide-output":true},"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":"2026-01-25T17:39:46.580564Z","iopub.execute_input":"2026-01-25T17:39:46.580749Z","iopub.status.idle":"2026-01-25T17:39:46.583852Z","shell.execute_reply.started":"2026-01-25T17:39:46.580729Z","shell.execute_reply":"2026-01-25T17:39:46.583063Z"}},"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 import SGD, AdamW, Muon\nfrom keras.optimizers.schedules import CosineDecay, PolynomialDecay\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    NormalizeIntensity,\n    ScaleIntensityRange,\n    Resize,\n    RandShiftIntensity,\n    RandRotate90,\n    RandRotate,\n    RandFlip,\n    RandCutOut,\n    RandSpatialCrop\n)\nfrom medicai.layers import ResizingND\nfrom medicai.models import (\n    UNet, SegFormer, TransUNet, SwinUNETR, UPerNet, ConvNeXtV2Tiny, UNETRPlusPlus\n)\nfrom medicai.losses import (\n    SparseDiceCELoss, SparseTverskyLoss, SparseCenterlineDiceLoss\n)\nfrom medicai.metrics import SparseDiceMetric\nfrom medicai.callbacks import SlidingWindowInferenceCallback\nfrom medicai.utils import SlidingWindowInference\nfrom medicai.utils import soft_skeletonize","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-01-25T17:39:46.584326Z","iopub.execute_input":"2026-01-25T17:39:46.584484Z","iopub.status.idle":"2026-01-25T17:39:50.831276Z","shell.execute_reply.started":"2026-01-25T17:39:46.584468Z","shell.execute_reply":"2026-01-25T17:39:50.830474Z"}},"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":"2026-01-25T17:39:50.831811Z","iopub.execute_input":"2026-01-25T17:39:50.832147Z","iopub.status.idle":"2026-01-25T17:40:00.887778Z","shell.execute_reply.started":"2026-01-25T17:39:50.832128Z","shell.execute_reply":"2026-01-25T17:40:00.886904Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"keras.version(), keras.config.backend(), medicai.version()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-25T17:40:00.888307Z","iopub.execute_input":"2026-01-25T17:40:00.888487Z","iopub.status.idle":"2026-01-25T17:40:00.894478Z","shell.execute_reply.started":"2026-01-25T17:40:00.888469Z","shell.execute_reply":"2026-01-25T17:40:00.893752Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Data Loader","metadata":{}},{"cell_type":"code","source":"input_shape=(128, 128, 128)\nbatch_size=1 * total_device\nnum_classes=3\n\n# Each tfrecord contains 6 samples, total 786 samples.\nnum_samples = 780\nepochs = 200","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-25T17:40:00.894921Z","iopub.execute_input":"2026-01-25T17:40:00.895098Z","iopub.status.idle":"2026-01-25T17:40:00.905015Z","shell.execute_reply.started":"2026-01-25T17:40:00.895081Z","shell.execute_reply":"2026-01-25T17:40:00.904257Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**TFRecord Decoder**","metadata":{}},{"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":"2026-01-25T17:40:00.905451Z","iopub.execute_input":"2026-01-25T17:40:00.905630Z","iopub.status.idle":"2026-01-25T17:40:00.913745Z","shell.execute_reply.started":"2026-01-25T17:40:00.905614Z","shell.execute_reply":"2026-01-25T17:40:00.913156Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Preprocessing and Augmentation**","metadata":{"execution":{"iopub.status.busy":"2025-12-02T18:19:07.080480Z","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.091450Z"}}},{"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":"2026-01-25T17:40:00.914203Z","iopub.execute_input":"2026-01-25T17:40:00.914360Z","iopub.status.idle":"2026-01-25T17:40:00.925156Z","shell.execute_reply.started":"2026-01-25T17:40:00.914345Z","shell.execute_reply":"2026-01-25T17:40:00.924497Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_transformation(image, label):\n    data = {\"image\": image, \"label\": label}\n    pipeline = Compose([\n        ## Geometric transformation\n        RandSpatialCrop(\n            keys=[\"image\", \"label\"],\n            roi_size=input_shape,\n            random_center=True,\n            random_size=False,\n            invalid_label=2,         \n            min_valid_ratio=0.5,     \n            max_attempts=10\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        RandRotate(\n            keys=[\"image\", \"label\"], \n            factor=0.2, \n            prob=0.7, \n            fill_mode=\"crop\",\n        ),\n\n        ## Intensiry transformation\n        NormalizeIntensity(\n            keys=[\"image\"], \n            nonzero=True,\n            channel_wise=False\n        ),\n        RandShiftIntensity(\n            keys=[\"image\"], offsets=0.10, prob=0.5\n        ),\n        ## Spatial transformation \n        RandCutOut(\n            keys=[\"image\", \"label\"],\n            invalid_label=2, \n            mask_size=[\n                input_shape[1]//4,\n                input_shape[2]//4\n            ],\n            fill_mode=\"constant\", # \"constant\", \"gaussian\"\n            cutout_mode='volume', # \"slice\", \"volume\"\n            prob=0.8,\n            num_cuts=5,\n        ),\n    ])\n    result = pipeline(data)\n    return result[\"image\"], result[\"label\"]\n\n\ndef val_transformation(image, label):\n    data = {\"image\": image, \"label\": label}\n    pipeline = Compose([\n        NormalizeIntensity(\n            keys=[\"image\"], \n            nonzero=True,\n            channel_wise=False\n        ),\n    ])\n    result = pipeline(data)\n    return result[\"image\"], result[\"label\"]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-25T17:40:00.925649Z","iopub.execute_input":"2026-01-25T17:40:00.925805Z","iopub.status.idle":"2026-01-25T17:40:00.934448Z","shell.execute_reply.started":"2026-01-25T17:40:00.925790Z","shell.execute_reply":"2026-01-25T17:40:00.933789Z"}},"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    if shuffle:\n        dataset = dataset.map(\n            train_transformation,\n            num_parallel_calls=tf.data.AUTOTUNE\n        )\n    else:\n        dataset = dataset.map(\n            val_transformation,\n            num_parallel_calls=tf.data.AUTOTUNE\n        )\n    dataset = dataset.batch(batch_size, drop_remainder=shuffle)\n    dataset = dataset.prefetch(tf.data.AUTOTUNE)\n    return dataset","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-25T17:40:00.934888Z","iopub.execute_input":"2026-01-25T17:40:00.935064Z","iopub.status.idle":"2026-01-25T17:40:00.948330Z","shell.execute_reply.started":"2026-01-25T17:40:00.935048Z","shell.execute_reply":"2026-01-25T17:40:00.947700Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"all_tfrec = sorted(\n    glob.glob(\"/kaggle/input/vesuvius-tfrecords/*.tfrec\"),\n    key=lambda x: int(x.split(\"_\")[-1].replace(\".tfrec\", \"\"))\n)\n\nval_idx = -1\nval_patterns = [all_tfrec[val_idx]]\ntrain_patterns = [\n    f for i, f in enumerate(all_tfrec) if i != len(all_tfrec) + val_idx\n]\n\ntrain_ds = tfrecord_loader(\n    train_patterns, batch_size=batch_size, shuffle=True\n)\nval_ds = tfrecord_loader(\n    val_patterns, batch_size=1, shuffle=False\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-25T17:40:00.948786Z","iopub.execute_input":"2026-01-25T17:40:00.948942Z","iopub.status.idle":"2026-01-25T17:40:04.232425Z","shell.execute_reply.started":"2026-01-25T17:40:00.948927Z","shell.execute_reply":"2026-01-25T17:40:04.231160Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"x, y = next(iter(val_ds))\nx.shape, y.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-25T17:40:04.233217Z","iopub.execute_input":"2026-01-25T17:40:04.233406Z","iopub.status.idle":"2026-01-25T17:40:05.683735Z","shell.execute_reply.started":"2026-01-25T17:40:04.233387Z","shell.execute_reply":"2026-01-25T17:40:05.682531Z"}},"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,"execution":{"iopub.status.busy":"2026-01-25T17:40:05.684610Z","iopub.execute_input":"2026-01-25T17:40:05.684794Z","iopub.status.idle":"2026-01-25T17:40:05.689766Z","shell.execute_reply.started":"2026-01-25T17:40:05.684775Z","shell.execute_reply":"2026-01-25T17:40:05.688857Z"},"_kg_hide-input":true,"jupyter":{"source_hidden":true}},"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":"2026-01-25T17:40:05.690320Z","iopub.execute_input":"2026-01-25T17:40:05.690486Z","iopub.status.idle":"2026-01-25T17:40:05.703277Z","shell.execute_reply.started":"2026-01-25T17:40:05.690471Z","shell.execute_reply":"2026-01-25T17:40:05.702391Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_sample(\n    x, y, sample_idx=0, max_slices=4\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-25T17:40:05.703971Z","iopub.execute_input":"2026-01-25T17:40:05.704150Z","iopub.status.idle":"2026-01-25T17:40:06.145704Z","shell.execute_reply.started":"2026-01-25T17:40:05.704135Z","shell.execute_reply":"2026-01-25T17:40:06.144427Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_planes(\n    np.squeeze(x[0]), # picking one sample\n    np.squeeze(y[0])  # picking one sample\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-25T17:40:06.146386Z","iopub.execute_input":"2026-01-25T17:40:06.146584Z","iopub.status.idle":"2026-01-25T17:40:06.602105Z","shell.execute_reply.started":"2026-01-25T17:40:06.146566Z","shell.execute_reply":"2026-01-25T17:40:06.600997Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"soft_skel = soft_skeletonize(\n    ops.cast(y == 1, 'float32'),\n    iters=10\n)\nsoft_skel.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-25T17:43:56.767342Z","iopub.execute_input":"2026-01-25T17:43:56.767646Z","iopub.status.idle":"2026-01-25T17:43:59.706607Z","shell.execute_reply.started":"2026-01-25T17:43:56.767623Z","shell.execute_reply":"2026-01-25T17:43:59.705668Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_sample(\n    y, soft_skel, sample_idx=0, max_slices=4\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-25T17:43:59.707117Z","iopub.execute_input":"2026-01-25T17:43:59.707291Z","iopub.status.idle":"2026-01-25T17:44:00.466693Z","shell.execute_reply.started":"2026-01-25T17:43:59.707273Z","shell.execute_reply":"2026-01-25T17:44:00.465813Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"## check available models (classification + segmentation)\n# medicai.models.list_models()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-25T17:44:03.387047Z","iopub.execute_input":"2026-01-25T17:44:03.387297Z","iopub.status.idle":"2026-01-25T17:44:03.390332Z","shell.execute_reply.started":"2026-01-25T17:44:03.387277Z","shell.execute_reply":"2026-01-25T17:44:03.389386Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Pre-build encoder\nmodel = SegFormer(\n    input_shape=input_shape + (1,),\n    encoder_name='mit_b0',\n    classifier_activation='softmax',\n    num_classes=num_classes,\n)\n\n# model = UPerNet(\n#     input_shape=input_shape + (1,),\n#     encoder_name=\"convnext_base\",\n#     classifier_activation='softmax',\n#     num_classes=num_classes,\n# )\n\n# model = UNETRPlusPlus(\n#     input_shape=input_shape + (1,),\n#     encoder_name='unetr_plusplus_encoder',\n#     classifier_activation='softmax',\n#     num_classes=num_classes,\n# )\n\n# model = TransUNet(\n#     encoder_name='seresnext50',\n#     input_shape=input_shape + (1,),\n#     num_classes=num_classes,\n#     classifier_activation='softmax'\n# )\nmodel.count_params() / 1e6","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-25T17:44:04.779204Z","iopub.execute_input":"2026-01-25T17:44:04.779435Z","iopub.status.idle":"2026-01-25T17:44:11.972950Z","shell.execute_reply.started":"2026-01-25T17:44:04.779409Z","shell.execute_reply":"2026-01-25T17:44:11.971832Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ## Custom encoder\n# # backbone with 4 skip connection\n# backbone = ConvNeXtV2Tiny(\n#     input_shape=input_shape + (1,),\n#     include_top=False\n# )\n# # print(backbone.pyramid_outputs)\n\n# # segmentator\n# segmentor = TransUNet(\n#     encoder=backbone,\n#     encoder_depth=4,\n#     num_classes=num_classes,\n#     classifier_activation='softmax'\n# )\n# inputs = keras.Input(shape=input_shape + (1,))\n# x = segmentor(inputs)\n\n# # final supsampling, 2x tmes.\n# outputs = ResizingND(\n#     target_shape=input_shape,\n#     interpolation='trilinear',\n#     align_corners=False\n# )(x)\n# model = keras.Model(inputs=inputs, outputs=outputs)\n# model.count_params() / 1e6","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-25T17:44:11.973808Z","iopub.execute_input":"2026-01-25T17:44:11.974001Z","iopub.status.idle":"2026-01-25T17:44:11.977220Z","shell.execute_reply.started":"2026-01-25T17:44:11.973982Z","shell.execute_reply":"2026-01-25T17:44:11.976365Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ALERT: This attributes only available in medicai (not in core keras)\ntry:\n    print(model.instance_describe())\nexcept AttributeError:\n    pass","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-25T17:44:11.977808Z","iopub.execute_input":"2026-01-25T17:44:11.977985Z","iopub.status.idle":"2026-01-25T17:44:11.987949Z","shell.execute_reply.started":"2026-01-25T17:44:11.977969Z","shell.execute_reply":"2026-01-25T17:44:11.987154Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# LR Schedules and Optimizer","metadata":{}},{"cell_type":"code","source":"steps_per_epoch = num_samples // batch_size\ntotal_steps = steps_per_epoch * epochs\nwarmup_steps = int(total_steps * 0.05)\ndecay_steps = max(1, total_steps - warmup_steps)\nlr_schedule = CosineDecay(\n    initial_learning_rate=1e-6,\n    decay_steps=decay_steps,\n    warmup_target=min(3e-4, 1e-4 * (batch_size / 2)),\n    warmup_steps=warmup_steps,\n    alpha=0.1,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-25T17:44:11.988353Z","iopub.execute_input":"2026-01-25T17:44:11.988530Z","iopub.status.idle":"2026-01-25T17:44:11.995681Z","shell.execute_reply.started":"2026-01-25T17:44:11.988494Z","shell.execute_reply":"2026-01-25T17:44:11.994521Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# define optomizer, loss, metrics\noptim = keras.optimizers.AdamW(\n    learning_rate=lr_schedule,\n    weight_decay=1e-5,\n)\n\ndice_ce_loss_fn = SparseDiceCELoss(\n    from_logits=False, \n    num_classes=num_classes,\n    ignore_class_ids=2,\n)\ncldice_loss_fn = SparseCenterlineDiceLoss(\n    from_logits=False, \n    num_classes=num_classes,\n    target_class_ids=1,\n    ignore_class_ids=2,\n    iters=50\n)\ncombined_loss_fn = lambda y_true, y_pred: (\n    dice_ce_loss_fn(y_true, y_pred) + cldice_loss_fn(y_true, y_pred)\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=combined_loss_fn,\n    metrics=metrics,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-25T17:44:11.996252Z","iopub.execute_input":"2026-01-25T17:44:11.996440Z","iopub.status.idle":"2026-01-25T17:44:12.032721Z","shell.execute_reply.started":"2026-01-25T17:44:11.996420Z","shell.execute_reply":"2026-01-25T17:44:12.031786Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"swi_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_ds,\n    metrics=swi_callback_metric,\n    num_classes=num_classes,\n    interval=5,\n    overlap=0.5,\n    mode='gaussian',\n    roi_size=input_shape,\n    sw_batch_size=1 * total_device,\n    save_path=\"model.weights.h5\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-25T17:44:12.033277Z","iopub.execute_input":"2026-01-25T17:44:12.033454Z","iopub.status.idle":"2026-01-25T17:44:12.039764Z","shell.execute_reply.started":"2026-01-25T17:44:12.033436Z","shell.execute_reply":"2026-01-25T17:44:12.038972Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ALERT: Starting may take time.\nmodel.fit(\n    train_ds,\n    epochs=epochs,\n    callbacks=[\n        swi_callback\n    ]\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-25T17:44:12.219787Z","iopub.execute_input":"2026-01-25T17:44:12.220007Z","iopub.status.idle":"2026-01-25T22:04:50.550408Z","shell.execute_reply.started":"2026-01-25T17:44:12.219988Z","shell.execute_reply":"2026-01-25T22:04:50.549012Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Eval","metadata":{}},{"cell_type":"code","source":"model.load_weights(\n    \"model.weights.h5\"\n)\nswi = SlidingWindowInference(\n    model,\n    num_classes=num_classes,\n    roi_size=input_shape,\n    mode='gaussian',\n    sw_batch_size=1 * total_device,\n    overlap=0.5,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-25T22:05:57.952315Z","iopub.execute_input":"2026-01-25T22:05:57.952623Z","iopub.status.idle":"2026-01-25T22:05:58.828682Z","shell.execute_reply.started":"2026-01-25T22:05:57.952602Z","shell.execute_reply":"2026-01-25T22:05:58.827408Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dice = SparseDiceMetric(\n    from_logits=False,\n    num_classes=num_classes,\n    ignore_class_ids=2,\n    name='dice',\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-25T22:06:00.509970Z","iopub.execute_input":"2026-01-25T22:06:00.510284Z","iopub.status.idle":"2026-01-25T22:06:00.517998Z","shell.execute_reply.started":"2026-01-25T22:06:00.510264Z","shell.execute_reply":"2026-01-25T22:06:00.517033Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for sample in val_ds:\n    x, y = sample\n    output = swi(x)\n    y = ops.convert_to_tensor(y)\n    output = ops.convert_to_tensor(output)\n    dice.update_state(y, output)\n\ndice_score = float(ops.convert_to_numpy(dice.result()))\nprint(f\"Dice Score: {dice_score:.4f}\")\ndice.reset_state()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-25T22:06:01.919653Z","iopub.execute_input":"2026-01-25T22:06:01.919871Z","iopub.status.idle":"2026-01-25T22:06:19.297303Z","shell.execute_reply.started":"2026-01-25T22:06:01.919854Z","shell.execute_reply":"2026-01-25T22:06:19.295986Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"x, y = next(iter(val_ds))\nx.shape, y.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-25T22:06:21.486733Z","iopub.execute_input":"2026-01-25T22:06:21.486998Z","iopub.status.idle":"2026-01-25T22:06:22.879939Z","shell.execute_reply.started":"2026-01-25T22:06:21.486979Z","shell.execute_reply":"2026-01-25T22:06:22.878615Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"y_pred = swi(x)\ny_pred.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-25T22:06:24.463862Z","iopub.execute_input":"2026-01-25T22:06:24.464134Z","iopub.status.idle":"2026-01-25T22:06:27.009065Z","shell.execute_reply.started":"2026-01-25T22:06:24.464112Z","shell.execute_reply":"2026-01-25T22:06:27.007792Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"segment = y_pred.argmax(-1).astype(np.uint8)\nsegment.shape, np.unique(segment)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-25T22:06:27.390824Z","iopub.execute_input":"2026-01-25T22:06:27.391099Z","iopub.status.idle":"2026-01-25T22:06:27.663571Z","shell.execute_reply.started":"2026-01-25T22:06:27.391076Z","shell.execute_reply":"2026-01-25T22:06:27.662302Z"}},"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":"2026-01-25T22:06:29.072818Z","iopub.execute_input":"2026-01-25T22:06:29.073083Z","iopub.status.idle":"2026-01-25T22:06:29.492194Z","shell.execute_reply.started":"2026-01-25T22:06:29.073062Z","shell.execute_reply":"2026-01-25T22:06:29.490870Z"}},"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- [Centerline Dice Loss](https://github.com/jocpae/clDice)\n- [Skeleton-Recall Loss](https://github.com/MIC-DKFZ/Skeleton-Recall)\n- [Virtually Unrolling the Herculaneum Papyri](https://arxiv.org/pdf/2512.04927v1)","metadata":{}},{"cell_type":"markdown","source":"## Deep Supervision Recipes\n\n**Deep supervision** in segmentation models improves training by adding auxiliary loss functions to hidden layers, mitigating gradient vanishing, and speeding up convergence. It forces intermediate layers to learn more discriminative features, enhancing overall performance.\n\n**Preface**\n\nSay, we have an segmentation model with `5` level pyramid encoder blocks. So, typically we would have `5` upsampling stages. Some are deep, some are high level. For example, model like `TransUNet`, we can inspect:\n\n```python\nmodel = TransUNet(\n    encoder_name='seresnext50',\n    input_shape=target_shape + (1,),\n    classifier_activation='softmax',\n    num_classes=num_classes,\n)\n\ndecoder_layers = [\n    (\"DS5\", \"decoder_proj_0\"),\n    (\"DS4\", \"decoder_conv_4\"),\n    (\"DS3\", \"decoder_conv_3\"),\n    (\"DS2\", \"decoder_conv_2\"),\n    (\"DS1\", \"decoder_conv_1\"),\n]\n\nfor name, layer_name in decoder_layers:\n    feat = model.get_layer(layer_name).output\n    print(name, feat.shape)\n    \nDS5 (None, 3, 3, 3, 256)\nDS4 (None, 6, 6, 6, 128)\nDS3 (None, 12, 12, 12, 64)\nDS2 (None, 24, 24, 24, 32)\nDS1 (None, 48, 48, 48, 16)\n```\n\nIdeally, we would pick higher level features, i.e. `DS3`, `DS2`, `DS1`. However for sake for demonstration, let's pick all. Now, in Keras, there are many ways we can model deep supervision. Here is one traditional way.\n\n\n**Step 1**\n\n```python\ndef prepare_multi_outputs(x, y):\n    y_dict = {\n        \"final\": y,\n        \"DS1\": y,\n        \"DS2\": y,\n        \"DS3\": y,\n        \"DS4\": y,\n        \"DS5\": y,\n    }\n    return x, y_dict\n\ndef load_tfrecord_dataset(tfrecord_pattern, ...):\n    ....\n    if shuffle:\n        dataset = dataset.map(\n            train_transformation,\n            num_parallel_calls=tf.data.AUTOTUNE\n        )\n        dataset = dataset.map( # < --- HERE\n            prepare_multi_outputs, \n            num_parallel_calls=tf.data.AUTOTUNE\n        )\n    else:\n        dataset = dataset.map(\n            val_transformation,\n            num_parallel_calls=tf.data.AUTOTUNE\n        )\n    dataset = dataset.batch(\n        batch_size, drop_remainder=shuffle\n    ).prefetch(tf.data.AUTOTUNE)\n    return dataset\n```\n\n**Step 2**\n\nBuild the multi-output model.\n\n```python\naux_outputs = {}\nfor name, layer_name in decoder_layers:\n    feat = model.get_layer(layer_name).output\n\n    # project features → class logits\n    logits = keras.layers.Conv3D(\n        filters=num_classes,\n        kernel_size=1,\n        padding=\"same\",\n        name=f\"aux_logits_{name}\",\n    )(feat)\n\n    # resize logits (trilinear)\n    resized = ResizingND(\n        target_shape=input_size,\n        interpolation=\"trilinear\",\n        name=f\"resize_{name}\",\n    )(logits)\n    \n    # softmax after resizing\n    aux = keras.layers.Activation(\n        \"softmax\",\n        name=name,\n        dtype=\"float32\",\n    )(resized)\n\n    aux_outputs[name] = aux\n\ndeep_supervised_model = keras.Model(\n    inputs=model.input,\n    outputs={\n        \"final\": model.output,\n        **aux_outputs,\n    },\n)\n```\n\nNow, we compile this model.\n\n```python\nlosses = {\n    \"final\": SparseDiceCELoss | AnyLossWeWant,\n    \"DS1\": SparseDiceCELoss | AnyLossWeWant,\n    \"DS2\": SparseDiceCELoss | AnyLossWeWant,\n    \"DS3\": SparseDiceCELoss | AnyLossWeWant,\n    \"DS4\": SparseDiceCELoss | AnyLossWeWant,\n    \"DS5\": SparseDiceCELoss | AnyLossWeWant,\n}\n\nloss_weights = {\n    \"final\":  0.5079,\n    \"DS1\":    0.2539,\n    \"DS2\":    0.1269,\n    \"DS3\":    0.0635,\n    \"DS4\":    0.0317,\n    \"DS5\":    0.0158,\n}\n\ndeep_supervised_model.compile(\n    optimizer=optim,\n    loss=losses,\n    loss_weights=loss_weights,\n)\n```\n\nWe use loss weights according to [nnUNet-tf](https://github.com/NVIDIA/DeepLearningExamples/blob/729963dd47e7c8bd462ad10bfac7a7b0b604e6dd/TensorFlow2/Segmentation/nnUNet/models/nn_unet.py#L82-L93).\n\n**Step 3**\n\nFor validation or inference, just use raw model. These are Keras functional API. So, weights are already shared. \n\n```python\nswi_callback = SlidingWindowInferenceCallback(\n    model,\n    ...\n)\n```\n\n**Step 4**\n\nFit the model.\n\n```python\ndeep_supervised_model.fit(\n    train_ds,\n    epochs=epochs,\n    callbacks=[\n        swi_callback,\n    ]\n)\n```","metadata":{}},{"cell_type":"markdown","source":"## Resume Training from Interruption\n\nSay, we set 500 epoch and saving checkpoint every 5 epochs. Now, the training got interrupted after epoch 350 before finising 351. In that case, we can resume training as follows:\n\n```python\nmodel = Model()\nmodel.load_weights(...) # last saved checkpoint\n\n# optim, lr sched etc, same as before, i.e.,\nsteps_per_epoch = num_samples // batch_size\nlr_schedule = CosineDecay(\n    initial_learning_rate=1e-6,\n    decay_steps=decay_steps,\n    warmup_target=min(3e-4, 1e-4 * (batch_size / 2)),\n    warmup_steps=warmup_steps,\n    alpha=0.1,\n)\noptim = keras.optimizers.AdamW(\n    learning_rate=lr_schedule,\n    weight_decay=1e-5,\n)\n\n# Important\nalready_trained_steps = 350 * steps_per_epoch\noptim.iterations.assign(already_trained_steps)\n\nmodel.fit(\n    ...\n    initial_epoch=350,\n)\n```\n\nFYI, if we could save the **compiled model** with `.keras` format, it would be much easy.","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}