{"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":"# [UW-Madison GI Tract Image Segmentation](https://www.kaggle.com/competitions/uw-madison-gi-tract-image-segmentation/)\n> Track healthy organs in medical scans to improve cancer treatment\n\n<img src=\"https://storage.googleapis.com/kaggle-competitions/kaggle/27923/logos/header.png?t=2021-06-02-20-30-25\">","metadata":{}},{"cell_type":"markdown","source":"# ⚽ Methodlogy\n* In this notebook I'll demonstrate how to train **Unet** model using PyTorch.\n* For mask I'll be using pre-computed mask from [here](https://www.kaggle.com/datasets/awsaf49/uwmgi-mask-dataset)\n* As there are overlaps between **Stomach**, **Large Bowel** & **Small Bowel** classes, this is a **MultiLabel Segmentation** task, so final activaion should be `sigmoid` instead of `softmax`.\n* For data split I'll be using **StratifiedGroupFold** to avoid data leakage due to `case` and to stratify `empty` and `non-empty` mask cases.\n* You can play with different models and losses.","metadata":{}},{"cell_type":"markdown","source":"# 🚩 Version Info","metadata":{}},{"cell_type":"markdown","source":"# 📒 Notebooks\n📌 **UNet**:\n* Train: [UWMGI: Unet [Train] [PyTorch]](https://www.kaggle.com/code/awsaf49/uwmgi-unet-train-pytorch/)\n* Infer: [UWMGI: Unet [Infer] [PyTorch]](https://www.kaggle.com/code/awsaf49/uwmgi-unet-infer-pytorch/)\n\n📌 **Data/Dataset**:\n* Data: [UWMGI: Mask Data](https://www.kaggle.com/datasets/awsaf49/uwmgi-mask-data)\n* Dataset: [UWMGI: Mask Dataset](https://www.kaggle.com/datasets/awsaf49/uwmgi-mask-dataset)","metadata":{}},{"cell_type":"markdown","source":"## Please Upvote if you Find this Useful :)","metadata":{}},{"cell_type":"markdown","source":"# 🛠 Install Libraries","metadata":{}},{"cell_type":"code","source":"!pip install -q ../input/pytorch-segmentation-models-lib/pretrainedmodels-0.7.4/pretrainedmodels-0.7.4\n!pip install -q ../input/pytorch-segmentation-models-lib/efficientnet_pytorch-0.6.3/efficientnet_pytorch-0.6.3\n!pip install -q ../input/pytorch-segmentation-models-lib/timm-0.4.12-py3-none-any.whl\n!pip install -q ../input/pytorch-segmentation-models-lib/segmentation_models_pytorch-0.2.0-py3-none-any.whl","metadata":{"_kg_hide-output":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-04-26T09:02:49.022747Z","iopub.execute_input":"2022-04-26T09:02:49.023628Z","iopub.status.idle":"2022-04-26T09:04:42.980579Z","shell.execute_reply.started":"2022-04-26T09:02:49.023521Z","shell.execute_reply":"2022-04-26T09:04:42.979774Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 📚 Import Libraries ","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\npd.options.plotting.backend = \"plotly\"\nimport random\nfrom glob import glob\nimport os, shutil\nfrom tqdm import tqdm\ntqdm.pandas()\nimport time\nimport copy\nimport joblib\nfrom collections import defaultdict\nimport gc\nfrom IPython import display as ipd\n\n# visualization\nimport cv2\nimport matplotlib.pyplot as plt\n\n# Sklearn\nfrom sklearn.model_selection import StratifiedKFold, KFold\n\n# PyTorch \nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda import amp\nimport torch.nn.functional as F\n\nimport timm\n\n# Albumentations for augmentations\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\n# For colored terminal text\nfrom colorama import Fore, Back, Style\nc_  = Fore.GREEN\nsr_ = Style.RESET_ALL\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\n# For descriptive error messages\nos.environ['CUDA_LAUNCH_BLOCKING'] = \"1\"","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-04-26T09:04:42.982432Z","iopub.execute_input":"2022-04-26T09:04:42.982703Z","iopub.status.idle":"2022-04-26T09:04:51.179421Z","shell.execute_reply.started":"2022-04-26T09:04:42.982669Z","shell.execute_reply":"2022-04-26T09:04:51.178544Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ⚙️ Configuration ","metadata":{}},{"cell_type":"code","source":"class CFG:\n    seed          = 101\n    debug         = False # set debug=False for Full Training\n    exp_name      = 'Baseline'\n    comment       = 'unet-efficientnet_b1-224x224'\n    model_name    = 'Unet'\n    backbone      = 'efficientnet-b1'\n    train_bs      = 64\n    valid_bs      = train_bs*2\n    img_size      = [224, 224]\n    epochs        = 15\n    lr            = 2e-3\n    scheduler     = 'CosineAnnealingLR'\n    min_lr        = 1e-6\n    T_max         = int(30000/train_bs*epochs)+50\n    T_0           = 25\n    warmup_epochs = 0\n    wd            = 1e-6\n    n_accumulate  = max(1, 32//train_bs)\n    n_fold        = 5\n    num_classes   = 3\n    device        = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n    thr           = 0.45\n    ttas          = [0]","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-04-26T09:04:51.180827Z","iopub.execute_input":"2022-04-26T09:04:51.181736Z","iopub.status.idle":"2022-04-26T09:04:51.245163Z","shell.execute_reply.started":"2022-04-26T09:04:51.181699Z","shell.execute_reply":"2022-04-26T09:04:51.244053Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ❗ Reproducibility","metadata":{}},{"cell_type":"code","source":"def set_seed(seed = 42):\n    '''Sets the seed of the entire notebook so results are the same every time we run.\n    This is for REPRODUCIBILITY.'''\n    np.random.seed(seed)\n    random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    # When running on the CuDNN backend, two further options must be set\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    # Set a fixed value for the hash seed\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    print('> SEEDING DONE')\n    \nset_seed(CFG.seed)","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-04-26T09:04:51.247564Z","iopub.execute_input":"2022-04-26T09:04:51.248083Z","iopub.status.idle":"2022-04-26T09:04:51.262541Z","shell.execute_reply.started":"2022-04-26T09:04:51.248039Z","shell.execute_reply":"2022-04-26T09:04:51.261688Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🔨 Utility","metadata":{}},{"cell_type":"code","source":"def get_metadata(row):\n    data = row['id'].split('_')\n    case = int(data[0].replace('case',''))\n    day = int(data[1].replace('day',''))\n    slice_ = int(data[-1])\n    row['case'] = case\n    row['day'] = day\n    row['slice'] = slice_\n    return row\n\ndef path2info(row):\n    path = row['image_path']\n    data = path.split('/')\n    slice_ = int(data[-1].split('_')[1])\n    case = int(data[-3].split('_')[0].replace('case',''))\n    day = int(data[-3].split('_')[1].replace('day',''))\n    width = int(data[-1].split('_')[2])\n    height = int(data[-1].split('_')[3])\n    row['height'] = height\n    row['width'] = width\n    row['case'] = case\n    row['day'] = day\n    row['slice'] = slice_\n#     row['id'] = f'case{case}_day{day}_slice_{slice_}'\n    return row","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-04-26T09:04:51.263898Z","iopub.execute_input":"2022-04-26T09:04:51.264272Z","iopub.status.idle":"2022-04-26T09:04:51.277361Z","shell.execute_reply.started":"2022-04-26T09:04:51.264213Z","shell.execute_reply":"2022-04-26T09:04:51.276472Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_img(path):\n    img = cv2.imread(path, cv2.IMREAD_UNCHANGED)\n    img = np.tile(img[...,None], [1, 1, 3]) # gray to rgb\n    img = img.astype('float32') # original is uint16\n    mx = np.max(img)\n    if mx:\n        img/=mx # scale image to [0, 1]\n    return img\n\ndef load_msk(path):\n    msk = np.load(path)\n    msk = msk.astype('float32')\n    msk/=255.0\n    return msk\n\ndef show_img(img, mask=None):\n    clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8,8))\n    img = clahe.apply(img)\n    plt.imshow(img, cmap='bone')\n    \n    if mask is not None:\n        # plt.imshow(np.ma.masked_where(mask!=1, mask), alpha=0.5, cmap='autumn')\n        plt.imshow(mask, alpha=0.5)\n        handles = [Rectangle((0,0),1,1, color=_c) for _c in [(0.667,0.0,0.0), (0.0,0.667,0.0), (0.0,0.0,0.667)]]\n        labels = [\"Large Bowel\", \"Small Bowel\", \"Stomach\"]\n        plt.legend(handles,labels)\n    plt.axis('off')","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-04-26T09:04:51.278905Z","iopub.execute_input":"2022-04-26T09:04:51.279198Z","iopub.status.idle":"2022-04-26T09:04:51.292173Z","shell.execute_reply.started":"2022-04-26T09:04:51.279158Z","shell.execute_reply":"2022-04-26T09:04:51.291383Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ref: https://www.kaggle.com/paulorzp/run-length-encode-and-decode\ndef rle_decode(mask_rle, shape):\n    '''\n    mask_rle: run-length as string formated (start length)\n    shape: (height,width) of array to return \n    Returns numpy array, 1 - mask, 0 - background\n\n    '''\n    s = mask_rle.split()\n    starts, lengths = [np.asarray(x, dtype=int) for x in (s[0:][::2], s[1:][::2])]\n    starts -= 1\n    ends = starts + lengths\n    img = np.zeros(shape[0]*shape[1], dtype=np.uint8)\n    for lo, hi in zip(starts, ends):\n        img[lo:hi] = 1\n    return img.reshape(shape)  # Needed to align to RLE direction\n\n\n# ref.: https://www.kaggle.com/stainsby/fast-tested-rle\ndef rle_encode(img):\n    '''\n    img: numpy array, 1 - mask, 0 - background\n    Returns run length as string formated\n    '''\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    return ' '.join(str(x) for x in runs)","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-04-26T09:04:51.293629Z","iopub.execute_input":"2022-04-26T09:04:51.294094Z","iopub.status.idle":"2022-04-26T09:04:51.30445Z","shell.execute_reply.started":"2022-04-26T09:04:51.294052Z","shell.execute_reply":"2022-04-26T09:04:51.303638Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 📖 Meta Data","metadata":{}},{"cell_type":"code","source":"BASE_PATH  = '/kaggle/input/uw-madison-gi-tract-image-segmentation'\nCKPT_DIR = '/kaggle/input/uwmgi-unet-train-pytorch-ds'","metadata":{"execution":{"iopub.status.busy":"2022-04-26T09:06:28.738413Z","iopub.execute_input":"2022-04-26T09:06:28.739214Z","iopub.status.idle":"2022-04-26T09:06:28.743996Z","shell.execute_reply.started":"2022-04-26T09:06:28.739168Z","shell.execute_reply":"2022-04-26T09:06:28.743219Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train","metadata":{}},{"cell_type":"code","source":"# df = pd.read_csv('../input/uwmgi-mask-dataset/uw-madison-gi-tract-image-segmentation/train.csv')\n# df['empty'] = df.segmentation.map(lambda x: int(pd.isna(x)))\n\n# df2 = df.groupby(['id'])['class'].agg(list).to_frame().reset_index()\n# df2 = df2.merge(df.groupby(['id'])['segmentation'].agg(list), on=['id'])\n# # df = df[['id','case','day','image_path','mask_path','height','width', 'empty']]\n\n# df = df.drop(columns=['segmentation', 'class'])\n# df = df.groupby(['id']).head(1).reset_index(drop=True)\n# df = df.merge(df2, on=['id'])\n# df.head()","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-04-24T21:32:53.951891Z","iopub.execute_input":"2022-04-24T21:32:53.952496Z","iopub.status.idle":"2022-04-24T21:32:53.956302Z","shell.execute_reply.started":"2022-04-24T21:32:53.952458Z","shell.execute_reply":"2022-04-24T21:32:53.955631Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Test","metadata":{}},{"cell_type":"code","source":"sub_df = pd.read_csv('../input/uw-madison-gi-tract-image-segmentation/sample_submission.csv')\nif not len(sub_df):\n    debug = True\n    sub_df = pd.read_csv('../input/uw-madison-gi-tract-image-segmentation/train.csv')[:1000*3]\n    sub_df = sub_df.drop(columns=['class','segmentation']).drop_duplicates()\nelse:\n    debug = False\n    sub_df = sub_df.drop(columns=['class','predicted']).drop_duplicates()\nsub_df = sub_df.progress_apply(get_metadata,axis=1)","metadata":{"execution":{"iopub.status.busy":"2022-04-26T09:06:31.67199Z","iopub.execute_input":"2022-04-26T09:06:31.672693Z","iopub.status.idle":"2022-04-26T09:06:33.56803Z","shell.execute_reply.started":"2022-04-26T09:06:31.672652Z","shell.execute_reply":"2022-04-26T09:06:33.567248Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if debug:\n    paths = glob(f'/kaggle/input/uw-madison-gi-tract-image-segmentation/train/**/*png',recursive=True)\n#     paths = sorted(paths)\nelse:\n    paths = glob(f'/kaggle/input/uw-madison-gi-tract-image-segmentation/test/**/*png',recursive=True)\n#     paths = sorted(paths)\npath_df = pd.DataFrame(paths, columns=['image_path'])\npath_df = path_df.progress_apply(path2info, axis=1)\npath_df.head()","metadata":{"execution":{"iopub.status.busy":"2022-04-26T09:06:35.333781Z","iopub.execute_input":"2022-04-26T09:06:35.33429Z","iopub.status.idle":"2022-04-26T09:08:12.805346Z","shell.execute_reply.started":"2022-04-26T09:06:35.334225Z","shell.execute_reply":"2022-04-26T09:08:12.804602Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = sub_df.merge(path_df, on=['case','day','slice'], how='left')\ntest_df.head()","metadata":{"execution":{"iopub.status.busy":"2022-04-26T09:08:12.806963Z","iopub.execute_input":"2022-04-26T09:08:12.808669Z","iopub.status.idle":"2022-04-26T09:08:12.834845Z","shell.execute_reply.started":"2022-04-26T09:08:12.808627Z","shell.execute_reply":"2022-04-26T09:08:12.834146Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🍚 Dataset","metadata":{}},{"cell_type":"code","source":"class BuildDataset(torch.utils.data.Dataset):\n    def __init__(self, df, label=False, transforms=None):\n        self.df         = df\n        self.label      = label\n        self.img_paths  = df['image_path'].tolist()\n        self.ids        = df['id'].tolist()\n        if 'msk_path' in df.columns:\n            self.msk_paths  = df['mask_path'].tolist()\n        else:\n            self.msk_paths = None\n        self.transforms = transforms\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        img_path  = self.img_paths[index]\n        id_       = self.ids[index]\n        img = []\n        img = load_img(img_path)\n        h, w = img.shape[:2]\n        if self.label:\n            msk_path = self.msk_paths[index]\n            msk = load_msk(msk_path)\n            if self.transforms:\n                data = self.transforms(image=img, mask=msk)\n                img  = data['image']\n                msk  = data['mask']\n            img = np.transpose(img, (2, 0, 1))\n            msk = np.transpose(msk, (2, 0, 1))\n            return torch.tensor(img), torch.tensor(msk)\n        else:\n            if self.transforms:\n                data = self.transforms(image=img)\n                img  = data['image']\n            img = np.transpose(img, (2, 0, 1))\n            return torch.tensor(img), id_, h, w","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-04-26T09:10:28.756669Z","iopub.execute_input":"2022-04-26T09:10:28.756939Z","iopub.status.idle":"2022-04-26T09:10:28.768629Z","shell.execute_reply.started":"2022-04-26T09:10:28.756909Z","shell.execute_reply":"2022-04-26T09:10:28.76791Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🌈 Augmentations","metadata":{}},{"cell_type":"code","source":"data_transforms = {\n    \"train\": A.Compose([\n        A.Resize(*CFG.img_size, interpolation=cv2.INTER_NEAREST),\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n#         A.ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.05, rotate_limit=5, p=0.5),\n        A.OneOf([\n            A.GridDistortion(num_steps=5, distort_limit=0.05, p=1.0),\n# #             A.OpticalDistortion(distort_limit=0.05, shift_limit=0.05, p=1.0),\n            A.ElasticTransform(alpha=1, sigma=50, alpha_affine=50, p=1.0)\n        ], p=0.25),\n#         A.CoarseDropout(max_holes=8, max_height=CFG.img_size[0]//20, max_width=CFG.img_size[1]//20,\n#                          min_holes=5, fill_value=0, mask_fill_value=0, p=0.5),\n        ], p=1.0),\n    \n    \"valid\": A.Compose([\n        A.Resize(*CFG.img_size, interpolation=cv2.INTER_NEAREST),\n        ], p=1.0)\n}","metadata":{"_kg_hide-output":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-04-26T09:10:32.904985Z","iopub.execute_input":"2022-04-26T09:10:32.905765Z","iopub.status.idle":"2022-04-26T09:10:32.91313Z","shell.execute_reply.started":"2022-04-26T09:10:32.905724Z","shell.execute_reply":"2022-04-26T09:10:32.912255Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🍰 DataLoader","metadata":{}},{"cell_type":"code","source":"# test_dataset = BuildDataset(test_df, transforms=data_transforms['valid'])\n# test_loader  = DataLoader(test_dataset, batch_size=64, \n#                           num_workers=4, shuffle=False, pin_memory=True)","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-04-24T21:38:10.446474Z","iopub.execute_input":"2022-04-24T21:38:10.447023Z","iopub.status.idle":"2022-04-24T21:38:10.450459Z","shell.execute_reply.started":"2022-04-24T21:38:10.446983Z","shell.execute_reply":"2022-04-24T21:38:10.449589Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# imgs, ids, (h, w) = next(iter(test_loader))\n# imgs = imgs.permute((0, 2, 3, 1))\n# imgs.size()","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-04-24T21:38:10.848369Z","iopub.execute_input":"2022-04-24T21:38:10.848899Z","iopub.status.idle":"2022-04-24T21:38:10.852381Z","shell.execute_reply.started":"2022-04-24T21:38:10.84886Z","shell.execute_reply":"2022-04-24T21:38:10.851563Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 📦 Model\n","metadata":{}},{"cell_type":"markdown","source":"## UNet\n\n<img src=\"https://developers.arcgis.com/assets/img/python-graphics/unet.png\" width=\"600\">\n\n📌 **Pros**:\n* Performs well even with smaller data\n* Can be used with `imagenet` pretrain models\n\n📌 **Cons**:\n* Struggles with **edge** cases\n* Semantic Difference in **Skip Connection**","metadata":{}},{"cell_type":"code","source":"import segmentation_models_pytorch as smp\n\ndef build_model():\n    model = smp.Unet(\n        encoder_name=CFG.backbone,      # choose encoder, e.g. mobilenet_v2 or efficientnet-b7\n        encoder_weights=None,     # use `imagenet` pre-trained weights for encoder initialization\n        in_channels=3,                  # model input channels (1 for gray-scale images, 3 for RGB, etc.)\n        classes=CFG.num_classes,        # model output channels (number of classes in your dataset)\n        activation=None,\n    )\n    model.to(CFG.device)\n    return model\n\ndef load_model(path):\n    model = build_model()\n    model.load_state_dict(torch.load(path))\n    model.eval()\n    return model","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-04-26T09:10:36.996191Z","iopub.execute_input":"2022-04-26T09:10:36.996739Z","iopub.status.idle":"2022-04-26T09:10:38.350007Z","shell.execute_reply.started":"2022-04-26T09:10:36.996699Z","shell.execute_reply":"2022-04-26T09:10:38.349146Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # test\n# img = torch.randn(1, 1, *CFG.img_size).to(CFG.device)\n# img = (img - img.min())/(img.max() - img.min())\n# model = build_model()\n# _ = model(img)","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-04-24T21:38:13.848239Z","iopub.execute_input":"2022-04-24T21:38:13.849713Z","iopub.status.idle":"2022-04-24T21:38:13.853181Z","shell.execute_reply.started":"2022-04-24T21:38:13.849666Z","shell.execute_reply":"2022-04-24T21:38:13.852585Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🔨 Helper","metadata":{}},{"cell_type":"code","source":"import cupy as cp\n\ndef mask2rle(msk, thr=0.5):\n    '''\n    img: numpy array, 1 - mask, 0 - background\n    Returns run length as string formated\n    '''\n    msk    = cp.array(msk)\n    pixels = msk.flatten()\n    pad    = cp.array([0])\n    pixels = cp.concatenate([pad, pixels, pad])\n    runs   = cp.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)\n\ndef masks2rles(msks, ids, heights, widths):\n    pred_strings = []; pred_ids = []; pred_classes = [];\n    for idx in range(msks.shape[0]):\n        height = heights[idx].item()\n        width = widths[idx].item()\n        msk = cv2.resize(msks[idx], \n                         dsize=(width, height), \n                         interpolation=cv2.INTER_NEAREST) # back to original shape\n        rle = [None]*3\n        for midx in [0, 1, 2]:\n            rle[midx] = mask2rle(msk[...,midx])\n        pred_strings.extend(rle)\n        pred_ids.extend([ids[idx]]*len(rle))\n        pred_classes.extend(['large_bowel', 'small_bowel', 'stomach'])\n    return pred_strings, pred_ids, pred_classes","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-04-26T09:10:40.898404Z","iopub.execute_input":"2022-04-26T09:10:40.898669Z","iopub.status.idle":"2022-04-26T09:10:42.047486Z","shell.execute_reply.started":"2022-04-26T09:10:40.89864Z","shell.execute_reply":"2022-04-26T09:10:42.046729Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🔭 Inference","metadata":{}},{"cell_type":"code","source":"@torch.no_grad()\ndef infer(model_paths, test_loader, num_log=1, thr=CFG.thr):\n    msks = []; imgs = [];\n    pred_strings = []; pred_ids = []; pred_classes = [];\n    for idx, (img, ids, heights, widths) in enumerate(tqdm(test_loader, total=len(test_loader), desc='Infer ')):\n        img = img.to(CFG.device, dtype=torch.float) # .squeeze(0)\n        size = img.size()\n        msk = []\n        msk = torch.zeros((size[0], 3, size[2], size[3]), device=CFG.device, dtype=torch.float32)\n        for path in model_paths:\n            model = load_model(path)\n            out   = model(img) # .squeeze(0) # removing batch axis\n            out   = nn.Sigmoid()(out) # removing channel axis\n            msk+=out/len(model_paths)\n        msk = (msk.permute((0,2,3,1))>thr).to(torch.uint8).cpu().detach().numpy() # shape: (n, h, w, c)\n        result = masks2rles(msk, ids, heights, widths)\n        pred_strings.extend(result[0])\n        pred_ids.extend(result[1])\n        pred_classes.extend(result[2])\n        if idx<num_log:\n            img = img.permute((0,2,3,1)).cpu().detach().numpy()\n            imgs.append(img[:10])\n            msks.append(msk[:10])\n        del img, msk, out, model, result\n        gc.collect()\n        torch.cuda.empty_cache()\n    return pred_strings, pred_ids, pred_classes, imgs, msks","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-04-26T09:10:42.723494Z","iopub.execute_input":"2022-04-26T09:10:42.724227Z","iopub.status.idle":"2022-04-26T09:10:42.738635Z","shell.execute_reply.started":"2022-04-26T09:10:42.724186Z","shell.execute_reply":"2022-04-26T09:10:42.73587Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset = BuildDataset(test_df, transforms=data_transforms['valid'])\ntest_loader  = DataLoader(test_dataset, batch_size=CFG.valid_bs, \n                          num_workers=4, shuffle=False, pin_memory=False)\nmodel_paths  = glob(f'{CKPT_DIR}/best_epoch*.bin')\npred_strings, pred_ids, pred_classes, imgs, msks = infer(model_paths, test_loader)","metadata":{"execution":{"iopub.status.busy":"2022-04-26T09:11:03.507508Z","iopub.execute_input":"2022-04-26T09:11:03.50776Z","iopub.status.idle":"2022-04-26T09:11:53.025729Z","shell.execute_reply.started":"2022-04-26T09:11:03.507731Z","shell.execute_reply":"2022-04-26T09:11:53.024937Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 📈 Visualization","metadata":{}},{"cell_type":"code","source":"for img, msk in zip(imgs[0][:5], msks[0][:5]):\n    plt.figure(figsize=(12, 7))\n    plt.subplot(1, 3, 1); plt.imshow(img, cmap='bone');\n    plt.axis('OFF'); plt.title('image')\n    plt.subplot(1, 3, 2); plt.imshow(msk*255); plt.axis('OFF'); plt.title('mask')\n    plt.subplot(1, 3, 3); plt.imshow(img, cmap='bone'); plt.imshow(msk*255, alpha=0.4);\n    plt.axis('OFF'); plt.title('overlay')\n    plt.tight_layout()\n    plt.show()","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-04-26T09:12:59.30086Z","iopub.execute_input":"2022-04-26T09:12:59.301132Z","iopub.status.idle":"2022-04-26T09:13:01.160858Z","shell.execute_reply.started":"2022-04-26T09:12:59.3011Z","shell.execute_reply":"2022-04-26T09:13:01.160165Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del imgs, msks\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-04-26T09:13:03.42923Z","iopub.execute_input":"2022-04-26T09:13:03.42951Z","iopub.status.idle":"2022-04-26T09:13:03.684973Z","shell.execute_reply.started":"2022-04-26T09:13:03.42948Z","shell.execute_reply":"2022-04-26T09:13:03.684298Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 📝 Submission","metadata":{}},{"cell_type":"code","source":"pred_df = pd.DataFrame({\n    \"id\":pred_ids,\n    \"class\":pred_classes,\n    \"predicted\":pred_strings\n})\nif not debug:\n    sub_df = pd.read_csv('../input/uw-madison-gi-tract-image-segmentation/sample_submission.csv')\n    del sub_df['predicted']\nelse:\n    sub_df = pd.read_csv('../input/uw-madison-gi-tract-image-segmentation/train.csv')[:1000*3]\n    del sub_df['segmentation']\n    \nsub_df = sub_df.merge(pred_df, on=['id','class'])\nsub_df.to_csv('submission.csv',index=False)\ndisplay(sub_df.head(5))","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-04-26T09:13:04.651599Z","iopub.execute_input":"2022-04-26T09:13:04.651861Z","iopub.status.idle":"2022-04-26T09:13:04.980926Z","shell.execute_reply.started":"2022-04-26T09:13:04.65183Z","shell.execute_reply":"2022-04-26T09:13:04.980118Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}