{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":9988,"databundleVersionId":868324,"sourceType":"competition"}],"dockerImageVersionId":30805,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport pandas as pd\nimport numpy as np\nimport cv2\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import train_test_split\nfrom torch.utils.data import Dataset, DataLoader\nimport torch\nimport torch.nn as nn\nfrom torchvision import datasets, transforms\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-20T19:42:31.158568Z","iopub.execute_input":"2025-01-20T19:42:31.158978Z","iopub.status.idle":"2025-01-20T19:42:34.885102Z","shell.execute_reply.started":"2025-01-20T19:42:31.158932Z","shell.execute_reply":"2025-01-20T19:42:34.884379Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"images_path = '/kaggle/input/airbus-ship-detection/train_v2'\ncsv_path = '/kaggle/input/airbus-ship-detection/train_ship_segmentations_v2.csv'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-20T19:42:34.886659Z","iopub.execute_input":"2025-01-20T19:42:34.887064Z","iopub.status.idle":"2025-01-20T19:42:34.891103Z","shell.execute_reply.started":"2025-01-20T19:42:34.887036Z","shell.execute_reply":"2025-01-20T19:42:34.890205Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_csv(csv_path)\nprint(\"Dataframe size:\", df.shape)\nprint(df.head())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-20T19:42:34.892251Z","iopub.execute_input":"2025-01-20T19:42:34.892632Z","iopub.status.idle":"2025-01-20T19:42:36.195106Z","shell.execute_reply.started":"2025-01-20T19:42:34.892590Z","shell.execute_reply":"2025-01-20T19:42:36.194196Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport matplotlib.pyplot as plt\n\n# Download CSV with data\ndf = pd.read_csv('/kaggle/input/airbus-ship-detection/train_ship_segmentations_v2.csv')\n\n# Clear data from rows with NaN in the EncodedPixels column\ndf_cleaned = df.dropna(subset=['EncodedPixels'])\n\n# Counting the number of ships in each image\nship_counts = df_cleaned.groupby('ImageId').size()\n\n# Display the number of ships for each image\nprint(ship_counts)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-20T19:42:36.196340Z","iopub.execute_input":"2025-01-20T19:42:36.196777Z","iopub.status.idle":"2025-01-20T19:42:36.847591Z","shell.execute_reply.started":"2025-01-20T19:42:36.196731Z","shell.execute_reply":"2025-01-20T19:42:36.846655Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_cleaned","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-20T19:42:36.849553Z","iopub.execute_input":"2025-01-20T19:42:36.849821Z","iopub.status.idle":"2025-01-20T19:42:36.861565Z","shell.execute_reply.started":"2025-01-20T19:42:36.849795Z","shell.execute_reply":"2025-01-20T19:42:36.860693Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df['ships_count'] = df.groupby('ImageId')['EncodedPixels'].transform('count')\n\n# Visualization of the distribution of the number of ships\nship_counts = df.groupby('ImageId')['ships_count'].max()\nplt.figure(figsize=(10, 6))\nship_counts.hist(bins=10)\nplt.title('Distribution of the number of ships per image (before cleaning)')\nplt.xlabel('Number of ships')\nplt.ylabel('Number of images')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-20T19:42:36.862515Z","iopub.execute_input":"2025-01-20T19:42:36.862770Z","iopub.status.idle":"2025-01-20T19:42:37.420707Z","shell.execute_reply.started":"2025-01-20T19:42:36.862745Z","shell.execute_reply":"2025-01-20T19:42:37.419799Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Clear data from rows with NaN in the EncodedPixels column\ndf_cleaned = df.dropna(subset=['EncodedPixels']).reset_index(drop=True)\n\n# Number of unique images with ships\nunique_images_with_ships = df_cleaned['ImageId'].nunique()\n\n# Counting the number of ships in each image (group by ImageId)\nship_counts = df_cleaned.groupby('ImageId').size()\n\n# Output of number of unique images with ships\nprint(\"Number of images with ships:\", unique_images_with_ships)\n\n# Visualization of the distribution of the number of ships in each image\nplt.figure(figsize=(10, 6))\nship_counts.hist(bins=10)\nplt.title('Distribution of the number of ships per image (after data cleaning)')\nplt.xlabel('Number of ships')\nplt.ylabel('Number of images')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-20T19:42:37.421774Z","iopub.execute_input":"2025-01-20T19:42:37.422035Z","iopub.status.idle":"2025-01-20T19:42:37.672093Z","shell.execute_reply.started":"2025-01-20T19:42:37.422009Z","shell.execute_reply":"2025-01-20T19:42:37.671154Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install segmentation-models-pytorch","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-20T19:42:37.673337Z","iopub.execute_input":"2025-01-20T19:42:37.673713Z","iopub.status.idle":"2025-01-20T19:42:51.494470Z","shell.execute_reply.started":"2025-01-20T19:42:37.673667Z","shell.execute_reply":"2025-01-20T19:42:51.493331Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport cv2\nfrom torch.utils.data import Dataset\nimport matplotlib.pyplot as plt\nfrom torchvision import transforms\nfrom PIL import Image\nfrom segmentation_models_pytorch.encoders import get_preprocessing_fn\n\nclass ShipDataset(Dataset):\n    def __init__(self, images_path, df_cleaned, height=768, width=768, image_ids=None):\n        self.images_path = images_path\n        self.df_cleaned = df_cleaned\n        self.height = height\n        self.width = width\n        self.image_ids = image_ids if image_ids is not None else df_cleaned['ImageId'].unique()\n        \n    def __len__(self):\n        return len(self.image_ids)\n    \n    def rle_to_mask(self, rle, height, width):\n        \"\"\"\n        Converts Run-Length Encoding to a height x width mask.\n        :param rle: string containing Run-Length Encoding (pairs start, length)\n        :param height: height of the mask\n        :param width: width of the mask\n        :return: mask in the form of a numpy array\n        \"\"\"\n        mask = np.zeros(height * width, dtype=np.uint8)  # the mask is initially empty\n        rle_values = list(map(int, rle.split()))  # divide the string into numbers\n        \n        for i in range(0, len(rle_values), 2):\n            start = rle_values[i] - 1  # the starting pixel\n            length = rle_values[i+1]  # segment length\n            \n            # translate the linear index into two-dimensional coordinates (row, column)\n            start_row = start // width\n            start_col = start % width\n            \n            # Fill the mask with the appropriate pixels\n            for j in range(length):\n                row = (start + j) // width  \n                col = (start + j) % width \n                \n                # Checking whether we are not going beyond the mask\n                if row < height and col < width:\n                    mask[row * width + col] = 1  # set the pixel to 1\n\n        return mask.reshape((height, width)).T  # convert the mask into 2D format (height, width)\n\n    def combine_masks(self, masks, height, width):\n        \"\"\"\n        Combines masks for several ships into one\n        \"\"\"\n        if masks:\n            combined_mask = np.zeros((height, width), dtype=np.uint8)\n            for mask in masks:\n                combined_mask = np.maximum(combined_mask, mask)\n        else:\n            combined_mask = np.zeros((height, width), dtype=np.uint8)\n        return combined_mask\n    \n    def split_image(self, image, mask, part_size=256):\n        \"\"\"\n        Cuts the image and mask into 9 parts 256x256\n        Selects the part with the most ship pixels\n        \"\"\"\n        best_part_idx = None\n        max_ship_pixels = 0\n\n        # Cut the image into 9 parts (3x3)\n        image_parts = []\n        mask_parts = []\n        \n        for i in range(3):\n            for j in range(3):\n                start_row = i * part_size\n                start_col = j * part_size\n                end_row = start_row + part_size\n                end_col = start_col + part_size\n\n                image_part = image[start_row:end_row, start_col:end_col]\n                mask_part = mask[start_row:end_row, start_col:end_col]\n\n                image_parts.append(image_part)\n                mask_parts.append(mask_part)\n\n                # count the number of pixels of ships (1 in the mask)\n                ship_pixels = np.sum(mask_part)\n                if ship_pixels > max_ship_pixels:\n                    max_ship_pixels = ship_pixels\n                    best_part_idx = len(image_parts) - 1\n\n        # return best part of the image and the mask\n        return image_parts[best_part_idx], mask_parts[best_part_idx]\n    \n    def __getitem__(self, idx):\n        image_id = self.image_ids[idx]\n        \n        image_path = f\"{self.images_path}/{image_id}\"\n        image = Image.open(image_path).convert('RGB')\n        image = np.array(image.resize((self.width, self.height))) \n\n        # load masks for the current image\n        masks = []\n        encoded_pixels = self.df_cleaned[self.df_cleaned['ImageId'] == image_id]['EncodedPixels']\n        \n        for rle in encoded_pixels:\n            mask = self.rle_to_mask(rle, self.height, self.width)\n            masks.append(mask)\n        \n        # combine masks for all ships\n        combined_mask = self.combine_masks(masks, self.height, self.width)\n        \n        # cut the image and the mask into 9 parts and choose the best part\n        best_image_part, best_mask_part = self.split_image(image, combined_mask)\n\n        preprocess_input = get_preprocessing_fn('resnet34', pretrained='imagenet')\n    \n        # apply preprocess_input to the image\n        best_image_part = preprocess_input(best_image_part)\n        \n        # Converting an image into a tensor\n        transform = transforms.ToTensor()\n        best_image_part = transform(best_image_part)\n        \n        return best_image_part, best_mask_part\n\nimages_path = '/kaggle/input/airbus-ship-detection/train_v2'\ndataset = ShipDataset(images_path=images_path, df_cleaned=df_cleaned)\n\n# display the image and mask for checking\nimage, mask = dataset[2]\n\n# display the image\nplt.figure(figsize=(12, 6))\nplt.subplot(1, 2, 1)\nplt.imshow(image.permute(1, 2, 0))  # convert from (C, H, W) to (H, W, C)\nplt.title('Best Image Part')\n\n# display the mask\nplt.subplot(1, 2, 2)\nplt.imshow(mask, cmap='gray')\nplt.title('Best Mask Part')\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-20T20:03:45.137902Z","iopub.execute_input":"2025-01-20T20:03:45.138239Z","iopub.status.idle":"2025-01-20T20:03:45.580470Z","shell.execute_reply.started":"2025-01-20T20:03:45.138212Z","shell.execute_reply":"2025-01-20T20:03:45.579606Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\n\nmain_ids, small_ids = train_test_split(dataset.image_ids, test_size=0.1, random_state=42)\n\n# divide the 10% part into training and test\ntrain_ids, test_ids = train_test_split(small_ids, test_size=0.2, random_state=42)\n\n# Creating a DataFrame for training and test data\ntrain_df = df_cleaned[df_cleaned['ImageId'].isin(train_ids)]\ntest_df = df_cleaned[df_cleaned['ImageId'].isin(test_ids)]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-20T19:42:55.699904Z","iopub.execute_input":"2025-01-20T19:42:55.700184Z","iopub.status.idle":"2025-01-20T19:42:55.727677Z","shell.execute_reply.started":"2025-01-20T19:42:55.700156Z","shell.execute_reply":"2025-01-20T19:42:55.726928Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f\"Train size: {train_df.shape[0]}\")\nprint(f\"Test size: {test_df.shape[0]}\") ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-20T19:42:55.728710Z","iopub.execute_input":"2025-01-20T19:42:55.728990Z","iopub.status.idle":"2025-01-20T19:42:55.734091Z","shell.execute_reply.started":"2025-01-20T19:42:55.728963Z","shell.execute_reply":"2025-01-20T19:42:55.733106Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import DataLoader\n\ntrain_dataset = ShipDataset(image_ids=train_ids, df_cleaned=train_df, images_path=images_path)\ntest_dataset = ShipDataset(image_ids=test_ids, df_cleaned=test_df, images_path=images_path)\n\ntrain_loader = DataLoader(train_dataset, batch_size=128, shuffle=True)\ntest_loader = DataLoader(test_dataset, batch_size=128, shuffle=False)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-20T19:42:55.735233Z","iopub.execute_input":"2025-01-20T19:42:55.735520Z","iopub.status.idle":"2025-01-20T19:42:55.747513Z","shell.execute_reply.started":"2025-01-20T19:42:55.735493Z","shell.execute_reply":"2025-01-20T19:42:55.746729Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"len(train_loader), len(test_loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-20T19:42:55.748462Z","iopub.execute_input":"2025-01-20T19:42:55.748723Z","iopub.status.idle":"2025-01-20T19:42:55.763272Z","shell.execute_reply.started":"2025-01-20T19:42:55.748697Z","shell.execute_reply":"2025-01-20T19:42:55.762429Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import segmentation_models_pytorch as smp\nimport torch\n\n# Create model UNet\nmodel = smp.Unet(\n    encoder_name=\"resnet34\", \n    encoder_weights=\"imagenet\",  \n    in_channels=3,\n    classes=1,\n)\n\nmodel\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-20T19:42:55.764397Z","iopub.execute_input":"2025-01-20T19:42:55.764754Z","iopub.status.idle":"2025-01-20T19:42:56.840251Z","shell.execute_reply.started":"2025-01-20T19:42:55.764716Z","shell.execute_reply":"2025-01-20T19:42:56.839403Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Freeze the encoder (set requires_grad = False for its parameters)\nfor param in model.encoder.parameters():\n    param.requires_grad = False\n\n# check whether the encoder is really frozen\nfor name, param in model.named_parameters():\n    print(f\"{name}: requires_grad={param.requires_grad}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-20T19:42:56.841395Z","iopub.execute_input":"2025-01-20T19:42:56.841664Z","iopub.status.idle":"2025-01-20T19:42:56.848444Z","shell.execute_reply.started":"2025-01-20T19:42:56.841637Z","shell.execute_reply":"2025-01-20T19:42:56.847602Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -q torchsummary","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-20T19:42:56.849336Z","iopub.execute_input":"2025-01-20T19:42:56.849677Z","iopub.status.idle":"2025-01-20T19:43:05.675390Z","shell.execute_reply.started":"2025-01-20T19:42:56.849640Z","shell.execute_reply":"2025-01-20T19:43:05.674110Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torchsummary import summary\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = model.to(device)\n\nsummary(model, input_size=(3, 256, 256))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-20T19:43:05.680057Z","iopub.execute_input":"2025-01-20T19:43:05.680414Z","iopub.status.idle":"2025-01-20T19:43:06.527792Z","shell.execute_reply.started":"2025-01-20T19:43:05.680377Z","shell.execute_reply":"2025-01-20T19:43:06.526851Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !pip install kornia","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-20T20:15:37.428244Z","iopub.execute_input":"2025-01-20T20:15:37.429127Z","iopub.status.idle":"2025-01-20T20:15:37.432817Z","shell.execute_reply.started":"2025-01-20T20:15:37.429091Z","shell.execute_reply":"2025-01-20T20:15:37.431910Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"loss_fn = smp.losses.DiceLoss(mode=\"binary\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-20T19:43:14.755420Z","iopub.execute_input":"2025-01-20T19:43:14.755943Z","iopub.status.idle":"2025-01-20T19:43:14.761100Z","shell.execute_reply.started":"2025-01-20T19:43:14.755897Z","shell.execute_reply":"2025-01-20T19:43:14.760319Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"optimizer = torch.optim.Adam(filter(lambda p: p.requires_grad, model.parameters()), lr=0.0005)\n\ntrainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\nnon_trainable_params = sum(p.numel() for p in model.parameters() if not p.requires_grad)\n\nprint(f\"Trainable parameters: {trainable_params}\")\nprint(f\"Non-trainable parameters: {non_trainable_params}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-20T19:43:14.762450Z","iopub.execute_input":"2025-01-20T19:43:14.762791Z","iopub.status.idle":"2025-01-20T19:43:15.029566Z","shell.execute_reply.started":"2025-01-20T19:43:14.762753Z","shell.execute_reply":"2025-01-20T19:43:15.028503Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import gc\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = model.to(device)\n\nnum_epochs = 15\ntrain_loss_history = []\nval_loss_history = []\n\nfor epoch in range(num_epochs):\n    model.train()\n    train_loss = 0.0\n    \n    # Training step\n    for images, masks in train_loader:\n\n        images = images.to(device).float()\n        masks = masks.to(device).float()\n\n        # Forward pass calculation\n        optimizer.zero_grad()\n        outputs = model(images)\n        loss = loss_fn(outputs, masks)\n        loss.backward()\n        optimizer.step()\n\n        train_loss += loss.item()\n\n    train_loss_avg = train_loss / len(train_loader)\n    train_loss_history.append(train_loss_avg)\n\n    torch.cuda.empty_cache()\n    gc.collect()\n\n    # Assessment on validation data\n    model.eval()\n    val_loss = 0.0\n    with torch.no_grad():\n        for images, masks in test_loader:\n            # images = images.to(device)\n            # masks = masks.to(device)\n            images = images.to(device).float()\n            masks = masks.to(device).float()\n\n            outputs = model(images)\n            loss = loss_fn(outputs, masks)\n            val_loss += loss.item()\n\n    val_loss_avg = val_loss / len(test_loader)\n    val_loss_history.append(val_loss_avg)\n\n    torch.cuda.empty_cache()\n    gc.collect()\n\n    # Output of the results of each stage\n    print(f\"Epoch [{epoch+1}/{num_epochs}], \"\n          f\"Train Loss: {train_loss_avg:.4f}, \"\n          f\"Validation Loss: {val_loss_avg:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-20T19:43:15.030832Z","iopub.execute_input":"2025-01-20T19:43:15.031120Z","iopub.status.idle":"2025-01-20T20:01:51.823993Z","shell.execute_reply.started":"2025-01-20T19:43:15.031090Z","shell.execute_reply":"2025-01-20T20:01:51.823169Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nnum_epochs = 15\nplt.plot(range(1, num_epochs+1), train_loss_history, label=\"Train Loss\")\nplt.plot(range(1, num_epochs+1), val_loss_history, label=\"Validation Loss\")\nplt.xlabel('Epochs')\nplt.ylabel('Loss')\nplt.title('Loss Curve')\nplt.legend()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-20T20:01:51.825613Z","iopub.execute_input":"2025-01-20T20:01:51.825983Z","iopub.status.idle":"2025-01-20T20:01:52.083084Z","shell.execute_reply.started":"2025-01-20T20:01:51.825942Z","shell.execute_reply":"2025-01-20T20:01:52.082250Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def denormalize_image(image):\n    \"\"\"\n    Denormalizes the image after `preprocess_input` on the GPU or CPU.\n    \"\"\"\n    mean = [0.485, 0.456, 0.406]\n    std = [0.229, 0.224, 0.225]\n    \n    # transfer mean and std to the same device where the images are located\n    mean = torch.tensor(mean).reshape(1, 3, 1, 1).to(image.device)\n    std = torch.tensor(std).reshape(1, 3, 1, 1).to(image.device)\n\n    # Denormalization\n    image = image * std + mean  # Reverse scaling\n    return image\n\n\ndef visualize_results(model, test_loader, device, num_images=15):\n    \"\"\"\n    Visualizes model results on a test dataset using denormalization for images.\n    \"\"\"\n    model.eval()  # transfer the model to the evaluation mode\n    test_iter = iter(test_loader)\n    \n    images_shown = 0  \n    \n    with torch.no_grad():  \n        while images_shown < num_images:\n            images, true_masks = next(test_iter)\n            \n            # Convert to float and transfer to GPU\n            images = images.to(device).float()  \n            true_masks = true_masks.to(device).float() \n            \n            # Prediction of masks\n            pred_masks = model(images)  # Model call\n            pred_masks = torch.sigmoid(pred_masks)  # Apply sigmoid to translate to [0, 1]\n            pred_masks = (pred_masks > 0.5).float()  # Let's binarize\n            \n            # Denormalization of images for visualization\n            images_denormalized = denormalize_image(images)\n            \n            # Visualization\n            for j in range(len(images)):\n                if images_shown >= num_images:\n                    break\n                \n                fig, ax = plt.subplots(1, 3, figsize=(15, 5))\n                \n                # Original image (denormalized)\n                ax[0].imshow(images_denormalized[j].cpu().permute(1, 2, 0).clip(0, 1))  # Transfer from GPU to CPU\n                ax[0].set_title(\"Original Image (Denormalized)\")\n                ax[0].axis(\"off\")\n                \n                # Real mask\n                ax[1].imshow(true_masks[j].cpu().squeeze(), cmap=\"gray\")  # Transfer from GPU to CPU                \n                ax[1].set_title(\"True Mask\")\n                ax[1].axis(\"off\")\n                \n                # A mask is provided\n                ax[2].imshow(pred_masks[j].cpu().squeeze(), cmap=\"gray\")  # Transfer from GPU to CPU\n                ax[2].set_title(\"Predicted Mask\")\n                ax[2].axis(\"off\")\n                \n                images_shown += 1 \n                plt.show()\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel.to(device)\nvisualize_results(model, test_loader, device, num_images=15)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-20T20:01:52.084477Z","iopub.execute_input":"2025-01-20T20:01:52.084850Z","iopub.status.idle":"2025-01-20T20:01:58.920474Z","shell.execute_reply.started":"2025-01-20T20:01:52.084809Z","shell.execute_reply":"2025-01-20T20:01:58.919610Z"}},"outputs":[],"execution_count":null}]}