{"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":"# FTUs⚕️Segm: EDA🔎 & baseline Lightning⚡Flash on tiled images\n\nThis is derived from Flash docs and paralele competition: https://www.kaggle.com/code/jirkaborovec/tract-segm-eda-flash-deeplab-albumentation","metadata":{}},{"cell_type":"code","source":"!pip uninstall -y torchtext\n# !pip install -q --upgrade torch torchvision\n!mkdir -p frozen_packages\n!cp ../input/starter-flash-semantic-segmentation/frozen_packages/* frozen_packages/\n!cp ../input/ftus-segm-eda-viewer/frozen_packages/* frozen_packages/\n!pip install -q \"lightning-flash[image]\" \"torchmetrics<0.8\" --no-index --find-links frozen_packages/\n!pip install -q -U timm segmentation-models-pytorch --no-index --find-links frozen_packages/\n!pip install -q 'kaggle-image-segmentation' --no-index --find-links frozen_packages/\n\n! pip list | grep -e torch -e lightning\n! nvidia-smi -L","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-08-13T23:36:55.571162Z","iopub.execute_input":"2022-08-13T23:36:55.571747Z","iopub.status.idle":"2022-08-13T23:38:33.831255Z","shell.execute_reply.started":"2022-08-13T23:36:55.571650Z","shell.execute_reply":"2022-08-13T23:38:33.830115Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Loading dataset\n\nIn this case we are using generated segmentation mask exported in this dataset: https://www.kaggle.com/datasets/jirkaborovec/hacking-the-human-body-annotation-masks\n\nand generated from following EDA kernel: https://www.kaggle.com/code/jirkaborovec/ftus-segm-eda-export-rle-mask","metadata":{}},{"cell_type":"code","source":"import os, glob\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nDATASET_FOLDER = \"/kaggle/input/hubmap-organ-segmentation\"\nDATASET_HPA = \"/kaggle/input/hacking-the-human-body-masks-png-images\"\nDATASET_HBMAP = \"/kaggle/input/hacking-the-kidney-annotation-masks-png-images\"\nDATASET_TILES = \"/kaggle/input/ftus-segm-decompose-large-images-tiles\"\npath_csv = os.path.join(DATASET_FOLDER, \"train.csv\")\ndf_train = pd.read_csv(path_csv)\ndisplay(df_train.head())","metadata":{"execution":{"iopub.status.busy":"2022-08-13T23:38:33.834545Z","iopub.execute_input":"2022-08-13T23:38:33.835547Z","iopub.status.idle":"2022-08-13T23:38:34.192887Z","shell.execute_reply.started":"2022-08-13T23:38:33.835506Z","shell.execute_reply":"2022-08-13T23:38:34.191770Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test = pd.read_csv(os.path.join(DATASET_FOLDER, \"test.csv\"))\n\ndisplay(df_test.head())","metadata":{"execution":{"iopub.status.busy":"2022-08-13T23:38:34.194332Z","iopub.execute_input":"2022-08-13T23:38:34.194663Z","iopub.status.idle":"2022-08-13T23:38:34.212568Z","shell.execute_reply.started":"2022-08-13T23:38:34.194628Z","shell.execute_reply":"2022-08-13T23:38:34.211741Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ls = glob.glob(os.path.join(DATASET_FOLDER, 'test_images', '*'))\nWITH_SUBMISSION = len(ls) > 1\n\nfor fname in ls[:2]:\n    plt.imshow(plt.imread(fname))","metadata":{"execution":{"iopub.status.busy":"2022-08-13T23:38:34.215257Z","iopub.execute_input":"2022-08-13T23:38:34.215694Z","iopub.status.idle":"2022-08-13T23:38:35.163580Z","shell.execute_reply.started":"2022-08-13T23:38:34.215658Z","shell.execute_reply":"2022-08-13T23:38:35.162656Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Make a grid/tiles\n\nfor more details and sourced tiles with annotation see:\n>https://www.kaggle.com/code/jirkaborovec/ftus-segm-decompose-large-images-tiles","metadata":{}},{"cell_type":"code","source":"import numpy as np\nfrom PIL import Image\n\n!mkdir -p /kaggle/temp/images\n!mkdir -p /kaggle/temp/masks\n\ndef extract_image_tiles(p_img, folder, size: int = 768) -> list:\n    im = np.array(Image.open(p_img))\n    # https://stackoverflow.com/a/47581978/4521646\n    w = h = size\n    idxs = [(i, (i + h), j, (j + w)) for i in range(0, im.shape[0], h) for j in range(0, im.shape[1], w)]\n    name, _ = os.path.splitext(os.path.basename(p_img))\n    files = []\n    for k, (i, i2, j, j2) in enumerate(idxs):\n        tile = im[i:i2, j:j2, ...]\n        if tile.shape[:2] != (h, w):\n            tile_ = tile\n            tile_size = (h, w) if tile.ndim == 2 else (h, w, tile.shape[2])\n            tile = np.zeros(tile_size, dtype=tile.dtype)\n            tile[:tile_.shape[0], :tile_.shape[1], ...] = tile_\n        p_img = os.path.join(folder, f\"{name}_{k:02}.png\")\n        Image.fromarray(tile).save(p_img)\n        files.append(p_img)\n    return files, idxs\n\n\ntiles_img, _ = extract_image_tiles(\n    os.path.join(DATASET_FOLDER, \"train_images\", \"12233.tiff\"),\n    \"/kaggle/temp/images\", size=1024\n)\ntiles_seg, idxs = extract_image_tiles(\n    os.path.join(DATASET_HPA, \"train_binary_masks\", \"12233.png\"),\n    \"/kaggle/temp/masks\", size=1024\n)\n\n!ls -lh /kaggle/temp/images\n!ls -lh /kaggle/temp/masks","metadata":{"_kg_hide-output":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-08-13T23:40:21.170825Z","iopub.execute_input":"2022-08-13T23:40:21.171669Z","iopub.status.idle":"2022-08-13T23:40:29.312387Z","shell.execute_reply.started":"2022-08-13T23:40:21.171628Z","shell.execute_reply":"2022-08-13T23:40:29.311210Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Lightning⚡Flash & Unet++\n\nlets follow the Semantinc segmentation example: https://lightning-flash.readthedocs.io/en/stable/reference/semantic_segmentation.html","metadata":{}},{"cell_type":"code","source":"import json\nimport torch\n\nimport flash\nimport numpy as np\nfrom flash.core.data.utils import download_data\nfrom flash.image import SemanticSegmentation, SemanticSegmentationData\n\n\nwith open(os.path.join(DATASET_HPA, 'labels.json')) as fp:\n    LABELS = json.load(fp)\nTILE_SIZE = 1024","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-08-13T23:43:10.320784Z","iopub.execute_input":"2022-08-13T23:43:10.321615Z","iopub.status.idle":"2022-08-13T23:43:10.334048Z","shell.execute_reply.started":"2022-08-13T23:43:10.321555Z","shell.execute_reply":"2022-08-13T23:43:10.333055Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 1. Create the DataModule","metadata":{}},{"cell_type":"code","source":"from dataclasses import dataclass\nfrom typing import Any, Callable, Dict, Mapping, Sequence, Tuple, Union\nimport albumentations as alb\n\nfrom flash.core.data.io.input_transform import InputTransform\nfrom flash.image.segmentation.input_transform import prepare_target, remove_extra_dimensions\nfrom kaggle_imsegm.transform import FlashAlbumentationsAdapter\n\n@dataclass\nclass SemanticSegmentationInputTransform(InputTransform):\n    # https://albumentations.ai/docs/examples/pytorch_semantic_segmentation\n\n    image_size: Tuple[int, int] = (224, 224)\n\n    def train_per_sample_transform(self) -> Callable:\n        return FlashAlbumentationsAdapter([\n            alb.Resize(*self.image_size),\n            alb.VerticalFlip(p=0.5),\n            alb.HorizontalFlip(p=0.5),\n            alb.RandomRotate90(p=0.5),\n            alb.ShiftScaleRotate(shift_limit=0.15, scale_limit=0.05, rotate_limit=5, p=1.),\n            alb.GaussNoise(var_limit=(0.001, 0.01), mean=0, per_channel=False, p=1.0),\n            # alb.OneOf([\n            #     alb.GridDistortion(num_steps=5, distort_limit=0.05, p=1.0),\n            #     alb.ElasticTransform(alpha=1, sigma=50, alpha_affine=50, p=1.0),\n            # ], p=0.25),\n            alb.RandomBrightnessContrast(brightness_limit=0.2, contrast_limit=0.2, p=0.8),\n        ])\n\n    def per_sample_transform(self) -> Callable:\n        return FlashAlbumentationsAdapter([alb.Resize(*self.image_size)])\n\n    def target_per_batch_transform(self) -> Callable:\n        return prepare_target\n\n    def predict_per_batch_transform(self) -> Callable:\n        return remove_extra_dimensions\n\n    def serve_per_batch_transform(self) -> Callable:\n        return remove_extra_dimensions","metadata":{"execution":{"iopub.status.busy":"2022-08-13T23:43:10.358112Z","iopub.execute_input":"2022-08-13T23:43:10.358378Z","iopub.status.idle":"2022-08-13T23:43:10.370052Z","shell.execute_reply.started":"2022-08-13T23:43:10.358354Z","shell.execute_reply":"2022-08-13T23:43:10.369048Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"IMAGE_SIZE = (384, 384)\nnb_labels = len(LABELS) + 1\n\ndatamodule = SemanticSegmentationData.from_folders(\n    train_folder=os.path.join(DATASET_TILES, \"train_images\"),\n    train_target_folder=os.path.join(DATASET_TILES, \"train_masks\"),\n    val_split=0.01 if WITH_SUBMISSION else 0.2,\n    predict_folder=os.path.join(DATASET_FOLDER, 'test_images'),\n    train_transform=SemanticSegmentationInputTransform,\n    val_transform=SemanticSegmentationInputTransform,\n    predict_transform=SemanticSegmentationInputTransform,\n    transform_kwargs=dict(image_size=IMAGE_SIZE),\n    num_classes=nb_labels,\n    batch_size=9,\n    num_workers=2,\n)","metadata":{"execution":{"iopub.status.busy":"2022-08-13T23:45:09.140104Z","iopub.execute_input":"2022-08-13T23:45:09.140470Z","iopub.status.idle":"2022-08-13T23:45:09.212898Z","shell.execute_reply.started":"2022-08-13T23:45:09.140440Z","shell.execute_reply":"2022-08-13T23:45:09.211797Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for batch in datamodule.train_dataloader():\n    nb_samples = min(5, len(batch['input']))\n    fig, axarr = plt.subplots(ncols=2, nrows=nb_samples, figsize=(8, 4 * nb_samples))\n    for i in range(nb_samples):\n        segm = batch['target'][i].numpy()\n        img = np.rollaxis(batch['input'][i].cpu().numpy(), 0, 3)\n        axarr[i, 0].imshow(img)\n        seg = axarr[i, 1].imshow(segm, vmin=0, vmax=nb_labels)\n        plt.colorbar(seg, ax=axarr[i, 1])\n    break","metadata":{"execution":{"iopub.status.busy":"2022-08-13T23:45:09.216392Z","iopub.execute_input":"2022-08-13T23:45:09.216667Z","iopub.status.idle":"2022-08-13T23:45:14.044957Z","shell.execute_reply.started":"2022-08-13T23:45:09.216642Z","shell.execute_reply":"2022-08-13T23:45:14.043531Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 2. Build the task","metadata":{}},{"cell_type":"code","source":"from pprint import pprint\n\npprint(SemanticSegmentation.available_heads())\npprint(SemanticSegmentation.available_backbones()['unetplusplus'])","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-08-13T23:45:14.047517Z","iopub.execute_input":"2022-08-13T23:45:14.047895Z","iopub.status.idle":"2022-08-13T23:45:14.056381Z","shell.execute_reply.started":"2022-08-13T23:45:14.047860Z","shell.execute_reply":"2022-08-13T23:45:14.055226Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import segmentation_models_pytorch as smp\n\nmodel = SemanticSegmentation(\n    backbone=\"resnext50_32x4d\",\n    head=\"unetplusplus\",\n    pretrained=False,\n    optimizer=\"Adamax\",\n    learning_rate=0.01,\n    lr_scheduler=(\"StepLR\", {\"step_size\": 2500}),\n    loss_fn=smp.losses.LovaszLoss(mode='multiclass'),\n    num_classes=datamodule.num_classes,\n)","metadata":{"execution":{"iopub.status.busy":"2022-08-13T23:45:14.057859Z","iopub.execute_input":"2022-08-13T23:45:14.059213Z","iopub.status.idle":"2022-08-13T23:45:14.946865Z","shell.execute_reply.started":"2022-08-13T23:45:14.059174Z","shell.execute_reply":"2022-08-13T23:45:14.945878Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 3. Create the trainer and finetune the model","metadata":{}},{"cell_type":"code","source":"import pytorch_lightning as pl\n\ntrainer = flash.Trainer(\n    max_epochs=20 if WITH_SUBMISSION else 10,\n    logger=pl.loggers.CSVLogger(save_dir='logs/'),\n    gpus=torch.cuda.device_count(),\n    precision=16 if torch.cuda.is_available() else 32,\n    accumulate_grad_batches=8,\n    gradient_clip_val=0.01,\n    limit_train_batches=1.0 if WITH_SUBMISSION else 0.5,\n    limit_val_batches=1.0 if WITH_SUBMISSION else 0.5,\n)","metadata":{"execution":{"iopub.status.busy":"2022-08-13T23:45:14.949673Z","iopub.execute_input":"2022-08-13T23:45:14.950075Z","iopub.status.idle":"2022-08-13T23:45:14.961299Z","shell.execute_reply.started":"2022-08-13T23:45:14.950039Z","shell.execute_reply":"2022-08-13T23:45:14.960263Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import warnings\nwarnings.filterwarnings(\"ignore\", category=DeprecationWarning) \n\n# Train the model\ntrainer.finetune(model, datamodule=datamodule, strategy=\"no_freeze\")\n\n# Save the model!\ntrainer.save_checkpoint(\"semantic_segmentation_model.pt\")","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-08-13T23:45:14.963144Z","iopub.execute_input":"2022-08-13T23:45:14.963917Z","iopub.status.idle":"2022-08-13T23:49:25.577485Z","shell.execute_reply.started":"2022-08-13T23:45:14.963878Z","shell.execute_reply":"2022-08-13T23:49:25.576341Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Show training progress","metadata":{}},{"cell_type":"code","source":"import seaborn as sn\n\nmetrics = pd.read_csv(f'{trainer.logger.log_dir}/metrics.csv')\ndel metrics[\"step\"]\nmetrics.set_index(\"epoch\", inplace=True)\n# display(metrics.dropna(axis=1, how=\"all\").head())\ng = sn.relplot(data=metrics, kind=\"line\")\nplt.gcf().set_size_inches(12, 4)\nplt.gca().set_yscale('log')\nplt.grid()","metadata":{"execution":{"iopub.status.busy":"2022-08-13T23:49:25.581538Z","iopub.execute_input":"2022-08-13T23:49:25.581876Z","iopub.status.idle":"2022-08-13T23:49:26.641217Z","shell.execute_reply.started":"2022-08-13T23:49:25.581824Z","shell.execute_reply":"2022-08-13T23:49:26.640138Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 4. Segment a few images!","metadata":{}},{"cell_type":"code","source":"sample_imgs = tiles_img[:5]\n\ndm = SemanticSegmentationData.from_files(\n    predict_files=sample_imgs,\n    predict_transform=SemanticSegmentationInputTransform,\n    transform_kwargs=dict(image_size=IMAGE_SIZE),\n    batch_size=3,\n)","metadata":{"execution":{"iopub.status.busy":"2022-08-13T23:49:26.643002Z","iopub.execute_input":"2022-08-13T23:49:26.643410Z","iopub.status.idle":"2022-08-13T23:49:26.654069Z","shell.execute_reply.started":"2022-08-13T23:49:26.643373Z","shell.execute_reply":"2022-08-13T23:49:26.652992Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from itertools import chain\n\nnrows = max(2, len(sample_imgs))\nfig, axarr = plt.subplots(nrows=nrows, ncols=nb_labels + 1, figsize=(15, 3 * nrows))\n\npreds = trainer.predict(model, datamodule=dm)\npreds = list(chain(*preds))\nfor i, pred in enumerate(preds):\n    # print(pred.keys())\n    img = np.rollaxis(pred['input'].cpu().numpy(), 0, 3)\n    print(img.dtype, img.min(), img.max())\n    axarr[i, 0].imshow(img)\n    for j, seg in enumerate(pred['preds'].cpu().numpy()):\n        p = axarr[i, j + 1].imshow(seg, vmin=-10, vmax=10)\n        plt.colorbar(p, ax=axarr[i, j + 1])","metadata":{"execution":{"iopub.status.busy":"2022-08-13T23:49:26.655922Z","iopub.execute_input":"2022-08-13T23:49:26.656454Z","iopub.status.idle":"2022-08-13T23:49:37.577643Z","shell.execute_reply.started":"2022-08-13T23:49:26.656417Z","shell.execute_reply":"2022-08-13T23:49:37.576814Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference 🔥","metadata":{}},{"cell_type":"code","source":"model = SemanticSegmentation.load_from_checkpoint(\n    \"semantic_segmentation_model.pt\"\n)\ntest_images = glob.glob(os.path.join(DATASET_FOLDER, \"test_images\", \"*.tiff\"))\nprint(f\"images: {len(test_images)}\")\n\n!rm /kaggle/temp/images/*\n!rm /kaggle/temp/masks/*","metadata":{"execution":{"iopub.status.busy":"2022-08-13T23:49:37.579285Z","iopub.execute_input":"2022-08-13T23:49:37.579892Z","iopub.status.idle":"2022-08-13T23:49:41.282140Z","shell.execute_reply.started":"2022-08-13T23:49:37.579843Z","shell.execute_reply":"2022-08-13T23:49:41.280072Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import cv2\nimport numpy as np\nfrom itertools import chain\nfrom kaggle_imsegm.mask import rle_encode\nfrom torch.utils.data import DataLoader\nfrom scipy.ndimage import binary_opening\nfrom skimage.morphology import disk\nfrom skimage.color import label2rgb\n\ndf_test['pixel_size'] =  df_test['pixel_size'].fillna(0.4)\n\npreds = []\nfor _, row in df_test.iterrows():\n    scale = row[\"pixel_size\"] / 0.4\n    test_img = os.path.join(DATASET_FOLDER, \"test_images\", f\"{row['id']}.tiff\")\n    im = plt.imread(test_img)\n    \n    # perform scaling on level tiles as the input is scaled to the CNN input size anyway\n    tiles_img, idxs = extract_image_tiles(test_img, \"/kaggle/temp/images\", size=int(TILE_SIZE / scale))\n    dm = SemanticSegmentationData.from_files(\n        predict_files=tiles_img,\n        predict_transform=SemanticSegmentationInputTransform,\n        transform_kwargs=dict(image_size=IMAGE_SIZE),\n        num_classes=2,\n        batch_size=3,\n        num_workers=2,\n    )\n    pred = trainer.predict(model, datamodule=dm, output=\"labels\")\n    pred = list(chain(*pred))\n    \n    seg = np.zeros(im.shape[:2], dtype=np.uint8)\n    for tile, (y, y2, x, x2) in zip(pred, idxs):\n        y2 = min(y2, im.shape[0])\n        x2 = min(x2, im.shape[1])\n        seg[y:y2, x:x2] = np.array(tile, dtype=np.uint8)[:(y2 - y), :(x2 - x)]\n    # seg = resize(seg * 255, img.shape[:2], order=0) / 255\n    seg = (seg >= 1).astype(np.uint8)  # binary mask\n    seg = binary_opening(seg, structure=disk(6)).astype(np.uint8)\n\n    if not WITH_SUBMISSION:\n        fig, axarr = plt.subplots(ncols=3, figsize=(12, 4))\n        axarr[0].imshow(im)\n        axarr[1].imshow(label2rgb(seg, im, bg_label=0, bg_color=(1.,1.,1.), alpha=0.25))\n        ax_im = axarr[2].imshow(seg)\n        plt.colorbar(ax_im, ax=axarr[2])\n\n    rle = rle_encode(seg.T) if np.sum(seg) > 1 else {}\n    name, _ = os.path.splitext(os.path.basename(test_img))\n    preds.append({\"id\": row['id'], \"rle\": rle.get(1, \"\")})\n\ndf_pred = pd.DataFrame(preds)\ndisplay(df_pred[df_pred[\"rle\"] != \"\"].head())","metadata":{"execution":{"iopub.status.busy":"2022-08-13T23:50:50.392847Z","iopub.execute_input":"2022-08-13T23:50:50.393286Z","iopub.status.idle":"2022-08-13T23:51:16.457639Z","shell.execute_reply.started":"2022-08-13T23:50:50.393251Z","shell.execute_reply":"2022-08-13T23:51:16.456450Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Finalize submissions","metadata":{}},{"cell_type":"code","source":"df_ssub = pd.read_csv(os.path.join(DATASET_FOLDER, \"sample_submission.csv\"))\ndel df_ssub['rle']\ndf_pred = df_ssub.merge(df_pred, on='id')\n\ndf_pred[['id', 'rle']].to_csv(\"submission.csv\", index=False)\n\n!head submission.csv","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-08-13T23:51:16.460160Z","iopub.execute_input":"2022-08-13T23:51:16.460533Z","iopub.status.idle":"2022-08-13T23:51:17.584932Z","shell.execute_reply.started":"2022-08-13T23:51:16.460498Z","shell.execute_reply":"2022-08-13T23:51:17.583699Z"},"trusted":true},"execution_count":null,"outputs":[]}]}