{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":5506221,"sourceType":"datasetVersion","datasetId":1108926},{"sourceId":8055343,"sourceType":"datasetVersion","datasetId":4750955}],"dockerImageVersionId":30673,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"import libraries","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-04-07T15:23:27.049495Z","iopub.execute_input":"2024-04-07T15:23:27.050405Z","iopub.status.idle":"2024-04-07T15:23:35.141000Z","shell.execute_reply.started":"2024-04-07T15:23:27.050364Z","shell.execute_reply":"2024-04-07T15:23:35.140221Z"}}},{"cell_type":"code","source":"from torch.utils.data.dataset import Dataset\nimport os\nimport matplotlib.pyplot as plt\nimport matplotlib.image as mpimg\nfrom PIL import Image \nfrom torchvision import transforms\nimport zipfile\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import random_split, DataLoader\nfrom torch import optim \nfrom tqdm import tqdm\nfrom torchmetrics import Dice","metadata":{"execution":{"iopub.status.busy":"2024-04-08T06:02:54.921063Z","iopub.execute_input":"2024-04-08T06:02:54.921576Z","iopub.status.idle":"2024-04-08T06:03:03.320776Z","shell.execute_reply.started":"2024-04-08T06:02:54.921528Z","shell.execute_reply":"2024-04-08T06:03:03.319959Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"dataset for pics","metadata":{}},{"cell_type":"code","source":"class MyDataset(Dataset):\n    def __init__(self, root_path):\n        self.root_path = root_path\n        \n        \n        self.images = sorted([root_path+'/images/' + i for i in os.listdir(root_path+'/images')])\n        self.masks = sorted([root_path+\"/masks/\"+i for i in os.listdir(root_path+'/masks')])\n        self.transform = transforms.Compose([\n            \n            transforms.ToTensor()])\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)\n\n    def __len__(self):\n        return len(self.images)","metadata":{"execution":{"iopub.status.busy":"2024-04-08T06:03:07.648555Z","iopub.execute_input":"2024-04-08T06:03:07.649028Z","iopub.status.idle":"2024-04-08T06:03:07.658259Z","shell.execute_reply.started":"2024-04-08T06:03:07.648999Z","shell.execute_reply":"2024-04-08T06:03:07.657384Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nDATA_PATH = \"/kaggle/input/segmentation-full-body-mads-dataset/segmentation_full_body_mads_dataset_1192_img/segmentation_full_body_mads_dataset_1192_img\"","metadata":{"execution":{"iopub.status.busy":"2024-04-08T06:03:11.326689Z","iopub.execute_input":"2024-04-08T06:03:11.327478Z","iopub.status.idle":"2024-04-08T06:03:11.331899Z","shell.execute_reply.started":"2024-04-08T06:03:11.327437Z","shell.execute_reply":"2024-04-08T06:03:11.331066Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = MyDataset(DATA_PATH)","metadata":{"execution":{"iopub.status.busy":"2024-04-08T06:03:13.023671Z","iopub.execute_input":"2024-04-08T06:03:13.024398Z","iopub.status.idle":"2024-04-08T06:03:13.582461Z","shell.execute_reply.started":"2024-04-08T06:03:13.024367Z","shell.execute_reply":"2024-04-08T06:03:13.581659Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img,mask = dataset.__getitem__(0)","metadata":{"execution":{"iopub.status.busy":"2024-04-08T06:03:15.496035Z","iopub.execute_input":"2024-04-08T06:03:15.496859Z","iopub.status.idle":"2024-04-08T06:03:15.564677Z","shell.execute_reply.started":"2024-04-08T06:03:15.496829Z","shell.execute_reply":"2024-04-08T06:03:15.563744Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img.shape","metadata":{"execution":{"iopub.status.busy":"2024-04-08T06:03:17.056277Z","iopub.execute_input":"2024-04-08T06:03:17.057145Z","iopub.status.idle":"2024-04-08T06:03:17.063529Z","shell.execute_reply.started":"2024-04-08T06:03:17.057112Z","shell.execute_reply":"2024-04-08T06:03:17.062604Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"test photo","metadata":{}},{"cell_type":"code","source":"plt.imshow(img.permute(1, 2, 0)) ","metadata":{"execution":{"iopub.status.busy":"2024-04-08T06:03:20.504972Z","iopub.execute_input":"2024-04-08T06:03:20.505941Z","iopub.status.idle":"2024-04-08T06:03:20.922112Z","shell.execute_reply.started":"2024-04-08T06:03:20.505899Z","shell.execute_reply":"2024-04-08T06:03:20.921221Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"test mask","metadata":{}},{"cell_type":"code","source":"plt.imshow(mask[0]) \n","metadata":{"execution":{"iopub.status.busy":"2024-04-08T06:03:27.554596Z","iopub.execute_input":"2024-04-08T06:03:27.554938Z","iopub.status.idle":"2024-04-08T06:03:27.800636Z","shell.execute_reply.started":"2024-04-08T06:03:27.554912Z","shell.execute_reply":"2024-04-08T06:03:27.799734Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"create blocks for UNet","metadata":{}},{"cell_type":"code","source":"class DoubleConv(nn.Module):\n    def __init__(self,in_channels,out_channels):\n        super().__init__()\n        self.convs = nn.Sequential(\n        \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    def forward(self,x):\n        return self.convs(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    def forward(self,x):\n        down = self.conv(x)\n        p = self.pool(down)\n        \n        return down,p\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    def forward(self,x1,x2):\n        x1 = self.up(x1)\n        x = torch.cat([x1,x2],1)\n        return self.conv(x)\n        \n     ","metadata":{"execution":{"iopub.status.busy":"2024-04-08T06:03:34.720598Z","iopub.execute_input":"2024-04-08T06:03:34.720971Z","iopub.status.idle":"2024-04-08T06:03:34.731211Z","shell.execute_reply.started":"2024-04-08T06:03:34.720942Z","shell.execute_reply":"2024-04-08T06:03:34.730297Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"create UNet","metadata":{}},{"cell_type":"code","source":"class UNet(nn.Module):\n    def __init__(self, in_channels):\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=1, 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":{"execution":{"iopub.status.busy":"2024-04-08T06:03:37.557675Z","iopub.execute_input":"2024-04-08T06:03:37.558054Z","iopub.status.idle":"2024-04-08T06:03:37.568373Z","shell.execute_reply.started":"2024-04-08T06:03:37.558026Z","shell.execute_reply":"2024-04-08T06:03:37.567396Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"checking the UNet","metadata":{}},{"cell_type":"code","source":"input_image = torch.rand(1,3,384,512)\nmodel = UNet(3)","metadata":{"execution":{"iopub.status.busy":"2024-04-08T06:03:40.785236Z","iopub.execute_input":"2024-04-08T06:03:40.785891Z","iopub.status.idle":"2024-04-08T06:03:41.098339Z","shell.execute_reply.started":"2024-04-08T06:03:40.785858Z","shell.execute_reply":"2024-04-08T06:03:41.097369Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"out = model(input_image)\nprint(out.shape)","metadata":{"execution":{"iopub.status.busy":"2024-04-08T06:03:42.382838Z","iopub.execute_input":"2024-04-08T06:03:42.383191Z","iopub.status.idle":"2024-04-08T06:03:44.724638Z","shell.execute_reply.started":"2024-04-08T06:03:42.383163Z","shell.execute_reply":"2024-04-08T06:03:44.723591Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"parameters for the UNet","metadata":{}},{"cell_type":"code","source":"lr = 3e-4\nbatch_size = 15\nepochs = 23\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'\n","metadata":{"execution":{"iopub.status.busy":"2024-04-08T06:03:46.088797Z","iopub.execute_input":"2024-04-08T06:03:46.089152Z","iopub.status.idle":"2024-04-08T06:03:46.115575Z","shell.execute_reply.started":"2024-04-08T06:03:46.089125Z","shell.execute_reply":"2024-04-08T06:03:46.114304Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"create test val and train datasets","metadata":{}},{"cell_type":"code","source":"dataset_size = len(dataset)\ntest_size = int(dataset_size * 0.15)\nval_size = int(dataset_size * 0.15)\ntrain_size = dataset_size - (val_size + test_size)\n\ntrain_dataset, val_dataset, test_dataset = random_split(dataset, [train_size, val_size, test_size])\n","metadata":{"execution":{"iopub.status.busy":"2024-04-08T06:03:48.969759Z","iopub.execute_input":"2024-04-08T06:03:48.970403Z","iopub.status.idle":"2024-04-08T06:03:48.980161Z","shell.execute_reply.started":"2024-04-08T06:03:48.970373Z","shell.execute_reply":"2024-04-08T06:03:48.979383Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"create dataloaders","metadata":{}},{"cell_type":"code","source":"train_loader = DataLoader(dataset = train_dataset,batch_size=batch_size,shuffle=True)\n\ntest_loader = DataLoader(dataset=test_dataset,batch_size=batch_size,shuffle=True)\n\nval_loader = DataLoader(dataset=val_dataset,batch_size=batch_size,shuffle=True)","metadata":{"execution":{"iopub.status.busy":"2024-04-08T06:03:51.179451Z","iopub.execute_input":"2024-04-08T06:03:51.180138Z","iopub.status.idle":"2024-04-08T06:03:51.185247Z","shell.execute_reply.started":"2024-04-08T06:03:51.180110Z","shell.execute_reply":"2024-04-08T06:03:51.184088Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = UNet(in_channels=3).to(device)","metadata":{"execution":{"iopub.status.busy":"2024-04-08T06:03:53.493425Z","iopub.execute_input":"2024-04-08T06:03:53.493772Z","iopub.status.idle":"2024-04-08T06:03:53.940144Z","shell.execute_reply.started":"2024-04-08T06:03:53.493746Z","shell.execute_reply":"2024-04-08T06:03:53.939250Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"optimizer = optim.AdamW(model.parameters(),lr=lr)\ncriterion = nn.BCEWithLogitsLoss()","metadata":{"execution":{"iopub.status.busy":"2024-04-08T06:03:57.357631Z","iopub.execute_input":"2024-04-08T06:03:57.358223Z","iopub.status.idle":"2024-04-08T06:03:57.363339Z","shell.execute_reply.started":"2024-04-08T06:03:57.358193Z","shell.execute_reply":"2024-04-08T06:03:57.362225Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The Dice coefficient, a.k.a. Dice score, is chosen here as it provides a more robust measure \nfor the similarity between the predicted and true segmentation masks, especially in cases \nwhere the classes are imbalanced. It accounts for both the false positives and false negatives, \noffering a better performance metric than simple accuracy in the context of image segmentation tasks.\n","metadata":{}},{"cell_type":"code","source":"def dice_coefficient(pred, target):\n    smooth = 1.0  \n    intersection = (pred * target).sum()\n    dice = (2. * intersection + smooth) / (pred.sum() + target.sum() + smooth)\n    return dice","metadata":{"execution":{"iopub.status.busy":"2024-04-08T06:04:00.264678Z","iopub.execute_input":"2024-04-08T06:04:00.265014Z","iopub.status.idle":"2024-04-08T06:04:00.270222Z","shell.execute_reply.started":"2024-04-08T06:04:00.264990Z","shell.execute_reply":"2024-04-08T06:04:00.269163Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"train train train train!!!","metadata":{}},{"cell_type":"code","source":"for epoch in tqdm(range(epochs)):\n    model.train()\n    train_running_loss = 0\n    for index, img_mask in enumerate(tqdm(train_loader)):\n        img, mask = img_mask[0].to(device), img_mask[1].to(device)\n        \n        y_pred = model(img)\n        optimizer.zero_grad()\n        \n        loss = criterion(y_pred,mask)\n        train_running_loss += loss.item()\n        \n        loss.backward()\n        optimizer.step()\n        \n    train_loss = train_running_loss/index+1\n    \n    model.eval()\n    val_running_loss = 0\n    val_running_dice = 0\n   \n\n    with torch.no_grad():\n        for index, img_mask in enumerate(tqdm(val_loader)):\n            img, mask = img_mask[0].to(device), img_mask[1].to(device)\n    \n            y_pred = model(img)\n            loss = criterion(y_pred, mask)\n    \n            val_running_loss += loss.item()\n    \n            preds = torch.sigmoid(y_pred) > 0.5\n            dice_score = dice_coefficient(preds.float(), mask.float())\n            \n            val_running_dice += dice_score.item()\n    \n    val_loss = val_running_loss / len(val_loader)\n    val_dice = val_running_dice / len(val_loader)\n    print(f'epoch: {epoch+1}')\n    print(f'train loss: {train_loss:.4f}')\n    print(f'val loss: {val_loss:.4f}')\n    print(f'dice coef: {val_dice:.4f}')\n\n        ","metadata":{"execution":{"iopub.status.busy":"2024-04-08T06:04:02.176348Z","iopub.execute_input":"2024-04-08T06:04:02.176693Z","iopub.status.idle":"2024-04-08T06:28:55.883994Z","shell.execute_reply.started":"2024-04-08T06:04:02.176667Z","shell.execute_reply":"2024-04-08T06:28:55.883090Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"test model","metadata":{}},{"cell_type":"code","source":"\nmodel.eval()\ntest_running_dice = 0\nwith torch.no_grad():\n        for index, img_mask in enumerate(tqdm(test_loader)):\n            img, mask = img_mask[0].to(device), img_mask[1].to(device)\n    \n            y_pred = model(img)\n            loss = criterion(y_pred, mask)\n    \n            val_running_loss += loss.item()\n    \n            preds = torch.sigmoid(y_pred) > 0.5\n            dice_score = dice_coefficient(preds.float(), mask.float())\n            \n            test_running_dice += dice_score.item()\ntest_dice = test_running_dice/len(test_loader)\nprint(test_dice)\n\n\n\n\n","metadata":{"execution":{"iopub.status.busy":"2024-04-08T06:47:34.338786Z","iopub.execute_input":"2024-04-08T06:47:34.339470Z","iopub.status.idle":"2024-04-08T06:47:44.776254Z","shell.execute_reply.started":"2024-04-08T06:47:34.339438Z","shell.execute_reply":"2024-04-08T06:47:44.775373Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"some visualization","metadata":{}},{"cell_type":"code","source":"img,mask = test_dataset.__getitem__(0)\n\n\n","metadata":{"execution":{"iopub.status.busy":"2024-04-08T06:47:51.852493Z","iopub.execute_input":"2024-04-08T06:47:51.853200Z","iopub.status.idle":"2024-04-08T06:47:51.873901Z","shell.execute_reply.started":"2024-04-08T06:47:51.853168Z","shell.execute_reply":"2024-04-08T06:47:51.873133Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_test = img.unsqueeze(0).to(device)","metadata":{"execution":{"iopub.status.busy":"2024-04-08T06:47:54.529789Z","iopub.execute_input":"2024-04-08T06:47:54.530167Z","iopub.status.idle":"2024-04-08T06:47:54.535548Z","shell.execute_reply.started":"2024-04-08T06:47:54.530136Z","shell.execute_reply":"2024-04-08T06:47:54.534623Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_for_demo = img.permute(1, 2, 0)\n","metadata":{"execution":{"iopub.status.busy":"2024-04-08T06:47:56.180088Z","iopub.execute_input":"2024-04-08T06:47:56.180453Z","iopub.status.idle":"2024-04-08T06:47:56.185226Z","shell.execute_reply.started":"2024-04-08T06:47:56.180424Z","shell.execute_reply":"2024-04-08T06:47:56.184213Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"out = model(img_test)","metadata":{"execution":{"iopub.status.busy":"2024-04-08T06:47:58.537694Z","iopub.execute_input":"2024-04-08T06:47:58.538056Z","iopub.status.idle":"2024-04-08T06:47:58.599465Z","shell.execute_reply.started":"2024-04-08T06:47:58.538027Z","shell.execute_reply":"2024-04-08T06:47:58.598747Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.imshow(img_for_demo)","metadata":{"execution":{"iopub.status.busy":"2024-04-08T06:48:00.225364Z","iopub.execute_input":"2024-04-08T06:48:00.225705Z","iopub.status.idle":"2024-04-08T06:48:00.647201Z","shell.execute_reply.started":"2024-04-08T06:48:00.225680Z","shell.execute_reply":"2024-04-08T06:48:00.646368Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"out1 = out[0].cpu().permute(1, 2, 0).detach().numpy()\nplt.imshow(out1)","metadata":{"execution":{"iopub.status.busy":"2024-04-08T06:48:04.632034Z","iopub.execute_input":"2024-04-08T06:48:04.632785Z","iopub.status.idle":"2024-04-08T06:48:04.981218Z","shell.execute_reply.started":"2024-04-08T06:48:04.632751Z","shell.execute_reply":"2024-04-08T06:48:04.980386Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.imshow(mask[0])","metadata":{"execution":{"iopub.status.busy":"2024-04-08T06:48:09.159727Z","iopub.execute_input":"2024-04-08T06:48:09.160451Z","iopub.status.idle":"2024-04-08T06:48:09.445246Z","shell.execute_reply.started":"2024-04-08T06:48:09.160418Z","shell.execute_reply":"2024-04-08T06:48:09.444366Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"    out2 = (out1 >= 0.7)\nplt.imshow(out2)","metadata":{"execution":{"iopub.status.busy":"2024-04-08T06:49:07.332339Z","iopub.execute_input":"2024-04-08T06:49:07.332721Z","iopub.status.idle":"2024-04-08T06:49:07.626561Z","shell.execute_reply.started":"2024-04-08T06:49:07.332690Z","shell.execute_reply":"2024-04-08T06:49:07.625630Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"human = img_for_demo*out2","metadata":{"execution":{"iopub.status.busy":"2024-04-08T06:49:10.092184Z","iopub.execute_input":"2024-04-08T06:49:10.092914Z","iopub.status.idle":"2024-04-08T06:49:10.097202Z","shell.execute_reply.started":"2024-04-08T06:49:10.092880Z","shell.execute_reply":"2024-04-08T06:49:10.096380Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.imshow(human)","metadata":{"execution":{"iopub.status.busy":"2024-04-08T06:49:11.580128Z","iopub.execute_input":"2024-04-08T06:49:11.580495Z","iopub.status.idle":"2024-04-08T06:49:11.888837Z","shell.execute_reply.started":"2024-04-08T06:49:11.580465Z","shell.execute_reply":"2024-04-08T06:49:11.887826Z"},"trusted":true},"execution_count":null,"outputs":[]}]}