{"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":"# Image Masking","metadata":{}},{"cell_type":"markdown","source":"## Preprocessing","metadata":{}},{"cell_type":"code","source":"import cv2\nimport matplotlib.pyplot as pyplot \nimport numpy\nimport os\nimport tensorflow\n\nos.environ['TF_CPP_MIN_LOG_LEVEL'] = '3'\n\nclass Data:\n    \"\"\"\n    TODO: docstring\n    \"\"\"\n    def __init__(self, batch_size):\n        \"\"\"\n        TODO: docstring\n        \"\"\"\n        self.batch_size = batch_size\n\n    def data_generator(self, image_list, mask_list, split='train'):\n        \"\"\"\n        TODO: docstring\n        \"\"\"\n        dataset = tensorflow.data.Dataset.from_tensor_slices((image_list, mask_list))\n        dataset = dataset.shuffle(8 * self.batch_size) if split == 'train' else dataset \n        dataset = dataset.map(self.load_data, num_parallel_calls=tensorflow.data.AUTOTUNE)\n        dataset = dataset.batch(self.batch_size, drop_remainder=True)\n        dataset = dataset.prefetch(tensorflow.data.AUTOTUNE)\n        return dataset\n\n    def load_data(self, image_list, mask_list):\n        \"\"\"\n        TODO: docstring\n        \"\"\"\n        image = self.read_image(image_list)\n        mask = self.read_image(mask_list, mask=True)\n        return image, mask\n\n    def read_image(self, image_path, mask=False):\n        \"\"\"\n        TODO: docstring\n        \"\"\"\n        image = tensorflow.io.read_file(image_path)\n        if mask:\n            image = tensorflow.image.decode_jpeg(image, channels=3)\n            image.set_shape([None, None, 3])\n            image = tensorflow.image.resize(images=image, size=[IMAGE_SIZE, IMAGE_SIZE])\n            #image = tensorflow.cast(image, tensorflow.int32)\n            image = image[:,:,:1]    \n            image = tensorflow.math.sign(image)\n        else:\n            image = tensorflow.image.decode_jpeg(image, channels=3)\n            image.set_shape([None, None, 3])\n            image = tensorflow.image.resize(images=image, size=[IMAGE_SIZE, IMAGE_SIZE])\n            #image = image / 255.0\n            image = tensorflow.cast(image, tensorflow.float32) / 255.0\n        return image\n    \ndef visualize(**images):\n    \"\"\"\n    Plot images in one row.\n    \"\"\"\n    n = len(images)\n    pyplot.figure(figsize=(16, 5))\n    for i, (name, image) in enumerate(images.items()):\n        pyplot.subplot(1, n, i + 1)\n        pyplot.xticks([])\n        pyplot.yticks([])\n        pyplot.title(' '.join(name.split('_')).title())\n        pyplot.imshow(image)\n    pyplot.show()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"exts = ('jpg', 'JPG', 'png', 'PNG', 'tif', 'gif', 'ppm')\ninput_data = '/content/drive/MyDrive/Projects/Kaggle/Image Masking/image_masking/data/train'\nimages = sorted([\n    os.path.join(input_data, fname)\n    for fname in os.listdir(input_data)\n    if fname.endswith(exts) and not fname.startswith('.')])\ntarget_data = '/content/drive/MyDrive/Projects/Kaggle/Image Masking/image_masking/data/train_masks'\nmasks = sorted([\n    os.path.join(target_data, fname)\n    for fname in os.listdir(target_data)\n    if fname.endswith(exts) and not fname.startswith('.')])\nprint('number of samples:', len(images), len(masks))\nfor input_path, target_path in zip(images[:10], masks[:10]):\n    print(input_path[-32:], '|', target_path[-31:], '|', numpy.unique(\n        cv2.imread(target_path)))\nIMAGE_SIZE = 128\nBATCH_SIZE = 86\ndata = Data(BATCH_SIZE)\ntrain_dataset = data.data_generator(images, masks)\nimage, mask = next(iter(train_dataset.take(1)))\nprint(image.shape, mask.shape)\nfor (img, msk) in zip(image[:5], mask[:5]):\n    print(mask.numpy().min(), mask.numpy().max())\n    #visualize(image=img.numpy(), gt_mask=msk.numpy())","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"code","source":"class DisplayCallback(tensorflow.keras.callbacks.Callback):\n    \"\"\"\n    TODO: docstring\n    \"\"\"\n    def __init__(self, dataset, epoch_interval=5):\n        \"\"\"\n        TODO: docstring\n        \"\"\"\n        self.dataset = dataset\n        self.epoch_interval = epoch_interval\n        \n    def create_mask(self, pred_mask):\n        \"\"\"\n        TODO: docstring\n        \"\"\"\n        pred_mask = (pred_mask > 0.5).astype('int32')\n        return pred_mask[0]\n            \n    def display(self, display_list, extra_title=''):\n        \"\"\"\n        TODO: docstring\n        \"\"\"\n        pyplot.figure(figsize=(15, 15))\n        title = ['Input Image', 'True Mask', 'Predicted Mask']\n        if len(display_list) > len(title):\n            title.append(extra_title)\n        for i in range(len(display_list)):\n            pyplot.subplot(1, len(display_list), i+1)\n            pyplot.title(title[i])\n            pyplot.imshow(display_list[i])\n            pyplot.axis('off')\n        pyplot.show()\n\n    def on_epoch_end(self, epoch, logs=None):\n        \"\"\"\n        TODO: docstring\n        \"\"\"\n        if epoch and epoch % self.epoch_interval == 0:\n            self.show_predictions(self.dataset)\n            print ('\\nsample prediction after epoch {}\\n'.format(epoch+1))\n    \n    def show_predictions(self, dataset, num=1):\n        \"\"\"\n        TODO: docstring\n        \"\"\"\n        for image, mask in dataset.take(num):\n            pred_mask = model.predict(image)\n            self.display([image[0], mask[0], self.create_mask(pred_mask)])\n            \ndef get_model(img_size, num_classes=1):\n    \"\"\"\n    TODO: docstring\n    \"\"\"\n    inputs = tensorflow.keras.Input(shape=img_size+(3,))\n    # downsampling\n    x = tensorflow.keras.layers.Conv2D(32, 3, strides=2, padding='same')(inputs)\n    x = tensorflow.keras.layers.BatchNormalization()(x)\n    x = tensorflow.keras.layers.Activation('relu')(x)\n    previous_block_activation = x\n    for filters in [64, 128, 256]:\n        x = tensorflow.keras.layers.Activation('relu')(x)\n        x = tensorflow.keras.layers.SeparableConv2D(filters, 3, padding='same')(x)\n        x = tensorflow.keras.layers.BatchNormalization()(x)\n        x = tensorflow.keras.layers.Activation('relu')(x)\n        x = tensorflow.keras.layers.SeparableConv2D(filters, 3, padding='same')(x)\n        x = tensorflow.keras.layers.BatchNormalization()(x)\n        x = tensorflow.keras.layers.MaxPooling2D(3, strides=2, padding='same')(x)\n        residual = tensorflow.keras.layers.Conv2D(\n            filters, 1, strides=2, padding='same')(previous_block_activation)\n        x = tensorflow.keras.layers.add([x, residual])\n        previous_block_activation = x\n    # upsampling\n    for filters in [256, 128, 64, 32]:\n        x = tensorflow.keras.layers.Activation('relu')(x)\n        x = tensorflow.keras.layers.Conv2DTranspose(filters, 3, padding='same')(x)\n        x = tensorflow.keras.layers.BatchNormalization()(x)\n        x = tensorflow.keras.layers.Activation('relu')(x)\n        x = tensorflow.keras.layers.Conv2DTranspose(filters, 3, padding='same')(x)\n        x = tensorflow.keras.layers.BatchNormalization()(x)\n        x = tensorflow.keras.layers.UpSampling2D(2)(x)\n        residual = tensorflow.keras.layers.UpSampling2D(2)(previous_block_activation)\n        residual = tensorflow.keras.layers.Conv2D(filters, 1, padding='same')(residual)\n        x = tensorflow.keras.layers.add([x, residual])\n        previous_block_activation = x\n    outputs = tensorflow.keras.layers.Conv2D(\n        num_classes, 3, activation='sigmoid', padding='same')(x)\n    model = tensorflow.keras.Model(inputs, outputs)\n    return model","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tensorflow.keras.backend.clear_session()\nmodel = get_model((IMAGE_SIZE, IMAGE_SIZE))\ntensorflow.keras.utils.plot_model(model, show_shapes=True)\nmodel.compile(\n    tensorflow.keras.optimizers.Adam(0.001),\n    tensorflow.keras.losses.BinaryCrossentropy(), ['accuracy'])\nmodel.fit(train_dataset, epochs=20)\nmodel.save('/content/drive/MyDrive/Projects/Kaggle/Image Masking/image_masking/model/{}'.format(IMAGE_SIZE))","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Testing","metadata":{}},{"cell_type":"code","source":"import cv2\nimport matplotlib.pyplot as pyplot\nimport tensorflow\n\n%matplotlib inline","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = tensorflow.keras.models.load_model('/content/drive/MyDrive/Projects/Kaggle/Image Masking/image_masking/model/128')\nimage_path = '/content/drive/MyDrive/Projects/Kaggle/Image Masking/image_masking/data/train/0cdf5b5d0ce1_01.jpg'\nimage = tensorflow.io.read_file(image_path)\nimage = tensorflow.image.decode_jpeg(image, channels=3)\nimage.set_shape([None, None, 3])\nimage = tensorflow.image.resize(images=image, size=[128, 128])\nimage = tensorflow.cast(image, tensorflow.float32) / 255.0\nimage = tensorflow.expand_dims(image, axis=0)\nimage = model.predict(image)\nimage = tensorflow.squeeze(image)\npyplot.imshow(image)\npyplot.show()\npyplot.savefig('/content/prediction.jpg')","metadata":{},"execution_count":null,"outputs":[]}]}