{"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":"gpu","dataSources":[{"sourceType":"competition","sourceId":117682,"databundleVersionId":15062069},{"sourceType":"datasetVersion","sourceId":14266755,"datasetId":8766236,"databundleVersionId":15067313},{"sourceType":"kernelVersion","sourceId":290917305}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# About\n\n- Implement a `torch.utils.data.DataLoader`.\n- Implement a `pytorch_lightning.LightningDataModule`\n- Build a **3D** segmentation model from [medicai](https://github.com/innat/medic-ai) using the **PyTorch** backend\n- Implement a `pytorch_lightning.LightningModule`\n\n---\n\n**Note**: This is just a fun experiment. We’ll be using **Keras 3**–based 3D segmentation models (running on the **PyTorch** backend), while the data loading and training loop are handled entirely through **PyTorch Lightning**. Also, please note that the **PyTorch Lightning** website is **not accessible** in my region, so I’m unable to check the official API documentation. If you notice any implementation mismatches related to **Lightning**, such as model checkpointing, learning rate schedulers, or callbacks; please feel free to adjust them or mention it in a comment. I’ll review and address those if possible.\n\n---\n\n**Caveat**: The loss and metrics (Dice-CE loss and Dice metric) come from **Keras 3** via `medicai`, because PyTorch does not provide these specialized components out of the box. They could be replaced by pure-PyTorch implementations at any time. Now, in Keras, the conventions differ slightly:\n- Loss functions follow `fn(y_true, y_pred)` - as opposed to PyTorch’s `fn(y_pred, y_true)`\n- Metrics follow: \n    - `metric.update_state(y_true, y_pred)` per batch\n    - `metric.result()` when retrieving current value\n    - `metric.reset_state()` at the end of each epoch\n\nIf you switch to native **PyTorch** loss functions and metrics, you can ignore these differences.","metadata":{}},{"cell_type":"code","source":"from IPython.display import clear_output\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\n!pip uninstall -y protobuf -q\n!pip install protobuf==5.26.1 -q\n!pip install git+https://github.com/innat/medic-ai.git -q\nclear_output()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-01-30T11:50:42.086748Z","iopub.execute_input":"2026-01-30T11:50:42.087261Z","iopub.status.idle":"2026-01-30T11:51:04.106871Z","shell.execute_reply.started":"2026-01-30T11:50:42.087210Z","shell.execute_reply":"2026-01-30T11:51:04.105791Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, warnings\n\nos.environ[\"KERAS_BACKEND\"] = \"torch\"\nos.environ[\"XLA_FLAGS\"] = \"--xla_gpu_force_compilation_parallelism=1\"\nos.environ[\"TF_CPP_MIN_LOG_LEVEL\"] = \"3\" \nwarnings.filterwarnings('ignore')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-30T11:51:04.109386Z","iopub.execute_input":"2026-01-30T11:51:04.111123Z","iopub.status.idle":"2026-01-30T11:51:04.115734Z","shell.execute_reply.started":"2026-01-30T11:51:04.111090Z","shell.execute_reply":"2026-01-30T11:51:04.114973Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import glob\nimport numpy as np\nimport matplotlib.pyplot as plt","metadata":{"trusted":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2026-01-30T11:51:04.116718Z","iopub.execute_input":"2026-01-30T11:51:04.117007Z","iopub.status.idle":"2026-01-30T11:51:04.133856Z","shell.execute_reply.started":"2026-01-30T11:51:04.116978Z","shell.execute_reply":"2026-01-30T11:51:04.133061Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"root_dir = \"/kaggle/input/vesuvius-npy\"\nimages_dir = f\"{root_dir}/train_images\"\nlabels_dir = f\"{root_dir}/train_labels\"\nall_image_files = sorted(glob.glob(os.path.join(images_dir, \"*.npy\")))\nall_label_files = sorted(glob.glob(os.path.join(labels_dir, \"*.npy\")))\n\n# number of validation samples\nval_count = 6\n\n# Split by slicing\ntrain_imgs = all_image_files[:-val_count]\nval_imgs   = all_image_files[-val_count:]\n\ntrain_lbls = all_label_files[:-val_count]\nval_lbls   = all_label_files[-val_count:]\n\nprint(\"Train images:\", len(train_imgs), len(train_lbls))\nprint(\"Val images:  \", len(val_imgs), len(val_lbls))","metadata":{"trusted":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2026-01-30T11:51:04.134606Z","iopub.execute_input":"2026-01-30T11:51:04.134817Z","iopub.status.idle":"2026-01-30T11:51:04.196088Z","shell.execute_reply.started":"2026-01-30T11:51:04.134797Z","shell.execute_reply":"2026-01-30T11:51:04.195299Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import keras\nfrom keras import ops\n\nimport torch\nimport pytorch_lightning as pl\n\nfrom pytorch_lightning.callbacks import ModelCheckpoint, LearningRateMonitor\nfrom pytorch_lightning.callbacks import Callback\n\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":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-30T11:51:04.198646Z","iopub.execute_input":"2026-01-30T11:51:04.198881Z","iopub.status.idle":"2026-01-30T11:51:32.465768Z","shell.execute_reply.started":"2026-01-30T11:51:04.198859Z","shell.execute_reply":"2026-01-30T11:51:32.465091Z"},"_kg_hide-output":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"keras.version(), keras.config.backend(), pl.__version__, medicai.version()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-30T11:51:32.466674Z","iopub.execute_input":"2026-01-30T11:51:32.467217Z","iopub.status.idle":"2026-01-30T11:51:32.473310Z","shell.execute_reply.started":"2026-01-30T11:51:32.467172Z","shell.execute_reply":"2026-01-30T11:51:32.472588Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Data Loader","metadata":{}},{"cell_type":"code","source":"input_shape=(128, 128, 128)\nbatch_size=1\nnum_classes=3\n\n# Total npy file - validation npy file\nnum_samples = 786 - val_count\nepochs = 10","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-30T11:51:32.474441Z","iopub.execute_input":"2026-01-30T11:51:32.474823Z","iopub.status.idle":"2026-01-30T11:51:32.500169Z","shell.execute_reply.started":"2026-01-30T11:51:32.474779Z","shell.execute_reply":"2026-01-30T11:51:32.499331Z"}},"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\n        ## Z-score norm\n        NormalizeIntensity(\n            keys=[\"image\"], \n            nonzero=True,\n            channel_wise=False\n        ),\n\n        ## Intensiry transformation\n        RandShiftIntensity(\n            keys=[\"image\"], offsets=0.10, prob=0.5\n        ),\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.2,\n            num_cuts=2,\n        ),\n    ])\n    result = pipeline(data)\n    return (\n        ops.convert_to_numpy(result[\"image\"]), \n        ops.convert_to_numpy(result[\"label\"])\n    )\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 (\n        ops.convert_to_numpy(result[\"image\"]), \n        ops.convert_to_numpy(result[\"label\"])\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-30T11:51:32.501328Z","iopub.execute_input":"2026-01-30T11:51:32.501626Z","iopub.status.idle":"2026-01-30T11:51:32.515742Z","shell.execute_reply.started":"2026-01-30T11:51:32.501588Z","shell.execute_reply":"2026-01-30T11:51:32.514826Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class VesuviusDataset(torch.utils.data.Dataset):\n    def __init__(self, image_paths, label_paths, training=True):\n        self.images = image_paths\n        self.labels = label_paths\n        self.training = training\n\n    def __len__(self):\n        return len(self.images)\n\n    def load_npy(self, path):\n        array = np.load(path)   \n        array = array[..., None]\n        array = array.astype(np.float32)\n        return array\n\n    def __getitem__(self, idx):\n        # load volumes\n        image = self.load_npy(self.images[idx])    \n        label = self.load_npy(self.labels[idx])\n\n        # apply transforms\n        if self.training:\n            image, label = train_transformation(image, label)\n        else:\n            image, label = val_transformation(image, label)\n\n        # convert to tensors\n        image = torch.from_numpy(image).float()\n        label = torch.from_numpy(label).float()\n        return image, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-30T11:51:32.516720Z","iopub.execute_input":"2026-01-30T11:51:32.517002Z","iopub.status.idle":"2026-01-30T11:51:32.534297Z","shell.execute_reply.started":"2026-01-30T11:51:32.516977Z","shell.execute_reply":"2026-01-30T11:51:32.533332Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"check_out_ds = VesuviusDataset(val_imgs, val_lbls, training=False)\ncheck_out_loader = torch.utils.data.DataLoader(\n    check_out_ds,\n    batch_size=1,\n    shuffle=False,\n    num_workers=4,\n)","metadata":{"trusted":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2026-01-30T11:51:32.535450Z","iopub.execute_input":"2026-01-30T11:51:32.535876Z","iopub.status.idle":"2026-01-30T11:51:32.554577Z","shell.execute_reply.started":"2026-01-30T11:51:32.535832Z","shell.execute_reply":"2026-01-30T11:51:32.553777Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"x, y = next(iter(check_out_loader))\nx.shape, y.shape","metadata":{"trusted":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2026-01-30T11:51:32.555525Z","iopub.execute_input":"2026-01-30T11:51:32.556324Z","iopub.status.idle":"2026-01-30T11:51:39.550439Z","shell.execute_reply.started":"2026-01-30T11:51:32.556287Z","shell.execute_reply":"2026-01-30T11:51:39.549574Z"}},"outputs":[],"execution_count":null},{"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()","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2026-01-30T11:51:39.551989Z","iopub.execute_input":"2026-01-30T11:51:39.552725Z","iopub.status.idle":"2026-01-30T11:51:39.559511Z","shell.execute_reply.started":"2026-01-30T11:51:39.552686Z","shell.execute_reply":"2026-01-30T11:51:39.558675Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_planes(image, mask, alpha=0.4):\n    image = np.squeeze(image)\n    mask = np.squeeze(mask)\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()","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2026-01-30T11:51:39.560985Z","iopub.execute_input":"2026-01-30T11:51:39.561331Z","iopub.status.idle":"2026-01-30T11:51:39.574674Z","shell.execute_reply.started":"2026-01-30T11:51:39.561303Z","shell.execute_reply":"2026-01-30T11:51:39.573910Z"}},"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-30T11:51:39.577959Z","iopub.execute_input":"2026-01-30T11:51:39.578200Z","iopub.status.idle":"2026-01-30T11:51:40.296358Z","shell.execute_reply.started":"2026-01-30T11:51:39.578175Z","shell.execute_reply":"2026-01-30T11:51:40.295436Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_planes(\n    x[0], # picking one sample\n    y[0]  # picking one sample\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-30T11:51:40.297324Z","iopub.execute_input":"2026-01-30T11:51:40.297618Z","iopub.status.idle":"2026-01-30T11:51:41.183626Z","shell.execute_reply.started":"2026-01-30T11:51:40.297591Z","shell.execute_reply":"2026-01-30T11:51:41.182651Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Lightning DataModule","metadata":{}},{"cell_type":"code","source":"class VesuviusDataModule(pl.LightningDataModule):\n    def __init__(\n        self,\n        train_imgs, \n        train_lbls, \n        val_imgs, \n        val_lbls,\n        train_batch_size=2,\n        val_batch_size=1,\n        num_workers=4,\n    ):\n        super().__init__()\n        self.train_imgs = train_imgs\n        self.train_lbls = train_lbls\n        self.val_imgs = val_imgs\n        self.val_lbls = val_lbls\n        self.train_batch_size = train_batch_size\n        self.val_batch_size = val_batch_size\n        self.num_workers = num_workers\n\n    def setup(self, stage=None):\n        self.train_ds = VesuviusDataset(\n            self.train_imgs, self.train_lbls, training=True\n        )\n        self.val_ds = VesuviusDataset(\n            self.val_imgs, self.val_lbls, training=False\n        )\n\n    def train_dataloader(self):\n        return  torch.utils.data.DataLoader(\n            self.train_ds, \n            batch_size=self.train_batch_size, \n            shuffle=True,\n            num_workers=self.num_workers, \n            pin_memory=True\n        )\n\n    def val_dataloader(self):\n        return  torch.utils.data.DataLoader(\n            self.val_ds, \n            batch_size=self.val_batch_size, \n            shuffle=False,\n            num_workers=self.num_workers, \n            pin_memory=True\n        )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-30T11:51:41.184872Z","iopub.execute_input":"2026-01-30T11:51:41.185148Z","iopub.status.idle":"2026-01-30T11:51:41.192290Z","shell.execute_reply.started":"2026-01-30T11:51:41.185121Z","shell.execute_reply":"2026-01-30T11:51:41.191643Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"model = SegFormer(\n    input_shape=input_shape + (1,),\n    encoder_name='mit_b0',\n    classifier_activation='softmax',\n    num_classes=num_classes,\n    dropout=0.2,\n)\nmodel.compile(jit_compile=False)\nsum(p.numel() for p in model.parameters() if p.requires_grad) / 1e6","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-30T11:51:41.193435Z","iopub.execute_input":"2026-01-30T11:51:41.194336Z","iopub.status.idle":"2026-01-30T11:51:42.854213Z","shell.execute_reply.started":"2026-01-30T11:51:41.194308Z","shell.execute_reply":"2026-01-30T11:51:42.853481Z"}},"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-30T11:51:42.855143Z","iopub.execute_input":"2026-01-30T11:51:42.855485Z","iopub.status.idle":"2026-01-30T11:51:42.860488Z","shell.execute_reply.started":"2026-01-30T11:51:42.855420Z","shell.execute_reply":"2026-01-30T11:51:42.859530Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Lightning Module","metadata":{}},{"cell_type":"code","source":"# loss function\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=1, # ideal to set 20-50 - computationally expensive\n)\n\ndef combined_loss(dice_ce_fn, cldice_fn):\n    def loss_fn(y_true, y_pred):\n        loss_dice_ce = dice_ce_fn(y_true, y_pred)\n        loss_cldice = cldice_fn(y_true, y_pred)\n        return loss_dice_ce + loss_cldice\n    return loss_fn","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-30T11:51:42.861531Z","iopub.execute_input":"2026-01-30T11:51:42.861891Z","iopub.status.idle":"2026-01-30T11:51:42.876474Z","shell.execute_reply.started":"2026-01-30T11:51:42.861857Z","shell.execute_reply":"2026-01-30T11:51:42.875577Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class VesuviusModel(pl.LightningModule):\n    def __init__(\n        self,\n        model,\n        learning_rate=1e-4,\n        sw_batch_size=4,\n        sw_overlap=0.5\n    ):\n        super().__init__()\n        self.model = model\n        self.num_classes = num_classes\n        self.learning_rate = learning_rate\n        self.sw_batch_size = sw_batch_size\n        self.sw_overlap = sw_overlap\n        \n        # loss function\n        # [NOTE]: Loss method should same for both training and validaiton\n        # But due to high computational cost, we opted out clDice for validation.\n        # Enable it with high resources.\n        self.train_criterion = combined_loss(\n            dice_ce_loss_fn, cldice_loss_fn\n        )\n        self.val_criterion = dice_ce_loss_fn\n\n        # metrics function\n        self.train_dice_score = SparseDiceMetric(\n            from_logits=False,\n            num_classes=num_classes,\n            ignore_class_ids=2,\n            name='dice',\n        )\n        self.val_dice_score = SparseDiceMetric(\n            from_logits=False,\n            num_classes=num_classes,\n            ignore_class_ids=2,\n            name='val_dice',\n        )\n        \n        # sliding window inferer for validation\n        self.sliding_window_inferer = SlidingWindowInference(\n            self.model,\n            num_classes=self.num_classes,\n            roi_size=input_shape,\n            sw_batch_size=self.sw_batch_size,\n            overlap=self.sw_overlap,\n        )\n        self.save_hyperparameters(ignore=['model'])\n\n    def forward(self, x):\n        return self.model(x)\n\n    def training_step(self, batch, batch_idx):\n        images, masks = batch\n        \n        # forward pass\n        prob = self.forward(images)\n        \n        # calculate loss\n        loss = self.train_criterion(masks, prob)\n        \n        # calculate metrics\n        self.train_dice_score.update_state(masks, prob)\n \n        # log metrics\n        current_dice_score = self.train_dice_score.result()\n        self.log(\n            'train_loss', \n            loss, \n            on_step=True, \n            on_epoch=True, \n            prog_bar=True\n        )\n        self.log(\n            'train_dice', \n            current_dice_score, \n            on_step=True, \n            on_epoch=True, \n            prog_bar=True\n        )\n        return loss\n\n    def validation_step(self, batch, batch_idx):\n        image, mask = batch\n        image = ops.convert_to_numpy(image)\n        prob = self.sliding_window_inferer(\n            image\n        )\n        self.val_dice_score.update_state(mask, prob)\n        \n        # calculate loss / metrics\n        loss = self.val_criterion(mask, prob)\n\n        # log loss\n        current_dice_score = self.val_dice_score.result()\n        self.log(\n            'val_loss', loss, on_step=False, on_epoch=True, prog_bar=True\n        )\n        self.log(\n            'val_dice', \n            current_dice_score, \n            on_step=False, \n            on_epoch=True, \n            prog_bar=True\n        )\n        return loss\n\n    def configure_optimizers(self):\n        optimizer = torch.optim.AdamW(\n            self.parameters(), \n            lr=self.learning_rate,\n            weight_decay=1e-5\n        )\n\n        scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n            optimizer,\n            mode='max',\n            factor=0.5,\n            patience=2,\n        )\n        return {\n            'optimizer': optimizer,\n            'lr_scheduler': {\n                'scheduler': scheduler,\n                'monitor': 'val_dice',\n                'frequency': 1\n            }\n        }\n\n    def on_train_epoch_end(self):\n        # reset metrics at the end of each epoch\n        self.train_dice_score.reset_state()\n\n    def on_validation_epoch_end(self):\n        # reset metrics at the end of each epoch\n        self.val_dice_score.reset_state()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-30T11:51:42.877619Z","iopub.execute_input":"2026-01-30T11:51:42.877886Z","iopub.status.idle":"2026-01-30T11:51:42.897050Z","shell.execute_reply.started":"2026-01-30T11:51:42.877859Z","shell.execute_reply":"2026-01-30T11:51:42.896283Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Custom Print Callback**","metadata":{}},{"cell_type":"code","source":"class MetricLoggerCallback(Callback):\n    def __init__(self, print_every_n_steps=20):\n        super().__init__()\n        self.print_every_n_steps = print_every_n_steps\n\n    def on_train_batch_end(\n        self, trainer, pl_module, outputs, batch, batch_idx\n    ):\n        global_step = trainer.global_step\n        if global_step > 0 and (global_step % self.print_every_n_steps == 0):\n            log_metrics = trainer.callback_metrics \n            loss = log_metrics.get('train_loss_step')\n            dice = log_metrics.get('train_dice_step')\n            print(\n                f\"[step {global_step}] batch log | \"\n                f\"loss: {loss:.4f}, dice: {dice:.4f}\"\n            )\n\n    def on_train_epoch_end(self, trainer, pl_module):\n        log_metrics = trainer.callback_metrics \n        loss = log_metrics.get('train_loss_epoch')\n        dice = log_metrics.get('train_dice_epoch')\n        print(\n            f\"train loss: {loss:.4f}, train dice: {dice:.4f}\"\n        )\n        print(\"-\" * 50)\n\n    def on_validation_epoch_end(self, trainer, pl_module):\n        val_loss_key = 'val_loss' \n        val_dice_key = 'val_dice'\n        val_loss = (\n            trainer.callback_metrics.get(val_loss_key) or \n            trainer.callback_metrics.get(f'{val_loss_key}_epoch')\n        )\n        val_dice = (\n            trainer.callback_metrics.get(val_dice_key) or \n            trainer.callback_metrics.get(f'{val_dice_key}_epoch')\n        )\n        print(\n            f\"epoch {trainer.current_epoch + 1:02d} completed: \"\n            f\"val loss: {val_loss:.4f}, val dice: {val_dice:.4f}\"\n        )","metadata":{"trusted":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2026-01-30T11:51:42.897970Z","iopub.execute_input":"2026-01-30T11:51:42.898283Z","iopub.status.idle":"2026-01-30T11:51:42.915230Z","shell.execute_reply.started":"2026-01-30T11:51:42.898238Z","shell.execute_reply":"2026-01-30T11:51:42.914630Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"callbacks = [\n    MetricLoggerCallback(),\n    ModelCheckpoint(\n        monitor='val_dice',\n        mode='max',\n        save_top_k=2,\n        save_last=True,\n        filename='best-{epoch:02d}-{val_dice:.3f}'\n    ),\n    LearningRateMonitor(logging_interval='epoch')\n]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-30T11:51:42.916212Z","iopub.execute_input":"2026-01-30T11:51:42.916499Z","iopub.status.idle":"2026-01-30T11:51:42.936682Z","shell.execute_reply.started":"2026-01-30T11:51:42.916461Z","shell.execute_reply":"2026-01-30T11:51:42.935812Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# create data module\ndata_module = VesuviusDataModule(\n    train_imgs=train_imgs,\n    train_lbls=train_lbls,\n    val_imgs=val_imgs,\n    val_lbls=val_lbls,\n    train_batch_size=batch_size,\n    val_batch_size=1,\n    num_workers=4,\n)\n\n# create the Lightning model\nlightning_model = VesuviusModel(\n    model=model,\n    learning_rate=1e-4,\n    sw_batch_size=2,\n    sw_overlap=0.5\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-30T11:51:42.937840Z","iopub.execute_input":"2026-01-30T11:51:42.938165Z","iopub.status.idle":"2026-01-30T11:51:42.951494Z","shell.execute_reply.started":"2026-01-30T11:51:42.938126Z","shell.execute_reply":"2026-01-30T11:51:42.950810Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trainer = pl.Trainer(\n    max_epochs=epochs,\n    precision=\"16-mixed\",\n    callbacks=callbacks,\n    accelerator='auto',\n    devices='auto',\n    log_every_n_steps=1,\n    accumulate_grad_batches=4,\n    check_val_every_n_epoch=1,\n    enable_progress_bar=False,\n    num_sanity_val_steps=0\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-30T11:51:42.952632Z","iopub.execute_input":"2026-01-30T11:51:42.952982Z","iopub.status.idle":"2026-01-30T11:51:43.030096Z","shell.execute_reply.started":"2026-01-30T11:51:42.952948Z","shell.execute_reply":"2026-01-30T11:51:43.029444Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trainer.fit(lightning_model, datamodule=data_module)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-30T11:51:43.030936Z","iopub.execute_input":"2026-01-30T11:51:43.031259Z","iopub.status.idle":"2026-01-30T13:17:08.647968Z","shell.execute_reply.started":"2026-01-30T11:51:43.031231Z","shell.execute_reply":"2026-01-30T13:17:08.646456Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.save(model.state_dict(), \"model.pth\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-30T13:17:08.653077Z","iopub.execute_input":"2026-01-30T13:17:08.653374Z","iopub.status.idle":"2026-01-30T13:17:08.698300Z","shell.execute_reply.started":"2026-01-30T13:17:08.653343Z","shell.execute_reply":"2026-01-30T13:17:08.697442Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# for name, param in model.named_parameters():\n#     print(name, param.dtype)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-30T13:17:08.699390Z","iopub.execute_input":"2026-01-30T13:17:08.699744Z","iopub.status.idle":"2026-01-30T13:17:08.703312Z","shell.execute_reply.started":"2026-01-30T13:17:08.699703Z","shell.execute_reply":"2026-01-30T13:17:08.702650Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Quick Inference**","metadata":{}},{"cell_type":"code","source":"x, y = next(iter(check_out_loader))\nx.shape, y.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-30T13:17:08.704303Z","iopub.execute_input":"2026-01-30T13:17:08.704619Z","iopub.status.idle":"2026-01-30T13:17:14.384612Z","shell.execute_reply.started":"2026-01-30T13:17:08.704588Z","shell.execute_reply":"2026-01-30T13:17:14.383640Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"swi = SlidingWindowInference(\n    model.to(\"cuda\"),\n    num_classes=num_classes,\n    roi_size=input_shape,\n    sw_batch_size=2,\n    overlap=0.5,\n    mode='gaussian',\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-30T13:17:14.386689Z","iopub.execute_input":"2026-01-30T13:17:14.387071Z","iopub.status.idle":"2026-01-30T13:17:14.410751Z","shell.execute_reply.started":"2026-01-30T13:17:14.387036Z","shell.execute_reply":"2026-01-30T13:17:14.410050Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"y_pred = swi(x)\ny_pred.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-30T13:17:14.411820Z","iopub.execute_input":"2026-01-30T13:17:14.412050Z","iopub.status.idle":"2026-01-30T13:17:32.028489Z","shell.execute_reply.started":"2026-01-30T13:17:14.412026Z","shell.execute_reply":"2026-01-30T13:17:32.027862Z"}},"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-30T13:17:32.029520Z","iopub.execute_input":"2026-01-30T13:17:32.029828Z","iopub.status.idle":"2026-01-30T13:17:33.133315Z","shell.execute_reply.started":"2026-01-30T13:17:32.029800Z","shell.execute_reply":"2026-01-30T13:17:33.132591Z"}},"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-30T13:17:33.134418Z","iopub.execute_input":"2026-01-30T13:17:33.134831Z","iopub.status.idle":"2026-01-30T13:17:33.802046Z","shell.execute_reply.started":"2026-01-30T13:17:33.134801Z","shell.execute_reply":"2026-01-30T13:17:33.801219Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}