{"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":"import cv2\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\nfrom pathlib import Path\nfrom collections import Counter\n\nimport seaborn as sns\nimport matplotlib.pyplot as plt\n%matplotlib inline\n\ndata_path = Path('/kaggle/input/hubmap-organ-segmentation')\nimg_path = data_path / 'train_images'\nfiles = [f.name for f in list(data_path.glob('*'))]\nprint(files)","metadata":{"execution":{"iopub.status.busy":"2022-08-08T10:31:20.228321Z","iopub.execute_input":"2022-08-08T10:31:20.229004Z","iopub.status.idle":"2022-08-08T10:31:20.926613Z","shell.execute_reply.started":"2022-08-08T10:31:20.228912Z","shell.execute_reply":"2022-08-08T10:31:20.925593Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Install Libraries","metadata":{}},{"cell_type":"code","source":"!pip install -U segmentation-models-pytorch -q","metadata":{"execution":{"iopub.status.busy":"2022-08-08T10:31:22.967737Z","iopub.execute_input":"2022-08-08T10:31:22.968334Z","iopub.status.idle":"2022-08-08T10:31:42.045915Z","shell.execute_reply.started":"2022-08-08T10:31:22.968271Z","shell.execute_reply":"2022-08-08T10:31:42.044763Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# View Data","metadata":{}},{"cell_type":"markdown","source":"## DataFrames","metadata":{}},{"cell_type":"code","source":"all_df = pd.read_csv(data_path / 'train.csv')\nall_df.head()","metadata":{"execution":{"iopub.status.busy":"2022-08-08T10:31:42.048213Z","iopub.execute_input":"2022-08-08T10:31:42.048561Z","iopub.status.idle":"2022-08-08T10:31:42.369654Z","shell.execute_reply.started":"2022-08-08T10:31:42.048531Z","shell.execute_reply":"2022-08-08T10:31:42.368490Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\n\ntrain_df, test_df = train_test_split(all_df, test_size=0.20, random_state=42, stratify=all_df['organ'])\n\nprint(len(train_df), len(test_df))","metadata":{"execution":{"iopub.status.busy":"2022-08-08T10:31:42.371339Z","iopub.execute_input":"2022-08-08T10:31:42.371774Z","iopub.status.idle":"2022-08-08T10:31:42.490177Z","shell.execute_reply.started":"2022-08-08T10:31:42.371736Z","shell.execute_reply":"2022-08-08T10:31:42.489142Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, axes = plt.subplots(1, 3, figsize=(20, 6))\n\nfor ax, df, title in zip(\n    axes.ravel(),\n    [all_df, train_df, test_df],\n    ['All Data', 'Train Data', 'Test Data']\n):\n    \n    ax.set_title(title)\n    organ_cnt = Counter(df['organ'])\n    keys = sorted(list(organ_cnt.keys()))\n    values = [organ_cnt[k] for k in keys]\n    sns.barplot(\n        x = keys,\n        y = values,\n        ax=ax\n    )\n    ax.set_ylabel(\"Number of images\")","metadata":{"execution":{"iopub.status.busy":"2022-08-08T10:31:47.050107Z","iopub.execute_input":"2022-08-08T10:31:47.050487Z","iopub.status.idle":"2022-08-08T10:31:47.482648Z","shell.execute_reply.started":"2022-08-08T10:31:47.050455Z","shell.execute_reply":"2022-08-08T10:31:47.481555Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Images","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(10,6))\nplt.title(\"Image Sizes\")\nwidth_cnt = Counter(all_df['img_width'])\nkeys = sorted(list(width_cnt.keys()))\nvalues = [width_cnt[k] for k in keys]\nsns.barplot(\n    x=keys,\n    y=values,\n)\nplt.ylabel(\"Number of images\")","metadata":{"execution":{"iopub.status.busy":"2022-08-08T10:31:51.064903Z","iopub.execute_input":"2022-08-08T10:31:51.065458Z","iopub.status.idle":"2022-08-08T10:31:51.431804Z","shell.execute_reply.started":"2022-08-08T10:31:51.065413Z","shell.execute_reply":"2022-08-08T10:31:51.430818Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def view_images(\n    imgs, titles,\n    rows, cols,\n    masks = None, color=(255, 0, 0), br=0.5,\n    figsize=(20, 20)\n):\n    assert(len(imgs) == len(titles) == rows*cols)\n    \n    fig, axes = plt.subplots(rows, cols, figsize=figsize)\n    for idx, (ax, img, title) in enumerate(\n        zip(axes.ravel(), imgs, titles)\n    ):\n        ax.set_title(title)\n        draw_img = img.copy()\n        if masks:\n            mask = masks[idx].astype(bool)\n            color_arr = np.zeros(draw_img.shape) + color\n            draw_img[mask] = draw_img[mask]*(1 - br) + color_arr[mask]*br\n            draw_img.astype(np.uint8)\n        ax.imshow(draw_img)\n\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-08-08T10:31:55.172967Z","iopub.execute_input":"2022-08-08T10:31:55.173697Z","iopub.status.idle":"2022-08-08T10:31:55.182462Z","shell.execute_reply.started":"2022-08-08T10:31:55.173660Z","shell.execute_reply":"2022-08-08T10:31:55.181357Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"imgs, titles = [], []\nlabels = ['kidney', 'largeintestine', 'lung', 'prostate', 'spleen']\nrows, cols = len(labels), 4\n\nfor r in range(rows):\n    for c in range(cols):\n        _df = all_df[all_df['organ'] == labels[r]].iloc[c]\n        \n        im = img_path / f\"{_df.id}.tiff\"\n        img = cv2.imread(str(im))\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        \n        imgs.append(img)\n        titles.append(_df.organ)\n        \nview_images(\n    imgs, titles,\n    rows, cols\n)","metadata":{"execution":{"iopub.status.busy":"2022-08-08T10:16:45.512801Z","iopub.execute_input":"2022-08-08T10:16:45.513911Z","iopub.status.idle":"2022-08-08T10:17:13.309055Z","shell.execute_reply.started":"2022-08-08T10:16:45.513866Z","shell.execute_reply":"2022-08-08T10:17:13.307737Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rle2mask(rle, size):\n    rle = np.array(list(map(int, rle.split())))\n    label = np.zeros((size*size), dtype=np.uint8)\n    for start, end in zip(rle[::2], rle[1::2]):\n        label[start:start+end] = 1\n    return label.reshape(size, size).T\n\ndef mask2rle(mask):\n    pixels = mask.T.flatten()   \n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0]\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)\n\nstr1 = all_df.iloc[0].rle\nsize = all_df.iloc[0].img_width\nstr2 = mask2rle(rle2mask(str1, size))\nprint(str1 == str2)","metadata":{"execution":{"iopub.status.busy":"2022-08-08T10:32:02.087516Z","iopub.execute_input":"2022-08-08T10:32:02.088098Z","iopub.status.idle":"2022-08-08T10:32:02.149531Z","shell.execute_reply.started":"2022-08-08T10:32:02.088064Z","shell.execute_reply":"2022-08-08T10:32:02.148435Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"imgs, titles, masks = [], [], []\nlabels = ['kidney', 'largeintestine', 'lung', 'prostate', 'spleen']\nrows, cols = len(labels), 4\n\nfor r in range(rows):\n    for c in range(cols):\n        _df = all_df[all_df['organ'] == labels[r]].iloc[c]\n        \n        im = img_path / f\"{_df.id}.tiff\"\n        img = cv2.imread(str(im))\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        \n        mask = rle2mask(_df.rle, _df.img_width)\n        \n        imgs.append(img)\n        titles.append(_df.organ)\n        masks.append(mask)\n        \nview_images(\n    imgs, titles,\n    rows, cols,\n    masks=masks\n)","metadata":{"execution":{"iopub.status.busy":"2022-08-08T10:17:13.377492Z","iopub.execute_input":"2022-08-08T10:17:13.378705Z","iopub.status.idle":"2022-08-08T10:17:41.549795Z","shell.execute_reply.started":"2022-08-08T10:17:13.378662Z","shell.execute_reply":"2022-08-08T10:17:41.547579Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Pytorch Dataset","metadata":{}},{"cell_type":"code","source":"from torch.utils.data import Dataset as BaseDataset\nfrom torch.utils.data import DataLoader\n\nclass Dataset(BaseDataset):\n    def __init__(\n        self,\n        df: pd.DataFrame, img_path: Path,\n        transform: callable = None, return_class: bool = False\n    ):\n        self.df = df\n        self.img_path = img_path\n        self.transform = transform\n        self.return_class = return_class\n        \n        self.images = []\n        for i, row in tqdm(df.iterrows(), desc='Prepare data'):\n            im = self.img_path / f\"{row['id']}.tiff\"\n            img = cv2.imread(str(im))\n            img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n            self.images.append(img)\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        \n#         im = self.img_path / f\"{self.df['id'].iloc[idx]}.tiff\"\n#         img = cv2.imread(str(im))\n#         img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        img = self.images[idx]\n        \n        mask = rle2mask(rle=self.df['rle'].iloc[idx], size=self.df['img_width'].iloc[idx])\n#         mask = np.expand_dims(mask, axis=2)\n        \n        if self.transform:\n            augment = self.transform(image=img, mask=mask)\n            img, mask = augment['image'], augment['mask']\n            \n        if self.return_class:\n            return img, mask, self.df['organ'].iloc[idx]\n        return img, mask","metadata":{"execution":{"iopub.status.busy":"2022-08-08T10:32:07.863651Z","iopub.execute_input":"2022-08-08T10:32:07.864246Z","iopub.status.idle":"2022-08-08T10:32:09.457468Z","shell.execute_reply.started":"2022-08-08T10:32:07.864211Z","shell.execute_reply":"2022-08-08T10:32:09.456500Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import albumentations as albu\nfrom albumentations.pytorch import ToTensorV2\n\ncrop = 448\ntest_crop = 2016\n\n# dataset for view images\nview_dataset = Dataset(\n    df = train_df[:20],\n    img_path = img_path,\n    transform = albu.Compose([\n        albu.PadIfNeeded(min_height=crop, min_width=crop),\n        albu.CropNonEmptyMaskIfExists(height=crop, width=crop),\n        albu.Flip(),\n    ]),\n    return_class = True\n)\n# train dataset\ntrain_dataset = Dataset(\n    df = train_df,\n    img_path = img_path,\n    transform = albu.Compose([\n        albu.PadIfNeeded(min_height=crop, min_width=crop),\n        albu.CropNonEmptyMaskIfExists(height=crop, width=crop),\n        albu.Flip(),\n        albu.Normalize(),\n        ToTensorV2(transpose_mask=True)\n    ])\n)\n# test dataset\ntest_dataset = Dataset(\n    df = test_df,\n    img_path = img_path,\n    transform = albu.Compose([\n        albu.PadIfNeeded(\n            min_height=None,\n            min_width=None,\n            pad_height_divisor=32,\n            pad_width_divisor=32,\n        ),\n        albu.Normalize(),\n        ToTensorV2(transpose_mask=True)\n    ])\n)","metadata":{"execution":{"iopub.status.busy":"2022-08-08T10:32:12.283822Z","iopub.execute_input":"2022-08-08T10:32:12.284402Z","iopub.status.idle":"2022-08-08T10:34:31.501330Z","shell.execute_reply.started":"2022-08-08T10:32:12.284359Z","shell.execute_reply":"2022-08-08T10:34:31.500261Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"imgs, titles, masks = [], [], []\nrows, cols = 4, 4\n\nfor r in range(rows):\n    for c in range(cols):\n        \n        img, mask, label = view_dataset[r*rows + c]\n        \n        imgs.append(img)\n        titles.append(label)\n        masks.append(mask)\n        \nview_images(\n    imgs, titles,\n    rows, cols,\n    masks=masks\n)","metadata":{"execution":{"iopub.status.busy":"2022-08-08T10:34:34.286904Z","iopub.execute_input":"2022-08-08T10:34:34.287351Z","iopub.status.idle":"2022-08-08T10:34:38.866629Z","shell.execute_reply.started":"2022-08-08T10:34:34.287307Z","shell.execute_reply":"2022-08-08T10:34:38.862357Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Trainer","metadata":{}},{"cell_type":"code","source":"import torch\nimport segmentation_models_pytorch as smp\nfrom segmentation_models_pytorch.losses import DiceLoss\nfrom torchmetrics import MeanMetric, Dice\nfrom torch.optim import Adam\nfrom torch.optim.lr_scheduler import CosineAnnealingLR","metadata":{"execution":{"iopub.status.busy":"2022-08-08T10:34:45.815049Z","iopub.execute_input":"2022-08-08T10:34:45.815492Z","iopub.status.idle":"2022-08-08T10:34:49.874218Z","shell.execute_reply.started":"2022-08-08T10:34:45.815454Z","shell.execute_reply":"2022-08-08T10:34:49.872867Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Trainer:\n    def __init__(\n        self, \n        model,\n        train_loader, test_loader,\n        loss, metric,\n        optimizer, sheduler\n    ):\n        self.model = model\n        self.train_loader = train_loader\n        self.test_loader = test_loader\n        self.loss = loss\n        self.metric = metric\n        self.optimizer = optimizer\n        self.sheduler = sheduler\n        \n    def fit(self, epochs, device):\n        \n        best_dice = 0\n        history = {\n            'epochs': [],\n            'lr': [],\n            'train_loss': [],\n            'train_dice': [],\n            'test_loss': [],\n            'test_dice': []\n        }\n        \n        self.model.to(device)\n        \n        for epoch in range(epochs):\n            \n            history['epochs'].append(epoch)\n            \n            # Train epoch\n            self.model.train()\n            mean_loss = MeanMetric().to(device)\n            mean_dice = MeanMetric().to(device)\n            for x, y in tqdm(self.train_loader, desc=f\"Train epoch {epoch}\"):\n                x, y = x.to(device), y.to(device)\n                \n                optimizer.zero_grad()\n                pred = model(x)\n                \n                mean_dice.update(self.metric(pred, y.long()))\n                \n                l = self.loss(pred, y.long())\n                mean_loss.update(l)\n                \n                l.backward()\n                optimizer.step()\n                \n            history['train_loss'].append(mean_loss.compute().item())\n            history['train_dice'].append(mean_dice.compute().item())\n                \n            # Test epoch\n            self.model.eval()\n            mean_loss = MeanMetric().to(device)\n            mean_dice = MeanMetric().to(device)\n            with torch.no_grad():\n                for x, y in tqdm(self.test_loader, desc=f\"Test epoch {epoch}\"):\n                    x, y = x.to(device), y.to(device)\n                    pred = model(x)\n\n                    mean_dice.update(self.metric(pred, y.long()))\n\n                    l = self.loss(pred, y.long())\n                    mean_loss.update(l)\n            \n            history['test_loss'].append(mean_loss.compute().item())\n            history['test_dice'].append(mean_dice.compute().item())\n            \n            if best_dice < history['test_dice'][-1]:\n                best_dice = history['test_dice'][-1]\n                print(f\"save best model with dice {best_dice}\")\n                torch.save(self.model, 'model_final.pth')\n                \n            history['lr'].append(self.sheduler.get_last_lr())\n            self.sheduler.step()\n            \n        return history","metadata":{"execution":{"iopub.status.busy":"2022-08-08T10:34:52.827735Z","iopub.execute_input":"2022-08-08T10:34:52.828105Z","iopub.status.idle":"2022-08-08T10:34:52.843879Z","shell.execute_reply.started":"2022-08-08T10:34:52.828073Z","shell.execute_reply":"2022-08-08T10:34:52.842248Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\n\nmodel = smp.Unet(\n    encoder_name = 'resnet18',\n    encoder_weights = None,\n    classes = 1,\n    activation = None,\n)\n\ntrain_loader = DataLoader(train_dataset, batch_size=16, shuffle=True)\ntest_loader = DataLoader(test_dataset, batch_size=1, shuffle=False)\n\nloss = DiceLoss(mode='binary')\nmetric = Dice(\n    average='samples',\n    ignore_index=0,\n).to(device)\n\nlr = 3e-4\nepochs = 300\noptimizer = Adam(model.parameters(), lr=lr)\nsheduler = CosineAnnealingLR(optimizer, T_max=epochs)","metadata":{"execution":{"iopub.status.busy":"2022-08-08T10:37:18.922603Z","iopub.execute_input":"2022-08-08T10:37:18.923302Z","iopub.status.idle":"2022-08-08T10:37:19.259752Z","shell.execute_reply.started":"2022-08-08T10:37:18.923248Z","shell.execute_reply":"2022-08-08T10:37:19.258542Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer = Trainer(\n    model=model,\n    train_loader=train_loader, test_loader=test_loader,\n    loss=loss, metric=metric,\n    optimizer=optimizer, sheduler=sheduler,\n)","metadata":{"execution":{"iopub.status.busy":"2022-08-08T10:37:21.706730Z","iopub.execute_input":"2022-08-08T10:37:21.707106Z","iopub.status.idle":"2022-08-08T10:37:21.712779Z","shell.execute_reply.started":"2022-08-08T10:37:21.707073Z","shell.execute_reply":"2022-08-08T10:37:21.711670Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history = trainer.fit(epochs, device)","metadata":{"execution":{"iopub.status.busy":"2022-08-08T10:37:25.916335Z","iopub.execute_input":"2022-08-08T10:37:25.916941Z","iopub.status.idle":"2022-08-08T10:38:53.970249Z","shell.execute_reply.started":"2022-08-08T10:37:25.916907Z","shell.execute_reply":"2022-08-08T10:38:53.968484Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Plot History","metadata":{}},{"cell_type":"code","source":"fig, axes = plt.subplots(1, 3, figsize=(20, 6))\n\n# losses\naxes[0].set_title('Loss')\naxes[0].plot(history['epochs'], history['train_loss'], 'r', label='train')\naxes[0].plot(history['epochs'], history['test_loss'], 'g', label='test')\n\n# losses\naxes[1].set_title('Dice')\naxes[1].plot(history['epochs'], history['train_dice'], 'r', label='train')\naxes[1].plot(history['epochs'], history['test_dice'], 'g', label='test')\n\n# losses\naxes[2].set_title('Learning Rate')\naxes[2].plot(history['epochs'], history['lr'], 'b')\n\nplt.plot()","metadata":{"execution":{"iopub.status.busy":"2022-08-07T20:19:50.194410Z","iopub.execute_input":"2022-08-07T20:19:50.195158Z","iopub.status.idle":"2022-08-07T20:19:50.644356Z","shell.execute_reply.started":"2022-08-07T20:19:50.195119Z","shell.execute_reply":"2022-08-07T20:19:50.643385Z"},"trusted":true},"execution_count":null,"outputs":[]}]}