{"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":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\n# for dirname, _, filenames in os.walk('/kaggle/input'):\n#     for filename in filenames:\n#         print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-09-15T02:20:53.283982Z","iopub.execute_input":"2022-09-15T02:20:53.284442Z","iopub.status.idle":"2022-09-15T02:20:53.290633Z","shell.execute_reply.started":"2022-09-15T02:20:53.284401Z","shell.execute_reply":"2022-09-15T02:20:53.289249Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Setup\n\n### Install external library","metadata":{}},{"cell_type":"code","source":"!pip install /kaggle/input/library3/torch_summary-1.4.5-py3-none-any.whl\n!cp /kaggle/input/library3/EfficientUnet-PyTorch ./ -r;\n!pip install ./EfficientUnet-PyTorch/EfficientUnet-PyTorch\n!rm -r ./EfficientUnet-PyTorch","metadata":{"execution":{"iopub.status.busy":"2022-09-15T02:20:53.293133Z","iopub.execute_input":"2022-09-15T02:20:53.293625Z","iopub.status.idle":"2022-09-15T02:22:03.036394Z","shell.execute_reply.started":"2022-09-15T02:20:53.293587Z","shell.execute_reply":"2022-09-15T02:22:03.035065Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Import library","metadata":{}},{"cell_type":"code","source":"import types\nimport random\nfrom collections import defaultdict\nimport os\nfrom pprint import pprint\nfrom glob import glob\nimport shutil\n\nimport cv2\nimport tifffile\nfrom imgaug import augmenters as iaa\n\nfrom matplotlib import pyplot as plt\nimport seaborn as sns\n\nimport torch\nimport torchvision\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.transforms as T\nfrom torch.utils.data import Dataset, random_split, DataLoader\n\nfrom tqdm.notebook import tqdm","metadata":{"execution":{"iopub.status.busy":"2022-09-15T02:22:03.038329Z","iopub.execute_input":"2022-09-15T02:22:03.039030Z","iopub.status.idle":"2022-09-15T02:22:04.923936Z","shell.execute_reply.started":"2022-09-15T02:22:03.038989Z","shell.execute_reply":"2022-09-15T02:22:04.922234Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from efficientunet import *\nfrom torchsummary import summary","metadata":{"execution":{"iopub.status.busy":"2022-09-15T02:22:04.925607Z","iopub.execute_input":"2022-09-15T02:22:04.926234Z","iopub.status.idle":"2022-09-15T02:22:04.937339Z","shell.execute_reply.started":"2022-09-15T02:22:04.926176Z","shell.execute_reply":"2022-09-15T02:22:04.936460Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import logging\n\nlogger = logging.getLogger()\nlogger.setLevel(logging.WARNING)","metadata":{"execution":{"iopub.status.busy":"2022-09-15T02:22:04.941587Z","iopub.execute_input":"2022-09-15T02:22:04.941884Z","iopub.status.idle":"2022-09-15T02:22:04.951567Z","shell.execute_reply.started":"2022-09-15T02:22:04.941857Z","shell.execute_reply":"2022-09-15T02:22:04.950516Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Config","metadata":{}},{"cell_type":"code","source":"class CFG:\n    # dir\n    train_img_dir = \"/kaggle/input/hubmap-organ-segmentation/train_images/\"\n    train_ann_dir = \"/kaggle/input/hubmap-organ-segmentation/train_annotations/\"\n    test_img_dir = \"/kaggle/input/hubmap-organ-segmentation/test_images/\"\n    data_path = \"../input/hubmap-organ-segmentation\"\n    # mini batch\n    epochs = 30\n    batch_size= 4\n    img_size = 224 * 3\n    # optim\n    lr=1e-3\n    weight_decay=1e-8\n    momentum=0.9\n    # cuda\n    device = torch.device( 'cuda' if torch.cuda.is_available() else 'cpu' )\n    # test\n    test=True\n    analysis=False\n    # pretrained model\n    use_pretrained=True # Check on training\n    train_more=False\n    pretrained_lib_dir=\"../input/pretrained-per-organ-v28\"\n    \n    aug_rate=0.8\n\n# if CFG.test:\nprint('>>>>> config >>>>>')\npprint(CFG.__dict__)\nprint('<<<<< config <<<<<')\n","metadata":{"execution":{"iopub.status.busy":"2022-09-15T02:22:04.952919Z","iopub.execute_input":"2022-09-15T02:22:04.953540Z","iopub.status.idle":"2022-09-15T02:22:05.019098Z","shell.execute_reply.started":"2022-09-15T02:22:04.953512Z","shell.execute_reply":"2022-09-15T02:22:05.018106Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Copy pretrained model","metadata":{}},{"cell_type":"code","source":"# !pwd\n# !cp ../input/pretrained-perorgan/* ./ -rv\n\nignore_dir = './ignored'\n\nif CFG.use_pretrained:\n    source_dir = CFG.pretrained_lib_dir\n    if os.path.exists(ignore_dir):\n        source_dir = ignore_dir\n        files = os.listdir(source_dir)\n        for fname in files:\n            _from = os.path.join(source_dir, fname)\n            _to = os.path.join(\".\", fname)\n            shutil.move(\n                _from, \n                _to\n            )\n            print(f'Move from {_from} -> {_to}')\n        os.rmdir(ignore_dir)\n    else:\n        files = os.listdir(source_dir)\n        for fname in files:\n            input_fname = os.path.join('.', fname)\n            if not os.path.exists(input_fname):\n                _from = os.path.join(source_dir, fname)\n                _to_dir = \".\"\n                _to = os.path.join(os.path.join('.', fname))\n                shutil.copy(\n                    _from, \n                    _to_dir\n                )\n                print(f'Copy from {_from} to {_to}')\n            else:\n                print(f'Exists: {fname} in {os.curdir}')\nelse:\n    if not os.path.exists(ignore_dir):\n        os.makedirs(ignore_dir)\n    files = os.listdir('.')\n    for fname in files:\n        if fname.endswith('.pth'):\n            _from = fname\n            _to = os.path.join(ignore_dir, fname)\n            shutil.move(\n                _from, \n                _to\n            )\n            print(f'Move from {_from} -> {_to}')\n            \nprint(f'CFG.use_pretrained: {CFG.use_pretrained}')\nprint(f'ls: {os.listdir(\".\")}')","metadata":{"execution":{"iopub.status.busy":"2022-09-15T02:22:05.020626Z","iopub.execute_input":"2022-09-15T02:22:05.021261Z","iopub.status.idle":"2022-09-15T02:22:09.053433Z","shell.execute_reply.started":"2022-09-15T02:22:05.021225Z","shell.execute_reply":"2022-09-15T02:22:09.052423Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !rm ignored -rf","metadata":{"execution":{"iopub.status.busy":"2022-09-15T02:22:09.054754Z","iopub.execute_input":"2022-09-15T02:22:09.055417Z","iopub.status.idle":"2022-09-15T02:22:09.063090Z","shell.execute_reply.started":"2022-09-15T02:22:09.055377Z","shell.execute_reply":"2022-09-15T02:22:09.061733Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f'{ignore_dir}:')\n!ls {ignore_dir}\nprint()\nprint(f'{CFG.pretrained_lib_dir}:')\n!ls {CFG.pretrained_lib_dir}\nprint()\nprint(f'{\"./\"}')\n!ls ./","metadata":{"execution":{"iopub.status.busy":"2022-09-15T02:22:09.065433Z","iopub.execute_input":"2022-09-15T02:22:09.065807Z","iopub.status.idle":"2022-09-15T02:22:12.076885Z","shell.execute_reply.started":"2022-09-15T02:22:09.065770Z","shell.execute_reply":"2022-09-15T02:22:12.075698Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataframe\n\n### Load dataframe","metadata":{}},{"cell_type":"code","source":"train_df = pd.read_csv(os.path.join(CFG.data_path, \"train.csv\"))\ntest_df = pd.read_csv(os.path.join(CFG.data_path, \"test.csv\"))","metadata":{"execution":{"iopub.status.busy":"2022-09-15T02:22:12.078991Z","iopub.execute_input":"2022-09-15T02:22:12.079368Z","iopub.status.idle":"2022-09-15T02:22:12.434937Z","shell.execute_reply.started":"2022-09-15T02:22:12.079330Z","shell.execute_reply":"2022-09-15T02:22:12.433902Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Analysis","metadata":{}},{"cell_type":"code","source":"if CFG.analysis:\n    train_df.head()","metadata":{"execution":{"iopub.status.busy":"2022-09-15T02:22:12.436358Z","iopub.execute_input":"2022-09-15T02:22:12.436752Z","iopub.status.idle":"2022-09-15T02:22:12.443872Z","shell.execute_reply.started":"2022-09-15T02:22:12.436715Z","shell.execute_reply":"2022-09-15T02:22:12.441523Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.analysis:\n    train_df.info()","metadata":{"execution":{"iopub.status.busy":"2022-09-15T02:22:12.446941Z","iopub.execute_input":"2022-09-15T02:22:12.447910Z","iopub.status.idle":"2022-09-15T02:22:12.452730Z","shell.execute_reply.started":"2022-09-15T02:22:12.447881Z","shell.execute_reply":"2022-09-15T02:22:12.451764Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.analysis:\n    train_df.describe()","metadata":{"execution":{"iopub.status.busy":"2022-09-15T02:22:12.454380Z","iopub.execute_input":"2022-09-15T02:22:12.455281Z","iopub.status.idle":"2022-09-15T02:22:12.462649Z","shell.execute_reply.started":"2022-09-15T02:22:12.455231Z","shell.execute_reply":"2022-09-15T02:22:12.461686Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Change dataset dtype\n- Categorical: organ, data_source, sex","metadata":{}},{"cell_type":"code","source":"cols_to_categorical = ['organ', 'data_source', 'sex']\nfor col in cols_to_categorical:\n    train_df[col].astype(\"category\")","metadata":{"execution":{"iopub.status.busy":"2022-09-15T02:22:12.467711Z","iopub.execute_input":"2022-09-15T02:22:12.468072Z","iopub.status.idle":"2022-09-15T02:22:12.483001Z","shell.execute_reply.started":"2022-09-15T02:22:12.468047Z","shell.execute_reply":"2022-09-15T02:22:12.482053Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.analysis:\n    train_df.describe(include ='object')","metadata":{"execution":{"iopub.status.busy":"2022-09-15T02:22:12.485585Z","iopub.execute_input":"2022-09-15T02:22:12.485999Z","iopub.status.idle":"2022-09-15T02:22:12.492847Z","shell.execute_reply.started":"2022-09-15T02:22:12.485957Z","shell.execute_reply":"2022-09-15T02:22:12.492020Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.analysis:\n    fig, ax = plt.subplots(1, 3)\n    sns.histplot(train_df, x=\"organ\", ax=ax.flat[0])\n    sns.histplot(train_df, x=\"sex\", ax=ax.flat[1])\n    sns.histplot(train_df, x=\"age\", ax=ax.flat[2])","metadata":{"execution":{"iopub.status.busy":"2022-09-15T02:22:12.494248Z","iopub.execute_input":"2022-09-15T02:22:12.494784Z","iopub.status.idle":"2022-09-15T02:22:12.503393Z","shell.execute_reply.started":"2022-09-15T02:22:12.494750Z","shell.execute_reply":"2022-09-15T02:22:12.502460Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Utils\n\n- Get data","metadata":{}},{"cell_type":"code","source":"# Utils fn\ndef _resize_image(img, size=CFG.img_size, dtype='uint8'):\n    return cv2.resize(img, [size, size], interpolation=cv2.INTER_CUBIC).astype(dtype)","metadata":{"execution":{"iopub.status.busy":"2022-09-15T02:22:12.504906Z","iopub.execute_input":"2022-09-15T02:22:12.505747Z","iopub.status.idle":"2022-09-15T02:22:12.512722Z","shell.execute_reply.started":"2022-09-15T02:22:12.505712Z","shell.execute_reply":"2022-09-15T02:22:12.511691Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rle_encode_less_memory(img):\n    #the image should be transposed\n    pixels = img.T.flatten()\n    \n    # This simplified method requires first and last pixel to be zero\n    pixels[0] = 0\n    pixels[-1] = 0\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 2\n    runs[1::2] -= runs[::2]\n    \n    return ' '.join(str(x) for x in runs)\n\n\nif CFG.test:\n    pe_img = torch.ones(9, 9)\n    rle_encode_less_memory(pe_img)","metadata":{"execution":{"iopub.status.busy":"2022-09-15T02:22:12.514136Z","iopub.execute_input":"2022-09-15T02:22:12.514636Z","iopub.status.idle":"2022-09-15T02:22:12.525736Z","shell.execute_reply.started":"2022-09-15T02:22:12.514601Z","shell.execute_reply":"2022-09-15T02:22:12.524803Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ref: https://www.kaggle.com/paulorzp/run-length-encode-and-decode\ndef get_mask(image_id):\n    row = train_df.loc[train_df['id'] == image_id].squeeze()\n    h, w = row[['img_height', 'img_width']]\n    mask = np.zeros(shape=[h * w], dtype=np.uint8)\n#     def rle_decode(mask_rle, shape):\n    s = row['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    for lo, hi in zip(starts, ends):\n        mask[lo : hi] = 1\n        \n    mask = mask.reshape([h, w]).T\n    mask = _resize_image(mask)\n    mask = np.expand_dims(mask, axis=2)\n        \n    return mask\n\nif CFG.test:\n    mask = get_mask(10488)\n    plt.imshow(mask)","metadata":{"execution":{"iopub.status.busy":"2022-09-15T02:22:12.527729Z","iopub.execute_input":"2022-09-15T02:22:12.528143Z","iopub.status.idle":"2022-09-15T02:22:12.919082Z","shell.execute_reply.started":"2022-09-15T02:22:12.528109Z","shell.execute_reply":"2022-09-15T02:22:12.918240Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ref: https://www.kaggle.com/code/masterray/hubmap-inference-tf-tpu-efficientnet-b8-640-640/notebook\ndef get_image(id: int, dir=\"train\", negative=True):\n    assert dir in [\"train\", \"test\"]\n    img_path = f\"/kaggle/input/hubmap-organ-segmentation/{dir}_images/{id}.tiff\"\n    image = tifffile.imread(img_path)\n    if len(image.shape) == 5:\n        image = image.squeeze().transpose(1, 2, 0)\n    # Image Size\n    img_size, _, _ = image.shape\n    # Reverse pixels to make tissue colored and background black\n    if negative:\n        image = image - image.min()\n        image = image / (image.max() - image.min())\n        image = image * 255\n        image = 255 - image.astype(np.uint8)  \n    # Resize\n    image = _resize_image(image)\n    return image, img_size\n    \nif CFG.test:\n    img, size = get_image(10488)\n    print(img.shape)\n    plt.imshow(img)\n    plt.show()\n    img.max()","metadata":{"execution":{"iopub.status.busy":"2022-09-15T02:22:12.921232Z","iopub.execute_input":"2022-09-15T02:22:12.921873Z","iopub.status.idle":"2022-09-15T02:22:13.786273Z","shell.execute_reply.started":"2022-09-15T02:22:12.921834Z","shell.execute_reply":"2022-09-15T02:22:13.785261Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"- Visualization","metadata":{}},{"cell_type":"code","source":"def mask_vis(img, mask, plt=plt, double=False):\n    assert img.shape[2] == 3\n    assert mask.shape[2] == 1\n    if double:\n        plt.subplot(1, 2, 1)\n        plt.imshow(img)\n        plt.subplot(1, 2, 2)\n        plt.imshow(mask)\n        plt.show()\n    else:\n        plt.imshow(img)\n        plt.imshow(mask, alpha=0.4)\n        plt.show()\n\nif CFG.test:\n    id = 10392\n    img, size = get_image(id)\n    mask = get_mask(id)\n    mask_vis(img, mask, double=True)\n    mask.max()","metadata":{"execution":{"iopub.status.busy":"2022-09-15T02:22:13.787354Z","iopub.execute_input":"2022-09-15T02:22:13.788007Z","iopub.status.idle":"2022-09-15T02:22:14.767891Z","shell.execute_reply.started":"2022-09-15T02:22:13.787959Z","shell.execute_reply":"2022-09-15T02:22:14.766988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Defining dataset","metadata":{}},{"cell_type":"markdown","source":"### Transform and augmentation","metadata":{}},{"cell_type":"code","source":"preprocess = T.Compose([\n    T.ToTensor(),\n    T.Normalize(mean=0, std=1),\n])\n\ntarget_transform = T.Compose([\n#     lambda x: torch.tensor(x)\n    T.ToTensor(),\n    T.Normalize(mean=0, std=1),\n    lambda x: x * 255,\n    lambda x: x.type(torch.LongTensor)\n])\n\nt = iaa.Sequential([]).augment_images\n\ndef rand_aug_transform ():\n    rfliplr = random.randint(0, 1)\n    rflipud = random.randint(0, 1)\n    rtx = random.randint(-10, 10)\n    rty = random.randint(-10, 10)\n    rsx = random.uniform(0.8, 1.8)\n    rsy = random.uniform(0.8, 1.8)\n    \n    return T.Compose([\n        iaa.Sequential([\n#             iaa.Sometimes(\n#                 CFG.aug_rate,\n#                 iaa.Sequential([\n                    iaa.Fliplr(rfliplr),\n                    iaa.Flipud(rflipud),\n                    iaa.TranslateX(px=(rtx, rtx)),\n                    iaa.TranslateY(px=(rty, rty)),\n                    iaa.ScaleX((rsx, rsx)),\n                    iaa.ScaleY((rsy, rsy)),\n#                 ])\n#             )\n        ]).augment_image,\n        lambda t: t.transpose(2, 0, 1),\n        lambda t: torch.tensor(t.copy())\n    ])\n\n# img, mask = train_ds.__getitem__(1)\n\ndef aug(img, mask):\n    transform = rand_aug_transform()\n    img = transform(img.permute(1, 2, 0).numpy())\n    mask = transform(mask.permute(1, 2, 0).float().numpy())\n    return img, mask\n\ndef decision(probability):\n    return random.random() < probability\n\n# img, mask = aug(img, mask)\n# # plt.imshow(img.permute(1, 2, 0))\n# mask_vis(img.permute(1, 2, 0), mask.permute(1, 2, 0), double=True)\n# # type(img)\n# img.shape","metadata":{"execution":{"iopub.status.busy":"2022-09-15T02:22:14.769217Z","iopub.execute_input":"2022-09-15T02:22:14.770406Z","iopub.status.idle":"2022-09-15T02:22:14.784634Z","shell.execute_reply.started":"2022-09-15T02:22:14.770365Z","shell.execute_reply":"2022-09-15T02:22:14.783513Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Dataset","metadata":{}},{"cell_type":"code","source":"class ImageDataset(Dataset):\n    def __init__(self, dataframe=train_df, transform=preprocess, target_transform=target_transform, dir=\"train\", aug=aug):\n        super().__init__()\n        self.df = dataframe\n        self.dir = dir\n        self.aug = aug\n        self.transform = transform\n        self.target_transform = target_transform\n        \n    def get_info(self, idx):\n        return self.df.loc[idx]\n    \n    def __getitem__(self, idx):\n        info = self.get_info(idx)\n        \n        img_index = info['id']\n        img, size = get_image(img_index, self.dir)\n        if self.transform:\n            img = self.transform(img)\n        \n        mask = None\n        if self.dir=='train':\n            mask = get_mask(img_index)\n            if self.target_transform:\n                mask = self.target_transform(mask)\n            if self.aug and decision(CFG.aug_rate):\n                img, mask = self.aug(img, mask)\n\n        return img, mask\n        \n        \n    def __len__(self):\n        return len(self.df)\n    \n#     def _draw(self, f):\n#         pass\n    \n#     def show_batch(self, from_idx=0, max_n=6):\n#         fig, ax = plt.subplots(max_n)\n#         for n in range(max_n):\n#             index = n + from_idx\n#             idx = from_idx + n\n#             img, size = self.__getitem__(idx)\n#             mask_vis(img, mask, double=True)","metadata":{"execution":{"iopub.status.busy":"2022-09-15T02:22:14.786149Z","iopub.execute_input":"2022-09-15T02:22:14.786546Z","iopub.status.idle":"2022-09-15T02:22:14.799637Z","shell.execute_reply.started":"2022-09-15T02:22:14.786510Z","shell.execute_reply":"2022-09-15T02:22:14.798649Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\n\n# a, b = train_test_split(train_df, test_size=0.2)\n","metadata":{"execution":{"iopub.status.busy":"2022-09-15T02:22:43.242937Z","iopub.execute_input":"2022-09-15T02:22:43.243343Z","iopub.status.idle":"2022-09-15T02:22:43.252309Z","shell.execute_reply.started":"2022-09-15T02:22:43.243302Z","shell.execute_reply":"2022-09-15T02:22:43.251354Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# spleen_df = train_df[train_df['organ']=='spleen'].copy().reset_index()\norgan_categories = ['kidney', 'prostate', 'largeintestine', 'spleen', 'lung'] # << list(train_df['organ'].value_counts().keys())\n\ndf_per_organ = {}\nval_df_per_organ = {}\nfor organ in organ_categories:\n    new_df = train_df[train_df['organ']==organ].copy()\n    _train, _val = train_test_split(new_df, test_size=0.2)\n    df_per_organ[organ] = _train.reset_index() \n    val_df_per_organ[organ] = _val.reset_index() \n    \ndf_per_organ.keys()","metadata":{"execution":{"iopub.status.busy":"2022-09-15T02:24:05.912795Z","iopub.execute_input":"2022-09-15T02:24:05.913727Z","iopub.status.idle":"2022-09-15T02:24:05.935057Z","shell.execute_reply.started":"2022-09-15T02:24:05.913688Z","shell.execute_reply":"2022-09-15T02:24:05.934030Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_df_per_organ['spleen']","metadata":{"execution":{"iopub.status.busy":"2022-09-15T02:25:38.696468Z","iopub.execute_input":"2022-09-15T02:25:38.696836Z","iopub.status.idle":"2022-09-15T02:25:38.715312Z","shell.execute_reply.started":"2022-09-15T02:25:38.696806Z","shell.execute_reply":"2022-09-15T02:25:38.714136Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# dataset\ndataset = ImageDataset(aug=aug)\ntrain_size = int(len(dataset) * 0.9)\nval_size = len(dataset) - train_size\ntrain_ds, val_ds = random_split(dataset, [train_size, val_size])\n\nds_per_organ = {}\nval_ds_per_organ = {}\nfor k in df_per_organ.keys():\n    ds_per_organ[k] = ImageDataset(dataframe=df_per_organ[k], dir='train')\n    val_ds_per_organ[k] = ImageDataset(dataframe=val_df_per_organ[k], dir='train')\n\n# data loaser\ntrain_dl = DataLoader(ds_per_organ['spleen'], batch_size=CFG.batch_size, shuffle=True)\nval_dl = DataLoader(val_ds, batch_size=CFG.batch_size, shuffle=False)\n\n# if True:\nif CFG.test:\n    img, mask = dataset.__getitem__(1)\n#     print(mask.max())\n    print(img.max())\n    print(img.shape)\n    print(img.dtype)\n#     print(mask.shape)\n    # len(val_ds)","metadata":{"execution":{"iopub.status.busy":"2022-09-14T18:09:31.932436Z","iopub.execute_input":"2022-09-14T18:09:31.933042Z","iopub.status.idle":"2022-09-14T18:09:32.285402Z","shell.execute_reply.started":"2022-09-14T18:09:31.933009Z","shell.execute_reply":"2022-09-14T18:09:32.284060Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Defining model","metadata":{}},{"cell_type":"markdown","source":"### Create model","metadata":{}},{"cell_type":"code","source":"net = get_efficientunet_b3(out_channels=1, concat_input=True, pretrained=False)\n\nif CFG.analysis:\n    summary(net)","metadata":{"execution":{"iopub.status.busy":"2022-09-14T18:09:32.287175Z","iopub.execute_input":"2022-09-14T18:09:32.287920Z","iopub.status.idle":"2022-09-14T18:09:32.503015Z","shell.execute_reply.started":"2022-09-14T18:09:32.287871Z","shell.execute_reply":"2022-09-14T18:09:32.501818Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load model\ndef load_model(model, name):\n    assert name, 'Require name'\n    if glob(f'{name}.pth'):\n        print(f'Loaded: {name}.pth')\n        net.load_state_dict(torch.load(f'{name}.pth'))\n        return True\n    else:\n        print(f'Not found: {name}.pth')\n        return False\n\ndef save_model(model, name):\n    assert name, 'Require name'\n    f_name = f\"{name}.pth\"\n    print(f\"Saved: {f_name}\")\n    torch.save(model.state_dict(), f_name)\n    \nif CFG.test:\n    load_model(net, 'a')","metadata":{"execution":{"iopub.status.busy":"2022-09-14T18:09:32.505741Z","iopub.execute_input":"2022-09-14T18:09:32.506510Z","iopub.status.idle":"2022-09-14T18:09:32.516157Z","shell.execute_reply.started":"2022-09-14T18:09:32.506460Z","shell.execute_reply":"2022-09-14T18:09:32.514894Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Loss fn","metadata":{}},{"cell_type":"code","source":"def dice_loss(pred, target, smooth = 1.):\n    pred = pred.contiguous()\n    target = target.contiguous()\n    \n    intersection = (pred * target).sum(dim=2).sum(dim=2)\n    loss = (1 - ((2. * intersection + smooth) / (pred.sum(dim=2).sum(dim=2) + target.sum(dim=2).sum(dim=2) + smooth)))\n    \n    return loss.mean()\n\ndef calc_loss(pred, target, metrics=None, bce_weight=0.5):\n    bce = F.binary_cross_entropy_with_logits(pred, target)\n\n    pred = torch.sigmoid(pred)\n    dice = dice_loss(pred, target)\n    \n    loss = bce * bce_weight + dice * (1 - bce_weight)\n    \n    if metrics:\n        metrics['bce'] += bce.data.cpu().numpy() * target.size(0)\n        metrics['dice'] += dice.data.cpu().numpy() * target.size(0)\n        metrics['loss'] += loss.data.cpu().numpy() * target.size(0)\n    \n    return loss\n\n\nif CFG.test:\n    pred = torch.rand(3, 1, 4, 4)\n    target = torch.ones(3, 1, 4, 4)\n    metrics = defaultdict(float)\n    calc_loss(pred, target, metrics)","metadata":{"execution":{"iopub.status.busy":"2022-09-14T18:09:32.517733Z","iopub.execute_input":"2022-09-14T18:09:32.518180Z","iopub.status.idle":"2022-09-14T18:09:32.530898Z","shell.execute_reply.started":"2022-09-14T18:09:32.518147Z","shell.execute_reply":"2022-09-14T18:09:32.529433Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Training","metadata":{}},{"cell_type":"code","source":"def weight_reset(m):\n    reset_parameters = getattr(m, \"reset_parameters\", None)\n#     count = 0\n    if callable(reset_parameters):\n        m.reset_parameters()\n#         count+=1\n#     print(count)\n\n# reset param\nif CFG.test:\n    net.apply(weight_reset)\n","metadata":{"execution":{"iopub.status.busy":"2022-09-14T18:09:32.537798Z","iopub.execute_input":"2022-09-14T18:09:32.538205Z","iopub.status.idle":"2022-09-14T18:09:32.702202Z","shell.execute_reply.started":"2022-09-14T18:09:32.538169Z","shell.execute_reply":"2022-09-14T18:09:32.700954Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\n\ngc.collect()\nif CFG.device == 'cuda':\n    torch.cuda.empty_cache()\n    ","metadata":{"execution":{"iopub.status.busy":"2022-09-14T18:09:32.703718Z","iopub.execute_input":"2022-09-14T18:09:32.704322Z","iopub.status.idle":"2022-09-14T18:09:32.994822Z","shell.execute_reply.started":"2022-09-14T18:09:32.704275Z","shell.execute_reply":"2022-09-14T18:09:32.993342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"optimizer = torch.optim.RMSprop(net.parameters(), lr=CFG.lr, weight_decay=CFG.weight_decay, momentum=CFG.momentum)\ncriterion = calc_loss\n\nepochs = CFG.epochs\n\ndef train(net, optimizer, criterion, epochs, dataloader, organ:str = None):\n    if organ:\n        load_model(net, organ)\n    data_size = len(dataloader.dataset)\n    \n    net.to(CFG.device)\n    net.train()\n    for epoch in range(1, epochs + 1):\n        epoch_loss = 0\n        with tqdm(total=data_size, desc=f'Epoch {epoch}/{epochs}', unit='img') as pbar:\n            logger.debug(f'Epoch {epoch}/{epochs}')        \n            for imgs, masks in dataloader:\n                imgs, masks = imgs.to(CFG.device), masks.to(CFG.device)\n                masks_pred = net(imgs)\n                loss = criterion(\n                    masks_pred,\n                    masks.float()\n                )\n                optimizer.zero_grad(set_to_none=True)\n                loss.backward()\n                optimizer.step()\n                pbar.update(imgs.shape[0])\n                epoch_loss += loss.item()\n                pbar.set_postfix(**{'loss (batch)': loss.item()})\n            logger.debug(f'loss (epoch): {epoch_loss}')\n    net.cpu()\n    \n    if organ:\n        save_model(net, organ)\n\n# Training\nfor organ in organ_categories:\n    if not CFG.train_more:\n        if load_model(net, organ):\n            print('|-> Found pretrained model -> ignore train')\n    else:\n        # Train if not exist model\n        print(f\"Train: {organ}\")\n        net.apply(weight_reset)\n        train(\n            net, \n            optimizer, \n            criterion, \n            epochs, \n            DataLoader(ds_per_organ[organ], batch_size=CFG.batch_size, shuffle=True),\n            organ\n        )","metadata":{"execution":{"iopub.status.busy":"2022-09-14T18:09:32.996754Z","iopub.execute_input":"2022-09-14T18:09:32.997506Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### vis","metadata":{}},{"cell_type":"code","source":"def extract_pred_mask(pred_mask, threshold=0.7, out='probability'):\n    \"\"\"\n    pred_mask: torch.Tensor [1 x img_size x img_size]\n    threshold: number or lambda mask: number\n    out: 'probability' or 'category'\n    return: torch.Tensor same size\n    \"\"\"\n    assert out in ['probability', 'category']\n    assert threshold is not float or not callable(threshold)\n    pred_mask = torch.sigmoid(pred_mask)\n    \n    # normalization\n    min = pred_mask.min()\n    max = pred_mask.max()\n#     print(f'max: {max}')\n#     print(f'min: {min}')\n\n    new_pred = (pred_mask - min) / (max - min)\n    \n    if out == 'category':\n        # filter\n        if callable(threshold):\n            thr = threshold(pred_mask)\n        else:\n            thr = threshold\n        pred_mask[pred_mask >= thr] = 1\n        pred_mask[pred_mask < thr] = 0\n\n        pred_mask = pred_mask.long()\n        \n    return pred_mask\n\nif CFG.test:\n    X = torch.rand(1, 3, 3)\n    print(X)\n    extract_pred_mask(X, out='category', threshold=lambda x: x.mean())","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def vis_result(img, mask, pred, organ):\n    plt.subplot(1, 3, 1)\n    plt.imshow(img.permute(1, 2, 0))\n    plt.subplot(1, 3, 2)\n    plt.imshow(mask[0])\n    plt.subplot(1, 3, 3)\n#     min = pred[0].detach().min()\n#     max = pred[0].detach().max()\n#     new_pred = (pred[0].detach() - min) / (max - min)\n#     new_pred[new_pred > 0.8] = 1\n    new_pred = extract_pred_mask(\n        pred[0].detach(), \n        out='category', \n        threshold=lambda x: (x.mean() + x.max()) / 2 if organ=='lung' else 0.7,\n    )\n    \n    plt.imshow(new_pred)\n    plt.show()\n    return\n\n@torch.no_grad()\ndef test_and_vis (dataset: ImageDataset, organ):\n#     imgs, masks = next(iter(val_dl))\n#     net.cpu()\n#     pred = net(imgs)\n\n#     idx = 2\n#     for idx in range(val_dl.batch_size):\n#         vis_result(imgs, masks, pred, idx)\n#     net.eval()\n\n    print(f\"========== {organ} ==========\")\n    if load_model(net, organ):\n\n        net.to(CFG.device)\n        for i in range(len(dataset)):\n            img, mask = dataset.__getitem__(i)\n            img, mask = img.to(CFG.device), mask.to(CFG.device)\n            info = dataset.get_info(i)\n            pred = net(img.unsqueeze(0))\n            pred = pred[0]\n            print(f'ID: {info[\"id\"]}')\n#             print(f'Organ: {info[\"organ\"]}')\n    #         print(info)\n            vis_result(img.cpu(), mask.cpu(), pred.cpu(), organ)\n    #         print(torch.sigmoid(mask).shape)\n            if i > 8:\n                break\n\nif CFG.test:\n    print(organ_categories)\n    for organ in organ_categories:\n        test_and_vis(val_ds_per_organ[organ], organ)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # /kaggle/input/hubmap-organ-segmentation/test_images/10078.tiff\n# test_img, _ = get_image(10078, \"test\")\n# test_img = preprocess(test_img)\n# plt.imshow(test_img.permute(1, 2, 0))\n# # test_img","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Submit","metadata":{}},{"cell_type":"markdown","source":"### Define test df","metadata":{}},{"cell_type":"code","source":"test_df = pd.read_csv(os.path.join(CFG.data_path, \"test.csv\"))\nif CFG.analysis:\n    test_df.head()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_ds = ImageDataset(test_df, dir='test')\n\nif CFG.test:\n    print(f\"test ds len: {len(test_ds)}\")\n\n    img, _ = test_ds.__getitem__(0)\n    \n    img = img.to(CFG.device)\n    net.to(CFG.device)\n    \n    imgs = img.unsqueeze(0)\n    mask_pred = net(imgs)\n    plt.imshow(mask_pred[0][0].detach().cpu())\n#     print(mask)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@torch.no_grad()\ndef test_one_epoch(dataset, model, debug=False):\n    pred_ids = []\n    ids_rles = {}\n    pred_rles = []\n    \n    model = model.to(CFG.device)\n    \n    for oid, organ in enumerate(organ_categories):\n        # load model\n        if organ and glob(f'{organ}.pth'):\n            print(f'Load: {organ}.pth')\n            net.load_state_dict(torch.load(f'{organ}.pth'))\n\n        for idx in range(len(dataset)):\n            # load info\n            info = dataset.get_info(idx)\n            img_organ = info['organ']\n            if oid == 0:\n                pred_ids.append(info[\"id\"])\n\n            if img_organ == organ:\n                # load img and infer\n                img, _ = dataset.__getitem__(idx)\n                imgs = img.unsqueeze(0).to(CFG.device)\n                preds = model(imgs)\n                pred_img = preds[0]\n                size = imgs.size()\n                mask = torch.zeros((1, size[2], size[3]), device=CFG.device, dtype=torch.float32)\n                mask += pred_img\n                mask = mask.cpu()\n\n#                 new_pred = mask\n#                 print(new_pred.shape)\n#                 new_pred = torch.sigmoid(new_pred)\n                \n#                 min = new_pred.min()\n#                 max = new_pred.max()\n#                 new_pred = (new_pred - min) / (max - min)\n# #                 print(new_pred.max())\n\n#                 new_pred[new_pred >= 0.5] = 1\n#         #         print(new_pred.mean())\n#         #         new_pred[new_pred <= new_pred.mean()] = 0\n#         #         print(new_pred.min())\n\n#                 new_pred = (new_pred).long().numpy()\n\n\n                new_pred = extract_pred_mask(\n                    mask, \n                    out='category', \n                    threshold=lambda x: (x.mean() + x.max()) / 2 if organ=='lung' else 0.5,\n                )\n#                 print(new_pred)\n                new_pred = new_pred.numpy()\n\n                info = dataset.get_info(idx)\n                height = info['img_height']\n                width = info['img_width']\n                msk = cv2.resize(new_pred.squeeze(), dsize=(width, height), interpolation=cv2.INTER_NEAREST)\n                rle = rle_encode_less_memory(msk)\n                ids_rles[info['id']] = rle\n                \n                if debug:\n                    fig, ax = plt.subplots(1, 3 if _ is not None else 2) \n                    ax[0].imshow(img.cpu().permute(1, 2, 0))\n\n                    if _ is not None:\n                        ax[1].imshow(_[0])\n                        ax[2].imshow(new_pred[0])\n                        print(_.shape)\n                    else:\n                        ax[1].imshow(new_pred[0])\n                    plt.show()\n                    break\n    if not debug:\n        for idx in pred_ids:\n            pred_rles.append(ids_rles[idx])\n    \n    return pred_ids, pred_rles\n\npred_ids, pred_rles = test_one_epoch(test_ds, net, debug=False)\n\nprint(pred_ids)\n# print(pred_rles)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(pred_rles))\n_min = 1000000\n_max = 0\nfor rle in pred_rles:\n    length = len(rle)\n    if _min > length:\n        _min = length\n    if _max < length:\n        _max = length\n        \nprint(_min)\nprint(_max)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Create submit.csv","metadata":{}},{"cell_type":"code","source":"submit_df = pd.DataFrame({\n    \"id\":pred_ids,\n    \"rle\":pred_rles\n})\nsubmit_df.to_csv('submission.csv',index=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submit_df","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}