{"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":"In this tutorial, \n\n1. We will build an **EfficientNet Unet** model with segmentation models from GitHub [here](https://github.com/qubvel/segmentation_models). \n2. We will write a **custom data generator** with keras.utils.Sequence.\n3. We will add **data augmentations** with albumentations (refer to this GitHub [repo](https://github.com/albumentations-team/albumentations)) to our dataset.\n4. At last, we will submit results in our [**Inference notebook**](https://www.kaggle.com/code/dingyan/hubmap-efficientnet-unet-inference).\n\nIf you feel this notebook is helpful to you, please **upvote**! Thank you.\n\n*This is a relatively advanced tutorial but if you are interested in basic and much simpler implementations, you could check my another [**notebook**](https://www.kaggle.com/code/dingyan/hubmap-unet-like-model-baseline-with-keras) which can help you easier to understand preprocessing data and building a segmentation model.*","metadata":{}},{"cell_type":"markdown","source":"Let's get started!\n\n**1. Install segmentation models and albumentations.**","metadata":{}},{"cell_type":"code","source":"!pip install -q -U segmentation-models\n!pip install -q -U albumentations","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-08-16T04:13:46.630543Z","iopub.execute_input":"2022-08-16T04:13:46.631303Z","iopub.status.idle":"2022-08-16T04:14:05.533707Z","shell.execute_reply.started":"2022-08-16T04:13:46.631258Z","shell.execute_reply":"2022-08-16T04:14:05.532466Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**2. Import packages we need and set the parameters.**","metadata":{}},{"cell_type":"code","source":"import tensorflow as tf\nfrom tensorflow import keras\nfrom tensorflow.keras import layers\nfrom PIL import Image\nimport os\nimport pandas as pd\nfrom pathlib import Path\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport segmentation_models as sm\nimport random\nimport math\nimport albumentations as A\n\nsm.set_framework('tf.keras')\ntrain_df = pd.read_csv('../input/hubmap-organ-segmentation/train.csv')\ntrain_img_dir = '../input/hubmap-organ-segmentation/train_images'\nids = train_df['id']\n\ninput_img_paths = sorted([os.path.join(train_img_dir, f\"{id}.tiff\") for id in ids])\nrandom.Random(1337).shuffle(input_img_paths)\n\nbatch_size = 4\nimg_size = 512\nthreshold = 0.3","metadata":{"execution":{"iopub.status.busy":"2022-08-16T04:14:05.539223Z","iopub.execute_input":"2022-08-16T04:14:05.539621Z","iopub.status.idle":"2022-08-16T04:14:13.324616Z","shell.execute_reply.started":"2022-08-16T04:14:05.539579Z","shell.execute_reply":"2022-08-16T04:14:13.323618Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**3. Writing a custom dataset with keras.utils.Sequence.**","metadata":{}},{"cell_type":"code","source":"class HubmapOrgan(keras.utils.Sequence):\n    def __init__(self, batch_size, img_size, input_img_paths, df, shuffle=True, augment=None):\n        super().__init__()\n        self.batch_size = batch_size\n        self.img_size = img_size\n        self.input_img_paths = input_img_paths\n        self.df = df\n        self.shuffle = shuffle\n        self.augment = augment\n       \n        \n    def __len__(self):\n        return math.ceil(len(self.input_img_paths) / self.batch_size)\n    \n    def __getitem__(self, idx):\n        i = idx * self.batch_size\n        batch_img_paths = self.input_img_paths[i : i + self.batch_size]\n        \n        x = np.zeros((self.batch_size, self.img_size, self.img_size, 3), dtype='uint8')\n        y = np.zeros((self.batch_size, self.img_size, self.img_size, 1), dtype='uint8')\n        for j, img_path in enumerate(batch_img_paths):\n            \n            # first load images to target size [512, 512]\n            img = keras.utils.load_img(img_path, \n                                      target_size=(self.img_size, self.img_size))\n            img_array = keras.utils.img_to_array(img, dtype='uint8')\n            x[j] = img_array\n            \n            # calculate mask\n            img_id = Path(img_path).stem\n            h, w= self.df[self.df['id']==int(img_id)]['img_height'].iloc[-1], self.df[self.df['id']==int(img_id)]['img_width'].iloc[-1]\n            rle = self.df[self.df['id']==int(img_id)]['rle'].iloc[-1]\n            s = rle.split()\n            starts, lengths = [np.asarray(t, dtype='int') for t in (s[0:][::2], s[1:][::2])]\n            starts = starts - 1\n            original_mask = np.zeros(h*w, dtype=np.uint8)\n            for s, l in zip(starts, lengths):\n                original_mask[s:s+l] = 1\n            original_mask = original_mask.reshape((h, w)).T\n            # we need at least three channels to save an image so here expand mask to (3000, 3000, 1)\n            original_mask = keras.utils.array_to_img(original_mask[:, :, tf.newaxis], scale=False)\n            \n            # resize mask\n            mask = original_mask.resize((self.img_size, self.img_size))\n            mask_array = keras.utils.img_to_array(mask, dtype='uint8')\n            y[j] = mask_array\n        if self.augment is None:\n            return x.astype('float32')/255, y\n        else:\n            aug_x, aug_y = [], []\n            for image, mask in zip(x, y):\n                transformed = self.augment(image=image, mask=mask)\n                aug_x.append(transformed['image'])\n                aug_y.append(transformed['mask'])\n            return np.array(aug_x), np.array(aug_y)\n    def on_epoch_end(self):\n        if self.shuffle == True:\n            np.random.shuffle(self.input_img_paths)","metadata":{"execution":{"iopub.status.busy":"2022-08-16T04:14:13.327382Z","iopub.execute_input":"2022-08-16T04:14:13.327781Z","iopub.status.idle":"2022-08-16T04:14:13.347830Z","shell.execute_reply.started":"2022-08-16T04:14:13.327744Z","shell.execute_reply":"2022-08-16T04:14:13.345421Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**4. Adding data augmentations with Albumentations.**\n\n\nYou could also try different augmentations and check the result. My suggestion of choosing augmentations is whether augmented data make sense in your scenerio. For example, there is an agumentation named \"RandomSunFlare\" which will add sun flare in images. But in this competition, it may not work because these medical data seem not to be exposed in sun.","metadata":{}},{"cell_type":"code","source":"AUGMENTATION_TRAIN = A.Compose([\n    A.HorizontalFlip(),\n    A.VerticalFlip(),\n    A.ShiftScaleRotate(rotate_limit=0),\n    A.RGBShift(),\n    A.ChannelShuffle(),\n    A.GaussNoise(),\n    A.ToFloat(max_value=255)\n])\nAUGMENTATION_TEST = A.Compose([\n    A.ToFloat(max_value=255)\n])","metadata":{"execution":{"iopub.status.busy":"2022-08-16T04:14:13.354558Z","iopub.execute_input":"2022-08-16T04:14:13.354851Z","iopub.status.idle":"2022-08-16T04:14:13.385338Z","shell.execute_reply.started":"2022-08-16T04:14:13.354824Z","shell.execute_reply":"2022-08-16T04:14:13.384220Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let's take a look at augmentated images and masks.","metadata":{}},{"cell_type":"code","source":"a = HubmapOrgan(30, img_size, input_img_paths, train_df, shuffle=False, augment=AUGMENTATION_TRAIN)\nimages, masks = a.__getitem__(0)\n\nrows = 5\ncols = 2\n\nplt.figure(figsize=((cols*5, rows*5)))\nfor i in range(0, rows*cols, 2):\n    plt.subplot(rows, cols, i+1)\n    plt.title('Image')\n    plt.imshow(images[i])\n\n    plt.subplot(rows, cols, i+2)\n    plt.title('Masks')\n    plt.imshow(images[i])\n    plt.imshow(masks[i], cmap='hot', alpha=0.4)","metadata":{"execution":{"iopub.status.busy":"2022-08-16T04:14:13.386775Z","iopub.execute_input":"2022-08-16T04:14:13.387296Z","iopub.status.idle":"2022-08-16T04:14:28.268886Z","shell.execute_reply.started":"2022-08-16T04:14:13.387260Z","shell.execute_reply":"2022-08-16T04:14:28.266449Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**5. Spliting training and validation data.** We only reserve 30 images for validation as we don't have a lot of images.","metadata":{}},{"cell_type":"code","source":"val_num = 30\ntrain_img_paths = input_img_paths[:-val_num]\nval_img_paths = input_img_paths[-val_num:]\n\ntrain_gen = HubmapOrgan(batch_size, img_size, train_img_paths, train_df, augment=AUGMENTATION_TRAIN)\nval_gen = HubmapOrgan(batch_size, img_size, val_img_paths, train_df, augment=AUGMENTATION_TEST)","metadata":{"execution":{"iopub.status.busy":"2022-08-16T04:14:28.269985Z","iopub.execute_input":"2022-08-16T04:14:28.270333Z","iopub.status.idle":"2022-08-16T04:14:28.276996Z","shell.execute_reply.started":"2022-08-16T04:14:28.270302Z","shell.execute_reply":"2022-08-16T04:14:28.276105Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**6. Finetuning pretrained EfficientNet B0 UNet model.**","metadata":{}},{"cell_type":"markdown","source":"We'd better first freeze the encoder layers and train decoder layers for some epochs. Otherwise pretrained weights will be destroyed due to the randomly initialized decoder layers weights updating and propagating through the network.\n\nLet's first train 10 epochs.","metadata":{}},{"cell_type":"code","source":"tf.keras.backend.clear_session()\nBACKBONE = 'efficientnetb0'\nmodel = sm.Unet(BACKBONE, encoder_weights='imagenet', encoder_freeze=True, classes=1, activation='sigmoid', input_shape=(img_size, img_size, 3))\nmodel.compile(optimizer=keras.optimizers.Adam(1e-3),\n             loss=keras.losses.BinaryCrossentropy(),\n             metrics='accuracy')\nmodel.fit(train_gen, epochs=10, validation_data=val_gen)","metadata":{"execution":{"iopub.status.busy":"2022-08-16T04:14:28.278585Z","iopub.execute_input":"2022-08-16T04:14:28.279357Z","iopub.status.idle":"2022-08-16T04:18:04.668477Z","shell.execute_reply.started":"2022-08-16T04:14:28.279322Z","shell.execute_reply":"2022-08-16T04:18:04.667320Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Then we unfreeze the encoder layers and train the whole model with a relatively **smaller learning rate**. We will reduce the learning rate when the loss doesn't go down for 3 epochs. And we also add early stopping.","metadata":{}},{"cell_type":"code","source":"for layer in model.layers:\n    layer.trainable = True\n\nmodel.compile(optimizer=keras.optimizers.Adam(5e-4),\n             loss=keras.losses.BinaryCrossentropy(),\n             metrics='accuracy')\n# model.summary()\ncallbacks = [keras.callbacks.EarlyStopping(monitor='val_loss', patience=6), \n             keras.callbacks.ModelCheckpoint('efficient_unet_model', save_best_only=True),\n            keras.callbacks.ReduceLROnPlateau(monitor='val_loss',\n                                             factor=0.2,\n                                             patience=3,\n                                             verbose=1)]\nhistory = model.fit(train_gen, epochs=50, validation_data=val_gen, callbacks=callbacks)","metadata":{"execution":{"iopub.status.busy":"2022-08-16T04:18:04.670431Z","iopub.execute_input":"2022-08-16T04:18:04.670813Z","iopub.status.idle":"2022-08-16T04:20:36.509329Z","shell.execute_reply.started":"2022-08-16T04:18:04.670776Z","shell.execute_reply":"2022-08-16T04:20:36.508248Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**7. Visualize the training and validation loss.**","metadata":{}},{"cell_type":"code","source":"epochs = range(len(history.history['loss']))\nplt.plot(epochs, history.history['loss'], label='train loss')\nplt.plot(epochs, history.history['val_loss'], label='val loss')\nplt.legend()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-08-16T04:20:36.529073Z","iopub.execute_input":"2022-08-16T04:20:36.529419Z","iopub.status.idle":"2022-08-16T04:20:36.766277Z","shell.execute_reply.started":"2022-08-16T04:20:36.529387Z","shell.execute_reply":"2022-08-16T04:20:36.765291Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = keras.models.load_model('./efficient_unet_model')\nmodel.evaluate(val_gen)","metadata":{"execution":{"iopub.status.busy":"2022-08-16T04:20:36.768455Z","iopub.execute_input":"2022-08-16T04:20:36.769136Z","iopub.status.idle":"2022-08-16T04:20:58.131386Z","shell.execute_reply.started":"2022-08-16T04:20:36.769095Z","shell.execute_reply":"2022-08-16T04:20:58.130239Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**8. Compare groundtruth and prediction masks**","metadata":{}},{"cell_type":"code","source":"def visualize_predictions(images, masks):\n    plt.figure(figsize=(cols*5, rows*5))\n    for i in range(0, rows*cols, 2):\n        plt.subplot(rows, cols, i+1)\n        plt.axis('off')\n        plt.title('Groundtruth')\n        plt.imshow(images[i])\n        plt.imshow(masks[i], alpha=0.3)\n        \n        plt.subplot(rows, cols, i+2)\n        plt.title('Prediction')\n        plt.axis('off')\n        preds = model.predict(np.expand_dims(images[i], 0))[0]\n        pred_mask = np.where(preds > threshold, 1, 0)\n        plt.imshow(images[i])\n        plt.imshow(pred_mask, alpha=0.3)\n\nrows = 10\ncols = 2\n\nval_images, val_masks = HubmapOrgan(30, img_size, val_img_paths, train_df).__getitem__(0)\nvisualize_predictions(val_images, val_masks)","metadata":{"execution":{"iopub.status.busy":"2022-08-16T04:20:58.133195Z","iopub.execute_input":"2022-08-16T04:20:58.133600Z","iopub.status.idle":"2022-08-16T04:21:07.993740Z","shell.execute_reply.started":"2022-08-16T04:20:58.133563Z","shell.execute_reply":"2022-08-16T04:21:07.990440Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We will submit the results in our [**Inference notebook**](https://www.kaggle.com/code/dingyan/hubmap-efficientnet-unet-inference).\n\nAgain, if you feel this notebook is helpful, please **upvote**! Thank you.","metadata":{}}]}