{"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: inference Lightning⚡Flash on tiled images\n\nThis is derived from Flash training baseline: https://www.kaggle.com/code/jirkaborovec/ftus-segm-baseline-flash-unet-tiled-aug-images","metadata":{}},{"cell_type":"code","source":"!pip uninstall -y torchtext\n# !pip install -q --upgrade torch torchvision\n!pip install -q \"lightning-flash[image]\" \"torchmetrics<0.8\" --no-index --find-links ../input/ftus-segm-baseline-flash-unet-tiled-aug-images/frozen_packages\n!pip install -q -U timm segmentation-models-pytorch --no-index --find-links ../input/ftus-segm-baseline-flash-unet-tiled-aug-images/frozen_packages\n!pip install -q 'kaggle-image-segmentation' --no-index --find-links ../input/ftus-segm-baseline-flash-unet-tiled-aug-images/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-10T10:44:05.186292Z","iopub.execute_input":"2022-08-10T10:44:05.186767Z","iopub.status.idle":"2022-08-10T10:45:27.218584Z","shell.execute_reply.started":"2022-08-10T10:44:05.186670Z","shell.execute_reply":"2022-08-10T10:45:27.216440Z"},"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\"\nANNOT_DATASET = \"/kaggle/input/hacking-the-human-body-annotation-masks\"\n\ndf_test = pd.read_csv(os.path.join(DATASET_FOLDER, \"test.csv\"))\ndisplay(df_test.head())","metadata":{"execution":{"iopub.status.busy":"2022-08-10T10:45:27.222048Z","iopub.execute_input":"2022-08-10T10:45:27.222881Z","iopub.status.idle":"2022-08-10T10:45:27.262954Z","shell.execute_reply.started":"2022-08-10T10:45:27.222833Z","shell.execute_reply":"2022-08-10T10:45:27.261753Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Make a grid/tiles","metadata":{}},{"cell_type":"code","source":"import numpy as np\nfrom PIL import Image\n\ndef tile_image(p_img, folder, size: int = 768) -> list:\n    w = h = size\n    im = np.array(Image.open(p_img))\n    # https://stackoverflow.com/a/47581978/4521646\n    tiles = [im[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    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, tile in enumerate(tiles):\n        if tile.shape[:2] != (h, w):\n            tile_ = tile\n            tile = np.zeros_like(tiles[0])\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","metadata":{"_kg_hide-output":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-08-10T10:45:27.264338Z","iopub.execute_input":"2022-08-10T10:45:27.264742Z","iopub.status.idle":"2022-08-10T10:45:27.276088Z","shell.execute_reply.started":"2022-08-10T10:45:27.264700Z","shell.execute_reply":"2022-08-10T10:45:27.274970Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Porting Lighnting ⚡ training config","metadata":{}},{"cell_type":"code","source":"import torch\n\nimport flash\nimport numpy as np\nfrom flash.core.data.utils import download_data\nfrom flash.image import SemanticSegmentation, SemanticSegmentationData\n\nimport warnings\nwarnings.filterwarnings(\"ignore\", category=DeprecationWarning) ","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-08-10T10:45:27.279787Z","iopub.execute_input":"2022-08-10T10:45:27.280821Z","iopub.status.idle":"2022-08-10T10:45:43.964716Z","shell.execute_reply.started":"2022-08-10T10:45:27.280776Z","shell.execute_reply":"2022-08-10T10:45:43.963592Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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 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-10T10:45:43.965910Z","iopub.execute_input":"2022-08-10T10:45:43.966561Z","iopub.status.idle":"2022-08-10T10:45:43.981572Z","shell.execute_reply.started":"2022-08-10T10:45:43.966529Z","shell.execute_reply":"2022-08-10T10:45:43.980454Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import json\n\nwith open(os.path.join(ANNOT_DATASET, 'labels.json')) as fp:\n    LABELS = json.load(fp)\n\nTILE_SIZE = 1024\nIMAGE_SIZE = (384, 384)\nnb_labels = len(LABELS) + 1\n\n!mkdir -p /kaggle/temp/images/","metadata":{"execution":{"iopub.status.busy":"2022-08-10T10:45:43.983403Z","iopub.execute_input":"2022-08-10T10:45:43.983757Z","iopub.status.idle":"2022-08-10T10:45:45.143049Z","shell.execute_reply.started":"2022-08-10T10:45:43.983728Z","shell.execute_reply":"2022-08-10T10:45:45.141752Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pytorch_lightning as pl\n\ntrainer = flash.Trainer(\n    gpus=torch.cuda.device_count(),\n    logger=pl.loggers.CSVLogger(save_dir='logs/'),\n    precision=16 if torch.cuda.is_available() else 32,\n)","metadata":{"execution":{"iopub.status.busy":"2022-08-10T10:45:45.145450Z","iopub.execute_input":"2022-08-10T10:45:45.146216Z","iopub.status.idle":"2022-08-10T10:45:45.163991Z","shell.execute_reply.started":"2022-08-10T10:45:45.146164Z","shell.execute_reply":"2022-08-10T10:45:45.163120Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load model and sample prediction","metadata":{}},{"cell_type":"code","source":"model = SemanticSegmentation.load_from_checkpoint(\n    \"/kaggle/input/ftus-segm-baseline-flash-unet-tiled-aug-images/semantic_segmentation_model.pt\"\n)\n\ntest_images = glob.glob(os.path.join(DATASET_FOLDER, \"test_images\", \"*.tiff\"))\nprint(f\"images: {len(test_images)}\")\nWITH_SUBMISSION = len(test_images) > 1","metadata":{"execution":{"iopub.status.busy":"2022-08-10T10:45:45.165399Z","iopub.execute_input":"2022-08-10T10:45:45.166029Z","iopub.status.idle":"2022-08-10T10:45:53.154884Z","shell.execute_reply.started":"2022-08-10T10:45:45.165994Z","shell.execute_reply":"2022-08-10T10:45:53.153777Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tiles_img, _ = tile_image(test_images[0], \"/kaggle/temp/images\", size=TILE_SIZE)\ndm = 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)","metadata":{"execution":{"iopub.status.busy":"2022-08-10T10:45:53.156178Z","iopub.execute_input":"2022-08-10T10:45:53.156532Z","iopub.status.idle":"2022-08-10T10:45:54.364277Z","shell.execute_reply.started":"2022-08-10T10:45:53.156499Z","shell.execute_reply":"2022-08-10T10:45:54.363160Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference 🔥","metadata":{}},{"cell_type":"code","source":"from itertools import chain\n\nnrows = max(2, len(tiles_img))\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-10T10:45:54.367299Z","iopub.execute_input":"2022-08-10T10:45:54.367638Z","iopub.status.idle":"2022-08-10T10:46:12.777086Z","shell.execute_reply.started":"2022-08-10T10:45:54.367608Z","shell.execute_reply":"2022-08-10T10:46:12.775954Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\nimport 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 import color\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 = tile_image(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, (i1, i2, j1, j2) in zip(pred, idxs):\n        i2 = min(i2, im.shape[0])\n        j2 = min(j2, im.shape[1])\n        seg[i1:i2, j1:j2] = np.array(tile, dtype=np.uint8)[:(i2 - i1), :(j2 - j1)]\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    del pred\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(color.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    del im\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    del seg\n    \n    !rm /kaggle/temp/images/*\n\ndf_pred = pd.DataFrame(preds)\ndisplay(df_pred[df_pred[\"rle\"] != \"\"].head())","metadata":{"execution":{"iopub.status.busy":"2022-08-10T10:49:46.655527Z","iopub.execute_input":"2022-08-10T10:49:46.658793Z","iopub.status.idle":"2022-08-10T10:50:44.817379Z","shell.execute_reply.started":"2022-08-10T10:49:46.658719Z","shell.execute_reply":"2022-08-10T10:50:44.816288Z"},"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']\n\ndf_pred = df_ssub.merge(df_pred, on='id')\ndisplay(df_pred.head())","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-08-10T10:50:44.819746Z","iopub.execute_input":"2022-08-10T10:50:44.820993Z","iopub.status.idle":"2022-08-10T10:50:44.860325Z","shell.execute_reply.started":"2022-08-10T10:50:44.820955Z","shell.execute_reply":"2022-08-10T10:50:44.859184Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_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-10T10:50:44.862424Z","iopub.execute_input":"2022-08-10T10:50:44.863127Z","iopub.status.idle":"2022-08-10T10:50:46.055169Z","shell.execute_reply.started":"2022-08-10T10:50:44.863081Z","shell.execute_reply":"2022-08-10T10:50:46.053598Z"},"trusted":true},"execution_count":null,"outputs":[]}]}