{"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":"\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2021-07-21T19:54:16.771889Z","iopub.execute_input":"2021-07-21T19:54:16.772326Z","iopub.status.idle":"2021-07-21T19:54:16.778557Z","shell.execute_reply.started":"2021-07-21T19:54:16.772255Z","shell.execute_reply":"2021-07-21T19:54:16.777796Z"}}},{"cell_type":"markdown","source":"# Image segmentation in Tensorflow","metadata":{}},{"cell_type":"markdown","source":"## Importing Library","metadata":{}},{"cell_type":"code","source":"!pip install git+https://github.com/tensorflow/examples.git\n!pip install -U tfds-nightly","metadata":{"execution":{"iopub.status.busy":"2021-07-23T12:15:10.496299Z","iopub.execute_input":"2021-07-23T12:15:10.496678Z","iopub.status.idle":"2021-07-23T12:15:33.134340Z","shell.execute_reply.started":"2021-07-23T12:15:10.496598Z","shell.execute_reply":"2021-07-23T12:15:33.133360Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import tensorflow as tf\nfrom tensorflow_examples.models.pix2pix import pix2pix\nfrom tensorflow.keras.models import Sequential\nfrom tensorflow.keras.layers import Dense, Flatten, Conv2D, MaxPooling2D\nfrom tensorflow.keras.optimizers import Adam\nimport numpy as np\nfrom tensorflow.keras.layers import Input, Conv2D,Dropout\nfrom tensorflow.keras.models import Sequential\nfrom tensorflow.keras.callbacks import EarlyStopping\nfrom IPython.display import clear_output\nimport random, re, math\nimport pandas as pd \nimport keras.backend as K\nfrom sklearn.model_selection import train_test_split\nimport matplotlib.pyplot as plt\nimport os\nimport zipfile","metadata":{"execution":{"iopub.status.busy":"2021-07-23T12:15:33.136192Z","iopub.execute_input":"2021-07-23T12:15:33.136535Z","iopub.status.idle":"2021-07-23T12:15:38.229283Z","shell.execute_reply.started":"2021-07-23T12:15:33.136495Z","shell.execute_reply":"2021-07-23T12:15:38.228378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Extracting the files","metadata":{}},{"cell_type":"code","source":"\ntrain_img_zip = '/kaggle/input/carvana-image-masking-challenge/train.zip'\ntrain_mask_zip='/kaggle/input/carvana-image-masking-challenge/train_masks.zip'\n\ntrain_img = zipfile.ZipFile(train_img_zip, 'r')\ntrain_img.extractall('/kaggle/working')\ntrain_img.close()\n\ntrain_mask = zipfile.ZipFile(train_mask_zip, 'r')\ntrain_mask.extractall('/kaggle/working')\ntrain_mask.close()","metadata":{"execution":{"iopub.status.busy":"2021-07-23T12:15:38.231010Z","iopub.execute_input":"2021-07-23T12:15:38.231372Z","iopub.status.idle":"2021-07-23T12:15:50.224168Z","shell.execute_reply.started":"2021-07-23T12:15:38.231344Z","shell.execute_reply":"2021-07-23T12:15:50.223288Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ntrain_dir = os.path.join('/kaggle/working/train')\ntrain_mask_dir = os.path.join('/kaggle/working/train_masks')\n","metadata":{"execution":{"iopub.status.busy":"2021-07-23T12:15:50.225614Z","iopub.execute_input":"2021-07-23T12:15:50.225954Z","iopub.status.idle":"2021-07-23T12:15:50.231103Z","shell.execute_reply.started":"2021-07-23T12:15:50.225918Z","shell.execute_reply":"2021-07-23T12:15:50.229475Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Let's take a look at some image examples and their correponding mask from the dataset.","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport matplotlib.image as mpimg\n\nplt.subplots(2, 5, figsize=(60, 20))\npic_index = 0\n\n\npic_index += 4\nnext_car_pix = [os.path.join(train_dir, fname) \n                for fname in sorted(os.listdir(train_dir))[pic_index-4:pic_index]]\nnext_mask_pix = [os.path.join(train_mask_dir, fname) \n                for fname in sorted(os.listdir(train_mask_dir))[pic_index-4:pic_index]]\n\nfor i, img_path in enumerate(next_car_pix+next_mask_pix):\n  # Set up subplot; subplot indices start at 1\n  sp = plt.subplot(2, 4, i + 1)\n  sp.axis('Off') # Don't show axes (or gridlines)\n\n  img = mpimg.imread(img_path)\n  plt.imshow(img)\n\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2021-07-23T12:15:50.232581Z","iopub.execute_input":"2021-07-23T12:15:50.233051Z","iopub.status.idle":"2021-07-23T12:15:53.322308Z","shell.execute_reply.started":"2021-07-23T12:15:50.233009Z","shell.execute_reply":"2021-07-23T12:15:53.321075Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## the shape of one random image","metadata":{}},{"cell_type":"code","source":"img = mpimg.imread(train_dir+'/28d9a149cb02_12.jpg')\n\nimg.shape","metadata":{"execution":{"iopub.status.busy":"2021-07-23T12:15:53.323693Z","iopub.execute_input":"2021-07-23T12:15:53.324025Z","iopub.status.idle":"2021-07-23T12:15:53.398795Z","shell.execute_reply.started":"2021-07-23T12:15:53.323975Z","shell.execute_reply":"2021-07-23T12:15:53.397728Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Making a dataframe to use in tensorflow dataset","metadata":{}},{"cell_type":"markdown","source":"#### We should make the image_path and mask_path have the same id:","metadata":{}},{"cell_type":"code","source":"def creat_dataframe(img_path,mask_path):\n\n    car_ids = []\n    car_paths = []\n    mask_ids=[]\n    mask_paths=[]\n    for p in (img_path,mask_path):\n        for dirname, _, filenames in os.walk(p):\n            for filename in filenames:\n                path = os.path.join(dirname, filename)  \n                if p==img_path:\n                    car_paths.append(path)\n                    car_id = filename.split(\".\")[0]\n                    car_ids.append(car_id)\n                    df=pd.DataFrame(data = {\"id\": car_ids, \"img_path\": car_paths}).set_index('id')\n                else:\n                    mask_paths.append(path)\n                    mask_id = filename.split(\".\")[0]\n                    mask_id = mask_id.split(\"_mask\")[0]\n                    mask_ids.append(mask_id)\n                    df_mask=pd.DataFrame(data = {\"id\": mask_ids, \"mask_path\": mask_paths}).set_index('id')\n                    \n    df[\"mask_path\"] = df_mask[\"mask_path\"]\n         \n    return df","metadata":{"execution":{"iopub.status.busy":"2021-07-23T12:15:53.400861Z","iopub.execute_input":"2021-07-23T12:15:53.401556Z","iopub.status.idle":"2021-07-23T12:15:53.609043Z","shell.execute_reply.started":"2021-07-23T12:15:53.401513Z","shell.execute_reply":"2021-07-23T12:15:53.607803Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df=creat_dataframe('/kaggle/working/train','/kaggle/working/train_masks')\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2021-07-23T12:15:53.612649Z","iopub.execute_input":"2021-07-23T12:15:53.613381Z","iopub.status.idle":"2021-07-23T12:16:10.173188Z","shell.execute_reply.started":"2021-07-23T12:15:53.613336Z","shell.execute_reply":"2021-07-23T12:16:10.172289Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.info()","metadata":{"execution":{"iopub.status.busy":"2021-07-23T12:16:10.174885Z","iopub.execute_input":"2021-07-23T12:16:10.175158Z","iopub.status.idle":"2021-07-23T12:16:10.189342Z","shell.execute_reply.started":"2021-07-23T12:16:10.175132Z","shell.execute_reply":"2021-07-23T12:16:10.188224Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Configuration","metadata":{}},{"cell_type":"code","source":"AUTO = tf.data.experimental.AUTOTUNE\n\n# Configuration\nIMAGE_SIZE = [256, 256]\nEPOCHS = 40\nSEED = 777\nBATCH_SIZE = 16 ","metadata":{"execution":{"iopub.status.busy":"2021-07-23T12:16:10.190907Z","iopub.execute_input":"2021-07-23T12:16:10.191328Z","iopub.status.idle":"2021-07-23T12:16:10.197103Z","shell.execute_reply.started":"2021-07-23T12:16:10.191286Z","shell.execute_reply":"2021-07-23T12:16:10.195921Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset  functions","metadata":{"execution":{"iopub.status.busy":"2021-07-22T20:33:33.271407Z","iopub.execute_input":"2021-07-22T20:33:33.27183Z","iopub.status.idle":"2021-07-22T20:33:33.276408Z","shell.execute_reply.started":"2021-07-22T20:33:33.271783Z","shell.execute_reply":"2021-07-22T20:33:33.275237Z"}}},{"cell_type":"code","source":"def flip(image,mask):\n    \n    image = tf.image.flip_left_right(image)\n    mask = tf.image.flip_left_right(mask)\n    \n    return image,mask\n\ndef get_mat(rotation, shear, height_zoom, width_zoom, height_shift, width_shift):\n        \n    rotation = math.pi * rotation / 180.\n    shear = math.pi * shear / 180.\n    # ROTATION MATRIX\n    c1 = tf.math.cos(rotation)\n    s1 = tf.math.sin(rotation)\n    one = tf.constant([1],dtype='float32')\n    zero = tf.constant([0],dtype='float32')\n    rotation_matrix = tf.reshape( tf.concat([c1,s1,zero, -s1,c1,zero, zero,zero,one],axis=0),[3,3] )\n        \n    # SHEAR MATRIX\n    c2 = tf.math.cos(shear)\n    s2 = tf.math.sin(shear)\n    shear_matrix = tf.reshape( tf.concat([one,s2,zero, zero,c2,zero, zero,zero,one],axis=0),[3,3] )    \n    \n    # ZOOM MATRIX\n    zoom_matrix = tf.reshape( tf.concat([one/height_zoom,zero,zero, zero,one/width_zoom,zero, zero,zero,one],axis=0),[3,3] )\n    \n    # SHIFT MATRIX\n    shift_matrix = tf.reshape( tf.concat([one,zero,height_shift, zero,one,width_shift, zero,zero,one],axis=0),[3,3] )\n    \n    return K.dot(K.dot(rotation_matrix, shear_matrix), K.dot(zoom_matrix, shift_matrix))\n  \n    \ndef transform(image,mask):\n    # is borrowed from https://www.kaggle.com/cdeotte/rotation-augmentation-gpu-tpu-0-96\n    DIM = IMAGE_SIZE[0]\n    XDIM = DIM%2 \n    \n    rot = 10. * tf.random.normal([1],dtype='float32')\n    shr = 2. * tf.random.normal([1],dtype='float32') \n    h_zoom = 1.0 + tf.random.normal([1],dtype='float32')/10.\n    w_zoom = 1.0 + tf.random.normal([1],dtype='float32')/10.\n    h_shift = 10. * tf.random.normal([1],dtype='float32') \n    w_shift = 10. * tf.random.normal([1],dtype='float32') \n  \n    m = get_mat(rot,shr,h_zoom,w_zoom,h_shift,w_shift) \n\n    x = tf.repeat( tf.range(DIM//2,-DIM//2,-1), DIM )\n    y = tf.tile( tf.range(-DIM//2,DIM//2),[DIM] )\n    z = tf.ones([DIM*DIM],dtype='int32')\n    idx = tf.stack( [x,y,z] )\n    \n    idx2 = K.dot(m,tf.cast(idx,dtype='float32'))\n    idx2 = K.cast(idx2,dtype='int32')\n    idx2 = K.clip(idx2,-DIM//2+XDIM+1,DIM//2)\n    \n    idx3 = tf.stack( [DIM//2-idx2[0,], DIM//2-1+idx2[1,]] )\n    d = tf.gather_nd(image,tf.transpose(idx3))\n    m = tf.gather_nd(mask,tf.transpose(idx3))\n        \n    return tf.reshape(d,[DIM,DIM,3]),tf.reshape(m,[DIM,DIM,1])\n\n\n\n\ndef read_image_and_mask(image_path, mask_path=None,resize=IMAGE_SIZE):\n    \n    image=tf.io.read_file(image_path)\n    image=tf.image.decode_jpeg(image, channels=3)\n    image=tf.image.resize(image, IMAGE_SIZE)\n    image = tf.cast(image, dtype=tf.float32)/255.\n    if not mask_path is None:\n        mask=tf.io.read_file(mask_path)\n        mask=tf.image.decode_jpeg(mask, channels=3)\n        mask=mask[:,:,:1]\n        mask=tf.image.resize(mask, IMAGE_SIZE)\n        mask = tf.cast(mask, dtype=tf.float32)/255.\n        return image, mask\n    return image\n\n\n\n\ndef get_training_dataset(df):\n    \n    training_dataset = tf.data.Dataset.from_tensor_slices((df[\"img_path\"].values, df[\"mask_path\"].values))\n    training_dataset = training_dataset.map(read_image_and_mask,num_parallel_calls=AUTO)\n    training_dataset = training_dataset.map(flip,num_parallel_calls=AUTO)\n    training_dataset = training_dataset.map(transform,num_parallel_calls=AUTO)\n    training_dataset = training_dataset.shuffle(512, reshuffle_each_iteration=True)\n    training_dataset = training_dataset.batch(BATCH_SIZE)\n    training_dataset = training_dataset.repeat()\n    training_dataset = training_dataset.prefetch(AUTO)\n\n    return training_dataset\n\n\ndef get_validation_dataset(df):\n  \n  validation_dataset = tf.data.Dataset.from_tensor_slices((df[\"img_path\"].values, df[\"mask_path\"].values))\n  validation_dataset = validation_dataset.map(read_image_and_mask)\n  validation_dataset = validation_dataset.batch(BATCH_SIZE)\n  \n\n  return validation_dataset\n\n\ndef get_test_dataset(images):\n  \n  test_dataset = tf.data.Dataset.from_tensor_slices((images))\n  test_dataset = test_dataset.map(read_image_and_mask)\n  test_dataset = test_dataset.batch(10, drop_remainder=True)\n\n  return test_dataset","metadata":{"execution":{"iopub.status.busy":"2021-07-23T12:16:10.199022Z","iopub.execute_input":"2021-07-23T12:16:10.199710Z","iopub.status.idle":"2021-07-23T12:16:10.229306Z","shell.execute_reply.started":"2021-07-23T12:16:10.199667Z","shell.execute_reply":"2021-07-23T12:16:10.228077Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tr_df, val_df = train_test_split(creat_dataframe('/kaggle/working/train','/kaggle/working/train_masks'), random_state=SEED, test_size=.25)\ntr_dataset = get_training_dataset(tr_df)\nval_dataset = get_validation_dataset(val_df)","metadata":{"execution":{"iopub.status.busy":"2021-07-23T12:16:10.230989Z","iopub.execute_input":"2021-07-23T12:16:10.231517Z","iopub.status.idle":"2021-07-23T12:16:28.900113Z","shell.execute_reply.started":"2021-07-23T12:16:10.231470Z","shell.execute_reply":"2021-07-23T12:16:28.899250Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Display Example Augmentation¶with mask\n","metadata":{}},{"cell_type":"code","source":"row = 2; col = 4;\nall_elements = tr_dataset.unbatch()\none_element = tf.data.Dataset.from_tensors( next(iter(all_elements)) )\naugmented_element = one_element.repeat().map(transform).batch(row*col)\n\nfor (img,mask) in augmented_element:\n    plt.figure(figsize=(15,int(15*row/col)))\n    for j in range(row*col):\n        plt.subplot(row,col,j+1)\n        plt.axis('off')\n        plt.imshow(img[j,])\n        plt.imshow(mask[j,:,:,0],alpha=.5,cmap='Reds')\n    plt.show()\n    break","metadata":{"execution":{"iopub.status.busy":"2021-07-23T12:16:28.901478Z","iopub.execute_input":"2021-07-23T12:16:28.901820Z","iopub.status.idle":"2021-07-23T12:16:44.011908Z","shell.execute_reply.started":"2021-07-23T12:16:28.901783Z","shell.execute_reply":"2021-07-23T12:16:44.010860Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Visualization function ","metadata":{}},{"cell_type":"code","source":"def display(display_list):\n    plt.figure(figsize=(15, 15))\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()","metadata":{"execution":{"iopub.status.busy":"2021-07-23T12:16:44.013409Z","iopub.execute_input":"2021-07-23T12:16:44.013729Z","iopub.status.idle":"2021-07-23T12:16:44.021806Z","shell.execute_reply.started":"2021-07-23T12:16:44.013697Z","shell.execute_reply":"2021-07-23T12:16:44.020397Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def batch_predict(image,model):\n    preds=model.predict(image)  \n    threshold = 0.5\n    preds[preds > threshold] = 1.0\n    preds[preds <= threshold] = 0.0 \n    return preds\n\ndef vis_compare(dataset=val_dataset,num_case=1):\n       \n    for sample in dataset.take(1):\n        image, label = sample[0].numpy(), sample[1].numpy()\n    preds=batch_predict(image,model)\n    if num_case>1:\n        cases=[j for j in np.random.choice(image.shape[0],size=num_case,replace=False)]    \n        for i in cases:\n            truth=(image[i],label[i])\n            pred=(image[i],preds[i])\n            print(f\"case_number_{i}\")\n            display([image,label,preds])\n            print('\\n')\n            print(464*'*')\n            print('\\n')\n    else:\n        truth=(image[0],label[0])\n        pred=(image[0],preds[0])\n        display([image[0],label[0],preds[0]])\n            \n    \n    \n    plt.show() ","metadata":{"execution":{"iopub.status.busy":"2021-07-23T12:16:44.023890Z","iopub.execute_input":"2021-07-23T12:16:44.024354Z","iopub.status.idle":"2021-07-23T12:16:44.222698Z","shell.execute_reply.started":"2021-07-23T12:16:44.024314Z","shell.execute_reply":"2021-07-23T12:16:44.220431Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## useful callback class","metadata":{}},{"cell_type":"markdown","source":"#### Let's observe how the model improves while it is training. To accomplish this task, a callback function is defined below. \n","metadata":{"execution":{"iopub.status.busy":"2021-07-22T19:54:24.08352Z","iopub.execute_input":"2021-07-22T19:54:24.08402Z","iopub.status.idle":"2021-07-22T19:54:24.11263Z","shell.execute_reply.started":"2021-07-22T19:54:24.083983Z","shell.execute_reply":"2021-07-22T19:54:24.110539Z"}}},{"cell_type":"code","source":"class DisplayCallback(tf.keras.callbacks.Callback):\n  def on_epoch_end(self, epoch, logs=None):\n    clear_output(wait=True)\n    vis_compare()\n    print ('\\nSample Prediction after epoch {}\\n'.format(epoch+1))\n","metadata":{"execution":{"iopub.status.busy":"2021-07-23T12:16:44.224029Z","iopub.execute_input":"2021-07-23T12:16:44.224412Z","iopub.status.idle":"2021-07-23T12:16:44.231935Z","shell.execute_reply.started":"2021-07-23T12:16:44.224362Z","shell.execute_reply":"2021-07-23T12:16:44.230834Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class EarlyStoppingAtMinLoss(tf.keras.callbacks.Callback):\n    \"\"\"Stop training when the loss is at its min, i.e. the loss stops decreasing.\n\n  Arguments:\n      patience: Number of epochs to wait after min has been hit. After this\n      number of no improvement, training stops.\n  \"\"\"\n\n    def __init__(self, patience=3):\n        super(EarlyStoppingAtMinLoss, self).__init__()\n        self.patience = patience\n        # best_weights to store the weights at which the minimum loss occurs.\n        self.best_weights = None\n\n    def on_train_begin(self, logs=None):\n        # The number of epoch it has waited when loss is no longer minimum.\n        self.wait = 0\n        # The epoch the training stops at.\n        self.stopped_epoch = 0\n        # Initialize the best as infinity.\n        self.best = np.Inf\n\n    def on_epoch_end(self, epoch, logs=None):\n        current = logs.get(\"val_loss\")\n        if np.less(current, self.best):\n            self.best = current\n            self.wait = 0\n            # Record the best weights if current results is better (less).\n            self.best_weights = self.model.get_weights()\n        else:\n            self.wait += 1\n            if self.wait >= self.patience:\n                self.stopped_epoch = epoch\n                self.model.stop_training = True\n                print(\"Restoring model weights from the end of the best epoch.\")\n                self.model.set_weights(self.best_weights)\n\n    def on_train_end(self, logs=None):\n        if self.stopped_epoch > 0:\n            print(\"Epoch %05d: early stopping\" % (self.stopped_epoch + 1))\n            \n","metadata":{"execution":{"iopub.status.busy":"2021-07-23T12:16:44.233519Z","iopub.execute_input":"2021-07-23T12:16:44.234037Z","iopub.status.idle":"2021-07-23T12:16:44.246079Z","shell.execute_reply.started":"2021-07-23T12:16:44.233983Z","shell.execute_reply":"2021-07-23T12:16:44.244785Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Loss function","metadata":{}},{"cell_type":"code","source":"def dice_coef(y_true, y_pred, smooth=1):\n    intersection = K.sum(y_true * y_pred, axis=[1,2,3])\n    union = K.sum(y_true, axis=[1,2,3]) + K.sum(y_pred, axis=[1,2,3])\n    return K.mean( (2. * intersection + smooth) / (union + smooth), axis=0)\n\ndef dice_loss(in_gt, in_pred):\n    return 1-dice_coef(in_gt, in_pred)","metadata":{"execution":{"iopub.status.busy":"2021-07-23T12:16:44.247528Z","iopub.execute_input":"2021-07-23T12:16:44.248342Z","iopub.status.idle":"2021-07-23T12:16:44.255936Z","shell.execute_reply.started":"2021-07-23T12:16:44.248300Z","shell.execute_reply":"2021-07-23T12:16:44.255060Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define the model","metadata":{"execution":{"iopub.status.busy":"2021-07-21T23:49:38.616018Z","iopub.execute_input":"2021-07-21T23:49:38.616426Z","iopub.status.idle":"2021-07-21T23:49:38.620698Z","shell.execute_reply.started":"2021-07-21T23:49:38.616392Z","shell.execute_reply":"2021-07-21T23:49:38.619838Z"}}},{"cell_type":"markdown","source":"\nThe model being used here is a modified U-Net. A U-Net consists of an encoder (downsampler) and decoder (upsampler). In-order to learn robust features, and reduce the number of trainable parameters, a pretrained model can be used as the encoder. Thus, the encoder for this task will be a pretrained MobileNetV2 model, whose intermediate outputs will be used, and the decoder will be the upsample block already implemented in TensorFlow Examples in the Pix2pix tutorial.","metadata":{}},{"cell_type":"code","source":"OUTPUT_CHANNELS = 1\nbase_model = tf.keras.applications.MobileNetV2(input_shape=[256, 256, 3], include_top=False)\n\n\n\nlayer_names = [\n    'block_1_expand_relu',   # 64x64\n    'block_3_expand_relu',   # 32x32\n    'block_6_expand_relu',   # 16x16\n    'block_13_expand_relu',  # 8x8\n    'block_16_project',      # 4x4\n]\n\n\nbase_model_outputs = [base_model.get_layer(name).output for name in layer_names]\n\n\ndown_stack = tf.keras.Model(inputs=base_model.input, outputs=base_model_outputs)\n\ndown_stack.trainable = False\n\nup_stack = [\n    pix2pix.upsample(512, 3),  # 4x4 -> 8x8\n    pix2pix.upsample(256, 3),  # 8x8 -> 16x16\n    pix2pix.upsample(128, 3),  # 16x16 -> 32x32\n    pix2pix.upsample(64, 3),   # 32x32 -> 64x64\n]\n\ndef unet_model(output_channels):\n    inputs = tf.keras.layers.Input(shape=[256, 256, 3])\n\n    skips = down_stack(inputs)\n    x = skips[-1]\n    skips = reversed(skips[:-1])\n\n    for up, skip in zip(up_stack, skips):\n        x = up(x)\n        concat = tf.keras.layers.Concatenate()\n        x = concat([x, skip])\n\n    last = tf.keras.layers.Conv2DTranspose(\n      output_channels, 3, strides=2, activation='sigmoid',\n      padding='same')  \n\n    x = last(x)\n    model = tf.keras.Model(inputs=inputs, outputs=x)\n    model.compile(optimizer='adam',loss = dice_loss,\n              metrics=[dice_coef,'binary_accuracy'])\n\n    return model\n\n\n\n\n\nmodel = unet_model(1)\n\n\n\n\n\n\n","metadata":{"execution":{"iopub.status.busy":"2021-07-23T12:16:44.257258Z","iopub.execute_input":"2021-07-23T12:16:44.257719Z","iopub.status.idle":"2021-07-23T12:16:45.871770Z","shell.execute_reply.started":"2021-07-23T12:16:44.257665Z","shell.execute_reply":"2021-07-23T12:16:45.870915Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train the model","metadata":{}},{"cell_type":"code","source":"steps_per_epoch = len(tr_df)//BATCH_SIZE\n\n\nhistory = model.fit(tr_dataset,\n                          steps_per_epoch=steps_per_epoch, validation_data=val_dataset, epochs=40,\n                          callbacks=[EarlyStoppingAtMinLoss()]) #,DisplayCallback()","metadata":{"execution":{"iopub.status.busy":"2021-07-23T12:16:45.873051Z","iopub.execute_input":"2021-07-23T12:16:45.873396Z","iopub.status.idle":"2021-07-23T12:51:45.386199Z","shell.execute_reply.started":"2021-07-23T12:16:45.873360Z","shell.execute_reply":"2021-07-23T12:51:45.384226Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Evaluation","metadata":{}},{"cell_type":"code","source":"model.evaluate(val_dataset)","metadata":{"execution":{"iopub.status.busy":"2021-07-23T12:51:45.388111Z","iopub.execute_input":"2021-07-23T12:51:45.388425Z","iopub.status.idle":"2021-07-23T12:52:19.127549Z","shell.execute_reply.started":"2021-07-23T12:51:45.388364Z","shell.execute_reply":"2021-07-23T12:52:19.126701Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def vis_compare(dataset=val_dataset,num_case=1):\n       \n    for sample in dataset.take(1):\n        image, label = sample[0].numpy(), sample[1].numpy()\n    preds=model.predict(image)\n    if num_case>1:\n        cases=[j for j in np.random.choice(image.shape[0],size=num_case,replace=False)]    \n        for i in cases:\n            truth=(image[i],label[i])\n            pred=(image[i],preds[i])\n            print(f\"case_number_{i}\")\n            display([image[i],label[i],preds[i]])\n            print('\\n')\n            print(464*'*')\n            print('\\n')\n    else:\n        truth=(image[0],label[0])\n        pred=(image[0],preds[0])\n        display([image[0],label[0],preds[0]])\n            \n    \n    \n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2021-07-23T12:52:19.128864Z","iopub.execute_input":"2021-07-23T12:52:19.129294Z","iopub.status.idle":"2021-07-23T12:52:19.215419Z","shell.execute_reply.started":"2021-07-23T12:52:19.129255Z","shell.execute_reply":"2021-07-23T12:52:19.214408Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"vis_compare(dataset=val_dataset,num_case=10)","metadata":{"execution":{"iopub.status.busy":"2021-07-23T12:52:19.221472Z","iopub.execute_input":"2021-07-23T12:52:19.221775Z","iopub.status.idle":"2021-07-23T12:52:23.564291Z","shell.execute_reply.started":"2021-07-23T12:52:19.221727Z","shell.execute_reply":"2021-07-23T12:52:23.563383Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}