{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":117682,"databundleVersionId":14443416,"sourceType":"competition"}],"dockerImageVersionId":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"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!pip install imagecodecs -q","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-15T19:40:26.272056Z","iopub.execute_input":"2025-11-15T19:40:26.272396Z","iopub.status.idle":"2025-11-15T19:40:39.444147Z","shell.execute_reply.started":"2025-11-15T19:40:26.272366Z","shell.execute_reply":"2025-11-15T19:40:39.443311Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os, warnings\n\nos.environ['TF_CPP_MIN_LOG_LEVEL'] = '3'\nwarnings.filterwarnings('ignore')\n\nimport tensorflow as tf\nimport keras\nfrom keras import ops\n\nfrom sklearn.model_selection import train_test_split\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nimport tifffile","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-11-15T23:07:49.737399Z","iopub.execute_input":"2025-11-15T23:07:49.737641Z","iopub.status.idle":"2025-11-15T23:07:49.759541Z","shell.execute_reply.started":"2025-11-15T23:07:49.737615Z","shell.execute_reply":"2025-11-15T23:07:49.758903Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"keras.version(), keras.config.backend()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-15T19:40:57.115763Z","iopub.execute_input":"2025-11-15T19:40:57.116357Z","iopub.status.idle":"2025-11-15T19:40:57.122138Z","shell.execute_reply.started":"2025-11-15T19:40:57.116333Z","shell.execute_reply":"2025-11-15T19:40:57.121364Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"root_dir = \"/kaggle/input/vesuvius-challenge-surface-detection\"\nimages_dir = f\"{root_dir}/train_images\"\nlabels_dir = f\"{root_dir}/train_labels\"\nall_image_files = sorted(tf.io.gfile.glob(os.path.join(images_dir, \"*.tif\")))\nall_label_files = sorted(tf.io.gfile.glob(os.path.join(labels_dir, \"*.tif\")))\n\n# Pick 500 sample for quick prototyping.\nall_image_files = all_image_files[:500]\nall_label_files = all_label_files[:500]\n\ntrain_imgs, val_imgs, train_lbls, val_lbls = train_test_split(\n    all_image_files,\n    all_label_files,\n    test_size=5,\n    random_state=42,\n    shuffle=True,\n)\n\nprint(\"Train images:\", len(train_imgs), len(train_lbls))\nprint(\"Val images:  \", len(val_imgs), len(val_lbls))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-15T19:40:57.124175Z","iopub.execute_input":"2025-11-15T19:40:57.124577Z","iopub.status.idle":"2025-11-15T19:40:57.940391Z","shell.execute_reply.started":"2025-11-15T19:40:57.124546Z","shell.execute_reply":"2025-11-15T19:40:57.939733Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_csv('/kaggle/input/vesuvius-challenge-surface-detection/train.csv')\nprint(df.id.nunique(), df.scroll_id.nunique())\ndf.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-15T19:40:57.941090Z","iopub.execute_input":"2025-11-15T19:40:57.941651Z","iopub.status.idle":"2025-11-15T19:40:57.975874Z","shell.execute_reply.started":"2025-11-15T19:40:57.941630Z","shell.execute_reply":"2025-11-15T19:40:57.975324Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Quick Look**","metadata":{}},{"cell_type":"code","source":"for path in train_imgs:\n    img = Image.open(path).convert('RGB')\n    arr = np.array(img)\n    print(\"Shape:\", arr.shape)\n    print(\"Dtype:\", arr.dtype)\n    print(\"Min:\", arr.min())\n    print(\"Max:\", arr.max())\n    break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-15T19:40:57.976648Z","iopub.execute_input":"2025-11-15T19:40:57.976908Z","iopub.status.idle":"2025-11-15T19:40:58.201696Z","shell.execute_reply.started":"2025-11-15T19:40:57.976877Z","shell.execute_reply":"2025-11-15T19:40:58.200834Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for path in train_lbls:\n    img = Image.open(path).convert('RGB')\n    arr = np.array(img)\n    print(\"Shape:\", arr.shape)\n    print(\"Dtype:\", arr.dtype)\n    print(\"Min:\", arr.min())\n    print(\"Max:\", arr.max())\n    print(\"Unique values:\", np.unique(arr))\n    break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-15T19:40:58.202560Z","iopub.execute_input":"2025-11-15T19:40:58.202876Z","iopub.status.idle":"2025-11-15T19:40:58.219706Z","shell.execute_reply.started":"2025-11-15T19:40:58.202850Z","shell.execute_reply":"2025-11-15T19:40:58.219036Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Data Loader","metadata":{}},{"cell_type":"code","source":"from medicai.transforms import (\n    Compose,\n    ScaleIntensityRange,\n    Resize,\n    RandShiftIntensity,\n    RandRotate90,\n    RandFlip,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-15T19:40:58.220510Z","iopub.execute_input":"2025-11-15T19:40:58.220754Z","iopub.status.idle":"2025-11-15T19:40:58.230204Z","shell.execute_reply.started":"2025-11-15T19:40:58.220736Z","shell.execute_reply":"2025-11-15T19:40:58.229509Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_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            clip = True,\n        ),\n        Resize(\n            keys=[\"image\", \"label\"],\n            spatial_shape=(64, 128, 128),\n            mode=(\"trilinear\", \"nearest\")\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        RandShiftIntensity(\n            keys=[\"image\"],\n            offsets=0.10,\n            prob=1.0\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        ScaleIntensityRange(\n            keys=[\"image\"],\n            a_min = 0,\n            a_max = 255,\n            clip = True,\n        ),\n    ])\n    result = pipeline(data)\n    return result[\"image\"], result[\"label\"]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-15T19:40:58.230973Z","iopub.execute_input":"2025-11-15T19:40:58.231248Z","iopub.status.idle":"2025-11-15T19:40:58.238871Z","shell.execute_reply.started":"2025-11-15T19:40:58.231219Z","shell.execute_reply":"2025-11-15T19:40:58.238290Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DataLoader(keras.utils.Sequence):\n    def __init__(\n        self, \n        image_paths, \n        labels,\n        batch_size=1, \n        dim=(128, 128, 128), \n        shuffle=True, \n        training=True,\n        **kwargs\n    ):\n        super().__init__(**kwargs)\n        self.image_paths = image_paths\n        self.labels = labels\n        self.batch_size = batch_size\n        self.dim = dim  # (D, H, W)\n        self.shuffle = shuffle\n        self.training = training\n        self.on_epoch_end()\n\n    def __len__(self):\n        return int(np.floor(len(self.image_paths) / self.batch_size))\n\n    def __getitem__(self, index):\n        # issue: https://github.com/keras-team/keras/issues/20001\n        if index >= self.__len__():\n            raise StopIteration\n            \n        # Generate batch indices\n        indices = self.indices[index * self.batch_size : (index + 1) * self.batch_size]\n        image_paths_batch = [self.image_paths[k] for k in indices]\n        labels_batch = [self.labels[k] for k in indices]\n\n        # Initialize arrays\n        X = []\n        y = []\n\n        # Load and preprocess batch\n        for i, (img_path, label_path) in enumerate(zip(image_paths_batch, labels_batch)):\n            image = tifffile.imread(img_path) # shape: (D, H, W)\n            label = tifffile.imread(label_path)\n            label = (label == 1)\n\n            # Add channel dimension\n            image = np.expand_dims(image, axis=-1)  # (D, H, W, 1)\n            label = np.expand_dims(label, axis=-1)  # (D, H, W, 1)\n\n            image = image.astype(np.float32)\n            label = label.astype(np.float32)\n\n            # Apply transformations\n            if self.training:\n                image, label = train_transformation(image, label)\n            else:\n                image, label = val_transformation(image, label)\n\n            X.append(image)\n            y.append(label)\n\n        X = np.stack(X, axis=0)\n        y = np.stack(y, axis=0)\n        return X, y\n\n    def on_epoch_end(self):\n        # Shuffle indices after each epoch\n        self.indices = np.arange(len(self.image_paths))\n        if self.shuffle:\n            np.random.shuffle(self.indices)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-15T19:40:58.241373Z","iopub.execute_input":"2025-11-15T19:40:58.241737Z","iopub.status.idle":"2025-11-15T19:40:58.252830Z","shell.execute_reply.started":"2025-11-15T19:40:58.241721Z","shell.execute_reply":"2025-11-15T19:40:58.252079Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"input_shape=(64, 128, 128)\nbatch_size=6\nnum_classes=1\n\ntrain_loader = DataLoader(\n    image_paths=train_imgs,\n    labels=train_lbls,\n    batch_size=batch_size,\n    dim=input_shape,\n    shuffle=True,\n    training=True\n)\n\nval_loader = DataLoader(\n    image_paths=val_imgs,\n    labels=val_lbls,\n    batch_size=1,\n    dim=input_shape,\n    shuffle=False,\n    training=False\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-15T19:40:58.253473Z","iopub.execute_input":"2025-11-15T19:40:58.253671Z","iopub.status.idle":"2025-11-15T19:40:58.269846Z","shell.execute_reply.started":"2025-11-15T19:40:58.253655Z","shell.execute_reply":"2025-11-15T19:40:58.269045Z"}},"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-11-15T23:08:00.580162Z","iopub.execute_input":"2025-11-15T23:08:00.580352Z","iopub.status.idle":"2025-11-15T23:08:05.912886Z","shell.execute_reply.started":"2025-11-15T23:08:00.580330Z","shell.execute_reply":"2025-11-15T23:08:05.912049Z"}},"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":"2025-11-15T19:41:09.308235Z","iopub.execute_input":"2025-11-15T19:41:09.308736Z","iopub.status.idle":"2025-11-15T19:41:09.314888Z","shell.execute_reply.started":"2025-11-15T19:41:09.308715Z","shell.execute_reply":"2025-11-15T19:41:09.314124Z"},"_kg_hide-input":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":"2025-11-15T19:41:09.315642Z","iopub.execute_input":"2025-11-15T19:41:09.315928Z","iopub.status.idle":"2025-11-15T19:41:09.328255Z","shell.execute_reply.started":"2025-11-15T19:41:09.315899Z","shell.execute_reply":"2025-11-15T19:41:09.327649Z"}},"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-11-15T19:41:09.329361Z","iopub.execute_input":"2025-11-15T19:41:09.329599Z","iopub.status.idle":"2025-11-15T19:41:10.877482Z","shell.execute_reply.started":"2025-11-15T19:41:09.329583Z","shell.execute_reply":"2025-11-15T19:41:10.876843Z"}},"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":"2025-11-15T19:41:10.878197Z","iopub.execute_input":"2025-11-15T19:41:10.878382Z","iopub.status.idle":"2025-11-15T19:41:11.794394Z","shell.execute_reply.started":"2025-11-15T19:41:10.878367Z","shell.execute_reply":"2025-11-15T19:41:11.793568Z"}},"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":"2025-11-15T19:41:11.795270Z","iopub.execute_input":"2025-11-15T19:41:11.795646Z","iopub.status.idle":"2025-11-15T19:41:12.499242Z","shell.execute_reply.started":"2025-11-15T19:41:11.795617Z","shell.execute_reply":"2025-11-15T19:41:12.498263Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"import medicai\nfrom medicai.models import UNet\nfrom medicai.losses import BinaryDiceCELoss\nfrom medicai.metrics import BinaryDiceMetric\nfrom medicai.callbacks import SlidingWindowInferenceCallback","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-15T19:41:12.500125Z","iopub.execute_input":"2025-11-15T19:41:12.500335Z","iopub.status.idle":"2025-11-15T19:41:12.537391Z","shell.execute_reply.started":"2025-11-15T19:41:12.500318Z","shell.execute_reply":"2025-11-15T19:41:12.536824Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# medicai.models.list_models()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-15T19:41:12.538099Z","iopub.execute_input":"2025-11-15T19:41:12.538405Z","iopub.status.idle":"2025-11-15T19:41:12.541812Z","shell.execute_reply.started":"2025-11-15T19:41:12.538389Z","shell.execute_reply":"2025-11-15T19:41:12.541222Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = UNet(\n    input_shape=(64, 128, 128, 1),\n    encoder_name='efficientnet_b0',\n    encoder_depth=4,\n    classifier_activation='sigmoid',\n    num_classes=num_classes,\n)\nmodel.count_params() / 1e6","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-15T19:41:12.542744Z","iopub.execute_input":"2025-11-15T19:41:12.543555Z","iopub.status.idle":"2025-11-15T19:41:14.311750Z","shell.execute_reply.started":"2025-11-15T19:41:12.543529Z","shell.execute_reply":"2025-11-15T19:41:14.311027Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# define optomizer, loss, metrics\noptim = keras.optimizers.AdamW(\n    learning_rate=1e-4,\n    weight_decay=1e-5,\n)\n\nloss_fn = BinaryDiceCELoss(\n    from_logits=False, \n    num_classes=num_classes\n)\n\nmetrics = [\n    BinaryDiceMetric(\n        from_logits=False, \n        num_classes=num_classes, \n        name='dice'\n    ),\n]\n\nmodel.compile(\n    optimizer=optim,\n    loss=loss_fn,\n    metrics=metrics\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-15T19:41:14.312603Z","iopub.execute_input":"2025-11-15T19:41:14.312920Z","iopub.status.idle":"2025-11-15T19:41:14.329831Z","shell.execute_reply.started":"2025-11-15T19:41:14.312902Z","shell.execute_reply":"2025-11-15T19:41:14.329276Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"swi_callback_metric = BinaryDiceMetric(\n    from_logits=False,\n    ignore_empty=True,\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=5,\n    overlap=0.5,\n    roi_size=(64, 128, 128),\n    sw_batch_size=2,\n    save_path=\"model.weights.h5\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-15T19:41:14.330489Z","iopub.execute_input":"2025-11-15T19:41:14.330724Z","iopub.status.idle":"2025-11-15T19:41:14.340130Z","shell.execute_reply.started":"2025-11-15T19:41:14.330707Z","shell.execute_reply":"2025-11-15T19:41:14.339469Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.fit(\n    train_loader,\n    epochs=20,\n    callbacks=[\n        swi_callback\n    ]\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-15T19:41:14.340915Z","iopub.execute_input":"2025-11-15T19:41:14.341349Z","iopub.status.idle":"2025-11-15T23:07:49.729507Z","shell.execute_reply.started":"2025-11-15T19:41:14.341326Z","shell.execute_reply":"2025-11-15T23:07:49.728853Z"}},"outputs":[],"execution_count":null}]}