{"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":"# Load data","metadata":{"_uuid":"5443792a-f37a-48b1-b239-2b8ddb341b01","_cell_guid":"39a9716b-1e8e-4084-ab67-4908252e605b","trusted":true}},{"cell_type":"code","source":"## This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n \n# Input data files are available in the read-only \"../input/\" directory\n\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\nbase_path = '/kaggle/input/11-785-fall-20-homework-4-part-2/hw4p2'\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"execution":{"iopub.status.busy":"2021-10-24T22:32:21.464187Z","iopub.execute_input":"2021-10-24T22:32:21.464570Z","iopub.status.idle":"2021-10-24T22:32:21.479188Z","shell.execute_reply.started":"2021-10-24T22:32:21.464538Z","shell.execute_reply":"2021-10-24T22:32:21.477730Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import math\n\nfrom dataclasses import dataclass\nfrom typing import List, Dict, Text\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\n\nimport torch\nimport torch.nn as nn\nfrom tqdm.notebook import tqdm\nfrom torch.utils import data\nfrom torch.nn.utils.rnn import pad_sequence\nfrom torch.nn.utils.rnn import pack_padded_sequence\nfrom torch.nn.utils.rnn import pad_packed_sequence\nfrom torch import utils\nimport torch.nn.functional as F\n\nprint(torch.cuda.is_available())\ncuda = torch.cuda.is_available()\nif cuda:  \n  dev = \"cuda:0\" \nelse:  \n  dev = \"cpu\" \n\nDEVICE = torch.device(dev)\nprint(\"Device:\", DEVICE)","metadata":{"_uuid":"35b3dcc5-3b76-4e4f-8e0a-46a01aebf159","_cell_guid":"7f9bb77c-fcef-4c4a-8675-2d87831fea9c","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2021-10-24T22:32:22.979999Z","iopub.execute_input":"2021-10-24T22:32:22.980334Z","iopub.status.idle":"2021-10-24T22:32:24.521594Z","shell.execute_reply.started":"2021-10-24T22:32:22.980280Z","shell.execute_reply":"2021-10-24T22:32:24.520678Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''\nLoading all the numpy files containing the utterance information and text information\n'''\ndef load_data(base_path = ''):\n    speech_train = np.load(os.path.join(base_path, 'train.npy'), allow_pickle=True, encoding='bytes')\n    speech_valid = np.load(os.path.join(base_path, 'dev.npy'), allow_pickle=True, encoding='bytes')\n    speech_test = np.load(os.path.join(base_path, 'test.npy'), allow_pickle=True, encoding='bytes')\n\n    transcript_train = np.load(os.path.join(base_path, 'train_transcripts.npy'), allow_pickle=True,encoding='bytes')\n    transcript_valid = np.load(os.path.join(base_path, './dev_transcripts.npy'), allow_pickle=True,encoding='bytes')\n\n    return speech_train, speech_valid, speech_test, transcript_train, transcript_valid\n\nspeech_train, speech_valid, speech_test, transcript_train, transcript_valid = load_data(base_path)","metadata":{"_uuid":"a88e08d5-dba3-4a26-b201-8d3155d906bb","_cell_guid":"1b33d30d-9bf1-4094-80d9-fece2967127c","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2021-10-24T22:31:06.803185Z","iopub.execute_input":"2021-10-24T22:31:06.803555Z","iopub.status.idle":"2021-10-24T22:32:00.200135Z","shell.execute_reply.started":"2021-10-24T22:31:06.803524Z","shell.execute_reply":"2021-10-24T22:32:00.199318Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load checkpoint\ncheckpoint = torch.load('/kaggle/input/decoder-pretrained-as-lm/checkpoint.pth')\nprint(checkpoint.keys())","metadata":{"execution":{"iopub.status.busy":"2021-10-24T22:32:25.221462Z","iopub.execute_input":"2021-10-24T22:32:25.221779Z","iopub.status.idle":"2021-10-24T22:32:29.899370Z","shell.execute_reply.started":"2021-10-24T22:32:25.221751Z","shell.execute_reply":"2021-10-24T22:32:29.898535Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Transform transcripts from letter to index","metadata":{"_uuid":"29051533-b36b-42c6-bd86-0543132b4889","_cell_guid":"ee70d7da-dbca-4690-b79b-cca4a6b07fa8","trusted":true}},{"cell_type":"code","source":"# TODO: see if unk is necessary\nLETTER_LIST = ['<pad>', 'a', 'b', 'c', 'd', 'e', 'f', 'g', 'h', 'i', 'j', 'k', 'l', 'm', 'n', 'o', 'p', 'q', \\\n               'r', 's', 't', 'u', 'v', 'w', 'x', 'y', 'z', '-', \"'\", '.', '_', '+', ' ','<sos>','<eos>']\n\ndef create_dictionaries(letter_list):\n    '''\n    Create dictionaries for letter2index and index2letter transformations\n    '''\n    index2letter = dict(enumerate(letter_list))\n    letter2index = dict()\n    for k, v in index2letter.items():\n        letter2index[v] = k\n    \n    return letter2index, index2letter\n\nletter2index, index2letter = create_dictionaries(LETTER_LIST)\n\n'''\nTransforms alphabetical input to numerical input, replace each letter by its corresponding \nindex from letter_list\n'''\ndef transform_letter_to_index(transcript, letter_list) -> List:\n    '''\n    Transform letter to index. Adds <sos> and <eos> indexes.\n    Receives transcript in batch.\n    \n    :param transcript :(N, ) Transcripts are the text input\n    :param letter_list: Letter list defined above\n    :return letter_to_index_list: Returns a list for all the transcript sentence to index\n    '''\n    \n    start_idx = letter2index['<sos>']\n    end_idx = letter2index['<eos>']\n    \n    new_transcript = []\n    \n    for sequence in transcript:\n        words = [word.decode('UTF-8') for word in sequence]\n        joined_words = ' '.join(words)\n        \n        # TODO: check if first element should be start_idx\n        new_seq = [start_idx]\n        for char in joined_words:\n            # If char is not there, then use the unk idx\n            # TODO: check if this is necessary\n            new_char = letter2index[char]\n            new_seq.append(new_char) \n        new_seq.append(end_idx)\n        \n        new_seq = np.array(new_seq)\n        new_transcript.append(new_seq)\n    \n             \n    return new_transcript\n\nINPUT_PADDING_VALUE = 0\nLABEL_PADDING_VALUE = letter2index['<pad>'] # It is also 0\nVOCAB_SIZE = len(LETTER_LIST)\nprint(\"Vocab size:\", VOCAB_SIZE)\nprint(\"Input padding value:\", INPUT_PADDING_VALUE)\nprint(\"Label padding value:\", LABEL_PADDING_VALUE)\n\nprint()\ncharacter_text_train = transform_letter_to_index(transcript_train, LETTER_LIST)\ncharacter_text_valid = transform_letter_to_index(transcript_valid, LETTER_LIST)\nprint(\"Sample transcript indexes:\", character_text_train[9])","metadata":{"_uuid":"3a457098-baf1-4fcb-a8cc-fe6056a0924c","_cell_guid":"36d6fcbe-f81a-481f-bfb4-3b29ace3288f","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2021-10-24T22:32:37.255827Z","iopub.execute_input":"2021-10-24T22:32:37.256148Z","iopub.status.idle":"2021-10-24T22:32:39.762620Z","shell.execute_reply.started":"2021-10-24T22:32:37.256119Z","shell.execute_reply":"2021-10-24T22:32:39.761568Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define Dataset class and Collate functions","metadata":{"_uuid":"87ab349f-cd30-49f4-ba25-4593b63eee2f","_cell_guid":"cd800e2f-9ab7-4dd8-a6e0-176751192f05","trusted":true}},{"cell_type":"code","source":"class Speech2TextDataset(Dataset):\n    '''\n    Dataset class for the speech to text data, this may need some tweaking in the\n    getitem method as your implementation in the collate function may be different from\n    ours. \n    '''\n    def __init__(self, speech, text=None, isTrain=True):\n        \"\"\"\n        :param speech Numpy.Array(num_samples): Each sample is a list of coefficients.\n        :param text List(num_samples): List of numpy arrays. Each array is composed of the letter indexes.\n        :isTrain bool: If train mode, getting an item returns the label too, otherwise it returns just the input.\n        \"\"\"\n        self.isTrain = isTrain     \n        self.speech = speech\n        if (text is not None):\n            self.text = text\n\n    def __len__(self):\n        return self.speech.shape[0]\n\n    def __getitem__(self, index):\n        if (self.isTrain == True):\n            return torch.tensor(self.speech[index].astype(np.float32)), torch.tensor(self.text[index])\n        else:\n            return torch.tensor(self.speech[index].astype(np.float32))\n\n\ndef collate_train(batch_data):\n    ### Return the padded speech and text data, and the length of utterance and transcript ###\n    \"\"\"\n    Args:\n        :batch_data List[Tuple]: List of (input, label)\n    Returns\n        :batch_input_padded Tensor(batch_size, input_max_len, input_dim):\n        :batch_label_padded Tensor(batch_size, label_max_len):\n        :input_lens Tensor(batch_size):\n        :label_lens Tensor(batch_size):\n    \"\"\"\n    batch_input, batch_label = zip(*batch_data)\n    \n    input_lens = torch.Tensor([len(seq) for seq in batch_input])\n    label_lens = torch.Tensor([len(seq) for seq in batch_label])\n    \n    batch_input_padded = torch.as_tensor(pad_sequence(batch_input, batch_first=True, padding_value=INPUT_PADDING_VALUE))\n    batch_label_padded = torch.as_tensor(pad_sequence(batch_label, batch_first=True, padding_value=LABEL_PADDING_VALUE))\n    \n    return batch_input_padded, batch_label_padded, input_lens, label_lens\n\n\ndef collate_test(batch_data):\n    ### Return padded speech and length of utterance ###\n    \"\"\"\n    Args:\n        :batch_data List[Tensor]: List of inputs\n    Returns\n        :batch_input_padded Tensor(batch_size, input_max_len, input_dim):\n        :input_lens Tensor(batch_size):\n    \"\"\"\n    batch_input = batch_data\n    input_lens = torch.Tensor([len(seq) for seq in batch_input])    \n    batch_input_padded = torch.as_tensor(pad_sequence(batch_input, batch_first=True, padding_value=INPUT_PADDING_VALUE))\n    \n    return batch_input_padded, input_lens\n\ntrain_dataset = Speech2TextDataset(speech_train, character_text_train)\nval_dataset = Speech2TextDataset(speech_valid, character_text_valid)","metadata":{"_uuid":"4fad7d16-ab0c-4efa-9baa-dd38242b2a7b","_cell_guid":"c8d0c202-1146-4ef7-bc72-6c2e491de7b6","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2021-10-24T22:32:41.921769Z","iopub.execute_input":"2021-10-24T22:32:41.922087Z","iopub.status.idle":"2021-10-24T22:32:41.933546Z","shell.execute_reply.started":"2021-10-24T22:32:41.922058Z","shell.execute_reply":"2021-10-24T22:32:41.932700Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define Models","metadata":{"_uuid":"cd994499-e282-456c-9d01-d30da65b2dab","_cell_guid":"dc105dae-5be1-45b7-9b5a-c962bd498346","trusted":true}},{"cell_type":"markdown","source":"## Listener: Encoder\n\nWe define a Pyramidal BLSTM of 4 layers and two final linear layers to obtain the Key and Value (useful for the Decoding part).","metadata":{"_uuid":"38e45b68-a48f-44d2-90b2-817bff05bf96","_cell_guid":"2edac509-bd08-4c59-9c19-03a756a6c2e5","trusted":true}},{"cell_type":"code","source":"class Encoder(nn.Module):\n    '''\n    Encoder takes the utterances as inputs and returns the key and value.\n    Key and value are nothing but simple projections of the output from pBLSTM network.\n    '''\n    def __init__(self, input_dim, hidden_dim, value_size=128, key_size=128, p_layers=3):\n        super(Encoder, self).__init__()\n        self.lstm = nn.LSTM(input_size=input_dim, hidden_size=hidden_dim, num_layers=1, bidirectional=True, batch_first=True)\n        \n        ### Add code to define the blocks of pBLSTMs! ###\n        pmodule_list = []\n        \n        for l in range(p_layers):\n            # Input dim is hidden_dim*2 due to bidirectionality of BLSTM\n            # Second *2 is due to the concatenation of features from pairs of consecutive timesteps\n            module = pBLSTM(hidden_dim*2*2, hidden_dim)\n            pmodule_list.append(module)\n            \n        self.p_layers = nn.ModuleList(pmodule_list)\n        \n        # Input dim is hidden_dim*2 due to bidirectionality of BLSTM\n        self.key_network = nn.Linear(hidden_dim*2, key_size)\n        self.value_network = nn.Linear(hidden_dim*2, value_size)\n\n    def forward(self, x, lens):\n        rnn_inp = pack_padded_sequence(x, lengths=lens, batch_first=True, enforce_sorted=False)\n        \n        outputs, _ = self.lstm(rnn_inp) # Packed Sequence\n\n        ### Use the outputs and pass it through the pBLSTM blocks! ###\n        for layer in self.p_layers:\n            outputs, _ = layer(outputs)\n            \n        linear_input, _ = pad_packed_sequence(outputs, batch_first=True)\n        keys = self.key_network(linear_input)\n        value = self.value_network(linear_input)\n\n        return keys, value\n    \nclass pBLSTM(nn.Module):\n    '''\n    Pyramidal BiLSTM\n    The length of utterance (speech input) can be hundereds to thousands of frames long.\n    The Paper reports that a direct LSTM implementation as Encoder resulted in slow convergence,\n    and inferior results even after extensive training.\n    The major reason is inability of AttendAndSpell operation to extract relevant information\n    from a large number of input steps.\n    '''\n    def __init__(self, input_dim, hidden_dim, sample_rate=2):\n        super(pBLSTM, self).__init__()\n        # TODO: what about batch first!\n        self.blstm = nn.LSTM(input_size=input_dim, \n                             hidden_size=hidden_dim, \n                             num_layers=1, \n                             bidirectional=True,\n                             batch_first=True)\n        self.sample_rate = sample_rate\n\n    def forward(self, x):\n        '''\n        :param x :(N, T) input to the pBLSTM: Packed Sequence.\n        :return output: (N, T, H) encoded sequence from pyramidal Bi-LSTM. Packed sequence.\n        '''\n        # Unpack input: (B, T, feature_dim)\n        padded_x, x_lens = pad_packed_sequence(x, \n                                               batch_first=True, \n                                               padding_value=INPUT_PADDING_VALUE)\n        \n        batch_size, max_len, feature_dim = padded_x.shape\n        \n        # Drop extra frames at the end\n        if max_len % self.sample_rate != 0:\n            padded_x = padded_x[:, :-(max_len % self.sample_rate), :]\n        \n        new_len = math.floor(max_len/self.sample_rate)  # New length is halved and rounded down\n        new_dim = feature_dim * self.sample_rate  # New dimension is doubled\n        reshaped_x = padded_x.contiguous().view((batch_size, new_len, new_dim))\n        \n        # Compute new lens for each sequence\n        new_x_lens = [math.floor(length/self.sample_rate) for length in x_lens]   \n        \n        # Pack the reshaped batch sequences\n        new_x = pack_padded_sequence(reshaped_x, lengths=new_x_lens, batch_first=True, enforce_sorted=False)\n        output = self.blstm(new_x)\n        \n        return output","metadata":{"_uuid":"daa5c0a0-9e94-45b3-8c22-06abba055481","_cell_guid":"7f70422d-0e63-48f2-83ce-39fc97ec7f1b","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2021-10-24T22:54:22.895606Z","iopub.execute_input":"2021-10-24T22:54:22.895982Z","iopub.status.idle":"2021-10-24T22:54:22.910020Z","shell.execute_reply.started":"2021-10-24T22:54:22.895932Z","shell.execute_reply":"2021-10-24T22:54:22.909219Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Attend and Spell: Decoder","metadata":{"_uuid":"95f38532-5eaa-472a-9d2a-81d612917c46","_cell_guid":"0a8cf184-4c65-47af-8aa6-306f3b6b9c96","trusted":true}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n%matplotlib inline\nfrom IPython.display import display\n\nclass Decoder(nn.Module):\n    '''\n    As mentioned in a previous recitation, each forward call of decoder deals with just one time step, \n    thus we use LSTMCell instead of LSLTM here.\n    The output from the second LSTMCell can be used as query here for attention module.\n    In place of value that we get from the attention, this can be replace by context we get from the attention.\n    Methods like Gumble noise and teacher forcing can also be incorporated for improving the performance.\n    '''\n    def __init__(self, vocab_size, hidden_dim, value_size=128, key_size=128, \n                 isAttended=True, teacher_force=None, with_context=True):\n        super(Decoder, self).__init__()\n        self.hidden_dim = hidden_dim\n        self.embedding = nn.Embedding(vocab_size, hidden_dim, padding_idx=0) # <pad> char is index 0\n        self.lstm1 = nn.LSTMCell(input_size=hidden_dim + value_size, hidden_size=hidden_dim)\n        self.lstm2 = nn.LSTMCell(input_size=hidden_dim, hidden_size=key_size)\n        self.teacher_force = teacher_force\n        self.isAttended = isAttended\n        self.with_context = with_context\n        if (isAttended == True):\n            self.attention = Attention()\n\n        self.character_prob = nn.Linear(key_size + value_size, vocab_size)\n        # LogSoftmax is not needed \n\n    def forward(self, key, values, text=None, isTrain=True):\n        '''\n        Takes the key and values from input, the text only if it is in training mode \n        and produces a character prediction probability.\n        \n        If only_decoder is set, then context is zero padded to ignore it.\n        \n        :param key :(N, T, key_size) Output of the Encoder Key projection layer\n        :param values: (N, T, value_size) Output of the Encoder Value projection layer\n        :param text: (N, text_len) Batch input of text with text_length\n        :param isTrain: Train or eval mode\n        :param teacher_force Float [0, 1]:\n        \n        :return predictions: Returns the character prediction probability \n        '''\n        batch_size = key.shape[0]\n        \n        if (isTrain == True):\n            max_len =  text.shape[1]\n            embeddings = self.embedding(text) # [batch_size, max_len, hidden_dim]\n        else:\n            # If it has text, then I can use the length of the text.\n            if text is None:\n                max_len = 250\n            else:\n                max_len = text.shape[1]\n        predictions = []\n        hidden_states = [None, None] # For each of the LSTMCells\n        prediction = torch.zeros(batch_size,1).to(DEVICE)#(torch.ones(batch_size, 1)*33).to(DEVICE)\n        \n        # Initialize attention context\n        context = torch.zeros((batch_size, self.hidden_dim)).to(DEVICE)\n        \n        # Attention masks of all output timesteps\n        attn_masks = []\n        \n        for i in range(max_len):\n            # * Implement Gumble noise and teacher forcing techniques \n            # * When attention is True, replace values[i,:,:] with the context you get from attention.\n            # * If you haven't implemented attention yet, then you may want to check the index and break \n            #   out of the loop so you do not get index out of range errors. \n            \n            # If not attended, just input the values projection from the encoder, and the output of the\n            # second LSTMCell is used for the character prob.\n            \n            # If not with_context, don't even pass the projection from encoder, just pass zero padding.\n            \n            if self.with_context:\n                if not self.isAttended:\n                    # Pass projection from encoder from the same timestep\n                    if i >= values.shape[1]:\n                        context = torch.zeros((batch_size, self.hidden_dim)).to(DEVICE)\n                    else:\n                        context = values[:,i,:]\n            else:\n                # Ignore utterance context, just pad with zeros\n                context = torch.zeros((batch_size, self.hidden_dim)).to(DEVICE)\n            \n                    \n            if (isTrain):\n                # Take the correct character embedding from previous output timestep\n                if self.teacher_force is not None:\n                    p_previous = self.teacher_force\n                    p_truth = 1 - self.teacher_force\n                    is_previous = int(np.random.choice([0,1], p=[p_truth, p_previous]))\n                    if is_previous:\n                        # Choose previous predicted char\n                        # TODO: if i == 0, just handle outside the teacher force conditional\n                        if i == 0:\n                            predicted_char = (torch.ones(batch_size)*letter2index['<sos>']).long().to(DEVICE)\n                        else:\n                            predicted_char = prediction.argmax(dim=1) # TODO: check what I am passing, the whole sequence or just last? \n                        \n                        char_embed = self.embedding(predicted_char)\n                    else: \n                        # Choose golden truth\n                        char_embed = embeddings[:,i,:] # [batch_size, hidden_dim]\n            else:\n                # If not training, just use previous char\n                if i == 0:\n                    predicted_char = (torch.ones(batch_size)*letter2index['<sos>']).long().to(DEVICE)\n                else:\n                    predicted_char = prediction.argmax(dim=1) # TODO: check what I am passing, the whole sequence or just last? \n                    # Just passing last...\n                char_embed = self.embedding(predicted_char)\n            \n            inp = torch.cat([char_embed, context], dim=1)\n            hidden_states[0] = self.lstm1(inp, hidden_states[0]) # (h, c)\n\n            inp_2 = hidden_states[0][0] # Just take the hidden state\n            hidden_states[1] = self.lstm2(inp_2, hidden_states[1])\n\n            ### Compute attention from the output of the second LSTM Cell ###\n            # Output has dim Tensor[batch_size, key_size]\n            output = hidden_states[1][0] # The query, a.k.a. the decoder state\n            \n            if self.with_context:\n                # Use attention context\n                if self.isAttended:\n                    context, attn_mask = self.attention(output, key, values) # Tensor[batch_size, key_size], Tensor[batch_size, max_len]\n                    attn_masks.append(attn_mask)\n                # Else, just initialize context at the beginning of loop\n            \n            prediction = self.character_prob(torch.cat([output, context], dim=1))  # Tensor[batch_size, vocab_size]\n            \n            # Add timestep dimension in second dim so that I can stack it in the end\n            predictions.append(prediction)\n        \n        predictions_cat = torch.stack(predictions, dim=1)  # Tensor[batch_size, max_len, vocab_size]\n        \n        if self.with_context:\n            if self.isAttended:\n                attn_masks = torch.stack(attn_masks, dim=1) # Tensor[batch_size, max_len, T]\n            return predictions_cat, attn_masks\n        else:\n            return predictions_cat\n    \nclass Attention(nn.Module):\n    '''\n    Attention is calculated using key, value and query from Encoder and decoder.\n    Below are the set of operations you need to perform for computing attention:\n        energy = bmm(key, query)\n        attention = softmax(energy)\n        context = bmm(attention, value)\n    '''\n    def __init__(self):\n        super(Attention, self).__init__()\n    \n    def forward(self, query, key, value): #lens):\n        '''\n        :param query :(batch_size, hidden_size) Query is the output of LSTMCell from Decoder\n        :param keys: (batch_size, max_len, encoder_size) Key Projection from Encoder\n        :param values: (batch_size, max_len, encoder_size) Value Projection from Encoder\n        \n        :return context: (batch_size, encoder_size) Attended Context\n        :return attention_mask: (batch_size, max_len) Attention mask that can be plotted \n        '''\n        # If it is learning well, it will simply apply a very low score to the padded parts\n        \n        # Batch Matrix to Matrix multiplication (b x n x m) @ (b x m x p) = (b x n x p)\n        # Adjust dims of query and key so that they are:\n        # Transform Query to: Tensor[batch_size, 1, key_size]\n        # Transform Key to: Tensor[batch_size, key_size, max_len]\n        energy = torch.bmm(query.unsqueeze(dim=1), key.transpose(1,2)) # Tensor[batch_size, 1, max_len]\n        \n        # Normalize energy to obtain the attention values\n        attention = F.softmax(energy, dim=2) # Tensor[batch_size, 1, max_len]\n        \n        # Obtain context with a weigthed sum of the values over all timesteps\n        context = torch.bmm(attention, value) # Tensor[batch_size, 1, encoder_size]\n        \n        # Remove the extra dimension in the middle\n        context = context.squeeze(dim=1)\n        attention = attention.squeeze(dim=1)\n        \n        return context, attention","metadata":{"_uuid":"68f047be-5454-4af5-bcab-9919560e906d","_cell_guid":"9a745819-5251-4545-81b2-3286027b0e23","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2021-10-24T22:54:26.240225Z","iopub.execute_input":"2021-10-24T22:54:26.240610Z","iopub.status.idle":"2021-10-24T22:54:26.281998Z","shell.execute_reply.started":"2021-10-24T22:54:26.240578Z","shell.execute_reply":"2021-10-24T22:54:26.280801Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Complete model: Encoder + Decoder","metadata":{"_uuid":"82dfd2db-c231-4b44-866e-3911c0a2357b","_cell_guid":"6aecb8fe-c418-42d2-b877-3641052efe60","trusted":true}},{"cell_type":"code","source":"class Seq2Seq(nn.Module):\n    '''\n    We train an end-to-end sequence to sequence model comprising of Encoder and Decoder.\n    This is simply a wrapper \"model\" for your encoder and decoder.\n    '''\n    def __init__(self, input_dim, vocab_size, hidden_dim, value_size=128, key_size=128, isAttended=False, teacher_force=None, with_context=True):\n        super(Seq2Seq, self).__init__()\n        self.with_context = with_context\n        self.is_attended = isAttended\n        self.encoder = Encoder(input_dim, hidden_dim)\n        self.decoder = Decoder(vocab_size, hidden_dim, isAttended=isAttended, teacher_force=teacher_force, with_context=with_context)\n\n    def forward(self, speech_input, speech_len, text_input=None, isTrain=True):\n        if self.with_context:\n            key, value = self.encoder(speech_input, speech_len)\n        \n            # If validating, I pass the text input just to check the length of the text \n            # and stop decoding\n            # Although I could just use a fixed number\n            predictions, attn_masks = self.decoder(key, value, text=text_input, isTrain=isTrain)\n        else:\n            # Do not pass input through encoder\n            # TODO: see what to return\n            predictions = self.decoder(speech_input, speech_len, text=text_input, isTrain=isTrain)\n            attn_masks = None\n               \n        return predictions, attn_masks","metadata":{"_uuid":"6cf87764-bf19-46aa-b3e3-e5ac10d4c253","_cell_guid":"f373d596-a54b-46b9-b1f6-0da9c49ebb66","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2021-10-24T22:44:58.199315Z","iopub.execute_input":"2021-10-24T22:44:58.199651Z","iopub.status.idle":"2021-10-24T22:44:58.210301Z","shell.execute_reply.started":"2021-10-24T22:44:58.199622Z","shell.execute_reply":"2021-10-24T22:44:58.209380Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{"_uuid":"cd0eb6c6-4f70-45f3-bfd7-453995deffb1","_cell_guid":"37543eb6-eff3-4acb-ab90-909ccf0f834f","trusted":true}},{"cell_type":"markdown","source":"### Plotting functions to monitor training process","metadata":{"_uuid":"45c19c0f-f1d5-4ad9-8fa6-350933df2f02","_cell_guid":"95de7cb7-9b9f-4522-beb8-b5d728d9ecb4","trusted":true}},{"cell_type":"code","source":"def plot_grad_flow(named_parameters):\n    '''Plots the gradients flowing through different layers in the net during training.\n    Can be used for checking for possible gradient vanishing / exploding problems.\n    \n    Usage: Plug this function in Trainer class after loss.backwards() as \n    \"plot_grad_flow(self.model.named_parameters())\" to visualize the gradient flow'''\n    \n    plt.figure()\n    ave_grads = []\n    max_grads = []\n    layers = []\n    for n, p in named_parameters:\n        if(p.requires_grad) and (\"bias\" not in n):\n            if(p is not None) and (p.grad is not None):\n                layers.append(n)\n                ave_grads.append(p.grad.abs().mean())\n                max_grads.append(p.grad.abs().max())\n    plt.bar(np.arange(len(max_grads)), max_grads, alpha=0.6, lw=1, color=\"c\")\n    plt.bar(np.arange(len(max_grads)), ave_grads, alpha=0.9, lw=1, color=\"b\")\n    plt.hlines(0, 0, len(ave_grads)+1, lw=2, color=\"k\" )\n    plt.xticks(range(0,len(ave_grads), 1), layers, rotation=\"vertical\")\n    plt.xlim(left=0, right=len(ave_grads))\n    #plt.ylim(bottom = -0.001, top=0.02) # zoom in on the lower gradient regions\n    plt.xlabel(\"Layers\")\n    plt.ylabel(\"Average gradient\")\n    plt.title(\"Gradient flow\")\n    #plt.tight_layout()\n    plt.grid(True)\n    plt.legend([Line2D([0], [0], color=\"c\", lw=4),\n                Line2D([0], [0], color=\"b\", lw=4),\n                Line2D([0], [0], color=\"k\", lw=4)], ['max-gradient', 'mean-gradient', 'zero-gradient'])\n    plt.show()\n    return plt, max_grads\n\ndef plot_attention(attention_mask):\n    plt.figure()\n    plt.imshow(attention_mask, cmap='cool', aspect='auto')\n    plt.title(\"Attention mask\")\n    plt.xlabel(\"Encoder Output Time\")\n    plt.ylabel(\"Decoder Output Label\")\n    plt.colorbar()\n    plt.show()","metadata":{"_uuid":"690c8bf9-d61f-450f-860a-4a683787a38b","_cell_guid":"27ed03ad-66e0-499e-90a1-229794537090","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2021-10-24T22:44:59.030801Z","iopub.execute_input":"2021-10-24T22:44:59.031152Z","iopub.status.idle":"2021-10-24T22:44:59.043345Z","shell.execute_reply.started":"2021-10-24T22:44:59.031121Z","shell.execute_reply":"2021-10-24T22:44:59.042086Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import time\nfrom matplotlib.lines import Line2D\n\ndef train(model, train_loader, criterion, optimizer, prints_per_epoch=2, grad_clip_threshold=2):\n    model.train()\n    model.to(DEVICE)\n    start = time.time()\n    batch_count = 0\n    epoch_losses = []\n    \n    total_samples = len(train_loader.dataset)\n    n_batches = math.floor(total_samples/train_loader.batch_size)\n    batches_print = math.floor(n_batches/prints_per_epoch)\n    batches_print = max(batches_print, 1)\n    \n    # 1) Iterate through your loader\n    for batch_input, batch_labels, input_lens, label_lens in tqdm(train_loader):\n        batch_count += 1\n        optimizer.zero_grad() # Remove active gradients\n        \n        # 2) Use torch.autograd.set_detect_anomaly(True) to get notices about gradient explosion\n        torch.autograd.set_detect_anomaly(True)\n        batch_size = batch_input.shape[0]\n        # 3) Set the inputs to the device.\n        batch_input = batch_input.to(DEVICE)\n        #input_lens = input_lens.to(DEVICE)\n        batch_labels = batch_labels.to(DEVICE)\n        # TODO: should I pass label lens?\n        \n        # 4) Pass your inputs, and length of speech into the model.\n        out, attn_masks = model(batch_input, input_lens, batch_labels)\n        \n        # Calculate Loss\n        # -- Alternative 1:\n        # 5) Generate a mask based on the lengths of the text to create a masked loss. \n        # 5.1) Ensure the mask is on the device and is the correct shape.\n        # 6) If necessary, reshape your predictions and origianl text input \n        # 6.1) Use .contiguous() if you need to. \n        # 7) Use the criterion to get the loss.\n        # 8) Use the mask to calculate a masked loss.\n        \n        # -- Alternative 2:\n        # Here I just use Packed Sequences which does the job\n        # I pass the length of the texts\n        out_pack = pack_padded_sequence(out, label_lens, batch_first=True, enforce_sorted=False)\n        labels_pack = pack_padded_sequence(batch_labels, label_lens, batch_first=True, enforce_sorted=False)\n        loss = criterion(out_pack.data, labels_pack.data)\n        loss_mean = loss / batch_size\n        # 9) Run the backward pass on the masked loss. \n        loss_mean.backward()\n        \n        # 10) Use torch.nn.utils.clip_grad_norm(model.parameters(), 2)\n        torch.nn.utils.clip_grad_norm_(model.parameters(), grad_clip_threshold)\n        \n        # 11) Take a step with your optimizer\n        optimizer.step()\n        \n        # 12) Normalize the loss\n        epoch_losses.append(loss_mean.item())\n        \n        # 13) Optionally print the training loss after every N batches\n        # TODO: parameterize according to total batches\n        if batch_count % batches_print == 0 or batch_count == 1:\n            print(f\"Batch {batch_count} loss (mean):\", loss_mean.item())\n            print(\"Batch input max len is\", batch_input.shape[1])\n            print(\"Reduced by 8\", batch_input.shape[1]/8)\n            print(\"First sample has len\", input_lens[0])\n            print(\"Reduced by 8 would be\", input_lens[0]/8)\n            \n            \n            #print(\"Stacked attn\", attn_masks.shape)\n            print(\"Showing attention mask for first sample in batch\")\n            #fig = plt.figure()\n            \n            #print(\"Max attention value for all text time steps for first sample in batch:\")\n            \n            \n            #print(\"Argmax\", arg_max.shape, arg_max)\n            #print(\"First batch attention\", attn_masks[0].shape)\n            #the_values = np.max(attn_masks[0], axis=0)\n            #print(\"The values\", the_values.shape, the_values)\n            #print(\"All attention values for one text time step for first sample in batch\")\n            #print(\"Out time 0\", attn_masks[0, :, 0])\n            #print(\"Out time 50\", attn_masks[0, :, 50])\n            #print(\"Out time 100\", attn_masks[0, :, 100])\n            #print(\"Out time 150\", attn_masks[0, :, 150])\n            \n            plot_grad_flow(model.named_parameters())\n            \n            if model.with_context and model.is_attended:\n                attn_masks = attn_masks.detach().cpu().numpy()  \n                plot_attention(attn_masks[0])\n            \n    \n    # Epoch average loss in a sequence \n    epoch_mean_loss = np.mean(np.array(epoch_losses))\n    \n    end = time.time()\n    print(\"Training epoch Duration:\", end-start)\n    return epoch_mean_loss\n\n\ndef val(model, val_loader, criterion, prints_per_epoch=2):\n    model.eval()\n    model.to(DEVICE)\n    start = time.time()\n    batch_count = 0\n    epoch_losses = []\n    \n    total_samples = len(val_loader.dataset)\n    n_batches = math.ceil(total_samples/val_loader.batch_size)\n    batches_print = math.floor(n_batches/prints_per_epoch)\n    \n    # 1) Iterate through your loader\n    for batch_input, batch_labels, input_lens, label_lens in tqdm(val_loader):\n        batch_count += 1\n        \n        batch_size = batch_input.shape[0]\n        # 3) Set the inputs to the device.\n        batch_input = batch_input.to(DEVICE)\n        #input_lens = input_lens.to(DEVICE)\n        batch_labels = batch_labels.to(DEVICE)\n        out, attn_masks = model(batch_input, input_lens, batch_labels, isTrain=False)\n        \n        out_pack = pack_padded_sequence(out, label_lens, batch_first=True, enforce_sorted=False)\n        labels_pack = pack_padded_sequence(batch_labels, label_lens, batch_first=True, enforce_sorted=False)\n        loss = criterion(out_pack.data, labels_pack.data)\n        loss_mean = loss / batch_size  \n        \n        # 12) Normalize the loss\n        epoch_losses.append(loss_mean.item())\n        \n        # 13) Optionally print the training loss after every N batches\n        if batch_count % batches_print == 0:\n            print(f\"Batch {batch_count} loss (mean):\", loss_mean.item())\n            \n            if model.with_context and model.is_attended:\n                print(\"Batch input max len is\", batch_input.shape[1])\n                print(\"Reduced by 8\", batch_input.shape[1]/8)\n                print(\"First sample has len\", input_lens[0])\n                print(\"Reduced by 8 would be\", input_lens[0]/8)\n                attn_masks = attn_masks.detach().cpu().numpy()  \n                #print(\"Stacked attn\", attn_masks.shape)\n                #print(\"Showing attention mask for first sample in batch\")\n                #print(\"Max attention value for all text time steps for first sample in batch:\")\n                #arg_max = np.argmax(attn_masks[0], axis=1)\n            \n                #print(\"Argmax\", arg_max.shape, arg_max)\n                #print(\"First batch attention\", attn_masks[0].shape)\n                #the_values = np.max(attn_masks[0], axis=0)\n                #print(\"The values\", the_values.shape, the_values)\n                #print(\"All attention values for one text time step for first sample in batch\")\n                #print(\"Out time 0\", attn_masks[0, 0, :])\n                #print(\"Out time 50\", attn_masks[0, 50, :])\n                #print(\"Out time 100\", attn_masks[0, 100, :])\n                #print(\"Out time 150\", attn_masks[0, 150, :])\n                plot_attention(attn_masks[0])\n            \n    \n    # Epoch average loss in a sequence \n    epoch_mean_loss = np.mean(np.array(epoch_losses))\n    \n    end = time.time()\n    print(\"Val epoch Duration:\", end-start)\n    return epoch_mean_loss","metadata":{"_uuid":"27dd7840-26d7-46b9-83ab-6372b86b6cbf","_cell_guid":"63478d72-6833-497d-a575-f3199ee94354","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2021-10-24T22:45:00.460667Z","iopub.execute_input":"2021-10-24T22:45:00.461011Z","iopub.status.idle":"2021-10-24T22:45:00.483839Z","shell.execute_reply.started":"2021-10-24T22:45:00.460981Z","shell.execute_reply":"2021-10-24T22:45:00.482616Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Define parameters and hyperparameters","metadata":{"_uuid":"3901ece4-620d-4828-93e2-5b8e20e1ea41","_cell_guid":"35b6dafa-209e-421a-a6e7-e517da68fe40","trusted":true}},{"cell_type":"code","source":"import torch.optim as optim\n\nPARAMS = dict(\n    input_dim=40,\n    hidden_dim=128,\n    teacher_force=0.1,\n    nepochs=8, # TODO: aumentar a 25\n    batch_size_train=128,\n    batch_size_val=64, # Use biggest possible batch size allowed by memory\n    optimizer=optim.Adam,\n    optimizer_params=dict(\n      lr=0.01  # Original 0.001\n    ),\n    grad_clip_threshold=2,\n    # Use this to explore data at the beggining\n    train_subset_size=None,\n    val_subset_size=None,\n    with_context=True,\n    is_attended=True\n)","metadata":{"_uuid":"190fdc86-9a76-4965-a741-55b1ad20bde1","_cell_guid":"cf0b6f48-3944-416f-8796-d4223cba9729","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2021-10-24T23:21:18.512171Z","iopub.execute_input":"2021-10-24T23:21:18.512522Z","iopub.status.idle":"2021-10-24T23:21:18.518005Z","shell.execute_reply.started":"2021-10-24T23:21:18.512491Z","shell.execute_reply":"2021-10-24T23:21:18.517025Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pprint\nfrom torch.utils.data import Subset\n\nmodel = Seq2Seq(input_dim=PARAMS['input_dim'], \n                vocab_size=VOCAB_SIZE, \n                hidden_dim=PARAMS['hidden_dim'], \n                isAttended=PARAMS['is_attended'], \n                teacher_force=PARAMS['teacher_force'],\n                with_context=PARAMS['with_context'])\nmodel.load_state_dict(checkpoint)\n\noptimizer = PARAMS['optimizer'](model.parameters(), **PARAMS['optimizer_params'])\ncriterion = nn.CrossEntropyLoss(reduction='sum') # It was reduction none\n    \nprint(\"Loading datasets\")\nprint(\"Training dataset: Originally {} samples\".format(len(train_dataset)))\nprint(\"Val dataset: Originally {} samples\".format(len(val_dataset)))\ntrain_subdataset = Subset(train_dataset, list(range(PARAMS['train_subset_size']))) if PARAMS['train_subset_size'] else train_dataset\nval_subdataset = Subset(val_dataset, list(range(PARAMS['val_subset_size']))) if PARAMS['val_subset_size'] else val_dataset\n\nprint(\"Training dataset: {} samples\".format(len(train_subdataset)))\nprint(\"Val dataset: {} samples\".format(len(val_subdataset)))\n    \nprint(\"Creating loaders\")\ntrain_loader = DataLoader(train_subdataset, \n                          batch_size=PARAMS['batch_size_train'], \n                          shuffle=True, \n                          collate_fn=collate_train)\n\nval_loader = DataLoader(val_subdataset, \n                        batch_size=PARAMS['batch_size_val'], \n                        shuffle=True, \n                        collate_fn=collate_train)\n\nprint(\"Parameters\", pprint.pprint(PARAMS))\nprint(\"Model\", model)\n\ntrain_history = []\nval_history = []\nfor epoch in range(PARAMS['nepochs']):\n    print(\"EPOCH:\", epoch)\n    \n    print(\"-- Training --\")\n    train_epoch_mean_loss = train(model, train_loader, criterion, optimizer, prints_per_epoch=2, grad_clip_threshold=PARAMS['grad_clip_threshold'])\n    train_history.append(train_epoch_mean_loss)\n    print(\"Train epoch mean loss:\", train_epoch_mean_loss)\n    \n    print(\"-- Validation --\")\n    val_epoch_mean_loss = val(model, val_loader, criterion)\n    val_history.append(val_epoch_mean_loss)\n    print(\"Val epoch mean loss:\", val_epoch_mean_loss)\n    print()\n    \n    # Save model\n    torch.save(model.state_dict(), 'checkpoint.pth')\n    torch.save(model.decoder.state_dict(), 'checkpoint_decoder.pth')\n    torch.save(model.encoder.state_dict(), 'checkpoint_encoder.pth')","metadata":{"_uuid":"cd3a450e-c175-4998-aa57-8bbda70bdbf4","_cell_guid":"f1352ca2-ce64-443b-9efa-295f6590d39a","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2021-10-24T23:21:19.770034Z","iopub.execute_input":"2021-10-24T23:21:19.770363Z","iopub.status.idle":"2021-10-24T23:21:27.471835Z","shell.execute_reply.started":"2021-10-24T23:21:19.770331Z","shell.execute_reply":"2021-10-24T23:21:27.469526Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\ndef visualize_results(val_losses, train_losses):\n    \n    plt.style.use('ggplot')\n    f, (ax1, ax2) = plt.subplots(1, 2, figsize=(12,6), dpi=80)\n    \n    ax2.plot(val_losses, label='Validation Loss', marker='o')\n    ax2.plot(train_losses, label='Training Loss', marker='x')\n    ax2.set_ylabel('Loss')\n    ax2.set_xlabel('Epoch')\n    ax2.set_title('Loss vs. Epochs')\n    ax2.legend()\n\n    plt.show()\n\nvisualize_results(val_history, train_history)","metadata":{"_uuid":"58031ba8-caa7-4b81-8742-b27de1e03214","_cell_guid":"4027354d-dfb1-414a-977e-c5bbba494f7f","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2021-10-24T23:08:26.373450Z","iopub.status.idle":"2021-10-24T23:08:26.374093Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Use test set to infere\nfrom tqdm.notebook import tqdm\n\ndef decode_sequence(predictions, vocab):\n    \"\"\"\n    Receives a batch of raw predictions and computes decoding using \n    Random Choice for each timestemp. Then, finds the first <eos>\n    and determines the sequence length for each sample\n    Args:\n        :predictions Tensor[batch_size, max_len, vocab_size]:\n        :vocab List: list of vocab\n    Returns:\n        :batch_decoded Tensor[batch_size, max_len]:\n        :pred_lengths List[batch_size]: List of sequence lengths\n    \"\"\"\n    # TODO: return also the length (cutting at <eos>)\n    \n    timesteps = predictions.shape[1]\n    batch_size = predictions.shape[0]\n    predictions = F.softmax(predictions, dim=2)\n    predictions = predictions.detach().cpu().numpy()\n    vocab_len = len(vocab)\n    \n    start = time.time()\n    batch_predicted_sequences = []\n    \n    for b in range(batch_size):\n        time_predicted_chars = []\n        for t in range(timesteps):\n            time_predicted_char = np.random.choice(vocab_len, p=predictions[b, t, :])\n            time_predicted_chars.append(time_predicted_char)\n        time_predicted_chars = torch.LongTensor(time_predicted_chars)\n        batch_predicted_sequences.append(time_predicted_chars)\n    batch_decoded = torch.stack(batch_predicted_sequences, dim=0)\n    \n    # Find the length\n    end_index = letter2index['<eos>']\n    pred_lengths = []\n    for seq in batch_decoded:\n        pred_length = len(seq)\n        eos_occurences = (seq == end_index).nonzero(as_tuple=True)[0] \n        if len(eos_occurences):\n            # Take first occurence index\n            first_eos_idx = eos_occurences[0].item()\n            pred_length = first_eos_idx + 1\n        pred_lengths.append(pred_length)   \n    end = time.time()\n    \n    return batch_decoded, pred_lengths\n\n\n\ndef random_search_inference(model, batch_input, input_lens, criterion, num_samples=10):\n    \"\"\"\n    Takes a batch of input utterances and outputs the decoding for each utterance.\n    \n    :return batch_decodings Tensor[batch_size, decoding_max_len]: \n    :return batch_best_decodings_len Tensor[batch_size]:\n    :return batch_best_losses List[batch_size]:\n    \"\"\"\n    # TODO: make number of samples a parameter\n    batch_size = batch_input.shape[0]\n    batch_losses_list = []\n    batch_decoded_iter_list = []\n    batch_decoded_lens_iter_list = []\n    \n    # TODO: maybe I can just always keep only the best, thus saving memory (although it would be more computing during the loops)\n    for i in tqdm(range(num_samples)):\n        \n        batch_out, _ = model(batch_input, input_lens, isTrain=False)\n        \n        batch_decoded, decoded_lengths = decode_sequence(batch_out, LETTER_LIST) \n        second_out, _ = model(batch_input, input_lens, isTrain=False)\n        \n        # NOTE: here I compute the losses only until the lengths of the \"new ground truth\"\n        # So can pack the data with decoded_lengths\n        batch_decoded = batch_decoded.to(DEVICE)\n        batch_decoded_iter_list.append(batch_decoded)\n        batch_decoded_lens_iter_list.append(decoded_lengths)\n        out_pack = pack_padded_sequence(second_out, decoded_lengths, batch_first=True, enforce_sorted=False)\n        labels_pack = pack_padded_sequence(batch_decoded, decoded_lengths, batch_first=True, enforce_sorted=False)\n        \n        # TODO: check if <sos> should be removed from ground truth?\n        # Here we are using NO reduction\n        loss = criterion(out_pack.data, labels_pack.data)\n        \n        # We want a Loss value per sample inside the batch\n        # First we get the losses \n        batch_losses_per_sample = list(torch.split(loss, decoded_lengths))\n        # We sum up the losses in each sequence (all timesteps of a sequence)\n        batch_losses_reduced = [torch.sum(seq_loss).item() for seq_loss in batch_losses_per_sample]\n        batch_losses_reduced = torch.Tensor(batch_losses_reduced)\n        batch_losses_list.append(batch_losses_reduced)\n        \n    # For each sample in the batch, pick the one with the lowest loss\n    batch_losses_list = torch.stack(batch_losses_list, dim=1)\n    min_loss_index = torch.argmin(batch_losses_list, dim=1)\n\n    \n    # Get best decoding for each sample in batch\n    batch_best_losses = []\n    batch_best_decodings = []\n    batch_best_decodings_len = []\n    \n    for sample in range(batch_size):\n        sample_min_index = min_loss_index[sample]\n        \n        best_loss = batch_losses_list[sample][sample_min_index]\n        decoded_seq = batch_decoded_iter_list[sample_min_index][sample]\n        decoded_len = batch_decoded_lens_iter_list[sample_min_index][sample]\n        \n        batch_best_losses.append(best_loss)\n        batch_best_decodings.append(decoded_seq)\n        batch_best_decodings_len.append(decoded_len)\n        \n    batch_decodings = torch.stack(batch_best_decodings, dim=0)\n    batch_best_losses = torch.Tensor(batch_best_losses)\n    \n    return batch_decodings, batch_best_decodings_len, batch_best_losses\n\ndef test(model, test_loader, criterion):\n    model.eval()\n    model.to(DEVICE)\n    batch_count = 0\n    \n    epoch_losses = []\n    \n    # TODO: make batch in validation bigger (as big as memory allows it)\n    batch_decodings = []\n    batch_decoding_lens = []\n    for batch_input, input_lens in test_loader:\n        batch_count += 1\n        print(\"batch count\", batch_count)\n        \n        batch_size = batch_input.shape[0]\n        batch_input = batch_input.to(DEVICE)\n        \n        start = time.time()\n        # TODO: make num sumples a parameter\n        batch_decoding, batch_lens, batch_loss = random_search_inference(model, batch_input, input_lens, criterion, num_samples=2)\n        end = time.time()\n        duration_decode = end-start\n        \n        batch_decodings.append(batch_decoding)\n        batch_decoding_lens = batch_decoding_lens + batch_lens\n        batch_avg_loss = batch_loss.mean().item()\n        epoch_losses.append(batch_avg_loss)\n        \n    \n    epoch_avg_loss = torch.Tensor(epoch_losses).mean().item()\n    batch_decodings = torch.cat(batch_decodings, dim=0)\n    print(\"Test loss\", epoch_avg_loss)\n    print(\"Batch decodings\", batch_decodings.shape)\n    \n    # TODO return decodings\n    return epoch_avg_loss, batch_decodings, batch_decoding_lens\n\nbatch_size = 128\ntest_dataset = Speech2TextDataset(speech_test, None, isTrain=False)\n# Should not shuffle so that decodings preserve the order\ntest_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False, collate_fn=collate_test)\n# TODO: call it test_criterion\n# epoch_avg_loss, test_decodings, test_decoding_lens = test(model, test_loader, val_criterion)","metadata":{"_uuid":"d4e6678d-bf4e-454f-96fa-28aa1d1f6f46","_cell_guid":"29355387-34a2-453a-936f-e6cbe834c35c","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2021-10-17T18:40:45.89767Z","iopub.execute_input":"2021-10-17T18:40:45.897998Z","iopub.status.idle":"2021-10-17T18:40:45.919426Z","shell.execute_reply.started":"2021-10-17T18:40:45.897968Z","shell.execute_reply":"2021-10-17T18:40:45.918599Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def transform_index_to_letter(decoding, decoding_len) -> List:\n    '''\n    Transform index_to_letter. Removes <sos> and <eos>?\n    Receives transcript in batch.\n    \n    :param decoding :(N, decoding_max_len) Decodings that are in index\n    :param decoding :(N ) Decodings lengths\n    :return index_to_letter_list: Returns a list for all the transcript sentence to index\n    '''\n    \n    unk_idx = letter2index['<unk>']\n    start_idx = letter2index['<sos>']\n    end_idx = letter2index['<eos>']\n    \n    decoding_transcript = []\n    \n    for sequence, sequence_len in zip(decoding, decoding_len):\n        sequence_letter = []\n        for index in sequence:\n            index = index.item()\n            if index in [start_idx, end_idx]:\n                continue\n            char = index2letter[index]\n            sequence_letter.append(char)\n        sequence_letter = ''.join(sequence_letter)\n        decoding_transcript.append(sequence_letter[:sequence_len])\n    \n    return decoding_transcript\n\n# hi = transform_index_to_letter(test_decodings, test_decoding_lens)","metadata":{"_uuid":"47b2d522-41ea-4e1a-92c4-5efefa5f0b54","_cell_guid":"ef7cf681-0fb1-4e44-9c3a-34c9b6bbceed","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2021-10-17T18:40:49.87929Z","iopub.execute_input":"2021-10-17T18:40:49.879639Z","iopub.status.idle":"2021-10-17T18:40:49.886145Z","shell.execute_reply.started":"2021-10-17T18:40:49.879609Z","shell.execute_reply":"2021-10-17T18:40:49.885102Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# df = pd.DataFrame(hi, columns=['label'])\n# df.index.rename('id', inplace=True)\n# df.to_csv('test_labels.csv')\n\n# df","metadata":{"_uuid":"87046f50-e128-4213-9d4c-ea04f3acfb78","_cell_guid":"e75347c8-0dc5-4416-9fae-0da1108168ff","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2021-10-17T18:40:51.214478Z","iopub.execute_input":"2021-10-17T18:40:51.214792Z","iopub.status.idle":"2021-10-17T18:40:51.218776Z","shell.execute_reply.started":"2021-10-17T18:40:51.214762Z","shell.execute_reply":"2021-10-17T18:40:51.217579Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def print_distances(distances, token1Length, token2Length):\n    for t1 in range(token1Length + 1):\n        for t2 in range(token2Length + 1):\n            print(int(distances[t1][t2]), end=\" \")\n        print()\n    \ndef levenshtein_distance(token1, token2, debug=False):\n    \"\"\"\n      Computes the Edit Distance between two sequences.\n      \"\"\"\n    # TODO: check if passing to GPU is better\n    distances = np.zeros((len(token1) + 1, len(token2) + 1))\n    for t1 in range(len(token1) + 1):\n        distances[t1][0] = t1\n\n    for t2 in range(len(token2) + 1):\n        distances[0][t2] = t2\n        \n    a = 0\n    b = 0\n    c = 0\n    \n    for t1 in range(1, len(token1) + 1):\n        for t2 in range(1, len(token2) + 1):\n            if (token1[t1-1] == token2[t2-1]):\n                distances[t1][t2] = distances[t1 - 1][t2 - 1]\n            else:\n                a = distances[t1][t2 - 1]\n                b = distances[t1 - 1][t2]\n                c = distances[t1 - 1][t2 - 1]             \n                min_dist = min((a,b,c))\n\n                distances[t1][t2] = min_dist + 1\n\n    if debug:\n        print_distances(distances, len(token1), len(token2))\n    \n    return distances[len(token1)][len(token2)]\n\ndef test_2(model, val_loader, epoch, criterion):\n    # TODO: DO NOT TOUCH\n    model.eval()\n    model.to(DEVICE)\n    batch_count = 0\n    \n    dataset_distances = []\n    epoch_losses = []\n    \n    # TODO: make batch in validation bigger (as big as memory allows it)\n    for batch_input, batch_labels, input_lens, label_lens in val_loader:\n        batch_count += 1\n        \n        batch_size = batch_input.shape[0]\n        batch_input = batch_input.to(DEVICE)\n        batch_labels = batch_labels.to(DEVICE)\n        \n        start = time.time()\n        batch_decoding, batch_lens, batch_loss = random_search_inference(model, batch_input, input_lens, criterion, num_samples=10)\n        end = time.time()\n        duration_decode = end-start\n        \n        batch_avg_loss = batch_loss.mean().item()\n        epoch_losses.append(batch_avg_loss)\n        \n        # Compute edit distance for all the samples in the batch\n        batch_distances = []\n        for i in range(batch_size):\n            batch_len = batch_lens[i]\n            pred_seq = batch_decoding[i][:batch_len]\n            \n            label_len = int(label_lens[i].item())\n            label_seq = batch_labels[i][:label_len]\n            \n            dist = levenshtein_distance(pred_seq, label_seq)\n            dataset_distances.append(dist)          \n            if i % 5 == 0:\n                print(\"Validation distance batch {}: {}\".format(batch_count, dist))\n    \n    epoch_avg_loss = torch.Tensor(epoch_losses).mean().item()\n    dataset_distances = np.array(dataset_distances)\n    dataset_mean_distance = np.mean(dataset_distances)\n    print(\"Validation loss\", epoch_avg_loss)\n    print(\"Dataset mean distance\", dataset_mean_distance)\n    \n    return epoch_avg_loss, dataset_mean_distance","metadata":{"_uuid":"76671f97-bad0-40ae-bb30-06cc5c1db3cf","_cell_guid":"7da03586-a03e-4f6b-9f6b-4a4e01669e77","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2021-10-17T18:40:54.707335Z","iopub.execute_input":"2021-10-17T18:40:54.707714Z","iopub.status.idle":"2021-10-17T18:40:54.725082Z","shell.execute_reply.started":"2021-10-17T18:40:54.707667Z","shell.execute_reply":"2021-10-17T18:40:54.724181Z"},"trusted":true},"execution_count":null,"outputs":[]}]}