{"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":10129395,"sourceType":"datasetVersion","datasetId":6251207}],"dockerImageVersionId":30787,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport random\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torchvision.transforms as transforms\nfrom PIL import Image\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\nimport gc\ngc.collect()\ntorch.cuda.empty_cache()\n\n\nWORKING_DIR = '/kaggle/working/'\nimage_directory = \"/kaggle/input/dataset-0-3-0/0.3.0/Images\"\nmask_directory = \"/kaggle/input/dataset-0-3-0/0.3.0/Masks\"\nSIZE = (128,256)\nLEARNING_RATE = 3e-4\nBATCH_SIZE = 16\nNUM_CLASS = 2\nEPOCHS = 1\n\nif (os.path.exists(WORKING_DIR + 'checkpoint') == False):\n    os.mkdir(WORKING_DIR + 'checkpoint')\nif (os.path.exists(WORKING_DIR + 'final') == False):\n    os.mkdir(WORKING_DIR + 'final')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import glob\nimport os\nimport re\nimport torch\nfrom torch.utils.data import Dataset, random_split\nfrom torchvision.transforms import v2 as transforms\nfrom PIL import Image\nimport numpy as np\n            \nclass SimDataset(Dataset):\n  def __init__(self, image_paths, mask_paths_left, mask_paths_right, transform_img=None, transform_mask=None):\n    self.image_paths = image_paths\n    self.mask_paths_left = mask_paths_left\n    self.mask_paths_right = mask_paths_right\n    self.transform_img = transform_img\n    self.transform_mask = transform_mask\n\n  def __len__(self):\n    return len(self.image_paths)\n\n  def __getitem__(self, idx):\n    img_path = self.image_paths[idx]\n    mask_path_left = self.mask_paths_left[idx]\n    mask_path_right = self.mask_paths_right[idx]\n\n    mask_left = Image.open(mask_path_left).convert(\"L\") # convert(\"L\") will convert the image to grayscale\n    mask_right = Image.open(mask_path_right).convert(\"L\") # convert(\"L\") will convert the image to grayscale\n    if self.transform_mask:\n        mask_left = self.transform_mask(mask_left)\n        mask_right = self.transform_mask(mask_right)\n        mask = torch.tensor(np.array([mask_left[0], mask_right[0]]))\n        mask[mask > 0.0] = 1\n\n    image = Image.open(img_path).convert(\"RGB\") # we might not need to do this because the images are loaded as RGB by default\n    if self.transform_img:\n      image = self.transform_img(image)\n    return [image, mask]\n\n\nimage_paths = sorted(glob.glob(os.path.join(image_directory, \"*.png\")), key=lambda x:float(re.findall(\"(\\d+)\",x)[-1]))\nmask_paths_left = sorted(glob.glob(os.path.join(mask_directory, \"*-left.png\")), key=lambda x:float(re.findall(\"(\\d+)\",x)[-1]))\nmask_paths_right = sorted(glob.glob(os.path.join(mask_directory, \"*-right.png\")), key=lambda x:float(re.findall(\"(\\d+)\",x)[-1]))\nassert len(mask_paths_left) == len(mask_paths_right), \"Amount of left and right mask is no\"\nassert len(image_paths) == len(mask_paths_left), \"Amount of images and label are different\"\ndataset = SimDataset(image_paths,\n                     mask_paths_left,\n                     mask_paths_right,\n                     transforms.Compose([\n                                         transforms.Resize(SIZE),\n                                         transforms.ToImage(),\n                                         transforms.ToDtype(torch.float32, scale=True),\n                                         # transforms.Normalize(mean=[0.0],\n                                         #                      std=[1.0])\n                                                              ]),\n                     transforms.Compose([\n                                         transforms.Resize(SIZE),\n                                         transforms.ToImage(),\n                                         transforms.ToDtype(torch.float32, scale=True),\n                                         # transforms.Normalize(mean=[0.0],\n                                         #                       std=[1.0])\n                                                              ]))\ngenerator = torch.Generator().manual_seed(25)\ntrain_dataset, test_dataset = random_split(dataset, [0.8, 0.2], generator=generator)\ntest_dataset, val_dataset = random_split(test_dataset, [0.5, 0.5], generator=generator)\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2025-01-30T11:12:03.578949Z","iopub.execute_input":"2025-01-30T11:12:03.579302Z","iopub.status.idle":"2025-01-30T11:12:03.754892Z","shell.execute_reply.started":"2025-01-30T11:12:03.579263Z","shell.execute_reply":"2025-01-30T11:12:03.754212Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom torch import nn\n\n# https://medium.com/@fernandopalominocobo/mastering-u-net-a-step-by-step-guide-to-segmentation-from-scratch-with-pytorch-6a17c5916114\nclass 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\nclass 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\n\nclass 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)\n\nclass 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-01-30T11:12:03.755903Z","iopub.execute_input":"2025-01-30T11:12:03.756165Z","iopub.status.idle":"2025-01-30T11:12:03.768104Z","shell.execute_reply.started":"2025-01-30T11:12:03.756140Z","shell.execute_reply":"2025-01-30T11:12:03.767147Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom torch import optim, nn\nfrom torch.utils.data import DataLoader\nfrom tqdm import tqdm\n\ntorch.cuda.empty_cache()\n\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nif device == \"cuda\":\n    num_workers = torch.cuda.device_count() * 4\nelse:\n    num_workers = 1\n\ntrain_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)","metadata":{"execution":{"iopub.status.busy":"2025-01-30T11:12:03.769962Z","iopub.execute_input":"2025-01-30T11:12:03.770770Z","iopub.status.idle":"2025-01-30T11:12:03.788248Z","shell.execute_reply.started":"2025-01-30T11:12:03.770742Z","shell.execute_reply":"2025-01-30T11:12:03.787443Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nmodel = UNet(in_channels=3, num_classes=NUM_CLASS).to(device)\n\n# model_pth = '/kaggle/working/final/final_epoch30.pth'\n# model.load_state_dict(torch.load(model_pth,weights_only=True, map_location=torch.device(device)))\noptimizer = optim.AdamW(model.parameters(), lr=LEARNING_RATE)\ncriterion = nn.BCEWithLogitsLoss().to(device=device)\n\n\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\n\n\ntorch.cuda.empty_cache()\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    if (epoch > 9 and epoch % 10 == 0):\n        torch.save(model.state_dict(), f\"checkpoint/checkpoint_epoch{epoch}.pth\")\n\n# Saving the model\ntorch.save(model.state_dict(), f'final/final_epoch{EPOCHS}.pth')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"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')\n\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":{"_kg_hide-input":false,"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T11:12:10.474461Z","iopub.status.idle":"2025-01-30T11:12:10.474749Z","shell.execute_reply.started":"2025-01-30T11:12:10.474610Z","shell.execute_reply":"2025-01-30T11:12:10.474624Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_pth = '/kaggle/working/final/final_epoch1.pth'\ntrained_model = UNet(in_channels=3, num_classes=NUM_CLASS).to(device)\ntrained_model.load_state_dict(torch.load(model_pth,weights_only=True, map_location=torch.device(device)))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T11:12:10.476149Z","iopub.status.idle":"2025-01-30T11:12:10.476448Z","shell.execute_reply.started":"2025-01-30T11:12:10.476310Z","shell.execute_reply":"2025-01-30T11:12:10.476325Z"}},"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)\nprint(f\"{test_loss=}\")\nprint(f\"{test_dc=}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T11:12:10.478151Z","iopub.status.idle":"2025-01-30T11:12:10.478583Z","shell.execute_reply.started":"2025-01-30T11:12:10.478353Z","shell.execute_reply":"2025-01-30T11:12:10.478375Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def random_images_inference(image_tensors, mask_tensors, model_pth, device):\n    model = UNet(in_channels=3, num_classes=NUM_CLASS).to(device)\n    model.load_state_dict(torch.load(model_pth, weights_only=True, map_location=torch.device(device)))\n\n    transform = transforms.Compose([\n        transforms.Resize(SIZE)\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) \n        pred_mask = pred_mask.permute(1,2,0)\n        \n        # Load the mask to compare\n        mask = transform(mask_pth).permute(1, 2, 0).to(device)\n        mask = mask.permute(2, 0, 1)\n        merged_mask = mask[0] + mask[1]\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.permute(2,0,1)\n        pred_mask[0] = torch.sigmoid(pred_mask[0]) > 0.5\n        pred_mask[1] = torch.sigmoid(pred_mask[1]) > 0.5\n        merged_pred_mask = pred_mask[0] + pred_mask[1]\n\n    \n        plt.figure(figsize=(16, 4))\n        plt.subplot(2,4,1), plt.imshow(img), plt.title(\"original\"), plt.axis(\"off\")\n        plt.subplot(2,4,2), plt.imshow(pred_mask[0], cmap=\"gray\"), plt.title(\"predicted left\"), plt.axis(\"off\")\n        plt.subplot(2,4,3), plt.imshow(pred_mask[1], cmap=\"gray\"), plt.title(\"predicted right\"), plt.axis(\"off\")\n        plt.subplot(2,4,4), plt.imshow(merged_pred_mask, cmap=\"gray\"), plt.title(\"predicted merged\"), plt.axis(\"off\")\n        \n        plt.subplot(2,4,6), plt.imshow(mask[0], cmap=\"gray\"), plt.title(\"mask left\"), plt.axis(\"off\")\n        plt.subplot(2,4,7), plt.imshow(mask[1], cmap=\"gray\"), plt.title(\"mask right\"), plt.axis(\"off\")\n        plt.subplot(2,4,8), plt.imshow(merged_mask, cmap=\"gray\"), plt.title(\"mask merged\"), plt.axis(\"off\")\n        plt.show()\n\nn = 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\nmodel_path = '/kaggle/working/final/final_epoch1.pth'\n\nrandom_images_inference(image_tensors, mask_tensors, model_path, device=\"cpu\")\n\n\n\n\n#resize nearest neighbor pour les mask","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T11:12:10.480414Z","iopub.status.idle":"2025-01-30T11:12:10.480719Z","shell.execute_reply.started":"2025-01-30T11:12:10.480572Z","shell.execute_reply":"2025-01-30T11:12:10.480587Z"}},"outputs":[],"execution_count":null}]}