{"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":"#### This notebook continues my previous work.\n#### If you trained model with https://www.kaggle.com/code/alexeyolkhovikov/segformer-training, you can use this notebook for inference.\n#### If you find this notebook useful, please upvote!","metadata":{"_uuid":"cc0d1647-767c-4e90-83da-a839f1006957","_cell_guid":"97158408-2db2-4107-bc9d-56da632655d3","jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-08-23T14:48:21.566796Z","iopub.execute_input":"2022-08-23T14:48:21.567427Z","iopub.status.idle":"2022-08-23T14:49:15.197698Z","shell.execute_reply.started":"2022-08-23T14:48:21.567393Z","shell.execute_reply":"2022-08-23T14:49:15.196469Z"}}},{"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-24T15:13:21.267249Z","iopub.execute_input":"2022-08-24T15:13:21.267643Z","iopub.status.idle":"2022-08-24T15:14:14.972679Z","shell.execute_reply.started":"2022-08-24T15:13:21.267607Z","shell.execute_reply":"2022-08-24T15:14:14.971304Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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\nfrom PIL import Image\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\nfrom transformers import SegformerFeatureExtractor, SegformerForSemanticSegmentation\n\nimport pytorch_lightning as pl\nfrom pytorch_lightning import Trainer\nimport segmentation_models_pytorch as smp\nfrom torchmetrics import Dice","metadata":{"execution":{"iopub.status.busy":"2022-08-24T15:14:14.975403Z","iopub.execute_input":"2022-08-24T15:14:14.976037Z","iopub.status.idle":"2022-08-24T15:14:20.887275Z","shell.execute_reply.started":"2022-08-24T15:14:14.975994Z","shell.execute_reply":"2022-08-24T15:14:20.886086Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TEST_IMG_PATH = \"../input/hubmap-organ-segmentation/test_images/\"\n\nMODEL_PATH = [\n    \"../input/deeplabv3plus/\"\n]","metadata":{"execution":{"iopub.status.busy":"2022-08-24T15:14:20.889485Z","iopub.execute_input":"2022-08-24T15:14:20.890500Z","iopub.status.idle":"2022-08-24T15:14:20.895629Z","shell.execute_reply.started":"2022-08-24T15:14:20.890457Z","shell.execute_reply":"2022-08-24T15:14:20.894612Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def 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-24T15:14:39.826549Z","iopub.execute_input":"2022-08-24T15:14:39.827532Z","iopub.status.idle":"2022-08-24T15:14:39.836507Z","shell.execute_reply.started":"2022-08-24T15:14:39.827488Z","shell.execute_reply":"2022-08-24T15:14:39.835610Z"},"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        self.criterion = smp.losses.DiceLoss(\n            smp.losses.BINARY_MODE,\n            from_logits=True,\n            log_loss=True\n        )\n\n        self.metric = Dice(\n            num_classes=1,\n            average='none'\n        )\n        \n    def forward(self, image, size):\n        outputs = self.model(image)\n        \n        upsampled_logits = nn.functional.interpolate(\n            outputs.logits,\n            size=size, \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'] for item in outputs]).sigmoid()\n        mask = torch.cat([item['mask'] for item in outputs])\n        \n        pred_mask = logits_mask.argmax(dim=1).float()\n        \n        tp, fp, fn, tn = smp.metrics.get_stats(pred_mask.long(), mask.long(), mode=\"binary\")\n        per_image_iou = smp.metrics.iou_score(tp, fp, fn, tn, reduction=\"micro-imagewise\")\n        dataset_iou = smp.metrics.iou_score(tp, fp, fn, tn, reduction=\"micro\")\n        \n        log_parameters = {\n            \"loss_train\": np.mean(loss),\n            \"per_image_iou_train\": per_image_iou,\n            \"dataset_iou_train\": dataset_iou,\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'] for item in outputs]).sigmoid()\n        mask = torch.cat([item['mask'] for item in outputs])\n        \n        pred_mask = logits_mask.argmax(dim=1).float()\n        \n        tp, fp, fn, tn = smp.metrics.get_stats(pred_mask.long(), mask.long(), mode=\"binary\")\n        per_image_iou = smp.metrics.iou_score(tp, fp, fn, tn, reduction=\"micro-imagewise\")\n        dataset_iou = smp.metrics.iou_score(tp, fp, fn, tn, reduction=\"micro\")\n        \n        log_parameters = {\n            \"loss_valid\": torch.mean(loss),\n            \"per_image_iou_valid\": per_image_iou,\n            \"dataset_iou_valid\": dataset_iou,\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.25), \n            lr_decay=0.95,\n            lr_max=1e-03,\n            lr_min=1e-06\n        )\n        \n        return {\n            \"optimizer\": optimizer,\n            \"lr_scheduler\": {\"scheduler\": scheduler, \"interval\": \"epoch\", \"frequency\": 1}\n        }","metadata":{"execution":{"iopub.status.busy":"2022-08-24T15:14:42.418358Z","iopub.execute_input":"2022-08-24T15:14:42.418729Z","iopub.status.idle":"2022-08-24T15:14:42.437634Z","shell.execute_reply.started":"2022-08-24T15:14:42.418698Z","shell.execute_reply":"2022-08-24T15:14:42.436282Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"models = []\n\nfor p in MODEL_PATH:\n    net = SegformerForSemanticSegmentation.from_pretrained(p)\n    model = SegmentationModel(net)\n    model = model.to(\"cuda\")\n    models.append(model)\n    ","metadata":{"execution":{"iopub.status.busy":"2022-08-24T15:14:44.646089Z","iopub.execute_input":"2022-08-24T15:14:44.646509Z","iopub.status.idle":"2022-08-24T15:14:52.431530Z","shell.execute_reply.started":"2022-08-24T15:14:44.646474Z","shell.execute_reply":"2022-08-24T15:14:52.430505Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"IMG_SIZE = 768\n\nclass CustomDatasetTest(Dataset):\n    def __init__(\n        self, \n        paths: str = None,\n        img_size: int = None\n        ):\n        \n        self.paths = paths\n        \n        self.transform = A.Compose([\n            A.Resize(img_size, img_size),\n            A.Normalize()\n        ])\n\n        self.to_tensor = ToTensorV2()\n        self.img_size = img_size\n    \n    def __len__(self):\n        return len(self.paths)\n    \n    def __getitem__(\n        self, \n        idx: int = None\n    ):\n        img = np.asarray(Image.open(self.paths[idx]))\n        \n        H, W = img.shape[:2]\n        \n        transformed = self.transform(image=img)\n        \n        return {\n            'image': self.to_tensor(image=transformed['image'].copy())['image'],\n#             'image_small': self.to_tensor(image=A.Normalize()(image=A.Resize(512, 512)(image=img)['image'])['image'])['image'],\n#             'image_big': self.to_tensor(image=A.Normalize()(image=A.Resize(1024, 1024)(image=img)['image'])['image'])['image'],\n#             'image_blurred': self.to_tensor(image=A.Blur(p=1.)(image=transformed['image'].copy())['image'])['image'],\n#             'image_jittered': self.to_tensor(image=A.ColorJitter(p=1.)(image=transformed['image'].copy())['image'])['image'],\n            'image_num': self.paths[idx].split(\"/\")[-1].split(\".\")[0],\n            'initial_size': (H, W)\n        }","metadata":{"execution":{"iopub.status.busy":"2022-08-24T15:15:00.285838Z","iopub.execute_input":"2022-08-24T15:15:00.286199Z","iopub.status.idle":"2022-08-24T15:15:00.295632Z","shell.execute_reply.started":"2022-08-24T15:15:00.286169Z","shell.execute_reply":"2022-08-24T15:15:00.294611Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"paths = [os.path.join(TEST_IMG_PATH, img) for img in os.listdir(TEST_IMG_PATH)]\ndataset_test = CustomDatasetTest(paths, IMG_SIZE)\ndataloader_test = DataLoader(dataset_test, batch_size=1, drop_last=False)","metadata":{"execution":{"iopub.status.busy":"2022-08-24T15:15:02.100985Z","iopub.execute_input":"2022-08-24T15:15:02.101655Z","iopub.status.idle":"2022-08-24T15:15:02.116766Z","shell.execute_reply.started":"2022-08-24T15:15:02.101616Z","shell.execute_reply":"2022-08-24T15:15:02.115897Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rotate(img, k, model):\n    logits_rot = model.forward(\n        torch.rot90(img, k=k, dims=[2, 3]).to(\"cuda\", dtype=torch.float32),\n        size=batch['initial_size']\n    ).to(\"cpu\") \n            \n    logits_rot = torch.rot90(logits_rot, k=k, dims=[3, 2])\n    \n    return logits_rot  \n\ndef flip_vertical(img, model):\n    logits_flipped = model.forward(\n        torch.flip(img, dims=[3]).to(\"cuda\", dtype=torch.float32),\n        size=batch['initial_size']\n    ).to(\"cpu\") \n    \n    logits_flipped = torch.flip(logits_flipped, dims=[3])\n\n    return logits_flipped\n\ndef flip_horizontal(img, model):\n    logits_flipped = model.forward(\n        torch.flip(img, dims=[1, 2]).to(\"cuda\", dtype=torch.float32),\n        size=batch['initial_size']\n    ).to(\"cpu\") \n    \n    logits_flipped = torch.flip(logits_flipped, dims=[1, 2])\n\n    return logits_flipped","metadata":{"execution":{"iopub.status.busy":"2022-08-24T15:15:03.608138Z","iopub.execute_input":"2022-08-24T15:15:03.608718Z","iopub.status.idle":"2022-08-24T15:15:03.620147Z","shell.execute_reply.started":"2022-08-24T15:15:03.608679Z","shell.execute_reply":"2022-08-24T15:15:03.618973Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ids = []\nrle = []\n\nfor _, batch in enumerate(tqdm(dataloader_test)):\n    preds = []\n    with torch.no_grad():\n        for model in models:\n            model.eval()\n            \n            img = batch['image']\n            img_initial = batch['image']\n            \n            logits_image = model.forward(\n                img.to(\"cuda\", dtype=torch.float32),\n                size=batch['initial_size']\n            ).to(\"cpu\")\n            \n#             logits_rot_90 = rotate(img, 1, model)\n#             logits_rot_180 = rotate(img, 2, model)\n#             logits_rots_270 = rotate(img, 3, model)\n            \n            logits_flipped_H = flip_horizontal(img, model)\n            logits_flipped_V = flip_vertical(img, model)\n\n        \n            preds.extend([logits_image, logits_flipped_H, logits_flipped_V]) #, logits_rot_90, logits_rot_180, logits_rots_270\n        \n    preds = torch.mean(torch.cat(preds, dim=0), dim=0).unsqueeze(0)\n    \n    pr_masks = preds.softmax(1).argmax(1)\n    pr_masks = torch.where(pr_masks==5, 0, pr_masks)\n    pr_masks = torch.where(pr_masks!=0, 1, pr_masks)\n        \n    ids.append(batch['image_num'][0])\n    rle.append(mask_to_rle(pr_masks.numpy()))","metadata":{"execution":{"iopub.status.busy":"2022-08-24T15:21:21.122500Z","iopub.execute_input":"2022-08-24T15:21:21.123209Z","iopub.status.idle":"2022-08-24T15:21:23.202363Z","shell.execute_reply.started":"2022-08-24T15:21:21.123172Z","shell.execute_reply":"2022-08-24T15:21:23.201330Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = {\n    \"id\": ids,\n    \"rle\": rle\n}","metadata":{"execution":{"iopub.status.busy":"2022-08-24T15:15:39.881976Z","iopub.execute_input":"2022-08-24T15:15:39.882334Z","iopub.status.idle":"2022-08-24T15:15:39.887385Z","shell.execute_reply.started":"2022-08-24T15:15:39.882304Z","shell.execute_reply":"2022-08-24T15:15:39.886092Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = pd.DataFrame.from_dict(submission)\nsubmission.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2022-08-24T15:15:40.895820Z","iopub.execute_input":"2022-08-24T15:15:40.896510Z","iopub.status.idle":"2022-08-24T15:15:40.909989Z","shell.execute_reply.started":"2022-08-24T15:15:40.896474Z","shell.execute_reply":"2022-08-24T15:15:40.908970Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission","metadata":{"execution":{"iopub.status.busy":"2022-08-24T15:15:41.093855Z","iopub.execute_input":"2022-08-24T15:15:41.094175Z","iopub.status.idle":"2022-08-24T15:15:41.109010Z","shell.execute_reply.started":"2022-08-24T15:15:41.094148Z","shell.execute_reply":"2022-08-24T15:15:41.107911Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}