{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.7.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":6927,"databundleVersionId":45059,"sourceType":"competition"},{"sourceId":3009448,"sourceType":"datasetVersion","datasetId":1843391}],"dockerImageVersionId":30198,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Please refer to my [Medium article](https://medium.com/@fernandopalominocobo/mastering-u-net-a-step-by-step-guide-to-segmentation-from-scratch-with-pytorch-6a17c5916114) for code explanations!!","metadata":{}},{"cell_type":"markdown","source":"# Load the required libraries","metadata":{}},{"cell_type":"code","source":"import copy\nimport os\nimport random\nimport shutil\nimport zipfile\nfrom math import atan2, cos, sin, sqrt, pi, log\n\nimport cv2\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.transforms as transforms\nfrom PIL import Image\nfrom numpy import linalg as LA\nfrom torch import optim, nn\nfrom torch.utils.data import DataLoader, random_split\nfrom torch.utils.data.dataset import Dataset\nfrom torchvision import transforms\nfrom tqdm import tqdm","metadata":{"execution":{"iopub.status.busy":"2025-04-04T14:22:49.790141Z","iopub.execute_input":"2025-04-04T14:22:49.790670Z","iopub.status.idle":"2025-04-04T14:22:52.169725Z","shell.execute_reply.started":"2025-04-04T14:22:49.790554Z","shell.execute_reply":"2025-04-04T14:22:52.168669Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# U-Net architecture","metadata":{}},{"cell_type":"code","source":"class DoubleConv(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n        self.conv_op = nn.Sequential(\n            nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1),\n            nn.ReLU(inplace=True)\n        )\n\n    def forward(self, x):\n        return self.conv_op(x)\n","metadata":{"execution":{"iopub.status.busy":"2025-04-04T14:22:52.171683Z","iopub.execute_input":"2025-04-04T14:22:52.172455Z","iopub.status.idle":"2025-04-04T14:22:52.182663Z","shell.execute_reply.started":"2025-04-04T14:22:52.172410Z","shell.execute_reply":"2025-04-04T14:22:52.181914Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DownSample(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n        self.conv = DoubleConv(in_channels, out_channels)\n        self.pool = nn.MaxPool2d(kernel_size=2, stride=2)\n\n    def forward(self, x):\n        down = self.conv(x)\n        p = self.pool(down)\n\n        return down, p","metadata":{"execution":{"iopub.status.busy":"2025-04-04T14:22:52.183591Z","iopub.execute_input":"2025-04-04T14:22:52.183905Z","iopub.status.idle":"2025-04-04T14:22:52.197268Z","shell.execute_reply.started":"2025-04-04T14:22:52.183876Z","shell.execute_reply":"2025-04-04T14:22:52.196448Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class UpSample(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n        self.up = nn.ConvTranspose2d(in_channels, in_channels//2, kernel_size=2, stride=2)\n        self.conv = DoubleConv(in_channels, out_channels)\n\n    def forward(self, x1, x2):\n        x1 = self.up(x1)\n        x = torch.cat([x1, x2], 1)\n        return self.conv(x)","metadata":{"execution":{"iopub.status.busy":"2025-04-04T14:22:52.198404Z","iopub.execute_input":"2025-04-04T14:22:52.198784Z","iopub.status.idle":"2025-04-04T14:22:52.209357Z","shell.execute_reply.started":"2025-04-04T14:22:52.198750Z","shell.execute_reply":"2025-04-04T14:22:52.208545Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class UNet(nn.Module):\n    def __init__(self, in_channels, num_classes):\n        super().__init__()\n        self.down_convolution_1 = DownSample(in_channels, 64)\n        self.down_convolution_2 = DownSample(64, 128)\n        self.down_convolution_3 = DownSample(128, 256)\n        self.down_convolution_4 = DownSample(256, 512)\n\n        self.bottle_neck = DoubleConv(512, 1024)\n\n        self.up_convolution_1 = UpSample(1024, 512)\n        self.up_convolution_2 = UpSample(512, 256)\n        self.up_convolution_3 = UpSample(256, 128)\n        self.up_convolution_4 = UpSample(128, 64)\n\n        self.out = nn.Conv2d(in_channels=64, out_channels=num_classes, kernel_size=1)\n\n    def forward(self, x):\n        down_1, p1 = self.down_convolution_1(x)\n        down_2, p2 = self.down_convolution_2(p1)\n        down_3, p3 = self.down_convolution_3(p2)\n        down_4, p4 = self.down_convolution_4(p3)\n\n        b = self.bottle_neck(p4)\n\n        up_1 = self.up_convolution_1(b, down_4)\n        up_2 = self.up_convolution_2(up_1, down_3)\n        up_3 = self.up_convolution_3(up_2, down_2)\n        up_4 = self.up_convolution_4(up_3, down_1)\n\n        out = self.out(up_4)\n        return out\n","metadata":{"execution":{"iopub.status.busy":"2025-04-04T14:22:52.210634Z","iopub.execute_input":"2025-04-04T14:22:52.210982Z","iopub.status.idle":"2025-04-04T14:22:52.221614Z","shell.execute_reply.started":"2025-04-04T14:22:52.210948Z","shell.execute_reply":"2025-04-04T14:22:52.220826Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"input_image = torch.rand((1,3,512,512))\nmodel = UNet(3,10)\noutput = model(input_image)\nprint(output.size())\n# You should get torch.Size([1, 10, 512, 512]) as a result","metadata":{"execution":{"iopub.status.busy":"2025-04-04T14:22:52.222581Z","iopub.execute_input":"2025-04-04T14:22:52.222876Z","iopub.status.idle":"2025-04-04T14:22:56.207237Z","shell.execute_reply.started":"2025-04-04T14:22:52.222853Z","shell.execute_reply":"2025-04-04T14:22:56.206432Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Load the Carvana Dataset","metadata":{}},{"cell_type":"code","source":"class CarvanaDataset(Dataset):\n    def __init__(self, root_path, limit=None):\n        self.root_path = root_path\n        self.limit = limit\n        self.images = sorted([root_path + \"/train/\" + i for i in os.listdir(root_path + \"/train/\")])[:self.limit]\n        self.masks = sorted([root_path + \"/train_masks/\" + i for i in os.listdir(root_path + \"/train_masks/\")])[:self.limit]\n\n        self.transform = transforms.Compose([\n            transforms.Resize((512, 512)),\n            transforms.ToTensor()])\n        \n        if self.limit is None:\n            self.limit = len(self.images)\n\n    def __getitem__(self, index):\n        img = Image.open(self.images[index]).convert(\"RGB\")\n        mask = Image.open(self.masks[index]).convert(\"L\")\n\n        return self.transform(img), self.transform(mask), self.images[index]\n\n    def __len__(self):\n        \n        return min(len(self.images), self.limit)","metadata":{"execution":{"iopub.status.busy":"2025-04-04T14:22:56.209938Z","iopub.execute_input":"2025-04-04T14:22:56.210641Z","iopub.status.idle":"2025-04-04T14:22:56.217891Z","shell.execute_reply.started":"2025-04-04T14:22:56.210610Z","shell.execute_reply":"2025-04-04T14:22:56.216904Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nprint(os.listdir(\"../input/carvana-image-masking-challenge/\"))\n\nDATASET_DIR = '../input/carvana-image-masking-challenge/'\nWORKING_DIR = '/kaggle/working/'","metadata":{"execution":{"iopub.status.busy":"2025-04-04T14:22:56.218899Z","iopub.execute_input":"2025-04-04T14:22:56.219177Z","iopub.status.idle":"2025-04-04T14:22:56.238113Z","shell.execute_reply.started":"2025-04-04T14:22:56.219153Z","shell.execute_reply":"2025-04-04T14:22:56.237215Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if len(os.listdir(WORKING_DIR)) <= 1:\n\n    with zipfile.ZipFile(DATASET_DIR + 'train.zip', 'r') as zip_file:\n        zip_file.extractall(WORKING_DIR)\n\n    with zipfile.ZipFile(DATASET_DIR + 'train_masks.zip', 'r') as zip_file:\n        zip_file.extractall(WORKING_DIR)\n    \n    print(\n        len(os.listdir(WORKING_DIR + 'train')),\n        len(os.listdir(WORKING_DIR + 'train_masks'))\n    )","metadata":{"execution":{"iopub.status.busy":"2025-04-04T14:22:56.239168Z","iopub.execute_input":"2025-04-04T14:22:56.239569Z","iopub.status.idle":"2025-04-04T14:23:04.858631Z","shell.execute_reply.started":"2025-04-04T14:22:56.239528Z","shell.execute_reply":"2025-04-04T14:23:04.857950Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset = CarvanaDataset(WORKING_DIR)\n\ngenerator = torch.Generator().manual_seed(25)\n","metadata":{"execution":{"iopub.status.busy":"2025-04-04T14:23:04.859849Z","iopub.execute_input":"2025-04-04T14:23:04.860537Z","iopub.status.idle":"2025-04-04T14:23:04.872966Z","shell.execute_reply.started":"2025-04-04T14:23:04.860497Z","shell.execute_reply":"2025-04-04T14:23:04.872338Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import random_split\n\ndataset_len = len(train_dataset)  # 原始是5088\ntrain_len = int(0.8 * dataset_len)  # 4060\ntemp_len = dataset_len - train_len  # 1028\n\n# 第一步：80%训练 + 20%临时数据\ntrain_dataset, temp_dataset = random_split(train_dataset, [train_len, temp_len], generator=generator)\n\n# 第二步：把临时数据再分成50%验证 + 50%测试\nval_len = test_len = temp_len // 2  # 各 514，如果是奇数，可以调整\nval_dataset, test_dataset = random_split(temp_dataset, [val_len, test_len], generator=generator)","metadata":{"execution":{"iopub.status.busy":"2025-04-04T14:26:11.880882Z","iopub.execute_input":"2025-04-04T14:26:11.881783Z","iopub.status.idle":"2025-04-04T14:26:11.887653Z","shell.execute_reply.started":"2025-04-04T14:26:11.881747Z","shell.execute_reply":"2025-04-04T14:26:11.886847Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint(device)\n\nif device == \"cuda\":\n    num_workers = torch.cuda.device_count() * 4","metadata":{"execution":{"iopub.status.busy":"2025-04-04T14:28:01.032345Z","iopub.execute_input":"2025-04-04T14:28:01.033068Z","iopub.status.idle":"2025-04-04T14:28:01.037558Z","shell.execute_reply.started":"2025-04-04T14:28:01.033032Z","shell.execute_reply":"2025-04-04T14:28:01.036771Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"LEARNING_RATE = 3e-4\nBATCH_SIZE = 8","metadata":{"execution":{"iopub.status.busy":"2025-04-04T14:28:03.527369Z","iopub.execute_input":"2025-04-04T14:28:03.528197Z","iopub.status.idle":"2025-04-04T14:28:03.531716Z","shell.execute_reply.started":"2025-04-04T14:28:03.528161Z","shell.execute_reply":"2025-04-04T14:28:03.530977Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataloader = DataLoader(dataset=train_dataset,\n                              num_workers=num_workers, pin_memory=False,\n                              batch_size=BATCH_SIZE,\n                              shuffle=True)\nval_dataloader = DataLoader(dataset=val_dataset,\n                            num_workers=num_workers, pin_memory=False,\n                            batch_size=BATCH_SIZE,\n                            shuffle=True)\n\ntest_dataloader = DataLoader(dataset=test_dataset,\n                            num_workers=num_workers, pin_memory=False,\n                            batch_size=BATCH_SIZE,\n                            shuffle=True)\n\nmodel = UNet(in_channels=3, num_classes=1).to(device)\noptimizer = optim.AdamW(model.parameters(), lr=LEARNING_RATE)\ncriterion = nn.BCEWithLogitsLoss()","metadata":{"execution":{"iopub.status.busy":"2025-04-04T14:28:06.130544Z","iopub.execute_input":"2025-04-04T14:28:06.130908Z","iopub.status.idle":"2025-04-04T14:28:09.306955Z","shell.execute_reply.started":"2025-04-04T14:28:06.130877Z","shell.execute_reply":"2025-04-04T14:28:09.306333Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Running the experiment","metadata":{}},{"cell_type":"code","source":"\ndef dice_coefficient(prediction, target, epsilon=1e-07):\n    prediction_copy = prediction.clone()\n\n    prediction_copy[prediction_copy < 0] = 0\n    prediction_copy[prediction_copy > 0] = 1\n\n    intersection = abs(torch.sum(prediction_copy * target))\n    union = abs(torch.sum(prediction_copy) + torch.sum(target))\n    dice = (2. * intersection + epsilon) / (union + epsilon)\n    \n    return dice","metadata":{"execution":{"iopub.status.busy":"2025-04-04T14:28:13.135743Z","iopub.execute_input":"2025-04-04T14:28:13.136106Z","iopub.status.idle":"2025-04-04T14:28:13.141347Z","shell.execute_reply.started":"2025-04-04T14:28:13.136076Z","shell.execute_reply":"2025-04-04T14:28:13.140433Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2025-04-04T14:28:15.263617Z","iopub.execute_input":"2025-04-04T14:28:15.264232Z","iopub.status.idle":"2025-04-04T14:28:15.268066Z","shell.execute_reply.started":"2025-04-04T14:28:15.264194Z","shell.execute_reply":"2025-04-04T14:28:15.267241Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"EPOCHS = 10 # 10\n\ntrain_losses = []\ntrain_dcs = []\nval_losses = []\nval_dcs = []\n\nfor epoch in tqdm(range(EPOCHS)):\n    model.train()\n    train_running_loss = 0\n    train_running_dc = 0\n    \n    for idx, img_mask in enumerate(tqdm(train_dataloader, position=0, leave=True)):\n        img = img_mask[0].float().to(device)\n        mask = img_mask[1].float().to(device)\n        \n        y_pred = model(img)\n        optimizer.zero_grad()\n        \n        dc = dice_coefficient(y_pred, mask)\n        loss = criterion(y_pred, mask)\n        \n        train_running_loss += loss.item()\n        train_running_dc += dc.item()\n\n        loss.backward()\n        optimizer.step()\n\n    train_loss = train_running_loss / (idx + 1)\n    train_dc = train_running_dc / (idx + 1)\n    \n    train_losses.append(train_loss)\n    train_dcs.append(train_dc)\n\n    model.eval()\n    val_running_loss = 0\n    val_running_dc = 0\n    \n    with torch.no_grad():\n        for idx, img_mask in enumerate(tqdm(val_dataloader, position=0, leave=True)):\n            img = img_mask[0].float().to(device)\n            mask = img_mask[1].float().to(device)\n\n            y_pred = model(img)\n            loss = criterion(y_pred, mask)\n            dc = dice_coefficient(y_pred, mask)\n            \n            val_running_loss += loss.item()\n            val_running_dc += dc.item()\n\n        val_loss = val_running_loss / (idx + 1)\n        val_dc = val_running_dc / (idx + 1)\n    \n    val_losses.append(val_loss)\n    val_dcs.append(val_dc)\n\n    print(\"-\" * 30)\n    print(f\"Training Loss EPOCH {epoch + 1}: {train_loss:.4f}\")\n    print(f\"Training DICE EPOCH {epoch + 1}: {train_dc:.4f}\")\n    print(\"\\n\")\n    print(f\"Validation Loss EPOCH {epoch + 1}: {val_loss:.4f}\")\n    print(f\"Validation DICE EPOCH {epoch + 1}: {val_dc:.4f}\")\n    print(\"-\" * 30)\n\n# Guardar el modelo\ntorch.save(model.state_dict(), 'my_checkpoint.pth')\n","metadata":{"execution":{"iopub.status.busy":"2025-04-04T14:28:21.761017Z","iopub.execute_input":"2025-04-04T14:28:21.761416Z","iopub.status.idle":"2025-04-04T14:34:50.111524Z","shell.execute_reply.started":"2025-04-04T14:28:21.761382Z","shell.execute_reply":"2025-04-04T14:34:50.110658Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Results","metadata":{}},{"cell_type":"code","source":"epochs_list = list(range(1, EPOCHS + 1))\n\nplt.figure(figsize=(12, 5))\nplt.subplot(1, 2, 1)\nplt.plot(epochs_list, train_losses, label='Training Loss')\nplt.plot(epochs_list, val_losses, label='Validation Loss')\nplt.xticks(ticks=list(range(1, EPOCHS + 1, 1))) \nplt.title('Loss over epochs')\nplt.xlabel('Epochs')\nplt.ylabel('Loss')\nplt.grid()\nplt.tight_layout()\n\nplt.legend()\n\n\nplt.subplot(1, 2, 2)\nplt.plot(epochs_list, train_dcs, label='Training DICE')\nplt.plot(epochs_list, val_dcs, label='Validation DICE')\nplt.xticks(ticks=list(range(1, EPOCHS + 1, 1)))  \nplt.title('DICE Coefficient over epochs')\nplt.xlabel('Epochs')\nplt.ylabel('DICE')\nplt.grid()\nplt.legend()\n\nplt.tight_layout()\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2025-04-04T14:36:47.617575Z","iopub.execute_input":"2025-04-04T14:36:47.618409Z","iopub.status.idle":"2025-04-04T14:36:47.977354Z","shell.execute_reply.started":"2025-04-04T14:36:47.618368Z","shell.execute_reply":"2025-04-04T14:36:47.976530Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"epochs_list = list(range(1, EPOCHS + 1))\n\nplt.figure(figsize=(12, 5))\nplt.plot(epochs_list, train_losses, label='Training Loss')\nplt.plot(epochs_list, val_losses, label='Validation Loss')\nplt.xticks(ticks=list(range(1, EPOCHS + 1, 1))) \nplt.ylim(0, 0.05)\nplt.title('Loss over epochs (zoomed)')\nplt.xlabel('Epochs')\nplt.ylabel('Loss')\nplt.grid()\nplt.tight_layout()\n\nplt.legend()\nplt.show()\n\n","metadata":{"execution":{"iopub.status.busy":"2025-04-04T14:36:51.865429Z","iopub.execute_input":"2025-04-04T14:36:51.866145Z","iopub.status.idle":"2025-04-04T14:36:52.078627Z","shell.execute_reply.started":"2025-04-04T14:36:51.866112Z","shell.execute_reply":"2025-04-04T14:36:52.077836Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_pth = '/kaggle/working/my_checkpoint.pth'\ntrained_model = UNet(in_channels=3, num_classes=1).to(device)\ntrained_model.load_state_dict(torch.load(model_pth, map_location=torch.device(device)))","metadata":{"execution":{"iopub.status.busy":"2025-04-04T14:36:54.902547Z","iopub.execute_input":"2025-04-04T14:36:54.902921Z","iopub.status.idle":"2025-04-04T14:36:55.301911Z","shell.execute_reply.started":"2025-04-04T14:36:54.902890Z","shell.execute_reply":"2025-04-04T14:36:55.301051Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_running_loss = 0\ntest_running_dc = 0\n\nwith torch.no_grad():\n    for idx, img_mask in enumerate(tqdm(test_dataloader, position=0, leave=True)):\n        img = img_mask[0].float().to(device)\n        mask = img_mask[1].float().to(device)\n\n        y_pred = trained_model(img)\n        loss = criterion(y_pred, mask)\n        dc = dice_coefficient(y_pred, mask)\n\n        test_running_loss += loss.item()\n        test_running_dc += dc.item()\n\n    test_loss = test_running_loss / (idx + 1)\n    test_dc = test_running_dc / (idx + 1)\n","metadata":{"execution":{"iopub.status.busy":"2025-04-04T14:36:57.575197Z","iopub.execute_input":"2025-04-04T14:36:57.575578Z","iopub.status.idle":"2025-04-04T14:37:17.925193Z","shell.execute_reply.started":"2025-04-04T14:36:57.575548Z","shell.execute_reply":"2025-04-04T14:37:17.924225Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_loss","metadata":{"execution":{"iopub.status.busy":"2025-04-04T14:37:38.495556Z","iopub.execute_input":"2025-04-04T14:37:38.495952Z","iopub.status.idle":"2025-04-04T14:37:38.502012Z","shell.execute_reply.started":"2025-04-04T14:37:38.495915Z","shell.execute_reply":"2025-04-04T14:37:38.501195Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_dc","metadata":{"execution":{"iopub.status.busy":"2025-04-04T14:37:41.039874Z","iopub.execute_input":"2025-04-04T14:37:41.040655Z","iopub.status.idle":"2025-04-04T14:37:41.045658Z","shell.execute_reply.started":"2025-04-04T14:37:41.040608Z","shell.execute_reply":"2025-04-04T14:37:41.044868Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def random_images_inference(image_tensors, mask_tensors, image_paths, model_pth, device):\n    model = UNet(in_channels=3, num_classes=1).to(device)\n    model.load_state_dict(torch.load(model_pth, map_location=torch.device(device)))\n\n    transform = transforms.Compose([\n        transforms.Resize((512, 512))\n    ])\n\n    # Iterate for the images, masks and paths\n    for image_pth, mask_pth, image_paths in zip(image_tensors, mask_tensors, image_paths):\n        # Load the image\n        img = transform(image_pth)\n        \n        # Predict the imagen with the model\n        pred_mask = model(img.unsqueeze(0))\n        pred_mask = pred_mask.squeeze(0).permute(1,2,0)\n        \n        # Load the mask to compare\n        mask = transform(mask_pth).permute(1, 2, 0).to(device)\n        \n        print(f\"Image: {os.path.basename(image_paths)}, DICE coefficient: {round(float(dice_coefficient(pred_mask, mask)),5)}\")\n        \n        # Show the images\n        img = img.cpu().detach().permute(1, 2, 0)\n        pred_mask = pred_mask.cpu().detach()\n        pred_mask[pred_mask < 0] = 0\n        pred_mask[pred_mask > 0] = 1\n        \n        plt.figure(figsize=(15, 16))\n        plt.subplot(131), plt.imshow(img), plt.title(\"original\")\n        plt.subplot(132), plt.imshow(pred_mask, cmap=\"gray\"), plt.title(\"predicted\")\n        plt.subplot(133), plt.imshow(mask, cmap=\"gray\"), plt.title(\"mask\")\n        plt.show()\n","metadata":{"execution":{"iopub.status.busy":"2025-04-04T14:37:48.549439Z","iopub.execute_input":"2025-04-04T14:37:48.549804Z","iopub.status.idle":"2025-04-04T14:37:48.558897Z","shell.execute_reply.started":"2025-04-04T14:37:48.549773Z","shell.execute_reply":"2025-04-04T14:37:48.558099Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"n = 10 #10\n\nimage_tensors = []\nmask_tensors = []\nimage_paths = []\n\nfor _ in range(n):\n    random_index = random.randint(0, len(test_dataloader.dataset) - 1)\n    random_sample = test_dataloader.dataset[random_index]\n\n    image_tensors.append(random_sample[0])  \n    mask_tensors.append(random_sample[1]) \n    image_paths.append(random_sample[2]) \n\n","metadata":{"execution":{"iopub.status.busy":"2025-04-04T14:37:48.560046Z","iopub.execute_input":"2025-04-04T14:37:48.560359Z","iopub.status.idle":"2025-04-04T14:37:48.650415Z","shell.execute_reply.started":"2025-04-04T14:37:48.560333Z","shell.execute_reply":"2025-04-04T14:37:48.649700Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_path = '/kaggle/working/my_checkpoint.pth'\n\nrandom_images_inference(image_tensors, mask_tensors, image_paths, model_path, device=\"cpu\")","metadata":{"execution":{"iopub.status.busy":"2025-04-04T14:39:32.177241Z","iopub.execute_input":"2025-04-04T14:39:32.177685Z","iopub.status.idle":"2025-04-04T14:39:35.822128Z","shell.execute_reply.started":"2025-04-04T14:39:32.177654Z","shell.execute_reply":"2025-04-04T14:39:35.821343Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Optional: checking the masks in the training","metadata":{}},{"cell_type":"code","source":"n = 10 # 10\n\nimage_tensors = []\nmask_tensors = []\nimage_paths = []\n\nfor _ in range(n):\n    random_index = random.randint(0, len(train_dataloader.dataset) - 1)\n    random_sample = train_dataloader.dataset[random_index]\n\n    image_tensors.append(random_sample[0])  \n    mask_tensors.append(random_sample[1]) \n    image_paths.append(random_sample[2]) \n\nrandom_images_inference(image_tensors, mask_tensors, image_paths, model_path, device=\"cpu\")","metadata":{"execution":{"iopub.status.busy":"2025-04-04T14:39:51.801924Z","iopub.execute_input":"2025-04-04T14:39:51.802402Z","iopub.status.idle":"2025-04-04T14:40:24.039716Z","shell.execute_reply.started":"2025-04-04T14:39:51.802368Z","shell.execute_reply":"2025-04-04T14:40:24.038886Z"},"trusted":true},"outputs":[],"execution_count":null}]}