{"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":"# Exploratory Data Analysis 🔎","metadata":{}},{"cell_type":"code","source":"!pip uninstall -y torchtext\n!mkdir -p frozen_packages/\n!cp ../input/demo-flash-semantic-segmentation/frozen_packages/* frozen_packages/\n!cp ../input/tract-segm-eda-3d-interactive-viewer/frozen_packages/* frozen_packages/\n# !pip install -q --upgrade torch torchvision\n!pip install -q \"lightning-flash[image]\" \"torchmetrics<0.8\" --pre --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 \"https://github.com/PyTorchLightning/lightning-flash/archive/refs/heads/segm/multi-label.zip\"\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":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-06-24T20:47:24.481711Z","iopub.execute_input":"2022-06-24T20:47:24.48237Z","iopub.status.idle":"2022-06-24T20:48:44.443805Z","shell.execute_reply.started":"2022-06-24T20:47:24.482265Z","shell.execute_reply":"2022-06-24T20:48:44.443009Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os, glob\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nDATASET_FOLDER = \"/kaggle/input/uw-madison-gi-tract-image-segmentation\"\ndf_train = pd.read_csv(os.path.join(DATASET_FOLDER, \"train.csv\"))\ndisplay(df_train.head())\n\ndf_pred = pd.read_csv(os.path.join(DATASET_FOLDER, \"sample_submission.csv\"))\nWITH_SUBMISSION = not df_pred.empty","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-06-24T20:48:44.447161Z","iopub.execute_input":"2022-06-24T20:48:44.447396Z","iopub.status.idle":"2022-06-24T20:48:44.988726Z","shell.execute_reply.started":"2022-06-24T20:48:44.447369Z","shell.execute_reply":"2022-06-24T20:48:44.988051Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_imgs = glob.glob(os.path.join(DATASET_FOLDER, \"train\", \"case*\", \"case*_day*\", \"scans\", \"*.png\"))\nall_imgs = [p.replace(DATASET_FOLDER, \"\") for p in all_imgs]\n\nprint(f\"images: {len(all_imgs)}\")\nprint(f\"annotated: {len(df_train['id'].unique())}\")","metadata":{"execution":{"iopub.status.busy":"2022-06-24T20:48:44.990132Z","iopub.execute_input":"2022-06-24T20:48:44.99044Z","iopub.status.idle":"2022-06-24T20:48:48.788609Z","shell.execute_reply.started":"2022-06-24T20:48:44.9904Z","shell.execute_reply":"2022-06-24T20:48:48.78786Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from pprint import pprint\nfrom kaggle_imsegm.data_io import extract_tract_details\n\npprint(extract_tract_details(df_train['id'].iloc[0], DATASET_FOLDER))\n\ndf_train[['Case','Day','Slice', 'image', 'image_path', 'height', 'width']] = df_train['id'].apply(\n    lambda x: pd.Series(extract_tract_details(x, DATASET_FOLDER))\n)\ndisplay(df_train.head())","metadata":{"execution":{"iopub.status.busy":"2022-06-24T20:48:48.790614Z","iopub.execute_input":"2022-06-24T20:48:48.791031Z","iopub.status.idle":"2022-06-24T20:51:03.645533Z","shell.execute_reply.started":"2022-06-24T20:48:48.790994Z","shell.execute_reply":"2022-06-24T20:51:03.644808Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prepare custom dataset 💽","metadata":{}},{"cell_type":"code","source":"import os.path\nfrom typing import Callable, Tuple, Sequence\n\nimport torch\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nfrom torch.utils.data import Dataset\nfrom tqdm.auto import tqdm\n\nfrom kaggle_imsegm.dataset import TractDataset2D\n\nds = TractDataset2D(df_train, DATASET_FOLDER)\nprint(len(ds))","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-06-24T20:51:03.646654Z","iopub.execute_input":"2022-06-24T20:51:03.646895Z","iopub.status.idle":"2022-06-24T20:51:52.542436Z","shell.execute_reply.started":"2022-06-24T20:51:03.646859Z","shell.execute_reply":"2022-06-24T20:51:52.54177Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"spl = ds[255]\nimg, seg = spl[\"input\"], spl[\"target\"]\nprint(img.shape)\nfig, axarr = plt.subplots(ncols=4, figsize=(12, 3))\naxarr[0].imshow(np.rollaxis(img.numpy(), 0, 3), cmap=\"gray\")\nprint(np.argmax(seg, axis=0).shape)\nfor i in range(seg.shape[0]):\n    axarr[i + 1].imshow(seg[i, ...])","metadata":{"execution":{"iopub.status.busy":"2022-06-24T20:51:52.543653Z","iopub.execute_input":"2022-06-24T20:51:52.543909Z","iopub.status.idle":"2022-06-24T20:51:53.10018Z","shell.execute_reply.started":"2022-06-24T20:51:52.54387Z","shell.execute_reply":"2022-06-24T20:51:53.099478Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from typing import Any, Callable, Dict, Tuple, Type, Union\nfrom pytorch_lightning import LightningDataModule\nfrom torch.utils.data import DataLoader, Dataset\nimport albumentations as alb\n\nfrom kaggle_imsegm.transform import FlashAlbumentationsAdapter\nfrom kaggle_imsegm.dataset import TractData\n\nCOLOR_MEAN: float = 0.349977\nCOLOR_STD: float = 0.215829\nDEFAULT_TRANSFORM = FlashAlbumentationsAdapter(\n    [alb.Resize(224, 224), alb.Normalize(mean=COLOR_MEAN, std=COLOR_STD, max_pixel_value=255)]\n)\n    \ndm = TractData(df_train, DATASET_FOLDER, dataloader_kwargs=dict(batch_size=12, num_workers=3))\ndm.setup()\nprint(len(dm.train_dataloader()))\nprint(len(dm.val_dataloader()))","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-06-24T20:51:53.101701Z","iopub.execute_input":"2022-06-24T20:51:53.10219Z","iopub.status.idle":"2022-06-24T20:52:12.149894Z","shell.execute_reply.started":"2022-06-24T20:51:53.102151Z","shell.execute_reply":"2022-06-24T20:52:12.148086Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from kaggle_imsegm.visual import show_tract_datamodule_samples_2d\n\n_= show_tract_datamodule_samples_2d(dm.val_dataloader(), nb=3)","metadata":{"execution":{"iopub.status.busy":"2022-06-24T20:52:12.151112Z","iopub.status.idle":"2022-06-24T20:52:12.1516Z","shell.execute_reply.started":"2022-06-24T20:52:12.151353Z","shell.execute_reply":"2022-06-24T20:52:12.151379Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Lightning⚡Flash & UNet++ & albumentations\n\nlets follow the Semantinc segmentation example: https://lightning-flash.readthedocs.io/en/stable/reference/semantic_segmentation.html","metadata":{}},{"cell_type":"code","source":"import torch\n\nimport flash\nfrom flash.core.data.utils import download_data\nfrom flash.image import SemanticSegmentation, SemanticSegmentationData\n\nprint(flash.__version__)","metadata":{"execution":{"iopub.status.busy":"2022-06-24T20:52:12.153327Z","iopub.status.idle":"2022-06-24T20:52:12.153743Z","shell.execute_reply.started":"2022-06-24T20:52:12.15352Z","shell.execute_reply":"2022-06-24T20:52:12.153542Z"},"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\nfrom flash.core.data.io.input_transform import InputTransform\nfrom flash.image.segmentation.input_transform import prepare_target, remove_extra_dimensions\n# from kaggle_imsegm.augment import FlashAlbumentationsAdapter\n\nIMAGE_SIZE = (320, 320)\nTRAIN_TRANSFORM = FlashAlbumentationsAdapter([\n    alb.Resize(*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.2, scale_limit=0.08, rotate_limit=10, p=1.),\n    alb.GaussNoise(var_limit=(0.001, 0.02), 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.7),\n    alb.Normalize(mean=COLOR_MEAN, std=COLOR_STD, max_pixel_value=255),\n])\nVAL_TRANSFORM = FlashAlbumentationsAdapter([\n    alb.Resize(*IMAGE_SIZE), alb.Normalize(mean=COLOR_MEAN, std=COLOR_STD, max_pixel_value=255)\n])","metadata":{"execution":{"iopub.status.busy":"2022-06-24T20:52:12.154854Z","iopub.status.idle":"2022-06-24T20:52:12.155681Z","shell.execute_reply.started":"2022-06-24T20:52:12.155437Z","shell.execute_reply":"2022-06-24T20:52:12.155463Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_imgs = glob.glob(os.path.join(DATASET_FOLDER, \"test\", \"**\", \"*.png\"), recursive=True)\nif not sample_imgs:\n    sample_imgs = glob.glob(os.path.join(DATASET_FOLDER, \"train\", \"case123\", \"**\", \"*.png\"), recursive=True)\nprint(f\"images: {len(sample_imgs)}\")\nsample_imgs = [p.replace(DATASET_FOLDER + os.path.sep, \"\") for p in sample_imgs[70:75]]\ntab_preds = pd.DataFrame({\"image_path\": sample_imgs})","metadata":{"execution":{"iopub.status.busy":"2022-06-24T20:52:12.156875Z","iopub.status.idle":"2022-06-24T20:52:12.157811Z","shell.execute_reply.started":"2022-06-24T20:52:12.157572Z","shell.execute_reply":"2022-06-24T20:52:12.157597Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"datamodule = TractData(\n    df_train,\n    dataset_dir=DATASET_FOLDER,\n    df_predict=tab_preds,\n    train_transform=TRAIN_TRANSFORM,\n    input_transform=VAL_TRANSFORM,\n    dataloader_kwargs=dict(batch_size=18, num_workers=3),\n    val_split=0.01 if WITH_SUBMISSION else 0.1, \n)\ndatamodule.setup()\nLABELS = datamodule.labels\nassert len(LABELS) == 3","metadata":{"execution":{"iopub.status.busy":"2022-06-24T20:52:12.158888Z","iopub.status.idle":"2022-06-24T20:52:12.159763Z","shell.execute_reply.started":"2022-06-24T20:52:12.159536Z","shell.execute_reply":"2022-06-24T20:52:12.159558Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"_= show_tract_datamodule_samples_2d(datamodule.train_dataloader(), nb=5, skip_empty=True)","metadata":{"execution":{"iopub.status.busy":"2022-06-24T20:52:12.160824Z","iopub.status.idle":"2022-06-24T20:52:12.161673Z","shell.execute_reply.started":"2022-06-24T20:52:12.161436Z","shell.execute_reply":"2022-06-24T20:52:12.161461Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 2. Build the task","metadata":{}},{"cell_type":"code","source":"# import segmentation_models_pytorch as smp\nfrom kaggle_imsegm.model import MixedLoss\nfrom kaggle_imsegm.transform import SemanticSegmentationOutputTransform\n\nmodel = SemanticSegmentation(\n    backbone=\"efficientnet-b3\",\n    head=\"unetplusplus\",\n    pretrained=False,\n    optimizer=\"asgd\",\n    learning_rate=0.05,\n    loss_fn=MixedLoss(\"dice\", smooth=0.01),\n    lr_scheduler=(\"cosineannealinglr\", {\"T_max\": 500, \"eta_min\": 1_000_000}),\n    num_classes=3,\n    multi_label=True,\n    output_transform=SemanticSegmentationOutputTransform(),\n)","metadata":{"execution":{"iopub.status.busy":"2022-06-24T20:52:12.162804Z","iopub.status.idle":"2022-06-24T20:52:12.163636Z","shell.execute_reply.started":"2022-06-24T20:52:12.1634Z","shell.execute_reply":"2022-06-24T20:52:12.163423Z"},"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\nGPUs = torch.cuda.device_count()\n\ntrainer = flash.Trainer(\n    max_epochs=15 if WITH_SUBMISSION else 5,\n    logger=pl.loggers.CSVLogger(save_dir='logs/'),\n    gpus=GPUs,\n    precision=16 if GPUs else 32,\n    accumulate_grad_batches=24,\n    gradient_clip_val=0.01,\n    limit_train_batches=1.0 if WITH_SUBMISSION else 0.2,\n    limit_val_batches=1.0 if WITH_SUBMISSION else 0.3,\n)","metadata":{"execution":{"iopub.status.busy":"2022-06-24T20:52:12.164765Z","iopub.status.idle":"2022-06-24T20:52:12.165599Z","shell.execute_reply.started":"2022-06-24T20:52:12.16536Z","shell.execute_reply":"2022-06-24T20:52:12.165384Z"},"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-06-24T20:52:12.16674Z","iopub.status.idle":"2022-06-24T20:52:12.167569Z","shell.execute_reply.started":"2022-06-24T20:52:12.167342Z","shell.execute_reply":"2022-06-24T20:52:12.167364Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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.grid()","metadata":{"execution":{"iopub.status.busy":"2022-06-24T20:52:12.168687Z","iopub.status.idle":"2022-06-24T20:52:12.169507Z","shell.execute_reply.started":"2022-06-24T20:52:12.169277Z","shell.execute_reply":"2022-06-24T20:52:12.1693Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 4. Segment a few images!","metadata":{}},{"cell_type":"code","source":"from itertools import chain\n\npreds = trainer.predict(model, datamodule=datamodule)  #, output=\"preds\"\npreds = list(chain(*preds))","metadata":{"execution":{"iopub.status.busy":"2022-06-24T20:52:12.170726Z","iopub.status.idle":"2022-06-24T20:52:12.171582Z","shell.execute_reply.started":"2022-06-24T20:52:12.171346Z","shell.execute_reply":"2022-06-24T20:52:12.17137Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, axarr = plt.subplots(ncols=4, nrows=len(sample_imgs), figsize=(12, 3 * len(sample_imgs)))\nfor i, pred in enumerate(preds):\n    print(pred.keys())\n    img = pred['input']\n    print(img.shape, img.min(), img.max())\n    axarr[i, 0].imshow(img)\n    for j, seg in enumerate(pred['preds']):\n        print(seg.shape, seg.min(), seg.max())\n        im = axarr[i, j + 1].imshow(seg, vmin=-10, vmax=10)\n        plt.colorbar(im, ax=axarr[i, j + 1])","metadata":{"execution":{"iopub.status.busy":"2022-06-24T20:52:12.172723Z","iopub.status.idle":"2022-06-24T20:52:12.173575Z","shell.execute_reply.started":"2022-06-24T20:52:12.173336Z","shell.execute_reply":"2022-06-24T20:52:12.173361Z"},"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)","metadata":{"execution":{"iopub.status.busy":"2022-06-24T20:52:12.174684Z","iopub.status.idle":"2022-06-24T20:52:12.175652Z","shell.execute_reply.started":"2022-06-24T20:52:12.175377Z","shell.execute_reply":"2022-06-24T20:52:12.175405Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sfolder = \"test\" if WITH_SUBMISSION else \"train\"\nls_images = glob.glob(os.path.join(DATASET_FOLDER, sfolder, \"**\", \"*.png\"), recursive=True)\nls_images = [p.replace(DATASET_FOLDER + os.path.sep, \"\") for p in ls_images]\ncase_day = [os.path.dirname(p).split(os.path.sep)[-2] for p in ls_images]\ndf_pred = pd.DataFrame({'Case_Day': case_day, 'image_path': ls_images})\n\nif not WITH_SUBMISSION:\n    df_pred = df_pred[df_pred[\"Case_Day\"].str.startswith(\"case123_day\")]\ndisplay(df_pred.head())","metadata":{"execution":{"iopub.status.busy":"2022-06-24T20:52:12.176826Z","iopub.status.idle":"2022-06-24T20:52:12.177644Z","shell.execute_reply.started":"2022-06-24T20:52:12.17741Z","shell.execute_reply":"2022-06-24T20:52:12.177433Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Predictions for test scans","metadata":{}},{"cell_type":"code","source":"import numpy as np\nfrom itertools import chain\nfrom scipy.ndimage import binary_opening\nfrom skimage.morphology import disk\nfrom kaggle_imsegm.mask import rle_encode\n\npreds = []\nfor case_day, tab_preds in tqdm(df_pred.groupby(\"Case_Day\")):\n    dm = TractData(\n        df_train[df_train[\"id\"].str.startswith(\"case123_day\")],  # FAKE\n        dataset_dir=DATASET_FOLDER,\n        df_predict=tab_preds,\n        train_transform=TRAIN_TRANSFORM,\n        input_transform=VAL_TRANSFORM,\n        dataloader_kwargs=dict(batch_size=10, num_workers=3),\n    )\n    # dm.setup()\n    results = trainer.predict(model, datamodule=dm)\n    results = list(chain(*results))\n    assert len(tab_preds[\"image_path\"]) == len(results)\n    for img_path, spl in zip(tab_preds[\"image_path\"], results):\n        name, _ = os.path.splitext(os.path.basename(img_path))\n        id_ = f\"{case_day}_\" + \"_\".join(name.split(\"_\")[:2])\n        # print(spl.keys())\n        for i, mask in enumerate(spl[\"preds\"]):\n            mask = (mask >= 0).astype(np.uint8)\n            mask = binary_opening(mask, structure=disk(4)).astype(np.uint8)\n            # print(seg.shape)\n            rle = rle_encode(mask)[1] if np.sum(mask) > 1 else \"\"\n            preds.append({\"id\": id_, \"class\": LABELS[i], \"predicted\": rle})\n\nassert len(preds) == 3 * len(df_pred)\ndf_pred = pd.DataFrame(preds)","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-06-24T20:52:12.178774Z","iopub.status.idle":"2022-06-24T20:52:12.180167Z","shell.execute_reply.started":"2022-06-24T20:52:12.179916Z","shell.execute_reply":"2022-06-24T20:52:12.179942Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display(df_pred[df_pred[\"predicted\"] != \"\"].head())","metadata":{"execution":{"iopub.status.busy":"2022-06-24T20:52:12.181244Z","iopub.status.idle":"2022-06-24T20:52:12.18184Z","shell.execute_reply.started":"2022-06-24T20:52:12.181571Z","shell.execute_reply":"2022-06-24T20:52:12.181596Z"},"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['predicted']\nif WITH_SUBMISSION:\n    assert len(df_ssub) == len(df_pred)\ndf_pred = df_ssub.merge(df_pred, on=['id','class'])\n\ndf_pred[['id', 'class', 'predicted']].to_csv(\"submission.csv\", index=False)\n\n!head submission.csv","metadata":{"execution":{"iopub.status.busy":"2022-06-24T20:52:12.183592Z","iopub.status.idle":"2022-06-24T20:52:12.184403Z","shell.execute_reply.started":"2022-06-24T20:52:12.184141Z","shell.execute_reply":"2022-06-24T20:52:12.184166Z"},"trusted":true},"execution_count":null,"outputs":[]}]}