{"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":"## Description \nThis notebook contains a basic implementation of a Keras image segmentation model for the [Sartorius - Cell Instance Segmentation competition](https://www.kaggle.com/c/sartorius-cell-instance-segmentation/data) using \n[Segmentation Models: \"Python library with Neural Networks for Image Segmentation based on Keras and TensorFlow. \"](https://github.com/qubvel/segmentation_models). This library contains four models architectures for image segmentation: UNet, LinkNet, PSPNet, and FPN.\n\nHere a UNet with a DenseNet121 backbone is implemented, but you could try the other models (LinkNet, PSPNet, or FPN) using the 25 available backbones (such as ResNet, MobileNet, DenseNet, etc.). By now, this implementation is unable to achieve a good score (> 0.0), it only provides a way for test different segmentation models with the competition's dataset. It also uses Jaccard loss and IOU score for training. \n\n## Problem definition\nData: \n\n[Phase-contrast microscopy](https://en.wikipedia.org/wiki/Phase-contrast_microscopy) images of human neuronal cell  types along with annotations (labels) representing cell segmentations. \n\nAim: \n\nThe trained model should be able to predict the annotations for cell segmentation, including rare cell types (such as neuroblastoma cell line SH-SY5Y as discussed in the [competition description](https://www.kaggle.com/c/sartorius-cell-instance-segmentation/overview)). Annotations should be provided in run-length format (see functions rle_decode() and rle_encode below).\n\n## Approach \n[Segmentation models library](https://github.com/qubvel/segmentation_models): \n\n\"The main features of this library are:\n\n*     High level API (just two lines of code to create model for segmentation)\n*     4 models architectures for binary and multi-class image segmentation (including legendary Unet)\n*     25 available backbones for each architecture\n*     All backbones have pre-trained weights for faster and better convergence\n*     Helpful segmentation losses (Jaccard, Dice, Focal) and metrics (IoU, F-score)\"\n\n\nReference notebooks:\n1. [Cell Segmentation - Run Length Decoding](https://www.kaggle.com/ihelon/cell-segmentation-run-length-decoding).\n2. [Sartors TF starter](https://www.kaggle.com/barteksadlej123/sartors-tf-starter).\n3. [Positive score with Detectron 3/3 - Inference](https://www.kaggle.com/dragonzhang/positive-score-with-detectron-3-3-inference).\n\nOther references: \n1. [Tensorflow Tutorials: Image segmentation ](https://www.tensorflow.org/tutorials/images/segmentation).\n2. [Segmentation models documentation](https://segmentation-models.readthedocs.io/en/latest/).\n","metadata":{}},{"cell_type":"markdown","source":"## Workflow\n\n[1. Imports](#section-1)\n\n[2. Functions](#section-2)\n\n[3. Constants](#section-3)\n\n[4. Generate train and validation data sets](#section-4)\n\n[5. Define the model](#section-5)\n\n[6. Training](#section-6)\n\n[7. Test set predictions](#section-7)\n\n","metadata":{}},{"cell_type":"markdown","source":"<a id=\"section-1\"></a>\n### 1. Libraries and paths","metadata":{}},{"cell_type":"code","source":"# install Segmentation Models \n!pip install -U segmentation-models","metadata":{"execution":{"iopub.status.busy":"2021-11-08T18:45:06.448636Z","iopub.execute_input":"2021-11-08T18:45:06.448963Z","iopub.status.idle":"2021-11-08T18:45:15.944309Z","shell.execute_reply.started":"2021-11-08T18:45:06.448885Z","shell.execute_reply":"2021-11-08T18:45:15.943462Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# required libraries \nimport numpy as np \nimport pandas as pd\nimport os\nfrom pathlib import Path\nimport cv2\n\nimport tensorflow as tf\nfrom tensorflow import keras\nfrom tensorflow.keras.layers.experimental import preprocessing\n\nimport segmentation_models\nsegmentation_models.set_framework('tf.keras')\nfrom segmentation_models import Unet\nfrom segmentation_models import get_preprocessing\nfrom segmentation_models.losses import bce_jaccard_loss\nfrom segmentation_models.metrics import iou_score\n\nfrom IPython.display import clear_output\nimport matplotlib.pyplot as plt","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2021-11-08T18:45:15.946292Z","iopub.execute_input":"2021-11-08T18:45:15.946592Z","iopub.status.idle":"2021-11-08T18:45:21.449108Z","shell.execute_reply.started":"2021-11-08T18:45:15.946550Z","shell.execute_reply":"2021-11-08T18:45:21.448386Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"section-2\"></a>\n### 2. Functions","metadata":{}},{"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    ref: https://www.kaggle.com/inversion/run-length-decoding-quick-start\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]), dtype=np.float32)\n    for lo, hi in zip(starts, ends):\n        img[lo : hi] = color\n    return img.reshape(shape)\n\n\ndef rle_encode(img):\n    '''\n    img: numpy array, 1 - mask, 0 - background\n    Returns run length as string formated\n    ref: https://www.kaggle.com/dragonzhang/positive-score-with-detectron-3-3-inference\n    '''\n    pixels = img.flatten()\n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)\n\n\ndef get_mask(image_id, df):\n    '''\n    Uses rle_decode() to get ndarray from mask using image_id in dataframe (df).\n    ref: https://www.kaggle.com/barteksadlej123/sartors-tf-starter\n    '''\n    current = df[df[\"id\"] == image_id]\n    labels = current[\"annotation\"].tolist()\n    \n    mask = np.zeros((HEIGHT, WIDTH))\n    for label in labels:\n        mask += rle_decode(label, (HEIGHT, WIDTH))\n    mask = mask.clip(0, 1)\n    \n    return mask\n\n\n#  fix overlaps: \n\ndef check_overlap(msk):\n    '''\n    Checks if there are overlap in a mask (msk).\n    ref: https://www.kaggle.com/awsaf49/sartorius-fix-overlap\n    '''\n    msk = msk.astype(np.bool).astype(np.uint8)\n    return np.any(np.sum(msk, axis=-1)>1)\n\n\ndef fix_overlap(msk):\n    '''\n    Args:\n        mask: multi-channel mask, each channel is an instance of cell, shape:(520,704,None)\n    Returns:\n        multi-channel mask with non-overlapping values, shape:(520,704,None) \n    ref: https://www.kaggle.com/awsaf49/sartorius-fix-overlap\n    '''\n    msk = np.array(msk)\n    msk = np.pad(msk, [[0,0],[0,0],[1,0]])\n    ins_len = msk.shape[-1]\n    msk = np.argmax(msk,axis=-1)\n    msk = tf.keras.utils.to_categorical(msk, num_classes=ins_len)\n    msk = msk[...,1:]\n    msk = msk[...,np.any(msk, axis=(0,1))]\n    return msk\n\n\n# make predictions for test set: \n\ndef make_predictions(dataset, num, keras_model, check_overlaps=False):\n    '''\n    For a tf.Dataset, makes predictions for n=num (num =-1 or all_images takes all images in the dataset), \n    images using a keras_model. Returns a list of predicted masks, each as ndarray. \n    '''\n    predictions = []\n    if dataset:\n        for image in dataset.take(num):\n            image = image[None]\n            pred_mask = keras_model.predict(image)\n            # changes shape from (1,512,512,1) to (512,512)\n            pred_mask = pred_mask[0, :, :, 0]\n            # fix overlaps\n            if check_overlaps:\n                if check_overlap(msk=pred_mask):\n                    pred_mask = pred_mask[None]\n                    pred_mask = fix_overlap(msk=pred_mask)\n            # transforms ndarray values to 0s and 1s\n            pred_mask =  np.where( pred_mask > 0.5, 1, 0)\n            predictions.append(pred_mask)\n    return predictions\n\n\ndef display(display_list):\n    '''\n    Displays an example image and mask \n    '''\n    plt.figure(figsize=(20, 20))\n\n    title = ['Input Image', 'True Mask', 'Predicted Mask']\n\n    for i in range(len(display_list)):\n        plt.subplot(1, len(display_list), i+1)\n        plt.title(title[i])\n        plt.imshow(tf.keras.preprocessing.image.array_to_img(display_list[i]))\n        plt.axis('off')\n    plt.show()\n\n    \n# functions to visualize predictions:\n\ndef create_mask(pred_mask):\n    '''Converts predicted mask values to 0s and 1s'''\n    pred_mask = tf.where(pred_mask > 0.5,1,0)\n    return pred_mask\n\ndef show_predictions(keras_model, dataset=None, num=1):\n    '''\n    Shows N=num predictions examples from dataset and a keras_model \n    '''\n    if dataset:\n        for image, mask in dataset.take(num):\n            pred_mask = keras_model.predict(image)\n            display([image[0], mask[0], create_mask(pred_mask[0])])\n    else:\n        display([sample_image, sample_mask,\n                 create_mask(model.predict(sample_image[tf.newaxis, ...])[0])])\n","metadata":{"execution":{"iopub.status.busy":"2021-11-08T18:45:21.450466Z","iopub.execute_input":"2021-11-08T18:45:21.450968Z","iopub.status.idle":"2021-11-08T18:45:21.475958Z","shell.execute_reply.started":"2021-11-08T18:45:21.450929Z","shell.execute_reply":"2021-11-08T18:45:21.475256Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"section-3\"></a>\n### 3. Constants ","metadata":{}},{"cell_type":"code","source":"# constants\n\nDEBUG = False\n\nSEED = 123\nWIDTH, HEIGHT = 704, 520\nRESIZE_WIDTH, RESIZE_HEIGHT = 512, 512\nBATCH_SIZE = 4\nBUFFER_SIZE = 32\n\n## in a segmentation task each pixel is given a class \n## OUTPUT_CLASSES: number of classes that can be assigned to each pixel \nOUTPUT_CLASSES = 1\n\n# architecture backbone for segmentation model\nBACKBONE = 'densenet121'\n\nVAL_SPLIT = 0.2\n\nAUTO = tf.data.AUTOTUNE\n\nEPOCHS = 100","metadata":{"execution":{"iopub.status.busy":"2021-11-08T18:45:21.478490Z","iopub.execute_input":"2021-11-08T18:45:21.478866Z","iopub.status.idle":"2021-11-08T18:45:21.488249Z","shell.execute_reply.started":"2021-11-08T18:45:21.478830Z","shell.execute_reply":"2021-11-08T18:45:21.487443Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# paths\n\n# input\nDIR = '../input/sartorius-cell-instance-segmentation'\ntrain_csv = os.path.join(DIR,'train.csv') \ntrain_path =  os.path.join(DIR, 'train/')\ntest_path = os.path.join(DIR, 'test/')\n\n# output \ncsv_output = os.path.join('./', 'submission.csv') \nmodel_output = os.path.join('./', 'unet_keras_' + BACKBONE + '_backbone.h5')","metadata":{"execution":{"iopub.status.busy":"2021-11-08T18:45:21.489495Z","iopub.execute_input":"2021-11-08T18:45:21.489944Z","iopub.status.idle":"2021-11-08T18:45:21.498257Z","shell.execute_reply.started":"2021-11-08T18:45:21.489906Z","shell.execute_reply":"2021-11-08T18:45:21.497562Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"section-4\"></a>\n### 4. Generate train and validation data sets","metadata":{}},{"cell_type":"code","source":"# train and validation split\ntrain = pd.read_csv(train_csv)\ntrain.head()\n\nn_ids = train.id.nunique()\n\nif DEBUG:\n    unique_ids_train = list(set(train['id'].tolist()))[:BATCH_SIZE]\n    unique_ids_valid = list(set(train['id'].tolist()))[BATCH_SIZE:2*BATCH_SIZE]\nelse:\n    unique_ids_train = list(set(train['id'].tolist()))[:int(n_ids * (1 - VAL_SPLIT))]\n    unique_ids_valid = list(set(train['id'].tolist()))[int(n_ids * (1 - VAL_SPLIT)):]\n\n\ntemp = pd.DataFrame()\nfor sample_id in unique_ids_train:\n    query = train[train.id == sample_id]\n    temp = pd.concat([temp, query])\ntrain = temp\ntrain = train.reset_index(drop=True)\n\ntemp = pd.DataFrame()\nfor sample_id in unique_ids_valid:\n    query = train[train.id == sample_id]\n    temp = pd.concat([temp, query])\nvalid = temp\nvalid = train.reset_index(drop=True)\n    \nTRAIN_LENGTH = train['id'].nunique()\nSTEPS_PER_EPOCH = TRAIN_LENGTH // BATCH_SIZE\n\nVALID_LENGTH = valid['id'].nunique()\nVALIDATION_STEPS = VALID_LENGTH // BATCH_SIZE","metadata":{"execution":{"iopub.status.busy":"2021-11-08T18:45:21.499439Z","iopub.execute_input":"2021-11-08T18:45:21.500180Z","iopub.status.idle":"2021-11-08T18:45:29.952403Z","shell.execute_reply.started":"2021-11-08T18:45:21.500141Z","shell.execute_reply":"2021-11-08T18:45:29.951651Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# training/validation data generator \n\npreprocess_input = get_preprocessing(BACKBONE)\n\ndef train_generator(df):\n    image_ids = set(df['id'].tolist())\n    \n    for image_id in image_ids:\n        \n        image = cv2.imread(os.path.join(train_path, image_id) + '.png')\n        image = preprocess_input(image)\n        image = cv2.resize(image, (RESIZE_HEIGHT, RESIZE_WIDTH))\n\n        mask = get_mask(image_id, df)        \n        mask = cv2.resize(mask, (RESIZE_HEIGHT, RESIZE_WIDTH))\n        \n        mask = mask.reshape((*mask.shape, 1))\n        \n        image = image.astype(np.float32)\n        mask = mask.astype(np.float32)\n        \n        yield image, mask","metadata":{"execution":{"iopub.status.busy":"2021-11-08T18:45:29.953535Z","iopub.execute_input":"2021-11-08T18:45:29.953776Z","iopub.status.idle":"2021-11-08T18:45:29.960307Z","shell.execute_reply.started":"2021-11-08T18:45:29.953745Z","shell.execute_reply":"2021-11-08T18:45:29.959613Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# use the generator to get training and validation sets\ntrain_ds = tf.data.Dataset.from_generator(\n    lambda : train_generator(train), \n    output_types=(tf.float32, tf.float32),\n    output_shapes=((RESIZE_HEIGHT, RESIZE_WIDTH, 3), (RESIZE_HEIGHT, RESIZE_WIDTH, 1)))\n\nvalid_ds = tf.data.Dataset.from_generator(\n    lambda : train_generator(valid), \n    output_types=(tf.float32, tf.float32),\n    output_shapes=((RESIZE_HEIGHT, RESIZE_WIDTH, 3), (RESIZE_HEIGHT, RESIZE_WIDTH, 1)))\n","metadata":{"execution":{"iopub.status.busy":"2021-11-08T18:45:29.961495Z","iopub.execute_input":"2021-11-08T18:45:29.961903Z","iopub.status.idle":"2021-11-08T18:45:32.452521Z","shell.execute_reply.started":"2021-11-08T18:45:29.961866Z","shell.execute_reply":"2021-11-08T18:45:32.451829Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# \"the following class performs a simple augmentation by randomly-flipping an image\"\nclass Augment(tf.keras.layers.Layer):\n    def __init__(self, seed=SEED):\n        super().__init__()\n        \n        self.augment_inputs = preprocessing.RandomFlip('horizontal', seed=seed)\n        self.augment_labels = preprocessing.RandomFlip('horizontal', seed=seed)\n        \n    def call(self, inputs, labels):\n        inputs = self.augment_inputs(inputs)\n        labels = self.augment_labels(labels)\n        return inputs, labels","metadata":{"execution":{"iopub.status.busy":"2021-11-08T18:45:32.453741Z","iopub.execute_input":"2021-11-08T18:45:32.453989Z","iopub.status.idle":"2021-11-08T18:45:32.460359Z","shell.execute_reply.started":"2021-11-08T18:45:32.453957Z","shell.execute_reply":"2021-11-08T18:45:32.459562Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# \"build the input pipeline, applying the augmentation after batching the inputs\"\n# for augmentation only add .map(Augment()) after .repeat() \n\ntrain_ds = (\n    train_ds\n    .shuffle(BUFFER_SIZE)\n    .batch(BATCH_SIZE)\n    .repeat()\n    .map(Augment())\n    .prefetch(AUTO))\n\nvalid_ds = (\n    valid_ds\n    .batch(BATCH_SIZE)\n    .repeat()\n    .prefetch(AUTO))","metadata":{"execution":{"iopub.status.busy":"2021-11-08T18:45:32.463382Z","iopub.execute_input":"2021-11-08T18:45:32.463813Z","iopub.status.idle":"2021-11-08T18:45:32.617291Z","shell.execute_reply.started":"2021-11-08T18:45:32.463772Z","shell.execute_reply":"2021-11-08T18:45:32.616640Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# \"visualize an image example and its corresponding mask from the dataset\"    \nfor images, masks in train_ds.take(1):\n    sample_image, sample_mask = images[0], masks[0]\n    display([sample_image, sample_mask])","metadata":{"execution":{"iopub.status.busy":"2021-11-08T18:45:32.618542Z","iopub.execute_input":"2021-11-08T18:45:32.618817Z","iopub.status.idle":"2021-11-08T18:45:37.450703Z","shell.execute_reply.started":"2021-11-08T18:45:32.618785Z","shell.execute_reply":"2021-11-08T18:45:37.450061Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"section-5\"></a>\n### 5. Define the model ","metadata":{}},{"cell_type":"code","source":"# model defininition\n\n# \"encoder_freeze: if True set all layers of encoder (backbone model) as non-trainable\"  \nmodel = Unet(backbone_name=BACKBONE, \n                encoder_weights='imagenet', \n                activation='sigmoid',\n                classes=OUTPUT_CLASSES,\n                encoder_freeze=True)","metadata":{"execution":{"iopub.status.busy":"2021-11-08T18:45:37.452074Z","iopub.execute_input":"2021-11-08T18:45:37.452309Z","iopub.status.idle":"2021-11-08T18:45:40.621565Z","shell.execute_reply.started":"2021-11-08T18:45:37.452275Z","shell.execute_reply":"2021-11-08T18:45:40.620707Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# compile model \n\n# optimizer \nopt = keras.optimizers.Adam(learning_rate=1e-3)\n\n# loss from Segmentation Models\n# can be customized using sm.segmentation_models.losses.JaccardLoss()\n# metrics from Segmentation Models, \n# can be customized using sm.segmentation_models.metrics.IOUScore() function\nmodel.compile(optimizer=opt,\n              loss=bce_jaccard_loss,\n              metrics=[iou_score])","metadata":{"execution":{"iopub.status.busy":"2021-11-08T18:45:40.622948Z","iopub.execute_input":"2021-11-08T18:45:40.623190Z","iopub.status.idle":"2021-11-08T18:45:40.645245Z","shell.execute_reply.started":"2021-11-08T18:45:40.623157Z","shell.execute_reply":"2021-11-08T18:45:40.644551Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"section-6\"></a>\n### 6. Training","metadata":{}},{"cell_type":"code","source":"# training\n\n# \"observe how the model improves while it is training\"\nclass DisplayCallback(tf.keras.callbacks.Callback):\n    def __init__(self):\n        super().__init__()\n    def on_epoch_end(self, epoch, logs=None):\n        clear_output(wait=False)\n        show_predictions(keras_model=model)\n        print ('\\nSample Prediction after epoch {}\\n'.format(epoch+1))\n# display callback defined above\ndisplay_cb = DisplayCallback()\n\n\n# \"save the Keras model or model weights at some frequency\"\nmodel_checkpoint = tf.keras.callbacks.ModelCheckpoint(\n    model_output,\n    save_best_only=True,\n    save_weights_only=False,\n)\n\n# \"reduce learning rate when a metric has stopped improving\"\n# documentation: https://www.tensorflow.org/api_docs/python/tf/keras/callbacks/ReduceLROnPlateau\nlr_reduce = tf.keras.callbacks.ReduceLROnPlateau( monitor='val_loss', factor=0.1, patience=10, verbose=0,\n    mode='auto', min_delta=0.0001, cooldown=0, min_lr=0)\n\nmodel_history = model.fit(train_ds, epochs=EPOCHS,\n                          steps_per_epoch=STEPS_PER_EPOCH,\n                          validation_steps=VALIDATION_STEPS,\n                          validation_data=valid_ds,\n                          callbacks=[display_cb, model_checkpoint, lr_reduce])","metadata":{"execution":{"iopub.status.busy":"2021-11-08T18:45:40.647454Z","iopub.execute_input":"2021-11-08T18:45:40.647950Z","iopub.status.idle":"2021-11-08T21:21:12.278347Z","shell.execute_reply.started":"2021-11-08T18:45:40.647912Z","shell.execute_reply":"2021-11-08T21:21:12.277428Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plot training curve\nloss = model_history.history['loss']\nval_loss = model_history.history['val_loss']\nplt.figure()\nplt.plot(model_history.epoch, loss, 'r', label='Training loss')\nplt.plot(model_history.epoch, val_loss, 'bo', label='Validation loss')\nplt.title('Training and Validation Loss')\nplt.xlabel('Epoch')\nplt.ylabel('Loss Value')\nplt.ylim([0, 1])\nplt.legend()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2021-11-08T21:21:12.280943Z","iopub.execute_input":"2021-11-08T21:21:12.281319Z","iopub.status.idle":"2021-11-08T21:21:12.501003Z","shell.execute_reply.started":"2021-11-08T21:21:12.281263Z","shell.execute_reply":"2021-11-08T21:21:12.500336Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"section-7\"></a>\n### 7. Test set predictions","metadata":{}},{"cell_type":"code","source":"# test data generator \ntest_ids = [  os.path.join(test_path, each)  for each in os.listdir(test_path) if each.endswith('.png')]\ndef test_generator(image_ids):\n    for image_id in image_ids:\n        image = cv2.imread(image_id) \n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)        \n        image = cv2.resize(image, (RESIZE_HEIGHT, RESIZE_WIDTH))\n        image = image.astype(np.float32)\n        yield image","metadata":{"execution":{"iopub.status.busy":"2021-11-08T21:21:12.502053Z","iopub.execute_input":"2021-11-08T21:21:12.502281Z","iopub.status.idle":"2021-11-08T21:21:12.513393Z","shell.execute_reply.started":"2021-11-08T21:21:12.502249Z","shell.execute_reply":"2021-11-08T21:21:12.512726Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# test dataset from test data generator \ntest_ds = tf.data.Dataset.from_generator(\n    lambda : test_generator(test_ids), \n    output_types=(tf.float32),\n    output_shapes=((RESIZE_HEIGHT, RESIZE_WIDTH, 3)) )","metadata":{"execution":{"iopub.status.busy":"2021-11-08T21:21:12.514896Z","iopub.execute_input":"2021-11-08T21:21:12.515151Z","iopub.status.idle":"2021-11-08T21:21:12.539423Z","shell.execute_reply.started":"2021-11-08T21:21:12.515118Z","shell.execute_reply":"2021-11-08T21:21:12.538835Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# test image ids and predictions\ntest_predictions = make_predictions(dataset=test_ds, \n                                    num=len(test_ids), \n                                    keras_model=model,\n                                    check_overlaps=True)","metadata":{"execution":{"iopub.status.busy":"2021-11-08T21:21:12.542075Z","iopub.execute_input":"2021-11-08T21:21:12.542274Z","iopub.status.idle":"2021-11-08T21:21:13.181821Z","shell.execute_reply.started":"2021-11-08T21:21:12.542234Z","shell.execute_reply":"2021-11-08T21:21:13.181072Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# encode predections in the RL format\ntest_predictions = [rle_encode(mask) for mask in test_predictions] ","metadata":{"execution":{"iopub.status.busy":"2021-11-08T21:21:13.185101Z","iopub.execute_input":"2021-11-08T21:21:13.185309Z","iopub.status.idle":"2021-11-08T21:21:13.193019Z","shell.execute_reply.started":"2021-11-08T21:21:13.185284Z","shell.execute_reply":"2021-11-08T21:21:13.192353Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# transform full image paths to ids \ntest_ids = [Path(ID).stem for ID in test_ids]","metadata":{"execution":{"iopub.status.busy":"2021-11-08T21:21:13.194416Z","iopub.execute_input":"2021-11-08T21:21:13.194904Z","iopub.status.idle":"2021-11-08T21:21:13.202388Z","shell.execute_reply.started":"2021-11-08T21:21:13.194868Z","shell.execute_reply":"2021-11-08T21:21:13.201773Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# generate submission data frame \nsubmisssion = pd.DataFrame.from_dict({'id': test_ids, 'predicted': test_predictions} )\nsubmisssion = submisssion.sort_values( ['id'], ascending=True )\nprint(submisssion.head(), 'n')\nsubmisssion.to_csv(csv_output, index=False)","metadata":{"execution":{"iopub.status.busy":"2021-11-08T21:21:13.204120Z","iopub.execute_input":"2021-11-08T21:21:13.205643Z","iopub.status.idle":"2021-11-08T21:21:13.222450Z","shell.execute_reply.started":"2021-11-08T21:21:13.205604Z","shell.execute_reply":"2021-11-08T21:21:13.221418Z"},"trusted":true},"execution_count":null,"outputs":[]}]}