{"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":"# Install and Import Libraries","metadata":{}},{"cell_type":"markdown","source":"Install necessary packages","metadata":{}},{"cell_type":"code","source":"!pip install git+https://github.com/qubvel/segmentation_models.pytorch\n!pip install scikit-learn\n# !pip install scikit-image","metadata":{"execution":{"iopub.status.busy":"2022-07-12T07:50:04.411605Z","iopub.execute_input":"2022-07-12T07:50:04.412013Z","iopub.status.idle":"2022-07-12T07:50:39.207264Z","shell.execute_reply.started":"2022-07-12T07:50:04.411884Z","shell.execute_reply":"2022-07-12T07:50:39.206434Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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 collections\nimport gc\nfrom IPython import display as ipd\nfrom pathlib import Path\nfrom typing import List\nfrom typing import cast\n\n# visualization\nimport cv2\nimport matplotlib.pyplot as plt\nfrom matplotlib.ticker import MaxNLocator\nfrom matplotlib.patches import Rectangle\n\n# Sklearn\nfrom sklearn.model_selection import StratifiedKFold, KFold, StratifiedGroupKFold\nfrom sklearn.model_selection import train_test_split\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\nfrom torch import Tensor, einsum\n\nimport timm\n\n# Scipy\nfrom scipy.ndimage.morphology import distance_transform_edt as edt\nfrom scipy.ndimage import convolve\n# from skimage import metrics\n\n# Albumentations for augmentations\nimport albumentations as albu\nfrom albumentations.pytorch import ToTensorV2\n\nimport rasterio\nfrom joblib import Parallel, delayed\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":{"execution":{"iopub.status.busy":"2022-07-12T07:50:39.210945Z","iopub.execute_input":"2022-07-12T07:50:39.211167Z","iopub.status.idle":"2022-07-12T07:50:44.119936Z","shell.execute_reply.started":"2022-07-12T07:50:39.211142Z","shell.execute_reply":"2022-07-12T07:50:44.119161Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Configuration","metadata":{}},{"cell_type":"code","source":"class CFG:\n    seed          = 42\n    debug         = False # set debug=False for Full Training\n    exp_name      = '2.5D (steps: -2, 0, +2). New scaler, No pad, pseudo labels dataset. efficientnet-b4. Data Resampling.'\n    comment       = 'Unet-efficientnet-b4-384x384'\n    model_name    = 'Unet'\n    backbone      = 'efficientnet-b5'\n    train_bs      = 8\n    valid_bs      = train_bs\n    pad_img_size  = [320, 384]\n    img_size      = [384, 384]\n    epochs        = 15\n    lr            = 2e-3\n    scheduler     = 'CosineAnnealingLR'\n    min_lr        = 1e-6\n    T_max         = int(30000/train_bs*epochs)+10\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    folds         = [0]\n    num_classes   = 3\n    device        = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")","metadata":{"execution":{"iopub.status.busy":"2022-07-12T07:50:44.121513Z","iopub.execute_input":"2022-07-12T07:50:44.121736Z","iopub.status.idle":"2022-07-12T07:50:44.187524Z","shell.execute_reply.started":"2022-07-12T07:50:44.121703Z","shell.execute_reply":"2022-07-12T07:50:44.185616Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if torch.cuda.is_available():\n    torch.cuda.empty_cache()\n    print(\"> CUDA CACHE IS CLEANED\")\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-07-12T07:50:44.189664Z","iopub.execute_input":"2022-07-12T07:50:44.190150Z","iopub.status.idle":"2022-07-12T07:50:44.358069Z","shell.execute_reply.started":"2022-07-12T07:50:44.190102Z","shell.execute_reply":"2022-07-12T07:50:44.357314Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(seed = 42):\n    '''Sets the seed of the entire notebook so results are the same every time we run'''\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('> EVERTHING IS SEEDED')\n    \nseed_everything(CFG.seed)","metadata":{"execution":{"iopub.status.busy":"2022-07-12T07:50:44.359342Z","iopub.execute_input":"2022-07-12T07:50:44.361102Z","iopub.status.idle":"2022-07-12T07:50:44.370627Z","shell.execute_reply.started":"2022-07-12T07:50:44.361062Z","shell.execute_reply":"2022-07-12T07:50:44.369824Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Loading","metadata":{}},{"cell_type":"code","source":"IMAGE_PATH  = '/kaggle/input/uw-madison-gi-tract-image-segmentation'\nMASK_PATH = '/kaggle/input/uwmgi-mask-dataset'\nIMAGE_2DOT5_PATH = '/kaggle/input/uwmgi-25d-stride2-2-0-2-unchanged-data'","metadata":{"execution":{"iopub.status.busy":"2022-07-12T07:50:44.371709Z","iopub.execute_input":"2022-07-12T07:50:44.372071Z","iopub.status.idle":"2022-07-12T07:50:44.377306Z","shell.execute_reply.started":"2022-07-12T07:50:44.372035Z","shell.execute_reply":"2022-07-12T07:50:44.376607Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"path_df = pd.DataFrame(glob('/kaggle/input/uwmgi-25d-stride2-2-0-2-unchanged-data/images/images/*'), columns=['image_path'])\npath_df['mask_path'] = path_df.image_path.str.replace('image','mask')\npath_df['id'] = path_df.image_path.map(lambda x: x.split('/')[-1].replace('.npy',''))\npath_df.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-12T07:50:44.378427Z","iopub.execute_input":"2022-07-12T07:50:44.378793Z","iopub.status.idle":"2022-07-12T07:50:45.366575Z","shell.execute_reply.started":"2022-07-12T07:50:44.378757Z","shell.execute_reply":"2022-07-12T07:50:45.365818Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv('../input/uwmgi-mask-dataset/train.csv')\ndf['segmentation'] = df.segmentation.fillna('')\ndf['rle_len'] = df.segmentation.map(len) # length of each rle mask\n\ndf2 = df.groupby(['id'])['segmentation'].agg(list).to_frame().reset_index() # rle list of each id\ndf2 = df2.merge(df.groupby(['id'])['rle_len'].agg(sum).to_frame().reset_index()) # total length of all rles of each id\n\ndf = df.drop(columns=['segmentation', 'class', 'rle_len'])\ndf = df.groupby(['id']).head(1).reset_index(drop=True)\ndf = df.merge(df2, on=['id'])\ndf['empty'] = (df.rle_len==0) # empty masks\n\ndf = df.drop(columns=['image_path','mask_path'])\ndf = df.merge(path_df, on=['id'])\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-12T07:50:45.368006Z","iopub.execute_input":"2022-07-12T07:50:45.368261Z","iopub.status.idle":"2022-07-12T07:50:46.992583Z","shell.execute_reply.started":"2022-07-12T07:50:45.368223Z","shell.execute_reply":"2022-07-12T07:50:46.991832Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Clean Data","metadata":{}},{"cell_type":"code","source":"print(\"Size before cleaning:\", df.shape)\n\n# # Masks for case7_day0 and case81_day30 are incorrect \n# # ref: https://www.kaggle.com/competitions/uw-madison-gi-tract-image-segmentation/discussion/319963\n# df = df.drop(df[(df['case'] == 7) & (df['day'] == 0)].index)\n# df = df.drop(df[(df['case'] == 81) & (df['day'] == 30)].index)\n\n# df = df.reset_index(drop=True)\n# print(\"Size after cleaning:\", df.shape)","metadata":{"execution":{"iopub.status.busy":"2022-07-12T07:50:46.993796Z","iopub.execute_input":"2022-07-12T07:50:46.994077Z","iopub.status.idle":"2022-07-12T07:50:47.000296Z","shell.execute_reply.started":"2022-07-12T07:50:46.994040Z","shell.execute_reply":"2022-07-12T07:50:46.999440Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df['empty'].value_counts().plot.bar()","metadata":{"execution":{"iopub.status.busy":"2022-07-12T07:50:47.003825Z","iopub.execute_input":"2022-07-12T07:50:47.004141Z","iopub.status.idle":"2022-07-12T07:50:49.718089Z","shell.execute_reply.started":"2022-07-12T07:50:47.004093Z","shell.execute_reply":"2022-07-12T07:50:49.717338Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data Resampling preparation","metadata":{}},{"cell_type":"code","source":"resampling_window = 10\n\ndf['close_to_not_empty'] = False\n\nfor index, row in df.iterrows():\n    if row['empty'] is False:\n        df.loc[index,'close_to_not_empty'] = True\n        continue\n        \n    neighboring_rows = df[((df['case'] == row['case']) & \\\n                           (df['day'] == row['day']) & \\\n                           (df['slice'] >= (row['slice'] - resampling_window)) &\\\n                           (df['slice'] <= (row['slice'] + resampling_window)))]\n\n    if len(neighboring_rows) > 0 and np.any(neighboring_rows['empty'] == False):\n        df.loc[index,'close_to_not_empty'] = True\n        continue\n\ndisplay(df.tail())\ndf['close_to_not_empty'].value_counts().plot.bar()","metadata":{"execution":{"iopub.status.busy":"2022-07-12T07:50:49.719275Z","iopub.execute_input":"2022-07-12T07:50:49.719846Z","iopub.status.idle":"2022-07-12T07:51:25.645394Z","shell.execute_reply.started":"2022-07-12T07:50:49.719806Z","shell.execute_reply":"2022-07-12T07:51:25.644560Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helper Functions","metadata":{}},{"cell_type":"markdown","source":"## RLE","metadata":{}},{"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":{"execution":{"iopub.status.busy":"2022-07-12T07:51:25.646773Z","iopub.execute_input":"2022-07-12T07:51:25.647039Z","iopub.status.idle":"2022-07-12T07:51:25.655175Z","shell.execute_reply.started":"2022-07-12T07:51:25.647003Z","shell.execute_reply":"2022-07-12T07:51:25.654137Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Mask","metadata":{}},{"cell_type":"code","source":"def id2mask(id_):\n    idf = df[df['id']==id_]\n    wh = idf[['height','width']].iloc[0]\n    shape = (wh.height, wh.width, 3)\n    mask = np.zeros(shape, dtype=np.uint8)\n    for i, class_ in enumerate(['large_bowel', 'small_bowel', 'stomach']):\n        cdf = idf[idf['class']==class_]\n        rle = cdf.segmentation.squeeze()\n        if len(cdf) and not pd.isna(rle):\n            mask[..., i] = rle_decode(rle, shape[:2])\n    return mask\n\ndef rgb2gray(mask):\n    pad_mask = np.pad(mask, pad_width=[(0,0),(0,0),(1,0)])\n    gray_mask = pad_mask.argmax(-1)\n    return gray_mask\n\ndef gray2rgb(mask):\n    rgb_mask = tf.keras.utils.to_categorical(mask, num_classes=4)\n    return rgb_mask[..., 1:].astype(mask.dtype)","metadata":{"execution":{"iopub.status.busy":"2022-07-12T07:51:25.656493Z","iopub.execute_input":"2022-07-12T07:51:25.657028Z","iopub.status.idle":"2022-07-12T07:51:25.670032Z","shell.execute_reply.started":"2022-07-12T07:51:25.656991Z","shell.execute_reply":"2022-07-12T07:51:25.669286Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Image","metadata":{}},{"cell_type":"code","source":"# def add_pad(img):\n#     img_shape = np.array(img.shape[:2])\n#     max_img_shape = np.max(img_shape)\n#     resize = np.array([max_img_shape, max_img_shape])\n#     if np.any(img_shape != resize):\n#         diff = resize - img_shape\n#         pad0, pad1 = diff[0], diff[1]\n#         pady = [pad0//2, pad0//2 + pad0%2]\n#         padx = [pad1//2, pad1//2 + pad1%2]\n#         img = np.pad(img, [pady, padx, [0,0]])\n#         img = img.reshape((*resize, 3))\n#     return img\n\ndef load_img(path):\n    img = np.load(path)\n#     img = add_pad(img)\n\n    img = img.astype('float32') # original is uint16\n    mx = np.max(img)\n    if mx:\n        q13 = np.percentile(img[img > 0], [25, 75])\n        max_div = np.max(img[img < (q13[1] + 3*(q13[1]-q13[0]))])\n        if max_div > 0:\n            img = np.clip(img/max_div, 0, 1)\n        else:\n            img = np.clip(img/mx, 0, 1)\n    return img\n\ndef load_msk(path):\n    msk = np.load(path)\n#     msk = add_pad(msk)\n\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    plt.imshow(img, cmap='bone')\n    \n    if mask is not None:\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), \n                                                             (0.0,0.667,0.0), \n                                                             (0.0,0.0,0.667)]]\n        labels = [\"Large Bowel\", \"Small Bowel\", \"Stomach\"]\n        plt.legend(handles,labels)\n    plt.axis('off')","metadata":{"execution":{"iopub.status.busy":"2022-07-12T07:51:25.671432Z","iopub.execute_input":"2022-07-12T07:51:25.672071Z","iopub.status.idle":"2022-07-12T07:51:25.683744Z","shell.execute_reply.started":"2022-07-12T07:51:25.672031Z","shell.execute_reply":"2022-07-12T07:51:25.682967Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Plot","metadata":{}},{"cell_type":"code","source":"def plot_batch(imgs, msks, size=3):\n    plt.figure(figsize=(5*5, 5))\n    for idx in range(size):\n        plt.subplot(1, 5, idx+1)\n        img = imgs[idx,].permute((1, 2, 0)).numpy()*255.0\n        img = img.astype('uint8')\n        msk = msks[idx,].permute((1, 2, 0)).numpy()*255.0\n        show_img(img, msk)\n    plt.tight_layout()\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-07-12T07:51:25.684647Z","iopub.execute_input":"2022-07-12T07:51:25.685631Z","iopub.status.idle":"2022-07-12T07:51:25.696729Z","shell.execute_reply.started":"2022-07-12T07:51:25.685595Z","shell.execute_reply":"2022-07-12T07:51:25.695949Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_metrics(history, epochs=CFG.epochs):\n    x = np.linspace(1, epochs, epochs).astype(int)\n    \n    f = plt.figure(figsize=(15, 40))\n    for i, (key, value) in enumerate(history.items()):\n        ax = f.add_subplot(5, 1, i+1)\n        ax.plot(x, value)\n        ax.set_title(key)\n        ax.xaxis.set_major_locator(MaxNLocator(integer=True))\n        \n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-07-12T07:51:25.699472Z","iopub.execute_input":"2022-07-12T07:51:25.700021Z","iopub.status.idle":"2022-07-12T07:51:25.707862Z","shell.execute_reply.started":"2022-07-12T07:51:25.699985Z","shell.execute_reply":"2022-07-12T07:51:25.707170Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Create Training Folds","metadata":{}},{"cell_type":"code","source":"skf = StratifiedGroupKFold(n_splits=CFG.n_fold, shuffle=True, random_state=CFG.seed)\nfor fold, (train_idx, val_idx) in enumerate(skf.split(df, df['empty'], groups = df[\"case\"])):\n    df.loc[val_idx, 'fold'] = fold\ndisplay(df.groupby(['fold','empty'])['id'].count())","metadata":{"execution":{"iopub.status.busy":"2022-07-12T07:51:25.708682Z","iopub.execute_input":"2022-07-12T07:51:25.711129Z","iopub.status.idle":"2022-07-12T07:51:25.876320Z","shell.execute_reply.started":"2022-07-12T07:51:25.711038Z","shell.execute_reply":"2022-07-12T07:51:25.875388Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Augmentations, Dataset, Dataloader","metadata":{}},{"cell_type":"markdown","source":"## Augmentations","metadata":{}},{"cell_type":"code","source":"def pre_transforms(image_size=CFG.img_size):\n    return [\n        albu.Resize(*image_size, interpolation=cv2.INTER_LINEAR, p=1)\n    ]\n\ndef flip_transforms():\n    return [\n#         albu.VerticalFlip(p=0.5)\n        albu.HorizontalFlip(p=0.5)\n        \n        #torch.sigmoid(0.5*model(images) + 0.5*model(images.flip(3)).flip(3))\n    ]\n\n# def resize_transforms(image_size=CFG.img_size):\n#     min_size = int(image_size[0] * 0.75)\n#     return albu.OneOf([\n#         albu.RandomSizedCrop(min_max_height=(min_size, image_size[0]), \n#                              height=image_size[0], \n#                              width=image_size[1], \n#                              p=0.5),\n#         albu.PadIfNeeded(min_height=image_size[0], \n#                          min_width=image_size[1], \n#                          p=0.5)\n#     ],p=1)\n\ndef spatial_transforms():\n    distorsion = [      \n        albu.GridDistortion(always_apply=False,\n                            p=0.25, \n                            num_steps=6, \n                            distort_limit=(-0.06, 0.06)),\n        albu.CoarseDropout(max_holes=8, \n                           max_height=CFG.img_size[0]//20, \n                           max_width=CFG.img_size[1]//20,\n                           min_holes=5, fill_value=0, \n                           mask_fill_value=0, \n                           p=0.25),\n        ]\n\n    return distorsion\n\ndef pixel_transforms():\n    return [\n        albu.RandomBrightnessContrast(always_apply=False, p=0.25, \n                                      brightness_limit=(-0.05, 0.05), \n                                      contrast_limit=(-0.05, 0.05), \n                                      brightness_by_max=True)\n    ]\n\ndef rotate_transforms():\n    return [\n        albu.ShiftScaleRotate(shift_limit=0.06, scale_limit=0.05, rotate_limit=10, p=0.5)\n    ]\n\ndef post_transforms():\n    # Convert image to torch.Tensor format\n    return [ToTensorV2(transpose_mask=True)]\n\ndef compose(transforms_to_compose):\n    # combine all augmentations into single pipeline\n    result = albu.Compose([\n        item for sublist in transforms_to_compose for item in sublist\n    ])\n    return result","metadata":{"execution":{"iopub.status.busy":"2022-07-12T07:51:39.538507Z","iopub.execute_input":"2022-07-12T07:51:39.538765Z","iopub.status.idle":"2022-07-12T07:51:39.550019Z","shell.execute_reply.started":"2022-07-12T07:51:39.538734Z","shell.execute_reply":"2022-07-12T07:51:39.549312Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_transforms = compose([\n    pre_transforms(),\n    flip_transforms(),\n    spatial_transforms(),\n    pixel_transforms(),\n    rotate_transforms(),\n    post_transforms()\n])\n\nvalid_transforms = compose([\n    pre_transforms(), \n    post_transforms()\n])","metadata":{"execution":{"iopub.status.busy":"2022-07-12T07:51:44.663365Z","iopub.execute_input":"2022-07-12T07:51:44.663922Z","iopub.status.idle":"2022-07-12T07:51:45.001202Z","shell.execute_reply.started":"2022-07-12T07:51:44.663882Z","shell.execute_reply":"2022-07-12T07:51:44.999994Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset Builder","metadata":{}},{"cell_type":"code","source":"class DatasetBuilder(torch.utils.data.Dataset):\n    def __init__(self, df, label=True, transforms=None):\n        self.df         = df\n        self.label      = label\n        self.img_paths  = df['image_path'].tolist()\n        self.msk_paths  = df['mask_path'].tolist()\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        img = []\n        img = load_img(img_path)\n        \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            return img, msk\n\n        else:\n            if self.transforms:\n                data = self.transforms(image=img)\n                img  = data['image']            \n            return img","metadata":{"execution":{"iopub.status.busy":"2022-07-07T14:06:39.070143Z","iopub.execute_input":"2022-07-07T14:06:39.070362Z","iopub.status.idle":"2022-07-07T14:06:39.084144Z","shell.execute_reply.started":"2022-07-07T14:06:39.070335Z","shell.execute_reply":"2022-07-07T14:06:39.083166Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataloader","metadata":{}},{"cell_type":"code","source":"def prepare_loader(df, transforms, fold=0, loader_type=\"train\", resample=True, debug=False):\n    if loader_type == \"train\":\n        valid = False\n        query = \"fold!=@fold\"\n        bs = CFG.train_bs\n        shuffle = True\n        drop_last = True\n    elif loader_type == \"valid\":\n        valid = True\n        query = \"fold==@fold\"\n        bs = CFG.valid_bs\n        shuffle = False\n        drop_last = False \n    else:\n        raise ValueError\n        \n    if resample:\n        df = df.drop(df[(df['fold'] != fold) & (df['close_to_not_empty'] == False)].index)\n        df = df.reset_index(drop=True)\n    \n    df = df.query(query).reset_index(drop=True) \n\n    if valid:\n        df = df.drop(df[(df['case'] == 7) & (df['day'] == 0)].index)\n        df = df.drop(df[(df['case'] == 81) & (df['day'] == 30)].index)\n        df = df.reset_index(drop=True)\n    \n    if debug:\n        df = df.head(32*5).query(\"empty==False\")\n    \n    dataset = DatasetBuilder(df, transforms=transforms)\n    loader = DataLoader(dataset, batch_size=bs if not debug else 20, \n                        num_workers=4, shuffle=shuffle, pin_memory=True, drop_last=drop_last)\n    return loader","metadata":{"execution":{"iopub.status.busy":"2022-07-07T14:06:39.08633Z","iopub.execute_input":"2022-07-07T14:06:39.086532Z","iopub.status.idle":"2022-07-07T14:06:39.100326Z","shell.execute_reply.started":"2022-07-07T14:06:39.086508Z","shell.execute_reply":"2022-07-07T14:06:39.099356Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loader = prepare_loader(df=df.copy(), transforms=train_transforms, fold=0, loader_type=\"train\", \n                              debug=CFG.debug)\nvalid_loader = prepare_loader(df=df.copy(), transforms=valid_transforms, fold=0, loader_type=\"valid\", \n                              debug=CFG.debug)","metadata":{"execution":{"iopub.status.busy":"2022-07-07T14:06:39.101973Z","iopub.execute_input":"2022-07-07T14:06:39.102727Z","iopub.status.idle":"2022-07-07T14:06:39.168859Z","shell.execute_reply.started":"2022-07-07T14:06:39.102663Z","shell.execute_reply":"2022-07-07T14:06:39.168074Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Test that the data is loaded correctly","metadata":{}},{"cell_type":"code","source":"for i in range(5):\n    start = time.time()\n    imgs, msks = next(iter(train_loader))\n    end = time.time()\n    time_elapsed = end - start\n    print(\"One batch time: \", time_elapsed)\n    \n    plot_batch(imgs, msks, size=5)\n\nimgs.size(), msks.size()","metadata":{"execution":{"iopub.status.busy":"2022-07-07T14:06:39.170224Z","iopub.execute_input":"2022-07-07T14:06:39.170797Z","iopub.status.idle":"2022-07-07T14:06:52.406364Z","shell.execute_reply.started":"2022-07-07T14:06:39.170756Z","shell.execute_reply":"2022-07-07T14:06:52.405331Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Garbage Collector","metadata":{}},{"cell_type":"code","source":"import gc\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-07-07T14:06:52.408164Z","iopub.execute_input":"2022-07-07T14:06:52.409012Z","iopub.status.idle":"2022-07-07T14:06:52.771292Z","shell.execute_reply.started":"2022-07-07T14:06:52.408973Z","shell.execute_reply":"2022-07-07T14:06:52.770259Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model Setup","metadata":{}},{"cell_type":"markdown","source":"## Model initialization","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=\"imagenet\",     # 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 save_model(epoch, model, optimizer, loss, val_loss, path):\n    torch.save({\n        'epoch': epoch,\n        'model_state_dict': model.state_dict(),\n        'optimizer_state_dict': optimizer.state_dict(),\n        'loss': loss,\n        'val_loss': val_loss,\n    }, path)\n\ndef load_model(path, optimizer=None):\n    model = build_model()\n\n#     checkpoint = torch.load(path, map_location=CFG.device)\n    checkpoint = torch.load(path)\n    epoch = checkpoint['epoch']\n    loss = checkpoint['loss']\n    model.load_state_dict(checkpoint['model_state_dict'])\n    if optimizer is not None:\n        optimizer.load_state_dict(checkpoint['optimizer_state_dict'])\n\n    model.eval()\n    return model, optimizer, epoch, loss","metadata":{"execution":{"iopub.status.busy":"2022-07-07T14:06:52.773145Z","iopub.execute_input":"2022-07-07T14:06:52.773579Z","iopub.status.idle":"2022-07-07T14:06:52.786263Z","shell.execute_reply.started":"2022-07-07T14:06:52.773477Z","shell.execute_reply":"2022-07-07T14:06:52.785355Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Loss Function","metadata":{}},{"cell_type":"code","source":"# def hausdorff_format(tensor: torch.Tensor, to_prob: bool) -> torch.Tensor:\n#     if to_prob:\n#         tensor = nn.Sigmoid()(tensor)\n#     tensor = tensor.permute((0,2,3,1)).to(torch.float32)\n#     return tensor[:, None, :, :, :]","metadata":{"execution":{"iopub.status.busy":"2022-07-07T14:06:52.787413Z","iopub.execute_input":"2022-07-07T14:06:52.788253Z","iopub.status.idle":"2022-07-07T14:06:52.801626Z","shell.execute_reply.started":"2022-07-07T14:06:52.788211Z","shell.execute_reply":"2022-07-07T14:06:52.800827Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # ref: https://github.com/PatRyg99/HausdorffLoss/blob/9f580acd421af648e74b45d46555ccb7a876c27c/hausdorff_loss.py\n\n# class HausdorffERLoss(nn.Module):\n#     \"\"\"Binary Hausdorff loss based on morphological erosion\"\"\"\n\n#     def __init__(self, alpha=2.0, erosions=10, **kwargs):\n#         super(HausdorffERLoss, self).__init__()\n#         self.alpha = alpha\n#         self.erosions = erosions\n#         self.prepare_kernels()\n\n#     def prepare_kernels(self):\n#         cross = np.array([cv2.getStructuringElement(cv2.MORPH_CROSS, (3, 3))])\n#         bound = np.array([[[0, 0, 0], [0, 1, 0], [0, 0, 0]]])\n\n#         self.kernel2D = cross * 0.2\n#         self.kernel3D = np.array([bound, cross, bound]) * (1 / 7)\n\n#     @torch.no_grad()\n#     def perform_erosion(self, pred: np.ndarray, target: np.ndarray, thr=0.5) -> np.ndarray:\n#         bound = (pred - target) ** 2\n\n#         if bound.ndim == 5:\n#             kernel = self.kernel3D\n#         elif bound.ndim == 4:\n#             kernel = self.kernel2D\n#         else:\n#             raise ValueError(f\"Dimension {bound.ndim} is nor supported.\")\n\n#         eroted = np.zeros_like(bound)\n\n#         for batch in range(len(bound)):\n#             for k in range(self.erosions):\n\n#                 # compute convolution with kernel\n#                 dilation = convolve(bound[batch], kernel, mode=\"constant\", cval=0.0)\n\n#                 # apply soft thresholding at 0.5 and normalize\n#                 erosion = dilation - thr\n#                 erosion[erosion < 0] = 0\n\n#                 if erosion.ptp() != 0:\n#                     erosion = (erosion - erosion.min()) / erosion.ptp()\n\n#                 # save erosion and add to loss\n#                 bound[batch] = erosion\n#                 eroted[batch] += erosion * (k + 1) ** self.alpha\n#         return eroted\n\n#     def forward(self, pred: torch.Tensor, target: torch.Tensor) -> torch.Tensor:\n#         \"\"\"\n#         Uses one binary channel: 1 - fg, 0 - bg\n#         pred: (b, 1, x, y, z) or (b, 1, x, y)\n#         target: (b, 1, x, y, z) or (b, 1, x, y)\n#         \"\"\"\n#         assert pred.dim() == 4 or pred.dim() == 5, \"Only 2D and 3D supported\"\n#         assert (\n#             pred.dim() == target.dim()\n#         ), \"Prediction and target need to be of same dimension\"\n        \n#         losses = []\n#         for i in range(CFG.num_classes):\n#             pred_c = hausdorff_format(pred, to_prob=True).cpu().detach().numpy()[..., i]\n#             target_c = hausdorff_format(target, to_prob=False).cpu().detach().numpy()[..., i]\n            \n#             eroted = torch.from_numpy(\n#                 self.perform_erosion(pred_c, target_c)\n#             ).float()\n\n#             loss = eroted.mean()\n#             losses.append(loss.item())\n\n#         losses = torch.from_numpy(np.array(losses)).to(CFG.device).float()\n#         return losses.mean()","metadata":{"execution":{"iopub.status.busy":"2022-07-07T14:06:52.803204Z","iopub.execute_input":"2022-07-07T14:06:52.803697Z","iopub.status.idle":"2022-07-07T14:06:52.8145Z","shell.execute_reply.started":"2022-07-07T14:06:52.80366Z","shell.execute_reply":"2022-07-07T14:06:52.813596Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def hausdorff_coef(y_true, y_pred, thr=0.5):\n#     y_true = y_true.permute((0,2,3,1)).to(torch.uint8).cpu().detach().numpy()\n#     y_pred = (y_pred.permute((0,2,3,1))>=thr).to(torch.uint8).cpu().detach().numpy()\n#     y_true_bool = y_true > 0\n#     y_pred_bool = y_pred > 0\n\n#     hds = []\n#     for msk, pred in zip(y_true_bool, y_pred_bool):\n#         img_hds = [metrics.hausdorff_distance(msk[..., i], pred[..., i]) for i in range(CFG.num_classes)]\n#         hds.append(np.array(img_hds).mean())\n#     hausdorff = np.asarray(hds).mean()\n#     return hausdorff\n\ndef dice_coef(y_true, y_pred, thr=0.5, dim=(2,3), epsilon=0.001):\n    y_true = y_true.to(torch.float32)\n    y_pred = (y_pred>thr).to(torch.float32)\n    inter = (y_true*y_pred).sum(dim=dim)\n    den = y_true.sum(dim=dim) + y_pred.sum(dim=dim)\n    dice = ((2*inter+epsilon)/(den+epsilon)).mean(dim=(1,0))\n    return dice\n\ndef iou_coef(y_true, y_pred, thr=0.5, dim=(2,3), epsilon=0.001):\n    y_true = y_true.to(torch.float32)\n    y_pred = (y_pred>thr).to(torch.float32)\n    inter = (y_true*y_pred).sum(dim=dim)\n    union = (y_true + y_pred - y_true*y_pred).sum(dim=dim)\n    iou = ((inter+epsilon)/(union+epsilon)).mean(dim=(1,0))\n    return iou","metadata":{"execution":{"iopub.status.busy":"2022-07-07T14:06:52.820613Z","iopub.execute_input":"2022-07-07T14:06:52.821321Z","iopub.status.idle":"2022-07-07T14:06:52.831702Z","shell.execute_reply.started":"2022-07-07T14:06:52.821275Z","shell.execute_reply":"2022-07-07T14:06:52.830652Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# TverskyLoss alpha == beta == 0.5, this loss becomes equal DiceLoss\nTverskyLoss           = smp.losses.TverskyLoss(mode='multilabel', log_loss=False)\nBCELoss               = smp.losses.SoftBCEWithLogitsLoss()\n\ndef criterion(y_pred, y_true):\n    return 0.7*TverskyLoss(y_pred, y_true) + 0.3*BCELoss(y_pred, y_true)","metadata":{"execution":{"iopub.status.busy":"2022-07-07T14:06:52.833415Z","iopub.execute_input":"2022-07-07T14:06:52.833694Z","iopub.status.idle":"2022-07-07T14:06:52.84261Z","shell.execute_reply.started":"2022-07-07T14:06:52.833656Z","shell.execute_reply":"2022-07-07T14:06:52.841812Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Optimizer and Scheduler","metadata":{}},{"cell_type":"code","source":"def setup_scheduler(optimizer):\n    if CFG.scheduler == 'CosineAnnealingLR':\n        scheduler = lr_scheduler.CosineAnnealingLR(optimizer,T_max=CFG.T_max, \n                                                   eta_min=CFG.min_lr)\n    \n    elif CFG.scheduler == 'CosineAnnealingWarmRestarts':\n        scheduler = lr_scheduler.CosineAnnealingWarmRestarts(optimizer,T_0=CFG.T_0, \n                                                             eta_min=CFG.min_lr)\n    \n    elif CFG.scheduler == 'ReduceLROnPlateau':\n        scheduler = lr_scheduler.ReduceLROnPlateau(optimizer,\n                                                   mode='min',\n                                                   factor=0.1,\n                                                   patience=7,\n                                                   threshold=0.0001,\n                                                   min_lr=CFG.min_lr,)\n    \n    elif CFG.scheduer == 'ExponentialLR':\n        scheduler = lr_scheduler.ExponentialLR(optimizer, gamma=0.85)\n    \n    elif CFG.scheduler == None:\n        return None\n        \n    return scheduler","metadata":{"execution":{"iopub.status.busy":"2022-07-07T14:06:52.843895Z","iopub.execute_input":"2022-07-07T14:06:52.844789Z","iopub.status.idle":"2022-07-07T14:06:52.85838Z","shell.execute_reply.started":"2022-07-07T14:06:52.844654Z","shell.execute_reply":"2022-07-07T14:06:52.857429Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = build_model()\noptimizer = optim.Adamax(model.parameters(), \n                         lr=CFG.lr,\n                         betas=(0.9, 0.999), \n                         eps=1e-08,\n                         weight_decay=CFG.wd)\nscheduler = setup_scheduler(optimizer)","metadata":{"execution":{"iopub.status.busy":"2022-07-07T14:06:52.861543Z","iopub.execute_input":"2022-07-07T14:06:52.862219Z","iopub.status.idle":"2022-07-07T14:06:53.591311Z","shell.execute_reply.started":"2022-07-07T14:06:52.862184Z","shell.execute_reply":"2022-07-07T14:06:53.590393Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model Training","metadata":{}},{"cell_type":"markdown","source":"## Model training and validation functions","metadata":{}},{"cell_type":"code","source":"def train_epoch(model, optimizer, scheduler, dataloader, device, epoch):\n    model.train()\n    scaler = amp.GradScaler()\n    \n    dataset_size = 0\n    running_loss = 0.0\n    \n    pbar = tqdm(enumerate(dataloader), total=len(dataloader), desc='Train ')\n    for step, (images, masks) in pbar:         \n        images = images.to(device, dtype=torch.float)\n        masks  = masks.to(device, dtype=torch.float)\n        \n        batch_size = images.size(0)\n        \n        with amp.autocast(enabled=True):\n            y_pred = model(images)\n            loss   = criterion(y_pred, masks)\n            loss   = loss / CFG.n_accumulate\n            \n        scaler.scale(loss).backward()\n    \n        if (step + 1) % CFG.n_accumulate == 0:\n            scaler.step(optimizer)\n            scaler.update()\n\n            # zero the parameter gradients\n            optimizer.zero_grad()\n\n            if scheduler is not None:\n                scheduler.step()\n                \n        running_loss += (loss.item() * batch_size)\n        dataset_size += batch_size\n        \n        epoch_loss = running_loss / dataset_size\n        \n        mem = torch.cuda.memory_reserved() / 1E9 if torch.cuda.is_available() else 0\n        current_lr = optimizer.param_groups[0]['lr']\n        pbar.set_postfix(train_loss=f'{epoch_loss:0.4f}',\n                         lr=f'{current_lr:0.5f}',\n                         gpu_mem=f'{mem:0.2f} GB')\n        torch.cuda.empty_cache()\n        gc.collect()\n    \n    return epoch_loss","metadata":{"execution":{"iopub.status.busy":"2022-07-07T14:06:53.592959Z","iopub.execute_input":"2022-07-07T14:06:53.593253Z","iopub.status.idle":"2022-07-07T14:06:53.606139Z","shell.execute_reply.started":"2022-07-07T14:06:53.593214Z","shell.execute_reply":"2022-07-07T14:06:53.605148Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@torch.no_grad()\ndef valid_epoch(model, dataloader, device, epoch):\n    model.eval()\n    \n    dataset_size = 0\n    running_loss = 0.0\n    \n    val_scores = []\n    \n    pbar = tqdm(enumerate(dataloader), total=len(dataloader), desc='Valid ')\n    for step, (images, masks) in pbar:        \n        images  = images.to(device, dtype=torch.float)\n        masks   = masks.to(device, dtype=torch.float)\n        \n        batch_size = images.size(0)\n        \n        y_pred  = model(images)\n        loss    = criterion(y_pred, masks)\n        \n        running_loss += (loss.item() * batch_size)\n        dataset_size += batch_size\n        \n        epoch_loss = running_loss / dataset_size\n        \n        y_pred = nn.Sigmoid()(y_pred)\n        val_dice = dice_coef(masks, y_pred).cpu().detach().numpy()\n        val_jaccard = iou_coef(masks, y_pred).cpu().detach().numpy()\n        val_scores.append([val_dice, val_jaccard])\n        \n        mem = torch.cuda.memory_reserved() / 1E9 if torch.cuda.is_available() else 0\n        current_lr = optimizer.param_groups[0]['lr']\n        pbar.set_postfix(valid_loss=f'{epoch_loss:0.4f}',\n                         lr=f'{current_lr:0.5f}',\n                         gpu_memory=f'{mem:0.2f} GB')\n    val_scores  = np.mean(val_scores, axis=0)\n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    return epoch_loss, val_scores","metadata":{"execution":{"iopub.status.busy":"2022-07-07T14:06:53.607844Z","iopub.execute_input":"2022-07-07T14:06:53.608157Z","iopub.status.idle":"2022-07-07T14:06:53.623095Z","shell.execute_reply.started":"2022-07-07T14:06:53.608114Z","shell.execute_reply":"2022-07-07T14:06:53.622254Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def run_training(model, optimizer, scheduler, fold, device, num_epochs):\n    if torch.cuda.is_available():\n        print(\"cuda: {}\\n\".format(torch.cuda.get_device_name()))\n    \n    start = time.time()\n    best_model_wts = copy.deepcopy(model.state_dict())\n    best_dice      = -np.inf\n    best_epoch     = -1\n    history = defaultdict(list)\n    \n    for epoch in range(1, num_epochs + 1): \n        gc.collect()\n        print(f'Epoch {epoch}/{num_epochs}', end='')\n        train_loss = train_epoch(model, optimizer, scheduler, \n                                 dataloader=train_loader, \n                                 device=CFG.device, epoch=epoch)\n        \n        val_loss, val_scores = valid_epoch(model, valid_loader, \n                                           device=CFG.device, \n                                           epoch=epoch)\n        val_dice, val_jaccard = val_scores\n    \n        history['Train Loss'].append(train_loss)\n        history['Valid Loss'].append(val_loss)\n        history['Valid Dice'].append(val_dice)\n        history['Valid Jaccard'].append(val_jaccard)\n        \n        # Log the metrics\n        print(f'Valid Dice: {val_dice:0.4f} | Valid Jaccard: {val_jaccard:0.4f}')\n        \n        # deep copy the model\n        if val_dice >= best_dice:\n            print(f\"{c_}Valid Score Improved ({best_dice:0.4f} ---> {val_dice:0.4f})\")\n            best_dice    = val_dice\n            best_jaccard = val_jaccard\n            best_epoch   = epoch\n            best_model_wts = copy.deepcopy(model.state_dict())\n            \n            # Save a model file from the current directory\n            PATH = f\"{CFG.model_name}_{CFG.backbone}_best_epoch-{fold:02d}.pth\"\n            save_model(epoch, model, optimizer, train_loss, val_loss, PATH)\n            print(f\"Model Saved: {PATH}\")\n            \n        last_model_wts = copy.deepcopy(model.state_dict())\n        PATH = f\"{CFG.model_name}_{CFG.backbone}_epoch_{epoch}-{fold:02d}.pth\"\n        save_model(epoch, model, optimizer, train_loss, val_loss, PATH)\n            \n        print(\"\\n\")\n    \n    end = time.time()\n    time_elapsed = end - start\n    print('Training complete in {:.0f}h {:.0f}m {:.0f}s'.format(\n        time_elapsed // 3600, (time_elapsed % 3600) // 60, (time_elapsed % 3600) % 60))\n    print(\"Best Score (Dice): {:.4f}\".format(best_dice))\n    print(\"Best Score (Jaccard): {:.4f}\".format(best_jaccard))\n    print(f\"Best Epoch: {best_epoch}\")\n    \n    # load best model weights\n    model.load_state_dict(best_model_wts)\n    \n    return model, history","metadata":{"execution":{"iopub.status.busy":"2022-07-07T14:06:53.624807Z","iopub.execute_input":"2022-07-07T14:06:53.625597Z","iopub.status.idle":"2022-07-07T14:06:53.643738Z","shell.execute_reply.started":"2022-07-07T14:06:53.62555Z","shell.execute_reply":"2022-07-07T14:06:53.642923Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model Training","metadata":{}},{"cell_type":"code","source":"for fold in CFG.folds:\n    print(f'#'*15)\n    print(f'### Fold: {fold}')\n    print(f'#'*15)\n\n    train_loader = prepare_loader(df=df.copy(), transforms=train_transforms, fold=fold, loader_type=\"train\", \n                                  debug=CFG.debug)\n    valid_loader = prepare_loader(df=df.copy(), transforms=valid_transforms, fold=fold, loader_type=\"valid\", \n                                  debug=CFG.debug)\n    model     = build_model()\n    optimizer = optim.Adamax(model.parameters(), \n                             lr=CFG.lr,\n                             betas=(0.9, 0.999), \n                             eps=1e-08,\n                             weight_decay=CFG.wd)\n    scheduler = setup_scheduler(optimizer)\n\n    model, history = run_training(model, optimizer, scheduler,\n                                  fold=fold,\n                                  device=CFG.device,\n                                  num_epochs=CFG.epochs)\n    plot_metrics(history)","metadata":{"execution":{"iopub.status.busy":"2022-07-07T14:06:53.645349Z","iopub.execute_input":"2022-07-07T14:06:53.646242Z","iopub.status.idle":"2022-07-07T14:07:16.685961Z","shell.execute_reply.started":"2022-07-07T14:06:53.646201Z","shell.execute_reply":"2022-07-07T14:07:16.684747Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prediction","metadata":{}},{"cell_type":"code","source":"test_dataset = DatasetBuilder(df.query(\"fold==0 & empty==False\").sample(frac=1.0), \n                              label=True, transforms=valid_transforms)\ntest_loader  = DataLoader(test_dataset, batch_size=5, \n                          num_workers=4, shuffle=False, pin_memory=True)\nimgs, msks = next(iter(test_loader))\nimgs = imgs.to(CFG.device, dtype=torch.float)\n\npreds = []\nfor fold in CFG.folds:\n    PATH = f\"{CFG.model_name}_{CFG.backbone}_best_epoch-{fold:02d}.pth\"\n    model, _, epoch, loss = load_model(PATH)\n    print(f\"Checkpoint from {epoch} epoch (training loss: {loss})\")\n\n    with torch.no_grad():\n        pred = model(imgs)\n        pred = (nn.Sigmoid()(pred)>0.5).double()\n    preds.append(pred)\n    \nimgs  = imgs.cpu().detach()\npreds = torch.mean(torch.stack(preds, dim=0), dim=0).cpu().detach()\n\nmsks = msks.cpu().detach()","metadata":{"execution":{"iopub.status.busy":"2022-07-07T14:07:16.687335Z","iopub.status.idle":"2022-07-07T14:07:16.688477Z","shell.execute_reply.started":"2022-07-07T14:07:16.688203Z","shell.execute_reply":"2022-07-07T14:07:16.688232Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model Prediction","metadata":{}},{"cell_type":"code","source":"plot_batch(imgs, preds, size=5)","metadata":{"execution":{"iopub.status.busy":"2022-07-07T14:07:16.689694Z","iopub.status.idle":"2022-07-07T14:07:16.690658Z","shell.execute_reply.started":"2022-07-07T14:07:16.690409Z","shell.execute_reply":"2022-07-07T14:07:16.690436Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Ground Truth","metadata":{}},{"cell_type":"code","source":"plot_batch(imgs, msks, size=5)","metadata":{"execution":{"iopub.status.busy":"2022-07-07T14:07:16.691761Z","iopub.status.idle":"2022-07-07T14:07:16.692293Z","shell.execute_reply.started":"2022-07-07T14:07:16.692043Z","shell.execute_reply":"2022-07-07T14:07:16.692069Z"},"trusted":true},"execution_count":null,"outputs":[]}]}