{"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":"!cp -r /kaggle/input/segmentation-models-pytorch-021 /kaggle/working/segmentation-models-pytorch-021\n!pip install /kaggle/working/segmentation-models-pytorch-021/pretrainedmodels-0.7.4\n!pip install /kaggle/working/segmentation-models-pytorch-021/efficientnet_pytorch-0.6.3\n!pip install /kaggle/working/segmentation-models-pytorch-021/timm-0.4.12-py3-none-any.whl\n!pip install /kaggle/working/segmentation-models-pytorch-021/segmentation_models_pytorch-0.2.1-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2022-07-14T09:42:21.114261Z","iopub.execute_input":"2022-07-14T09:42:21.114583Z","iopub.status.idle":"2022-07-14T09:44:19.855264Z","shell.execute_reply.started":"2022-07-14T09:42:21.114502Z","shell.execute_reply":"2022-07-14T09:44:19.854325Z"},"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\n\n# visualization\nimport cv2\nimport matplotlib.pyplot as plt\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\n\nimport timm\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\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-14T09:44:19.857426Z","iopub.execute_input":"2022-07-14T09:44:19.857645Z","iopub.status.idle":"2022-07-14T09:44:25.163386Z","shell.execute_reply.started":"2022-07-14T09:44:19.857617Z","shell.execute_reply":"2022-07-14T09:44:25.162539Z"},"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). Pad non-square images, pseudo labels dataset. efficientnet-b4. Data Resampling. 0.7*Dice+0.3*BCE'\n    comment       = 'Unet-efficientnet-b4-320x384'\n    model_name    = 'Unet'\n    backbone      = 'efficientnet-b4'\n    threshold     = 0.5\n    train_bs      = 32\n    valid_bs      = 8\n    pad_img_size  = [320, 384]\n#     img_size      = [320, 384]\n    img_size      = [384, 384]\n    epochs        = 28\n    lr            = 2e-3\n    scheduler     = 'CosineAnnealingLR'\n    min_lr        = 1e-6\n    T_max         = int(10000/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    num_classes   = 3\n    device        = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")","metadata":{"execution":{"iopub.status.busy":"2022-07-14T09:44:25.164783Z","iopub.execute_input":"2022-07-14T09:44:25.165030Z","iopub.status.idle":"2022-07-14T09:44:25.390599Z","shell.execute_reply.started":"2022-07-14T09:44:25.164997Z","shell.execute_reply":"2022-07-14T09:44:25.388777Z"},"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-14T09:44:25.393223Z","iopub.execute_input":"2022-07-14T09:44:25.393724Z","iopub.status.idle":"2022-07-14T09:44:25.410509Z","shell.execute_reply.started":"2022-07-14T09:44:25.393669Z","shell.execute_reply":"2022-07-14T09:44:25.409721Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helper Functions","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    \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    \n    row['height'] = height\n    row['width'] = width\n    row['case'] = case\n    row['day'] = day\n    row['slice'] = slice_\n    return row","metadata":{"execution":{"iopub.status.busy":"2022-07-14T09:44:25.411912Z","iopub.execute_input":"2022-07-14T09:44:25.412269Z","iopub.status.idle":"2022-07-14T09:44:25.422761Z","shell.execute_reply.started":"2022-07-14T09:44:25.412231Z","shell.execute_reply":"2022-07-14T09:44:25.422028Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_img(path):\n    img = cv2.imread(path, cv2.IMREAD_UNCHANGED)\n    img_shape = np.array(img.shape[:2])\n    return img, img_shape\n\n# def load_img(path):\n#     img = cv2.imread(path, cv2.IMREAD_UNCHANGED)\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])\n#         img = img.reshape((*resize))\n#     return img, img_shape\n\ndef load_imgs(img_paths):    \n    for i, img_path in enumerate(img_paths):\n        if i==0:\n            img, img_shape = load_img(img_path)\n            imgs = np.zeros((*np.array(img.shape[:2]), len(img_paths)), dtype=np.float32)\n        else:\n            img, _ = load_img(img_path)\n        img = img.astype('float32') # original is uint16\n        imgs[..., i]+=img\n        \n    mx = np.max(imgs)\n    if mx:\n        q13 = np.percentile(imgs[imgs > 0], [25, 75])\n        max_div = np.max(imgs[imgs < (q13[1] + 3*(q13[1]-q13[0]))])\n        if max_div:\n            imgs = np.clip(imgs/max_div, 0, 1)\n        else:\n            imgs = np.clip(imgs/mx, 0, 1)\n\n    return imgs, img_shape\n\ndef load_msk(path):\n    msk = np.load(path)\n    img_shape = np.array(msk.shape[:2])\n    msk = msk.astype('float32')\n    msk/=255.0\n    return msk\n\n# def load_msk(path, size=CFG.pad_img_size):\n#     msk = np.load(path)\n#     img_shape = np.array(msk.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#         msk = np.pad(msk, [pady, padx, [0,0]])\n#         msk = msk.reshape((*resize, 3))\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(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-14T09:44:25.424463Z","iopub.execute_input":"2022-07-14T09:44:25.424993Z","iopub.status.idle":"2022-07-14T09:44:25.441569Z","shell.execute_reply.started":"2022-07-14T09:44:25.424953Z","shell.execute_reply":"2022-07-14T09:44:25.440776Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train metadata","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-14T09:44:25.444248Z","iopub.execute_input":"2022-07-14T09:44:25.444423Z","iopub.status.idle":"2022-07-14T09:44:25.455850Z","shell.execute_reply.started":"2022-07-14T09:44:25.444401Z","shell.execute_reply":"2022-07-14T09:44:25.455014Z"},"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-14T09:44:25.457327Z","iopub.execute_input":"2022-07-14T09:44:25.457564Z","iopub.status.idle":"2022-07-14T09:44:26.409652Z","shell.execute_reply.started":"2022-07-14T09:44:25.457523Z","shell.execute_reply":"2022-07-14T09:44:26.408941Z"},"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'])","metadata":{"execution":{"iopub.status.busy":"2022-07-14T09:44:26.410856Z","iopub.execute_input":"2022-07-14T09:44:26.412231Z","iopub.status.idle":"2022-07-14T09:44:28.129705Z","shell.execute_reply.started":"2022-07-14T09:44:26.412190Z","shell.execute_reply":"2022-07-14T09:44:28.128888Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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","metadata":{"execution":{"iopub.status.busy":"2022-07-14T09:44:28.133011Z","iopub.execute_input":"2022-07-14T09:44:28.133316Z","iopub.status.idle":"2022-07-14T09:45:04.565940Z","shell.execute_reply.started":"2022-07-14T09:44:28.133281Z","shell.execute_reply":"2022-07-14T09:45:04.565120Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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-14T09:45:04.567225Z","iopub.execute_input":"2022-07-14T09:45:04.567480Z","iopub.status.idle":"2022-07-14T09:45:04.730761Z","shell.execute_reply.started":"2022-07-14T09:45:04.567446Z","shell.execute_reply":"2022-07-14T09:45:04.729653Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display(df.head())","metadata":{"execution":{"iopub.status.busy":"2022-07-14T09:45:04.732378Z","iopub.execute_input":"2022-07-14T09:45:04.732726Z","iopub.status.idle":"2022-07-14T09:45:04.749944Z","shell.execute_reply.started":"2022-07-14T09:45:04.732681Z","shell.execute_reply":"2022-07-14T09:45:04.748882Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"folds_cases = df.groupby(['fold'], sort=True)['case'].unique()\ndisplay(folds_cases)","metadata":{"execution":{"iopub.status.busy":"2022-07-14T09:45:04.751544Z","iopub.execute_input":"2022-07-14T09:45:04.751867Z","iopub.status.idle":"2022-07-14T09:45:04.770134Z","shell.execute_reply.started":"2022-07-14T09:45:04.751815Z","shell.execute_reply":"2022-07-14T09:45:04.769344Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Test metadata","metadata":{}},{"cell_type":"code","source":"IMAGE_PATH  = '/kaggle/input/uw-madison-gi-tract-image-segmentation'\nMASK_PATH = '/kaggle/input/uwmgi-mask-dataset'\nCKPT_DIR = '/kaggle/input/uwmgi-smp-unet-25d-weights-folds'","metadata":{"execution":{"iopub.status.busy":"2022-07-14T09:45:04.771618Z","iopub.execute_input":"2022-07-14T09:45:04.772126Z","iopub.status.idle":"2022-07-14T09:45:04.776785Z","shell.execute_reply.started":"2022-07-14T09:45:04.772089Z","shell.execute_reply":"2022-07-14T09:45:04.775771Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df = pd.read_csv('/kaggle/input/uw-madison-gi-tract-image-segmentation/sample_submission.csv')\ndebug = True if not len(sub_df) else False\nprint(\"Debug mode: \", debug)","metadata":{"execution":{"iopub.status.busy":"2022-07-14T09:45:04.778540Z","iopub.execute_input":"2022-07-14T09:45:04.778842Z","iopub.status.idle":"2022-07-14T09:45:04.796900Z","shell.execute_reply.started":"2022-07-14T09:45:04.778806Z","shell.execute_reply":"2022-07-14T09:45:04.796015Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display(sub_df.head(9))","metadata":{"execution":{"iopub.status.busy":"2022-07-14T09:45:04.798602Z","iopub.execute_input":"2022-07-14T09:45:04.798868Z","iopub.status.idle":"2022-07-14T09:45:04.807107Z","shell.execute_reply.started":"2022-07-14T09:45:04.798832Z","shell.execute_reply":"2022-07-14T09:45:04.806363Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if debug:\n    sub_df = pd.read_csv('/kaggle/input/uw-madison-gi-tract-image-segmentation/train.csv')[:1000*CFG.num_classes]\n    sub_df.columns = sub_df.columns.str.replace('segmentation', 'predicted')\n\ntest_df = sub_df.drop(columns=['class','predicted']).drop_duplicates()\ntest_df = test_df.progress_apply(get_metadata,axis=1)\ndisplay(test_df.head())","metadata":{"execution":{"iopub.status.busy":"2022-07-14T09:45:04.808921Z","iopub.execute_input":"2022-07-14T09:45:04.809535Z","iopub.status.idle":"2022-07-14T09:45:06.826712Z","shell.execute_reply.started":"2022-07-14T09:45:04.809484Z","shell.execute_reply":"2022-07-14T09:45:06.825968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display(sub_df.head(9))","metadata":{"execution":{"iopub.status.busy":"2022-07-14T09:45:06.828253Z","iopub.execute_input":"2022-07-14T09:45:06.828500Z","iopub.status.idle":"2022-07-14T09:45:06.840050Z","shell.execute_reply.started":"2022-07-14T09:45:06.828465Z","shell.execute_reply":"2022-07-14T09:45:06.839021Z"},"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)\nelse:\n    paths = glob(f'/kaggle/input/uw-madison-gi-tract-image-segmentation/test/**/*png', recursive=True)\n    \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-07-14T09:45:06.841497Z","iopub.execute_input":"2022-07-14T09:45:06.841727Z","iopub.status.idle":"2022-07-14T09:46:49.166407Z","shell.execute_reply.started":"2022-07-14T09:45:06.841698Z","shell.execute_reply":"2022-07-14T09:46:49.165648Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = test_df.merge(path_df, on=['case','day','slice'])\ndisplay(test_df.head())","metadata":{"execution":{"iopub.status.busy":"2022-07-14T09:46:49.167791Z","iopub.execute_input":"2022-07-14T09:46:49.168502Z","iopub.status.idle":"2022-07-14T09:46:49.194750Z","shell.execute_reply.started":"2022-07-14T09:46:49.168456Z","shell.execute_reply":"2022-07-14T09:46:49.193944Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"channels=3\nstride=2\n\n# for i in range(channels):\n#     test_df[f'image_path_{i:02}'] = test_df.groupby(['case','day'])['image_path'].shift(-i*stride).fillna(method=\"ffill\")\n\ntest_df[f'image_path_{0:02}'] = test_df.groupby(['case','day'])['image_path'].shift(stride).fillna(method=\"bfill\")\ntest_df[f'image_path_{1:02}'] = test_df.groupby(['case','day'])['image_path'].fillna(method=\"ffill\")\ntest_df[f'image_path_{2:02}'] = test_df.groupby(['case','day'])['image_path'].shift(-stride).fillna(method=\"ffill\")\ntest_df['image_paths'] = test_df[[f'image_path_{i:02d}' for i in range(channels)]].values.tolist()\n\ntest_df.image_paths[0]","metadata":{"execution":{"iopub.status.busy":"2022-07-14T09:46:49.196154Z","iopub.execute_input":"2022-07-14T09:46:49.196472Z","iopub.status.idle":"2022-07-14T09:46:49.223290Z","shell.execute_reply.started":"2022-07-14T09:46:49.196432Z","shell.execute_reply":"2022-07-14T09:46:49.222503Z"},"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 post_transforms():\n    # Convert image to torch.Tensor\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-14T09:46:49.224603Z","iopub.execute_input":"2022-07-14T09:46:49.224926Z","iopub.status.idle":"2022-07-14T09:46:49.232761Z","shell.execute_reply.started":"2022-07-14T09:46:49.224889Z","shell.execute_reply":"2022-07-14T09:46:49.232013Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"valid_transforms = compose([\n    pre_transforms(), \n    post_transforms()\n])","metadata":{"execution":{"iopub.status.busy":"2022-07-14T09:46:49.234054Z","iopub.execute_input":"2022-07-14T09:46:49.234400Z","iopub.status.idle":"2022-07-14T09:46:49.241320Z","shell.execute_reply.started":"2022-07-14T09:46:49.234350Z","shell.execute_reply":"2022-07-14T09:46:49.240455Z"},"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=False, transforms=None):\n        self.df         = df\n        self.label      = label\n        self.img_paths  = df['image_paths'].tolist()\n        self.ids        = df['id'].tolist()\n        self.cases      = df['case'].tolist()\n        if 'mask_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        case      = self.cases[index]\n        img = []\n        img, img_shape = load_imgs(img_path)\n        h, w = img_shape\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        else:\n            if self.transforms:\n                data = self.transforms(image=img)\n                img  = data['image']\n            return img, id_, case, (h, w)","metadata":{"execution":{"iopub.status.busy":"2022-07-14T09:46:49.242543Z","iopub.execute_input":"2022-07-14T09:46:49.242877Z","iopub.status.idle":"2022-07-14T09:46:49.254824Z","shell.execute_reply.started":"2022-07-14T09:46:49.242842Z","shell.execute_reply":"2022-07-14T09:46:49.253873Z"},"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-14T09:46:49.256322Z","iopub.execute_input":"2022-07-14T09:46:49.256792Z","iopub.status.idle":"2022-07-14T09:46:49.422608Z","shell.execute_reply.started":"2022-07-14T09:46:49.256756Z","shell.execute_reply":"2022-07-14T09:46:49.421772Z"},"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=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, optimizer=None):\n    model = build_model()\n\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-14T09:46:49.424000Z","iopub.execute_input":"2022-07-14T09:46:49.424338Z","iopub.status.idle":"2022-07-14T09:46:50.612345Z","shell.execute_reply.started":"2022-07-14T09:46:49.424310Z","shell.execute_reply":"2022-07-14T09:46:50.611579Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference","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\n# def masks2rles(msks, ids, heights, widths):\n#     pred_strings = []\n#     pred_ids = []\n#     pred_classes = []\n\n#     for idx in range(msks.shape[0]):\n#         msk = msks[idx]\n#         height = heights[idx].item()\n#         width = widths[idx].item()\n#         pad_img_size = max(height, width)\n#         msk = cv2.resize(msk,\n#                          dsize=(pad_img_size, pad_img_size), \n#                          interpolation=cv2.INTER_LINEAR)\n#         img_shape = np.array([height, width])\n#         resize = np.array([pad_img_size, pad_img_size])\n\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# #             msk = msk[pady[0]:-pady[1], padx[0]:-padx[1], :]\n#             msk = msk[pady[0]:-pady[1], :, :]\n#             msk = msk.reshape((*img_shape, 3))\n        \n#         rle = [None]*3\n#         for midx in [0, 1, 2]:\n#             rle[midx] = mask2rle(msk[...,midx])\n        \n#         pred_strings.extend(rle)\n#         pred_ids.extend([ids[idx]]*len(rle))\n#         pred_classes.extend(['large_bowel', 'small_bowel', 'stomach'])\n\n#     return pred_strings, pred_ids, pred_classes\n\ndef masks2rles(msks, ids, heights, widths):\n    pred_strings = []\n    pred_ids = []\n    pred_classes = []\n\n    for idx in range(msks.shape[0]):\n        msk = msks[idx]\n        height = heights[idx].item()\n        width = widths[idx].item()\n        msk = cv2.resize(msk,\n                         dsize=(width, height), \n                         interpolation=cv2.INTER_LINEAR)\n        \n        rle = [None]*3\n        for midx in [0, 1, 2]:\n            rle[midx] = mask2rle(msk[...,midx])\n        \n        pred_strings.extend(rle)\n        pred_ids.extend([ids[idx]]*len(rle))\n        pred_classes.extend(['large_bowel', 'small_bowel', 'stomach'])\n\n    return pred_strings, pred_ids, pred_classes","metadata":{"execution":{"iopub.status.busy":"2022-07-14T09:46:50.615680Z","iopub.execute_input":"2022-07-14T09:46:50.615887Z","iopub.status.idle":"2022-07-14T09:46:51.711306Z","shell.execute_reply.started":"2022-07-14T09:46:50.615861Z","shell.execute_reply":"2022-07-14T09:46:51.710565Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Single-Fold Inference","metadata":{}},{"cell_type":"code","source":"# @torch.no_grad()\n# def inference(sub_df, model_path, test_loader, threshold=CFG.threshold):    \n#     model, _, _, _ = load_model(model_path)\n\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        \n#         size = img.size()\n#         msk = []\n#         msk = torch.zeros((size[0], 3, size[2], size[3]), device=CFG.device, dtype=torch.float32)\n        \n#         out   = model(img) # .squeeze(0) # removing batch axis\n#         out   = nn.Sigmoid()(out) # removing channel axis\n        \n#         msk+= out\n#         msk = (msk.permute((0,2,3,1))>=threshold).to(torch.uint8).cpu().detach().numpy() # shape: (n, h, w, c)\n        \n#         pred_strings, pred_ids, pred_classes = masks2rles(msk, ids, heights, widths)\n        \n#         for pred_string, pred_id, pred_class in zip(pred_strings, pred_ids, pred_classes):\n#             sub_df.loc[(sub_df[\"id\"] == pred_id) & (sub_df[\"class\"] == pred_class), 'predicted'] = pred_string\n\n#         del img, msk, pred_string, pred_id, pred_class\n        \n#         gc.collect()\n#         gc.collect()\n#         torch.cuda.empty_cache()\n#         gc.collect()\n\n#     return sub_df","metadata":{"execution":{"iopub.status.busy":"2022-07-14T09:46:51.715307Z","iopub.execute_input":"2022-07-14T09:46:51.715502Z","iopub.status.idle":"2022-07-14T09:46:51.719440Z","shell.execute_reply.started":"2022-07-14T09:46:51.715477Z","shell.execute_reply":"2022-07-14T09:46:51.718646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# test_dataset = DatasetBuilder(test_df, transforms=valid_transforms)\n# test_loader  = DataLoader(test_dataset, batch_size=CFG.valid_bs, \n#                           num_workers=4, shuffle=False, pin_memory=True)\n\n# # model_path = f'{CKPT_DIR}/{CFG.model_name}_{CFG.backbone}_best_epoch-00.pth'\n# # model_path = f'{CKPT_DIR}/{CFG.model_name}_{CFG.backbone}_epoch_23-00.pth' # for 'Version 4'\n# model_path = f'{CKPT_DIR}/{CFG.model_name}_{CFG.backbone}_epoch_24-00.pth' # for 'Version 7'\n\n# print(\"Model path: \", model_path)\n# sub_df = inference(sub_df, model_path, test_loader)","metadata":{"execution":{"iopub.status.busy":"2022-07-14T09:46:51.720707Z","iopub.execute_input":"2022-07-14T09:46:51.721115Z","iopub.status.idle":"2022-07-14T09:46:51.732603Z","shell.execute_reply.started":"2022-07-14T09:46:51.721078Z","shell.execute_reply":"2022-07-14T09:46:51.731843Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Multi-Fold Inference","metadata":{}},{"cell_type":"code","source":"def get_mul_factors(cases):\n    cases = cases.to(torch.uint8).cpu().detach().numpy()\n\n    mul_factors = np.ones((CFG.n_fold, CFG.valid_bs)) * 0.2\n    for fold, fold_cases in enumerate(folds_cases):\n        ids = np.squeeze(np.argwhere(np.isin(cases, fold_cases)))\n        mul_factors[:, ids] = 0.22\n        mul_factors[fold, ids] = 0.12\n\n    return mul_factors","metadata":{"execution":{"iopub.status.busy":"2022-07-14T09:46:51.733983Z","iopub.execute_input":"2022-07-14T09:46:51.734290Z","iopub.status.idle":"2022-07-14T09:46:51.742610Z","shell.execute_reply.started":"2022-07-14T09:46:51.734254Z","shell.execute_reply":"2022-07-14T09:46:51.741860Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#             mul_factor = torch.from_numpy(mul_factors[fold]).to(CFG.device, dtype=torch.float)\n#             msk += torch.mul(out, mul_factors[fold])","metadata":{"execution":{"iopub.status.busy":"2022-07-14T09:46:51.743876Z","iopub.execute_input":"2022-07-14T09:46:51.744448Z","iopub.status.idle":"2022-07-14T09:46:51.754201Z","shell.execute_reply.started":"2022-07-14T09:46:51.744413Z","shell.execute_reply":"2022-07-14T09:46:51.753451Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@torch.no_grad()\ndef inference(sub_df, model_paths, test_loader, threshold=CFG.threshold): \n    models = []\n    for model_path in model_paths:\n        model, _, _, _ = load_model(model_path)\n        models.append(model)\n\n    for idx, (img, ids, cases, (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\n        #mul_factors = get_mul_factors(cases)\n\n        size = img.size()\n        msk = []\n        msk = torch.zeros((size[0], 3, size[2], size[3]), device=CFG.device, dtype=torch.float32)\n\n        for fold, model in enumerate(models):\n            #out   = model(img) # .squeeze(0) # removing batch axis\n            #out   = nn.Sigmoid()(out) # removing channel axis\n            out = 0.5*nn.Sigmoid()(model(img)) + 0.5*nn.Sigmoid()(model(img.flip(3)).flip(3))\n\n            #for i, mul_factor in enumerate(mul_factors[fold]):\n                #out[i] = torch.mul(out[i], mul_factor)\n            msk += out * 0.2\n\n            gc.collect()\n            gc.collect()\n            torch.cuda.empty_cache()\n            gc.collect()\n\n        msk = (msk.permute((0,2,3,1))>=threshold).to(torch.uint8).cpu().detach().numpy() # shape: (n, h, w, c)\n\n        pred_strings, pred_ids, pred_classes = masks2rles(msk, ids, heights, widths)\n\n        for pred_string, pred_id, pred_class in zip(pred_strings, pred_ids, pred_classes):\n            sub_df.loc[(sub_df[\"id\"] == pred_id) & (sub_df[\"class\"] == pred_class), 'predicted'] = pred_string\n\n        del img, msk, pred_string, pred_id, pred_class\n\n        gc.collect()\n        gc.collect()\n        torch.cuda.empty_cache()\n        gc.collect()\n\n    return sub_df","metadata":{"execution":{"iopub.status.busy":"2022-07-14T09:46:51.755489Z","iopub.execute_input":"2022-07-14T09:46:51.755837Z","iopub.status.idle":"2022-07-14T09:46:51.770587Z","shell.execute_reply.started":"2022-07-14T09:46:51.755800Z","shell.execute_reply":"2022-07-14T09:46:51.769702Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset = DatasetBuilder(test_df, transforms=valid_transforms)\ntest_loader  = DataLoader(test_dataset, batch_size=CFG.valid_bs, \n                          num_workers=4, shuffle=False, pin_memory=True)\ntorch.cuda.empty_cache()\n\n# mixed best dice/best validation loss\n# model_paths = [\n#     f'{CKPT_DIR}/{CFG.model_name}_{CFG.backbone}_epoch_27-00.pth',\n#     f'{CKPT_DIR}/{CFG.model_name}_{CFG.backbone}_epoch_27-01.pth',\n#     f'{CKPT_DIR}/{CFG.model_name}_{CFG.backbone}_epoch_27-02.pth',\n#     f'{CKPT_DIR}/{CFG.model_name}_{CFG.backbone}_epoch_26-03.pth',\n#     f'{CKPT_DIR}/{CFG.model_name}_{CFG.backbone}_epoch_28-04.pth',\n# ]\n\n# best validation loss [1] (!)\nmodel_paths = [\n    #f'{CKPT_DIR}/{CFG.model_name}_{CFG.backbone}_epoch_15-00.pth',\n    \n    f'{CKPT_DIR}/{CFG.model_name}_{CFG.backbone}_epoch_23-00.pth',\n    f'{CKPT_DIR}/{CFG.model_name}_{CFG.backbone}_epoch_25-01.pth',\n    f'{CKPT_DIR}/{CFG.model_name}_{CFG.backbone}_epoch_27-02.pth',\n    f'{CKPT_DIR}/{CFG.model_name}_{CFG.backbone}_epoch_26-03.pth',\n    f'{CKPT_DIR}/{CFG.model_name}_{CFG.backbone}_epoch_25-04.pth',\n]\n\n# best validation loss [2]\n# model_paths = [\n#     f'{CKPT_DIR}/{CFG.model_name}_{CFG.backbone}_epoch_27-00.pth',\n#     f'{CKPT_DIR}/{CFG.model_name}_{CFG.backbone}_epoch_25-01.pth',\n#     f'{CKPT_DIR}/{CFG.model_name}_{CFG.backbone}_epoch_27-02.pth',\n#     f'{CKPT_DIR}/{CFG.model_name}_{CFG.backbone}_epoch_26-03.pth',\n#     f'{CKPT_DIR}/{CFG.model_name}_{CFG.backbone}_epoch_25-04.pth',\n# ]\n\nprint(\"Model path:\\n\", model_paths)\nsub_df = inference(sub_df, model_paths, test_loader)","metadata":{"execution":{"iopub.status.busy":"2022-07-14T09:46:51.772206Z","iopub.execute_input":"2022-07-14T09:46:51.772509Z","iopub.status.idle":"2022-07-14T09:50:18.781959Z","shell.execute_reply.started":"2022-07-14T09:46:51.772472Z","shell.execute_reply":"2022-07-14T09:50:18.780706Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submission","metadata":{}},{"cell_type":"code","source":"sub_df.to_csv('submission.csv', index=False)\ndisplay(sub_df)","metadata":{"execution":{"iopub.status.busy":"2022-07-14T09:50:18.783315Z","iopub.status.idle":"2022-07-14T09:50:18.784417Z","shell.execute_reply.started":"2022-07-14T09:50:18.784130Z","shell.execute_reply":"2022-07-14T09:50:18.784162Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}