{"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":"markdown","source":"<div style=\"height:200px;width:100%;margin: 0;\">\n    <img src=\"https://storage.googleapis.com/kaggle-competitions/kaggle/34547/logos/header.png?t=2022-02-15-22-37-27\" style=\"width:100%;\" />\n</div>","metadata":{}},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"goal\"><center>Goal</center></h3>\n\nThe scope of this notebook is to train people on how to train the baseline of a CoAt for this dataset.<br>\nCoAT currently has the highest submission, so give this notebook a try.<br>\nIf you really want a medal, you have to tweek hyperparameters and work hard.<br>\nThe [Inference notebook](https://www.kaggle.com/code/alincijov/inference-hubmap-coat) for this notebook.<br>\nIf you appreciate the work, don't forget to <b>upvote</b> !","metadata":{}},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"install\"><center>Installation</center></h3>","metadata":{}},{"cell_type":"code","source":"!pip install -qq torch==1.7.1+cu110 torchvision==0.8.2+cu110 torchaudio==0.7.2 -f https://download.pytorch.org/whl/torch_stable.html\n!pip install -qq git+https://github.com/qubvel/segmentation_models.pytorch\n!pip install -qq timm==0.4.12\n!pip install -qq einops","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-09-05T11:26:56.519794Z","iopub.execute_input":"2022-09-05T11:26:56.520778Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"libraries\"><center>Libraries</center></h3>","metadata":{}},{"cell_type":"code","source":"import os\nimport gc\nimport sys\nimport glob\n\nimport torch\nimport torch.nn as nn\nimport albumentations as A\n\nimport cv2\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\n\nimport tqdm\nimport segmentation_models_pytorch as smp\nfrom sklearn.model_selection import StratifiedKFold\n\nimport tifffile as tiff\n\ntorch.backends.cudnn.benchmark = True","metadata":{"execution":{"iopub.status.busy":"2022-09-05T10:31:52.270058Z","iopub.execute_input":"2022-09-05T10:31:52.27055Z","iopub.status.idle":"2022-09-05T10:31:57.236306Z","shell.execute_reply.started":"2022-09-05T10:31:52.270504Z","shell.execute_reply":"2022-09-05T10:31:57.235205Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"config\"><center>Configuration</center></h3>","metadata":{}},{"cell_type":"code","source":"fold = 0\nnfolds = 3\nimsize = 512\ntrain_csv = '../input/hubmap-organ-segmentation/train.csv'\nBATCH_SIZE = 2\nDEVICE = ('cuda' if torch.cuda.is_available() else 'cpu')\nEPOCHS = 200\nNUM_WORKERS = 4\nSEED = 24\nTRAIN_PATH = '../input/ori-r2-512/train/'\nMASK_PATH = '../input/ori-r2-512/masks/'","metadata":{"execution":{"iopub.status.busy":"2022-09-05T10:48:05.62653Z","iopub.execute_input":"2022-09-05T10:48:05.62722Z","iopub.status.idle":"2022-09-05T10:48:05.633641Z","shell.execute_reply.started":"2022-09-05T10:48:05.62718Z","shell.execute_reply":"2022-09-05T10:48:05.632424Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_seed(seed=12):\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    np.random.seed(seed)\n    torch.backends.cudnn.deterministic = True\nset_seed(12)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-09-05T10:48:07.199301Z","iopub.execute_input":"2022-09-05T10:48:07.200111Z","iopub.status.idle":"2022-09-05T10:48:07.207308Z","shell.execute_reply.started":"2022-09-05T10:48:07.200068Z","shell.execute_reply":"2022-09-05T10:48:07.206107Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"dataset\"><center>Dataset</center></h3>","metadata":{}},{"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        \n    def img2tensor(self, 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))\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        img = cv2.cvtColor(cv2.imread(TRAIN_PATH + fname), cv2.COLOR_BGR2RGB)\n        mask = cv2.imread((MASK_PATH + fname),cv2.IMREAD_GRAYSCALE)\n        if self.tfms is not None:\n            augmented = self.tfms(image=img,mask=mask)\n            img,mask = augmented['image'],augmented['mask']\n        return self.img2tensor(self.resize(img , cv2.INTER_NEAREST)) , self.img2tensor(self.resize(mask , cv2.INTER_NEAREST))","metadata":{"execution":{"iopub.status.busy":"2022-09-05T10:48:08.771907Z","iopub.execute_input":"2022-09-05T10:48:08.772263Z","iopub.status.idle":"2022-09-05T10:48:08.785674Z","shell.execute_reply.started":"2022-09-05T10:48:08.772233Z","shell.execute_reply":"2022-09-05T10:48:08.7841Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"transformer\"><center>Transformer</center></h3>","metadata":{}},{"cell_type":"code","source":"def transformer(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-09-05T10:48:10.343715Z","iopub.execute_input":"2022-09-05T10:48:10.344465Z","iopub.status.idle":"2022-09-05T10:48:10.351567Z","shell.execute_reply.started":"2022-09-05T10:48:10.344426Z","shell.execute_reply":"2022-09-05T10:48:10.35032Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"view\"><center>Viewing Data</center></h3>","metadata":{}},{"cell_type":"code","source":"# ds = HuBMAPDataset(tfms=transformer())\nds = HuBMAPDataset()\ndl = torch.utils.data.DataLoader(ds,batch_size=8,shuffle=False,num_workers=NUM_WORKERS)\nit = iter(dl)\nimgs,masks = next(it)","metadata":{"execution":{"iopub.status.busy":"2022-09-05T10:48:11.917396Z","iopub.execute_input":"2022-09-05T10:48:11.918495Z","iopub.status.idle":"2022-09-05T10:48:13.323318Z","shell.execute_reply.started":"2022-09-05T10:48:11.918447Z","shell.execute_reply":"2022-09-05T10:48:13.321146Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.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-09-05T10:48:16.654622Z","iopub.execute_input":"2022-09-05T10:48:16.656033Z","iopub.status.idle":"2022-09-05T10:48:17.857576Z","shell.execute_reply.started":"2022-09-05T10:48:16.655971Z","shell.execute_reply":"2022-09-05T10:48:17.856598Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"coat\"><center>Importing CoAT</center></h3>","metadata":{}},{"cell_type":"code","source":"sys.path.append('../input/hubmap-coat/')\n\nfrom coat import *\nfrom daformer import *\nfrom helper import *","metadata":{"execution":{"iopub.status.busy":"2022-09-05T10:48:51.90547Z","iopub.execute_input":"2022-09-05T10:48:51.905863Z","iopub.status.idle":"2022-09-05T10:48:51.943633Z","shell.execute_reply.started":"2022-09-05T10:48:51.905814Z","shell.execute_reply":"2022-09-05T10:48:51.942387Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"modify\"><center>Modifying CoAT</center></h3>","metadata":{}},{"cell_type":"code","source":"class Net(nn.Module):\n    \n    def __init__(self,\n                 encoder=coat_lite_medium,\n                 decoder=daformer_conv3x3,\n                 encoder_cfg={},\n                 decoder_cfg={},\n                 ):\n        \n        super(Net, self).__init__()\n        decoder_dim = decoder_cfg.get('decoder_dim', 512)\n\n        self.encoder = encoder\n\n        self.rgb = RGB()\n\n        encoder_dim = self.encoder.embed_dims\n        # [64, 128, 320, 512]\n\n        self.decoder = decoder(\n            encoder_dim=encoder_dim,\n            decoder_dim=decoder_dim,\n        )\n        self.logit = nn.Sequential(\n            nn.Conv2d(decoder_dim, 1, kernel_size=1),\n            nn.Upsample(scale_factor = 4, mode='bilinear', align_corners=False),\n        )\n\n    def forward(self, batch):\n\n        x = self.rgb(batch)\n\n        B, C, H, W = x.shape\n        encoder = self.encoder(x)\n\n        last, decoder = self.decoder(encoder)\n        logit = self.logit(last)\n\n        output = {}\n        probability_from_logit = torch.sigmoid(logit)\n        output['probability'] = probability_from_logit\n\n        return output","metadata":{"execution":{"iopub.status.busy":"2022-09-05T10:49:01.265189Z","iopub.execute_input":"2022-09-05T10:49:01.267482Z","iopub.status.idle":"2022-09-05T10:49:01.284786Z","shell.execute_reply.started":"2022-09-05T10:49:01.267435Z","shell.execute_reply":"2022-09-05T10:49:01.283251Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def init_model():\n    encoder = coat_lite_medium()\n    checkpoint = '../input/coat-medium/coat_lite_medium_a750cd63.pth'\n    checkpoint = torch.load(checkpoint, map_location=lambda storage, loc: storage)\n    state_dict = checkpoint['model']\n    encoder.load_state_dict(state_dict,strict=False)\n    \n    net = Net(encoder=encoder).cuda()\n    \n    return net","metadata":{"execution":{"iopub.status.busy":"2022-09-05T10:49:04.342287Z","iopub.execute_input":"2022-09-05T10:49:04.343209Z","iopub.status.idle":"2022-09-05T10:49:04.350959Z","shell.execute_reply.started":"2022-09-05T10:49:04.343171Z","shell.execute_reply":"2022-09-05T10:49:04.350004Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"loss\"><center>Loss functions</center></h3>","metadata":{}},{"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-09-05T10:49:07.820516Z","iopub.execute_input":"2022-09-05T10:49:07.821175Z","iopub.status.idle":"2022-09-05T10:49:07.827618Z","shell.execute_reply.started":"2022-09-05T10:49:07.821137Z","shell.execute_reply":"2022-09-05T10:49:07.826553Z"},"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-09-05T10:49:10.310439Z","iopub.execute_input":"2022-09-05T10:49:10.31124Z","iopub.status.idle":"2022-09-05T10:49:10.319008Z","shell.execute_reply.started":"2022-09-05T10:49:10.311201Z","shell.execute_reply":"2022-09-05T10:49:10.317969Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"add\"><center>Additionals</center></h3>","metadata":{}},{"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-09-05T10:49:12.831492Z","iopub.execute_input":"2022-09-05T10:49:12.831849Z","iopub.status.idle":"2022-09-05T10:49:12.839577Z","shell.execute_reply.started":"2022-09-05T10:49:12.831819Z","shell.execute_reply":"2022-09-05T10:49:12.838377Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"training\"><center>Training</center></h3>","metadata":{}},{"cell_type":"code","source":"print(f\"Running on device :  {DEVICE}\" )\nfor fold in range(1):\n    \n    val_losses = []\n    losses = []\n    train_scores=[]\n    val_scores = []\n    best_loss = 999\n    best_score = 0\n    \n    ds_train = HuBMAPDataset(fold=fold, train=True, tfms=transformer())\n    ds_val = HuBMAPDataset(fold=fold, train=False)\n    \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    \n    model = init_model().to(DEVICE)\n    \n    optimizer = torch.optim.AdamW([\n        {'params': model.decoder.parameters(), 'lr': 5e-5}, \n        {'params': model.encoder.parameters(), 'lr': 8e-5},  \n    ])\n    \n    scheduler = torch.optim.lr_scheduler.OneCycleLR(optimizer=optimizer, pct_start=0.1, div_factor=1e3, \n                                              max_lr=1e-5, epochs=EPOCHS, steps_per_epoch=len(dataloader_train))\n    \n    loss_func = CustomLoss()\n    dice_coe = DiceCoef()\n    \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        \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)['probability']    \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        \n        with torch.no_grad():\n            \n            valid_loss = 0\n            val_score = 0\n            \n            for data in dataloader_val:\n                \n                img, mask = data\n                img = img.to(DEVICE)\n                mask = mask.to(DEVICE)\n\n                outputs = model(img)['probability']\n\n                loss = loss_func(outputs, mask)\n                valid_loss += loss.item()\n                val_score += dice_coe(outputs,mask).item()\n                \n            valid_loss /= len(dataloader_val)\n            val_losses.append(valid_loss)\n            \n            val_score /= len(dataloader_val)\n            val_scores.append(val_score)\n            \n            print(f\"FOLD: {fold}, EPOCH: {epoch + 1}, valid_loss: {valid_loss} , Val Dice COE : {val_score}\") #\n            \n            gc.collect()\n            torch.cuda.empty_cache()\n            \n        if val_score > best_score:\n            best_score = val_score\n            torch.save(model.state_dict(), f\"/kaggle/working/FOLD{fold}_best_score.pth\")\n            print(f\"Saved model for best score : FOLD{fold}_best_score.pth\")\n            \n        if valid_loss < best_loss:\n            best_loss = valid_loss\n            torch.save(model.state_dict(), f\"/kaggle/working/FOLD{fold}_best_loss.pth\")\n            print(f\"Saved model for best loss : FOLD{fold}_best_loss.pth\")    \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-09-05T10:49:15.426077Z","iopub.execute_input":"2022-09-05T10:49:15.426881Z","iopub.status.idle":"2022-09-05T11:22:11.664012Z","shell.execute_reply.started":"2022-09-05T10:49:15.426843Z","shell.execute_reply":"2022-09-05T11:22:11.662393Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}