{"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":"!cp -r ../input/pytorch-segmentation-models-lib/ ./\n!cp -r ../input/torchmetrics/ ./\n\n!pip config set global.disable-pip-version-check true\n\n!pip install -q ./pytorch-segmentation-models-lib/pretrainedmodels-0.7.4/pretrainedmodels-0.7.4\n!pip install -q ./pytorch-segmentation-models-lib/efficientnet_pytorch-0.6.3/efficientnet_pytorch-0.6.3\n!pip install -q ./pytorch-segmentation-models-lib/timm-0.4.12-py3-none-any.whl\n!pip install -q ./pytorch-segmentation-models-lib/segmentation_models_pytorch-0.2.0-py3-none-any.whl\n!pip install -q ./torchmetrics/torchmetrics-0.9.1-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2022-08-22T13:51:22.396058Z","iopub.execute_input":"2022-08-22T13:51:22.396501Z","iopub.status.idle":"2022-08-22T13:52:20.822688Z","shell.execute_reply.started":"2022-08-22T13:51:22.396469Z","shell.execute_reply":"2022-08-22T13:52:20.821463Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### https://arxiv.org/pdf/2105.15203.pdf","metadata":{}},{"cell_type":"markdown","source":"## Library imports","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\n\nimport os\nimport json\nfrom tqdm.auto import tqdm\nimport gc\n\nfrom skimage import io\nfrom skimage.transform import resize\n\nimport matplotlib.pyplot as plt\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\nfrom albumentations.pytorch.transforms import ToTensorV2\n\nimport pytorch_lightning as pl\nfrom pytorch_lightning import Trainer\n\nimport segmentation_models_pytorch as smp\nfrom torchmetrics.functional import dice\n\nfrom sklearn.model_selection import KFold\nfrom sklearn.model_selection import train_test_split\nfrom operator import itemgetter ","metadata":{"execution":{"iopub.status.busy":"2022-08-22T13:52:20.825404Z","iopub.execute_input":"2022-08-22T13:52:20.825815Z","iopub.status.idle":"2022-08-22T13:52:27.797373Z","shell.execute_reply.started":"2022-08-22T13:52:20.825773Z","shell.execute_reply":"2022-08-22T13:52:27.796340Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Paths definition","metadata":{}},{"cell_type":"code","source":"TRAIN_IMG_PATH = \"../input/hubmap-organ-segmentation/train_images\"\nTEST_IMG_PATH = \"../input/hubmap-organ-segmentation/test_images\"\n\nTRAIN_IMG_ANNOTATIONS = \"../input/hubmap-organ-segmentation/train_annotations\"\nTRAIN_IMG_INFO = \"../input/hubmap-organ-segmentation/train.csv\"","metadata":{"execution":{"iopub.status.busy":"2022-08-22T13:52:27.798767Z","iopub.execute_input":"2022-08-22T13:52:27.799628Z","iopub.status.idle":"2022-08-22T13:52:27.805852Z","shell.execute_reply.started":"2022-08-22T13:52:27.799578Z","shell.execute_reply":"2022-08-22T13:52:27.804403Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Auxiliary functions","metadata":{}},{"cell_type":"code","source":"def rle2mask(mask_rle: str, shape=None, label: int = 0):\n    \"\"\"\n    mask_rle: run-length as string formatted (start length)\n    shape: (height,width) of array to return\n    Returns numpy array, 1 - mask, 0 - background\n\n    \"\"\"\n    rle = np.array(list(map(int, mask_rle.split())))\n    labels = np.zeros(shape).flatten()\n    \n    for start, end in zip(rle[::2], rle[1::2]):\n        labels[start:start+end] = label\n\n    return labels.reshape(shape).T  # Needed to align to RLE direction\n\n\ndef mask_to_rle(mask):\n    #Rescale image to original size\n    size = int(len(mask.flatten())**.5)\n    n = Image.fromarray(mask.reshape((size, size))*255.0)\n    n = np.array(n).astype(np.float32)\n    #Get pixels to flatten\n    pixels = n.T.flatten()\n    #Round the pixels using the half of the range of pixel value\n    pixels = (pixels-min(pixels) > ((max(pixels)-min(pixels))/2)).astype(int)\n    pixels = np.nan_to_num(pixels) #incase of zero-div-error\n    \n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0]\n    runs[1::2] -= runs[::2]\n    \n    return ' '.join(str(x) for x in runs)","metadata":{"execution":{"iopub.status.busy":"2022-08-22T13:52:27.809006Z","iopub.execute_input":"2022-08-22T13:52:27.810353Z","iopub.status.idle":"2022-08-22T13:52:27.821235Z","shell.execute_reply.started":"2022-08-22T13:52:27.810312Z","shell.execute_reply":"2022-08-22T13:52:27.820015Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Reading training dataframe","metadata":{}},{"cell_type":"code","source":"df_train = pd.read_csv(TRAIN_IMG_INFO)","metadata":{"execution":{"iopub.status.busy":"2022-08-22T13:52:27.822584Z","iopub.execute_input":"2022-08-22T13:52:27.823045Z","iopub.status.idle":"2022-08-22T13:52:28.112365Z","shell.execute_reply.started":"2022-08-22T13:52:27.823011Z","shell.execute_reply":"2022-08-22T13:52:28.111406Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.head(2)","metadata":{"execution":{"iopub.status.busy":"2022-08-22T13:52:28.113826Z","iopub.execute_input":"2022-08-22T13:52:28.114199Z","iopub.status.idle":"2022-08-22T13:52:28.138143Z","shell.execute_reply.started":"2022-08-22T13:52:28.114163Z","shell.execute_reply":"2022-08-22T13:52:28.137058Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"organs = df_train['organ'].unique()","metadata":{"execution":{"iopub.status.busy":"2022-08-22T13:52:28.139864Z","iopub.execute_input":"2022-08-22T13:52:28.140292Z","iopub.status.idle":"2022-08-22T13:52:28.152107Z","shell.execute_reply.started":"2022-08-22T13:52:28.140250Z","shell.execute_reply":"2022-08-22T13:52:28.150980Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"organ_annotations = {}\n\nfor i, organ in enumerate(organs):\n    organ_annotations[organ] = i+1\n    \norgan_annotations","metadata":{"execution":{"iopub.status.busy":"2022-08-22T13:52:28.153803Z","iopub.execute_input":"2022-08-22T13:52:28.154298Z","iopub.status.idle":"2022-08-22T13:52:28.163083Z","shell.execute_reply.started":"2022-08-22T13:52:28.154236Z","shell.execute_reply":"2022-08-22T13:52:28.161864Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Network training config","metadata":{}},{"cell_type":"code","source":"class Config:\n    BATCH_SIZE = 16\n    EPOCHS = 1","metadata":{"execution":{"iopub.status.busy":"2022-08-22T13:52:28.164815Z","iopub.execute_input":"2022-08-22T13:52:28.165523Z","iopub.status.idle":"2022-08-22T13:52:28.171687Z","shell.execute_reply.started":"2022-08-22T13:52:28.165484Z","shell.execute_reply":"2022-08-22T13:52:28.170440Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Scheduler","metadata":{}},{"cell_type":"code","source":"class PeakScheduler(torch.optim.lr_scheduler._LRScheduler):\n        def __init__(\n                self, optimizer,\n                epoch_size=-1,\n                lr_start   = 0.000001,\n                lr_max     = 0.000005 * Config.BATCH_SIZE,\n                lr_min     = 0.000001,\n                lr_ramp_ep = 4,\n                lr_sus_ep  = 0,\n                lr_decay   = 0.8,\n                verbose = True\n            ):\n            self.epoch_size = epoch_size\n            self.optimizer= optimizer\n            self.lr_start = lr_start\n            self.lr_max = lr_max\n            self.lr_min = lr_min\n            self.lr_ramp_ep = lr_ramp_ep\n            self.lr_sus_ep = lr_sus_ep\n            self.lr_decay = lr_decay\n            self.is_plotting = True\n            \n            epochs = list(range(Config.EPOCHS))\n            learning_rates = []\n            for i in epochs:\n                self.epoch = i\n                learning_rates.append(self.get_lr())\n            self.is_plotting = False\n            self.epoch = 0\n            plt.scatter(epochs,learning_rates)\n            plt.show()\n            super(PeakScheduler, self).__init__(optimizer, verbose=verbose)\n\n        def get_lr(self):\n            if not self.is_plotting:\n                if self.epoch_size == -1:\n                    self.epoch = self._step_count - 1\n                else:\n                    self.epoch = (self._step_count - 1) / self.epoch_size\n                    \n            if self.epoch < self.lr_ramp_ep:\n                lr = (self.lr_max - self.lr_start) / self.lr_ramp_ep * self.epoch + self.lr_start\n\n            elif self.epoch < self.lr_ramp_ep + self.lr_sus_ep:\n                lr = self.lr_max\n            else:\n                lr = (self.lr_max - self.lr_min) * self.lr_decay**(self.epoch - self.lr_ramp_ep - self.lr_sus_ep) + self.lr_min\n            return [lr for _ in self.optimizer.param_groups]\n","metadata":{"execution":{"iopub.status.busy":"2022-08-22T13:52:28.176190Z","iopub.execute_input":"2022-08-22T13:52:28.176627Z","iopub.status.idle":"2022-08-22T13:52:28.189239Z","shell.execute_reply.started":"2022-08-22T13:52:28.176601Z","shell.execute_reply":"2022-08-22T13:52:28.187155Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Hugging face model","metadata":{}},{"cell_type":"code","source":"from transformers import SegformerForSemanticSegmentation","metadata":{"execution":{"iopub.status.busy":"2022-08-22T13:52:28.191282Z","iopub.execute_input":"2022-08-22T13:52:28.191983Z","iopub.status.idle":"2022-08-22T13:52:28.275192Z","shell.execute_reply.started":"2022-08-22T13:52:28.191950Z","shell.execute_reply":"2022-08-22T13:52:28.274306Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from PIL import Image\n\nimgs = []\n\nfor i, img in enumerate(tqdm(os.listdir(TRAIN_IMG_PATH))):\n    \n    imgs.append(os.path.join(TRAIN_IMG_PATH, img))","metadata":{"execution":{"iopub.status.busy":"2022-08-22T13:52:28.276618Z","iopub.execute_input":"2022-08-22T13:52:28.276977Z","iopub.status.idle":"2022-08-22T13:52:28.373852Z","shell.execute_reply.started":"2022-08-22T13:52:28.276943Z","shell.execute_reply":"2022-08-22T13:52:28.372887Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_size = 512\n\nclass CustomDataset(Dataset):\n    def __init__(self, imgs, df, stage):\n        self.imgs = imgs\n        self.df = df\n        \n        if stage == 'train':\n            self.transforms = A.Compose([\n                A.augmentations.crops.RandomResizedCrop(height=img_size, width=img_size),\n                A.augmentations.Rotate(limit=90, p=0.5),\n                A.augmentations.HorizontalFlip(p=0.5),\n                A.augmentations.VerticalFlip(p=0.5),\n                A.augmentations.transforms.ColorJitter(p=0.5),\n                A.OneOf([\n                    A.OpticalDistortion(p=0.5),\n                    A.GridDistortion(p=.5),\n                    A.PiecewiseAffine(p=0.5),\n                ], p=0.5),\n                A.OneOf([\n                    A.HueSaturationValue(10, 15, 10),\n                    A.CLAHE(clip_limit=4),\n                    A.RandomBrightnessContrast(),            \n                ], p=0.5),                \n                A.Normalize()\n            ])\n        else:\n            self.transforms = A.Compose([\n                A.Resize(img_size, img_size),\n                A.Normalize() \n            ])\n            \n        \n        \n    def __len__(self):\n        return len(self.imgs)\n    \n    def __getitem__(self, idx):\n        img = self.imgs[idx]\n        \n        img_number = int(img.split(\"/\")[-1].split(\".\")[0])\n    \n        rle = self.df[self.df['id'] == img_number]['rle'].values[0]\n        height = self.df[self.df['id'] == img_number]['img_height'].values[0]\n        width = self.df[self.df['id'] == img_number]['img_width'].values[0]\n        organ = self.df[self.df['id'] == img_number]['organ'].values[0]\n        \n        img = np.asarray(Image.open(img))\n        mask = rle2mask(rle, shape=(height, width), label=organ_annotations[organ])\n        \n        transformed = self.transforms(image=img, mask=np.expand_dims(mask, axis=2))\n        \n        return np.transpose(transformed['image'], (2, 0, 1)).astype(np.float32), np.squeeze(transformed['mask'], 2).astype(np.int16)","metadata":{"execution":{"iopub.status.busy":"2022-08-22T13:52:28.375512Z","iopub.execute_input":"2022-08-22T13:52:28.375909Z","iopub.status.idle":"2022-08-22T13:52:28.393959Z","shell.execute_reply.started":"2022-08-22T13:52:28.375872Z","shell.execute_reply":"2022-08-22T13:52:28.392812Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SegmentationModel(pl.LightningModule):\n    def __init__(\n        self,\n        model\n        ):\n        super(SegmentationModel, self).__init__()\n        \n        self.model = model\n        \n        \n    def forward(self, image):\n        outputs = self.model(pixel_values=image)\n        \n        upsampled_logits = nn.functional.interpolate(\n            outputs.logits,\n            size=mask.shape[-2:], \n            mode=\"bilinear\",\n            align_corners=False\n        )\n        \n        return upsampled_logits\n        \n\n    def training_step(self, batch, batch_idx):\n        image, mask = batch[0], batch[1]\n        outputs = self.model(pixel_values=image, labels=mask.long())\n        \n        upsampled_logits = nn.functional.interpolate(\n            outputs.logits,\n            size=mask.shape[-2:], \n            mode=\"bilinear\",\n            align_corners=False\n        )\n        \n        loss = outputs.loss\n        \n        return {'loss': loss, 'logits_mask': upsampled_logits, 'mask': mask}\n    \n    def training_epoch_end(self, outputs):\n        loss = [item['loss'].item() for item in outputs]\n        logits_mask = torch.cat([item['logits_mask'].cpu() for item in outputs]).softmax(1)\n        mask = torch.cat([item['mask'].cpu() for item in outputs])\n        \n        pred_mask = logits_mask.argmax(dim=1).float()\n        \n        dice_score = dice(\n            logits_mask, \n            mask\n        ).item()\n        \n        log_parameters = {\n            \"loss_train\": np.mean(loss),\n            \"dice_score_train\": dice_score\n        }\n        \n        self.log_dict(log_parameters)\n    \n    def validation_step(self, batch, batch_idx):\n        image, mask = batch[0], batch[1]\n        outputs = self.model(pixel_values=image, labels=mask.long())\n        \n        upsampled_logits = nn.functional.interpolate(\n            outputs.logits,\n            size=mask.shape[-2:], \n            mode=\"bilinear\",\n            align_corners=False\n        )\n        \n        loss = outputs.loss\n        \n        return {'loss': loss, 'logits_mask': upsampled_logits, 'mask': mask}\n        \n    def validation_epoch_end(self, outputs):\n        loss = torch.from_numpy(np.array([item['loss'].item() for item in outputs]))\n        logits_mask = torch.cat([item['logits_mask'].cpu() for item in outputs]).softmax(1)\n        mask = torch.cat([item['mask'].cpu() for item in outputs]).long()\n        \n        pred_mask = logits_mask.argmax(dim=1).long()\n        \n        dice_score = dice(\n            logits_mask, \n            mask\n        ).item()\n        \n        log_parameters = {\n            \"loss_valid\": torch.mean(loss),\n            \"dice_score_valid\": dice_score\n        }\n        \n        self.log_dict(log_parameters)        \n\n    def configure_optimizers(self):\n        optimizer = torch.optim.Adam(self.parameters(), lr=1e-3)\n        scheduler = PeakScheduler(\n            optimizer,\n            lr_ramp_ep=int(Config.EPOCHS * 0.3), \n            lr_decay=0.95,\n            lr_max=1e-03,\n            lr_min=1e-08\n        )\n        \n        return {\n            \"optimizer\": optimizer,\n            \"lr_scheduler\": {\"scheduler\": scheduler, \"interval\": \"epoch\", \"frequency\": 1}\n        }","metadata":{"execution":{"iopub.status.busy":"2022-08-22T13:52:28.395695Z","iopub.execute_input":"2022-08-22T13:52:28.396801Z","iopub.status.idle":"2022-08-22T13:52:28.415746Z","shell.execute_reply.started":"2022-08-22T13:52:28.396762Z","shell.execute_reply":"2022-08-22T13:52:28.414614Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"kf = KFold(n_splits=4, shuffle=True, random_state=42)","metadata":{"execution":{"iopub.status.busy":"2022-08-22T13:52:28.419100Z","iopub.execute_input":"2022-08-22T13:52:28.419429Z","iopub.status.idle":"2022-08-22T13:52:28.427257Z","shell.execute_reply.started":"2022-08-22T13:52:28.419404Z","shell.execute_reply":"2022-08-22T13:52:28.426300Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fold = 1\n\nfor train_idx, test_idx in tqdm(kf.split(range(len(imgs)))):\n    train_imgs = itemgetter(*train_idx)(imgs)\n    valid_imgs = itemgetter(*test_idx)(imgs)\n    \n    train_dataset = CustomDataset(train_imgs, df_train, 'train')\n    valid_dataset = CustomDataset(valid_imgs, df_train, 'valid')\n\n    train_dataloader = DataLoader(train_dataset, batch_size=Config.BATCH_SIZE, shuffle=True, drop_last=False)\n    valid_dataloader = DataLoader(valid_dataset, batch_size=Config.BATCH_SIZE, shuffle=False, drop_last=False)\n    \n    net = SegformerForSemanticSegmentation.from_pretrained(\n        \"nvidia/mit-b0\",\n        ignore_mismatched_sizes=True, \n        num_labels=len(organ_annotations)+1, \n        reshape_last_stage=True\n    )\n    \n    trainer = Trainer(\n        max_epochs=Config.EPOCHS,\n        accelerator=\"cuda\",\n        gpus=1,\n        precision=32, \n        auto_lr_find=True,\n        accumulate_grad_batches=8,\n        auto_scale_batch_size=True,\n        check_val_every_n_epoch=5\n    )\n\n    model = SegmentationModel(net)\n\n    trainer.fit(\n        model,\n        train_dataloader,\n        valid_dataloader,\n\n    )\n    \n    MODEL_NAME = f\"Segformer-fold-{fold}-epochs-{Config.EPOCHS}\"\n\n    if not os.path.exists(os.path.join(\"./\", MODEL_NAME)):\n        os.mkdir(os.path.join(\"./\", MODEL_NAME))\n\n    model.model.save_pretrained(f\"./{MODEL_NAME}\")  \n    \n    fold += 1\n    \n    break","metadata":{"execution":{"iopub.status.busy":"2022-08-22T13:52:28.428794Z","iopub.execute_input":"2022-08-22T13:52:28.429248Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### You may try to create an inference notebook by yourself or wait a little bit. I will add it to the public! Even with test-time augmentations ;)","metadata":{}},{"cell_type":"markdown","source":"### Segformer model with mit-b5 backbone, 3 folds gave me 0.70 Dice score on a public LB.","metadata":{}},{"cell_type":"markdown","source":"### If you find this notebook useful, please upvote it.","metadata":{}},{"cell_type":"markdown","source":"### If you have any questions, feel free to start discussions in comments =)","metadata":{}},{"cell_type":"markdown","source":"#### >>> Work is in progress. Notebook will be updated!","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}