{"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 os\nimport tensorflow as tf\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport pandas as pd\nimport cv2\nimport tensorflow.keras.backend as K\nfrom sklearn.model_selection import train_test_split\nfrom IPython.display import clear_output\n\n\nimport warnings\n\nwarnings.filterwarnings(\"ignore\")","metadata":{"execution":{"iopub.status.busy":"2021-12-25T13:37:49.54301Z","iopub.execute_input":"2021-12-25T13:37:49.543378Z","iopub.status.idle":"2021-12-25T13:37:56.011878Z","shell.execute_reply.started":"2021-12-25T13:37:49.543261Z","shell.execute_reply":"2021-12-25T13:37:56.010832Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Config:\n    '''\n    Config in which I hold all hyperparameters and frequently used variables such as image shape, train directory path etc.\n    '''\n    def __init__(self, DEBUG=False):\n        self.DEBUG = DEBUG\n        \n    TRAIN_CSV = '../input/sartorius-cell-instance-segmentation/train.csv'\n    TRAIN_DIR = '../input/sartorius-cell-instance-segmentation/train/'\n    TEST_DIR = '../input/sartorius-cell-instance-segmentation/test/'\n    \n    IMG_SHAPE = (512, 512)\n    \n    LR = 1e-3\n    \n    EPOCHS = 100\n    \n    N_FILTERS = 32\n\n    \n    BATCH_SIZE = 4\n    AUTOTUNE = tf.data.AUTOTUNE\n    \n    N_CLASSES = 1\n    BUFFER_SIZE = 2\n    \n    val_size = 0.1\n    \n    WEIGHTS_PATH = os.path.join('./', 'model.h5')\n        \n    seed = 123\n    \n\nconfig = Config()","metadata":{"execution":{"iopub.status.busy":"2021-12-25T13:37:56.015968Z","iopub.execute_input":"2021-12-25T13:37:56.016222Z","iopub.status.idle":"2021-12-25T13:37:56.029129Z","shell.execute_reply.started":"2021-12-25T13:37:56.016192Z","shell.execute_reply":"2021-12-25T13:37:56.027392Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_csv = pd.read_csv(Config.TRAIN_CSV)\ntrain_csv.head()","metadata":{"execution":{"iopub.status.busy":"2021-12-25T13:37:56.033564Z","iopub.execute_input":"2021-12-25T13:37:56.033834Z","iopub.status.idle":"2021-12-25T13:37:56.742708Z","shell.execute_reply.started":"2021-12-25T13:37:56.033793Z","shell.execute_reply":"2021-12-25T13:37:56.74151Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(15,15))\n\nplt.imshow(cv2.imread(config.TRAIN_DIR + '0030fd0e6378' + '.png'))\nplt.axis(\"off\")\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2021-12-25T13:37:56.745591Z","iopub.execute_input":"2021-12-25T13:37:56.745992Z","iopub.status.idle":"2021-12-25T13:37:57.229304Z","shell.execute_reply.started":"2021-12-25T13:37:56.745948Z","shell.execute_reply":"2021-12-25T13:37:57.228493Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_csv.shape","metadata":{"execution":{"iopub.status.busy":"2021-12-25T13:37:57.230412Z","iopub.execute_input":"2021-12-25T13:37:57.230695Z","iopub.status.idle":"2021-12-25T13:37:57.239043Z","shell.execute_reply.started":"2021-12-25T13:37:57.230658Z","shell.execute_reply":"2021-12-25T13:37:57.237938Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_csv[\"id\"].unique().shape","metadata":{"execution":{"iopub.status.busy":"2021-12-25T13:37:57.240737Z","iopub.execute_input":"2021-12-25T13:37:57.241411Z","iopub.status.idle":"2021-12-25T13:37:57.26636Z","shell.execute_reply.started":"2021-12-25T13:37:57.241357Z","shell.execute_reply":"2021-12-25T13:37:57.265246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_csv[train_csv['id'] == '0030fd0e6378']","metadata":{"execution":{"iopub.status.busy":"2021-12-25T13:37:57.268272Z","iopub.execute_input":"2021-12-25T13:37:57.268834Z","iopub.status.idle":"2021-12-25T13:37:57.30736Z","shell.execute_reply.started":"2021-12-25T13:37:57.268791Z","shell.execute_reply":"2021-12-25T13:37:57.306498Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rle_decode(mask_rle, shape, color=1):\n    '''\n    mask_rle: run-length as string formated (start length)\n    shape: (height,width) of array to return \n    Returns numpy array, 1 - mask, 0 - background\n\n    '''\n    s = mask_rle.split()\n    starts, lengths = [np.asarray(x, dtype=int) for x in (s[0:][::2], s[1:][::2])]\n    starts -= 1\n    ends = starts + lengths\n    img = np.zeros((shape[0] * shape[1], shape[2]), dtype=np.float32)\n    for lo, hi in zip(starts, ends):\n        img[lo : hi] = color\n    return img.reshape(shape)\n\n\ndef build_masks(labels, input_shape, colors=True):\n    height, width = input_shape\n    if colors:\n        mask = np.zeros((height, width, 3))\n        for label in labels:\n            mask += rle_decode(label, shape=(height, width , 3), color=np.random.rand(3))\n    else:\n        mask = np.zeros((height, width, 1))\n        for label in labels:\n            mask += rle_decode(label, shape=(height, width, 1))\n    mask = mask.clip(0, 1)\n    return mask","metadata":{"execution":{"iopub.status.busy":"2021-12-25T13:37:57.308855Z","iopub.execute_input":"2021-12-25T13:37:57.309166Z","iopub.status.idle":"2021-12-25T13:37:57.320523Z","shell.execute_reply.started":"2021-12-25T13:37:57.309129Z","shell.execute_reply":"2021-12-25T13:37:57.319219Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_csv[\"cell_type\"].unique()","metadata":{"execution":{"iopub.status.busy":"2021-12-25T13:37:57.32266Z","iopub.execute_input":"2021-12-25T13:37:57.32299Z","iopub.status.idle":"2021-12-25T13:37:57.343517Z","shell.execute_reply.started":"2021-12-25T13:37:57.322948Z","shell.execute_reply":"2021-12-25T13:37:57.3423Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cell_types = train_csv[\"cell_type\"].value_counts()\n\nplt.figure(figsize=(10, 6), tight_layout=True)\n\nplt.bar(cell_types.index, cell_types.values)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2021-12-25T13:37:57.348671Z","iopub.execute_input":"2021-12-25T13:37:57.348961Z","iopub.status.idle":"2021-12-25T13:37:57.592641Z","shell.execute_reply.started":"2021-12-25T13:37:57.34893Z","shell.execute_reply":"2021-12-25T13:37:57.591441Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"shy5y_sample = train_csv[train_csv['cell_type'] == 'shsy5y'].sample(2)['id']\ncort_sample = train_csv[train_csv['cell_type'] == 'cort'].sample(2)['id']\nastro_sample = train_csv[train_csv['cell_type'] == 'astro'].sample(2)['id']","metadata":{"execution":{"iopub.status.busy":"2021-12-25T13:37:57.594312Z","iopub.execute_input":"2021-12-25T13:37:57.594688Z","iopub.status.idle":"2021-12-25T13:37:57.643481Z","shell.execute_reply.started":"2021-12-25T13:37:57.594644Z","shell.execute_reply":"2021-12-25T13:37:57.642571Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def display_sample(sample_ids, n_samples=2, hspace=-0.6):\n    '''\n    Function to visualize images and their annotations\n    \n    sample_ids - list of ids\n    n_samples - number of samples to display\n    hspace - parameter for matplotlib, it contorls spacing between images \n    '''\n    fig, axs = plt.subplots(n_samples, 3, figsize=(22, 25))\n    \n    for idx, sample_id in enumerate(sample_ids):\n        sample_image = cv2.imread(os.path.join(config.TRAIN_DIR + sample_id + '.png'))\n        \n        sample_rles = train_csv.loc[train_csv['id'] == sample_id]['annotation'].values\n        \n        sample_mask_colors = build_masks(sample_rles, (520, 704), colors=True)\n        sample_mask = build_masks(sample_rles, (520, 704), colors=False)\n        \n        axs[idx][0].imshow(sample_image)\n        axs[idx][0].axis('off')        \n        \n        axs[idx][1].imshow(sample_mask_colors)        \n        axs[idx][1].axis('off')        \n        \n        axs[idx][2].imshow(sample_mask)\n        axs[idx][2].axis('off')        \n\n    axs[0][0].set_title(\"Image\", fontsize=16)\n    axs[0][1].set_title(\"Mask Color\", fontsize=16)\n    axs[0][2].set_title(\"Mask\", fontsize=16)\n\n    fig.subplots_adjust(hspace=hspace)\n    plt.show()       ","metadata":{"execution":{"iopub.status.busy":"2021-12-25T13:37:57.644954Z","iopub.execute_input":"2021-12-25T13:37:57.645327Z","iopub.status.idle":"2021-12-25T13:37:57.657365Z","shell.execute_reply.started":"2021-12-25T13:37:57.645281Z","shell.execute_reply":"2021-12-25T13:37:57.655872Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display_sample(shy5y_sample)","metadata":{"execution":{"iopub.status.busy":"2021-12-25T13:37:57.659806Z","iopub.execute_input":"2021-12-25T13:37:57.660407Z","iopub.status.idle":"2021-12-25T13:38:00.442487Z","shell.execute_reply.started":"2021-12-25T13:37:57.660306Z","shell.execute_reply":"2021-12-25T13:38:00.438577Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display_sample(cort_sample)","metadata":{"execution":{"iopub.status.busy":"2021-12-25T13:38:00.444751Z","iopub.execute_input":"2021-12-25T13:38:00.445177Z","iopub.status.idle":"2021-12-25T13:38:01.601043Z","shell.execute_reply.started":"2021-12-25T13:38:00.445135Z","shell.execute_reply":"2021-12-25T13:38:01.596816Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display_sample(astro_sample)","metadata":{"execution":{"iopub.status.busy":"2021-12-25T13:38:01.602899Z","iopub.execute_input":"2021-12-25T13:38:01.603465Z","iopub.status.idle":"2021-12-25T13:38:02.998373Z","shell.execute_reply.started":"2021-12-25T13:38:01.603423Z","shell.execute_reply":"2021-12-25T13:38:02.997509Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ids = train_csv['id'].unique()\n\ntrain_ids, val_ids = train_test_split(ids, test_size=config.val_size, random_state=config.seed)","metadata":{"execution":{"iopub.status.busy":"2021-12-25T13:38:02.999621Z","iopub.execute_input":"2021-12-25T13:38:03.000011Z","iopub.status.idle":"2021-12-25T13:38:03.016731Z","shell.execute_reply.started":"2021-12-25T13:38:02.999974Z","shell.execute_reply":"2021-12-25T13:38:03.015499Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ids.shape","metadata":{"execution":{"iopub.status.busy":"2021-12-25T13:38:03.018541Z","iopub.execute_input":"2021-12-25T13:38:03.020186Z","iopub.status.idle":"2021-12-25T13:38:03.028169Z","shell.execute_reply.started":"2021-12-25T13:38:03.020138Z","shell.execute_reply":"2021-12-25T13:38:03.026843Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_train_ds():\n    '''\n    This function creates a generator for train dataset\n    '''\n    for image_id in train_ids:\n        rows = train_csv.loc[train_csv['id'] == image_id]\n        image = tf.io.read_file(config.TRAIN_DIR + image_id + '.png')\n        image = tf.image.decode_image(image, channels=3, dtype=tf.float32)\n        rles = rows['annotation'].values\n\n        mask = build_masks(rles, (520, 704), colors=False)\n        mask = tf.cast(tf.image.resize(mask, config.IMG_SHAPE), tf.int32)\n\n        image = tf.image.resize(image, config.IMG_SHAPE)\n        image /= 255.0\n    \n        yield image, mask\n        \n        \ndef load_val_ds():\n    '''\n    This function creates a generator for train dataset\n    '''\n    for image_id in val_ids:\n        rows = train_csv.loc[train_csv['id'] == image_id]\n        image = tf.io.read_file(config.TRAIN_DIR + image_id + '.png')\n        image = tf.image.decode_image(image, channels=3, dtype=tf.float32)\n        rles = rows['annotation'].values\n\n        mask = build_masks(rles, (520, 704), colors=False)\n        mask = tf.cast(tf.image.resize(mask, config.IMG_SHAPE), tf.int32)\n\n        image = tf.image.resize(image, config.IMG_SHAPE)\n        image /= 255.0\n    \n        yield image, mask","metadata":{"execution":{"iopub.status.busy":"2021-12-25T13:38:03.030607Z","iopub.execute_input":"2021-12-25T13:38:03.031437Z","iopub.status.idle":"2021-12-25T13:38:03.046795Z","shell.execute_reply.started":"2021-12-25T13:38:03.031302Z","shell.execute_reply":"2021-12-25T13:38:03.045514Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ds = tf.data.Dataset.from_generator(\n    load_train_ds, \n    output_types=(tf.float32, tf.int32)\n)\n\nval_ds = tf.data.Dataset.from_generator(\n    load_val_ds, \n    output_types=(tf.float32, tf.int32)\n)","metadata":{"execution":{"iopub.status.busy":"2021-12-25T13:38:03.049996Z","iopub.execute_input":"2021-12-25T13:38:03.050571Z","iopub.status.idle":"2021-12-25T13:38:05.811038Z","shell.execute_reply.started":"2021-12-25T13:38:03.050524Z","shell.execute_reply":"2021-12-25T13:38:05.810102Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def augment_ds(image, mask):\n    image = tf.image.random_flip_up_down(image, seed=config.seed)\n    mask = tf.image.random_flip_up_down(mask, seed=config.seed)\n    \n    image = tf.image.random_flip_left_right(image, seed=config.seed)\n    mask = tf.image.random_flip_left_right(mask, seed=config.seed)\n    \n    return image, mask","metadata":{"execution":{"iopub.status.busy":"2021-12-25T13:38:05.812394Z","iopub.execute_input":"2021-12-25T13:38:05.812958Z","iopub.status.idle":"2021-12-25T13:38:05.820942Z","shell.execute_reply.started":"2021-12-25T13:38:05.812899Z","shell.execute_reply":"2021-12-25T13:38:05.819824Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ds = (\n    train_ds\n    .shuffle(config.BUFFER_SIZE)\n    .map(augment_ds)\n    .batch(config.BATCH_SIZE)    \n    .prefetch(Config.AUTOTUNE)\n)\n\nval_ds = val_ds.batch(config.BATCH_SIZE)","metadata":{"execution":{"iopub.status.busy":"2021-12-25T13:38:05.822776Z","iopub.execute_input":"2021-12-25T13:38:05.823491Z","iopub.status.idle":"2021-12-25T13:38:06.148663Z","shell.execute_reply.started":"2021-12-25T13:38:05.823446Z","shell.execute_reply":"2021-12-25T13:38:06.147455Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_batch = next(iter(train_ds))\n\nimages, masks = sample_batch\n\nfig, ax = plt.subplots(config.BATCH_SIZE, 2, figsize=(20, 20))\n\nfor i in range(config.BATCH_SIZE):\n    ax[i][0].imshow(images[i] * 255)\n    ax[i][0].axis('off')        \n    \n    ax[i][1].imshow(masks[i])    \n    ax[i][1].axis('off')        \n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2021-12-25T13:38:06.151532Z","iopub.execute_input":"2021-12-25T13:38:06.152441Z","iopub.status.idle":"2021-12-25T13:38:08.12834Z","shell.execute_reply.started":"2021-12-25T13:38:06.152392Z","shell.execute_reply":"2021-12-25T13:38:08.127498Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def attention(input_tensor, g, inter_shape):    \n    input_shapes = input_tensor.shape\n    g_shapes = g.shape\n    \n    x = tf.keras.layers.Conv2D(inter_shape, 1, 2, padding=\"same\")(input_tensor)\n    g = tf.keras.layers.Conv2D(inter_shape, 1, padding=\"same\")(g)\n\n\n    add = tf.keras.layers.add([x, g])\n    relu = tf.keras.layers.Activation('relu')(x)\n    \n    psi = tf.keras.layers.Conv2D(1, 1, padding=\"same\")(relu)\n    sigmoid = tf.keras.layers.Activation('sigmoid')(psi)\n    \n    upsample = tf.keras.layers.UpSampling2D(size=(2, 2))(sigmoid)\n    \n    att = tf.keras.layers.multiply([upsample, input_tensor])\n    \n    output = tf.keras.layers.Conv2D(input_shapes[3], 1, padding=\"same\")(att)\n    return tf.keras.layers.BatchNormalization()(output)\n\ndef conv_block(input_tensor, n_filters, dropout=0.5, batch_norm=True):\n    x_save = tf.keras.layers.Conv2D(n_filters, 3, activation=\"relu\", padding=\"same\")(input_tensor)\n    if batch_norm:\n        x = tf.keras.layers.BatchNormalization()(x_save)\n    \n    x = tf.keras.layers.Conv2D(n_filters, 3, activation=\"relu\", padding=\"same\")(x)\n    if batch_norm:\n        x = tf.keras.layers.BatchNormalization()(x)\n    \n    if dropout:\n        x = tf.keras.layers.Dropout(dropout)(x)\n            \n    x = tf.keras.layers.add([x, x_save])\n    x = tf.keras.layers.Activation(\"relu\")(x)\n\n    return x\n\ndef downsample(x, n_filters, dropout=0.5, batch_norm=True):\n    res_conn = conv_block(x, n_filters, dropout=dropout, batch_norm=batch_norm)\n    \n    x = tf.keras.layers.MaxPool2D((2, 2), strides=(2, 2))(res_conn)\n    \n    return x, res_conn\n\ndef upsample(x, n_filters, skip_conn, dropout=0.5, batch_norm=True):\n    att = attention(skip_conn, x, n_filters)\n    x = tf.keras.layers.Conv2DTranspose(n_filters, (2, 2), strides=2, padding=\"same\", activation=\"relu\")(x)\n    x = tf.keras.layers.Concatenate()([x, att])\n    x = conv_block(x, n_filters)\n    \n    if dropout:\n        x = tf.keras.layers.Dropout(dropout)(x)\n\n    if batch_norm:\n        x = tf.keras.layers.BatchNormalization()(x)\n        \n    return x\n    \ndef create_model(n_filters):\n    inputs = tf.keras.layers.Input(shape=(*Config.IMG_SHAPE, 3))\n    \n    x, skip_conn1 = downsample(inputs, n_filters)\n    x, skip_conn2 = downsample(x, n_filters * 2)\n    x, skip_conn3 = downsample(x, n_filters * 4)\n    x, skip_conn4 = downsample(x, n_filters * 8)\n    \n    x = conv_block(x, n_filters * 16)\n\n    x = upsample(x, n_filters * 8, skip_conn4)\n    x = upsample(x, n_filters * 4, skip_conn3)    \n    x = upsample(x, n_filters * 2, skip_conn2)    \n    x = upsample(x, n_filters, skip_conn1)\n    \n    outputs = tf.keras.layers.Conv2D(config.N_CLASSES, 3, activation=\"sigmoid\", padding=\"same\")(x)\n    \n    return tf.keras.Model(inputs=inputs, outputs=outputs)","metadata":{"execution":{"iopub.status.busy":"2021-12-25T13:38:08.130123Z","iopub.execute_input":"2021-12-25T13:38:08.13086Z","iopub.status.idle":"2021-12-25T13:38:08.155519Z","shell.execute_reply.started":"2021-12-25T13:38:08.130816Z","shell.execute_reply":"2021-12-25T13:38:08.154657Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = create_model(config.N_FILTERS)\nmodel.summary()","metadata":{"execution":{"iopub.status.busy":"2021-12-25T13:38:08.156766Z","iopub.execute_input":"2021-12-25T13:38:08.157991Z","iopub.status.idle":"2021-12-25T13:38:09.307129Z","shell.execute_reply.started":"2021-12-25T13:38:08.157936Z","shell.execute_reply":"2021-12-25T13:38:09.306075Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def dice_loss(y_true, y_pred, smooth=1.0):\n    y_true = tf.cast(y_true, tf.float32)\n    y_true_f = K.flatten(y_true)\n    y_pred_f = K.flatten(y_pred)\n    intersection = K.sum(y_true_f * y_pred_f)\n    return 1 - (2. * intersection + smooth) / (K.sum(K.square(y_true_f)) + K.sum(K.square(y_pred_f)) + smooth)\n\ndef iou_coef(y_true, y_pred, smooth=1):\n    intersection = K.sum(K.abs(y_true * y_pred), axis=[1,2,3])\n    union = K.sum(y_true,[1,2,3]) + K.sum(y_pred,[1,2,3]) - intersection\n    iou = K.mean((intersection + smooth) / (union + smooth), axis=0)\n    return iou","metadata":{"execution":{"iopub.status.busy":"2021-12-25T13:38:09.310677Z","iopub.execute_input":"2021-12-25T13:38:09.311042Z","iopub.status.idle":"2021-12-25T13:38:09.319952Z","shell.execute_reply.started":"2021-12-25T13:38:09.310999Z","shell.execute_reply":"2021-12-25T13:38:09.318508Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"optimizer = tf.keras.optimizers.Adam(learning_rate=config.LR)\n\nmetrics = [iou_coef]\nmodel.compile(optimizer=optimizer, loss=dice_loss, metrics=metrics)","metadata":{"execution":{"iopub.status.busy":"2021-12-25T13:38:09.321627Z","iopub.execute_input":"2021-12-25T13:38:09.321917Z","iopub.status.idle":"2021-12-25T13:38:09.346846Z","shell.execute_reply.started":"2021-12-25T13:38:09.321878Z","shell.execute_reply":"2021-12-25T13:38:09.345434Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cp_callback = tf.keras.callbacks.ModelCheckpoint(\n    config.WEIGHTS_PATH,\n    save_best_only=True,\n    save_weights_only=True,\n    verbose=1,\n    monitor=\"val_loss\",\n    mode=\"min\"\n)\n\nrlr_callback = tf.keras.callbacks.ReduceLROnPlateau(\n    monitor='val_loss', \n    factor=0.01, \n    patience=5, \n    min_delta=1e-2\n)\n\nes_callback = tf.keras.callbacks.EarlyStopping(\n    monitor='val_loss', \n    min_delta=1e-2, \n    patience=15, \n    verbose=1,\n    mode='min',\n)","metadata":{"execution":{"iopub.status.busy":"2021-12-25T13:38:09.348686Z","iopub.execute_input":"2021-12-25T13:38:09.349009Z","iopub.status.idle":"2021-12-25T13:38:09.355901Z","shell.execute_reply.started":"2021-12-25T13:38:09.348978Z","shell.execute_reply":"2021-12-25T13:38:09.354664Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history = model.fit(\n    train_ds, \n    epochs=config.EPOCHS, \n    validation_data=val_ds,\n    callbacks=[cp_callback, rlr_callback, es_callback]\n)","metadata":{"execution":{"iopub.status.busy":"2021-12-25T13:38:09.363176Z","iopub.execute_input":"2021-12-25T13:38:09.363956Z","iopub.status.idle":"2021-12-25T14:47:51.474501Z","shell.execute_reply.started":"2021-12-25T13:38:09.363813Z","shell.execute_reply":"2021-12-25T14:47:51.473388Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history_dict = history.history\n\nfig, ax = plt.subplots(1, 2, figsize=(15, 5), tight_layout=True)\n\nax[0].plot(history_dict['loss'], label=\"Training loss\", linewidth=3)\nax[0].plot(history_dict['val_loss'], label=\"Validation loss\", linewidth=3)\nax[0].set_xlabel(\"Epoch\")\nax[0].set_ylabel(\"Loss\")\nax[0].set_title(\"Loss\")\nax[0].legend()\n\nax[1].plot(history_dict['iou_coef'], label=\"Training IOU\", linewidth=3)\nax[1].plot(history_dict['val_iou_coef'], label=\"Validation IOU\", linewidth=3)\nax[1].set_xlabel(\"Epoch\")\nax[1].set_ylabel(\"IOU\")\nax[1].set_title(\"IOU\")\nax[1].legend()\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2021-12-25T14:47:51.476975Z","iopub.execute_input":"2021-12-25T14:47:51.477326Z","iopub.status.idle":"2021-12-25T14:47:52.062477Z","shell.execute_reply.started":"2021-12-25T14:47:51.477279Z","shell.execute_reply":"2021-12-25T14:47:52.061425Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Get the best model\n\nmodel.load_weights(config.WEIGHTS_PATH)","metadata":{"execution":{"iopub.status.busy":"2021-12-25T14:47:52.063976Z","iopub.execute_input":"2021-12-25T14:47:52.064881Z","iopub.status.idle":"2021-12-25T14:47:52.276498Z","shell.execute_reply.started":"2021-12-25T14:47:52.064803Z","shell.execute_reply":"2021-12-25T14:47:52.275595Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_ids = os.listdir(config.TEST_DIR)\n\ndef load_test_ds():\n    for image_id in test_ids:\n        image = tf.io.read_file(os.path.join(config.TEST_DIR, image_id))         \n        image = tf.image.decode_image(image, channels=3, dtype=tf.float32)\n        image = tf.image.resize(image, config.IMG_SHAPE)\n        image /= 255.0\n        yield image\n        \ntest_ds = (\n    tf.data.Dataset.from_generator(\n        load_test_ds, \n        output_types=tf.float32\n    )\n    .batch(3)\n)","metadata":{"execution":{"iopub.status.busy":"2021-12-25T14:47:52.278083Z","iopub.execute_input":"2021-12-25T14:47:52.278377Z","iopub.status.idle":"2021-12-25T14:47:52.310924Z","shell.execute_reply.started":"2021-12-25T14:47:52.278308Z","shell.execute_reply":"2021-12-25T14:47:52.310065Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds = model.predict(test_ds)","metadata":{"execution":{"iopub.status.busy":"2021-12-25T14:47:52.312596Z","iopub.execute_input":"2021-12-25T14:47:52.312943Z","iopub.status.idle":"2021-12-25T14:47:54.499064Z","shell.execute_reply.started":"2021-12-25T14:47:52.312899Z","shell.execute_reply":"2021-12-25T14:47:54.49801Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds = (preds > 0.5).astype(np.int32)","metadata":{"execution":{"iopub.status.busy":"2021-12-25T14:47:54.500603Z","iopub.execute_input":"2021-12-25T14:47:54.50096Z","iopub.status.idle":"2021-12-25T14:47:54.507287Z","shell.execute_reply.started":"2021-12-25T14:47:54.500919Z","shell.execute_reply":"2021-12-25T14:47:54.506022Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_iter = next(iter(test_ds))\n\nfig, ax = plt.subplots(3, 2, figsize=(20, 20))\n\nfor i in range(3):\n    ax[i][0].imshow(test_iter[i] * 255)\n    ax[i][0].axis('off')        \n    \n    ax[i][1].imshow(preds[i])\n    ax[i][1].axis('off')        \n        \nplt.show()","metadata":{"execution":{"iopub.status.busy":"2021-12-25T14:47:54.509145Z","iopub.execute_input":"2021-12-25T14:47:54.510367Z","iopub.status.idle":"2021-12-25T14:47:55.554435Z","shell.execute_reply.started":"2021-12-25T14:47:54.510313Z","shell.execute_reply":"2021-12-25T14:47:55.553612Z"},"trusted":true},"execution_count":null,"outputs":[]}]}