{"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":"# Augmentations\n\n**[Info]**. For larger augmentation pipeline support, please refer to use [KerasCV](https://keras.io/api/keras_cv/layers/augmentation/). It offers augmentation for classification, detection, segmentation, and many more.\n\n---\n\n![image](https://user-images.githubusercontent.com/17668390/169665594-608f7468-7323-41f8-9ba0-400fe9eb828f.gif)\n","metadata":{}},{"cell_type":"code","source":"import random\nimport numpy as np\nimport pandas as pd\nimport seaborn as sns\nfrom glob import glob\nfrom pylab import rcParams\nimport matplotlib.pyplot as plt\nimport plotly.graph_objects as go\nimport os, gc, cv2, random, warnings, math, sys, json, pprint\n\n# sklearn\nfrom sklearn.utils import class_weight\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics import accuracy_score, balanced_accuracy_score\n\n# tf \nimport tensorflow as tf\nfrom tensorflow import keras \nfrom tensorflow.keras import backend as K\nfrom tensorflow.keras.preprocessing.image import ImageDataGenerator\n\nos.environ['TF_CPP_MIN_LOG_LEVEL'] = '3' \nwarnings.simplefilter('ignore')\ntf.__version__","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-04-15T04:52:13.486806Z","iopub.execute_input":"2023-04-15T04:52:13.487771Z","iopub.status.idle":"2023-04-15T04:52:21.165325Z","shell.execute_reply.started":"2023-04-15T04:52:13.487711Z","shell.execute_reply":"2023-04-15T04:52:21.164292Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data","metadata":{}},{"cell_type":"code","source":"IMAGE_DIM = (100, 100, 3)\nBATCH_SIZ = 25\nSEED  = 101","metadata":{"execution":{"iopub.status.busy":"2023-04-15T04:52:21.168283Z","iopub.execute_input":"2023-04-15T04:52:21.170835Z","iopub.status.idle":"2023-04-15T04:52:21.175793Z","shell.execute_reply.started":"2023-04-15T04:52:21.170789Z","shell.execute_reply":"2023-04-15T04:52:21.174859Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TRAIN_IMG_PATH = '../input/flowers-recognition/flowers'\ntrain_datagen = ImageDataGenerator()\ndatagens = train_datagen.flow_from_directory(\n    TRAIN_IMG_PATH,\n    target_size=IMAGE_DIM[:2],\n    batch_size=BATCH_SIZ,\n    seed=SEED, \n    shuffle=True,\n    class_mode='categorical'\n)\nNUM_CLASSES = 5","metadata":{"execution":{"iopub.status.busy":"2023-04-15T04:52:21.177335Z","iopub.execute_input":"2023-04-15T04:52:21.177937Z","iopub.status.idle":"2023-04-15T04:52:23.096622Z","shell.execute_reply.started":"2023-04-15T04:52:21.177897Z","shell.execute_reply":"2023-04-15T04:52:23.095623Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"images, labels = next(\n    iter\n    (\n        datagens\n    )\n)\n\nimages.shape, labels.shape","metadata":{"execution":{"iopub.status.busy":"2023-04-15T04:52:23.099581Z","iopub.execute_input":"2023-04-15T04:52:23.100068Z","iopub.status.idle":"2023-04-15T04:52:23.310868Z","shell.execute_reply.started":"2023-04-15T04:52:23.100038Z","shell.execute_reply":"2023-04-15T04:52:23.309865Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## CutMix Augmentation","metadata":{}},{"cell_type":"code","source":"# keras_cv.layers.CutMix\nclass CutMix(keras.layers.Layer):\n    \"\"\"Original implementation: https://github.com/keras-team/keras-cv.\n    The original implementaiton provide more interface to apply mixup on\n    various CV related task, i.e. object detection etc. It also provides\n    many effective validation check.\n    \n    Derived and modified for simpler usages: M.Innat.\n    Ref. https://gist.github.com/innat/0524ee77de17f0601f0dee69aa52c713\n    \"\"\"\n\n    def __init__(self, alpha=1.0, seed=None, **kwargs):\n        super().__init__(**kwargs)\n        self.alpha = alpha\n        self.seed = seed\n\n    @staticmethod\n    def _sample_from_beta(alpha, beta, shape):\n        sample_alpha = tf.random.gamma(shape, 1.0, beta=alpha)\n        sample_beta = tf.random.gamma(shape, 1.0, beta=beta)\n        return sample_alpha / (sample_alpha + sample_beta)\n\n    def _cutmix_labels(self, labels, lambda_sample, permutation_order):\n        cutout_labels = tf.gather(labels, permutation_order)\n\n        lambda_sample = tf.reshape(lambda_sample, [-1, 1])\n        labels = lambda_sample * labels + (1.0 - lambda_sample) * cutout_labels\n        return labels\n\n    def _cutmix_samples(self, images):\n        input_shape = tf.shape(images)\n        batch_size, image_height, image_width = (\n            input_shape[0],\n            input_shape[1],\n            input_shape[2],\n        )\n\n        permutation_order = tf.random.shuffle(tf.range(0, batch_size), seed=self.seed)\n        lambda_sample = CutMix._sample_from_beta(self.alpha, self.alpha, (batch_size,))\n\n        ratio = tf.math.sqrt(1 - lambda_sample)\n\n        cut_height = tf.cast(\n            ratio * tf.cast(image_height, dtype=tf.float32), dtype=tf.int32\n        )\n        cut_width = tf.cast(\n            ratio * tf.cast(image_height, dtype=tf.float32), dtype=tf.int32\n        )\n\n        random_center_height = tf.random.uniform(\n            shape=[batch_size], minval=0, maxval=image_height, dtype=tf.int32\n        )\n        random_center_width = tf.random.uniform(\n            shape=[batch_size], minval=0, maxval=image_width, dtype=tf.int32\n        )\n\n        bounding_box_area = cut_height * cut_width\n        lambda_sample = 1.0 - bounding_box_area / (image_height * image_width)\n        lambda_sample = tf.cast(lambda_sample, dtype=tf.float32)\n\n        images = self.fill_rectangle(\n            images,\n            random_center_width,\n            random_center_height,\n            cut_width,\n            cut_height,\n            tf.gather(images, permutation_order),\n        )\n\n        return images, lambda_sample, permutation_order\n\n    def call(self, batch_inputs, training=None):\n        bs_images = tf.cast(batch_inputs[0], dtype=tf.float32)  \n        bs_labels = tf.cast(batch_inputs[1], dtype=tf.float32)  \n\n        cutmix_images, lambda_sample, permutation_order = self._cutmix_samples(\n            bs_images\n        )\n        cutmix_labels = self._cutmix_labels(bs_labels, lambda_sample, permutation_order)\n\n        return [cutmix_images, cutmix_labels]\n\n    def fill_rectangle(\n        self, images, centers_x, centers_y, widths, heights, fill_values\n    ):\n        images_shape = tf.shape(images)\n        images_height = images_shape[1]\n        images_width = images_shape[2]\n\n        xywh = tf.stack([centers_x, centers_y, widths, heights], axis=1)\n        xywh = tf.cast(xywh, tf.float32)\n        corners = self.convert_format(xywh)\n        mask_shape = (images_width, images_height)\n\n        is_rectangle = self.corners_to_mask(corners, mask_shape)\n        is_rectangle = tf.expand_dims(is_rectangle, -1)\n        images = tf.where(is_rectangle, fill_values, images)\n        return images\n\n    def convert_format(self, boxes):\n        boxes = tf.cast(boxes, dtype=tf.float32)\n        x, y, width, height, rest = tf.split(boxes, [1, 1, 1, 1, -1], axis=-1)\n        results = tf.concat(\n            [\n                x - width / 2.0,\n                y - height / 2.0,\n                x + width / 2.0,\n                y + height / 2.0,\n                rest,\n            ],\n            axis=-1,\n        )\n        return results\n\n    def _axis_mask(self, starts, ends, mask_len):\n        # index range of axis\n        batch_size = tf.shape(starts)[0]\n        axis_indices = tf.range(mask_len, dtype=starts.dtype)\n        axis_indices = tf.expand_dims(axis_indices, 0)\n        axis_indices = tf.tile(axis_indices, [batch_size, 1])\n\n        # mask of index bounds\n        axis_mask = tf.greater_equal(axis_indices, starts) & tf.less(axis_indices, ends)\n        return axis_mask\n\n    def corners_to_mask(self, bounding_boxes, mask_shape):\n        mask_width, mask_height = mask_shape\n        x0, y0, x1, y1 = tf.split(bounding_boxes, [1, 1, 1, 1], axis=-1)\n\n        w_mask = self._axis_mask(x0, x1, mask_width)\n        h_mask = self._axis_mask(y0, y1, mask_height)\n\n        w_mask = tf.expand_dims(w_mask, axis=1)\n        h_mask = tf.expand_dims(h_mask, axis=2)\n        masks = tf.logical_and(w_mask, h_mask)\n        return masks","metadata":{"execution":{"iopub.status.busy":"2023-04-15T04:52:23.312280Z","iopub.execute_input":"2023-04-15T04:52:23.312649Z","iopub.status.idle":"2023-04-15T04:52:23.338109Z","shell.execute_reply.started":"2023-04-15T04:52:23.312607Z","shell.execute_reply":"2023-04-15T04:52:23.336963Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with tf.device('/device:GPU:0'):\n    cutmix_image, cutmix_label = CutMix()(\n        [images, labels]\n    )","metadata":{"execution":{"iopub.status.busy":"2023-04-15T04:52:25.177223Z","iopub.execute_input":"2023-04-15T04:52:25.178144Z","iopub.status.idle":"2023-04-15T04:52:27.922806Z","shell.execute_reply.started":"2023-04-15T04:52:25.178091Z","shell.execute_reply":"2023-04-15T04:52:27.921796Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig = plt.figure(figsize=(20, 20))\ncolumns = 4\nrows = 5\n\nfor i in range(1, columns*rows +1):\n    img = cutmix_image[i].numpy().astype('int')\n    lbl = cutmix_label[i].numpy()\n    fig.add_subplot(rows, columns, i)\n    plt.imshow(img)\n    plt.axis(\"off\")\nplt.show()","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-04-15T04:52:34.438803Z","iopub.execute_input":"2023-04-15T04:52:34.439196Z","iopub.status.idle":"2023-04-15T04:52:36.035577Z","shell.execute_reply.started":"2023-04-15T04:52:34.439161Z","shell.execute_reply":"2023-04-15T04:52:36.034222Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## MixUp Augmentation","metadata":{}},{"cell_type":"code","source":"# keras_cv.layers.MixUp\nclass MixUp(keras.layers.Layer):\n    \"\"\"Original implementation: https://github.com/keras-team/keras-cv.\n    The original implementaiton provide more interface to apply mixup on\n    various CV related task, i.e. object detection etc. It also provides\n    many effective validation check.\n\n    Derived and modified for simpler usages: M.Innat.\n    Ref. https://gist.github.com/innat/0ee2b6155d663aac2617fe596e1d8d49\n    \"\"\"\n\n    def __init__(self, alpha=0.2, seed=None, **kwargs):\n        super().__init__(**kwargs)\n        self.alpha = alpha\n        self.seed = seed\n\n    @staticmethod\n    def _sample_from_beta(alpha, beta, shape):\n        sample_alpha = tf.random.gamma(shape, 1.0, beta=alpha)\n        sample_beta = tf.random.gamma(shape, 1.0, beta=beta)\n        return sample_alpha / (sample_alpha + sample_beta)\n\n    def _mixup_samples(self, images):\n        batch_size = tf.shape(images)[0]\n        permutation_order = tf.random.shuffle(tf.range(0, batch_size), seed=self.seed)\n\n        lambda_sample = MixUp._sample_from_beta(self.alpha, self.alpha, (batch_size,))\n        lambda_sample = tf.reshape(lambda_sample, [-1, 1, 1, 1])\n\n        mixup_images = tf.gather(images, permutation_order)\n        images = lambda_sample * images + (1.0 - lambda_sample) * mixup_images\n\n        return images, tf.squeeze(lambda_sample), permutation_order\n\n    def _mixup_labels(self, labels, lambda_sample, permutation_order):\n        labels_for_mixup = tf.gather(labels, permutation_order)\n\n        lambda_sample = tf.reshape(lambda_sample, [-1, 1])\n        labels = lambda_sample * labels + (1.0 - lambda_sample) * labels_for_mixup\n\n        return labels\n\n    def call(self, batch_inputs):\n        bs_images = tf.cast(batch_inputs[0], dtype=tf.float32)  \n        bs_labels = tf.cast(batch_inputs[1], dtype=tf.float32)  \n\n        mixup_images, lambda_sample, permutation_order = self._mixup_samples(bs_images)\n        mixup_labels = self._mixup_labels(bs_labels, lambda_sample, permutation_order)\n\n        return [mixup_images, mixup_labels]\n\n    def get_config(self):\n        config = super().get_config()\n        config.update(\n            {\n                \"alpha\": self.alpha,\n                \"seed\": self.seed,\n            }\n        )\n        return config","metadata":{"execution":{"iopub.status.busy":"2023-04-15T04:53:19.311278Z","iopub.execute_input":"2023-04-15T04:53:19.312279Z","iopub.status.idle":"2023-04-15T04:53:19.326562Z","shell.execute_reply.started":"2023-04-15T04:53:19.312226Z","shell.execute_reply":"2023-04-15T04:53:19.324994Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with tf.device('/device:GPU:0'):\n    mixup_image, mixup_label = MixUp()(\n        [images, labels]\n    )","metadata":{"execution":{"iopub.status.busy":"2023-04-15T04:54:23.285379Z","iopub.execute_input":"2023-04-15T04:54:23.286067Z","iopub.status.idle":"2023-04-15T04:54:23.302340Z","shell.execute_reply.started":"2023-04-15T04:54:23.286023Z","shell.execute_reply":"2023-04-15T04:54:23.301209Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig = plt.figure(figsize=(20, 20))\ncolumns = 4\nrows = 5\n\nfor i in range(1, columns*rows +1):\n    img = mixup_image[i].numpy().astype('int')\n    lbl = mixup_label[i].numpy()\n    fig.add_subplot(rows, columns, i)\n    plt.imshow(img)\n    plt.axis(\"off\")\n    \nplt.show()","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-04-15T04:54:23.974397Z","iopub.execute_input":"2023-04-15T04:54:23.974765Z","iopub.status.idle":"2023-04-15T04:54:25.631079Z","shell.execute_reply.started":"2023-04-15T04:54:23.974731Z","shell.execute_reply":"2023-04-15T04:54:25.628903Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Channel Shuffle","metadata":{}},{"cell_type":"code","source":"# keras_cv.layers.ChannelShuffle\nclass ChannelShuffle(keras.layers.Layer):\n    def __init__(self, groups=3, seed=None, **kwargs):\n        super().__init__(**kwargs)\n        self.groups = groups\n        self.seed = seed\n\n    def _channel_shuffling(self, images):\n        height = tf.shape(images)[1]\n        width = tf.shape(images)[2]\n        num_channels = images.shape[3]\n        channels_per_group = num_channels // self.groups\n        \n        images = tf.reshape(\n            images, [-1, height, width, self.groups, channels_per_group]\n        )\n        images = tf.transpose(images, perm=[3, 1, 2, 4, 0])\n        images = tf.random.shuffle(images, seed=self.seed)\n        images = tf.transpose(images, perm=[4, 1, 2, 3, 0])\n        images = tf.reshape(images, [-1, height, width, num_channels])\n        return images\n\n    def call(self, images, training=True):\n        if training:\n            return self._channel_shuffling(images)\n        else:\n            return images\n\n    def get_config(self):\n        config = super().get_config()\n        config.update({\"groups\": self.groups, \"seed\": self.seed})\n        return config","metadata":{"execution":{"iopub.status.busy":"2023-04-15T04:55:04.371831Z","iopub.execute_input":"2023-04-15T04:55:04.372374Z","iopub.status.idle":"2023-04-15T04:55:04.385157Z","shell.execute_reply.started":"2023-04-15T04:55:04.372335Z","shell.execute_reply":"2023-04-15T04:55:04.383983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"chlshl_image = ChannelShuffle(groups=3)(images)","metadata":{"execution":{"iopub.status.busy":"2023-04-15T04:55:06.514208Z","iopub.execute_input":"2023-04-15T04:55:06.514706Z","iopub.status.idle":"2023-04-15T04:55:06.575981Z","shell.execute_reply.started":"2023-04-15T04:55:06.514653Z","shell.execute_reply":"2023-04-15T04:55:06.574437Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig = plt.figure(figsize=(20, 20))\ncolumns = 4\nrows = 5\n\nfor i in range(1, columns*rows +1):\n    img = chlshl_image[i].numpy().astype('int')\n    fig.add_subplot(rows, columns, i)\n    plt.imshow(img)\n    plt.axis(\"off\")\nplt.show()","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-04-15T04:55:07.705915Z","iopub.execute_input":"2023-04-15T04:55:07.706446Z","iopub.status.idle":"2023-04-15T04:55:09.949827Z","shell.execute_reply.started":"2023-04-15T04:55:07.706409Z","shell.execute_reply":"2023-04-15T04:55:09.948874Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Color Jitter Augmentation","metadata":{}},{"cell_type":"code","source":"# keras_cv.layers.RandomColorJitter\nclass ColorJitter(keras.layers.Layer):\n    def __init__(\n        self,\n        brightness_factor=0.5,\n        contrast_factor=(0.5, 0.9),\n        saturation_factor=(0.5, 0.9),\n        hue_factor=0.5,\n        seed=None,\n        **kwargs,\n    ):\n        super().__init__(**kwargs)\n        self.seed = seed\n        self.brightness_factor = self._check_factor_limit(\n            brightness_factor, name=\"brightness\"\n        )\n        self.contrast_factor = self._check_factor_limit(\n            contrast_factor, name=\"contrast\"\n        )\n        self.saturation_factor = self._check_factor_limit(\n            saturation_factor, name=\"saturation\"\n        )\n        self.hue_factor = self._check_factor_limit(hue_factor, name=\"hue\")\n\n    def _check_factor_limit(self, factor, name):\n        if isinstance(factor, (int, float)):\n            if factor < 0:\n                raise TypeError(\n                    \"The factor value should be non-negative scalar or tuple \"\n                    f\"or list of two upper and lower bound number. Received: {factor}\"\n                )\n            if name == \"brightness\" or name == \"hue\":\n                return abs(factor)\n            return (0, abs(factor))\n        elif isinstance(factor, (tuple, list)) and len(factor) == 2:\n            if name == \"brightness\" or name == \"hue\":\n                raise ValueError(\n                    \"The factor limit for brightness and hue, it should be a single \"\n                    f\"non-negative scaler. Received: {factor} for {name}\"\n                )\n            return sorted(factor)\n        else:\n            raise TypeError(\n                \"The factor value should be non-negative scalar or tuple \"\n                f\"or list of two upper and lower bound number. Received: {factor}\"\n            )\n\n    def _color_jitter(self, images):\n        original_dtype = images.dtype\n        images = tf.cast(images, dtype=tf.float32)\n\n        brightness = tf.image.random_brightness(\n            images, max_delta=self.brightness_factor * 255.0, seed=self.seed\n        )\n        brightness = tf.clip_by_value(brightness, 0.0, 255.0)\n\n        contrast = tf.image.random_contrast(\n            brightness,\n            lower=self.contrast_factor[0],\n            upper=self.contrast_factor[1],\n            seed=self.seed,\n        )\n        saturation = tf.image.random_saturation(\n            contrast,\n            lower=self.saturation_factor[0],\n            upper=self.saturation_factor[1],\n            seed=self.seed,\n        )\n        hue = tf.image.random_hue(saturation, max_delta=self.hue_factor, seed=self.seed)\n        return tf.cast(hue, original_dtype)\n\n    def call(self, images, training=True):\n        if training:\n            return self._color_jitter(images)\n        else:\n            return images\n\n    def get_config(self):\n        config = super().get_config()\n        config.update(\n            {\n                \"brightness_factor\": self.brightness_factor,\n                \"contrast_factor\": self.contrast_factor,\n                \"saturation_factor\": self.saturation_factor,\n                \"hue_factor\": self.hue_factor,\n                \"seed\": self.seed,\n            }\n        )\n        return config","metadata":{"execution":{"iopub.status.busy":"2023-04-15T04:56:22.467093Z","iopub.execute_input":"2023-04-15T04:56:22.467465Z","iopub.status.idle":"2023-04-15T04:56:22.484170Z","shell.execute_reply.started":"2023-04-15T04:56:22.467430Z","shell.execute_reply":"2023-04-15T04:56:22.482982Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cjit_image = ColorJitter()(images)","metadata":{"execution":{"iopub.status.busy":"2023-04-15T04:56:24.659208Z","iopub.execute_input":"2023-04-15T04:56:24.660152Z","iopub.status.idle":"2023-04-15T04:56:24.691052Z","shell.execute_reply.started":"2023-04-15T04:56:24.660097Z","shell.execute_reply":"2023-04-15T04:56:24.690008Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig = plt.figure(figsize=(20, 20))\ncolumns = 4\nrows = 5\n\nfor i in range(1, columns*rows +1):\n    img = cjit_image[i].numpy().astype('int')\n    fig.add_subplot(rows, columns, i)\n    plt.imshow(img)\n    plt.axis(\"off\")\nplt.show()","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-04-15T04:56:27.562588Z","iopub.execute_input":"2023-04-15T04:56:27.563291Z","iopub.status.idle":"2023-04-15T04:56:28.982778Z","shell.execute_reply.started":"2023-04-15T04:56:27.563251Z","shell.execute_reply":"2023-04-15T04:56:28.981887Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## RGBShift","metadata":{"execution":{"iopub.status.busy":"2022-04-01T07:17:54.292559Z","iopub.execute_input":"2022-04-01T07:17:54.292998Z","iopub.status.idle":"2022-04-01T07:17:54.297488Z","shell.execute_reply.started":"2022-04-01T07:17:54.29296Z","shell.execute_reply":"2022-04-01T07:17:54.296588Z"}}},{"cell_type":"code","source":"# keras_cv.layers.RandomChannelShift\nclass RGBShift(keras.layers.Layer):\n    \"\"\"RGBShift class randomly shift values for each channel of the input RGB image. \n    \"\"\"\n    def __init__(\n        self,\n        factor,\n        seed=None,\n        **kwargs\n    ):\n        super().__init__(**kwargs)\n        self.factor = self._set_shift_limit(factor)\n        self.seed = seed\n        \n    def _set_shift_limit(self, factor):\n        if isinstance(factor, (tuple, list)):\n            if len(factor) != 2: \n                raise ValueError(\n                    'The factor should be scalar'\n                    'tuple or list of two upper and lower' \n                    f'bound number. Got {factor}'\n                )\n            return self._check_factor_range(sorted(factor))\n        elif isinstance(factor, (int, float)):\n            factor = abs(factor)\n            return self._check_factor_range([-factor, factor])\n        else:\n            raise ValueError(\n                'The factor should be scalar'\n                f'tuple or list of two upper and lower bound umber. Got {factor}'\n            )\n            \n    @staticmethod\n    def _check_factor_range(factor):\n        if all(isinstance(each_elem, float) for each_elem in factor):\n            if factor[0] < -1.0 or factor[1] > 1.0:\n                raise ValueError(f\"Got {factor}\")\n            return factor\n        elif all(isinstance(each_elem, int) for each_elem in factor):\n            if factor[0] < -255 or factor[1] > 255:\n                raise ValueError(f\"Got {factor}\")\n            return factor\n        else:\n            raise ValueError(f'Both bound must be same dtype. Got {factor}')\n            \n    def _get_random_uniform(self, shift_limit, rgb_delta_shape):\n            if self.seed is not None:\n                _rand_uniform = tf.random.stateless_uniform(\n                    shape=rgb_delta_shape,\n                    seed=[0, self.seed],\n                    minval=shift_limit[0],\n                    maxval=shift_limit[1],\n                )\n            else:\n                _rand_uniform = tf.random.uniform(\n                    rgb_delta_shape, \n                    minval=shift_limit[0], \n                    maxval=shift_limit[1], \n                    dtype=tf.float32\n                )\n                \n            if all(isinstance(each_elem, float) for each_elem in shift_limit):\n                _rand_uniform = _rand_uniform * 85.0\n            \n            return _rand_uniform\n    \n    def _rgb_shifting(self, images):\n        rank = images.shape.rank\n        original_dtype = images.dtype\n\n        if rank == 3:\n            rgb_delta_shape = (1, 1)\n        elif rank == 4:\n            # Keep only the batch dim. This will ensure to have same adjustment\n            # with in one image, but different across the images.\n            rgb_delta_shape = [tf.shape(images)[0], 1, 1]\n        else:\n            raise ValueError(\n                f\"Expect the input image to be rank 3 or 4. Got {images.shape}\"\n            )\n        r_shift = self._get_random_uniform(self.factor, rgb_delta_shape)   \n        g_shift = self._get_random_uniform(self.factor, rgb_delta_shape)\n        b_shift = self._get_random_uniform(self.factor, rgb_delta_shape)\n        unstack_rgb = tf.unstack(tf.cast(images, dtype=tf.float32), axis=-1)\n        shifted_rgb = tf.stack(\n            [\n                tf.add(unstack_rgb[0], r_shift),\n                tf.add(unstack_rgb[1], g_shift),\n                tf.add(unstack_rgb[2], b_shift)\n            ], axis=-1\n        )\n        shifted_rgb = tf.clip_by_value(shifted_rgb, 0.0, 255.0)\n        return tf.cast(shifted_rgb, dtype=original_dtype)\n\n    def call(self, images, training=True):\n        return self._rgb_shifting(images)\n    \n    def get_config(self):\n        config = super().get_config()\n        config.update(\n            {\n                \"factor\": self.factor, \n                \"seed\": self.seed\n            }\n        )\n        return config \n    def compute_output_shape(self, input_shape):\n        return input_shape","metadata":{"execution":{"iopub.status.busy":"2023-04-15T04:57:46.306682Z","iopub.execute_input":"2023-04-15T04:57:46.307129Z","iopub.status.idle":"2023-04-15T04:57:46.326605Z","shell.execute_reply.started":"2023-04-15T04:57:46.307091Z","shell.execute_reply":"2023-04-15T04:57:46.325632Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rgbshift_images = RGBShift(factor=(-120, 120))(images)","metadata":{"execution":{"iopub.status.busy":"2023-04-15T04:57:46.868292Z","iopub.execute_input":"2023-04-15T04:57:46.868665Z","iopub.status.idle":"2023-04-15T04:57:46.888520Z","shell.execute_reply.started":"2023-04-15T04:57:46.868630Z","shell.execute_reply":"2023-04-15T04:57:46.887557Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig = plt.figure(figsize=(20, 20))\ncolumns = 4\nrows = 5\n\nfor i in range(1, columns*rows +1):\n    img = rgbshift_images[i].numpy().astype('int')\n    fig.add_subplot(rows, columns, i)\n    plt.imshow(img)\n    plt.axis(\"off\")\nplt.show()","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-04-15T04:57:47.787622Z","iopub.execute_input":"2023-04-15T04:57:47.788721Z","iopub.status.idle":"2023-04-15T04:57:49.363947Z","shell.execute_reply.started":"2023-04-15T04:57:47.788679Z","shell.execute_reply":"2023-04-15T04:57:49.360601Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# From Keras API","metadata":{}},{"cell_type":"code","source":"from tensorflow.keras.layers import Rescaling\nfrom tensorflow.keras.layers import Resizing\nfrom tensorflow.keras.layers import RandomCrop\nfrom tensorflow.keras.layers import RandomFlip\nfrom tensorflow.keras.layers import RandomZoom\nfrom tensorflow.keras.layers import RandomRotation","metadata":{"execution":{"iopub.status.busy":"2023-04-15T04:58:00.642793Z","iopub.execute_input":"2023-04-15T04:58:00.643516Z","iopub.status.idle":"2023-04-15T04:58:00.650593Z","shell.execute_reply.started":"2023-04-15T04:58:00.643476Z","shell.execute_reply":"2023-04-15T04:58:00.649278Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"keras_augmenter = keras.Sequential([\n    Resizing(\n        *[IMAGE_DIM[0] + 32] * 2,\n        interpolation=\"bilinear\"\n    ),\n    RandomCrop(\n        *[IMAGE_DIM[0]] * 2\n    ),\n    RandomFlip(\"horizontal\"),\n    RandomZoom(0.6, fill_mode='reflect'),\n    RandomRotation(0.4, fill_mode='reflect'),\n])","metadata":{"execution":{"iopub.status.busy":"2023-04-15T04:58:03.848809Z","iopub.execute_input":"2023-04-15T04:58:03.849811Z","iopub.status.idle":"2023-04-15T04:58:03.874279Z","shell.execute_reply.started":"2023-04-15T04:58:03.849771Z","shell.execute_reply":"2023-04-15T04:58:03.873065Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"keras_augmented = keras_augmenter(images, training=True)","metadata":{"execution":{"iopub.status.busy":"2023-04-15T04:58:04.518549Z","iopub.execute_input":"2023-04-15T04:58:04.519746Z","iopub.status.idle":"2023-04-15T04:58:10.885169Z","shell.execute_reply.started":"2023-04-15T04:58:04.519692Z","shell.execute_reply":"2023-04-15T04:58:10.884091Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig = plt.figure(figsize=(20, 20))\ncolumns = 4\nrows = 5\n\nfor i in range(1, columns*rows +1):\n    img = keras_augmented[i].numpy().astype('int')\n    fig.add_subplot(rows, columns, i)\n    plt.imshow(img)\n    plt.axis(\"off\")\nplt.show()","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-04-15T04:58:10.890092Z","iopub.execute_input":"2023-04-15T04:58:10.890397Z","iopub.status.idle":"2023-04-15T04:58:12.479643Z","shell.execute_reply.started":"2023-04-15T04:58:10.890368Z","shell.execute_reply":"2023-04-15T04:58:12.478746Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# All At Once","metadata":{}},{"cell_type":"code","source":"def plot_stuff(a, b, c, d, titles):\n    plt.figure(figsize=(25, 25))\n    \n    plt.subplot(1, 4, 1)\n    plt.axis('off')\n    plt.imshow(a.astype('int'))\n    plt.title(titles[0])\n    \n    plt.subplot(1, 4, 2)\n    plt.axis('off')\n    plt.imshow(b.astype('int'))\n    plt.title(titles[1])\n    \n    plt.subplot(1, 4, 3)\n    plt.axis('off')\n    plt.imshow(c.astype('int'))\n    plt.title(titles[2])\n    \n    plt.subplot(1, 4, 4)\n    plt.axis('off')\n    plt.imshow(d.astype('int'))\n    plt.title(titles[3])\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-04-15T04:58:31.363230Z","iopub.execute_input":"2023-04-15T04:58:31.363618Z","iopub.status.idle":"2023-04-15T04:58:31.372359Z","shell.execute_reply.started":"2023-04-15T04:58:31.363582Z","shell.execute_reply":"2023-04-15T04:58:31.371176Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rgbshift_images = RGBShift(factor=(-120, 120))(images)\ncjit_image = ColorJitter()(images)\nchlshl_image = ChannelShuffle(groups=3)(images)\nkes_image = keras_augmenter(images, training=True)\n\nmixup_image, mixup_label = MixUp()([images, labels])\ncutmix_image, cutmix_label = CutMix()([images, labels])","metadata":{"execution":{"iopub.status.busy":"2023-04-15T04:58:31.865344Z","iopub.execute_input":"2023-04-15T04:58:31.865713Z","iopub.status.idle":"2023-04-15T04:58:34.171207Z","shell.execute_reply.started":"2023-04-15T04:58:31.865679Z","shell.execute_reply":"2023-04-15T04:58:34.170193Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for i, (orig, rg, cj, ch, fm, mi, cu) in enumerate(\n    zip(\n        images, rgbshift_images, cjit_image, \n        chlshl_image, kes_image, mixup_image,cutmix_image\n    )\n):\n    plot_stuff(\n        orig,\n        rg.numpy(),\n        cj.numpy(), \n        ch.numpy(), \n        ['Input', 'RGBShift', 'ColorJitter', 'ChannelShuffle']\n    )\n    \n    plot_stuff(\n        orig,\n        fm.numpy(),\n        mi.numpy(),\n        cu.numpy(),\n        ['Input','KerasLayers', 'MixUp', 'CutMix']\n    )\n    \n    print('[INFO]: Iter ...............')","metadata":{"execution":{"iopub.status.busy":"2023-04-15T04:58:36.354565Z","iopub.execute_input":"2023-04-15T04:58:36.354972Z","iopub.status.idle":"2023-04-15T04:58:50.199419Z","shell.execute_reply.started":"2023-04-15T04:58:36.354933Z","shell.execute_reply":"2023-04-15T04:58:50.198410Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"---\n\n## Resource\n\nThis notebook was initially started with the following projects.\n\n- [[TF.Keras]: Cassava: Advanced Training Mechanism](https://www.kaggle.com/ipythonx/tf-keras-cassava-advanced-training-mechanism)\n- https://www.tensorflow.org/tutorials/images/data_augmentation","metadata":{}}]}