{"metadata":{"accelerator":"GPU","colab":{"provenance":[]},"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":62741,"databundleVersionId":6820148,"sourceType":"competition"}],"dockerImageVersionId":30588,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# HW4P2: Attention-based Speech Recognition\n\n<img src=\"https://cdn.shopify.com/s/files/1/0272/2080/3722/products/SmileBumperSticker_5400x.jpg\" alt=\"A cute cat\" width=\"600\">\n\n\nWelcome to the final assignment in 11785. In this HW, you will work on building a speech recognition system with <i>attention</i>. <br> <br>\n\n<center>\n<img src=\"https://popmn.org/wp-content/uploads/2020/03/pay-attention.jpg\" alt=\"A cute cat\" height=\"100\">\n</center>\n\nHW Writeup: [TODO] <br>\nKaggle Competition Link: https://www.kaggle.com/competitions/attention-based-speech-recognition <br>\nKaggle Dataset Link: https://www.kaggle.com/competitions/attention-based-speech-recognition/data\n<br>\nLAS Paper: https://arxiv.org/pdf/1508.01211.pdf <br>\nAttention is all you need:https://arxiv.org/pdf/1706.03762.pdf","metadata":{"id":"8XpNMS7Vk6Df"}},{"cell_type":"markdown","source":"# Read this section importantly!","metadata":{"id":"vwIdDTTmmZVe"}},{"cell_type":"markdown","source":"1. By now, we believe that you are already a great deep learning practitioner, Congratulations. 🎉\n\n2. You are allowed to use code from your previous homeworks for this homework. We will only provide, aspects that are necessary and new with this homework.\n\n3. There are a lot of resources provided in this notebook, that will help you check if you are running your implementations correctly.","metadata":{"id":"y9qsVrRemgh7"}},{"cell_type":"code","source":"!nvidia-smi","metadata":{"id":"8UK7J-dp5iN5","outputId":"66802957-853d-4169-8f19-d142f9a9e68f","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Install some required libraries\n# Feel free to add more if you want\n!pip install -q python-levenshtein torchsummaryX wandb kaggle pytorch-nlp","metadata":{"id":"nYgaLmgy5iqR","outputId":"b2a294d1-ecb1-4eac-96a6-e8681bbeecae","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torchaudio\nfrom torch import nn, Tensor\n# import torchsummary\n\nimport numpy as np\nimport os\n\nimport gc\nimport time\n\nimport pandas as pd\nfrom tqdm.notebook import tqdm as blue_tqdm\nimport matplotlib.pyplot as plt\nimport seaborn\nimport json\nfrom torch.nn.utils.rnn import pad_sequence, pack_padded_sequence, pad_packed_sequence\n\nimport math\nfrom typing import Optional, List\n\n\n#imports for decoding and distance calculation\ntry:\n    import wandb\n    import torchsummaryX\n    import Levenshtein\nexcept:\n    print(\"Didnt install some/all imports\")\n\nimport warnings\nwarnings.filterwarnings('ignore')\n\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\nprint(\"Device: \", DEVICE)","metadata":{"id":"0mkii-6Dsjr8","outputId":"2ff62419-c217-47a0-f5ab-8264b0cec846","execution":{"iopub.status.busy":"2023-12-01T21:12:04.097505Z","iopub.execute_input":"2023-12-01T21:12:04.098459Z","iopub.status.idle":"2023-12-01T21:12:10.110951Z","shell.execute_reply.started":"2023-12-01T21:12:04.098417Z","shell.execute_reply":"2023-12-01T21:12:10.109675Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{"id":"AIOBPQjzrx5n"}},{"cell_type":"code","source":"","metadata":{"id":"EPchlig7rxia","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Kaggle Dataset Download","metadata":{"id":"-njBvl2Opd6I"}},{"cell_type":"code","source":"# !pip install --upgrade --force-reinstall --no-deps kaggle==1.5.8\n# !mkdir /root/.kaggle\n\n# with open(\"/root/.kaggle/kaggle.json\", \"w+\") as f:\n#     f.write('{\"username\":\"hadizaumaryusuf\",\"key\":\"e1a8fb5be250ccf20574b175b274069a\"}')\n#     # Put your kaggle username & key here\n\n# !chmod 600 /root/.kaggle/kaggle.json","metadata":{"id":"PTyWR2sIp0Ns","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # # to download the dataset\n# !kaggle competitions download -c attention-based-speech-recognition\n\n# # # to unzip data quickly and quietly\n# !unzip -q attention-based-speech-recognition.zip -d ./data","metadata":{"id":"F581gjfnqE2C","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !unzip -q attention-based-speech-recognition.zip -d ./data ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"config = dict (\n\n    train_dataset       = 'train-clean-460', # train-clean-100, train-clean-360, train-clean-460\n    cepstral_norm       = True,\n    transforms          = dict(\n        TimeMasking     = 200,\n        FreqMasking     = 5,\n        GaussianNoise   = [0, 0.05],\n    ),\n    batch_size          = 96,\n\n    listener            = dict(\n        base_lstm_layers    = 1,\n        base_lstm_dp        = 0,\n        locked_dp           = 0,\n\n        n_pblstms           = 0\n    ),\n\n    attention           = dict(\n        attention_type  = 'dot-product'\n    ),\n\n    speller             = dict(\n        embedding_dp    = 0.0,\n        lstm_dp         = 0.0,\n    ),\n\n    epochs              = 100,\n\n    learning_rate       = 1e-4,\n    optimizer           = 'AdamW',\n    weight_decay        = 5e-3,\n\n    label_smoothing     = 0.01,\n    scheduler           = 'ReduceLR',\n    tf_scheduler        = 'Cosine'\n\n)","metadata":{"execution":{"iopub.status.busy":"2023-12-01T21:12:10.113145Z","iopub.execute_input":"2023-12-01T21:12:10.113975Z","iopub.status.idle":"2023-12-01T21:12:10.122497Z","shell.execute_reply.started":"2023-12-01T21:12:10.113935Z","shell.execute_reply":"2023-12-01T21:12:10.121392Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!wget -q https://cmu.box.com/shared/static/om4qpzd4tf1xo4h7230k4v1pbdyueghe --content-disposition --show-progress\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!unzip -q /kaggle/working/hw4p2_toy.zip -d ./data_toy","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Character-based LibriSpeech (HW4P2)\n\nIn terms of the dataset, the dataset structure for HW3P2 and HW4P2 dataset are very similar. Can you spot out the differences? What all will be required??\n\nHints:\n\n- Check how big is the dataset (do you require memory efficient loading techniques??)\n- How do we load mfccs? Do we need to normalise them?\n- Does the data have \\<SOS> and \\<EOS> tokens in each sequences? Do we remove them or do we not remove them? (Read writeup)\n- Would we want a collating function? Ask yourself: Why did we need a collate function last time?\n- Observe the VOCAB, is the dataset same as HW3P2?\n- Should you add augmentations, if yes which augmentations? When should you add augmentations? (Check bootcamp for answer)\n","metadata":{"id":"zUJyBBwIqQs6"}},{"cell_type":"code","source":"VOCAB = [\n    '<pad>', '<sos>', '<eos>',\n    'A',   'B',    'C',    'D',\n    'E',   'F',    'G',    'H',\n    'I',   'J',    'K',    'L',\n    'M',   'N',    'O',    'P',\n    'Q',   'R',    'S',    'T',\n    'U',   'V',    'W',    'X',\n    'Y',   'Z',    \"'\",    ' ',\n]\n\nVOCAB_MAP = {VOCAB[i]:i for i in range(0, len(VOCAB))}\n\nPAD_TOKEN = VOCAB_MAP[\"<pad>\"]\nSOS_TOKEN = VOCAB_MAP[\"<sos>\"]\nEOS_TOKEN = VOCAB_MAP[\"<eos>\"]\n\nprint(f\"Length of vocab : {len(VOCAB)}\")\nprint(f\"Vocab           : {VOCAB}\")\nprint(f\"PAD_TOKEN       : {PAD_TOKEN}\")\nprint(f\"SOS_TOKEN       : {SOS_TOKEN}\")\nprint(f\"EOS_TOKEN       : {EOS_TOKEN}\")","metadata":{"id":"MBMLGYX-kZcd","execution":{"iopub.status.busy":"2023-12-01T21:12:10.123809Z","iopub.execute_input":"2023-12-01T21:12:10.124786Z","iopub.status.idle":"2023-12-01T21:12:10.148053Z","shell.execute_reply.started":"2023-12-01T21:12:10.124751Z","shell.execute_reply":"2023-12-01T21:12:10.146897Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SpeechDataset(torch.utils.data.Dataset):\n\n    def __init__(self, root, partition= \"train-clean-100\", transforms = None, cepstral=True):\n\n        self.VOCAB  = VOCAB\n\n        if partition != \"train-clean-460\":\n            mfcc_dir       = root + \"/\" + partition + \"/mfcc/\"\n            transcript_dir = root + \"/\" + partition + \"/transcripts/\"\n\n            mfcc_files          = [mfcc_dir + mfcc              for mfcc in os.listdir(mfcc_dir)]\n            transcript_files    = [transcript_dir + transcript for transcript in os.listdir(transcript_dir)]\n\n        else:\n            mfcc_dir       = root + \"/train-clean-100/mfcc/\"\n            transcript_dir = root + \"/train-clean-100/transcripts/\"\n\n            mfcc_files          = [mfcc_dir + mfcc              for mfcc in os.listdir(mfcc_dir)]\n            transcript_files    = [transcript_dir + transcript for transcript in os.listdir(transcript_dir)]\n\n            mfcc_dir       = root + \"/train-clean-360/mfcc/\"\n            transcript_dir = root + \"/train-clean-360/transcripts/\"\n\n            mfcc_files.extend([mfcc_dir + mfcc                      for mfcc in os.listdir(mfcc_dir)])\n            transcript_files.extend([transcript_dir + transcript    for transcript in os.listdir(transcript_dir)])\n\n        assert len(mfcc_files) == len(transcript_files)\n        length = len(mfcc_files)\n\n        self.mfccs, self.transcripts = [], []\n        for i in blue_tqdm(range(length)):\n            mfcc        = np.load(mfcc_files[i])\n            transcript  = np.load(transcript_files[i])\n\n            mfcc                = (mfcc - mfcc.mean(axis=0))/mfcc.std(axis=0) if cepstral else mfcc\n            transcript_mapped   = np.array([self.VOCAB.index(i) for i in transcript])\n\n            self.mfccs.append(mfcc)\n            self.transcripts.append(transcript_mapped)\n\n        self.length = len(transcript_files)\n        print(\"Loaded: \", partition)\n\n\n    def __len__(self):\n        return self.length\n\n    def __getitem__(self, ind):\n        return torch.FloatTensor(self.mfccs[ind]), torch.LongTensor(self.transcripts[ind])\n\n    def collate_fn(self,batch):\n\n        batch_x, batch_y, lengths_x, lengths_y = [], [], [], []\n        for x, y in batch:\n            batch_x.append(x)\n            batch_y.append(y)\n            lengths_x.append(len(x))\n            lengths_y.append(len(y))\n\n        batch_x_pad = torch.nn.utils.rnn.pad_sequence(batch_x, batch_first=True)\n        batch_y_pad = torch.nn.utils.rnn.pad_sequence(batch_y, batch_first=True)\n\n        return batch_x_pad, batch_y_pad, torch.tensor(lengths_x), torch.tensor(lengths_y)\n\n\nclass SpeechDatasetME(torch.utils.data.Dataset): # Memory efficient\n    # Loades the data in get item to save RAM\n\n    def __init__(self, root, partition= \"train-clean-100\", transforms = None, cepstral=True):\n\n        self.VOCAB      = VOCAB\n        self.cepstral   = cepstral\n\n        if partition == \"train-clean-100\" or partition == \"train-clean-360\":\n            mfcc_dir       = root + \"/\" + partition + \"/mfcc/\"\n            transcript_dir = root + \"/\" + partition + \"/transcripts/\"\n\n            mfcc_files          = [mfcc_dir + mfcc              for mfcc in os.listdir(mfcc_dir)]\n            transcript_files    = [transcript_dir + transcript for transcript in os.listdir(transcript_dir)]\n\n        else:\n            mfcc_dir       = root + \"/train-clean-100/mfcc/\"\n            transcript_dir = root + \"/train-clean-100/transcripts/\"\n\n            mfcc_files          = [mfcc_dir + mfcc              for mfcc in os.listdir(mfcc_dir)]\n            transcript_files    = [transcript_dir + transcript for transcript in os.listdir(transcript_dir)]\n\n            mfcc_dir       = root + \"/train-clean-360/mfcc/\"\n            transcript_dir = root + \"/train-clean-360/transcripts/\"\n\n            mfcc_files.extend([mfcc_dir + mfcc                      for mfcc in os.listdir(mfcc_dir)])\n            transcript_files.extend([transcript_dir + transcript    for transcript in os.listdir(transcript_dir)])\n\n        assert len(mfcc_files) == len(transcript_files)\n        length = len(mfcc_files)\n\n        self.mfcc_files         = mfcc_files\n        self.transcript_files   = transcript_files\n        self.length             = len(transcript_files)\n        print(\"Loaded file paths ME: \", partition)\n\n\n    def __len__(self):\n        return self.length\n\n    def __getitem__(self, ind):\n\n        mfcc        = np.load(self.mfcc_files[ind])\n        transcript  = np.load(self.transcript_files[ind])\n\n        mfcc                = (mfcc - mfcc.mean(axis=0))/mfcc.std(axis=0) if self.cepstral else mfcc\n        transcript_mapped   = np.array([self.VOCAB.index(i) for i in transcript])\n\n        return torch.FloatTensor(mfcc), torch.LongTensor(transcript_mapped)\n\n    def collate_fn(self,batch):\n\n        batch_x, batch_y, lengths_x, lengths_y = [], [], [], []\n\n        for x, y in batch:\n            batch_x.append(x)\n            batch_y.append(y)\n            lengths_x.append(len(x))\n            lengths_y.append(len(y))\n\n        batch_x_pad = torch.nn.utils.rnn.pad_sequence(batch_x, batch_first=True)\n        batch_y_pad = torch.nn.utils.rnn.pad_sequence(batch_y, batch_first=True)\n\n        return batch_x_pad, batch_y_pad, torch.tensor(lengths_x), torch.tensor(lengths_y)\n","metadata":{"id":"VuneWaTStdF2","execution":{"iopub.status.busy":"2023-12-01T21:12:10.150742Z","iopub.execute_input":"2023-12-01T21:12:10.15118Z","iopub.status.idle":"2023-12-01T21:12:10.182472Z","shell.execute_reply.started":"2023-12-01T21:12:10.151143Z","shell.execute_reply":"2023-12-01T21:12:10.181401Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SpeechDatasetTest(torch.utils.data.Dataset):\n\n    def __init__(self, root, partition, cepstral=False):\n\n        self.mfcc_dir   = root + \"/\" + partition + \"/mfcc/\"\n        self.mfcc_files = sorted(os.listdir(self.mfcc_dir)) # list files in the mfcc directory\n\n        self.mfccs = []\n        for i, filename in enumerate(blue_tqdm(self.mfcc_files)):\n            mfcc = np.load(self.mfcc_dir + filename)\n            if cepstral:\n                mfcc = (mfcc - mfcc.mean(axis=0))/mfcc.std(axis=0)\n            self.mfccs.append(torch.FloatTensor(mfcc))\n\n        print(\"Loaded: \", partition)\n\n    def __len__(self):\n        return len(self.mfccs)\n\n    def __getitem__(self, ind):\n        return self.mfccs[ind]\n\n    def collate_fn(self,batch):\n\n        batch_x, lengths_x = [], []\n        for x in batch:\n            batch_x.append(x)\n            lengths_x.append(len(x))\n        batch_x_pad = torch.nn.utils.rnn.pad_sequence(batch_x, batch_first=True)\n\n        return batch_x_pad, torch.tensor(lengths_x)","metadata":{"id":"jUrqTkG4VZfJ","execution":{"iopub.status.busy":"2023-12-01T21:12:10.183575Z","iopub.execute_input":"2023-12-01T21:12:10.184857Z","iopub.status.idle":"2023-12-01T21:12:10.198633Z","shell.execute_reply.started":"2023-12-01T21:12:10.184823Z","shell.execute_reply":"2023-12-01T21:12:10.197597Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# DATA_DIR        = '/kaggle/input/attention-based-speech-recognition/11-785-f23-hw4p2'\n# PARTITION       = config['train_dataset']\n# CEPSTRAL        = config['cepstral_norm']\n\n# train_dataset   = SpeechDataset( # Or AudioDatasetME\n#     root        = DATA_DIR,\n#     partition   = PARTITION,\n#     cepstral    = CEPSTRAL\n# )\n# valid_dataset   = SpeechDataset(\n#     root        = DATA_DIR,\n#     partition   = 'dev-clean',\n#     cepstral    = CEPSTRAL\n# )\n# test_dataset    = SpeechDatasetTest(\n#     root        = DATA_DIR,\n#     partition   = 'test-clean',\n#     cepstral    = CEPSTRAL,\n# )\n\n# gc.collect()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# DATA_DIR        = 'data/11-785-f23-hw4p2'\nDATA_DIR        = '/kaggle/input/attention-based-speech-recognition/11-785-f23-hw4p2'\n\n# PARTITION       = config['train_dataset']\nPARTITION       = \"train-clean-100\"\n\nCEPSTRAL        = config['cepstral_norm']\n\ntrain_dataset   = SpeechDatasetME( # Or AudioDatasetME\n    root        = DATA_DIR,\n    partition   = PARTITION,\n    cepstral    = CEPSTRAL\n)\nvalid_dataset   = SpeechDataset(\n    root        = DATA_DIR,\n    partition   = 'dev-clean',\n    cepstral    = CEPSTRAL\n)\ntest_dataset    = SpeechDatasetTest(\n    root        = DATA_DIR,\n    partition   = 'test-clean',\n    cepstral    = CEPSTRAL,\n)\n\ngc.collect()","metadata":{"id":"rsl5Q1jLvsOL","execution":{"iopub.status.busy":"2023-12-01T21:12:10.200167Z","iopub.execute_input":"2023-12-01T21:12:10.200532Z","iopub.status.idle":"2023-12-01T21:13:04.154121Z","shell.execute_reply.started":"2023-12-01T21:12:10.200499Z","shell.execute_reply":"2023-12-01T21:13:04.153317Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loader    = torch.utils.data.DataLoader(\n    dataset     = train_dataset,\n    batch_size  = config['batch_size'],\n    shuffle     = True,\n    num_workers = 4,\n    pin_memory  = True,\n    collate_fn  = train_dataset.collate_fn\n)\n\nvalid_loader    = torch.utils.data.DataLoader(\n    dataset     = valid_dataset,\n    batch_size  = config['batch_size'],\n    shuffle     = False,\n    num_workers = 2,\n    pin_memory  = True,\n    collate_fn  = valid_dataset.collate_fn\n)\n\ntest_loader     = torch.utils.data.DataLoader(\n    dataset     = test_dataset,\n    batch_size  = config['batch_size'],\n    shuffle     = False,\n    num_workers = 2,\n    pin_memory  = True,\n    collate_fn  = test_dataset.collate_fn\n)\n\nprint(\"No. of train mfccs   : \", train_dataset.__len__())\nprint(\"Batch size           : \", config['batch_size'])\nprint(\"Train batches        : \", train_loader.__len__())\nprint(\"Valid batches        : \", valid_loader.__len__())\nprint(\"Test batches         : \", test_loader.__len__())","metadata":{"id":"OeqXHogpwFfa","execution":{"iopub.status.busy":"2023-12-01T21:13:04.155137Z","iopub.execute_input":"2023-12-01T21:13:04.155392Z","iopub.status.idle":"2023-12-01T21:13:04.164054Z","shell.execute_reply.started":"2023-12-01T21:13:04.155368Z","shell.execute_reply":"2023-12-01T21:13:04.163114Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"\\nChecking the shapes of the data...\")\nfor batch in train_loader:\n    x, y, x_len, y_len = batch\n    print(x.shape, y.shape, x_len.shape, y_len.shape)\n    print(y)\n    break","metadata":{"id":"tzuIXCyAuNvo","outputId":"633dcb9a-d1ef-46f2-a896-35e31f97543a","execution":{"iopub.status.busy":"2023-12-01T21:13:04.165635Z","iopub.execute_input":"2023-12-01T21:13:04.166032Z","iopub.status.idle":"2023-12-01T21:13:10.573591Z","shell.execute_reply.started":"2023-12-01T21:13:04.166001Z","shell.execute_reply":"2023-12-01T21:13:10.572479Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def verify_dataset(dataset, partition= 'train-clean-100'):\n#     print(\"\\nPartition loaded     : \", partition)\n#     if partition != 'test-clean':\n#         print(\"Max mfcc length          : \", np.max([data[0].shape[0] for data in dataset]))\n#         print(\"Avg mfcc length          : \", np.mean([data[0].shape[0] for data in dataset]))\n#         print(\"Max transcript length    : \", np.max([data[1].shape[0] for data in dataset]))\n#         print(\"Max transcript length    : \", np.mean([data[1].shape[0] for data in dataset]))\n#     else:\n#         print(\"Max mfcc length          : \", np.max([data.shape[0] for data in dataset]))\n#         print(\"Avg mfcc length          : \", np.mean([data.shape[0] for data in dataset]))\n\n# verify_dataset(train_dataset, partition= 'train-clean-100')\n# verify_dataset(valid_dataset, partition= 'dev-clean')\n# verify_dataset(test_dataset, partition= 'test-clean')\n# dataset_max_len  = max(\n#     np.max([data[0].shape[0] for data in train_dataset]),\n#     np.max([data[0].shape[0] for data in valid_dataset]),\n#     np.max([data.shape[0] for data in test_dataset])\n# )\n# print(\"\\nMax Length: \", dataset_max_len)","metadata":{"id":"8QFdYrM7xcI1","scrolled":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Check if you are loading the data correctly with the following:\n\n- Train Dataset\n```\nPartition loaded:  train-clean-100\nMax mfcc length:  2448\nAverage mfcc length:  1264.6258453344547\nMax transcript:  400\nAverage transcript length:  186.65321139493324\n```\n\n- Dev Dataset\n```\nPartition loaded:  dev-clean\nMax mfcc length:  3260\nAverage mfcc length:  713.3570107288198\nMax transcript:  518\nAverage transcript length:  108.71698113207547\n```\n\n- Test Dataset\n```\nPartition loaded:  test-clean\nMax mfcc length:  3491\nAverage mfcc length:  738.2206106870229\n```\n\nIf your values is not matching, read hints, think what could have gone wrong. Then approach TAs.","metadata":{"id":"i_n3pqt7ud4t"}},{"cell_type":"markdown","source":"# THE MODEL\n\n### Listen, Attend and Spell\nListen, Attend and Spell (LAS) is a neural network model used for speech recognition and synthesis tasks.\n\n- LAS is designed to handle long input sequences and is robust to noisy speech signals.\n- LAS is known for its high accuracy and ability to improve over time with additional training data.\n- It consists of an <b>listener, an attender and a speller</b>, which work together to convert an input speech signal into a corresponding output text.\n\n#### The Dataflow:\n<center>\n<img src=\"https://github.com/varunjain3/11785_s23_h4p2/raw/main/DataFlow.png\" alt=\"data flow\" height=\"100\">\n</center>\n\n#### The Listener:\n- converts the input speech signal into a sequence of hidden states.\n\n#### The Attender:\n- Decides how the sequence of Encoder hidden state is propogated to decoder.\n\n#### The Speller:\n- A language model, that incorporates the \"context of attender\"(output of attender) to predict sequence of words.\n\n\n\n\n","metadata":{"id":"M8q9wt4TwzPt"}},{"cell_type":"markdown","source":"## Utils\n","metadata":{"id":"wTZ-lv47XOj0"}},{"cell_type":"code","source":"class PermuteBlock(torch.nn.Module):\n    def forward(self, x):\n        return x.transpose(1, 2)\n\ndef plot_attention(attention):\n    # Function for plotting attention\n    # You need to get a diagonal plot\n    plt.clf()\n    seaborn.heatmap(attention, cmap='GnBu')\n    plt.show()\n\ndef save_model(model, optimizer, scheduler, tf_scheduler, metric, epoch, path):\n    torch.save(\n        {'model_state_dict'         : model.state_dict(),\n         'optimizer_state_dict'     : optimizer.state_dict(),\n         'scheduler_state_dict'     : scheduler.state_dict(),\n         'tf_scheduler'             : tf_scheduler,\n         metric[0]                  : metric[1],\n         'epoch'                    : epoch},\n         path\n    )\n\ndef load_model(best_path, epoch_path, model, mode= 'best', metric= 'valid_acc', optimizer= None, scheduler= None, tf_scheduler= None):\n\n\n    if mode == 'best':\n        checkpoint  = torch.load(best_path)\n        print(\"Loading best checkpoint: \", checkpoint[metric])\n    else:\n        checkpoint  = torch.load(epoch_path)\n        print(\"Loading epoch checkpoint: \", checkpoint[metric])\n\n    model.load_state_dict(checkpoint['model_state_dict'], strict= False)\n\n    if optimizer != None:\n        optimizer.load_state_dict(checkpoint['optimizer_state_dict'])\n        #optimizer.param_groups[0]['lr'] = 1.5e-3\n        optimizer.param_groups[0]['weight_decay'] = 1e-5\n    if scheduler != None:\n        scheduler.load_state_dict(checkpoint['scheduler_state_dict'])\n    if tf_scheduler != None:\n        tf_scheduler    = checkpoint['tf_scheduler']\n\n    epoch   = checkpoint['epoch']\n    metric  = torch.load(best_path)[metric]\n\n    return [model, optimizer, scheduler, tf_scheduler, epoch, metric]\n\nclass TimeElapsed():\n    def __init__(self):\n        self.start  = -1\n\n    def time_elapsed(self):\n        if self.start == -1:\n            self.start = time.time()\n        else:\n            end = time.time() - self.start\n            hrs, rem    = divmod(end, 3600)\n            min, sec    = divmod(rem, 60)\n            min         = min + 60*hrs\n            print(\"Time Elapsed: {:0>2}:{:02}\".format(int(min),int(sec)))\n            self.start  = -1","metadata":{"id":"FuRzPOaX0EtJ","execution":{"iopub.status.busy":"2023-12-01T21:13:10.57536Z","iopub.execute_input":"2023-12-01T21:13:10.576239Z","iopub.status.idle":"2023-12-01T21:13:10.58995Z","shell.execute_reply.started":"2023-12-01T21:13:10.576199Z","shell.execute_reply":"2023-12-01T21:13:10.589046Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Modules","metadata":{"id":"YiUnK0GMXTY6"}},{"cell_type":"markdown","source":"# Transformer Encoder","metadata":{"id":"nUQUwEHmCxeI"}},{"cell_type":"code","source":"import torch\nfrom torch import nn\n\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\nprint(\"Device: \", DEVICE)\n\nclass PermuteBlock(torch.nn.Module):\n    def forward(self, x):\n        return x.transpose(1, 2)\n\nclass SelfAttention(torch.nn.Module):\n    def __init__(self, d_model, dim_qk, dim_v, dropout= 0.0):\n        super(SelfAttention, self).__init__()\n\n        self.d_model = d_model\n        self.dim_qk = dim_qk\n\n        self.permute    = PermuteBlock()\n\n        self.Key = nn.Linear(d_model, dim_qk)\n        self.Query = nn.Linear(d_model, dim_qk)\n        self.Value = nn.Linear(d_model, dim_v)\n\n        # self.mask =\n        self.softmax = nn.Softmax(2)\n\n        #Check batch multimul\n\n    def forward(self, x_embeds):\n        Key = self.Key(x_embeds)\n        Query = self.Query(x_embeds)\n        Value = self.Value(x_embeds)\n\n        Key_perm = self.permute(Key)\n\n        scaled_att = torch.bmm(Query, Key_perm)/np.sqrt(self.dim_qk)\n        # mask = self.mask #TO-DO mask\n        att_wts = self.softmax(scaled_att)\n\n        mask = torch.tril(torch.ones(att_wts.shape[-1], att_wts.shape[-1])).to(DEVICE)\n\n        masked_att_wts  = att_wts * mask\n        print(\"Attention shape:\", masked_att_wts.shape)\n\n        context = torch.bmm(masked_att_wts, Value)\n\n        return context, masked_att_wts\n\n\nclass MultiHeadSelfAttention(torch.nn.Module):\n    def __init__(self, d_model, dim_qk, dim_v, num_heads, dropout= 0.0):\n        super(MultiHeadSelfAttention, self).__init__()\n\n        self.dim_qk = dim_qk\n        self.dim_v = dim_v\n\n        self.permute    = PermuteBlock()\n\n\n        self.multi_atts = []\n        for h in range(num_heads):\n            attent = SelfAttention(d_model, dim_qk//num_heads, dim_v//num_heads).to(DEVICE)\n            self.multi_atts.append(attent)\n        # self.multi_atts = [SelfAttention(d_model, dim_qk//num_heads, dim_v//num_heads)]*num_heads\n\n        self.mlp = nn.Linear(dim_v, d_model)\n\n    def forward(self, x_embeds):\n\n        contexts = []\n\n        for att in self.multi_atts:\n            context, _ = att(x_embeds)\n            contexts.append(context)\n\n        contexts_conc = torch.concat(contexts, dim=-1)\n        # print(contexts_conc.shape)\n        final_cont = self.mlp(contexts_conc)\n\n        return final_cont\n\nclass TransformerEncoder(torch.nn.Module):\n    def __init__(self, d_model, dim_qk, dim_v, num_heads, dropout= 0.0):\n        super(TransformerEncoder, self).__init__()\n\n        self.mult_attention = MultiHeadSelfAttention(d_model, dim_qk, dim_v, num_heads)\n\n        # self.layer_norm = LayerNorm() #TO-DO implement LayerNorm\n        self.layer_norm = nn.LayerNorm(d_model)\n        self.permute    = PermuteBlock()\n\n        self.mlp = nn.Sequential(\n            nn.Linear(d_model, d_model)\n        )\n\n\n    def forward(self, x):\n\n        context = self.mult_attention(x)\n\n        context_norm = self.layer_norm(context)\n\n        out_att = x + context_norm\n\n        # print(out_att.shape)\n\n        out_mlp = self.mlp(out_att)\n\n        out_norm = self.layer_norm(out_mlp)\n\n        out =  out_att + out_norm\n\n        return out\n\n\n# batch, sentence_length, embedding_dim = 8, 8, 16\n# embedding = torch.randn(batch, sentence_length, embedding_dim).to(DEVICE)\n\n# model = TransformerEncoder(28, 8, 8, 2).to(DEVICE)\n\n# out = model(x.to(DEVICE))\n","metadata":{"id":"yTvGkyPKH3qW","outputId":"ebf0d53b-5e73-4fb7-ece6-b0c1be4ab27d","execution":{"iopub.status.busy":"2023-12-01T21:13:10.593588Z","iopub.execute_input":"2023-12-01T21:13:10.594029Z","iopub.status.idle":"2023-12-01T21:13:10.613472Z","shell.execute_reply.started":"2023-12-01T21:13:10.594003Z","shell.execute_reply":"2023-12-01T21:13:10.612655Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x.shape","metadata":{"id":"B1ZCZvclDeyc","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# out_np = out.detach().cpu().numpy()\n# plot_attention(out_np[4])","metadata":{"id":"-RgDHSXwhPuj"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import math\n\nclass PositionalEncoding(torch.nn.Module):\n\n    def __init__(self, projection_size, max_seq_len= 3000):\n        super().__init__()\n        # Read the Attention Is All You Need paper to learn how to code code the positional encoding\n\n        # TODO\n        self.projection_size = projection_size\n        self.max_seq_len = max_seq_len\n\n        self.embeds =  torch.zeros([1, self.max_seq_len, self.projection_size]).to(DEVICE)\n\n        for k in range(self.max_seq_len):\n            for i in range(self.projection_size//2):\n                den = 10000**((2*i)/self.projection_size)\n                self.embeds[:, k, 2*i] = torch.sin(torch.tensor([k/den]))\n                self.embeds[:, k, 2*i+1] = torch.cos(torch.tensor([k/den]))\n\n    def forward(self, x):\n        # TODO\n        ln = x.shape[1]\n        self.embeds = self.embeds[:, :ln, :]\n        return x + self.embeds\n\nclass pBLSTM(torch.nn.Module):\n\n    def __init__(self, input_size, hidden_size):\n        super(pBLSTM, self).__init__()\n\n        self.blstm = nn.LSTM(input_size=2*input_size, hidden_size=hidden_size,num_layers=1, batch_first= True, bidirectional= True)# TODO: Initialize a single layer bidirectional LSTM with the given input_size and hidden_size\n\n    def forward(self, x_packed):\n\n        x_pad, x_seq_len = pad_packed_sequence(x_packed, batch_first= True)\n        x_trunc, x_len = self.trunc_reshape(x_pad,x_seq_len)\n        x_packed = pack_padded_sequence(x_trunc,x_len.cpu(), batch_first= True, enforce_sorted= False)\n        output, (h_n, c_n) = self.blstm(x_packed)\n\n        return output\n\n    def trunc_reshape(self, x, x_lens):\n        if x.shape[1]%2 != 0:\n          x = x[:,:-1,:]      # remove last element if time steps are odd\n          x_lens -= 1\n        x = x.reshape((x.shape[0],x.shape[1]//2,x.shape[2]*2))\n        x_lens = x_lens//2\n        return x, x_lens\n\nclass TransformerListener(torch.nn.Module):\n\n    def __init__(self,\n                 input_size,\n                 base_lstm_layers        = 1,\n                 pblstm_layers           = 1,\n                 listener_hidden_size    = 256,\n                 n_heads                 = 8,\n                 tf_blocks               = 1,\n                 d_model=512,\n                 dim_qk=64,\n                 dim_v=64):\n        super().__init__()\n\n        # create an lstm layer\n        self.base_lstm = nn.LSTM(input_size, listener_hidden_size, base_lstm_layers,  batch_first=True, dropout = 0.5, bidirectional=True)\n\n        self.pblstm = pBLSTM(listener_hidden_size*2, listener_hidden_size)\n\n        # create a sequence of Conv1d layers\n        self.permute    = PermuteBlock()\n\n        self.embedding = nn.Sequential(\n            nn.Conv1d(listener_hidden_size*2, 256, kernel_size = 5, stride=1),\n            nn.BatchNorm1d(256),\n            nn.GELU(),\n            # nn.MaxPool1d(kernel_size= 2, stride= 2),\n            nn.Dropout(0.5),\n            nn.Conv1d(256, d_model, kernel_size = 5, stride=1),\n            nn.BatchNorm1d(d_model),\n            nn.GELU(),\n            # nn.MaxPool1d(kernel_size= 2, stride= 2),\n            nn.Dropout(0.5)\n        )\n\n        # compute the postion encoding\n        self.positional_encoding    = PositionalEncoding(d_model, max_seq_len=3000)# TODO\n\n        # create a sequence of transformer blocks\n        blocks = []\n        for i in range(tf_blocks):\n            blocks.append(TransformerEncoder(d_model, dim_qk, dim_v, n_heads))\n            # TODO\n        self.transformer_encoder    = torch.nn.Sequential(*blocks)\n\n    def forward(self, x, x_len):\n\n        # pack the inputs before passing them to the LSTm\n        x_packed = pack_padded_sequence(x, batch_first=True, lengths=x_len, enforce_sorted=False)\n\n        # Pass the packed sequence through the lstm\n        lstm_out, _             = self.base_lstm(x_packed)# TODO\n        pblstm_out            = self.pblstm(lstm_out)# TODO\n\n        # Unpack the output of the lstm\n        output, output_lengths  = pad_packed_sequence(pblstm_out, batch_first=True)# TODO: Need to 'unpack' the LSTM output using pad_packed_sequence\n\n        output = self.permute(output)\n        # Pass the output through the embedding\n        output                  = self.embedding(output)# TODO\n        # calculate the new output length\n        output_lengths = torch.clamp(output_lengths, max=output.shape[2])\n\n        output = self.permute(output)\n\n        # calculate the position encoding\n        output  = self.positional_encoding(output) # TODO\n        # Pass the output of the positional encoding through the transformer encoder\n        print(output.shape)\n\n        output  = self.transformer_encoder(output) # TODO\n\n\n        return output, output_lengths\n\n\n# batch, sentence_length, embedding_dim = 8, 32, 28\n# x = torch.randn(batch, sentence_length, embedding_dim).to(DEVICE)\n# lens = torch.randint(5, 32, (8,))\n\n# model = TransformerListener(\n#                  input_size = 28,\n#                  base_lstm_layers        = 1,\n#                  pblstm_layers           = 1,\n#                  listener_hidden_size    = 28,\n#                  n_heads                 = 8,\n#                  tf_blocks               = 2,\n#                  d_model=512,\n#                  dim_qk=64,\n#                  dim_v=64\n# ).to(DEVICE)\n\n# # batch, sentence_length, embedding_dim = 8, 8, 16\n# # embedding = torch.randn(batch, sentence_length, embedding_dim).to(DEVICE)\n\n# out = model(x.to(DEVICE), x_len)\n","metadata":{"id":"F0opYqry_EGi","outputId":"907dbae0-5388-4a81-fd42-d789e975c4ae","execution":{"iopub.status.busy":"2023-12-01T21:13:10.614595Z","iopub.execute_input":"2023-12-01T21:13:10.614944Z","iopub.status.idle":"2023-12-01T21:13:10.637359Z","shell.execute_reply.started":"2023-12-01T21:13:10.614918Z","shell.execute_reply":"2023-12-01T21:13:10.636505Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x_len","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nbatch, sentence_length, embedding_dim = 8, 32, 28\nx = torch.randn(batch, sentence_length)","metadata":{"id":"QZdK0HaVWF6n"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model","metadata":{"id":"ar9hWGN2Em6S","outputId":"8ff58140-e574-4963-8494-ebac66cb6d45"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.nn.utils.rnn import pad_sequence\nimport torch\n# Example usage\nsequences = [torch.tensor([1, 2, 3]), torch.tensor([4, 5]), torch.tensor([6])]\npadded_sequence = pad_sequence(sequences, batch_first=True, padding_value=0)\npadded_sequence.shape","metadata":{"id":"33Xn8tJesXvO","outputId":"b583637f-0212-497b-8c89-35b379551269"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Attention\n\n### Different ways to compute Attention\n\n1. Dot-product attention\n    * raw_weights = bmm(key, query)\n    * Optional: Scaled dot-product by normalizing with sqrt key dimension\n    * Check \"Attention is All You Need\" Section 3.2.1\n    * 1st way is what most TAs are comfortable with, but if you want to explore, check out other methods below\n\n\n2. Cosine attention\n    * raw_weights = cosine(query, key) # almost the same as dot-product xD\n\n3. Bi-linear attention\n    * W = Linear transformation (learnable parameter): d_k -> d_q\n    * raw_weights = bmm(key @ W, query)\n\n4. Multi-layer perceptron\n    * Check \"Neural Machine Translation and Sequence-to-sequence Models: A Tutorial\" Section 8.4\n\n5. Multi-Head Attention\n    * Check \"Attention is All You Need\" Section 3.2.2\n    * h = Number of heads\n    * W_Q, W_K, W_V: Weight matrix for Q, K, V (h of them in total)\n    * W_O: d_v -> d_v\n    * Reshape K: (B, T, d_k) to (B, T, h, d_k // h) and transpose to (B, h, T, d_k // h)\n    * Reshape V: (B, T, d_v) to (B, T, h, d_v // h) and transpose to (B, h, T, d_v // h)\n    * Reshape Q: (B, d_q) to (B, h, d_q // h) `\n    * raw_weights = Q @ K^T\n    * masked_raw_weights = mask(raw_weights)\n    * attention = softmax(masked_raw_weights)\n    * multi_head = attention @ V\n    * multi_head = multi_head reshaped to (B, d_v)\n    * context = multi_head @ W_O","metadata":{"id":"5fG9jDZBVklL"}},{"cell_type":"markdown","source":"Pseudocode:\n\n```python\nclass Attention:\n    '''\n    Attention is calculated using the key, value (from encoder embeddings) and query from decoder.\n\n    After obtaining the raw weights, compute and return attention weights and context as follows.:\n\n    attention_weights   = softmax(raw_weights)\n    attention_context   = einsum(\"thinkwhatwouldbetheequationhere\",attention, value) #take hint from raw_weights calculation\n\n    At the end, you can pass context through a linear layer too.\n    '''\n\n    def init(listener_hidden_size,\n              speller_hidden_size,\n              projection_size):\n\n        VW = Linear(listener_hidden_size,projection_size)\n        KW = Linear(listener_hidden_size,projection_size)\n        QW = Linear(speller_hidden_size,projection_size)\n\n    def set_key_value(encoder_outputs):\n        '''\n        In this function we take the encoder embeddings and make key and values from it.\n        key.shape   = (batch_size, timesteps, projection_size)\n        value.shape = (batch_size, timesteps, projection_size)\n        '''\n        key = KW(encoder_outputs)\n        value = VW(encoder_outputs)\n      \n    def compute_context(decoder_context):\n        '''\n        In this function from decoder context, we make the query, and then we\n         multiply the queries with the keys to find the attention logits,\n         finally we take a softmax to calculate attention energy which gets\n         multiplied to the generted values and then gets summed.\n\n        key.shape   = (batch_size, timesteps, projection_size)\n        value.shape = (batch_size, timesteps, projection_size)\n        query.shape = (batch_size, projection_size)\n\n        You are also recomended to check out Abu's Lecture 19 to understand Attention better.\n        '''\n        query = QW(decoder_context) #(batch_size, projection_size)\n\n        raw_weights = #using bmm or einsum. We need to perform batch matrix multiplication. It is important you do this step correctly.\n        #What will be the shape of raw_weights?\n\n        attention_weights = #What makes raw_weights -> attention_weights\n\n        attention_context = #Multiply attention weights to values\n\n        return attention_context, attention_weights\n```","metadata":{"id":"wyv0Q65t5SDd"}},{"cell_type":"code","source":"class Attention(torch.nn.Module):\n  def __init__(self):\n    super().__init__()\n    pass\n\n  def set_key_value(self, encoder_outputs):\n    pass\n\n  def compute_context(self, decoder_context):\n    pass","metadata":{"id":"771TXxn7ViOW"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# The Speller\n\nSimilar to the language model that you coded up for HW4P1, you have to code a language model for HW4P2 as well. This time, we will also call the attention context step, within the decoder to get the attended-encoder-embeddings.\n\n\nWhat you have coded till now:\n\n<center>\n<img src=\"https://github.com/varunjain3/11785_s23_h4p2/raw/main/EncoderAttention.png\" alt=\"data flow\" height=\"400\">\n</center>\n\nFor the Speller, what we have to code:\n\n\n<center>\n<img src=\"https://github.com/varunjain3/11785_s23_h4p2/raw/main/Decoder.png\" alt=\"data flow\" height=\"400\">\n</center>","metadata":{"id":"4Sp1WywZmm1L"}},{"cell_type":"code","source":"\n# class CrossAttention(torch.nn.Module):\n#     def __init__(self, d_model, dim_qk, dim_v, dropout= 0.0):\n#         super().__init__()\n\n#         self.d_model = d_model\n#         self.dim_qk = dim_qk\n\n#         self.permute    = PermuteBlock()\n\n#         self.Key = nn.Linear(d_model, dim_qk)\n#         self.Query = nn.Linear(d_model, dim_qk)\n#         self.Value = nn.Linear(d_model, dim_v)\n\n#         # self.mask =\n#         self.softmax = nn.Softmax(2)\n\n#         #Check batch multimul\n\n#     def forward(self, encoder_embeds, decoder_embeds):\n#         Key = self.Key(encoder_embeds)\n#         Query = self.Query(decoder_embeds)\n#         Value = self.Value(encoder_embeds)\n\n#         Key_perm = self.permute(Key)\n\n#         scaled_att = torch.bmm(Query, Key_perm)/np.sqrt(self.dim_qk)\n#         # mask = self.mask #TO-DO mask\n#         att_wts = self.softmax(scaled_att)\n\n#         mask = torch.tril(torch.ones(att_wts.shape[-1], att_wts.shape[-1]))\n\n#         masked_att_wts  = att_wts * mask\n\n#         context = torch.bmm(masked_att_wts, Value)\n\n#         return context, masked_att_wts","metadata":{"id":"ffuJioVsRHVI"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CrossAttention(torch.nn.Module):\n\n    def __init__(self, listener_hidden_size, speller_hidden_size, projection_size):\n        super(). __init__()\n\n        self.projection_size = projection_size\n\n        self.VW = nn.Linear(listener_hidden_size, projection_size)\n        self.KW = nn.Linear(listener_hidden_size, projection_size)\n        self.QW = nn.Linear(speller_hidden_size, projection_size)\n        self.permute    = PermuteBlock()\n        self.softmax = nn.Softmax(1)\n\n    def set_key_value(self, encoder_outputs):\n        self.key = self.KW(encoder_outputs)\n        self.value = self.VW(encoder_outputs)\n#         print(self.key.shape)\n\n    def compute_context(self, decoder_context):\n\n        query = self.QW(decoder_context) #(batch_size, projection_size)\n#         key = self.permute(self.key)\n\n        # raw_weights = #using bmm or einsum. We need to perform batch matrix multiplication. It is important you do this step correctly.\n        raw_weights = torch.bmm(self.key, query.unsqueeze(2)).squeeze(2)/np.sqrt(self.projection_size)\n\n        #What will be the shape of raw_weights?\n\n        attention_weights = self.softmax(raw_weights) #What makes raw_weights -> attention_weights\n#         print(attention_weights.shape)\n        attention_context = torch.bmm(attention_weights.unsqueeze(1), self.value).squeeze(1) #Multiply attention weights to values\n\n        return attention_context, attention_weights\n    \nattention = CrossAttention(28, 28, 256).to(DEVICE)\nattention.set_key_value(x.to(DEVICE))\nattention_context, attention_weights = attention.compute_context(x[:,0,:].to(DEVICE))\nprint(attention_context.shape)","metadata":{"id":"EWsF0htTPHrL","execution":{"iopub.status.busy":"2023-12-01T21:13:10.638347Z","iopub.execute_input":"2023-12-01T21:13:10.638628Z","iopub.status.idle":"2023-12-01T21:13:12.430339Z","shell.execute_reply.started":"2023-12-01T21:13:10.638584Z","shell.execute_reply":"2023-12-01T21:13:12.429304Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.empty_cache()\ngc.collect()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\nclass Speller(torch.nn.Module):\n\n  # Refer to your HW4P1 implementation for help with setting up the language model.\n  # The only thing you need to implement on top of your HW4P1 model is the attention module and teacher forcing.\n\n  def __init__(self, input_size, d_model, dim_qk, dim_v,  attender:CrossAttention, lstm_hidden_size=512, max_timesteps=64, vocab_size=31):\n    super(). __init__()\n\n    # self.attend = CrossAttention(d_model, dim_qk, dim_v) # Attention object in speller\n    self.attend = attender\n    self.max_timesteps = max_timesteps # Max timesteps\n    self.log_softmax = nn.LogSoftmax(dim=2)\n    self.dim_v = dim_v\n    self.d_model = d_model\n    self.lstm_hidden_size = lstm_hidden_size\n    self.proj_size = attender.projection_size\n    self.batch_size =  config['batch_size']\n\n\n    # self.embedding =  # Embedding layer to convert token to latent space\n\n    self.embedding = torch.nn.Embedding(vocab_size, d_model)\n\n#     self.lstm_cells = [nn.LSTMCell(self.proj_size+self.d_model, lstm_hidden_size), nn.LSTMCell(lstm_hidden_size, lstm_hidden_size), nn.LSTMCell(lstm_hidden_size, lstm_hidden_size)] # Create a sequence of LSTM Cells\n\n    self.lstm_cells = torch.nn.Sequential(\n            torch.nn.LSTMCell(self.d_model+self.attend.projection_size, lstm_hidden_size),\n            torch.nn.LSTMCell(lstm_hidden_size, lstm_hidden_size),\n            torch.nn.LSTMCell(lstm_hidden_size, lstm_hidden_size)\n    )\n    \n    # For CDN (Feel free to change)\n    # Linear module to convert outputs to correct hidden size (Optional: TO make dimensions match)\n    self.output_to_char = nn.Sequential(\n            nn.Linear(lstm_hidden_size + self.proj_size, 2048),\n            nn.ReLU(),\n            nn.Dropout(0.3),\n            nn.Linear(2048, d_model),\n            nn.ReLU(),\n            nn.Dropout(0.3),\n        )\n    # self.activation = # Check which activation is suggested\n    self.char_prob = nn.Linear(d_model, vocab_size) # Linear layer to convert hidden space back to logits for token classification\n    self.char_prob.weight = self.embedding.weight# Weight tying (From embedding layer)\n    self.dropout = torch.nn.Dropout(0.25)\n    self.lstm_dropout = torch.nn.Dropout(0.25)\n\n  def lstm_step(self, input_word, hidden_states):\n        # raise NotImplementedError # Feed the input through each LSTM Cell\n\n        embedding = input_word.to(DEVICE)\n\n        for i in range(len(self.lstm_cells)):\n            hidden_states[i] = self.lstm_cells[i](embedding, hidden_states[i])\n            embedding = self.lstm_dropout(hidden_states[i][0])\n\n        return embedding, hidden_states # What information does forward() need?\n\n  def CDN(self,input_):\n    # Make the CDN here, you can add the output-to-char\n    out =  self.output_to_char(input_)\n    out_prob =  self.char_prob(out)\n\n    return out_prob\n\n\n  def forward (self, y=None, teacher_forcing_ratio=1):\n\n    attn_context =  torch.zeros(config['batch_size'], self.proj_size).to(DEVICE)# initial context tensor for time t = 0\n    output_symbol = torch.tensor([SOS_TOKEN]*self.batch_size).to(DEVICE) # Set it to SOS for time t = 0\n    raw_outputs = []\n    attention_plot = []\n\n    if y is None:\n      timesteps = self.max_timesteps\n      teacher_forcing_ratio = 0 #Why does it become zero?\n\n    else:\n      timesteps = y.shape[1] # How many timesteps are we predicting for?\n\n#     hidden_states_list = [torch.zeros(self.lstm_hidden_size, self.lstm_hidden_size) for lc in self.lstm_cells] # Initialize your hidden_states list here similar to HW4P1\n#     hidden_states_list[0] = torch.zeros(self.proj_size+self.d_model, self.lstm_hidden_size)\n    hidden_states_list = [None] * len(self.lstm_cells)\n#     hidden_states_list = hidden_states_list.to(DEVICE)\n    for t in range(timesteps):\n      p = np.random.random_sample() # generate a probability p between 0 and 1\n\n      if p < teacher_forcing_ratio and t > 0: # Why do we consider cases only when t > 0? What is considered when t == 0? Think.\n        output_symbol = y[:, t-1] # Take from y, else draw from probability distribution\n\n#       print(t, output_symbol.shape)\n      char_embed = self.embedding(output_symbol) # Embed the character symbol\n    \n#       print(\"attn_context shape\", attn_context.shape)  \n      # Concatenate the character embedding and context from attention, as shown in the diagram\n      lstm_input = torch.cat((char_embed, attn_context), dim=-1)\n\n      lstm_out, hidden_states = self.lstm_step(lstm_input, hidden_states_list) # Feed the input through LSTM Cells and attention.\n      # What should we retrieve from forward_step to prepare for the next timestep?\n\n      attn_context, attn_weights = self.attend.compute_context(lstm_out) # Feed the resulting hidden state into attention\n\n      cdn_input = torch.cat((attn_context, lstm_out), dim=-1) # TODO: You need to concatenate the context from the attention module with the LSTM output hidden state, as shown in the diagram\n\n      raw_pred = self.CDN(cdn_input) # call CDN with cdn_input\n\n      # Generate a prediction for this timestep and collect it in output_symbols\n      output_symbol = torch.argmax(raw_pred) # Draw correctly from raw_pred\n\n      raw_outputs.append(raw_pred) # for loss calculation\n      attention_plot.append(attn_weights) # for plotting attention plot\n\n\n    attention_plot = torch.stack(attention_plot, dim=1)\n    raw_outputs = torch.stack(raw_outputs, dim=1)\n\n    return raw_outputs, attention_plot\n\nspeller = Speller(input_size=28, \n                  d_model=64, \n                  dim_qk=64, \n                  dim_v=64,  \n                  attender=attention,\n                  lstm_hidden_size=28, \n                  max_timesteps=64, \n                  vocab_size=31).to(DEVICE)\n\nraw_outputs, attention_plot = speller(y.to(DEVICE))","metadata":{"id":"nFkc6MbnlUPu","execution":{"iopub.status.busy":"2023-12-01T13:21:31.811042Z","iopub.execute_input":"2023-12-01T13:21:31.811462Z","iopub.status.idle":"2023-12-01T13:21:32.263552Z","shell.execute_reply.started":"2023-12-01T13:21:31.81143Z","shell.execute_reply":"2023-12-01T13:21:32.262692Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y[:, 1]","metadata":{"execution":{"iopub.status.busy":"2023-12-01T02:27:32.395068Z","iopub.execute_input":"2023-12-01T02:27:32.395866Z","iopub.status.idle":"2023-12-01T02:27:32.403331Z","shell.execute_reply.started":"2023-12-01T02:27:32.395817Z","shell.execute_reply":"2023-12-01T02:27:32.402302Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nclass ASRModel(torch.nn.Module):\n  def __init__(self,): # add parameters\n    super().__init__()\n\n    # Pass the right parameters here\n    self.listener = TransformerListener(\n                 input_size = 28,\n                 base_lstm_layers        = 1,\n                 pblstm_layers           = 1,\n                 listener_hidden_size    = 256,\n                 n_heads                 = 8,\n                 tf_blocks               = 2,\n                 d_model=256,\n                 dim_qk=256,\n                 dim_v=256)\n\n#     self.attend = Attention()\n    self.attend = CrossAttention(listener_hidden_size = 256, \n                                 speller_hidden_size = 256, \n                                 projection_size = 256).to(DEVICE)\n    \n#     self.speller = Speller(self.attend)\n    self.speller = Speller(input_size=28, \n                  d_model=256, \n                  dim_qk=256, \n                  dim_v=256,  \n                  attender=self.attend,\n                  lstm_hidden_size=256, \n                  max_timesteps=128, \n                  vocab_size=31).to(DEVICE)\n  def forward(self, x, lx, y=None, teacher_forcing_ratio=1):\n    # Encode speech features\n    encoder_outputs, _ = self.listener(x, lx)\n\n    # We want to compute keys and values ahead of the decoding step, as they are constant for all timesteps\n    # Set keys and values using the encoder outputs\n    self.attend.set_key_value(encoder_outputs)\n\n    # Decode text with the speller using context from the attention\n    raw_outputs, attention_plots = self.speller(y=y,teacher_forcing_ratio=teacher_forcing_ratio)\n\n    return raw_outputs, attention_plots\n        \nmodel = ASRModel().to(DEVICE)\nmodel(x.to(DEVICE), x_len)","metadata":{"id":"scvB2cI-OSof","execution":{"iopub.status.busy":"2023-12-01T13:56:39.841508Z","iopub.execute_input":"2023-12-01T13:56:39.841944Z","iopub.status.idle":"2023-12-01T13:57:10.957717Z","shell.execute_reply.started":"2023-12-01T13:56:39.84191Z","shell.execute_reply":"2023-12-01T13:57:10.956115Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"id":"pHSOhT4tCnFj"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nconv =  torch.nn.Conv1d(28, 256, kernel_size = 5, stride=1)\nembedding          = torch.nn.Embedding(28, 256)\n\n\nlin = torch.nn.Linear(256, 28)\nlin.weight = embedding.weight\n","metadata":{"id":"4ZZB9MwHCnJD"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"lin(x)\n","metadata":{"id":"ev1YY5oNDEqB","outputId":"bde1af4a-edd7-4d2b-99bb-f565932fd937"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"conv.weight.shape","metadata":{"id":"QVTbaPV0DUxe","outputId":"f93d04b0-2d30-428e-ec3f-60f8690067af"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x = torch.randn(128,256)\n# lin = torch.nn.Linear(256, 64)\n","metadata":{"id":"0R9Cqy7pJjYe","outputId":"0452f83a-62b6-4492-afbc-521beeb1b104"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model Setup","metadata":{"id":"bPZD3vqdUisj"}},{"cell_type":"code","source":"torch.cuda.empty_cache()\ngc.collect()","metadata":{"id":"p2uX2P9YVnbk","outputId":"2063db54-51e0-47e9-f9b6-e36acf12540e"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = ASRModel(\n\n    # Initialize your model\n)\n\nmodel = model.to(DEVICE)\nprint(model)","metadata":{"id":"a9LN0l5VUk_s","scrolled":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss Function, Optimizers, Scheduler","metadata":{"id":"23DMfXsaU6kj"}},{"cell_type":"code","source":"optimizer   = # TODO\n\ncriterion   = # TODO\n\nscaler      = # TODO\n\nscheduler   = # TODO","metadata":{"id":"216ukmHbU-ol"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Levenshtein Distance","metadata":{"id":"ZWQnB8lUVY4f"}},{"cell_type":"code","source":"# We have given you this utility function which takes a sequence of indices and converts them to a list of characters\ndef indices_to_chars(indices, vocab):\n    tokens = []\n    for i in indices: # This loops through all the indices\n        if int(i) == SOS_TOKEN: # If SOS is encountered, dont add it to the final list\n            continue\n        elif int(i) == EOS_TOKEN: # If EOS is encountered, stop the decoding process\n            break\n        else:\n            tokens.append(vocab[i])\n    return tokens\n\n# To make your life more easier, we have given the Levenshtein distantce / Edit distance calculation code\ndef calc_edit_distance(predictions, y, y_len, vocab= VOCAB, print_example= False):\n\n    dist                = 0\n    batch_size, seq_len = predictions.shape\n\n    for batch_idx in range(batch_size):\n\n        y_sliced    = indices_to_chars(y[batch_idx,0:y_len[batch_idx]], vocab)\n        pred_sliced = indices_to_chars(predictions[batch_idx], vocab)\n\n        # Strings - When you are using characters from the AudioDataset\n        y_string    = ''.join(y_sliced)\n        pred_string = ''.join(pred_sliced)\n\n        #dist        += Levenshtein.distance(pred_string, y_string)\n        # Comment the above abd uncomment below for toy dataset\n        dist      += Levenshtein.distance(y_sliced, pred_sliced)\n\n    if print_example:\n        # Print y_sliced and pred_sliced if you are using the toy dataset\n        print(\"\\nGround Truth : \", y_string)\n        print(\"Prediction   : \", pred_string)\n\n    dist    /= batch_size\n    return dist","metadata":{"id":"rSsiCdxPVeZW"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train and Validation functions\n","metadata":{"id":"Pu4MrSMUUIyp"}},{"cell_type":"code","source":"def train(model, dataloader, criterion, optimizer, teacher_forcing_rate):\n\n    model.train()\n    batch_bar = tqdm(total=len(dataloader), dynamic_ncols=True, leave=False, position=0, desc='Train')\n\n    running_loss        = 0.0\n    running_perplexity  = 0.0\n\n    for i, (x, y, lx, ly) in enumerate(dataloader):\n\n        optimizer.zero_grad()\n\n        x, y, lx, ly = x.to(DEVICE), y.to(DEVICE), lx, ly\n\n        with torch.cuda.amp.autocast():\n\n            raw_predictions, attention_plot = model(x, lx, y= y, tf_rate= teacher_forcing_rate)\n\n            # Predictions are of Shape (batch_size, timesteps, vocab_size).\n            # Transcripts are of shape (batch_size, timesteps) Which means that you have batch_size amount of batches with timestep number of tokens.\n            # So in total, you have batch_size*timesteps amount of characters.\n            # Similarly, in predictions, you have batch_size*timesteps amount of probability distributions.\n            # How do you need to modify transcipts and predictions so that you can calculate the CrossEntropyLoss? Hint: Use Reshape/View and read the docs\n            # Also we recommend you plot the attention weights, you should get convergence in around 10 epochs, if not, there could be something wrong with\n            # your implementation\n            loss        =  # TODO: Cross Entropy Loss\n\n            perplexity  = torch.exp(loss) # Perplexity is defined the exponential of the loss\n\n            running_loss        += loss.item()\n            running_perplexity  += perplexity.item()\n\n        # Backward on the masked loss\n        scaler.scale(loss).backward()\n\n        # Optional: Use torch.nn.utils.clip_grad_norm to clip gradients to prevent them from exploding, if necessary\n        # If using with mixed precision, unscale the Optimizer First before doing gradient clipping\n\n        scaler.step(optimizer)\n        scaler.update()\n\n\n        batch_bar.set_postfix(\n            loss=\"{:.04f}\".format(running_loss/(i+1)),\n            perplexity=\"{:.04f}\".format(running_perplexity/(i+1)),\n            lr=\"{:.04f}\".format(float(optimizer.param_groups[0]['lr'])),\n            tf_rate='{:.02f}'.format(teacher_forcing_rate))\n        batch_bar.update()\n\n        del x, y, lx, ly\n        torch.cuda.empty_cache()\n\n    running_loss /= len(dataloader)\n    running_perplexity /= len(dataloader)\n    batch_bar.close()\n\n    return running_loss, running_perplexity, attention_plot","metadata":{"id":"gYRVKs9_2rsx"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def validate(model, dataloader):\n\n    model.eval()\n\n    batch_bar = tqdm(total=len(dataloader), dynamic_ncols=True, position=0, leave=False, desc=\"Val\")\n\n    running_lev_dist = 0.0\n\n    for i, (x, y, lx, ly) in enumerate(dataloader):\n\n        x, y, lx, ly = x.to(DEVICE), y.to(DEVICE), lx, ly\n\n        with torch.inference_mode():\n            raw_predictions, attentions = model(x, lx, y = None)\n\n        # Greedy Decoding\n        greedy_predictions   =  # TODO: How do you get the most likely character from each distribution in the batch?\n\n        # Calculate Levenshtein Distance\n        running_lev_dist    += calc_edit_distance(greedy_predictions, y, ly, VOCAB, print_example = False) # You can use print_example = True for one specific index i in your batches if you want\n\n        batch_bar.set_postfix(\n            dist=\"{:.04f}\".format(running_lev_dist/(i+1)))\n        batch_bar.update()\n\n        del x, y, lx, ly\n        torch.cuda.empty_cache()\n\n    batch_bar.close()\n    running_lev_dist /= len(dataloader)\n\n    return running_lev_dist","metadata":{"id":"uIx3tW7a2tze"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Wandb\n","metadata":{"id":"WhwhevgWQbDX"}},{"cell_type":"code","source":"# Login to Wandb\n# Initialize your Wandb Run Here\n# Save your model architecture in a txt file, and save the file to Wandb","metadata":{"id":"1Xbw_0eAQcoR"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_attention(attention):\n    # Function for plotting attention\n    # You need to get a diagonal plot\n    plt.clf()\n    sns.heatmap(attention, cmap='GnBu')\n    plt.show()","metadata":{"id":"vitGNx_O3MjW"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Experiment","metadata":{"id":"JmZhxhNseaIr"}},{"cell_type":"code","source":"best_lev_dist = float(\"inf\")\ntf_rate = 1.0\n\nfor epoch in range(0, config['epochs']):\n\n    print(\"\\nEpoch: {}/{}\".format(epoch+1, config['epochs']))\n\n    # Call train and validate, get attention weights from training\n\n    # Print your metrics\n\n    # Plot Attention for a single item in the batch\n    plot_attention(attention_plot[0].cpu().detach().numpy())\n\n    # Log metrics to Wandb\n\n    # Optional: Scheduler Step / Teacher Force Schedule Step\n\n\n    if valid_dist <= best_lev_dist:\n        best_lev_dist = valid_dist\n        # Save your model checkpoint here","metadata":{"id":"JcTFu-AH3m4e"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Testing","metadata":{"id":"hgFYFaBGeBqM"}},{"cell_type":"code","source":"# Optional: Load your best model Checkpoint here\n\n# TODO: Create a testing function similar to validation\n# TODO: Create a file with all predictions\n# TODO: Submit to Kaggle","metadata":{"id":"ndNCpcxkx2KG"},"execution_count":null,"outputs":[]}]}