{"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":"## Product Catptioning with H&M data\nIn this notebook, I follow a tutorial from Keras (https://keras.io/examples/vision/image_captioning) that build an image captioning model and use it on data from https://www.kaggle.com/competitions/h-and-m-personalized-fashion-recommendations\n\nImage captioning architecture consists of three models:\n\nA CNN: used to extract the image features <br>\nA TransformerEncoder: The extracted image features are then passed to a Transformer based encoder that generates a new representation of the inputs<br>\nA TransformerDecoder: This model takes the encoder output and the text data (sequences) as inputs and tries to learn to generate the caption.<br>","metadata":{"papermill":{"duration":0.008112,"end_time":"2022-07-14T12:37:52.033447","exception":false,"start_time":"2022-07-14T12:37:52.025335","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"input file is generated in another notebook and it is easy to generate using articles.csv file os package from python\n<br>\nOutput model weights are saved in output folder along with subclassed_model.txt which is necessary to save a subclassed model","metadata":{"papermill":{"duration":0.007277,"end_time":"2022-07-14T12:37:52.048238","exception":false,"start_time":"2022-07-14T12:37:52.040961","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import os\nimport re\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport pandas as pd\nimport tensorflow as tf\nfrom tensorflow import keras\nfrom tensorflow.keras import layers\nfrom tensorflow.keras.applications import efficientnet\nfrom tensorflow.keras.layers import TextVectorization\nfrom sklearn.model_selection import train_test_split\nimport random\nimport warnings\n\nwarnings.filterwarnings('ignore')\n\ndef sample_from_dict(d, sample=5):\n    keys = random.sample(list(d), sample)\n    values = [d[k] for k in keys]\n    return dict(zip(keys, values))\n\n\nseed = 111\ntf.random.set_seed(seed)","metadata":{"execution":{"iopub.execute_input":"2022-07-14T12:37:52.063379Z","iopub.status.busy":"2022-07-14T12:37:52.062939Z","iopub.status.idle":"2022-07-14T12:37:59.445224Z","shell.execute_reply":"2022-07-14T12:37:59.444190Z"},"id":"x4PfHLiTosIn","papermill":{"duration":7.393074,"end_time":"2022-07-14T12:37:59.448118","exception":false,"start_time":"2022-07-14T12:37:52.055044","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Desired image dimensions\nIMAGE_SIZE = (299, 299)\n\n# Vocabulary size\nVOCAB_SIZE = 100000\n\n# Fixed length allowed for any sequence\nSEQ_LENGTH = 10\n\n# Dimension for the image embeddings and token embeddings\nEMBED_DIM = 512\n\n# Per-layer units in the feed-forward network\nFF_DIM = 512\n\n# Other training parameters\nBATCH_SIZE = 64\nEPOCHS = 20\nAUTOTUNE = tf.data.AUTOTUNE","metadata":{"execution":{"iopub.execute_input":"2022-07-14T12:37:59.464189Z","iopub.status.busy":"2022-07-14T12:37:59.463596Z","iopub.status.idle":"2022-07-14T12:37:59.470546Z","shell.execute_reply":"2022-07-14T12:37:59.469684Z"},"id":"-dFJBmMNosIp","papermill":{"duration":0.017086,"end_time":"2022-07-14T12:37:59.472479","exception":false,"start_time":"2022-07-14T12:37:59.455393","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data = pd.read_csv('../input/hm-image-path-with-desc/article_with_paths_and_desc.csv', dtype=str)\ndata.rename({'article_id':'id', 'detail_desc':'desc', 'path':'path'}, inplace=True, axis=1)","metadata":{"execution":{"iopub.execute_input":"2022-07-14T12:37:59.488742Z","iopub.status.busy":"2022-07-14T12:37:59.487292Z","iopub.status.idle":"2022-07-14T12:38:00.125961Z","shell.execute_reply":"2022-07-14T12:38:00.125005Z"},"papermill":{"duration":0.648905,"end_time":"2022-07-14T12:38:00.128467","exception":false,"start_time":"2022-07-14T12:37:59.479562","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data.head()","metadata":{"execution":{"iopub.execute_input":"2022-07-14T12:38:00.144767Z","iopub.status.busy":"2022-07-14T12:38:00.144478Z","iopub.status.idle":"2022-07-14T12:38:00.160581Z","shell.execute_reply":"2022-07-14T12:38:00.159557Z"},"papermill":{"duration":0.02717,"end_time":"2022-07-14T12:38:00.163549","exception":false,"start_time":"2022-07-14T12:38:00.136379","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def map_list(item):\n    item = \"<start> \" + item.strip() + \" <end>\"\n    return [item]\n\ndata_dict = pd.Series(data.desc.map(map_list).values,index=data.path).to_dict()","metadata":{"execution":{"iopub.execute_input":"2022-07-14T12:38:00.179920Z","iopub.status.busy":"2022-07-14T12:38:00.179642Z","iopub.status.idle":"2022-07-14T12:38:00.542565Z","shell.execute_reply":"2022-07-14T12:38:00.541559Z"},"papermill":{"duration":0.373945,"end_time":"2022-07-14T12:38:00.544912","exception":false,"start_time":"2022-07-14T12:38:00.170967","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_from_dict(data_dict, 5)","metadata":{"execution":{"iopub.execute_input":"2022-07-14T12:38:00.561154Z","iopub.status.busy":"2022-07-14T12:38:00.560853Z","iopub.status.idle":"2022-07-14T12:38:00.571935Z","shell.execute_reply":"2022-07-14T12:38:00.570859Z"},"papermill":{"duration":0.021542,"end_time":"2022-07-14T12:38:00.574110","exception":false,"start_time":"2022-07-14T12:38:00.552568","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Preparing the dataset","metadata":{"id":"NN1OdJMxosIq","papermill":{"duration":0.007128,"end_time":"2022-07-14T12:38:00.589117","exception":false,"start_time":"2022-07-14T12:38:00.581989","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def train_val_split(caption_data, train_size=0.8, shuffle=True):\n\n    # 1. Get the list of all image names\n    all_images = list(caption_data.keys())\n\n    # 2. Shuffle if necessary\n    if shuffle:\n        np.random.shuffle(all_images)\n\n    # 3. Split into training and validation sets\n    train_size = int(len(caption_data) * train_size)\n\n    training_data = {\n        img_name: caption_data[img_name] for img_name in all_images[:train_size]\n    }\n    validation_data = {\n        img_name: caption_data[img_name] for img_name in all_images[train_size:]\n    }\n\n    # 4. Return the splits\n    return training_data, validation_data\n\n\ncaptions_mapping, text_data = data_dict, data.desc.values\n\n# Split the dataset into training and validation sets\ntrain_data, valid_data = train_val_split(captions_mapping)\n\nprint(\"Number of training samples: \", len(train_data))\nprint(\"Number of validation samples: \", len(valid_data))","metadata":{"execution":{"iopub.execute_input":"2022-07-14T12:38:00.605207Z","iopub.status.busy":"2022-07-14T12:38:00.604915Z","iopub.status.idle":"2022-07-14T12:38:00.660260Z","shell.execute_reply":"2022-07-14T12:38:00.659356Z"},"id":"48SX4BdKosIq","papermill":{"duration":0.065829,"end_time":"2022-07-14T12:38:00.662303","exception":false,"start_time":"2022-07-14T12:38:00.596474","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_from_dict(captions_mapping, 1)","metadata":{"execution":{"iopub.execute_input":"2022-07-14T12:38:00.678953Z","iopub.status.busy":"2022-07-14T12:38:00.678624Z","iopub.status.idle":"2022-07-14T12:38:00.685992Z","shell.execute_reply":"2022-07-14T12:38:00.684975Z"},"papermill":{"duration":0.018438,"end_time":"2022-07-14T12:38:00.688238","exception":false,"start_time":"2022-07-14T12:38:00.669800","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"text_data[:5]","metadata":{"execution":{"iopub.execute_input":"2022-07-14T12:38:00.703939Z","iopub.status.busy":"2022-07-14T12:38:00.703692Z","iopub.status.idle":"2022-07-14T12:38:00.709936Z","shell.execute_reply":"2022-07-14T12:38:00.708887Z"},"papermill":{"duration":0.016485,"end_time":"2022-07-14T12:38:00.712083","exception":false,"start_time":"2022-07-14T12:38:00.695598","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Vectorizing the text data\n\nWe'll use the `TextVectorization` layer to vectorize the text data,\nthat is to say, to turn the\noriginal strings into integer sequences where each integer represents the index of\na word in a vocabulary. We will use a custom string standardization scheme\n(strip punctuation characters except `<` and `>`) and the default\nsplitting scheme (split on whitespace).","metadata":{"id":"xlA7B3lMosIr","papermill":{"duration":0.007448,"end_time":"2022-07-14T12:38:00.727626","exception":false,"start_time":"2022-07-14T12:38:00.720178","status":"completed"},"tags":[]}},{"cell_type":"code","source":"\ndef custom_standardization(input_string):\n    lowercase = tf.strings.lower(input_string)\n    return tf.strings.regex_replace(lowercase, \"[%s]\" % re.escape(strip_chars), \"\")\n\n\nstrip_chars = \"!\\\"#$%&'()*+,-./:;<=>?@[\\]^_`{|}~\"\nstrip_chars = strip_chars.replace(\"<\", \"\")\nstrip_chars = strip_chars.replace(\">\", \"\")\n\nvectorization = TextVectorization(\n    max_tokens=VOCAB_SIZE,\n    output_mode=\"int\",\n    output_sequence_length=SEQ_LENGTH,\n    standardize=custom_standardization,\n)\nvectorization.adapt(text_data)\n\n# Data augmentation for image data\nimage_augmentation = keras.Sequential(\n    [\n        layers.RandomFlip(\"horizontal\"),\n        layers.RandomRotation(0.2),\n        layers.RandomContrast(0.3),\n    ]\n)\n","metadata":{"execution":{"iopub.execute_input":"2022-07-14T12:38:00.743729Z","iopub.status.busy":"2022-07-14T12:38:00.743478Z","iopub.status.idle":"2022-07-14T12:38:09.340720Z","shell.execute_reply":"2022-07-14T12:38:09.339754Z"},"id":"m31Ik5eEosIs","papermill":{"duration":8.607925,"end_time":"2022-07-14T12:38:09.343064","exception":false,"start_time":"2022-07-14T12:38:00.735139","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Building a `tf.data.Dataset` pipeline for training\n\nWe will generate pairs of images and corresponding captions using a `tf.data.Dataset` object.\nThe pipeline consists of two steps:\n\n1. Read the image from the disk\n2. Tokenize all the five captions corresponding to the image","metadata":{"id":"54yq7rniosIt","papermill":{"duration":0.008165,"end_time":"2022-07-14T12:38:09.360300","exception":false,"start_time":"2022-07-14T12:38:09.352135","status":"completed"},"tags":[]}},{"cell_type":"code","source":"\ndef decode_and_resize(img_path):\n    img = tf.io.read_file(img_path)\n    img = tf.image.decode_jpeg(img, channels=3)\n    img = tf.image.resize(img, IMAGE_SIZE)\n    img = tf.image.convert_image_dtype(img, tf.float32)\n    return img\n\n\ndef process_input(img_path, captions):\n    return decode_and_resize(img_path), vectorization(captions)\n\n\ndef make_dataset(images, captions):\n    dataset = tf.data.Dataset.from_tensor_slices((images, captions))\n    dataset = dataset.shuffle(len(images))\n    dataset = dataset.map(process_input, num_parallel_calls=AUTOTUNE)\n    dataset = dataset.batch(BATCH_SIZE).prefetch(AUTOTUNE)\n\n    return dataset\n\n\n# Pass the list of images and the list of corresponding captions\ntrain_dataset = make_dataset(list(train_data.keys()), list(train_data.values()))\n\nvalid_dataset = make_dataset(list(valid_data.keys()), list(valid_data.values()))","metadata":{"execution":{"iopub.execute_input":"2022-07-14T12:38:09.377492Z","iopub.status.busy":"2022-07-14T12:38:09.376944Z","iopub.status.idle":"2022-07-14T12:38:10.881688Z","shell.execute_reply":"2022-07-14T12:38:10.880719Z"},"id":"f8yE8MnsosIt","papermill":{"duration":1.516426,"end_time":"2022-07-14T12:38:10.884538","exception":false,"start_time":"2022-07-14T12:38:09.368112","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ex = next(train_dataset.take(1).as_numpy_iterator())\nex[0].shape, ex[1].shape","metadata":{"execution":{"iopub.execute_input":"2022-07-14T12:38:10.902041Z","iopub.status.busy":"2022-07-14T12:38:10.901752Z","iopub.status.idle":"2022-07-14T12:38:11.872876Z","shell.execute_reply":"2022-07-14T12:38:11.871910Z"},"papermill":{"duration":0.982495,"end_time":"2022-07-14T12:38:11.875167","exception":false,"start_time":"2022-07-14T12:38:10.892672","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Building the model\n\nOur image captioning architecture consists of three models:\n\n1. A CNN: used to extract the image features\n2. A TransformerEncoder: The extracted image features are then passed to a Transformer\n                    based encoder that generates a new representation of the inputs\n3. A TransformerDecoder: This model takes the encoder output and the text data\n                    (sequences) as inputs and tries to learn to generate the caption.","metadata":{"id":"x6LZtHfJosIt","papermill":{"duration":0.007701,"end_time":"2022-07-14T12:38:11.891344","exception":false,"start_time":"2022-07-14T12:38:11.883643","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def get_cnn_model():\n    base_model = efficientnet.EfficientNetB0(\n        input_shape=(*IMAGE_SIZE, 3), include_top=False, weights=\"imagenet\",\n    )\n    # We freeze our feature extractor\n    base_model.trainable = False\n    base_model_out = base_model.output\n    base_model_out = layers.Reshape((-1, base_model_out.shape[-1]))(base_model_out)\n    cnn_model = keras.models.Model(base_model.input, base_model_out)\n    return cnn_model\n\n\nclass TransformerEncoderBlock(layers.Layer):\n    def __init__(self, embed_dim, dense_dim, num_heads, **kwargs):\n        super().__init__(**kwargs)\n        self.embed_dim = embed_dim\n        self.dense_dim = dense_dim\n        self.num_heads = num_heads\n        self.attention_1 = layers.MultiHeadAttention(\n            num_heads=num_heads, key_dim=embed_dim, dropout=0.0\n        )\n        self.layernorm_1 = layers.LayerNormalization()\n        self.layernorm_2 = layers.LayerNormalization()\n        self.dense_1 = layers.Dense(embed_dim, activation=\"relu\")\n\n    def call(self, inputs, training, mask=None):\n        inputs = self.layernorm_1(inputs)\n        inputs = self.dense_1(inputs)\n\n        attention_output_1 = self.attention_1(\n            query=inputs,\n            value=inputs,\n            key=inputs,\n            attention_mask=None,\n            training=training,\n        )\n        out_1 = self.layernorm_2(inputs + attention_output_1)\n        return out_1\n\n\nclass PositionalEmbedding(layers.Layer):\n    def __init__(self, sequence_length, vocab_size, embed_dim, **kwargs):\n        super().__init__(**kwargs)\n        self.token_embeddings = layers.Embedding(\n            input_dim=vocab_size, output_dim=embed_dim\n        )\n        self.position_embeddings = layers.Embedding(\n            input_dim=sequence_length, output_dim=embed_dim\n        )\n        self.sequence_length = sequence_length\n        self.vocab_size = vocab_size\n        self.embed_dim = embed_dim\n        self.embed_scale = tf.math.sqrt(tf.cast(embed_dim, tf.float32))\n\n    def call(self, inputs):\n        length = tf.shape(inputs)[-1]\n        positions = tf.range(start=0, limit=length, delta=1)\n        embedded_tokens = self.token_embeddings(inputs)\n        embedded_tokens = embedded_tokens * self.embed_scale\n        embedded_positions = self.position_embeddings(positions)\n        return embedded_tokens + embedded_positions\n\n    def compute_mask(self, inputs, mask=None):\n        return tf.math.not_equal(inputs, 0)\n\n\nclass TransformerDecoderBlock(layers.Layer):\n    def __init__(self, embed_dim, ff_dim, num_heads, **kwargs):\n        super().__init__(**kwargs)\n        self.embed_dim = embed_dim\n        self.ff_dim = ff_dim\n        self.num_heads = num_heads\n        self.attention_1 = layers.MultiHeadAttention(\n            num_heads=num_heads, key_dim=embed_dim, dropout=0.1\n        )\n        self.attention_2 = layers.MultiHeadAttention(\n            num_heads=num_heads, key_dim=embed_dim, dropout=0.1\n        )\n        self.ffn_layer_1 = layers.Dense(ff_dim, activation=\"relu\")\n        self.ffn_layer_2 = layers.Dense(embed_dim)\n\n        self.layernorm_1 = layers.LayerNormalization()\n        self.layernorm_2 = layers.LayerNormalization()\n        self.layernorm_3 = layers.LayerNormalization()\n\n        self.embedding = PositionalEmbedding(\n            embed_dim=EMBED_DIM, sequence_length=SEQ_LENGTH, vocab_size=VOCAB_SIZE\n        )\n        self.out = layers.Dense(VOCAB_SIZE, activation=\"softmax\")\n\n        self.dropout_1 = layers.Dropout(0.3)\n        self.dropout_2 = layers.Dropout(0.5)\n        self.supports_masking = True\n\n    def call(self, inputs, encoder_outputs, training, mask=None):\n        inputs = self.embedding(inputs)\n        causal_mask = self.get_causal_attention_mask(inputs)\n\n        if mask is not None:\n            padding_mask = tf.cast(mask[:, :, tf.newaxis], dtype=tf.int32)\n            combined_mask = tf.cast(mask[:, tf.newaxis, :], dtype=tf.int32)\n            combined_mask = tf.minimum(combined_mask, causal_mask)\n\n        attention_output_1 = self.attention_1(\n            query=inputs,\n            value=inputs,\n            key=inputs,\n            attention_mask=combined_mask,\n            training=training,\n        )\n        out_1 = self.layernorm_1(inputs + attention_output_1)\n\n        attention_output_2 = self.attention_2(\n            query=out_1,\n            value=encoder_outputs,\n            key=encoder_outputs,\n            attention_mask=padding_mask,\n            training=training,\n        )\n        out_2 = self.layernorm_2(out_1 + attention_output_2)\n\n        ffn_out = self.ffn_layer_1(out_2)\n        ffn_out = self.dropout_1(ffn_out, training=training)\n        ffn_out = self.ffn_layer_2(ffn_out)\n\n        ffn_out = self.layernorm_3(ffn_out + out_2, training=training)\n        ffn_out = self.dropout_2(ffn_out, training=training)\n        preds = self.out(ffn_out)\n        return preds\n\n    def get_causal_attention_mask(self, inputs):\n        input_shape = tf.shape(inputs)\n        batch_size, sequence_length = input_shape[0], input_shape[1]\n        i = tf.range(sequence_length)[:, tf.newaxis]\n        j = tf.range(sequence_length)\n        mask = tf.cast(i >= j, dtype=\"int32\")\n        mask = tf.reshape(mask, (1, input_shape[1], input_shape[1]))\n        mult = tf.concat(\n            [tf.expand_dims(batch_size, -1), tf.constant([1, 1], dtype=tf.int32)],\n            axis=0,\n        )\n        return tf.tile(mask, mult)\n\n\nclass ImageCaptioningModel(keras.Model):\n    def __init__(\n        self, cnn_model, encoder, decoder, num_captions_per_image=1, image_aug=None,\n    ):\n        super().__init__()\n        self.cnn_model = cnn_model\n        self.encoder = encoder\n        self.decoder = decoder\n        self.loss_tracker = keras.metrics.Mean(name=\"loss\")\n        self.acc_tracker = keras.metrics.Mean(name=\"accuracy\")\n        self.num_captions_per_image = num_captions_per_image\n        self.image_aug = image_aug\n\n    def calculate_loss(self, y_true, y_pred, mask):\n        loss = self.loss(y_true, y_pred)\n        mask = tf.cast(mask, dtype=loss.dtype)\n        loss *= mask\n        return tf.reduce_sum(loss) / tf.reduce_sum(mask)\n\n    def calculate_accuracy(self, y_true, y_pred, mask):\n        accuracy = tf.equal(y_true, tf.argmax(y_pred, axis=2))\n        accuracy = tf.math.logical_and(mask, accuracy)\n        accuracy = tf.cast(accuracy, dtype=tf.float32)\n        mask = tf.cast(mask, dtype=tf.float32)\n        return tf.reduce_sum(accuracy) / tf.reduce_sum(mask)\n\n    def _compute_caption_loss_and_acc(self, img_embed, batch_seq, training=True):\n        encoder_out = self.encoder(img_embed, training=training)\n        batch_seq_inp = batch_seq[:, :-1]\n        batch_seq_true = batch_seq[:, 1:]\n        mask = tf.math.not_equal(batch_seq_true, 0)\n        batch_seq_pred = self.decoder(\n            batch_seq_inp, encoder_out, training=training, mask=mask\n        )\n        loss = self.calculate_loss(batch_seq_true, batch_seq_pred, mask)\n        acc = self.calculate_accuracy(batch_seq_true, batch_seq_pred, mask)\n        return loss, acc\n\n    def train_step(self, batch_data):\n        batch_img, batch_seq = batch_data\n        batch_loss = 0\n        batch_acc = 0\n\n        if self.image_aug:\n            batch_img = self.image_aug(batch_img)\n\n        # 1. Get image embeddings\n        img_embed = self.cnn_model(batch_img)\n\n        # 2. Pass each of the five captions one by one to the decoder\n        # along with the encoder outputs and compute the loss as well as accuracy\n        # for each caption.\n        for i in range(self.num_captions_per_image):\n            with tf.GradientTape() as tape:\n                loss, acc = self._compute_caption_loss_and_acc(\n                    img_embed, batch_seq[:, i, :], training=True\n                )\n\n                # 3. Update loss and accuracy\n                batch_loss += loss\n                batch_acc += acc\n\n            # 4. Get the list of all the trainable weights\n            train_vars = (\n                self.encoder.trainable_variables + self.decoder.trainable_variables\n            )\n\n            # 5. Get the gradients\n            grads = tape.gradient(loss, train_vars)\n\n            # 6. Update the trainable weights\n            self.optimizer.apply_gradients(zip(grads, train_vars))\n\n        # 7. Update the trackers\n        batch_acc /= float(self.num_captions_per_image)\n        self.loss_tracker.update_state(batch_loss)\n        self.acc_tracker.update_state(batch_acc)\n\n        # 8. Return the loss and accuracy values\n        return {\"loss\": self.loss_tracker.result(), \"acc\": self.acc_tracker.result()}\n\n    def test_step(self, batch_data):\n        batch_img, batch_seq = batch_data\n        batch_loss = 0\n        batch_acc = 0\n\n        # 1. Get image embeddings\n        img_embed = self.cnn_model(batch_img)\n\n        # 2. Pass each of the five captions one by one to the decoder\n        # along with the encoder outputs and compute the loss as well as accuracy\n        # for each caption.\n        for i in range(self.num_captions_per_image):\n            loss, acc = self._compute_caption_loss_and_acc(\n                img_embed, batch_seq[:, i, :], training=False\n            )\n\n            # 3. Update batch loss and batch accuracy\n            batch_loss += loss\n            batch_acc += acc\n\n        batch_acc /= float(self.num_captions_per_image)\n\n        # 4. Update the trackers\n        self.loss_tracker.update_state(batch_loss)\n        self.acc_tracker.update_state(batch_acc)\n\n        # 5. Return the loss and accuracy values\n        return {\"loss\": self.loss_tracker.result(), \"acc\": self.acc_tracker.result()}\n\n    @property\n    def metrics(self):\n        # We need to list our metrics here so the `reset_states()` can be\n        # called automatically.\n        return [self.loss_tracker, self.acc_tracker]","metadata":{"execution":{"iopub.execute_input":"2022-07-14T12:38:11.918815Z","iopub.status.busy":"2022-07-14T12:38:11.918518Z","iopub.status.idle":"2022-07-14T12:38:11.963898Z","shell.execute_reply":"2022-07-14T12:38:11.963037Z"},"id":"VvMHO8euosIu","papermill":{"duration":0.066659,"end_time":"2022-07-14T12:38:11.965837","exception":false,"start_time":"2022-07-14T12:38:11.899178","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cnn_model = get_cnn_model()\nencoder = TransformerEncoderBlock(embed_dim=EMBED_DIM, dense_dim=FF_DIM, num_heads=1)\ndecoder = TransformerDecoderBlock(embed_dim=EMBED_DIM, ff_dim=FF_DIM, num_heads=2)\ncaption_model = ImageCaptioningModel(\n    cnn_model=cnn_model, encoder=encoder, decoder=decoder, image_aug=image_augmentation)","metadata":{"execution":{"iopub.execute_input":"2022-07-14T12:38:11.983283Z","iopub.status.busy":"2022-07-14T12:38:11.982546Z","iopub.status.idle":"2022-07-14T12:38:13.720122Z","shell.execute_reply":"2022-07-14T12:38:13.719058Z"},"papermill":{"duration":1.748856,"end_time":"2022-07-14T12:38:13.722523","exception":false,"start_time":"2022-07-14T12:38:11.973667","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model training","metadata":{"id":"Ble5uiLCosIv","papermill":{"duration":0.007979,"end_time":"2022-07-14T12:38:13.739483","exception":false,"start_time":"2022-07-14T12:38:13.731504","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Define the loss function\ncross_entropy = keras.losses.SparseCategoricalCrossentropy(\n    from_logits=False, reduction=\"none\"\n)\n\n# EarlyStopping criteria\nearly_stopping = keras.callbacks.EarlyStopping(patience=3, restore_best_weights=True)\n\n\n# Learning Rate Scheduler for the optimizer\nclass LRSchedule(keras.optimizers.schedules.LearningRateSchedule):\n    def __init__(self, post_warmup_learning_rate, warmup_steps):\n        super().__init__()\n        self.post_warmup_learning_rate = post_warmup_learning_rate\n        self.warmup_steps = warmup_steps\n\n    def __call__(self, step):\n        global_step = tf.cast(step, tf.float32)\n        warmup_steps = tf.cast(self.warmup_steps, tf.float32)\n        warmup_progress = global_step / warmup_steps\n        warmup_learning_rate = self.post_warmup_learning_rate * warmup_progress\n        return tf.cond(\n            global_step < warmup_steps,\n            lambda: warmup_learning_rate,\n            lambda: self.post_warmup_learning_rate,\n        )\n\n\n# Create a learning rate schedule\nnum_train_steps = len(train_dataset) * EPOCHS\nnum_warmup_steps = num_train_steps // 15\nlr_schedule = LRSchedule(post_warmup_learning_rate=1e-4, warmup_steps=num_warmup_steps)\n\n# Compile the model\ncaption_model.compile(optimizer=keras.optimizers.Adam(lr_schedule), loss=cross_entropy)\n\n# Fit the model\nhistory = caption_model.fit(\n    train_dataset,\n    epochs=EPOCHS,\n    validation_data=valid_dataset,\n    callbacks=[early_stopping],\n)","metadata":{"execution":{"iopub.execute_input":"2022-07-14T12:38:13.758286Z","iopub.status.busy":"2022-07-14T12:38:13.756819Z","iopub.status.idle":"2022-07-14T19:22:33.030914Z","shell.execute_reply":"2022-07-14T19:22:33.028885Z"},"id":"ahdvIzfaosIv","papermill":{"duration":24259.285815,"end_time":"2022-07-14T19:22:33.033503","exception":false,"start_time":"2022-07-14T12:38:13.747688","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.plot(history.history['acc'])\nplt.plot(history.history['val_acc'])","metadata":{"execution":{"iopub.execute_input":"2022-07-14T19:22:36.321868Z","iopub.status.busy":"2022-07-14T19:22:36.320500Z","iopub.status.idle":"2022-07-14T19:22:36.554418Z","shell.execute_reply":"2022-07-14T19:22:36.553444Z"},"papermill":{"duration":1.821067,"end_time":"2022-07-14T19:22:36.556339","exception":false,"start_time":"2022-07-14T19:22:34.735272","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.plot(history.history['loss'])\nplt.plot(history.history['val_loss'])","metadata":{"execution":{"iopub.execute_input":"2022-07-14T19:22:39.836348Z","iopub.status.busy":"2022-07-14T19:22:39.835964Z","iopub.status.idle":"2022-07-14T19:22:40.007854Z","shell.execute_reply":"2022-07-14T19:22:40.006914Z"},"papermill":{"duration":1.769656,"end_time":"2022-07-14T19:22:40.009780","exception":false,"start_time":"2022-07-14T19:22:38.240124","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Check sample predictions","metadata":{"id":"CUOnH3R8osIv","papermill":{"duration":1.929292,"end_time":"2022-07-14T19:22:43.676815","exception":false,"start_time":"2022-07-14T19:22:41.747523","status":"completed"},"tags":[]}},{"cell_type":"code","source":"vocab = vectorization.get_vocabulary()\nindex_lookup = dict(zip(range(len(vocab)), vocab))\nmax_decoded_sentence_length = SEQ_LENGTH - 1\nvalid_images = list(valid_data.keys())\n\ndef generate_caption(model):\n    # Select a random image from the validation dataset\n    sample_img = np.random.choice(valid_images)\n    im_toshow = plt.imread(sample_img)\n    actual_caption = \"Actual Caption: \\n\" + data.desc[data.path == sample_img].values[0]\n\n    # Read the image from the disk\n    sample_img = decode_and_resize(sample_img)\n    img = sample_img.numpy().clip(0, 255).astype(np.uint8)\n    plt.imshow(im_toshow)\n    plt.axis('off')\n    plt.show()\n\n    # Pass the image to the CNN\n    img = tf.expand_dims(sample_img, 0)\n    img = model.cnn_model(img)\n\n    # Pass the image features to the Transformer encoder\n    encoded_img = model.encoder(img, training=False)\n\n    # Generate the caption using the Transformer decoder\n    decoded_caption = \"<start> \"\n    for i in range(max_decoded_sentence_length):\n        tokenized_caption = vectorization([decoded_caption])[:, :-1]\n        mask = tf.math.not_equal(tokenized_caption, 0)\n        predictions = model.decoder(\n            tokenized_caption, encoded_img, training=False, mask=mask\n        )\n        sampled_token_index = np.argmax(predictions[0, i, :])\n        sampled_token = index_lookup[sampled_token_index]\n        if sampled_token == \" <end>\":\n            break\n        decoded_caption += \" \" + sampled_token\n\n    decoded_caption = decoded_caption.replace(\"<start> \", \"\")\n    decoded_caption = decoded_caption.replace(\" <end>\", \"\").strip()\n   \n    print('*******************')\n    print(\"Predicted Caption: \\n \\n\", decoded_caption)\n    print('*******************')\n    print(actual_caption)\n    print('###############################################################')","metadata":{"execution":{"iopub.execute_input":"2022-07-14T19:22:47.164543Z","iopub.status.busy":"2022-07-14T19:22:47.163947Z","iopub.status.idle":"2022-07-14T19:22:47.203667Z","shell.execute_reply":"2022-07-14T19:22:47.202700Z"},"id":"FF_E2kLeosIw","papermill":{"duration":1.942976,"end_time":"2022-07-14T19:22:47.206762","exception":false,"start_time":"2022-07-14T19:22:45.263786","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Check predictions for a few samples\ngenerate_caption(caption_model)\ngenerate_caption(caption_model)\ngenerate_caption(caption_model)","metadata":{"execution":{"iopub.execute_input":"2022-07-14T19:22:50.777810Z","iopub.status.busy":"2022-07-14T19:22:50.777470Z","iopub.status.idle":"2022-07-14T19:22:52.914341Z","shell.execute_reply":"2022-07-14T19:22:52.912965Z"},"papermill":{"duration":3.834614,"end_time":"2022-07-14T19:22:52.916861","exception":false,"start_time":"2022-07-14T19:22:49.082247","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Saving subclassed model","metadata":{"papermill":{"duration":1.699575,"end_time":"2022-07-14T19:22:56.383278","exception":false,"start_time":"2022-07-14T19:22:54.683703","status":"completed"},"tags":[]}},{"cell_type":"code","source":"caption_model.save_weights('model_weights.h5')","metadata":{"execution":{"iopub.execute_input":"2022-07-14T19:22:59.648108Z","iopub.status.busy":"2022-07-14T19:22:59.647753Z","iopub.status.idle":"2022-07-14T19:23:00.760742Z","shell.execute_reply":"2022-07-14T19:23:00.759774Z"},"papermill":{"duration":2.803873,"end_time":"2022-07-14T19:23:00.763294","exception":false,"start_time":"2022-07-14T19:22:57.959421","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Reinstantiating the model and loading saved weights","metadata":{"papermill":{"duration":1.835331,"end_time":"2022-07-14T19:23:04.195686","exception":false,"start_time":"2022-07-14T19:23:02.360355","status":"completed"},"tags":[]}},{"cell_type":"code","source":"loaded_model = ImageCaptioningModel(\n    cnn_model=cnn_model, encoder=encoder, decoder=decoder, image_aug=image_augmentation,\n)\nloaded_model.built=True\nloaded_model.load_weights('./model_weights.h5')","metadata":{"execution":{"iopub.execute_input":"2022-07-14T19:23:07.647994Z","iopub.status.busy":"2022-07-14T19:23:07.646861Z","iopub.status.idle":"2022-07-14T19:23:08.511189Z","shell.execute_reply":"2022-07-14T19:23:08.510188Z"},"papermill":{"duration":2.598684,"end_time":"2022-07-14T19:23:08.513968","exception":false,"start_time":"2022-07-14T19:23:05.915284","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Check predictions for a few samples\ngenerate_caption(loaded_model)\ngenerate_caption(loaded_model)\ngenerate_caption(loaded_model)","metadata":{"execution":{"iopub.execute_input":"2022-07-14T19:23:11.837747Z","iopub.status.busy":"2022-07-14T19:23:11.837073Z","iopub.status.idle":"2022-07-14T19:23:13.891086Z","shell.execute_reply":"2022-07-14T19:23:13.889054Z"},"papermill":{"duration":3.777456,"end_time":"2022-07-14T19:23:13.893285","exception":false,"start_time":"2022-07-14T19:23:10.115829","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]}]}