{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","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"}],"dockerImageVersionId":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"ac8ce650-d480-4434-a777-5e34cbdbdffb","_cell_guid":"47d20e50-b192-4068-a072-a3416601e411","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:07:12.765675Z","iopub.execute_input":"2025-04-16T12:07:12.766505Z","iopub.status.idle":"2025-04-16T12:07:13.069650Z","shell.execute_reply.started":"2025-04-16T12:07:12.766471Z","shell.execute_reply":"2025-04-16T12:07:13.068856Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"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":{"_uuid":"810f7e77-5eb0-412b-8c56-07f8467ccd3d","_cell_guid":"60572985-52b2-4ef8-bbf2-e298e794f210","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:07:13.070660Z","iopub.execute_input":"2025-04-16T12:07:13.071054Z","iopub.status.idle":"2025-04-16T12:07:19.786384Z","shell.execute_reply.started":"2025-04-16T12:07:13.071005Z","shell.execute_reply":"2025-04-16T12:07:19.785829Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"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)","metadata":{"_uuid":"4e7a8d22-d926-4612-8b50-5cee65c4daa1","_cell_guid":"ea9adbb7-b4aa-4e66-8773-eda14648f06c","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-04-16T12:07:19.786939Z","iopub.execute_input":"2025-04-16T12:07:19.787244Z","iopub.status.idle":"2025-04-16T12:07:19.791848Z","shell.execute_reply.started":"2025-04-16T12:07:19.787225Z","shell.execute_reply":"2025-04-16T12:07:19.791078Z"}},"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":{"_uuid":"b28ef2f3-cb6a-4105-86c3-537364bedac3","_cell_guid":"395214f5-218d-4add-9494-0684cc128c36","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-04-16T12:07:19.793359Z","iopub.execute_input":"2025-04-16T12:07:19.793549Z","iopub.status.idle":"2025-04-16T12:07:19.817548Z","shell.execute_reply.started":"2025-04-16T12:07:19.793535Z","shell.execute_reply":"2025-04-16T12:07:19.816895Z"}},"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":{"_uuid":"114065c1-cfd0-42e8-9b99-d5fcbb39bf9d","_cell_guid":"755f879d-04da-464d-ab85-3a2649f2adee","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:07:19.818427Z","iopub.execute_input":"2025-04-16T12:07:19.818641Z","iopub.status.idle":"2025-04-16T12:07:19.832705Z","shell.execute_reply.started":"2025-04-16T12:07:19.818620Z","shell.execute_reply":"2025-04-16T12:07:19.832056Z"},"jupyter":{"outputs_hidden":false}},"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","metadata":{"_uuid":"9d287956-5a52-4625-9148-baedbbff8185","_cell_guid":"44574194-7bd2-4b1d-beef-f8ddd1f1fa9b","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:07:19.833408Z","iopub.execute_input":"2025-04-16T12:07:19.833632Z","iopub.status.idle":"2025-04-16T12:07:19.846872Z","shell.execute_reply.started":"2025-04-16T12:07:19.833595Z","shell.execute_reply":"2025-04-16T12:07:19.846186Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Check for correct output size","metadata":{"_uuid":"487b4707-d062-4af4-bb53-cd60889a9b52","_cell_guid":"d3ff1002-cbd4-4532-bc00-d45981c0bd5f","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"input_image = torch.rand((1,3,512,512))\nmodel = UNet(3,10)\noutput = model(input_image)\nprint(output.size())","metadata":{"_uuid":"5d91dce5-a513-446f-9cbf-45851fef23e1","_cell_guid":"f13d9bf5-8f26-46db-97d5-87d5a738c24b","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:07:19.847808Z","iopub.execute_input":"2025-04-16T12:07:19.848136Z","iopub.status.idle":"2025-04-16T12:07:23.260453Z","shell.execute_reply.started":"2025-04-16T12:07:19.848113Z","shell.execute_reply":"2025-04-16T12:07:23.259646Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset","metadata":{"_uuid":"6387e4ea-f6f5-4068-96f3-ce410db312c1","_cell_guid":"1a9195eb-1e31-4866-b2ff-e732cdb5b433","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"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\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\") # convert to grayscale\n\n        return self.transform(img), self.transform(mask)\n\n    def __len__(self):\n        return min(len(self.images), self.limit)","metadata":{"_uuid":"a990928b-2a66-4a3b-af36-8bae6902a94c","_cell_guid":"b058c6ac-9327-4f21-b4e7-6bfec9313e40","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:07:23.261315Z","iopub.execute_input":"2025-04-16T12:07:23.261574Z","iopub.status.idle":"2025-04-16T12:07:23.267634Z","shell.execute_reply.started":"2025-04-16T12:07:23.261547Z","shell.execute_reply":"2025-04-16T12:07:23.266907Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(os.listdir(\"../input/carvana-image-masking-challenge/\"))","metadata":{"_uuid":"2e567e07-65e2-4aba-bd43-18ca03915471","_cell_guid":"74e9f8aa-d40f-444d-8ba6-f67b0c82dee6","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:07:23.268526Z","iopub.execute_input":"2025-04-16T12:07:23.268762Z","iopub.status.idle":"2025-04-16T12:07:23.286159Z","shell.execute_reply.started":"2025-04-16T12:07:23.268748Z","shell.execute_reply":"2025-04-16T12:07:23.285544Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DATASET_DIR = '../input/carvana-image-masking-challenge/'\nWORKING_DIR = '/kaggle/working/'","metadata":{"_uuid":"a1363874-df33-49d7-9295-fd3f3e44d4a6","_cell_guid":"c77c7b3f-8e78-42e8-9be7-09eb005288bb","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:07:23.288412Z","iopub.execute_input":"2025-04-16T12:07:23.288940Z","iopub.status.idle":"2025-04-16T12:07:23.299664Z","shell.execute_reply.started":"2025-04-16T12:07:23.288923Z","shell.execute_reply":"2025-04-16T12:07:23.299054Z"},"jupyter":{"outputs_hidden":false}},"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":{"_uuid":"8f1f6404-bcfe-4d33-b72a-0f73e8ef200d","_cell_guid":"3641c096-9731-4a37-a4f2-c3383d9e82e9","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:07:23.300208Z","iopub.execute_input":"2025-04-16T12:07:23.300433Z","iopub.status.idle":"2025-04-16T12:07:34.103350Z","shell.execute_reply.started":"2025-04-16T12:07:23.300417Z","shell.execute_reply":"2025-04-16T12:07:34.102696Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset = CarvanaDataset(WORKING_DIR)\n\n# generator = torch.Generator().manual_seed(25)","metadata":{"_uuid":"3940b736-5727-4a1c-bd59-3de0ad698495","_cell_guid":"820ce447-7b4a-4c1d-b3d4-096e6ce7a64b","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:07:34.104087Z","iopub.execute_input":"2025-04-16T12:07:34.104323Z","iopub.status.idle":"2025-04-16T12:07:34.117731Z","shell.execute_reply.started":"2025-04-16T12:07:34.104298Z","shell.execute_reply":"2025-04-16T12:07:34.117057Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset, test_dataset = random_split(train_dataset, [0.8, 0.2])","metadata":{"_uuid":"6c6e3c9b-69a9-46bc-83d1-29318048d1b1","_cell_guid":"b56dd837-58b1-4b7e-99e5-455fc72b689d","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:07:34.118525Z","iopub.execute_input":"2025-04-16T12:07:34.119155Z","iopub.status.idle":"2025-04-16T12:07:34.147115Z","shell.execute_reply.started":"2025-04-16T12:07:34.119132Z","shell.execute_reply":"2025-04-16T12:07:34.146402Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_dataset, val_dataset = random_split(test_dataset, [0.5, 0.5])","metadata":{"_uuid":"32e130a2-dfba-47a9-aaa3-d52d61306685","_cell_guid":"0e33f3ca-5d08-406c-9775-a94f814f36ca","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:07:34.147880Z","iopub.execute_input":"2025-04-16T12:07:34.148099Z","iopub.status.idle":"2025-04-16T12:07:34.153721Z","shell.execute_reply.started":"2025-04-16T12:07:34.148076Z","shell.execute_reply":"2025-04-16T12:07:34.153030Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Sneak-Peek Data","metadata":{"_uuid":"81f70ada-621c-4fa3-ad39-8335fc280d6e","_cell_guid":"52312489-94d7-41c8-aec8-b07c480ae667","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"# Let's peek at some data, but only training data\nimg = train_dataset[0]\nprint(img)","metadata":{"_uuid":"f0c43e25-df76-4e9a-befb-38d67a7f87d9","_cell_guid":"9c01d07d-a8e0-4cf0-b73f-1d734fe2491a","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:07:34.154527Z","iopub.execute_input":"2025-04-16T12:07:34.154739Z","iopub.status.idle":"2025-04-16T12:07:34.285600Z","shell.execute_reply.started":"2025-04-16T12:07:34.154716Z","shell.execute_reply":"2025-04-16T12:07:34.285002Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img[0].size()","metadata":{"_uuid":"24edc35b-a5e4-4438-8e04-df255dc57406","_cell_guid":"f7d8813f-85e3-4cb3-ab91-3098d6e53a02","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:07:34.286210Z","iopub.execute_input":"2025-04-16T12:07:34.286452Z","iopub.status.idle":"2025-04-16T12:07:34.291715Z","shell.execute_reply.started":"2025-04-16T12:07:34.286427Z","shell.execute_reply":"2025-04-16T12:07:34.290974Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"imgPIL = img[0].permute(1,2,0)\nmaskPIL = img[1].permute(1,2,0)\nimgPIL.size()","metadata":{"_uuid":"d1b17406-181b-47be-9ac6-c5c1cf07c8f2","_cell_guid":"89d0d196-20ba-4612-8c9b-d7622fe62949","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:07:34.292492Z","iopub.execute_input":"2025-04-16T12:07:34.292685Z","iopub.status.idle":"2025-04-16T12:07:34.306451Z","shell.execute_reply.started":"2025-04-16T12:07:34.292672Z","shell.execute_reply":"2025-04-16T12:07:34.305938Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.imshow(imgPIL.numpy())","metadata":{"_uuid":"fbc75472-d741-4d54-9541-a0e21473603d","_cell_guid":"4f2965e0-ce0f-409b-b11c-82b146f32bad","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:07:34.307021Z","iopub.execute_input":"2025-04-16T12:07:34.307172Z","iopub.status.idle":"2025-04-16T12:07:34.574155Z","shell.execute_reply.started":"2025-04-16T12:07:34.307160Z","shell.execute_reply":"2025-04-16T12:07:34.573383Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.imshow(maskPIL.numpy(), cmap='gray')","metadata":{"_uuid":"438a5e05-3800-44ed-8ba9-936b2ec44992","_cell_guid":"4da92349-2c6d-4c35-94c8-52005c4c2f8d","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:07:34.574903Z","iopub.execute_input":"2025-04-16T12:07:34.575121Z","iopub.status.idle":"2025-04-16T12:07:34.740469Z","shell.execute_reply.started":"2025-04-16T12:07:34.575105Z","shell.execute_reply":"2025-04-16T12:07:34.739867Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Setup DataLoader and Hyperparameters","metadata":{"_uuid":"0d7104e1-2edf-4c9a-a6e6-d9c648e7dc6c","_cell_guid":"7e11272e-2d91-4a0b-b8c9-60085f8fc9ff","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"os.cpu_count()","metadata":{"_uuid":"bb692bcb-fdbc-4a2d-9d55-9eff51da3340","_cell_guid":"ab1c93be-9cd8-4999-8f41-8e5765faf8d9","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:07:34.741115Z","iopub.execute_input":"2025-04-16T12:07:34.741337Z","iopub.status.idle":"2025-04-16T12:07:34.745840Z","shell.execute_reply.started":"2025-04-16T12:07:34.741320Z","shell.execute_reply":"2025-04-16T12:07:34.745350Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nif device == \"cuda\":\n    NUM_WORKERS = torch.cuda.device_count() * min(4, os.cpu_count())\nelse: \n    NUM_WORKERS = min(4, os.cpu_count())","metadata":{"_uuid":"99c25501-dd54-449e-8da5-e08e27c1e3e5","_cell_guid":"9f524585-c175-41e8-96d0-711a33cf6979","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:07:34.746554Z","iopub.execute_input":"2025-04-16T12:07:34.747059Z","iopub.status.idle":"2025-04-16T12:07:34.825001Z","shell.execute_reply.started":"2025-04-16T12:07:34.747002Z","shell.execute_reply":"2025-04-16T12:07:34.824506Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"LEARNING_RATE = 3e-4\nBATCH_SIZE = 8\nPIN_MEMORY = False\nSHUFFLE = True","metadata":{"_uuid":"ba58512d-7d51-4d8b-925f-4d2cc5f3f06b","_cell_guid":"94e7b9de-3c7c-4939-b08d-afb495730f7f","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:07:34.825739Z","iopub.execute_input":"2025-04-16T12:07:34.825968Z","iopub.status.idle":"2025-04-16T12:07:34.829497Z","shell.execute_reply.started":"2025-04-16T12:07:34.825952Z","shell.execute_reply":"2025-04-16T12:07:34.828875Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from functools import partial\n\nsetup_dataloader = partial(DataLoader, \n                           num_workers=NUM_WORKERS,\n                           pin_memory=PIN_MEMORY,\n                           batch_size=BATCH_SIZE,\n                           shuffle=SHUFFLE\n                          )","metadata":{"_uuid":"de2206cc-8197-4f32-adbd-beecd1f287cb","_cell_guid":"527fd14a-bf9e-4899-a586-6772df8ebd22","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:07:34.830238Z","iopub.execute_input":"2025-04-16T12:07:34.830454Z","iopub.status.idle":"2025-04-16T12:07:34.843534Z","shell.execute_reply.started":"2025-04-16T12:07:34.830440Z","shell.execute_reply":"2025-04-16T12:07:34.842922Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataloader = setup_dataloader(dataset=train_dataset)\nval_dataloader = setup_dataloader(dataset=val_dataset)\ntest_dataloader = setup_dataloader(dataset=test_dataset)","metadata":{"_uuid":"7d1da930-834d-4135-9da5-7bc7216760df","_cell_guid":"3999f521-b1f5-4f4d-af5f-55906304316e","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:07:34.844311Z","iopub.execute_input":"2025-04-16T12:07:34.844660Z","iopub.status.idle":"2025-04-16T12:07:34.858361Z","shell.execute_reply.started":"2025-04-16T12:07:34.844639Z","shell.execute_reply":"2025-04-16T12:07:34.857734Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Setup Model","metadata":{"_uuid":"fa5dd677-cbbc-4169-9e97-e94ad6377283","_cell_guid":"b5233049-143b-44c4-9001-78f116c9074c","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"model = UNet(in_channels=3, num_classes=1).to(device)\noptimizer = optim.AdamW(model.parameters(), lr=LEARNING_RATE)\ncriterion = nn.BCEWithLogitsLoss()","metadata":{"_uuid":"ff777b1d-61e5-4487-a218-bf5864efe41c","_cell_guid":"fee31190-ecd2-4481-be29-16bc28e6b6be","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:07:34.858927Z","iopub.execute_input":"2025-04-16T12:07:34.859133Z","iopub.status.idle":"2025-04-16T12:07:35.283084Z","shell.execute_reply.started":"2025-04-16T12:07:34.859120Z","shell.execute_reply":"2025-04-16T12:07:35.282321Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def 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":{"_uuid":"93bfa911-3015-478f-ad04-dc63c2ed59ca","_cell_guid":"64ea0891-c29f-4764-ba31-9eee5af4b916","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:07:35.283894Z","iopub.execute_input":"2025-04-16T12:07:35.284226Z","iopub.status.idle":"2025-04-16T12:07:35.288501Z","shell.execute_reply.started":"2025-04-16T12:07:35.284206Z","shell.execute_reply":"2025-04-16T12:07:35.287842Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.cuda.empty_cache()","metadata":{"_uuid":"9fc19a9b-3553-4f0e-bc3a-0623c279cff9","_cell_guid":"8d9a39be-1612-4ffb-ae12-df5ffb70e365","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:07:35.289142Z","iopub.execute_input":"2025-04-16T12:07:35.289357Z","iopub.status.idle":"2025-04-16T12:07:35.304169Z","shell.execute_reply.started":"2025-04-16T12:07:35.289340Z","shell.execute_reply":"2025-04-16T12:07:35.303398Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training the model","metadata":{"_uuid":"86a8982c-7674-43af-a4f8-5554ebd28eb5","_cell_guid":"83660aa8-ea83-411f-9804-b8b85fcb128d","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"def trainUNet(epochs):\n    EPOCHS = epochs\n\n    train_losses = []\n    train_dcs = []\n    val_losses = []\n    val_dcs = []\n    \n    for 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    # Saving the model\n    torch.save(model.state_dict(), 'my_checkpoint.pth')\n    return train_losses, train_dcs, val_losses, val_dcs","metadata":{"_uuid":"f8bcb842-2f9d-4217-a2ed-1ea1e28983ef","_cell_guid":"74f45582-7760-412d-b096-10b8aa88380f","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:07:35.307615Z","iopub.execute_input":"2025-04-16T12:07:35.307803Z","iopub.status.idle":"2025-04-16T12:07:35.319435Z","shell.execute_reply.started":"2025-04-16T12:07:35.307789Z","shell.execute_reply":"2025-04-16T12:07:35.318805Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"EPOCHS = 1\ntrain_losses, train_dcs, val_losses, val_dcs = trainUNet(EPOCHS)","metadata":{"_uuid":"46188360-598b-4e28-b99d-ad71c0db1a49","_cell_guid":"88c6d55b-efd2-491f-adec-62858cff5cf2","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:07:35.319995Z","iopub.execute_input":"2025-04-16T12:07:35.320214Z","iopub.status.idle":"2025-04-16T12:14:28.820076Z","shell.execute_reply.started":"2025-04-16T12:07:35.320200Z","shell.execute_reply":"2025-04-16T12:14:28.819398Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Visualize Training Results","metadata":{"_uuid":"c751d111-3a84-45ff-ad97-734f4ea24e70","_cell_guid":"3df281b0-0dd9-4541-9d96-7ad5c4db841f","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"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()","metadata":{"_uuid":"34d4a372-8c38-4272-9a4d-26a9cc1d6fac","_cell_guid":"6e1bb0fe-0726-441f-8528-6445593fcc70","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:14:28.821143Z","iopub.execute_input":"2025-04-16T12:14:28.821421Z","iopub.status.idle":"2025-04-16T12:14:29.129378Z","shell.execute_reply.started":"2025-04-16T12:14:28.821396Z","shell.execute_reply":"2025-04-16T12:14:29.128649Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Evaluation with Test Set","metadata":{"_uuid":"0e55c682-f24e-4cbd-b4d8-6b8ac04952ae","_cell_guid":"16bcee79-9c38-469c-9ccc-a761e925a5bd","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"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":{"_uuid":"d9415734-7a58-401c-b2b3-8b5a9f935b32","_cell_guid":"a671c372-ea49-4302-9537-7d0994f81326","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:14:29.130173Z","iopub.execute_input":"2025-04-16T12:14:29.130821Z","iopub.status.idle":"2025-04-16T12:14:29.512415Z","shell.execute_reply.started":"2025-04-16T12:14:29.130801Z","shell.execute_reply":"2025-04-16T12:14:29.511730Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def testUNet():\n    test_running_loss = 0\n    test_running_dc = 0\n\n    with 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        print(f\"test loss: {test_loss}\")\n        print(f\"test dc: {test_dc}\")","metadata":{"_uuid":"f99f0c11-27e1-4459-8844-ce962d2d5811","_cell_guid":"9f912483-0fcc-4080-8560-4a4154550a16","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:14:29.513150Z","iopub.execute_input":"2025-04-16T12:14:29.513403Z","iopub.status.idle":"2025-04-16T12:14:29.518692Z","shell.execute_reply.started":"2025-04-16T12:14:29.513387Z","shell.execute_reply":"2025-04-16T12:14:29.517937Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"testUNet()","metadata":{"_uuid":"d3bbfdc4-1379-442c-a852-99368e43bc28","_cell_guid":"d368b42d-97cf-4bbd-b3e8-9b21d7ad8e23","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:14:29.519385Z","iopub.execute_input":"2025-04-16T12:14:29.519644Z","iopub.status.idle":"2025-04-16T12:14:49.551327Z","shell.execute_reply.started":"2025-04-16T12:14:29.519615Z","shell.execute_reply":"2025-04-16T12:14:49.550251Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Visualize Inference Results","metadata":{"_uuid":"4c913155-80f7-4254-bbc1-78f491917eb4","_cell_guid":"9d676c97-5ef6-4ef4-87c0-9782a8c5e0e7","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"def random_images_inference(image_tensors, mask_tensors, 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 in zip(image_tensors, mask_tensors):\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\"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()","metadata":{"_uuid":"65523837-7e03-4f25-bc98-843729dba989","_cell_guid":"072c85d2-5fe3-440f-834b-98b0df05217b","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:14:49.552330Z","iopub.execute_input":"2025-04-16T12:14:49.552560Z","iopub.status.idle":"2025-04-16T12:14:49.560207Z","shell.execute_reply.started":"2025-04-16T12:14:49.552537Z","shell.execute_reply":"2025-04-16T12:14:49.559405Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(len(train_dataloader.dataset))\nprint(len(val_dataloader.dataset))\nprint(len(test_dataloader.dataset))","metadata":{"_uuid":"981c0fef-a3bc-4623-a70b-37c20e463ea8","_cell_guid":"205b0951-2b92-4026-891e-15d5232b14df","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:14:49.560875Z","iopub.execute_input":"2025-04-16T12:14:49.561109Z","iopub.status.idle":"2025-04-16T12:14:49.581295Z","shell.execute_reply.started":"2025-04-16T12:14:49.561085Z","shell.execute_reply":"2025-04-16T12:14:49.580495Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"n = 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])","metadata":{"_uuid":"bf0e4fca-023a-47c4-863d-f561f653c2bb","_cell_guid":"c574f8cc-25f7-4f0f-9549-cbe08b0c5fdf","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:15:07.886892Z","iopub.execute_input":"2025-04-16T12:15:07.887508Z","iopub.status.idle":"2025-04-16T12:15:08.376501Z","shell.execute_reply.started":"2025-04-16T12:15:07.887481Z","shell.execute_reply":"2025-04-16T12:15:08.375791Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_path = '/kaggle/working/my_checkpoint.pth'\n\nrandom_images_inference(image_tensors, mask_tensors, model_path, device=\"cpu\")","metadata":{"_uuid":"1f8d7d12-0878-4134-a849-0e0bbec2e617","_cell_guid":"dca802cc-4484-4add-8c75-44b3f05df8be","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:15:08.377648Z","iopub.execute_input":"2025-04-16T12:15:08.377849Z","iopub.status.idle":"2025-04-16T12:15:38.514469Z","shell.execute_reply.started":"2025-04-16T12:15:08.377834Z","shell.execute_reply":"2025-04-16T12:15:38.513704Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def show_image_with_mask(image_tensors, mask_tensors, model_path, alpha=0.5, device='cpu'):\n    \"\"\"\n    image: 3x H x W (torch.tensor)\n    mask: 1x H x W (torch.tensor)\n    alpha: Opacity of the mask\n    mask_color: Color of the mask (e.g., 'r' for red, 'g' for green, etc.)\n    \"\"\"\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 in zip(image_tensors, mask_tensors):\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\"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.imshow(img)\n        plt.imshow(pred_mask, cmap='Reds', alpha=alpha)\n        plt.show()","metadata":{"_uuid":"6b45dd8a-c22f-4916-aebd-4f2f9cb1dc76","_cell_guid":"5ef094cd-a911-4530-8510-88261d383f95","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:15:38.515350Z","iopub.execute_input":"2025-04-16T12:15:38.515636Z","iopub.status.idle":"2025-04-16T12:15:38.521559Z","shell.execute_reply.started":"2025-04-16T12:15:38.515599Z","shell.execute_reply":"2025-04-16T12:15:38.521043Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_path = '/kaggle/working/my_checkpoint.pth'\n\nshow_image_with_mask(image_tensors, mask_tensors, model_path, device=\"cpu\")","metadata":{"_uuid":"0bb51079-38e8-42c5-8c3d-9ec3a16a8475","_cell_guid":"a2aaf728-b609-4e1c-b764-6edc65277693","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-16T12:15:38.522820Z","iopub.execute_input":"2025-04-16T12:15:38.523005Z","iopub.status.idle":"2025-04-16T12:16:11.426050Z","shell.execute_reply.started":"2025-04-16T12:15:38.522991Z","shell.execute_reply":"2025-04-16T12:16:11.425292Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"_uuid":"ba13709a-31cc-4e76-a5ac-1ca1a703c82a","_cell_guid":"369e376e-22ae-4d68-b763-98df5ca49d59","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}