{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install segmentation-models-pytorch ","metadata":{"execution":{"iopub.status.busy":"2021-12-19T22:13:02.641395Z","iopub.execute_input":"2021-12-19T22:13:02.642111Z","iopub.status.idle":"2021-12-19T22:13:18.684290Z","shell.execute_reply.started":"2021-12-19T22:13:02.641926Z","shell.execute_reply":"2021-12-19T22:13:18.683369Z"},"_kg_hide-input":false,"collapsed":true,"jupyter":{"outputs_hidden":true},"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import segmentation_models_pytorch  as smp\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport torch\nimport torchvision\nimport cv2\nimport math\nimport time\nfrom tqdm import tqdm\nfrom torch.nn import functional as F\nimport torch.backends.cudnn as cudnn\nimport torchmetrics\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nimport gc\nimport torch.nn as nn\nfrom torch.utils.data import Dataset\nfrom torch.utils.data import DataLoader\n\nfrom albumentations import (HorizontalFlip, VerticalFlip, \n                            ShiftScaleRotate, Normalize, Resize, \n                            Compose, GaussNoise)\nfrom albumentations.pytorch import ToTensorV2\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2021-12-19T22:16:28.420947Z","iopub.execute_input":"2021-12-19T22:16:28.421242Z","iopub.status.idle":"2021-12-19T22:16:28.428915Z","shell.execute_reply.started":"2021-12-19T22:16:28.421211Z","shell.execute_reply":"2021-12-19T22:16:28.428141Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv('../input/sartorius-cell-instance-segmentation/train.csv')\ntrain_df","metadata":{"execution":{"iopub.status.busy":"2021-12-19T22:16:32.340295Z","iopub.execute_input":"2021-12-19T22:16:32.340595Z","iopub.status.idle":"2021-12-19T22:16:32.715697Z","shell.execute_reply.started":"2021-12-19T22:16:32.340558Z","shell.execute_reply":"2021-12-19T22:16:32.714877Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df['cell_type'].unique()","metadata":{"execution":{"iopub.status.busy":"2021-12-19T22:13:37.385311Z","iopub.execute_input":"2021-12-19T22:13:37.385913Z","iopub.status.idle":"2021-12-19T22:13:37.405934Z","shell.execute_reply.started":"2021-12-19T22:13:37.385876Z","shell.execute_reply":"2021-12-19T22:13:37.404643Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(train_df['height'].unique())\nprint(train_df['width'].unique())","metadata":{"execution":{"iopub.status.busy":"2021-12-19T22:13:39.234396Z","iopub.execute_input":"2021-12-19T22:13:39.235358Z","iopub.status.idle":"2021-12-19T22:13:39.242555Z","shell.execute_reply.started":"2021-12-19T22:13:39.235313Z","shell.execute_reply":"2021-12-19T22:13:39.241675Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df = pd.read_csv('../input/sartorius-cell-instance-segmentation/sample_submission.csv')\nsub_df","metadata":{"execution":{"iopub.status.busy":"2021-12-19T22:13:39.921691Z","iopub.execute_input":"2021-12-19T22:13:39.922385Z","iopub.status.idle":"2021-12-19T22:13:39.942037Z","shell.execute_reply.started":"2021-12-19T22:13:39.922344Z","shell.execute_reply":"2021-12-19T22:13:39.941420Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TEST_IMGS_PATH = \"../input/sartorius-cell-instance-segmentation/test/\"\nTRAIN_IMGS_PATH = \"../input/sartorius-cell-instance-segmentation/train/\"\n\nIMGS_WIDTH = 704\nIMGS_HEIGHT = 520\n\nRESNET_MEAN = (0.485, 0.456, 0.406)\nRESNET_STD = (0.229, 0.224, 0.225)\n\nTARGET_IMGS_HEIGHT=512\nTARGET_IMGS_WIDTH=512\n\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint(\"Using : \",DEVICE)","metadata":{"execution":{"iopub.status.busy":"2021-12-19T22:13:41.139274Z","iopub.execute_input":"2021-12-19T22:13:41.139558Z","iopub.status.idle":"2021-12-19T22:13:41.146607Z","shell.execute_reply.started":"2021-12-19T22:13:41.139527Z","shell.execute_reply":"2021-12-19T22:13:41.145710Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rle_decode(x,shape,color=1):\n    \n    out = np.zeros((shape[0]*shape[1],shape[2]))\n    x=[int(i) for i in x.split(\" \")]\n    for i in range(0,len(x),2):\n        out[ x[i]:(x[i]+x[i+1]) ]=color\n\n    return np.reshape(out,shape)","metadata":{"execution":{"iopub.status.busy":"2021-12-19T22:13:42.911607Z","iopub.execute_input":"2021-12-19T22:13:42.911964Z","iopub.status.idle":"2021-12-19T22:13:42.920417Z","shell.execute_reply.started":"2021-12-19T22:13:42.911928Z","shell.execute_reply":"2021-12-19T22:13:42.918567Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rle_encode(x):\n    out=[]\n    x=x.flatten()\n    for i in range(0,x.shape[0]-1):\n        if(x[i]==1):\n            count=1\n            out.append(str(i))\n            i+=1\n            while(x[i]==1):\n                count+=1\n                i+=1\n            out.append(str(count))\n    return \" \".join(out)","metadata":{"execution":{"iopub.status.busy":"2021-12-19T22:13:43.651382Z","iopub.execute_input":"2021-12-19T22:13:43.651888Z","iopub.status.idle":"2021-12-19T22:13:43.659508Z","shell.execute_reply.started":"2021-12-19T22:13:43.651829Z","shell.execute_reply":"2021-12-19T22:13:43.658492Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def show_masks(img_id):\n    rle_masks = train_df[train_df[\"id\"]==img_id]['annotation'].tolist()\n    \n    cells_mask = np.zeros((IMGS_HEIGHT,IMGS_WIDTH,3))\n    for rle_mask in rle_masks:\n        cells_mask+=rle_decode(rle_mask,(IMGS_HEIGHT,IMGS_WIDTH,3),color=np.random.rand(3))\n    img = cv2.cvtColor(cv2.imread(TRAIN_IMGS_PATH+img_id+\".png\"),cv2.COLOR_BGR2RGB)\n    plt.figure(figsize=(16,32))\n    plt.imshow(img)\n    plt.imshow(cells_mask,alpha=0.3)\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2021-12-19T22:13:44.772909Z","iopub.execute_input":"2021-12-19T22:13:44.773544Z","iopub.status.idle":"2021-12-19T22:13:44.781067Z","shell.execute_reply.started":"2021-12-19T22:13:44.773504Z","shell.execute_reply":"2021-12-19T22:13:44.780079Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"show_masks('0030fd0e6378')\nshow_masks('ffdb3cc02eef')\nshow_masks('0df9d6419078')","metadata":{"execution":{"iopub.status.busy":"2021-12-19T20:17:34.247166Z","iopub.execute_input":"2021-12-19T20:17:34.247791Z","iopub.status.idle":"2021-12-19T20:17:37.427543Z","shell.execute_reply.started":"2021-12-19T20:17:34.247754Z","shell.execute_reply":"2021-12-19T20:17:37.426856Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SartoriusCellDataset(Dataset):\n    def __init__(self,train_df,train_imgs_path,transforms,isVal=False):\n        self.train_df=train_df\n        self.image_ids=np.unique(train_df['id']).tolist()\n        self.train_imgs_path=train_imgs_path\n        self.transforms = transforms\n    def __len__(self):\n        return len(self.image_ids)\n    def __getitem__(self,idx):\n        \n        image = cv2.cvtColor( cv2.imread( self.train_imgs_path +  self.image_ids[idx] + \".png\"),cv2.COLOR_BGR2RGB)\n        mask = np.zeros((image.shape[0],image.shape[1],1),dtype=np.float32)\n        rle_masks=self.train_df[train_df[\"id\"]==self.image_ids[idx]]['annotation'].tolist()\n        \n        for rle_mask in rle_masks:\n            mask+=rle_decode(rle_mask,(image.shape[0],image.shape[1],1)).astype(np.float32)\n        mask = mask.clip(0, 1)\n        \n        if self.transforms:\n            aug = self.transforms(image=image,mask=mask)\n            image,mask=aug['image'],aug['mask']\n            \n        return image,mask.reshape((1,image.shape[1],image.shape[2]))","metadata":{"execution":{"iopub.status.busy":"2021-12-19T22:13:47.878914Z","iopub.execute_input":"2021-12-19T22:13:47.879789Z","iopub.status.idle":"2021-12-19T22:13:47.891276Z","shell.execute_reply.started":"2021-12-19T22:13:47.879752Z","shell.execute_reply":"2021-12-19T22:13:47.890690Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transforms = Compose([Resize(TARGET_IMGS_HEIGHT,TARGET_IMGS_WIDTH),\n                    Normalize(mean=RESNET_MEAN,std=RESNET_STD),\n                    VerticalFlip(p=0.5),\n                    HorizontalFlip(p=0.5),\n                    ToTensorV2()])\ncell_dataset = SartoriusCellDataset(train_df,\n                          TRAIN_IMGS_PATH,\n                          transforms)","metadata":{"execution":{"iopub.status.busy":"2021-12-19T22:13:48.786324Z","iopub.execute_input":"2021-12-19T22:13:48.786819Z","iopub.status.idle":"2021-12-19T22:13:48.856315Z","shell.execute_reply.started":"2021-12-19T22:13:48.786784Z","shell.execute_reply":"2021-12-19T22:13:48.855196Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_split=0.2\nval_len=math.floor(len(cell_dataset)*val_split)\ntrain_len=len(cell_dataset) - val_len\ntrain_ds,val_ds = torch.utils.data.random_split(cell_dataset,[train_len,val_len],generator=torch.Generator().manual_seed(42))","metadata":{"execution":{"iopub.status.busy":"2021-12-19T22:13:50.413497Z","iopub.execute_input":"2021-12-19T22:13:50.414145Z","iopub.status.idle":"2021-12-19T22:13:50.427103Z","shell.execute_reply.started":"2021-12-19T22:13:50.414104Z","shell.execute_reply":"2021-12-19T22:13:50.426250Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_img,sample_mask=train_ds[1]\nprint(sample_img.shape,\"\\n\",sample_mask.shape)\nprint(sample_img.dtype)\nprint(sample_mask.dtype)","metadata":{"execution":{"iopub.status.busy":"2021-12-19T22:13:51.326609Z","iopub.execute_input":"2021-12-19T22:13:51.326946Z","iopub.status.idle":"2021-12-19T22:13:51.489042Z","shell.execute_reply.started":"2021-12-19T22:13:51.326910Z","shell.execute_reply":"2021-12-19T22:13:51.488125Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.imshow(sample_img[0],cmap=\"gray\")","metadata":{"execution":{"iopub.status.busy":"2021-12-19T22:13:52.908742Z","iopub.execute_input":"2021-12-19T22:13:52.909891Z","iopub.status.idle":"2021-12-19T22:13:53.207794Z","shell.execute_reply.started":"2021-12-19T22:13:52.909848Z","shell.execute_reply":"2021-12-19T22:13:53.206919Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.imshow(sample_mask[0])","metadata":{"execution":{"iopub.status.busy":"2021-12-19T22:13:55.053621Z","iopub.execute_input":"2021-12-19T22:13:55.053943Z","iopub.status.idle":"2021-12-19T22:13:55.263751Z","shell.execute_reply.started":"2021-12-19T22:13:55.053912Z","shell.execute_reply":"2021-12-19T22:13:55.262780Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BATCH_SIZE=32\ntrain_loader=DataLoader(\n    train_ds,\n    batch_size=BATCH_SIZE,\n    shuffle=False\n)\nval_loader=DataLoader(\n    val_ds,\n    batch_size=BATCH_SIZE,\n    shuffle=False\n)","metadata":{"execution":{"iopub.status.busy":"2021-12-19T22:13:56.709821Z","iopub.execute_input":"2021-12-19T22:13:56.710351Z","iopub.status.idle":"2021-12-19T22:13:56.715100Z","shell.execute_reply.started":"2021-12-19T22:13:56.710313Z","shell.execute_reply":"2021-12-19T22:13:56.714424Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nmodel = smp.DeepLabV3Plus('resnet34',\n                  encoder_weights=\"imagenet\",\n                )\nmodel","metadata":{"execution":{"iopub.status.busy":"2021-12-19T22:13:59.835774Z","iopub.execute_input":"2021-12-19T22:13:59.836564Z","iopub.status.idle":"2021-12-19T22:14:04.528423Z","shell.execute_reply.started":"2021-12-19T22:13:59.836520Z","shell.execute_reply":"2021-12-19T22:14:04.527328Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"model = torch.hub.load('pytorch/vision:v0.10.0', 'deeplabv3_resnet50', pretrained=True)\n","metadata":{"execution":{"iopub.status.busy":"2021-12-19T20:14:07.740119Z","iopub.status.idle":"2021-12-19T20:14:07.740652Z","shell.execute_reply.started":"2021-12-19T20:14:07.740426Z","shell.execute_reply":"2021-12-19T20:14:07.740451Z"}}},{"cell_type":"code","source":"class dice_bce_loss(nn.Module):\n    def __init__(self, weight=None, size_average=True):\n        super(dice_bce_loss, self).__init__()\n\n    def forward(self, inputs, targets, smooth=1):\n        inputs = inputs.view(-1)\n        targets = targets.view(-1)\n        \n        intersection = (inputs * targets).sum()                            \n        dice_loss = 1 - (2.*intersection + smooth)/(inputs.sum() + targets.sum() + smooth)  \n        bce = F.binary_cross_entropy(inputs, targets, reduction='mean')\n        dice_bce = bce + dice_loss\n        \n        return dice_bce","metadata":{"execution":{"iopub.status.busy":"2021-12-19T20:17:42.513472Z","iopub.execute_input":"2021-12-19T20:17:42.515699Z","iopub.status.idle":"2021-12-19T20:17:42.527697Z","shell.execute_reply.started":"2021-12-19T20:17:42.515657Z","shell.execute_reply":"2021-12-19T20:17:42.526280Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class IoULoss(nn.Module):\n    def __init__(self, weight=None, size_average=True):\n        super(IoULoss, self).__init__()\n\n    def forward(self, inputs, targets, smooth=1):\n        \n        #comment out if your model contains a sigmoid or equivalent activation layer\n        inputs = F.sigmoid(inputs)       \n        \n        #flatten label and prediction tensors\n        inputs = inputs.view(-1)\n        targets = targets.view(-1)\n        \n        #intersection is equivalent to True Positive count\n        #union is the mutually inclusive area of all labels & predictions \n        intersection = (inputs * targets).sum()\n        total = (inputs + targets).sum()\n        union = total - intersection \n        \n        IoU = (intersection + smooth)/(union + smooth)\n                \n        return 1 - IoU","metadata":{"execution":{"iopub.status.busy":"2021-12-19T20:17:42.578510Z","iopub.execute_input":"2021-12-19T20:17:42.578847Z","iopub.status.idle":"2021-12-19T20:17:42.592484Z","shell.execute_reply.started":"2021-12-19T20:17:42.578812Z","shell.execute_reply":"2021-12-19T20:17:42.591623Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"EPOCHS=50\nLEARNING_RATE=5e-4\n\nmodel.to(DEVICE)\n\nloss_fn=IoULoss(1)\noptimizer = torch.optim.Adam(model.parameters(),lr=LEARNING_RATE)\nscheduler=torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer,mode=\"min\",patience=5,verbose=True)\nbest_val_loss=float(1e6)\nsince = time.time()\nfor n_epoch in range(1,EPOCHS):\n    \n    print(\"EPOCH : \"+str(n_epoch)+\"/\"+str(EPOCHS))\n    \n    \n    running_train_loss=0.0\n    running_val_loss=0.0\n    \n    \n    model.train()\n    #TRAINING\n    for train_batch_idx,train_batch in enumerate(train_loader):\n        optimizer.zero_grad()\n\n        #PREDICT\n        images,masks = train_batch\n        images,masks=images.to(DEVICE),masks.to(DEVICE)\n        \n       \n        preds=model(images)\n        train_loss=loss_fn(preds,masks)\n        \n        gc.collect()\n        del train_batch\n        del images\n        del masks\n        \n        #BACKPROPAGATION\n        train_loss.backward()\n        optimizer.step()\n        running_train_loss += train_loss.item()\n        \n        gc.collect()\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n\n\n    model.eval()\n    #VALIDATION\n    with torch.no_grad():\n        for val_batch_idx,val_batch in enumerate (val_loader):\n\n            #Predict\n            images,masks = val_batch\n            images,masks=images.to(DEVICE),masks.to(DEVICE)\n            val_preds=model(images)\n            val_loss=loss_fn(val_preds,masks)\n\n            gc.collect()\n            del val_batch\n            del images\n            del masks\n\n            running_val_loss+=val_loss.item()\n        \n    running_train_loss /= train_batch_idx+1\n    running_val_loss /= val_batch_idx+1\n    \n    #Reduce LR on Plateau\n    scheduler.step(running_val_loss) \n\n    print(f\"EPOCH : {n_epoch} Train Loss : {running_train_loss:.5f}, Val Loss : {running_val_loss:.5f}\")\n    if(running_val_loss < best_val_loss):\n        torch.save(model.state_dict(), \"/kaggle/working/best_model.pth\")\n        print(\"Model Saved\")\n        best_val_loss=running_val_loss\ntime_elapsed = time.time() - since\nprint('Training complete in {:.0f}m {:.0f}s'.format(\ntime_elapsed // 60, time_elapsed % 60))\nprint('Best Val Loss: {:4f}'.format(best_val_loss))","metadata":{"execution":{"iopub.status.busy":"2021-12-19T20:17:42.597018Z","iopub.execute_input":"2021-12-19T20:17:42.597306Z","iopub.status.idle":"2021-12-19T21:20:00.737149Z","shell.execute_reply.started":"2021-12-19T20:17:42.597271Z","shell.execute_reply":"2021-12-19T21:20:00.736452Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Inference test on training data","metadata":{}},{"cell_type":"code","source":"model.load_state_dict(torch.load(\"../input/deeplabv3plus-resnet34-sartorius/best_model.pth\",map_location=torch.device('cpu')))\nmodel.to(DEVICE)","metadata":{"execution":{"iopub.status.busy":"2021-12-19T22:15:05.329273Z","iopub.execute_input":"2021-12-19T22:15:05.330154Z","iopub.status.idle":"2021-12-19T22:15:07.824642Z","shell.execute_reply.started":"2021-12-19T22:15:05.330114Z","shell.execute_reply":"2021-12-19T22:15:07.823639Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for batch_idx,batch in enumerate(train_loader):\n    images,masks=batch\n    if torch.cuda.is_available():\n        images,masks=images.cuda(),masks.cuda()\n    preds=model(images)\n    print(preds.shape)\n    fig,axs=plt.subplots(16,2,figsize=(10,80))\n    images,masks=images.cpu(),masks.cpu()\n    preds=preds.cpu().detach().numpy()\n    print(preds[0].max(),preds[0].min())\n    for i in range(16):\n        #axs[i][0].imshow(images[i].reshape(512,512,3))\n        axs[i][0].imshow(masks[i].reshape(512,512,1))\n        axs[i][0].title.set_text(\"Ground truth\")\n        #axs[i][1].imshow(images[i].reshape(512,512,3))\n        axs[i][1].imshow(preds[i].reshape(512,512,1))\n        axs[i][1].title.set_text(\"Prediction\")\n\n    plt.subplots_adjust(wspace=0.1)\n\n    plt.show()\n    break","metadata":{"execution":{"iopub.status.busy":"2021-12-19T22:15:14.480544Z","iopub.execute_input":"2021-12-19T22:15:14.480851Z","iopub.status.idle":"2021-12-19T22:15:53.935959Z","shell.execute_reply.started":"2021-12-19T22:15:14.480821Z","shell.execute_reply":"2021-12-19T22:15:53.934726Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# INFERENCE (PART 2) HERE:\n# [Sartorius-cell-segmentation-Deeplabv3-INFERENCE](https://www.kaggle.com/albertozorzetto/sartorius-cell-segmentation-deeplabv3-inference/edit)","metadata":{}}]}