{"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 git+https://github.com/qubvel/segmentation_models.pytorch","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-07-22T08:33:38.668053Z","iopub.execute_input":"2022-07-22T08:33:38.668817Z","iopub.status.idle":"2022-07-22T08:34:16.421276Z","shell.execute_reply.started":"2022-07-22T08:33:38.668713Z","shell.execute_reply":"2022-07-22T08:34:16.419959Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#  v5 - \n* Added code to save only best model based on score and loss\n* Reduced LR , it was too high and training was very unstable\n* Added custom loss \n\n\n","metadata":{}},{"cell_type":"markdown","source":"# v6 - \n* Playing with LR and loss weights\n* Fixed save model bug","metadata":{}},{"cell_type":"markdown","source":"Thanks to https://www.kaggle.com/code/nayuts/hubmap-pytorch-smp-unet/notebook. I have repeated some code from here.","metadata":{}},{"cell_type":"code","source":"import os\nimport torch\nimport torch.nn as nn\nimport albumentations as A\nimport cv2\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport tqdm\nimport segmentation_models_pytorch as smp\nfrom sklearn.model_selection import StratifiedKFold , KFold , train_test_split\nimport tifffile as tiff\nimport glob","metadata":{"execution":{"iopub.status.busy":"2022-07-22T08:34:16.423323Z","iopub.execute_input":"2022-07-22T08:34:16.423614Z","iopub.status.idle":"2022-07-22T08:34:22.042794Z","shell.execute_reply.started":"2022-07-22T08:34:16.423587Z","shell.execute_reply":"2022-07-22T08:34:22.041636Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_seed(seed=2):\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    np.random.seed(seed)\n    torch.backends.cudnn.deterministic = True\nset_seed(99)","metadata":{"execution":{"iopub.status.busy":"2022-07-22T08:34:22.044749Z","iopub.execute_input":"2022-07-22T08:34:22.045627Z","iopub.status.idle":"2022-07-22T08:34:22.054695Z","shell.execute_reply.started":"2022-07-22T08:34:22.045589Z","shell.execute_reply":"2022-07-22T08:34:22.053796Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fold = 0\nnfolds = 5\ntrain_folds = [0]\nimsize = 384\ntrain_csv = '../input/hubmap-organ-segmentation/train.csv'\nBATCH_SIZE = 8\nDEVICE = ('cuda' if torch.cuda.is_available() else 'cpu')\nEPOCHS = 20\nNUM_WORKERS = 2\nSEED = 99\nTRAIN_PATH = '../input/hubmap-2022-512x512/train/'\nMASK_PATH = '../input/hubmap-2022-512x512/masks/'","metadata":{"execution":{"iopub.status.busy":"2022-07-22T08:34:22.057788Z","iopub.execute_input":"2022-07-22T08:34:22.058183Z","iopub.status.idle":"2022-07-22T08:34:22.112555Z","shell.execute_reply.started":"2022-07-22T08:34:22.058128Z","shell.execute_reply":"2022-07-22T08:34:22.111400Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# glob.glob('../input/create-tiled-dataset-hacking-the-human-body/tiled_dataset/*.png')[0].split('/')[-1]","metadata":{"execution":{"iopub.status.busy":"2022-07-22T08:34:22.114432Z","iopub.execute_input":"2022-07-22T08:34:22.114949Z","iopub.status.idle":"2022-07-22T08:34:22.122169Z","shell.execute_reply.started":"2022-07-22T08:34:22.114868Z","shell.execute_reply":"2022-07-22T08:34:22.121201Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pd.read_csv(train_csv)","metadata":{"execution":{"iopub.status.busy":"2022-07-22T08:34:22.123612Z","iopub.execute_input":"2022-07-22T08:34:22.123955Z","iopub.status.idle":"2022-07-22T08:34:22.420722Z","shell.execute_reply.started":"2022-07-22T08:34:22.123920Z","shell.execute_reply":"2022-07-22T08:34:22.419631Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# int(os.listdir(TRAIN_PATH)[2].split('_')[0])","metadata":{"execution":{"iopub.status.busy":"2022-07-22T08:34:22.422343Z","iopub.execute_input":"2022-07-22T08:34:22.422771Z","iopub.status.idle":"2022-07-22T08:34:22.428815Z","shell.execute_reply.started":"2022-07-22T08:34:22.422732Z","shell.execute_reply":"2022-07-22T08:34:22.427835Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# glob.glob(TRAIN_PATH + '/*.png')[0].split('/')[-1].","metadata":{"execution":{"iopub.status.busy":"2022-07-22T08:34:22.430140Z","iopub.execute_input":"2022-07-22T08:34:22.430722Z","iopub.status.idle":"2022-07-22T08:34:22.437709Z","shell.execute_reply.started":"2022-07-22T08:34:22.430679Z","shell.execute_reply":"2022-07-22T08:34:22.436695Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def img2tensor(img,dtype:np.dtype=np.float32):\n    if img.ndim==2 : img = np.expand_dims(img,2)\n    img = np.transpose(img,(2,0,1)) # C , H , W\n    return torch.from_numpy(img.astype(dtype, copy=False))","metadata":{"execution":{"iopub.status.busy":"2022-07-22T08:34:22.439194Z","iopub.execute_input":"2022-07-22T08:34:22.440175Z","iopub.status.idle":"2022-07-22T08:34:22.447285Z","shell.execute_reply.started":"2022-07-22T08:34:22.440121Z","shell.execute_reply":"2022-07-22T08:34:22.446175Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HuBMAPDataset(torch.utils.data.Dataset):\n    def __init__(self, fold=fold, train=True, tfms=None):\n        self.train = train\n        ids = pd.read_csv(train_csv).id.values\n        labels = pd.read_csv(train_csv).organ.values\n        kf = StratifiedKFold(n_splits=nfolds,random_state=SEED,shuffle=True)\n        ids = (ids[list(kf.split(ids,labels))[fold][0 if train else 1]]).tolist()\n        self.fnames = [fname for fname in os.listdir(TRAIN_PATH) if int(fname.split('_')[0]) in ids]\n        self.image_size = imsize\n        self.tfms = tfms\n#         print(f' type - {type(ids)} IDS - {ids}')\n#         print(f'fnames- :{self.fnames}')\n    \n    def __len__(self):\n        return len(self.fnames)\n    \n    def resize(self, img, interp):\n        return  cv2.resize(\n            img, (self.image_size, self.image_size), interpolation=interp)\n    \n    def __getitem__(self, idx):\n        fname = self.fnames[idx]\n#         print(fname)\n        img = cv2.cvtColor(cv2.imread(TRAIN_PATH + fname), cv2.COLOR_BGR2RGB)\n        mask = cv2.imread((MASK_PATH + fname),cv2.IMREAD_GRAYSCALE)\n#         print(img.shape)\n#         print(mask.shape)\n        if self.tfms is not None:\n            augmented = self.tfms(image=img,mask=mask)\n            img,mask = augmented['image'],augmented['mask']\n        return img2tensor(self.resize(img , cv2.INTER_NEAREST)) , img2tensor(self.resize(mask , cv2.INTER_NEAREST))","metadata":{"execution":{"iopub.status.busy":"2022-07-22T08:34:22.452170Z","iopub.execute_input":"2022-07-22T08:34:22.452587Z","iopub.status.idle":"2022-07-22T08:34:22.467396Z","shell.execute_reply.started":"2022-07-22T08:34:22.452561Z","shell.execute_reply":"2022-07-22T08:34:22.466204Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def albu_aug(p=1.0):\n    return A.Compose([\n        A.HorizontalFlip(),\n        A.VerticalFlip(),\n        A.RandomRotate90(),\n        A.ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.2, rotate_limit=15, p=0.9, \n                         border_mode=cv2.BORDER_REFLECT),\n        A.OneOf([\n            A.OpticalDistortion(p=0.3),\n            A.GridDistortion(p=.1),\n            A.PiecewiseAffine(p=0.3),\n        ], p=0.3),\n        A.OneOf([\n            A.HueSaturationValue(10,15,10),\n            A.CLAHE(clip_limit=2),\n            A.RandomBrightnessContrast(),            \n        ], p=0.3),\n    ], p=p)","metadata":{"execution":{"iopub.status.busy":"2022-07-22T08:34:22.469589Z","iopub.execute_input":"2022-07-22T08:34:22.470428Z","iopub.status.idle":"2022-07-22T08:34:22.480319Z","shell.execute_reply.started":"2022-07-22T08:34:22.470366Z","shell.execute_reply":"2022-07-22T08:34:22.479148Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds = HuBMAPDataset(tfms=albu_aug())\ndl = torch.utils.data.DataLoader(ds,batch_size=64,shuffle=False,num_workers=NUM_WORKERS)\nit = iter(dl)\nimgs,masks = next(it)\n\nplt.figure(figsize=(16,16))\nfor i,(img,mask) in enumerate(zip(imgs,masks)):\n    img = ((img.permute(1,2,0))).numpy().astype(np.uint8)  # H , W , C\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 ds,dl,imgs,masks","metadata":{"execution":{"iopub.status.busy":"2022-07-22T08:34:22.481863Z","iopub.execute_input":"2022-07-22T08:34:22.482783Z","iopub.status.idle":"2022-07-22T08:34:41.740238Z","shell.execute_reply.started":"2022-07-22T08:34:22.482755Z","shell.execute_reply":"2022-07-22T08:34:41.739017Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def init_model():\n    model =  smp.Unet(\n                 encoder_name='efficientnet-b4',\n                 encoder_weights='imagenet',\n                 in_channels=3,\n                 classes=1)\n    return model","metadata":{"execution":{"iopub.status.busy":"2022-07-22T08:34:41.742042Z","iopub.execute_input":"2022-07-22T08:34:41.743245Z","iopub.status.idle":"2022-07-22T08:34:41.763741Z","shell.execute_reply.started":"2022-07-22T08:34:41.743207Z","shell.execute_reply":"2022-07-22T08:34:41.762627Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DiceCoef(nn.Module):\n    def __init__(self, weight=None, size_average=True):\n        super().__init__()\n\n    def forward(self, y_pred, y_true, smooth=1.):\n        y_true = y_true.view(-1)\n        y_pred = y_pred.view(-1)\n        \n        #Round off y_pred\n        y_pred = torch.round((y_pred - y_pred.min()) / (y_pred.max() - y_pred.min()))\n        \n        intersection = (y_true * y_pred).sum()\n        dice = (2.0*intersection + smooth)/(y_true.sum() + y_pred.sum() + smooth)\n        \n        return dice","metadata":{"execution":{"iopub.status.busy":"2022-07-22T08:34:41.765651Z","iopub.execute_input":"2022-07-22T08:34:41.766857Z","iopub.status.idle":"2022-07-22T08:34:41.788252Z","shell.execute_reply.started":"2022-07-22T08:34:41.766813Z","shell.execute_reply":"2022-07-22T08:34:41.787209Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\ngc.collect()\ntorch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2022-07-22T08:34:41.789775Z","iopub.execute_input":"2022-07-22T08:34:41.790771Z","iopub.status.idle":"2022-07-22T08:34:42.208358Z","shell.execute_reply.started":"2022-07-22T08:34:41.790734Z","shell.execute_reply":"2022-07-22T08:34:42.206517Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_df(df):\n    fig,ax = plt.subplots(1,2,figsize=(15,5))\n    ax[0].plot(df['Train_loss'])\n    ax[0].plot(df['Val_loss'])\n    ax[0].legend()\n    ax[0].set_title('Loss')\n    ax[1].plot(df['Train_Dice'])\n    ax[1].plot(df['Val_Dice'])\n    ax[1].legend()\n    ax[1].set_title('Dice')","metadata":{"execution":{"iopub.status.busy":"2022-07-22T08:34:42.213755Z","iopub.execute_input":"2022-07-22T08:34:42.214111Z","iopub.status.idle":"2022-07-22T08:34:42.230990Z","shell.execute_reply.started":"2022-07-22T08:34:42.214078Z","shell.execute_reply":"2022-07-22T08:34:42.229822Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CustomLoss(nn.Module):\n    def __init__(self):\n        super(CustomLoss,self).__init__()\n        self.diceloss = smp.losses.DiceLoss(mode='binary')\n        self.binloss = smp.losses.SoftBCEWithLogitsLoss(reduction = 'mean' , smooth_factor = 0.1)\n\n    def forward(self, output, mask):\n        dice = self.diceloss(outputs,mask)\n        bce = self.binloss(outputs , mask)\n        loss = dice * 0.7 + bce * 0.3\n        return loss","metadata":{"execution":{"iopub.status.busy":"2022-07-22T08:34:42.233462Z","iopub.execute_input":"2022-07-22T08:34:42.234763Z","iopub.status.idle":"2022-07-22T08:34:42.243872Z","shell.execute_reply.started":"2022-07-22T08:34:42.234725Z","shell.execute_reply":"2022-07-22T08:34:42.242766Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f\"Running on device :  {DEVICE}\" )\nfor fold in (train_folds):\n    val_losses = []\n    losses = []\n    train_scores=[]\n    val_scores = []\n    best_loss = 999\n    best_score = 0\n    ds_train = HuBMAPDataset(fold=fold, train=True, tfms=albu_aug())\n    ds_val = HuBMAPDataset(fold=fold, train=False)\n    dataloader_train = torch.utils.data.DataLoader(ds_train,batch_size=BATCH_SIZE, shuffle=True,num_workers=NUM_WORKERS)\n    dataloader_val = torch.utils.data.DataLoader(ds_val,batch_size=BATCH_SIZE, shuffle=False,num_workers=NUM_WORKERS)\n    model = init_model().to(DEVICE)\n    \n    optimizer = torch.optim.Adam([\n        {'params': model.decoder.parameters(), 'lr': 5e-5}, \n        {'params': model.encoder.parameters(), 'lr': 8e-5},  \n    ])\n    scheduler = torch.optim.lr_scheduler.OneCycleLR(optimizer=optimizer, pct_start=0.1, div_factor=1e3, \n                                              max_lr=1e-3, epochs=EPOCHS, steps_per_epoch=len(dataloader_train))\n    \n    loss_func = CustomLoss()\n    dice_coe = DiceCoef()\n    print(f\"########FOLD: {fold}##############\")\n    \n    for epoch in tqdm.notebook.tqdm(range(EPOCHS)):\n        \n        ### Train ###########################################################################################\n        \n        model.train()\n        train_loss = 0\n        score = 0\n        for data in tqdm.notebook.tqdm(dataloader_train ,total = len(dataloader_train)):\n            optimizer.zero_grad()\n            img, mask = data\n            img = img.to(DEVICE)\n            mask = mask.to(DEVICE)\n        \n            outputs = model(img)\n    \n            loss = loss_func(outputs, mask)\n            loss.backward()\n            optimizer.step()\n            scheduler.step()\n            train_loss += loss.item()\n            score += dice_coe(outputs,mask).item()\n            \n        train_loss /= len(dataloader_train)\n        score /= len(dataloader_train)\n        losses.append(train_loss)\n        train_scores.append(score)\n        print(f\"FOLD: {fold}, EPOCH: {epoch + 1}, train_loss: {train_loss} , Dice coe : {score} \") #\n        \n        \n        gc.collect()\n        torch.cuda.empty_cache()\n        \n        ### Validation ####################################################################################\n        \n        model.eval()\n        with torch.no_grad():\n            valid_loss = 0\n            val_score = 0\n            for data in dataloader_val:\n                img, mask = data\n                img = img.to(DEVICE)\n                mask = mask.to(DEVICE)\n\n                outputs = model(img)\n\n                loss = loss_func(outputs, mask)\n\n                valid_loss += loss.item()\n                val_score += dice_coe(outputs,mask).item()\n            valid_loss /= len(dataloader_val)\n            val_score /= len(dataloader_val)\n            val_losses.append(valid_loss)\n            val_scores.append(val_score)\n            print(f\"FOLD: {fold}, EPOCH: {epoch + 1}, valid_loss: {valid_loss} , Val Dice COE : {val_score}\") #\n            gc.collect()\n            torch.cuda.empty_cache()\n        if val_score > best_score:\n            best_score = val_score\n            torch.save(model.state_dict(), f\"FOLD{fold}_best_score.pth\")\n            print(f\"Saved model for best score : FOLD{fold}_best_score.pth\")\n        if valid_loss < best_loss:\n            best_loss = valid_loss\n            torch.save(model.state_dict(), f\"FOLD{fold}_best_loss.pth\")\n            print(f\"Saved model for best loss : FOLD{fold}_best_loss.pth\")    \n    ###Save model\n\n    \n    column_names = ['Train_loss','Val_loss','Train_Dice','Val_Dice']\n    df = pd.DataFrame(np.stack([losses,val_losses,train_scores,val_scores],axis=1),columns=column_names)\n    print(f\" ################# FOLD {fold} #####################\")\n    plot_df(df)\n    plt.show(block=False)\n    df.to_csv(f\"logs_fold{fold}.csv\")","metadata":{"execution":{"iopub.status.busy":"2022-07-22T08:38:46.432641Z","iopub.execute_input":"2022-07-22T08:38:46.432997Z","iopub.status.idle":"2022-07-22T08:43:04.159320Z","shell.execute_reply.started":"2022-07-22T08:38:46.432968Z","shell.execute_reply":"2022-07-22T08:43:04.157874Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}