{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":61446,"databundleVersionId":6962461,"sourceType":"competition"},{"sourceId":150248402,"sourceType":"kernelVersion"},{"sourceId":154513636,"sourceType":"kernelVersion"}],"dockerImageVersionId":30588,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# First of all\n\nThis notebook is a copy of [this great notebook](https://www.kaggle.com/code/aniketkolte04/sennet-hoa-seg-pytorch-attention-gated-unet), also some parts are taken from [my own previous work](https://github.com/korotaS/grad_project_2021) and [this notebook](https://www.kaggle.com/code/kashiwaba/sennet-hoa-inference-unet-simple-baseline/notebook). ","metadata":{}},{"cell_type":"markdown","source":"# Dependencies","metadata":{}},{"cell_type":"code","source":"# because we don't have acess to the internet\n!pip install --no-index --find-links /kaggle/input/pip-download-for-segmentation-models-pytorch/ segmentation_models_pytorch\n!pip install --no-index --find-links /kaggle/input/pip-download-clearml/ clearml","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2023-12-11T11:55:49.438983Z","iopub.execute_input":"2023-12-11T11:55:49.439615Z","iopub.status.idle":"2023-12-11T11:56:29.447805Z","shell.execute_reply.started":"2023-12-11T11:55:49.439568Z","shell.execute_reply":"2023-12-11T11:56:29.44617Z"},"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nfrom datetime import datetime\nfrom pathlib import Path\n\nimport albumentations as A\nimport albumentations.pytorch as AP\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision\nimport tifffile as tiff\nimport cv2\nimport matplotlib.pyplot as plt\nimport pytorch_lightning as pl\nimport segmentation_models_pytorch as smp\nimport pandas as pd\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nfrom PIL import Image\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nfrom pytorch_lightning.loggers import TensorBoardLogger\nfrom pytorch_lightning.callbacks import LearningRateMonitor, ModelCheckpoint\nfrom kaggle_secrets import UserSecretsClient\nfrom clearml import Task\n\n# first you need to paste your clearml config into kaggle secret\n# then you can use this kind of magic so that clearml init is working\n\n# user_secrets = UserSecretsClient()\n# clearml_config_str = user_secrets.get_secret(\"clearml_config\")\n# with open('/root/clearml.conf', 'w+') as w:\n#     w.write(clearml_config_str.replace('     ', '\\n').replace('} }', '}\\n}'))","metadata":{"execution":{"iopub.status.busy":"2023-12-11T11:56:29.451049Z","iopub.execute_input":"2023-12-11T11:56:29.452325Z","iopub.status.idle":"2023-12-11T11:56:46.103635Z","shell.execute_reply.started":"2023-12-11T11:56:29.452265Z","shell.execute_reply":"2023-12-11T11:56:46.102324Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!nvidia-smi","metadata":{"execution":{"iopub.status.busy":"2023-12-11T11:56:46.105492Z","iopub.execute_input":"2023-12-11T11:56:46.107466Z","iopub.status.idle":"2023-12-11T11:56:47.243162Z","shell.execute_reply.started":"2023-12-11T11:56:46.107412Z","shell.execute_reply":"2023-12-11T11:56:47.241966Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Some looking into the data","metadata":{}},{"cell_type":"code","source":"base_path = '/kaggle/input/blood-vessel-segmentation/train'  \n\ndataset = 'kidney_1_dense'\n\nimages_path = os.path.join(base_path, dataset, 'images')\nlabels_path = os.path.join(base_path, dataset, 'labels')\n\nimage_files = sorted([os.path.join(images_path, f) \n                      for f in os.listdir(images_path) \n                      if f.endswith('.tif')])\nlabel_files = sorted([os.path.join(labels_path, f) \n                      for f in os.listdir(labels_path) \n                      if f.endswith('.tif')])\n    \nfig, axes = plt.subplots(1, 2, figsize=(8, 8))\n\nfirst_image = tiff.imread(image_files[981])\naxes[0].imshow(first_image, cmap='gray')\naxes[0].set_title('First Image')\nfirst_label = tiff.imread(label_files[981])\naxes[1].imshow(first_label, cmap='gray')\naxes[1].set_title('First Label')\n\nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-12-11T11:57:56.024483Z","iopub.execute_input":"2023-12-11T11:57:56.024998Z","iopub.status.idle":"2023-12-11T11:57:56.852334Z","shell.execute_reply.started":"2023-12-11T11:57:56.02496Z","shell.execute_reply":"2023-12-11T11:57:56.851265Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# By the way... Why tifffile? \n\nThe tifffile is faster than opencv and uses less memory. Speed comparison below: ","metadata":{"_kg_hide-input":false}},{"cell_type":"code","source":"%%timeit -r10 -n100\n_ = cv2.imread(image_files[981])","metadata":{"execution":{"iopub.status.busy":"2023-12-11T12:03:57.942802Z","iopub.execute_input":"2023-12-11T12:03:57.943278Z","iopub.status.idle":"2023-12-11T12:04:04.107434Z","shell.execute_reply.started":"2023-12-11T12:03:57.943243Z","shell.execute_reply":"2023-12-11T12:04:04.106002Z"},"_kg_hide-input":false,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%timeit -r10 -n100\n_ = tiff.imread(image_files[981])","metadata":{"execution":{"iopub.status.busy":"2023-12-11T12:04:08.864148Z","iopub.execute_input":"2023-12-11T12:04:08.864584Z","iopub.status.idle":"2023-12-11T12:04:11.2665Z","shell.execute_reply.started":"2023-12-11T12:04:08.864552Z","shell.execute_reply":"2023-12-11T12:04:11.265118Z"},"_kg_hide-input":false,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"It may seem that ~5ms and ~1.5ms are about the same, but when the images are huge and the resources are limited - the choice if obvious. ","metadata":{"_kg_hide-input":false}},{"cell_type":"markdown","source":"# Datasets & dataloaders","metadata":{}},{"cell_type":"code","source":"class CustomDataset(Dataset):\n    def __init__(self, image_files, mask_files, transforms, size=256):\n        self.image_files = image_files\n        self.mask_files = mask_files\n        self.transforms = transforms\n        self.pre_transforms = A.Resize(size, size, interpolation=cv2.INTER_NEAREST)\n        self.post_transforms = AP.ToTensorV2()\n\n    def __len__(self):\n        return len(self.image_files)\n    \n    # simple preprocess - normalize\n    def preprocess_image(self, path):\n        img = tiff.imread(path)\n        img = img.astype('float32') \n        mx = np.max(img)\n        if mx:\n            img /= mx\n        \n        return img\n    \n    # simple preprocess - normalize\n    def preprocess_mask(self, path):\n        mask = tiff.imread(path)\n        mask = mask.astype('float32')\n        mask /= 255.\n\n        return mask\n    \n    def augment_image(self, image, mask):\n        pre_tr = self.pre_transforms(image=image, mask=mask)\n        pre_tr_image, pre_tr_mask = pre_tr['image'], pre_tr['mask']\n        \n        tr = self.transforms(image=pre_tr_image, mask=pre_tr_mask)\n        tr_image, tr_mask = tr['image'], tr['mask']\n\n        post_tr = self.post_transforms(image=tr_image, mask=tr_mask)\n        post_tr_image, post_tr_mask = post_tr['image'], post_tr['mask']\n        post_tr_mask = post_tr_mask.unsqueeze(0)  # HW to CHW\n        \n        # save not only tensor images and masks, but original too - for visualization\n        return pre_tr_image, tr_image, post_tr_image, tr_mask, post_tr_mask\n\n    def __getitem__(self, idx):\n       \n        image_path = self.image_files[idx]\n        mask_path = self.mask_files[idx]\n\n        image = self.preprocess_image(image_path)\n        mask = self.preprocess_mask(mask_path)\n\n        pre_tr_image, tr_image, post_tr_image, tr_mask, post_tr_mask = self.augment_image(image, mask)\n\n        return pre_tr_image, tr_image, post_tr_image, tr_mask, post_tr_mask, idx","metadata":{"execution":{"iopub.status.busy":"2023-12-10T18:59:26.247686Z","iopub.execute_input":"2023-12-10T18:59:26.247967Z","iopub.status.idle":"2023-12-10T18:59:26.259773Z","shell.execute_reply.started":"2023-12-10T18:59:26.247943Z","shell.execute_reply":"2023-12-10T18:59:26.258742Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_image_files, val_image_files, train_mask_files, val_mask_files = train_test_split(\n    image_files, label_files, test_size=0.2, random_state=42\n)","metadata":{"execution":{"iopub.status.busy":"2023-12-10T18:59:26.261082Z","iopub.execute_input":"2023-12-10T18:59:26.261538Z","iopub.status.idle":"2023-12-10T18:59:26.276096Z","shell.execute_reply.started":"2023-12-10T18:59:26.261464Z","shell.execute_reply":"2023-12-10T18:59:26.275162Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transforms = A.Compose([\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n#             A.ShiftScaleRotate(scale_limit=0.5, rotate_limit=0, shift_limit=0.1, p=1, border_mode=0),\n            A.RandomCrop(height=256, width=256, always_apply=True),\n            A.RandomBrightnessContrast(p=1),\n            A.OneOf(\n                [\n                    A.Blur(blur_limit=3, p=1),\n                    A.MotionBlur(blur_limit=3, p=1),\n                ],\n                p=0.9,\n            )\n])\n\ntrain_dataset = CustomDataset(train_image_files, train_mask_files, transforms=transforms, size=256)\nval_dataset = CustomDataset(val_image_files, val_mask_files, transforms=A.NoOp(), size=256)","metadata":{"execution":{"iopub.status.busy":"2023-12-10T18:59:26.277102Z","iopub.execute_input":"2023-12-10T18:59:26.277414Z","iopub.status.idle":"2023-12-10T18:59:26.287775Z","shell.execute_reply.started":"2023-12-10T18:59:26.277391Z","shell.execute_reply":"2023-12-10T18:59:26.286975Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# custom collate function because not all elements in batch are tensors\ndef custom_collate(batch):\n    our_raw_imgs = []\n    out_imgs = []\n    out_imgs_ten = []\n    out_masks = []\n    out_masks_ten = []\n    ids = []\n    \n    for raw_image, image, image_ten, mask, mask_ten, idx in batch:\n        our_raw_imgs.append(raw_image)\n        out_imgs.append(image)\n        out_imgs_ten.append(image_ten)\n        out_masks.append(mask)\n        out_masks_ten.append(mask_ten)\n        ids.append(idx)\n    \n    out_imgs_ten = torch.stack(out_imgs_ten)\n    out_masks_ten = torch.stack(out_masks_ten)\n    \n    return our_raw_imgs, out_imgs, out_imgs_ten, out_masks, out_masks_ten, ids\n\ntrain_dataloader = DataLoader(train_dataset, batch_size=32, shuffle=True, collate_fn=custom_collate, num_workers=3)\nval_dataloader = DataLoader(val_dataset, batch_size=32, shuffle=False, collate_fn=custom_collate, num_workers=3)","metadata":{"execution":{"iopub.status.busy":"2023-12-10T18:59:26.288967Z","iopub.execute_input":"2023-12-10T18:59:26.289693Z","iopub.status.idle":"2023-12-10T18:59:26.301323Z","shell.execute_reply.started":"2023-12-10T18:59:26.289667Z","shell.execute_reply":"2023-12-10T18:59:26.300433Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for batch_idx, (raw_images, images, images_ten, masks, masks_ten, ids) in enumerate(train_dataloader):\n    break","metadata":{"execution":{"iopub.status.busy":"2023-12-10T18:59:26.302344Z","iopub.execute_input":"2023-12-10T18:59:26.302664Z","iopub.status.idle":"2023-12-10T18:59:29.181655Z","shell.execute_reply.started":"2023-12-10T18:59:26.30263Z","shell.execute_reply":"2023-12-10T18:59:29.180589Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visualization for training","metadata":{}},{"cell_type":"code","source":"def draw(raw_images, images, masks, masks_pred, ids, metrics, num=4, figsize=(14, 15), cmap='gray'):\n    columns = 4\n    rows = min(num, len(ids))\n    fig, axs = plt.subplots(rows, columns, figsize=figsize)\n    \n    for i, (raw_image, image, mask, mask_pred, idx, metric) in enumerate(zip(\n        raw_images, images, masks, masks_pred, ids, metrics\n    )):\n        if i == num:\n            break\n        \n        axs[i, 0].set_title(f'{idx} orig image')\n        axs[i, 0].imshow(raw_image, cmap=cmap)\n\n        axs[i, 1].set_title(f'{idx} auged image')\n        axs[i, 1].imshow(image, cmap=cmap)\n\n        axs[i, 2].set_title(f'{idx} true mask')\n        axs[i, 2].imshow(mask, cmap=cmap)\n        \n        mask_pred_numpy = mask_pred.detach().cpu().numpy()[0]\n        axs[i, 3].set_title(f'{idx} pred mask, metric: {metric:.3f}')\n        axs[i, 3].imshow(mask_pred_numpy, cmap=cmap)\n\n    return fig\n    \nfig = draw(raw_images, images, masks, masks_ten, ids, [0] * len(ids))\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-12-10T18:59:29.185375Z","iopub.execute_input":"2023-12-10T18:59:29.185727Z","iopub.status.idle":"2023-12-10T18:59:31.84546Z","shell.execute_reply.started":"2023-12-10T18:59:29.185697Z","shell.execute_reply":"2023-12-10T18:59:31.844553Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# The Model (and trainer)","metadata":{}},{"cell_type":"code","source":"class ImageSegmentationModel(pl.LightningModule):\n    def __init__(self, model, criterion):\n        super().__init__()\n        self.model = model\n        self.criterion = criterion\n        self.metrics = {\n            'dice': dice_coeff\n        }\n\n    def forward(self, x):\n        return self.model(x)\n    \n    def _draw(self, raw_images, images, masks, masks_ten, ids, outputs, prefix):\n        metrics_draw = [self.metrics['dice'](pr, tr) for pr, tr in zip(outputs, masks_ten)]\n        fig = draw(raw_images, images, masks, outputs, ids, metrics_draw, num=4)\n        self.logger.experiment.add_figure(f'{prefix}_visualization', fig, self.current_epoch)\n    \n    def _step(self, batch, batch_idx, prefix):\n        raw_images, images, images_ten, masks, masks_ten, ids = batch\n        batch_size = len(ids)\n        outputs = self.model(images_ten)\n        loss = self.criterion(outputs, masks_ten)\n        tb_logs = {f'{prefix}_loss': loss.cpu()}\n        for metric_name, metric in self.metrics.items():\n            tb_logs[f'{prefix}_' + metric_name] = metric(\n                outputs.cpu().ravel(), masks_ten.cpu().ravel()\n            )\n        for key, value in tb_logs.items():\n            self.log(key, value, batch_size=batch_size)\n            \n        if self.current_epoch % 3 == 0 and batch_idx == 0:\n            self._draw(raw_images, images, masks, masks_ten, ids, outputs, prefix)\n        \n        return loss\n\n    def training_step(self, batch, batch_idx):\n        loss = self._step(batch, batch_idx, 'train')\n        return {'loss': loss.cpu()}\n\n    def validation_step(self, batch, batch_idx):\n        loss = self._step(batch, batch_idx, 'val')\n        return {'loss': loss.cpu()}\n\n    def configure_optimizers(self):\n        # for example\n        optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)\n        # for example\n        scheduler = torch.optim.lr_scheduler.CyclicLR(\n            optimizer, \n            base_lr=1e-4, \n            max_lr=1e-3, \n            step_size_up=1,\n            step_size_down=4,\n            cycle_momentum=False\n        )\n        scheduler_config = {\n            \"scheduler\": scheduler,\n            \"interval\": \"epoch\",\n            \"frequency\": 1,\n            \"monitor\": \"val_loss\",\n        }\n        return [[optimizer], [scheduler_config]]","metadata":{"execution":{"iopub.status.busy":"2023-12-10T18:59:31.846777Z","iopub.execute_input":"2023-12-10T18:59:31.847108Z","iopub.status.idle":"2023-12-10T18:59:31.867234Z","shell.execute_reply.started":"2023-12-10T18:59:31.847072Z","shell.execute_reply":"2023-12-10T18:59:31.866186Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# credits to https://www.kaggle.com/code/aniketkolte04/sennet-hoa-seg-pytorch-attention-gated-unet\nclass FocalLoss(nn.modules.loss._WeightedLoss):\n    def __init__(self, gamma=0, size_average=None, reduce=None, balance_param=1.0):\n        super(FocalLoss, self).__init__(size_average)\n        self.gamma = gamma\n        self.size_average = size_average\n        self.balance_param = balance_param\n\n    def forward(self, input, target):\n        assert len(input.shape) == len(target.shape)\n        assert input.size(0) == target.size(0)\n        assert input.size(1) == target.size(1)\n\n        logpt = - F.binary_cross_entropy_with_logits(input, target)\n        pt = torch.exp(logpt)\n\n        focal_loss = -((1 - pt) ** self.gamma) * logpt\n        balanced_focal_loss = self.balance_param * focal_loss\n        return balanced_focal_loss\n\n# also credits to https://www.kaggle.com/code/aniketkolte04/sennet-hoa-seg-pytorch-attention-gated-unet \ndef dice_coeff(prediction, target, thresh=0.5, epsilon=1e-6):\n\n    mask = torch.zeros_like(prediction)\n    mask[prediction >= thresh] = 1\n\n    inter = torch.sum(mask * target)\n    union = torch.sum(mask) + torch.sum(target)\n    result = torch.mean(2 * inter / (union + epsilon))\n    return result","metadata":{"execution":{"iopub.status.busy":"2023-12-10T18:59:31.868497Z","iopub.execute_input":"2023-12-10T18:59:31.869201Z","iopub.status.idle":"2023-12-10T18:59:31.880926Z","shell.execute_reply.started":"2023-12-10T18:59:31.869172Z","shell.execute_reply":"2023-12-10T18:59:31.879907Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_segm = smp.DeepLabV3(\n    encoder_name='mobilenet_v2', \n    in_channels=1, \n    activation='sigmoid', \n    encoder_weights=None  # without internet\n)\ncriterion = FocalLoss(gamma=2)\nmodel = ImageSegmentationModel(model_segm, criterion)","metadata":{"execution":{"iopub.status.busy":"2023-12-10T18:59:31.882039Z","iopub.execute_input":"2023-12-10T18:59:31.88231Z","iopub.status.idle":"2023-12-10T18:59:32.355776Z","shell.execute_reply.started":"2023-12-10T18:59:31.882287Z","shell.execute_reply":"2023-12-10T18:59:32.354775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"project = 'SenNet + HOA'\nmodel_name = 'smp_deeplabv3_mobilenetv2'\nexp = 'test'\nversion = datetime.now().strftime(\"%Y%m%dT%H%M%S\")\nfilename = 'epoch_{epoch}_val_loss_{val_loss:.4f}_'\nmetric = 'val_dice_{val_dice:.4f}_'\nfilename = filename + metric + version\n\nlr_logger = LearningRateMonitor()\ncheckpoint_callback = ModelCheckpoint(\n    dirpath=f'weights/{project}/{model_name}/{exp}/{version}',\n    filename=filename,\n    auto_insert_metric_name=False,\n    monitor='val_loss',\n    save_top_k=1,\n    save_last=True,\n    verbose=True\n)","metadata":{"execution":{"iopub.status.busy":"2023-12-10T18:59:32.357441Z","iopub.execute_input":"2023-12-10T18:59:32.357832Z","iopub.status.idle":"2023-12-10T18:59:32.366749Z","shell.execute_reply.started":"2023-12-10T18:59:32.357799Z","shell.execute_reply":"2023-12-10T18:59:32.365854Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# if you have internet and have access to clearml - uncomment\n\n# task = Task.init(\n#     project_name=f'{project}/{model_name}',\n#     task_name=exp\n# )\n# task.connect_configuration({'model architecture': '\\n' + repr(model_segm)}, name='model architecture')\n\nlogger = TensorBoardLogger(\n    save_dir=f'logs/{project}',\n    name=model_name,\n    version=f'{exp}/{version}'\n)\npl_trainer = pl.Trainer(\n    logger=logger,\n    log_every_n_steps=25,\n    num_sanity_val_steps=5,\n    gradient_clip_val=0.5,\n    max_epochs=1,\n    callbacks=[lr_logger, checkpoint_callback],\n    accelerator='gpu',\n    devices=[0]\n)\npl_trainer.fit(\n    model=model,\n    train_dataloaders=train_dataloader,\n    val_dataloaders=val_dataloader\n);","metadata":{"execution":{"iopub.status.busy":"2023-12-10T18:59:32.367739Z","iopub.execute_input":"2023-12-10T18:59:32.367992Z","iopub.status.idle":"2023-12-10T19:00:47.807668Z","shell.execute_reply.started":"2023-12-10T18:59:32.367969Z","shell.execute_reply":"2023-12-10T19:00:47.80663Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# task.close()","metadata":{"execution":{"iopub.status.busy":"2023-12-10T19:00:53.496465Z","iopub.execute_input":"2023-12-10T19:00:53.497371Z","iopub.status.idle":"2023-12-10T19:00:53.501242Z","shell.execute_reply.started":"2023-12-10T19:00:53.497339Z","shell.execute_reply":"2023-12-10T19:00:53.500318Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"cell_type":"code","source":"# can be loaded from another notebook's output\ndevice = 'cuda:0'\nmodel.load_state_dict(torch.load(checkpoint_callback.best_model_path, map_location=device)['state_dict'])\nmodel.to(device)\nmodel.eval();","metadata":{"execution":{"iopub.status.busy":"2023-12-10T19:00:54.892563Z","iopub.execute_input":"2023-12-10T19:00:54.892936Z","iopub.status.idle":"2023-12-10T19:00:55.181144Z","shell.execute_reply.started":"2023-12-10T19:00:54.892907Z","shell.execute_reply":"2023-12-10T19:00:55.180314Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def remove_small_objects(img, min_size):\n    # Find all connected components (labels)\n    num_labels, labels, stats, centroids = cv2.connectedComponentsWithStats(img, connectivity=8)\n\n    # Create a mask where small objects are removed\n    new_img = np.zeros_like(img)\n    for label in range(1, num_labels):\n        if stats[label, cv2.CC_STAT_AREA] >= min_size:\n            new_img[labels == label] = 1\n\n    return new_img\n\ndef batchings(arr, batch_size=1):\n    for s in range(0, len(arr), batch_size):\n        yield arr[s:s+batch_size]\n    \ndef rle_encode(img):\n    pixels = img.flatten()\n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    rle = ' '.join(str(x) for x in runs)\n    if rle=='':\n        rle = '1 0'\n    return rle","metadata":{"execution":{"iopub.status.busy":"2023-12-10T19:07:16.641661Z","iopub.execute_input":"2023-12-10T19:07:16.642561Z","iopub.status.idle":"2023-12-10T19:07:16.651081Z","shell.execute_reply.started":"2023-12-10T19:07:16.64253Z","shell.execute_reply":"2023-12-10T19:07:16.650191Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_data_folder = '/kaggle/input/blood-vessel-segmentation/test/'\ntest_image_files = list(Path(test_data_folder).rglob('*.tif'))\nprint('number of files:', len(test_image_files))\n\nbatches = list(batchings(test_image_files, 32))\nprint('number of batches:', len(batches))\nsize = 256\ntest_transforms = A.Compose([\n    A.Resize(size, size, interpolation=cv2.INTER_NEAREST),\n    AP.ToTensorV2()\n])","metadata":{"execution":{"iopub.status.busy":"2023-12-10T19:00:58.537504Z","iopub.execute_input":"2023-12-10T19:00:58.538412Z","iopub.status.idle":"2023-12-10T19:00:58.565814Z","shell.execute_reply.started":"2023-12-10T19:00:58.538378Z","shell.execute_reply":"2023-12-10T19:00:58.564868Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"result = {'id': [], 'rle': []}\n\nfor batch in tqdm(batches):\n    images = [tiff.imread(image_path) for image_path in batch]\n    shapes = [image.shape for image in images]\n    images = [image.astype('float32') / image.max() for image in images]\n    images = [test_transforms(image=image)['image'] for image in images]\n    batch_input = torch.stack(images).to(device)\n    batch_output = model(batch_input)\n    batch_output = (batch_output > 0.5).squeeze(1).detach().cpu().numpy().astype('uint8')\n    for one_output, shape, path in zip(batch_output, shapes, batch):\n        one_output = cv2.resize(one_output, (shape[1], shape[0]), cv2.INTER_NEAREST)\n        one_output = remove_small_objects(one_output, 10)\n        rle = rle_encode(one_output)\n        \n        _, _, _, _, _, dataset, _, part = path.parts\n        part = part.split('.')[0]\n        result['id'].append(f'{dataset}_{part}')\n        result['rle'].append(rle)","metadata":{"execution":{"iopub.status.busy":"2023-12-10T19:09:44.804189Z","iopub.execute_input":"2023-12-10T19:09:44.805119Z","iopub.status.idle":"2023-12-10T19:09:44.957184Z","shell.execute_reply.started":"2023-12-10T19:09:44.805083Z","shell.execute_reply":"2023-12-10T19:09:44.956311Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pd.DataFrame(result).to_csv('submission.csv', index=None)","metadata":{"execution":{"iopub.status.busy":"2023-12-10T19:09:47.08955Z","iopub.execute_input":"2023-12-10T19:09:47.089927Z","iopub.status.idle":"2023-12-10T19:09:47.099433Z","shell.execute_reply.started":"2023-12-10T19:09:47.089899Z","shell.execute_reply":"2023-12-10T19:09:47.098624Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}