{"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":"2022-07-18T20:25:13.581823Z","iopub.execute_input":"2022-07-18T20:25:13.582487Z","iopub.status.idle":"2022-07-18T20:25:27.227319Z","shell.execute_reply.started":"2022-07-18T20:25:13.582417Z","shell.execute_reply":"2022-07-18T20:25:27.225520Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom torch.nn import functional as F\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nfrom torch.utils.data import Dataset\nfrom torch.utils.data import DataLoader\nimport torchmetrics\nimport torchvision\nimport segmentation_models_pytorch as smp\n\nfrom albumentations import (HorizontalFlip, VerticalFlip, \n                            ShiftScaleRotate, Normalize, Resize, \n                            Compose, GaussNoise, RandomRotate90)\nfrom albumentations.pytorch import ToTensorV2\n\nimport matplotlib.pyplot as plt\nimport pandas as pd\nimport numpy as np\nimport cv2\n\nimport os\nimport math\nimport time\nimport gc\nimport random","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-07-18T20:25:27.231151Z","iopub.execute_input":"2022-07-18T20:25:27.231797Z","iopub.status.idle":"2022-07-18T20:25:27.247183Z","shell.execute_reply.started":"2022-07-18T20:25:27.231737Z","shell.execute_reply":"2022-07-18T20:25:27.246137Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ENV","metadata":{}},{"cell_type":"code","source":"class config():\n    # ENVIRONMENT\n    \n    DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    \n    KAGGLE = True\n    DS_PATH = \"../input/hubmap-organ-segmentation/\" if KAGGLE else \"./\"\n    TRAIN_IMGS_PATH = os.path.join(DS_PATH, \"train_images\")\n    TEST_IMGS_PATH = os.path.join(DS_PATH, \"test_images\")\n    \n    TILES_DS_PATH = \"../input/hubmap-2022-256x256\" if KAGGLE else \"./hubmap-2022-256x256\"\n    TILES_IMGS_PATH = os.path.join(TILES_DS_PATH, \"train\")\n    TILES_MASKS_PATH = os.path.join(TILES_DS_PATH, \"masks\")\n    \n    # DATA\n    #TARGET_SHAPE = (512, 512)\n    VAL_SPLIT = 0.06\n    TILES_SIZE = 256\n    \n    #Imagenet parameters\n    MEAN = np.array([0.485, 0.456, 0.406])\n    STD = np.array([0.229, 0.224, 0.225])\n    \n    # TRAINING\n    BATCH_SIZE=32\n    EPOCHS = 20\n    LR = 0.05\n    \ndef seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    #the following line gives ~10% speedup\n    #but may lead to some stochasticity in the results \n    torch.backends.cudnn.benchmark = True\n    \nconfig()\nseed_everything(353)","metadata":{"execution":{"iopub.status.busy":"2022-07-18T20:25:27.249120Z","iopub.execute_input":"2022-07-18T20:25:27.250590Z","iopub.status.idle":"2022-07-18T20:25:27.271289Z","shell.execute_reply.started":"2022-07-18T20:25:27.250528Z","shell.execute_reply":"2022-07-18T20:25:27.269666Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# EDA","metadata":{}},{"cell_type":"code","source":"train_df = pd.read_csv(os.path.join(config.DS_PATH, \"train.csv\"))\ntest_df = pd.read_csv(os.path.join(config.DS_PATH, \"test.csv\"))\nsub_df = pd.read_csv(os.path.join(config.DS_PATH, \"sample_submission.csv\"))","metadata":{"execution":{"iopub.status.busy":"2022-07-18T20:25:27.275036Z","iopub.execute_input":"2022-07-18T20:25:27.276467Z","iopub.status.idle":"2022-07-18T20:25:27.478286Z","shell.execute_reply.started":"2022-07-18T20:25:27.276405Z","shell.execute_reply":"2022-07-18T20:25:27.477090Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df","metadata":{"execution":{"iopub.status.busy":"2022-07-18T20:25:27.480122Z","iopub.execute_input":"2022-07-18T20:25:27.480554Z","iopub.status.idle":"2022-07-18T20:25:27.527493Z","shell.execute_reply.started":"2022-07-18T20:25:27.480518Z","shell.execute_reply":"2022-07-18T20:25:27.526049Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df","metadata":{"execution":{"iopub.status.busy":"2022-07-18T20:25:27.529800Z","iopub.execute_input":"2022-07-18T20:25:27.530318Z","iopub.status.idle":"2022-07-18T20:25:27.545154Z","shell.execute_reply.started":"2022-07-18T20:25:27.530266Z","shell.execute_reply":"2022-07-18T20:25:27.543256Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def read_img(img_path, resize_shape = None):\n    img = cv2.imread(img_path)\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    \n    if resize_shape:\n        img = cv2.resize(img, resize_shape, interpolation = cv2.INTER_NEAREST)\n    return img\n\ndef img2tensor(img, dtype = np.float32):\n    \n    if img.ndim==2 :\n        img = np.expand_dims(img,2)\n        \n    img = np.transpose(img,(2,0,1))\n    return torch.from_numpy(img.astype(dtype, copy=False))","metadata":{"execution":{"iopub.status.busy":"2022-07-18T20:25:27.549978Z","iopub.execute_input":"2022-07-18T20:25:27.551100Z","iopub.status.idle":"2022-07-18T20:25:27.563952Z","shell.execute_reply.started":"2022-07-18T20:25:27.551034Z","shell.execute_reply":"2022-07-18T20:25:27.562125Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rle_decode(mask_rle, shape, color=1):\n    s = np.array(mask_rle.split(), dtype=int)\n\n    starts = s[0::2] - 1\n    lengths = s[1::2]\n    ends = starts + lengths\n\n    if len(shape)==3:\n        h, w, d = shape\n        img = np.zeros((h * w, d), dtype=np.float32)\n    else:\n        h, w = shape\n        img = np.zeros((h * w,), dtype=np.float32)\n\n    for lo, hi in zip(starts, ends):\n        img[lo : hi] = color\n        \n    return img.reshape(shape).T\n\ndef 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":"2022-07-18T20:25:27.566322Z","iopub.execute_input":"2022-07-18T20:25:27.567829Z","iopub.status.idle":"2022-07-18T20:25:27.585828Z","shell.execute_reply.started":"2022-07-18T20:25:27.567766Z","shell.execute_reply":"2022-07-18T20:25:27.584266Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Sample Image","metadata":{}},{"cell_type":"code","source":"sample_img = read_img(os.path.join(config.TRAIN_IMGS_PATH, \"10610.tiff\"))\nprint(sample_img.shape)\nplt.imshow(sample_img)","metadata":{"execution":{"iopub.status.busy":"2022-07-18T20:25:27.587800Z","iopub.execute_input":"2022-07-18T20:25:27.591689Z","iopub.status.idle":"2022-07-18T20:25:29.005548Z","shell.execute_reply.started":"2022-07-18T20:25:27.591625Z","shell.execute_reply":"2022-07-18T20:25:29.003959Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def show_mask(img_id):\n    rle_mask, img_width, img_height = train_df[train_df.id == int(img_id)][['rle', 'img_width', 'img_height']].squeeze().tolist()\n    img = read_img(os.path.join(config.TRAIN_IMGS_PATH, str(img_id) + \".tiff\"))\n    mask = rle_decode(rle_mask, shape=(img_height, img_width), color=1)\n    plt.figure(figsize=(8,8))\n    plt.imshow(img)\n    plt.imshow(mask, alpha=0.3)\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-07-18T20:25:29.009776Z","iopub.execute_input":"2022-07-18T20:25:29.010222Z","iopub.status.idle":"2022-07-18T20:25:29.019335Z","shell.execute_reply.started":"2022-07-18T20:25:29.010184Z","shell.execute_reply":"2022-07-18T20:25:29.017928Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"show_mask('10044')\nshow_mask('10488')\nshow_mask('62')","metadata":{"execution":{"iopub.status.busy":"2022-07-18T20:25:29.021286Z","iopub.execute_input":"2022-07-18T20:25:29.021780Z","iopub.status.idle":"2022-07-18T20:25:36.405226Z","shell.execute_reply.started":"2022-07-18T20:25:29.021739Z","shell.execute_reply":"2022-07-18T20:25:36.404175Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Preprocessing","metadata":{}},{"cell_type":"markdown","source":"### Tiled Dataset","metadata":{}},{"cell_type":"code","source":"class HuBMAPDataset(Dataset):\n    def __init__(self,train_df, tiles_size = config.TILES_SIZE, normalize = True, transforms = None, isVal=False):\n        self.image_ids = train_df['id'].astype(str).values\n        self.imgs_names = [fname for fname in os.listdir(config.TILES_IMGS_PATH) if fname.split('_')[0] in self.image_ids]\n        self.transforms = transforms\n\n    def __len__(self):\n        return len(self.imgs_names)\n    \n    def __getitem__(self,idx):\n        \n        image = read_img(os.path.join(config.TILES_IMGS_PATH, str(self.imgs_names[idx])))\n        mask = mask = cv2.imread(os.path.join(config.TILES_MASKS_PATH, str(self.imgs_names[idx])),cv2.IMREAD_GRAYSCALE)\n        \n        if self.transforms:\n            aug = self.transforms(image=image,mask=mask)\n            image,mask=aug['image'],aug['mask']\n        \n        return img2tensor((image/255.0 - config.MEAN)/config.STD),img2tensor(mask)\n    ","metadata":{"execution":{"iopub.status.busy":"2022-07-18T20:25:36.407159Z","iopub.execute_input":"2022-07-18T20:25:36.407601Z","iopub.status.idle":"2022-07-18T20:25:36.422383Z","shell.execute_reply.started":"2022-07-18T20:25:36.407562Z","shell.execute_reply":"2022-07-18T20:25:36.420599Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transforms = Compose([\n                    #Resize(config.TARGET_SHAPE[0],config.TARGET_SHAPE[1]),\n                    VerticalFlip(p=0.5),\n                    HorizontalFlip(p=0.5),\n                    RandomRotate90(p=0.5),\n                    ShiftScaleRotate(shift_limit=0.15,scale_limit = 0.05, p=0.5)\n                    ])\n\nhubmap_dataset = HuBMAPDataset(train_df,\n                               transforms = transforms\n                                )","metadata":{"execution":{"iopub.status.busy":"2022-07-18T20:25:36.423878Z","iopub.execute_input":"2022-07-18T20:25:36.424698Z","iopub.status.idle":"2022-07-18T20:25:36.526204Z","shell.execute_reply.started":"2022-07-18T20:25:36.424652Z","shell.execute_reply":"2022-07-18T20:25:36.524636Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_len = math.floor(len(hubmap_dataset)*config.VAL_SPLIT)\ntrain_len = len(hubmap_dataset) - val_len\n\ntrain_ds, val_ds = torch.utils.data.random_split(hubmap_dataset,[train_len,val_len],generator=torch.Generator().manual_seed(42))","metadata":{"execution":{"iopub.status.busy":"2022-07-18T20:25:36.528283Z","iopub.execute_input":"2022-07-18T20:25:36.528940Z","iopub.status.idle":"2022-07-18T20:25:36.538160Z","shell.execute_reply.started":"2022-07-18T20:25:36.528880Z","shell.execute_reply":"2022-07-18T20:25:36.536607Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loader=DataLoader(\n    train_ds,\n    batch_size=config.BATCH_SIZE,\n    shuffle=False\n)\nval_loader=DataLoader(\n    val_ds,\n    batch_size=config.BATCH_SIZE,\n    shuffle=False\n)","metadata":{"execution":{"iopub.status.busy":"2022-07-18T20:25:36.540252Z","iopub.execute_input":"2022-07-18T20:25:36.540902Z","iopub.status.idle":"2022-07-18T20:25:36.551842Z","shell.execute_reply.started":"2022-07-18T20:25:36.540852Z","shell.execute_reply":"2022-07-18T20:25:36.550594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_ds = HuBMAPDataset(train_df, transforms = transforms)\nsample_loader = DataLoader(sample_ds, config.BATCH_SIZE, shuffle = False)\n\nsample_imgs,sample_masks = next(iter(sample_loader))\nprint(sample_imgs.shape, sample_masks.shape)\n\nplt.figure(figsize=(16,16))\nfor i,(img,mask) in enumerate(zip(sample_imgs,sample_masks)):\n    img = ((img.permute(1,2,0)*config.STD + config.MEAN)*255.0).numpy().astype(np.uint8)\n    plt.subplot(8,8,i+1)\n    plt.imshow(img,vmin=0,vmax=255)\n    plt.imshow(mask.squeeze().numpy(), alpha=0.2)\n    plt.axis('off')\n    plt.subplots_adjust(wspace=None, hspace=None)\n    \ndel sample_ds,sample_loader,sample_imgs,sample_masks","metadata":{"execution":{"iopub.status.busy":"2022-07-18T20:25:36.553780Z","iopub.execute_input":"2022-07-18T20:25:36.554572Z","iopub.status.idle":"2022-07-18T20:25:40.720574Z","shell.execute_reply.started":"2022-07-18T20:25:36.554495Z","shell.execute_reply":"2022-07-18T20:25:40.718688Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"model = smp.DeepLabV3Plus('resnet34',\n                  encoder_weights=\"imagenet\",\n                )\nmodel = model.to(config.DEVICE)\nmodel","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-07-18T20:25:40.722475Z","iopub.execute_input":"2022-07-18T20:25:40.723266Z","iopub.status.idle":"2022-07-18T20:25:41.405156Z","shell.execute_reply.started":"2022-07-18T20:25:40.723213Z","shell.execute_reply":"2022-07-18T20:25:41.402714Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss and Optimization","metadata":{}},{"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\n    \nclass 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":"2022-07-18T20:25:41.407044Z","iopub.execute_input":"2022-07-18T20:25:41.407615Z","iopub.status.idle":"2022-07-18T20:25:41.422051Z","shell.execute_reply.started":"2022-07-18T20:25:41.407558Z","shell.execute_reply":"2022-07-18T20:25:41.420674Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"loss_fn=IoULoss(1)\noptimizer = torch.optim.Adam(model.parameters(),lr=config.LR)\nscheduler=torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer,mode=\"min\",patience=5,verbose=True)\n\nbest_val_loss=float(1e6)\nsince = time.time()\n\nfor n_epoch in range(1,config.EPOCHS):\n    \n    print(\"EPOCH : \"+str(n_epoch)+\"/\"+str(config.EPOCHS))\n    \n    running_train_loss=0.0\n    running_val_loss=0.0\n    \n    #TRAINING\n    model.train()\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(config.DEVICE),masks.to(config.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    #VALIDATION\n    model.eval()\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(config.DEVICE),masks.to(config.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":"2022-07-17T11:41:16.897458Z","iopub.execute_input":"2022-07-17T11:41:16.897933Z","iopub.status.idle":"2022-07-17T18:39:28.898754Z","shell.execute_reply.started":"2022-07-17T11:41:16.897898Z","shell.execute_reply":"2022-07-17T18:39:28.897336Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = torch.load(\"../input/hubmap-dlv3plus-model/best_model.pth\")\nmodel.eval()","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-07-18T20:25:41.424068Z","iopub.execute_input":"2022-07-18T20:25:41.424571Z","iopub.status.idle":"2022-07-18T20:25:41.581596Z","shell.execute_reply.started":"2022-07-18T20:25:41.424520Z","shell.execute_reply":"2022-07-18T20:25:41.580024Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"imgs = []\nreal_masks = []\npreds_masks = []\n\nwith torch.no_grad():\n    for val_batch_idx,val_batch in enumerate(val_loader):\n        images,masks = val_batch\n        images,masks=images.to(config.DEVICE),masks.to(config.DEVICE)\n        val_preds=model(images)\n        \n        imgs.extend(images)\n        real_masks.extend(masks)\n        preds_masks.extend(val_preds)","metadata":{"execution":{"iopub.status.busy":"2022-07-18T20:25:41.583065Z","iopub.execute_input":"2022-07-18T20:25:41.583441Z","iopub.status.idle":"2022-07-18T20:26:14.674044Z","shell.execute_reply.started":"2022-07-18T20:25:41.583407Z","shell.execute_reply":"2022-07-18T20:26:14.672715Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fix, axs = plt.subplots(len(imgs), 2, figsize=(8, 700))\n\nfor i in range(len(imgs)):\n    #real\n    axs[i][0].imshow(imgs[i].reshape(config.TILES_SIZE, config.TILES_SIZE, 3))\n    axs[i][0].imshow(real_masks[i].numpy().reshape(config.TILES_SIZE, config.TILES_SIZE, 1))\n    axs[i][0].set_title(\"Real\")\n    \n    #pred\n    axs[i][1].imshow(imgs[i].reshape(config.TILES_SIZE, config.TILES_SIZE, 3))\n    axs[i][1].imshow(preds_masks[i].numpy().reshape(config.TILES_SIZE, config.TILES_SIZE, 1))\n    axs[i][1].set_title(\"Pred\")","metadata":{"execution":{"iopub.status.busy":"2022-07-18T20:26:14.676330Z","iopub.execute_input":"2022-07-18T20:26:14.676937Z","iopub.status.idle":"2022-07-18T20:27:08.194638Z","shell.execute_reply.started":"2022-07-18T20:26:14.676895Z","shell.execute_reply":"2022-07-18T20:27:08.192495Z"},"trusted":true},"execution_count":null,"outputs":[]}]}