{"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":"# Based on FAIR Paper : Masked AutoEncoders are scalable vision learners\n\n## Part (1/2) Pretraining on Imagenet1K dataset.\n\nConcept is very simple, simmilar to text transformers, you can pretrain Vit by masking the image and then reconstructing it.\n\n### But how?\n\nIn text if you mask ~10% input then you can drastically change the meaning and embedding. In Images masking ~10% does little  to remove information. Authors suggest masking *75* of input tokens or best downstream performance\n\n### Part (2/2) Linear Probing with triplet/ arcface and Fine Tuning Coming soon (under progress).\n\n#### This is not an attempt to win, but just for my learning. If you found this useful kindly upvote.\n\n# Imports","metadata":{}},{"cell_type":"code","source":"import tensorflow as tf\nfrom tensorflow.keras import layers\nfrom tensorflow import keras\n\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport random\nfrom tqdm.auto import tqdm, trange\n\n\nimport os\nimport glob\n\ntf.version.VERSION\n","metadata":{"execution":{"iopub.status.busy":"2022-07-16T16:52:38.046924Z","iopub.execute_input":"2022-07-16T16:52:38.047314Z","iopub.status.idle":"2022-07-16T16:52:44.731577Z","shell.execute_reply.started":"2022-07-16T16:52:38.047197Z","shell.execute_reply":"2022-07-16T16:52:44.730615Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# HyperParameters\n\nMost of them are used as it is from the paper.","metadata":{}},{"cell_type":"code","source":"# DATA\nBUFFER_SIZE = 1024\nBATCH_SIZE = 256\nAUTO = tf.data.AUTOTUNE\nINPUT_SHAPE = (224, 224, 3)\nNUM_CLASSES = 64\n\n# OPTIMIZER\nLEARNING_RATE = 5e-3\nWEIGHT_DECAY = 1e-4\n\n# PRETRAINING\nEPOCHS = 100\n\n# AUGMENTATION\nIMAGE_SIZE = 48  # We will resize input images to this size.\nPATCH_SIZE = 6  # Size of the patches to be extracted from the input images.\nNUM_PATCHES = (IMAGE_SIZE // PATCH_SIZE) ** 2\nMASK_PROPORTION = 0.75  # We have found 75% masking to give us the best results.\n\n# ENCODER and DECODER\nLAYER_NORM_EPS = 1e-6\nENC_PROJECTION_DIM = 128\nDEC_PROJECTION_DIM = 64\nENC_NUM_HEADS = 4\nENC_LAYERS = 6\nDEC_NUM_HEADS = 4\nDEC_LAYERS = (\n    2  # The decoder is lightweight but should be reasonably deep for reconstruction.\n)\nENC_TRANSFORMER_UNITS = [\n    ENC_PROJECTION_DIM * 2,\n    ENC_PROJECTION_DIM,\n]  # Size of the transformer layers.\nDEC_TRANSFORMER_UNITS = [\n    DEC_PROJECTION_DIM * 2,\n    DEC_PROJECTION_DIM,\n]","metadata":{"execution":{"iopub.status.busy":"2022-07-16T16:52:44.733522Z","iopub.execute_input":"2022-07-16T16:52:44.735119Z","iopub.status.idle":"2022-07-16T16:52:44.744215Z","shell.execute_reply.started":"2022-07-16T16:52:44.735079Z","shell.execute_reply":"2022-07-16T16:52:44.741678Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load imagenet1k dataset from TFrecords","metadata":{}},{"cell_type":"code","source":"train_files = sorted(glob.glob('/kaggle/input/**/train-*-of-01024', recursive=True))\nvalid_files = sorted(glob.glob('/kaggle/input/imagenet*/validation/*'))\ndict(valid=dict(count=len(valid_files), first_file=valid_files[0], last_file=valid_files[-1]),\n     train=dict(count=len(train_files), first_file=train_files[0], last_file=train_files[-1]))","metadata":{"execution":{"iopub.status.busy":"2022-07-16T16:52:44.746537Z","iopub.execute_input":"2022-07-16T16:52:44.747189Z","iopub.status.idle":"2022-07-16T16:52:45.945466Z","shell.execute_reply.started":"2022-07-16T16:52:44.747148Z","shell.execute_reply":"2022-07-16T16:52:45.944570Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_shape = (224, 224)\n\ndef decode(serialized_example):\n    features = tf.io.parse_single_example(\n        serialized_example,\n        features={\n            'image/encoded': tf.io.FixedLenFeature([], tf.string),\n            'image/class/label': tf.io.FixedLenFeature([], tf.int64),\n        })\n    image = tf.image.decode_jpeg(features['image/encoded'], channels=3)\n    image = tf.image.resize_with_pad(image, *image_shape)  # crop/augment instead\n    label = tf.cast(features['image/class/label'], tf.int64) - 1  # [0-999]\n    return image","metadata":{"execution":{"iopub.status.busy":"2022-07-16T16:52:45.948135Z","iopub.execute_input":"2022-07-16T16:52:45.948702Z","iopub.status.idle":"2022-07-16T16:52:45.957137Z","shell.execute_reply.started":"2022-07-16T16:52:45.948665Z","shell.execute_reply":"2022-07-16T16:52:45.956197Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_ds(tfrecords, batch_size=BATCH_SIZE):\n    dataset = tf.data.TFRecordDataset(tfrecords, num_parallel_reads=tf.data.AUTOTUNE)\n    dataset = dataset.map(decode, num_parallel_calls=tf.data.AUTOTUNE)\n    dataset = dataset.batch(batch_size)\n    dataset = dataset.prefetch(tf.data.AUTOTUNE)\n    return dataset","metadata":{"execution":{"iopub.status.busy":"2022-07-16T16:52:45.958630Z","iopub.execute_input":"2022-07-16T16:52:45.959049Z","iopub.status.idle":"2022-07-16T16:52:45.971180Z","shell.execute_reply.started":"2022-07-16T16:52:45.959012Z","shell.execute_reply":"2022-07-16T16:52:45.970202Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = create_ds(valid_files)\n\nfor images in dataset.take(1):\n    print('Image shape:', images.shape)\n\nconcat = np.concatenate(list(images), axis=1)\n\nplt.figure(figsize=(20, 3), dpi=80)\nplt.imshow(concat.astype(int))\nplt.axis('off')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-07-16T16:52:45.972991Z","iopub.execute_input":"2022-07-16T16:52:45.973392Z","iopub.status.idle":"2022-07-16T16:52:52.233226Z","shell.execute_reply.started":"2022-07-16T16:52:45.973341Z","shell.execute_reply":"2022-07-16T16:52:52.232119Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"valid_ds = create_ds(valid_files, batch_size=32)","metadata":{"execution":{"iopub.status.busy":"2022-07-16T16:52:52.234526Z","iopub.execute_input":"2022-07-16T16:52:52.235801Z","iopub.status.idle":"2022-07-16T16:52:52.463085Z","shell.execute_reply.started":"2022-07-16T16:52:52.235761Z","shell.execute_reply":"2022-07-16T16:52:52.462026Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ds = create_ds(train_files, batch_size=32)","metadata":{"execution":{"iopub.status.busy":"2022-07-16T16:52:52.464858Z","iopub.execute_input":"2022-07-16T16:52:52.465211Z","iopub.status.idle":"2022-07-16T16:52:52.547616Z","shell.execute_reply.started":"2022-07-16T16:52:52.465178Z","shell.execute_reply":"2022-07-16T16:52:52.546593Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Get the train augmentation model\n\n## Here I am using the augmentations used in Vit small model","metadata":{}},{"cell_type":"code","source":"def get_train_augmentation_model():\n    model = keras.Sequential(\n        [\n            layers.Rescaling(1 / 255.0),\n            layers.Resizing(INPUT_SHAPE[0] + 20, INPUT_SHAPE[0] + 20),\n            layers.RandomCrop(IMAGE_SIZE, IMAGE_SIZE),\n            layers.RandomFlip(\"horizontal\"),\n            layers.Normalization(mean= [\n                                        0.485,\n                                        0.456,\n                                        0.406\n                                      ],\n                                 variance= [\n                                        0.229,\n                                        0.224,\n                                        0.225\n                                     ], \n                                 axis = -1\n                                ),\n        ],\n        name=\"train_data_augmentation\",\n    )\n    return model\n\n\ndef get_test_augmentation_model():\n    model = keras.Sequential(\n        [\n            layers.Rescaling(1 / 255.0),\n            layers.Resizing(IMAGE_SIZE, IMAGE_SIZE),\n            layers.Normalization(mean= [\n                                        0.485,\n                                        0.456,\n                                        0.406\n                                      ],\n                                 variance= [\n                                        0.229,\n                                        0.224,\n                                        0.225\n                                     ], \n                                 axis = -1\n                                ),\n        ],\n        name=\"test_data_augmentation\",\n    )\n    return model","metadata":{"execution":{"iopub.status.busy":"2022-07-16T16:52:52.549216Z","iopub.execute_input":"2022-07-16T16:52:52.549982Z","iopub.status.idle":"2022-07-16T16:52:52.560773Z","shell.execute_reply.started":"2022-07-16T16:52:52.549942Z","shell.execute_reply":"2022-07-16T16:52:52.559538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Keras layer that will handle creation of patches with utils for visualisation\n#####  Inspired from keras website","metadata":{}},{"cell_type":"code","source":"class Patches(layers.Layer):\n    def __init__(self, patch_size=PATCH_SIZE, **kwargs):\n        super().__init__(**kwargs)\n        self.patch_size = patch_size\n\n        # Assuming the image has three channels each patch would be\n        # of size (patch_size, patch_size, 3).\n        self.resize = layers.Reshape((-1, patch_size * patch_size * 3))\n\n    def call(self, images):\n        # Create patches from the input images\n        patches = tf.image.extract_patches(\n            images=images,\n            sizes=[1, self.patch_size, self.patch_size, 1],\n            strides=[1, self.patch_size, self.patch_size, 1],\n            rates=[1, 1, 1, 1],\n            padding=\"VALID\",\n        )\n\n        # Reshape the patches to (batch, num_patches, patch_area) and return it.\n        patches = self.resize(patches)\n        return patches\n\n    def show_patched_image(self, images, patches):\n        # This is a utility function which accepts a batch of images and its\n        # corresponding patches and help visualize one image and its patches\n        # side by side.\n        idx = np.random.choice(patches.shape[0])\n        print(f\"Index selected: {idx}.\")\n\n        plt.figure(figsize=(4, 4))\n        plt.imshow(keras.utils.array_to_img(images[idx]))\n        plt.axis(\"off\")\n        plt.show()\n\n        n = int(np.sqrt(patches.shape[1]))\n        plt.figure(figsize=(4, 4))\n        for i, patch in enumerate(patches[idx]):\n            ax = plt.subplot(n, n, i + 1)\n            patch_img = tf.reshape(patch, (self.patch_size, self.patch_size, 3))\n            plt.imshow(keras.utils.img_to_array(patch_img))\n            plt.axis(\"off\")\n        plt.show()\n\n        # Return the index chosen to validate it outside the method.\n        return idx\n\n    # taken from https://stackoverflow.com/a/58082878/10319735\n    def reconstruct_from_patch(self, patch):\n        # This utility function takes patches from a *single* image and\n        # reconstructs it back into the image. This is useful for the train\n        # monitor callback.\n        num_patches = patch.shape[0]\n        n = int(np.sqrt(num_patches))\n        patch = tf.reshape(patch, (num_patches, self.patch_size, self.patch_size, 3))\n        rows = tf.split(patch, n, axis=0)\n        rows = [tf.concat(tf.unstack(x), axis=1) for x in rows]\n        reconstructed = tf.concat(rows, axis=0)\n        return reconstructed","metadata":{"execution":{"iopub.status.busy":"2022-07-16T16:52:52.565715Z","iopub.execute_input":"2022-07-16T16:52:52.566978Z","iopub.status.idle":"2022-07-16T16:52:52.584134Z","shell.execute_reply.started":"2022-07-16T16:52:52.566941Z","shell.execute_reply":"2022-07-16T16:52:52.583025Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Keras Layer to handle patch embedding and positional embeddings","metadata":{}},{"cell_type":"code","source":"class PatchEncoder(layers.Layer):\n    def __init__(\n        self,\n        patch_size=PATCH_SIZE,\n        projection_dim=ENC_PROJECTION_DIM,\n        mask_proportion=MASK_PROPORTION,\n        downstream=False,\n        **kwargs,\n    ):\n        super().__init__(**kwargs)\n        self.patch_size = patch_size\n        self.projection_dim = projection_dim\n        self.mask_proportion = mask_proportion\n        self.downstream = downstream\n\n        # This is a trainable mask token initialized randomly from a normal\n        # distribution.\n        self.mask_token = tf.Variable(\n            tf.random.normal([1, patch_size * patch_size * 3]), trainable=True\n        )\n\n    def build(self, input_shape):\n        (_, self.num_patches, self.patch_area) = input_shape\n\n        # Create the projection layer for the patches.\n        self.projection = layers.Dense(units=self.projection_dim)\n\n        # Create the positional embedding layer.\n        self.position_embedding = layers.Embedding(\n            input_dim=self.num_patches, output_dim=self.projection_dim\n        )\n\n        # Number of patches that will be masked.\n        self.num_mask = int(self.mask_proportion * self.num_patches)\n\n    def call(self, patches):\n        # Get the positional embeddings.\n        batch_size = tf.shape(patches)[0]\n        positions = tf.range(start=0, limit=self.num_patches, delta=1)\n        pos_embeddings = self.position_embedding(positions[tf.newaxis, ...])\n        pos_embeddings = tf.tile(\n            pos_embeddings, [batch_size, 1, 1]\n        )  # (B, num_patches, projection_dim)\n\n        # Embed the patches.\n        patch_embeddings = (\n            self.projection(patches) + pos_embeddings\n        )  # (B, num_patches, projection_dim)\n\n        if self.downstream:\n            return patch_embeddings\n        else:\n            mask_indices, unmask_indices = self.get_random_indices(batch_size)\n            # The encoder input is the unmasked patch embeddings. Here we gather\n            # all the patches that should be unmasked.\n            unmasked_embeddings = tf.gather(\n                patch_embeddings, unmask_indices, axis=1, batch_dims=1\n            )  # (B, unmask_numbers, projection_dim)\n\n            # Get the unmasked and masked position embeddings. We will need them\n            # for the decoder.\n            unmasked_positions = tf.gather(\n                pos_embeddings, unmask_indices, axis=1, batch_dims=1\n            )  # (B, unmask_numbers, projection_dim)\n            masked_positions = tf.gather(\n                pos_embeddings, mask_indices, axis=1, batch_dims=1\n            )  # (B, mask_numbers, projection_dim)\n\n            # Repeat the mask token number of mask times.\n            # Mask tokens replace the masks of the image.\n            mask_tokens = tf.repeat(self.mask_token, repeats=self.num_mask, axis=0)\n            mask_tokens = tf.repeat(\n                mask_tokens[tf.newaxis, ...], repeats=batch_size, axis=0\n            )\n\n            # Get the masked embeddings for the tokens.\n            masked_embeddings = self.projection(mask_tokens) + masked_positions\n            return (\n                unmasked_embeddings,  # Input to the encoder.\n                masked_embeddings,  # First part of input to the decoder.\n                unmasked_positions,  # Added to the encoder outputs.\n                mask_indices,  # The indices that were masked.\n                unmask_indices,  # The indices that were unmaksed.\n            )\n\n    def get_random_indices(self, batch_size):\n        # Create random indices from a uniform distribution and then split\n        # it into mask and unmask indices.\n        rand_indices = tf.argsort(\n            tf.random.uniform(shape=(batch_size, self.num_patches)), axis=-1\n        )\n        mask_indices = rand_indices[:, : self.num_mask]\n        unmask_indices = rand_indices[:, self.num_mask :]\n        return mask_indices, unmask_indices\n\n    def generate_masked_image(self, patches, unmask_indices):\n        # Choose a random patch and it corresponding unmask index.\n        idx = np.random.choice(patches.shape[0])\n        patch = patches[idx]\n        unmask_index = unmask_indices[idx]\n\n        # Build a numpy array of same shape as patch.\n        new_patch = np.zeros_like(patch)\n\n        # Iterate of the new_patch and plug the unmasked patches.\n        count = 0\n        for i in range(unmask_index.shape[0]):\n            new_patch[unmask_index[i]] = patch[unmask_index[i]]\n        return new_patch, idx","metadata":{"execution":{"iopub.status.busy":"2022-07-16T16:52:52.585868Z","iopub.execute_input":"2022-07-16T16:52:52.586409Z","iopub.status.idle":"2022-07-16T16:52:52.606372Z","shell.execute_reply.started":"2022-07-16T16:52:52.586372Z","shell.execute_reply":"2022-07-16T16:52:52.605379Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Simple Dense block with a dropout","metadata":{}},{"cell_type":"code","source":"def mlp(x, dropout_rate, hidden_units):\n    for units in hidden_units:\n        x = layers.Dense(units, activation=tf.nn.gelu)(x)\n        x = layers.Dropout(dropout_rate)(x)\n    return x","metadata":{"execution":{"iopub.status.busy":"2022-07-16T16:52:52.607638Z","iopub.execute_input":"2022-07-16T16:52:52.608188Z","iopub.status.idle":"2022-07-16T16:52:52.621414Z","shell.execute_reply.started":"2022-07-16T16:52:52.608151Z","shell.execute_reply":"2022-07-16T16:52:52.620419Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Using the standard VIT encoder block architecture.  \n\n![Encoder block © Machine Learning | Carnegie Mellon University](https://blog.ml.cmu.edu/wp-content/uploads/2020/02/transformer_encoder-2-283x625.png)","metadata":{}},{"cell_type":"code","source":"def create_encoder(num_heads=ENC_NUM_HEADS, num_layers=ENC_LAYERS):\n    inputs = layers.Input((None, ENC_PROJECTION_DIM))\n    x = inputs\n\n    for _ in range(num_layers):\n        # Layer normalization 1.\n        x1 = layers.LayerNormalization(epsilon=LAYER_NORM_EPS)(x)\n\n        # Create a multi-head attention layer.\n        attention_output = layers.MultiHeadAttention(\n            num_heads=num_heads, key_dim=ENC_PROJECTION_DIM, dropout=0.1\n        )(x1, x1)\n\n        # Skip connection 1.\n        x2 = layers.Add()([attention_output, x])\n\n        # Layer normalization 2.\n        x3 = layers.LayerNormalization(epsilon=LAYER_NORM_EPS)(x2)\n\n        # MLP.\n        x3 = mlp(x3, hidden_units=ENC_TRANSFORMER_UNITS, dropout_rate=0.1)\n\n        # Skip connection 2.\n        x = layers.Add()([x3, x2])\n\n    outputs = layers.LayerNormalization(epsilon=LAYER_NORM_EPS)(x)\n    return keras.Model(inputs, outputs, name=\"mae_encoder\")","metadata":{"execution":{"iopub.status.busy":"2022-07-16T16:52:52.622850Z","iopub.execute_input":"2022-07-16T16:52:52.623262Z","iopub.status.idle":"2022-07-16T16:52:52.634242Z","shell.execute_reply.started":"2022-07-16T16:52:52.623225Z","shell.execute_reply":"2022-07-16T16:52:52.633286Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#  As stated earlier, Decoder is lightwieght because we are only interested in encoder.","metadata":{}},{"cell_type":"code","source":"def create_decoder(\n    num_layers=DEC_LAYERS, num_heads=DEC_NUM_HEADS, image_size=IMAGE_SIZE\n):\n    inputs = layers.Input((NUM_PATCHES, ENC_PROJECTION_DIM))\n    x = layers.Dense(DEC_PROJECTION_DIM)(inputs)\n\n    for _ in range(num_layers):\n        # Layer normalization 1.\n        x1 = layers.LayerNormalization(epsilon=LAYER_NORM_EPS)(x)\n\n        # Create a multi-head attention layer.\n        attention_output = layers.MultiHeadAttention(\n            num_heads=num_heads, key_dim=DEC_PROJECTION_DIM, dropout=0.1\n        )(x1, x1)\n\n        # Skip connection 1.\n        x2 = layers.Add()([attention_output, x])\n\n        # Layer normalization 2.\n        x3 = layers.LayerNormalization(epsilon=LAYER_NORM_EPS)(x2)\n\n        # MLP.\n        x3 = mlp(x3, hidden_units=DEC_TRANSFORMER_UNITS, dropout_rate=0.1)\n\n        # Skip connection 2.\n        x = layers.Add()([x3, x2])\n\n    x = layers.LayerNormalization(epsilon=LAYER_NORM_EPS)(x)\n    x = layers.Flatten()(x)\n    pre_final = layers.Dense(units=image_size * image_size * 3, activation=\"sigmoid\")(x)\n    outputs = layers.Reshape((image_size, image_size, 3))(pre_final)\n\n    return keras.Model(inputs, outputs, name=\"mae_decoder\")","metadata":{"execution":{"iopub.status.busy":"2022-07-16T16:52:52.635934Z","iopub.execute_input":"2022-07-16T16:52:52.636389Z","iopub.status.idle":"2022-07-16T16:52:52.647338Z","shell.execute_reply.started":"2022-07-16T16:52:52.636348Z","shell.execute_reply":"2022-07-16T16:52:52.646358Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Putting everything together","metadata":{}},{"cell_type":"code","source":"class MaskedAutoencoder(keras.Model):\n    def __init__(\n        self,\n        train_augmentation_model,\n        test_augmentation_model,\n        patch_layer,\n        patch_encoder,\n        encoder,\n        decoder,\n        **kwargs,\n    ):\n        super().__init__(**kwargs)\n        self.train_augmentation_model = train_augmentation_model\n        self.test_augmentation_model = test_augmentation_model\n        self.patch_layer = patch_layer\n        self.patch_encoder = patch_encoder\n        self.encoder = encoder\n        self.decoder = decoder\n\n    def calculate_loss(self, images, test=False):\n        # Augment the input images.\n        if test:\n            augmented_images = self.test_augmentation_model(images)\n        else:\n            augmented_images = self.train_augmentation_model(images)\n\n        # Patch the augmented images.\n        patches = self.patch_layer(augmented_images)\n\n        # Encode the patches.\n        (\n            unmasked_embeddings,\n            masked_embeddings,\n            unmasked_positions,\n            mask_indices,\n            unmask_indices,\n        ) = self.patch_encoder(patches)\n\n        # Pass the unmaksed patche to the encoder.\n        encoder_outputs = self.encoder(unmasked_embeddings)\n\n        # Create the decoder inputs.\n        encoder_outputs = encoder_outputs + unmasked_positions\n        decoder_inputs = tf.concat([encoder_outputs, masked_embeddings], axis=1)\n\n        # Decode the inputs.\n        decoder_outputs = self.decoder(decoder_inputs)\n        decoder_patches = self.patch_layer(decoder_outputs)\n\n        loss_patch = tf.gather(patches, mask_indices, axis=1, batch_dims=1)\n        loss_output = tf.gather(decoder_patches, mask_indices, axis=1, batch_dims=1)\n\n        # Compute the total loss.\n        total_loss = self.compiled_loss(loss_patch, loss_output)\n\n        return total_loss, loss_patch, loss_output\n\n    def train_step(self, images):\n        with tf.GradientTape() as tape:\n            total_loss, loss_patch, loss_output = self.calculate_loss(images)\n\n        # Apply gradients.\n        train_vars = [\n            self.train_augmentation_model.trainable_variables,\n            self.patch_layer.trainable_variables,\n            self.patch_encoder.trainable_variables,\n            self.encoder.trainable_variables,\n            self.decoder.trainable_variables,\n        ]\n        grads = tape.gradient(total_loss, train_vars)\n        tv_list = []\n        for (grad, var) in zip(grads, train_vars):\n            for g, v in zip(grad, var):\n                tv_list.append((g, v))\n        self.optimizer.apply_gradients(tv_list)\n\n        # Report progress.\n        self.compiled_metrics.update_state(loss_patch, loss_output)\n        return {m.name: m.result() for m in self.metrics}\n\n    def test_step(self, images):\n        total_loss, loss_patch, loss_output = self.calculate_loss(images, test=True)\n\n        # Update the trackers.\n        self.compiled_metrics.update_state(loss_patch, loss_output)\n        return {m.name: m.result() for m in self.metrics}","metadata":{"execution":{"iopub.status.busy":"2022-07-16T16:52:52.649054Z","iopub.execute_input":"2022-07-16T16:52:52.649469Z","iopub.status.idle":"2022-07-16T16:52:52.668073Z","shell.execute_reply.started":"2022-07-16T16:52:52.649361Z","shell.execute_reply":"2022-07-16T16:52:52.666940Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Creating Model object","metadata":{}},{"cell_type":"code","source":"train_augmentation_model = get_train_augmentation_model()\ntest_augmentation_model = get_test_augmentation_model()\npatch_layer = Patches()\npatch_encoder = PatchEncoder()\nencoder = create_encoder()\ndecoder = create_decoder()\n\nmae_model = MaskedAutoencoder(\n    train_augmentation_model=train_augmentation_model,\n    test_augmentation_model=test_augmentation_model,\n    patch_layer=patch_layer,\n    patch_encoder=patch_encoder,\n    encoder=encoder,\n    decoder=decoder,\n)","metadata":{"execution":{"iopub.status.busy":"2022-07-16T16:52:52.669722Z","iopub.execute_input":"2022-07-16T16:52:52.670060Z","iopub.status.idle":"2022-07-16T16:52:53.481590Z","shell.execute_reply.started":"2022-07-16T16:52:52.670025Z","shell.execute_reply":"2022-07-16T16:52:53.480665Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Callbacks","metadata":{}},{"cell_type":"code","source":"# Taking a batch of test inputs to measure model's progress.\ntest_images = next(iter(valid_ds))\n\n\nclass TrainMonitor(keras.callbacks.Callback):\n    def __init__(self, epoch_interval=None):\n        self.epoch_interval = epoch_interval\n\n    def on_epoch_end(self, epoch, logs=None):\n        if self.epoch_interval and epoch % self.epoch_interval == 0:\n            test_augmented_images = self.model.test_augmentation_model(test_images)\n            test_patches = self.model.patch_layer(test_augmented_images)\n            (\n                test_unmasked_embeddings,\n                test_masked_embeddings,\n                test_unmasked_positions,\n                test_mask_indices,\n                test_unmask_indices,\n            ) = self.model.patch_encoder(test_patches)\n            test_encoder_outputs = self.model.encoder(test_unmasked_embeddings)\n            test_encoder_outputs = test_encoder_outputs + test_unmasked_positions\n            test_decoder_inputs = tf.concat(\n                [test_encoder_outputs, test_masked_embeddings], axis=1\n            )\n            test_decoder_outputs = self.model.decoder(test_decoder_inputs)\n\n            # Show a maksed patch image.\n            test_masked_patch, idx = self.model.patch_encoder.generate_masked_image(\n                test_patches, test_unmask_indices\n            )\n            print(f\"\\nIdx chosen: {idx}\")\n            original_image = test_augmented_images[idx]\n            masked_image = self.model.patch_layer.reconstruct_from_patch(\n                test_masked_patch\n            )\n            reconstructed_image = test_decoder_outputs[idx]\n\n            fig, ax = plt.subplots(nrows=1, ncols=3, figsize=(15, 5))\n            ax[0].imshow(original_image)\n            ax[0].set_title(f\"Original: {epoch:03d}\")\n\n            ax[1].imshow(masked_image)\n            ax[1].set_title(f\"Masked: {epoch:03d}\")\n\n            ax[2].imshow(reconstructed_image)\n            ax[2].set_title(f\"Resonstructed: {epoch:03d}\")\n\n            plt.show()\n            plt.close()","metadata":{"execution":{"iopub.status.busy":"2022-07-16T16:52:53.483152Z","iopub.execute_input":"2022-07-16T16:52:53.483488Z","iopub.status.idle":"2022-07-16T16:52:53.677629Z","shell.execute_reply.started":"2022-07-16T16:52:53.483453Z","shell.execute_reply":"2022-07-16T16:52:53.674594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# LR schedule as mentioned in the paper","metadata":{}},{"cell_type":"code","source":"\n\nclass WarmUpCosine(keras.optimizers.schedules.LearningRateSchedule):\n    def __init__(\n        self, learning_rate_base, total_steps, warmup_learning_rate, warmup_steps\n    ):\n        super(WarmUpCosine, self).__init__()\n\n        self.learning_rate_base = learning_rate_base\n        self.total_steps = total_steps\n        self.warmup_learning_rate = warmup_learning_rate\n        self.warmup_steps = warmup_steps\n        self.pi = tf.constant(np.pi)\n\n    def __call__(self, step):\n        if self.total_steps < self.warmup_steps:\n            raise ValueError(\"Total_steps must be larger or equal to warmup_steps.\")\n\n        cos_annealed_lr = tf.cos(\n            self.pi\n            * (tf.cast(step, tf.float32) - self.warmup_steps)\n            / float(self.total_steps - self.warmup_steps)\n        )\n        learning_rate = 0.5 * self.learning_rate_base * (1 + cos_annealed_lr)\n\n        if self.warmup_steps > 0:\n            if self.learning_rate_base < self.warmup_learning_rate:\n                raise ValueError(\n                    \"Learning_rate_base must be larger or equal to \"\n                    \"warmup_learning_rate.\"\n                )\n            slope = (\n                self.learning_rate_base - self.warmup_learning_rate\n            ) / self.warmup_steps\n            warmup_rate = slope * tf.cast(step, tf.float32) + self.warmup_learning_rate\n            learning_rate = tf.where(\n                step < self.warmup_steps, warmup_rate, learning_rate\n            )\n        return tf.where(\n            step > self.total_steps, 0.0, learning_rate, name=\"learning_rate\"\n        )","metadata":{"execution":{"iopub.status.busy":"2022-07-16T16:52:53.679175Z","iopub.execute_input":"2022-07-16T16:52:53.679587Z","iopub.status.idle":"2022-07-16T16:52:53.704824Z","shell.execute_reply.started":"2022-07-16T16:52:53.679534Z","shell.execute_reply":"2022-07-16T16:52:53.703324Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"total_steps = int((1281167  / BATCH_SIZE) * EPOCHS) # we know xtrain is 50000\nwarmup_epoch_percentage = 0.15\nwarmup_steps = int(total_steps * warmup_epoch_percentage)\nscheduled_lrs = WarmUpCosine(\n    learning_rate_base=LEARNING_RATE,\n    total_steps=total_steps,\n    warmup_learning_rate=0.0,\n    warmup_steps=warmup_steps,\n)\n# Assemble the callbacks.\ntrain_callbacks = [TrainMonitor(epoch_interval=5)]","metadata":{"execution":{"iopub.status.busy":"2022-07-16T16:52:53.706481Z","iopub.execute_input":"2022-07-16T16:52:53.706843Z","iopub.status.idle":"2022-07-16T16:52:53.721754Z","shell.execute_reply.started":"2022-07-16T16:52:53.706807Z","shell.execute_reply":"2022-07-16T16:52:53.716670Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"optimizer = tf.optimizers.Adam(learning_rate=scheduled_lrs)\n\n# Compile and pretrain the model.\nmae_model.compile(\n    optimizer=optimizer, loss=keras.losses.MeanSquaredError(), metrics=[\"mae\"]\n)\n","metadata":{"execution":{"iopub.status.busy":"2022-07-16T16:52:53.725174Z","iopub.execute_input":"2022-07-16T16:52:53.725789Z","iopub.status.idle":"2022-07-16T16:52:53.748503Z","shell.execute_reply.started":"2022-07-16T16:52:53.725751Z","shell.execute_reply":"2022-07-16T16:52:53.747165Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history = mae_model.fit(\n    train_ds, epochs=5, validation_data=valid_ds, callbacks=train_callbacks,\n)\n","metadata":{"execution":{"iopub.status.busy":"2022-07-16T16:52:53.753632Z","iopub.execute_input":"2022-07-16T16:52:53.754322Z","iopub.status.idle":"2022-07-16T16:53:08.870160Z","shell.execute_reply.started":"2022-07-16T16:52:53.754284Z","shell.execute_reply":"2022-07-16T16:53:08.867340Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n# Measure its performance.\nloss, mae = mae_model.evaluate(valid_ds)\nprint(f\"Loss: {loss:.2f}\")\nprint(f\"MAE: {mae:.2f}\")","metadata":{"execution":{"iopub.status.busy":"2022-07-16T16:53:11.904437Z","iopub.execute_input":"2022-07-16T16:53:11.905143Z","iopub.status.idle":"2022-07-16T16:53:13.674837Z","shell.execute_reply.started":"2022-07-16T16:53:11.905105Z","shell.execute_reply":"2022-07-16T16:53:13.672024Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Saving everything for linear probing later.","metadata":{}},{"cell_type":"code","source":"train_augmentation_model.save('train_augmentation_model.h5')\ntest_augmentation_model.save('test_augmentation_model.h5')\nencoder.save('encoder.h5')\ndecoder.save('decoder.h5')","metadata":{"execution":{"iopub.status.busy":"2022-07-16T16:53:59.147168Z","iopub.execute_input":"2022-07-16T16:53:59.147523Z","iopub.status.idle":"2022-07-16T16:53:59.661204Z","shell.execute_reply.started":"2022-07-16T16:53:59.147490Z","shell.execute_reply":"2022-07-16T16:53:59.658620Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Thank You!\nkindly refer to version 3 of this notebook if you want to see cell outputs, I ran out of GPU credits to re-run with all the markdown. :)","metadata":{}}]}