{"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":"# Installs","metadata":{"id":"UR4qfYrVoO4v"}},{"cell_type":"code","source":"%pip install torch==1.13.1+cu117 torchvision==0.14.1+cu117 torchtext==0.14.1 torchaudio==0.13.1 torchdata==0.5.1 --extra-index-url https://download.pytorch.org/whl/cu117 -q","metadata":{"id":"mA9qZoIDcx-h","execution":{"iopub.status.busy":"2023-11-10T03:34:10.502316Z","iopub.execute_input":"2023-11-10T03:34:10.502968Z","iopub.status.idle":"2023-11-10T03:35:56.889179Z","shell.execute_reply.started":"2023-11-10T03:34:10.502939Z","shell.execute_reply":"2023-11-10T03:35:56.88803Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"\nThis may take a while","metadata":{"id":"ONgAWhqdoYy-"}},{"cell_type":"code","source":"!pip install wandb --quiet\n!pip install python-Levenshtein -q\n!git clone --recursive https://github.com/parlance/ctcdecode.git\n!pip install wget -q\n%cd ctcdecode\n!pip install . -q\n%cd ..\n\n!pip install torchsummaryX -q","metadata":{"id":"SS7a7xeEoaV9","outputId":"1a238860-511a-4c1f-a6c6-6d0f50d13762","execution":{"iopub.status.busy":"2023-11-10T03:35:56.890412Z","iopub.execute_input":"2023-11-10T03:35:56.890717Z","iopub.status.idle":"2023-11-10T03:38:35.66285Z","shell.execute_reply.started":"2023-11-10T03:35:56.890686Z","shell.execute_reply":"2023-11-10T03:38:35.661765Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Imports","metadata":{"id":"IWVONJxCobPc"}},{"cell_type":"code","source":"import torch\nimport random\nimport numpy as np\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torchsummaryX import summary\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.nn.utils.rnn import pad_sequence, pack_padded_sequence, pad_packed_sequence\n\nimport torchaudio.transforms as tat\n\nfrom sklearn.metrics import accuracy_score\nimport gc\n\nimport zipfile\nimport pandas as pd\nfrom tqdm import tqdm\nimport os\nimport datetime\n\n# imports for decoding and distance calculation\nimport ctcdecode\nimport Levenshtein\nfrom ctcdecode import CTCBeamDecoder\n\nimport warnings\nwarnings.filterwarnings('ignore')\n\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'\nprint(\"Device: \", device)","metadata":{"id":"78ZTCIXoof2f","outputId":"ce5929a1-d13e-4e9f-ea63-ff4ab0c9f1a7","execution":{"iopub.status.busy":"2023-11-10T03:38:35.664314Z","iopub.execute_input":"2023-11-10T03:38:35.66465Z","iopub.status.idle":"2023-11-10T03:38:38.232489Z","shell.execute_reply.started":"2023-11-10T03:38:35.664619Z","shell.execute_reply":"2023-11-10T03:38:38.23151Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Kaggle Setup","metadata":{"id":"gg3-yJ8tok34"}},{"cell_type":"code","source":"","metadata":{"id":"Ty15EP2mDCTj","outputId":"13556bb6-16c5-48df-af06-71a3786b3dba"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"id":"hKHBrRByDRAd","outputId":"9e133452-b7d0-4763-8e78-59dcd7a423d7"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"id":"dSjBwfXeoq4B","outputId":"cfa30456-85ec-44a3-c546-9b6391f0acf5"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"id":"_ruxWP60LCQA","outputId":"2c549ef0-95bc-4d15-9fa7-90fcb92ed38b"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset and Dataloader","metadata":{"id":"2ORNHnSFroP0"}},{"cell_type":"code","source":"# ARPABET PHONEME MAPPING\n# DO NOT CHANGE\n\nCMUdict_ARPAbet = {\n    \"\" : \" \",\n    \"[SIL]\": \"-\", \"NG\": \"G\", \"F\" : \"f\", \"M\" : \"m\", \"AE\": \"@\",\n    \"R\"    : \"r\", \"UW\": \"u\", \"N\" : \"n\", \"IY\": \"i\", \"AW\": \"W\",\n    \"V\"    : \"v\", \"UH\": \"U\", \"OW\": \"o\", \"AA\": \"a\", \"ER\": \"R\",\n    \"HH\"   : \"h\", \"Z\" : \"z\", \"K\" : \"k\", \"CH\": \"C\", \"W\" : \"w\",\n    \"EY\"   : \"e\", \"ZH\": \"Z\", \"T\" : \"t\", \"EH\": \"E\", \"Y\" : \"y\",\n    \"AH\"   : \"A\", \"B\" : \"b\", \"P\" : \"p\", \"TH\": \"T\", \"DH\": \"D\",\n    \"AO\"   : \"c\", \"G\" : \"g\", \"L\" : \"l\", \"JH\": \"j\", \"OY\": \"O\",\n    \"SH\"   : \"S\", \"D\" : \"d\", \"AY\": \"Y\", \"S\" : \"s\", \"IH\": \"I\",\n    \"[SOS]\": \"[SOS]\", \"[EOS]\": \"[EOS]\"\n}\n\nCMUdict = list(CMUdict_ARPAbet.keys())\nARPAbet = list(CMUdict_ARPAbet.values())\n\n\nPHONEMES = CMUdict[:-2]\nLABELS = ARPAbet[:-2]","metadata":{"id":"k0v7wHRWrqH6","execution":{"iopub.status.busy":"2023-11-10T03:38:38.235811Z","iopub.execute_input":"2023-11-10T03:38:38.236322Z","iopub.status.idle":"2023-11-10T03:38:38.24407Z","shell.execute_reply.started":"2023-11-10T03:38:38.236288Z","shell.execute_reply":"2023-11-10T03:38:38.243167Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# You might want to play around with the mapping as a sanity check here","metadata":{"id":"eN2kcxwXLLBb"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Train Data","metadata":{"id":"agmNBKf4JrLV"}},{"cell_type":"code","source":"class AudioDataset(torch.utils.data.Dataset):\n\n    def __init__(self, root, partition=\"train-clean-100\", divider=1, context=0, phonemes=PHONEMES):\n        '''\n        Initializes the dataset.\n\n        INPUTS: What inputs do you need here?\n        '''\n        self.context = context\n        self.divider = divider\n        self.phonemes = phonemes\n\n        # Load the directory and all files in them\n        self.mfcc_dir = os.path.join(root, partition, 'mfcc/')\n        self.transcript_dir = os.path.join(root, partition, 'transcript/')\n\n        self.mfcc_files = sorted(os.listdir(self.mfcc_dir))\n        self.transcript_files = sorted(os.listdir(self.transcript_dir))\n\n        # assert len(self.mfcc_files) == len(self.transcript_files)\n        # TODO: List files in sefl.mfcc_dir using os.listdir in sorted order\n        mfcc_names          = sorted(os.listdir(self.mfcc_dir))\n        # TODO: List files in self.transcript_dir using os.listdir in sorted order\n        transcript_names    = sorted(os.listdir(self.transcript_dir))\n        self.mfccs = []\n        self.transcripts = []\n        # self.mfccsLength = []\n        '''\n        # Calculate the dataset length\n        total_timestamps = 0\n        for i in range(len(self.mfcc_files)//self.divider):\n            mfcc_file = self.mfcc_files[i]\n            total_timestamps += len(np.load(os.path.join(self.mfcc_dir, mfcc_file)))\n        self.length = total_timestamps + 2 * context * (len(self.mfcc_files)//self.divider)\n        '''\n        for i in range(len(mfcc_names)//divider):\n        #   Load a single mfcc\n            mfcc        = np.load(self.mfcc_dir + mfcc_names[i])\n            mfcc_mean = np.mean(mfcc, axis = 0)\n            mfcc_stddev =  np.std(mfcc, axis = 0)\n            cepstral_norm = (mfcc - mfcc_mean)/mfcc_stddev\n        #   Do Cepstral Normalization of mfcc (explained in writeup)\n        #   Load the corresponding transcript\n            transcript  = np.load(self.transcript_dir + transcript_names[i]) # Remove [SOS] and [EOS] from the transcript\n            transcript = transcript[1: -1]\n            # (Is there an efficient way to do this without traversing through the transcript?)\n            # Note that SOS will always be in the starting and EOS at end, as the name suggests.\n        #   Append each mfcc to self.mfcc, transcript to self.transcript\n            self.mfccs.append(cepstral_norm)\n            # self.mfccsLength.append(len(cepstral_norm))\n            self.transcripts.append(np.array(list(map(lambda x : self.phonemes.index(x), transcript))))\n            #self.mfccs.append(cepstral_norm)\n            #self.transcripts.append(transcript)\n        self.length = len(self.mfccs)\n        self.mfccs = np.array(self.mfccs)\n        self.transcripts = np.array(self.transcripts)\n        # Creating a phoneme to index mapping\n        #self.phoneme_to_index = {phoneme: idx for idx, phoneme in enumerate(self.phonemes)}\n        #self.index_to_phoneme = {idx: phoneme for phoneme, idx in self.phoneme_to_index.items()}\n\n\n\n    def __len__(self):\n        return self.length\n\n    def __getitem__(self, ind):\n        # mfcc = self.mfccs[ind: ind + 2*self.context + 1].flatten() # Get MFCCs with context\n        # transcript = torch.tensor(self.transcripts[ind])\n        # TODO: Based on context and offset, return a frame at given index with context frames to the left, and right.\n        # TODO: Based on context and offset, return a frame at given index with context frames to the left, and right.\n        frames = self.mfccs[ind]\n        # After slicing, you get an array of shape 2*context+1 x 28. But our MLP needs 1d data and not 2d.\n        # frames = frames.flatten() # TODO: Flatten to get 1d data\n\n        frames      = torch.FloatTensor(frames) # Convert to tensors\n        phonemes    = torch.tensor(self.transcripts[ind])\n\n        return frames, phonemes\n\n    def collate_fn(self, batch):\n        # print(*batch)\n        # print(*batch)\n        # print(zip(*batch))\n        batch_mfcc = [batch_item[0] for batch_item in batch]\n        batch_transcript = [batch_item[1] for batch_item in batch]\n        #batch_mfcc, batch_transcript = zip(*batch) (B = batch size, T differs, C = 28)\n\n        # Pad sequences\n        batch_mfcc_pad = pad_sequence(batch_mfcc, batch_first=True)\n        lengths_mfcc = [len(seq) for seq in batch_mfcc]\n\n        batch_transcript_pad = pad_sequence(batch_transcript, batch_first=True)\n        lengths_transcript = [len(seq) for seq in batch_transcript]\n\n        return batch_mfcc_pad, batch_transcript_pad, torch.tensor(lengths_mfcc), torch.tensor(lengths_transcript)","metadata":{"id":"isVjzCnuPeuM","execution":{"iopub.status.busy":"2023-11-10T03:38:38.245579Z","iopub.execute_input":"2023-11-10T03:38:38.246234Z","iopub.status.idle":"2023-11-10T03:38:38.264237Z","shell.execute_reply.started":"2023-11-10T03:38:38.246206Z","shell.execute_reply":"2023-11-10T03:38:38.263341Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Test Data","metadata":{"id":"hqDrxeHfJw4g"}},{"cell_type":"code","source":"class AudioDatasetTest(torch.utils.data.Dataset):\n\n    def __init__(self, root, partition=\"test-clean\", divider=1, context=0, phonemes=PHONEMES):\n        '''\n        Initializes the dataset.\n\n        INPUTS: What inputs do you need here?\n        '''\n        self.context = context\n        self.divider = divider\n        self.phonemes = phonemes\n\n        # Load the directory and all files in them\n        self.mfcc_dir = os.path.join(root, partition, 'mfcc/')\n\n        self.mfcc_files = sorted(os.listdir(self.mfcc_dir))\n\n        # assert len(self.mfcc_files) == len(self.transcript_files)\n        # TODO: List files in sefl.mfcc_dir using os.listdir in sorted order\n        mfcc_names          = sorted(os.listdir(self.mfcc_dir))\n        # TODO: List files in self.transcript_dir using os.listdir in sorted order\n        self.mfccs = []\n        # self.mfccsLength = []\n        '''\n        # Calculate the dataset length\n        total_timestamps = 0\n        for i in range(len(self.mfcc_files)//self.divider):\n            mfcc_file = self.mfcc_files[i]\n            total_timestamps += len(np.load(os.path.join(self.mfcc_dir, mfcc_file)))\n        self.length = total_timestamps + 2 * context * (len(self.mfcc_files)//self.divider)\n        '''\n        for i in range(len(mfcc_names)//divider):\n        #   Load a single mfcc\n            mfcc        = np.load(self.mfcc_dir + mfcc_names[i])\n            mfcc_mean = np.mean(mfcc, axis = 0)\n            mfcc_stddev =  np.std(mfcc, axis = 0)\n            cepstral_norm = (mfcc - mfcc_mean)/mfcc_stddev\n        #   Do Cepstral Normalization of mfcc (explained in writeup)\n        #   Load the corresponding transcript\n            # (Is there an efficient way to do this without traversing through the transcript?)\n            # Note that SOS will always be in the starting and EOS at end, as the name suggests.\n        #   Append each mfcc to self.mfcc, transcript to self.transcript\n            self.mfccs.append(cepstral_norm)\n            # self.mfccsLength.append(len(cepstral_norm))\n            #self.mfccs.append(cepstral_norm)\n            #self.transcripts.append(transcript)\n        self.length = len(self.mfccs)\n        self.mfccs = np.array(self.mfccs)\n        # Creating a phoneme to index mapping\n        #self.phoneme_to_index = {phoneme: idx for idx, phoneme in enumerate(self.phonemes)}\n        #self.index_to_phoneme = {idx: phoneme for phoneme, idx in self.phoneme_to_index.items()}\n\n\n\n    def __len__(self):\n        return self.length\n\n    def __getitem__(self, ind):\n        # mfcc = self.mfccs[ind: ind + 2*self.context + 1].flatten() # Get MFCCs with context\n        # transcript = torch.tensor(self.transcripts[ind])\n        # TODO: Based on context and offset, return a frame at given index with context frames to the left, and right.\n        # TODO: Based on context and offset, return a frame at given index with context frames to the left, and right.\n        frames = self.mfccs[ind]\n        # After slicing, you get an array of shape 2*context+1 x 28. But our MLP needs 1d data and not 2d.\n        # frames = frames.flatten() # TODO: Flatten to get 1d data\n\n        frames      = torch.FloatTensor(frames) # Convert to tensors\n\n        return frames\n\n    def collate_fn(self, batch):\n        batch_mfcc = batch\n\n        # Pad sequences\n        batch_mfcc_pad = pad_sequence(batch_mfcc, batch_first=True)\n        lengths_mfcc = [len(seq) for seq in batch_mfcc]\n\n\n        return batch_mfcc_pad, torch.tensor(lengths_mfcc)","metadata":{"id":"HrLS1wfVJppA","execution":{"iopub.status.busy":"2023-11-10T03:38:38.265711Z","iopub.execute_input":"2023-11-10T03:38:38.266057Z","iopub.status.idle":"2023-11-10T03:38:38.280492Z","shell.execute_reply.started":"2023-11-10T03:38:38.266025Z","shell.execute_reply":"2023-11-10T03:38:38.279605Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Config - Hyperparameters","metadata":{"id":"Pt-veYcdL6Fe"}},{"cell_type":"code","source":"root = '/kaggle/input/automatic-speech-recognition-asr/11-785-f23-hw3p2/'\n\n# Feel free to add more items here\nconfig = {\n    \"beam_width\" : 3,\n    \"lr\"         : 0.002,\n    \"epochs\"     : 150,\n    \"batch_size\" : 64  # Increase if your device can handle it\n}\n\n# You may pass this as a parameter to the dataset class above\n# This will help modularize your implementation\ntransforms = [] # set of tranformations","metadata":{"id":"MN82c3KpLup8","execution":{"iopub.status.busy":"2023-11-10T03:52:48.390799Z","iopub.execute_input":"2023-11-10T03:52:48.391176Z","iopub.status.idle":"2023-11-10T03:52:48.396465Z","shell.execute_reply.started":"2023-11-10T03:52:48.391145Z","shell.execute_reply":"2023-11-10T03:52:48.39552Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Data loaders","metadata":{"id":"NmuPk9J6L8dz"}},{"cell_type":"code","source":"# get me RAMMM!!!!\nimport gc\ngc.collect()","metadata":{"id":"3_kG0gU2x4hH","outputId":"0641b747-8c15-40bd-e337-b62a633cc0a8","execution":{"iopub.status.busy":"2023-11-10T03:52:50.137455Z","iopub.execute_input":"2023-11-10T03:52:50.138568Z","iopub.status.idle":"2023-11-10T03:52:50.313554Z","shell.execute_reply.started":"2023-11-10T03:52:50.138527Z","shell.execute_reply":"2023-11-10T03:52:50.312533Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import DataLoader\nroot_directory = '/kaggle/input/automatic-speech-recognition-asr/11-785-f23-hw3p2/'\n\n\n# Create objects for the dataset class\ntrain_data = AudioDataset(root=root_directory, partition=\"train-clean-100\") #TODO: Add other necessary arguments if needed\nval_data = AudioDataset(root=root_directory, partition=\"dev-clean\") # TODO: Add other necessary arguments if needed\ntest_data = AudioDatasetTest(root=root_directory, partition=\"test-clean\") #TODO: Add other necessary arguments if needed\n\n# Do NOT forget to pass in the collate function as parameter while creating the dataloader\ntrain_loader = DataLoader(train_data, batch_size=config['batch_size'], shuffle=True, collate_fn=train_data.collate_fn, num_workers = 12, pin_memory  = True)\nval_loader = DataLoader(val_data, batch_size=config['batch_size'], shuffle=False, collate_fn=train_data.collate_fn, num_workers = 12)  # usually we don't shuffle validation data\ntest_loader = DataLoader(test_data, batch_size=config['batch_size'], shuffle=False, collate_fn=test_data.collate_fn, num_workers = 12)  # usually we don't shuffle test data\n\nprint(\"Batch size: \", config['batch_size'])\nprint(\"Train dataset samples = {}, batches = {}\".format(len(train_data), len(train_loader)))\nprint(\"Val dataset samples = {}, batches = {}\".format(len(val_data), len(val_loader)))\nprint(\"Test dataset samples = {}, batches = {}\".format(len(test_data), len(test_loader)))\n","metadata":{"id":"4z9_8i6zQnnj","outputId":"92b8bd83-ca47-4aa6-abb7-b56efaa0db77","execution":{"iopub.status.busy":"2023-11-10T03:38:38.405962Z","iopub.execute_input":"2023-11-10T03:38:38.406562Z","iopub.status.idle":"2023-11-10T03:43:36.590284Z","shell.execute_reply.started":"2023-11-10T03:38:38.406525Z","shell.execute_reply":"2023-11-10T03:43:36.58928Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(type(train_data))","metadata":{"id":"zZM-0EGrbu7A"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# sanity check\ncounter = 0\nx_shapes = []\nfor data in train_loader:\n    x, y, lx, ly = data\n    x_shapes.append(x.shape)\n    counter += 1\n    if counter >= 5:\n        break\n\nprint(x_shapes)","metadata":{"id":"cXMtwyviKaxK","outputId":"7c552afe-0054-4cb6-ee47-19430f5a4ff5","execution":{"iopub.status.busy":"2023-11-10T03:52:55.913058Z","iopub.execute_input":"2023-11-10T03:52:55.913427Z","iopub.status.idle":"2023-11-10T03:52:57.775138Z","shell.execute_reply.started":"2023-11-10T03:52:55.913396Z","shell.execute_reply":"2023-11-10T03:52:57.774025Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# NETWORK","metadata":{"id":"wSexxhdfMUzx"}},{"cell_type":"markdown","source":"## Basic\n\nThis is a basic block for understanding, you can skip this and move to pBLSTM one","metadata":{"id":"HLad4pChcuvX"}},{"cell_type":"code","source":"torch.cuda.empty_cache()\n\nclass Network(nn.Module):\n\n    def __init__(self, input_size, num_classes):\n        \"\"\"\n        input_size: size of the input feature vector\n        num_classes: number of output classes for classification\n        \"\"\"\n        super(Network, self).__init__()\n\n        hidden_size = 256\n\n        # LSTM layer\n        self.lstm = nn.LSTM(input_size=28,\n                            hidden_size=hidden_size,\n                            num_layers=2,\n                            bidirectional=True)\n\n        # Classification layer\n        # For bidirectional LSTM, the output dimension is 2 * hidden_size\n        self.classification = nn.Linear(in_features=2 * hidden_size, out_features=num_classes)\n\n        # LogSoftmax layer\n        # The log softmax is applied on the last dimension where the class scores are\n        self.act = nn.LogSoftmax(dim=-1)\n\n    def forward(self, x, lx):\n        \"\"\"\n        x: the input tensor\n        lx: lengths of the sequences in the batch (useful for packing)\n        \"\"\"\n        # Packing sequences\n        x = nn.utils.rnn.pack_padded_sequence(x, lx, batch_first=True, enforce_sorted=False)\n\n        # LSTM layer\n        lstm_out, _ = self.lstm(x)\n\n        # Unpacking sequences\n        lstm_out, lh = nn.utils.rnn.pad_packed_sequence(lstm_out, batch_first=True)\n\n        # Classification layer\n        output = self.classification(lstm_out)\n\n        # Applying LogSoftmax\n        output = self.act(output)\n\n        return output, lh\n","metadata":{"id":"ak6CuQP-57aH","execution":{"iopub.status.busy":"2023-11-10T03:52:58.995087Z","iopub.execute_input":"2023-11-10T03:52:58.995468Z","iopub.status.idle":"2023-11-10T03:52:59.035887Z","shell.execute_reply.started":"2023-11-10T03:52:58.995435Z","shell.execute_reply":"2023-11-10T03:52:59.034951Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = Network(num_classes = 41, input_size = 28).to(device)","metadata":{"id":"FJ8L8oZjZEe9"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Initialize Basic Network\n(If trying out the basic Network)","metadata":{"id":"tUThsowyQdN7"}},{"cell_type":"markdown","source":"## ASR Network","metadata":{"id":"e-qb7wnAzCZl"}},{"cell_type":"markdown","source":"### Pyramid Bi-LSTM (pBLSTM)","metadata":{"id":"PB6eh3gnMUzy"}},{"cell_type":"code","source":"# Utils for network\ntorch.cuda.empty_cache()\n\nclass PermuteBlock(torch.nn.Module):\n    def forward(self, x):\n        return x.transpose(1, 2)","metadata":{"id":"qd4BEX_yMUzz","execution":{"iopub.status.busy":"2023-11-10T03:53:02.785526Z","iopub.execute_input":"2023-11-10T03:53:02.785922Z","iopub.status.idle":"2023-11-10T03:53:02.791533Z","shell.execute_reply.started":"2023-11-10T03:53:02.785888Z","shell.execute_reply":"2023-11-10T03:53:02.790529Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Encoder","metadata":{"id":"g3ZQ75OcMUz0"}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n\nclass pBLSTM(nn.Module):\n    def __init__(self, input_size, hidden_size):\n        super(pBLSTM, self).__init__()\n        self.blstm = nn.LSTM(input_size=input_size * 2, hidden_size=hidden_size, num_layers=1, bidirectional=True)\n\n    def forward(self, x_packed):\n        x, x_lens = torch.nn.utils.rnn.pad_packed_sequence(x_packed, batch_first=True)\n        x, x_lens = self.trunc_reshape(x, x_lens)\n        x_packed = torch.nn.utils.rnn.pack_padded_sequence(x, x_lens, batch_first=True, enforce_sorted=False)\n        output_packed, (h_n, c_n) = self.blstm(x_packed)\n        return output_packed, x_lens\n\n    def trunc_reshape(self, x, x_lens):\n        if x.size(1) % 2 != 0:\n            x = x[:, :-1, :]\n            x_lens = x_lens - (x_lens % 2 > 0).int()\n        batch_size, seq_len, feature_dim = x.size()\n        x = x.reshape(batch_size, seq_len // 2, feature_dim * 2)\n        x_lens = x_lens // 2\n        return x, x_lens\n\nclass PermuteBlock(nn.Module):\n    def forward(self, x):\n        return x.transpose(1, 2)\n\nclass Network(nn.Module):\n    def __init__(self, num_embeddings, embedding_dim):\n        super(Network, self).__init__()\n\n        # Embedding layer for categorical input\n\n        self.permuteBlock = PermuteBlock()\n        print(embedding_dim)\n\n        self.embedding0 = nn.Conv1d(in_channels=28,\n                              out_channels=embedding_dim//2,\n                              kernel_size=1,\n                              stride=1)\n\n        self.embedding1 = nn.Conv1d(in_channels=embedding_dim//2,\n                              out_channels=embedding_dim//1,\n                              kernel_size=1,\n                              stride=1)\n\n        # self.embedding2 = nn.Conv1d(in_channels=embedding_dim//2,\n        #                       out_channels=embedding_dim,\n        #                       kernel_size=1,\n        #                       stride=1)\n\n\n\n        hidden_size = 512\n        self.num_layers = 2 # Define the number of pBLSTM layers you want in your network\n\n        # The first LSTM layer\n        self.lstm = nn.LSTM(input_size=embedding_dim,\n                            hidden_size=hidden_size,\n                            num_layers=4,\n                            bidirectional=True)\n\n        # Subsequent pBLSTM layers\n        self.pblstm_layers = nn.ModuleList([\n            pBLSTM(input_size=hidden_size*2, hidden_size=hidden_size) for _ in range(self.num_layers)\n        ])\n\n    def forward(self, x, lx):\n\n        x = self.permuteBlock(x)\n\n        x = self.embedding0(x)\n        x = self.embedding1(x)\n        #x = self.embedding2(x)\n\n        x = self.permuteBlock(x)\n\n        # Packing sequences\n        x_packed = nn.utils.rnn.pack_padded_sequence(x, lx, batch_first=True, enforce_sorted=False)\n\n        # LSTM layer\n        x_packed, _ = self.lstm(x_packed)\n\n        # pBLSTM layers\n        for layer in self.pblstm_layers:\n            x_packed, lx = layer(x_packed)\n\n        # Unpacking sequences\n        x, _ = nn.utils.rnn.pad_packed_sequence(x_packed, batch_first=True)\n\n\n\n        return x, lx\n\n\n# Now, you can use the `network` instance to train on your data.\n\n\nclass Encoder(nn.Module):\n    '''\n    The Encoder takes utterances as inputs and returns latent feature representations\n    '''\n    def __init__(self, num_embeddings, embedding_dim):\n        super(Encoder, self).__init__()\n\n\n        #self.embedding = #TODO: You can use CNNs as Embedding layer to extract features. Keep in mind the Input dimensions and expected dimension of Pytorch CNN.\n\n        self.layers = torch.nn.Sequential( # How many pBLSTMs are required?\n            # TODO: Fill this up with pBLSTMs - What should the input_size be?\n            # Hint: You are downsampling timesteps by a factor of 2, upsampling features by a factor of 2 and the LSTM is bidirectional)\n            # Optional: Dropout/Locked Dropout after each pBLSTM (Not needed for early submission)\n            # https://github.com/salesforce/awd-lstm-lm/blob/dfd3cb0235d2caf2847a4d53e1cbd495b781b5d2/locked_dropout.py#L5\n            # ...\n            # ...\n            Network(num_embeddings, embedding_dim)\n\n        )\n\n    def forward(self, x, x_lens):\n        # Where are x and x_lens coming from? The dataloader\n        #TODO: Call the embedding layer\n        # TODO: Pack Padded Sequence\n        # TODO: Pass Sequence through the pyramidal Bi-LSTM layer\n        # TODO: Pad Packed Sequence\n\n        for layer in self.layers:\n            x, x_lens = layer.forward(x, x_lens)\n\n\n        # Remember the number of output(s) each function returns\n\n        return x, x_lens","metadata":{"id":"GEzw5_xmMUz0","execution":{"iopub.status.busy":"2023-11-10T03:53:04.834487Z","iopub.execute_input":"2023-11-10T03:53:04.834891Z","iopub.status.idle":"2023-11-10T03:53:04.855001Z","shell.execute_reply.started":"2023-11-10T03:53:04.834856Z","shell.execute_reply":"2023-11-10T03:53:04.854037Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Decoder","metadata":{"id":"kg82HXa3MUz1"}},{"cell_type":"code","source":"# This architecture will make you cross the very low cutoff\n# However, you need to run a lot of experiments to cross the medium or high cutoff\nclass MLP(torch.nn.Module):\n\n    def __init__(self, input_size, output_size):\n\n        super(MLP, self).__init__()\n\n        self.model = torch.nn.Sequential(\n\n\n\n            torch.nn.Linear(input_size, 1024),\n            torch.nn.GELU(), #logged softmax only for probs\n            torch.nn.Dropout(p=0.25),\n\n            torch.nn.Linear(1024, 2048),\n            PermuteBlock(),\n            torch.nn.BatchNorm1d(2048),\n            PermuteBlock(),\n            torch.nn.GELU(),\n            torch.nn.Dropout(p=0.25),\n\n\n            torch.nn.Linear(2048, 1024),\n            torch.nn.GELU(),\n            torch.nn.Dropout(p=0.25),\n\n            torch.nn.Linear(1024, output_size),\n        )\n\n    def forward(self, x):\n        out = self.model(x)\n\n        return out\n\nclass Decoder(nn.Module): #specify log probability\n\n    def __init__(self, embedding_dim, output_size= 41):\n        super().__init__()\n        print(embedding_dim)\n\n        self.mlp = torch.nn.Sequential(\n            PermuteBlock(), torch.nn.BatchNorm1d(1024), PermuteBlock(), #1024 is hidden size *2\n            MLP(input_size=1024, output_size = output_size) #because LSTM bidirectional and pBLSTM\n            #TODO define your MLP arch. Refer HW1P2\n            #Use Permute Block before and after BatchNorm1d() to match the size\n        )\n\n        self.softmax = torch.nn.LogSoftmax(dim=2)\n\n    def forward(self, encoder_out):\n        print(encoder_out.shape)\n\n        encoder_out = self.mlp(encoder_out)\n        encoder_out = self.softmax(encoder_out)\n\n        return encoder_out","metadata":{"id":"PQIRxdNTMUz1","execution":{"iopub.status.busy":"2023-11-10T03:53:09.118019Z","iopub.execute_input":"2023-11-10T03:53:09.119042Z","iopub.status.idle":"2023-11-10T03:53:09.129912Z","shell.execute_reply.started":"2023-11-10T03:53:09.119006Z","shell.execute_reply":"2023-11-10T03:53:09.129019Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ASRModel(torch.nn.Module):\n\n    def __init__(self, input_size, embed_size= 128, output_size= len(PHONEMES)):\n        super().__init__()\n\n        self.augmentations  = torch.nn.Sequential(\n            #TODO Add Time Masking/ Frequency Masking\n            #Hint: See how to use PermuteBlock() function defined above\n        )\n        self.encoder        = Encoder(num_embeddings=input_size, embedding_dim=embed_size)\n        self.decoder        = Decoder(embedding_dim=embed_size, output_size=num_classes)\n\n      # encoder self, num_embeddings, embedding_dim, input_size):\n\n    # def __init__(self, num_embeddings, embedding_dim, input_size, num_classe):\n    def forward(self, x, lengths_x):\n\n        if self.training:\n            x = self.augmentations(x)\n\n        encoder_out, encoder_lens   = self.encoder(x, lengths_x)\n        decoder_out                 = self.decoder(encoder_out)\n\n        return decoder_out, encoder_lens","metadata":{"id":"qmHf6pFiMUz1","execution":{"iopub.status.busy":"2023-11-10T03:53:12.197766Z","iopub.execute_input":"2023-11-10T03:53:12.198378Z","iopub.status.idle":"2023-11-10T03:53:12.205741Z","shell.execute_reply.started":"2023-11-10T03:53:12.198346Z","shell.execute_reply":"2023-11-10T03:53:12.204626Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.empty_cache()","metadata":{"id":"_LmiwS3rAaby","execution":{"iopub.status.busy":"2023-11-10T03:53:14.54088Z","iopub.execute_input":"2023-11-10T03:53:14.541332Z","iopub.status.idle":"2023-11-10T03:53:14.546114Z","shell.execute_reply.started":"2023-11-10T03:53:14.541294Z","shell.execute_reply":"2023-11-10T03:53:14.545081Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#torch.cuda.empty_cache()\n# THIS ONE\n\n\ninput_size = 28  # This should match the feature size of your MFCC input\nnum_classes = 41  # Example, replace with actual number of phonemes/classes\nembedding_dim = 128 # From video????\n\nmodel = ASRModel(input_size = input_size, embed_size = embedding_dim, output_size=41).to(device)\n#summary(model, x.to(device), lx) # x and lx come from the sanity check above :)","metadata":{"id":"maX-7zWqTAUK","outputId":"7c51db2f-a7db-4ffe-f48c-f26619c35776","execution":{"iopub.status.busy":"2023-11-10T03:56:28.46612Z","iopub.execute_input":"2023-11-10T03:56:28.46648Z","iopub.status.idle":"2023-11-10T03:56:28.877534Z","shell.execute_reply.started":"2023-11-10T03:56:28.466451Z","shell.execute_reply":"2023-11-10T03:56:28.876474Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training Config\nInitialize Loss Criterion, Optimizer, CTC Beam Decoder, Scheduler, Scaler (Mixed-Precision), etc.","metadata":{"id":"IBwunYpyugFg"}},{"cell_type":"code","source":"from torch.optim.lr_scheduler import ReduceLROnPlateau","metadata":{"id":"Z7ZnjHz2EsYH","execution":{"iopub.status.busy":"2023-11-10T03:56:33.496919Z","iopub.execute_input":"2023-11-10T03:56:33.497845Z","iopub.status.idle":"2023-11-10T03:56:33.501799Z","shell.execute_reply.started":"2023-11-10T03:56:33.497805Z","shell.execute_reply":"2023-11-10T03:56:33.500802Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#TODO\n\n\n# criterion = # Define CTC loss as the criterion. How would the losses be reduced?\n# CTC Loss: https://pytorch.org/docs/stable/generated/torch.nn.CTCLoss.html\n# Refer to the handout for hints\n\ncriterion = nn.CTCLoss(blank=0)\n\n\noptimizer = torch.optim.AdamW(model.parameters(), lr=config['lr'])\n\n# Declare the decoder. Use the CTC Beam Decoder to decode phonemes\n# CTC Beam Decoder Doc: https://github.com/parlance/ctcdecode\n# CTC Beam Decoder\nlabels = PHONEMES  # your list of phonemes\ndecoder = CTCBeamDecoder(labels=labels, beam_width=config['beam_width'], blank_id=0, log_probs_input=True)\n\nstep_size = 10  # number of epochs after which the lr is multiplied by gamma\ngamma = 0.1  # multiplicative factor of lr after each step_size epochs\nscheduler = ReduceLROnPlateau(optimizer, 'min', patience=3, factor=0.75, min_lr=0.00001)\n\n# Mixed Precision, if you need it\nscaler = torch.cuda.amp.GradScaler()","metadata":{"id":"iGoozH2nd6KB","execution":{"iopub.status.busy":"2023-11-10T03:56:34.601629Z","iopub.execute_input":"2023-11-10T03:56:34.602435Z","iopub.status.idle":"2023-11-10T03:56:34.609846Z","shell.execute_reply.started":"2023-11-10T03:56:34.602404Z","shell.execute_reply":"2023-11-10T03:56:34.608862Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Decode Prediction","metadata":{"id":"Jmc6_4eWL2Xp"}},{"cell_type":"code","source":"def decode_prediction(output, output_lens, decoder, PHONEME_MAP=LABELS):\n    # output = torch.permute(output, (1, 0, 2))\n\n    beam_results, beam_scores, timesteps, out_seq_len = decoder.decode(output, seq_lens=output_lens)\n\n    pred_strings = []\n\n    for i in range(out_seq_len.shape[0]):\n        pred = beam_results[i][0][:out_seq_len[i][0]] # correct?\n        #pred_string = [PHONEME_MAP[i] for i in pred.tolist()]\n        pred_string = [PHONEME_MAP[i] for i in pred]\n        #Note string vs strings\n        pred_strings.append(''.join(pred_string)) #should I be joining?\n\n    return pred_strings\n\n# print(calculate_levenshtein(h, y, lx, ly, decoder, LABELS))\n\n\ndef calculate_levenshtein(output, label, output_lens, label_lens, decoder, PHONEME_MAP=LABELS):\n\n    dist = 0\n    batch_size = label.shape[0]\n\n    pred_strings = decode_prediction(output, output_lens, decoder, PHONEME_MAP)\n\n    for i in range(batch_size):\n        pred_string = pred_strings[i]\n        #label_string = ''.join([PHONEME_MAP[i] for i in label[i][:label_lens[i]]])\n        label_string = ''.join([PHONEME_MAP[i] for i in label[i][:label_lens[i]]])\n        #label_string = [PHONEME_MAP[int(i)] for i in label[i][:label_lens[i]]]\n\n        dist += Levenshtein.distance(pred_string, label_string)\n\n    dist /= batch_size # Averaging the Levenshtein distance over the batch size\n\n    return dist\n","metadata":{"id":"AWj86UcDAcz7","execution":{"iopub.status.busy":"2023-11-10T03:56:36.677782Z","iopub.execute_input":"2023-11-10T03:56:36.678647Z","iopub.status.idle":"2023-11-10T03:56:36.687393Z","shell.execute_reply.started":"2023-11-10T03:56:36.678614Z","shell.execute_reply":"2023-11-10T03:56:36.686399Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Test Implementation","metadata":{"id":"0Qk9iZud1LXT"}},{"cell_type":"code","source":"# test code to check shapes\n\nmodel.eval()\nfor i, data in enumerate(val_loader, 0):\n    x, y, lx, ly = data\n    x, y = x.to(device), y.to(device)\n    h, lh = model(x, lx)\n\n    print(h.shape)\n    print(calculate_levenshtein(h, y, lx, ly, decoder, LABELS))\n\n    h = torch.permute(h, (1, 0, 2))\n    print(h.shape, y.shape)\n    loss = criterion(h, y, lh, ly)\n    print(loss)\n    print(lh)\n\n\n\n\n    break","metadata":{"id":"GnTLL-5gMBrY","outputId":"285b2d39-37b6-46ca-a9f0-dab47c7eb6fc","execution":{"iopub.status.busy":"2023-11-10T03:56:39.153442Z","iopub.execute_input":"2023-11-10T03:56:39.153821Z","iopub.status.idle":"2023-11-10T03:56:42.395433Z","shell.execute_reply.started":"2023-11-10T03:56:39.153791Z","shell.execute_reply":"2023-11-10T03:56:42.394175Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# WandB\n\nYou will need to fetch your api key from wandb.ai","metadata":{"id":"rd5aNaLVoR_g"}},{"cell_type":"code","source":"import wandb\nwandb.login(key=\"b78f5a0a316b8f09525ad5c671b527f4b95c1695\") #API Key is in your wandb account, under settings (wandb.ai/settings)","metadata":{"id":"PiDduMaDIARE","outputId":"f80cf04e-9ff8-4811-f047-a6403dcedeee","execution":{"iopub.status.busy":"2023-11-10T03:56:44.425665Z","iopub.execute_input":"2023-11-10T03:56:44.426091Z","iopub.status.idle":"2023-11-10T03:56:44.732583Z","shell.execute_reply.started":"2023-11-10T03:56:44.426052Z","shell.execute_reply":"2023-11-10T03:56:44.731487Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create your wandb run\nrun = wandb.init(\n    name = \"highcutoff1\", ## Wandb creates random run names if you skip this field\n    #reinit = True, ### Allows reinitalizing runs when you re-run this cell\n    id = \"n21m0kcv\", ### Insert specific run id here if you want to resume a previous run\n    resume = \"must\", ### You need this to resume previous runs, but comment out reinit = True when using this\n    project = \"hw3p2-ablations\", ### Project should be created in your wandb account\n    config = config ### Wandb Config for your run\n)","metadata":{"id":"NvjDy-fneTkX","outputId":"9b1a8dcc-f1bd-4ac0-fa28-2068b5693e0f"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train Functions","metadata":{"id":"6fLLj5KIMMOe"}},{"cell_type":"code","source":"from tqdm import tqdm\n\ndef train_model(model, train_loader, criterion, optimizer):\n\n    model.train()\n    batch_bar = tqdm(total=len(train_loader), dynamic_ncols=True, leave=False, position=0, desc='Train')\n\n    total_loss = 0\n\n    for i, data in enumerate(train_loader):\n        optimizer.zero_grad()\n\n        x, y, lx, ly = data\n        x, y = x.to(device), y.to(device)\n\n        with torch.cuda.amp.autocast():\n            h, lh = model(x, lx)\n            h = torch.permute(h, (1, 0, 2))\n            loss = criterion(h, y, lh, ly)\n\n        total_loss += loss.item()\n\n        batch_bar.set_postfix(\n            loss=\"{:.04f}\".format(float(total_loss / (i + 1))),\n            lr=\"{:.06f}\".format(float(optimizer.param_groups[0]['lr'])))\n\n        batch_bar.update() # Update tqdm bar\n\n        # Another couple things you need for FP16.\n        scaler.scale(loss).backward() # This is a replacement for loss.backward()\n        scaler.step(optimizer) # This is a replacement for optimizer.step()\n        scaler.update() # This is something added just for FP16\n\n        ### Backward Propagation\n        # loss.backward()\n\n        ### Gradient Descent\n        # optimizer.step()\n\n        del x, y, lx, ly, h, lh, loss\n        torch.cuda.empty_cache()\n\n    batch_bar.close() # You need this to close the tqdm bar\n\n    return total_loss / len(train_loader)\n\n\ndef validate_model(model, val_loader, decoder, phoneme_map= LABELS):\n\n    model.eval()\n    batch_bar = tqdm(total=len(val_loader), dynamic_ncols=True, position=0, leave=False, desc='Val')\n\n    total_loss = 0\n    vdist = 0\n\n    for i, data in enumerate(val_loader):\n\n        x, y, lx, ly = data\n        x, y = x.to(device), y.to(device)\n\n        with torch.inference_mode():\n            h, lh = model(x, lx)\n            h = torch.permute(h, (1, 0, 2))\n            loss = criterion(h, y, lh, ly)\n\n        total_loss += float(loss)\n        vdist += calculate_levenshtein(torch.permute(h, (1, 0, 2)), y, lh, ly, decoder, phoneme_map)\n\n        batch_bar.set_postfix(loss=\"{:.04f}\".format(float(total_loss / (i + 1))), dist=\"{:.04f}\".format(float(vdist / (i + 1))))\n\n        batch_bar.update()\n\n        del x, y, lx, ly, h, lh, loss\n        torch.cuda.empty_cache()\n\n    batch_bar.close()\n    total_loss = total_loss/len(val_loader)\n    val_dist = vdist/len(val_loader)\n    return total_loss, val_dist","metadata":{"id":"ri87MAdhMUz5","execution":{"iopub.status.busy":"2023-11-10T03:56:47.290344Z","iopub.execute_input":"2023-11-10T03:56:47.290737Z","iopub.status.idle":"2023-11-10T03:56:47.306378Z","shell.execute_reply.started":"2023-11-10T03:56:47.290701Z","shell.execute_reply":"2023-11-10T03:56:47.305407Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training Setup","metadata":{"id":"qpYExu4vT4_g"}},{"cell_type":"code","source":"def save_model(model, optimizer, 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         metric[0]                  : metric[1],\n         'epoch'                    : epoch},\n         path\n    )\n\ndef load_model(path, model, metric= 'valid_acc', optimizer= None, scheduler= None):\n\n    checkpoint = torch.load(path)\n    model.load_state_dict(checkpoint['model_state_dict'])\n\n    if optimizer != None:\n        optimizer.load_state_dict(checkpoint['optimizer_state_dict'])\n    if scheduler != None:\n        scheduler.load_state_dict(checkpoint['scheduler_state_dict'])\n\n    epoch   = checkpoint['epoch']\n\n    return [model, optimizer, scheduler, epoch, metric]","metadata":{"id":"husa5_EYMUz6","execution":{"iopub.status.busy":"2023-11-10T03:53:51.718268Z","iopub.execute_input":"2023-11-10T03:53:51.719139Z","iopub.status.idle":"2023-11-10T03:53:51.726493Z","shell.execute_reply.started":"2023-11-10T03:53:51.719107Z","shell.execute_reply":"2023-11-10T03:53:51.725458Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.load_state_dict(torch.load('/kaggle/input/modelhigh3/Model (4)')['model_state_dict'])\noptimizer.load_state_dict(torch.load('/kaggle/input/modelhigh3/Model (4)')['optimizer_state_dict'])\nscheduler.load_state_dict(torch.load('/kaggle/input/modelhigh3/Model (4)')['scheduler_state_dict'])\nvalidate_model(model, val_loader, decoder, phoneme_map=LABELS)","metadata":{"execution":{"iopub.status.busy":"2023-11-10T03:57:04.614366Z","iopub.execute_input":"2023-11-10T03:57:04.614732Z","iopub.status.idle":"2023-11-10T03:57:05.045492Z","shell.execute_reply.started":"2023-11-10T03:57:04.614701Z","shell.execute_reply":"2023-11-10T03:57:05.044513Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# This is for checkpointing, if you're doing it over multiple sessions\n\nlast_epoch_completed = 0\nstart = last_epoch_completed\nend = config[\"epochs\"]\nbest_lev_dist = float(\"inf\") # if you're restarting from some checkpoint, use what you saw there.\n# epoch_model_path = #TODO set the model path( Optional, you can just store best one. Make sure to make the changes below )\n# best_model_path = #TODO set best model path","metadata":{"id":"tExvyl1BIdMC"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.empty_cache()\ngc.collect()\n\n#TODO: Please complete the training loop\n\nfor epoch in range(0, config['epochs']):\n\n    print(\"\\nEpoch: {}/{}\".format(epoch+1, config['epochs']))\n\n    curr_lr = optimizer.param_groups[0]['lr']\n\n    train_loss              = train_model(model, train_loader, criterion, optimizer)\n    valid_loss, valid_dist  = validate_model(model, val_loader, decoder, phoneme_map=LABELS)\n    scheduler.step(valid_dist)\n\n    print(\"\\tTrain Loss {:.04f}\\t Learning Rate {:.07f}\".format(train_loss, curr_lr))\n    print(\"\\tVal Dist {:.04f}%\\t Val Loss {:.04f}\".format(valid_dist, valid_loss))\n\n\n    wandb.log({\n        'train_loss': train_loss,\n        'valid_dist': valid_dist,\n        'valid_loss': valid_loss,\n        'lr'        : curr_lr\n    })\n\n    # save_model(model, optimizer, scheduler, ['valid_dist', valid_dist], epoch, epoch_model_path)\n    # wandb.save(epoch_model_path)\n    # print(\"Saved epoch model\")\n\n    if valid_dist <= best_lev_dist:\n        best_lev_dist = valid_dist\n        save_model(model, optimizer, scheduler, ['valid_dist', valid_dist], epoch, 'Model')\n        model_artifact = wandb.Artifact('V4pBLSM', type='model')\n        model_artifact.add_file('Model')\n        run.log_artifact(model_artifact)\n        wandb.save('Model')\n        print(\"Saved best model\")\nrun.finish()","metadata":{"id":"JR43E28rM9Ak","outputId":"36a4d3de-5147-48f7-f02c-999e6d771113"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Generate Predictions and Submit to Kaggle","metadata":{"id":"M2H4EEj-sD32"}},{"cell_type":"code","source":"validate_model(model, val_loader, decoder, phoneme_map=LABELS)","metadata":{"id":"FpRqu_X8dYum","outputId":"3e6e7000-166c-447e-8b57-09d04d690a1b","execution":{"iopub.status.busy":"2023-11-10T03:57:14.295845Z","iopub.execute_input":"2023-11-10T03:57:14.296217Z","iopub.status.idle":"2023-11-10T03:58:08.008509Z","shell.execute_reply.started":"2023-11-10T03:57:14.296186Z","shell.execute_reply":"2023-11-10T03:58:08.007451Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#TODO: Make predictions\n\n# Follow the steps below:\n# 1. Create a new object for CTCBeamDecoder with larger (why?) number of beams\n# 2. Get prediction string by decoding the results of the beam decoder\n\nTEST_BEAM_WIDTH = config['beam_width']\n\ntest_decoder    = decoder\nresults = []\n\nmodel.eval()\nprint(\"Testing\")\nfor data in tqdm(test_loader):\n\n    x, lx   = data\n    x       = x.to(device)\n\n    with torch.no_grad():\n        h, lh = model(x, lx)\n\n    prediction_string = decode_prediction(h, lh, test_decoder, LABELS)\n    results.extend(prediction_string)  # Appending the results\n\n    del x, lx, h, lh\n    torch.cuda.empty_cache()","metadata":{"id":"2moYJhTWsOG-","outputId":"12299858-6091-4259-852c-a30c6e70f22d","execution":{"iopub.status.busy":"2023-11-10T03:58:08.010432Z","iopub.execute_input":"2023-11-10T03:58:08.01076Z","iopub.status.idle":"2023-11-10T03:58:58.521957Z","shell.execute_reply.started":"2023-11-10T03:58:08.010729Z","shell.execute_reply":"2023-11-10T03:58:58.520666Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_dir = f\"{root}/test-clean/random_submission.csv\"\ndf = pd.read_csv(data_dir)\n# print(len(results[0]))\n# print(results[0][111])\ndf.label = results\n# Test dataset samples = 28539, batches = 446\ndf.to_csv('submission.csv', index = False)","metadata":{"id":"d70dvu_lsMlv","execution":{"iopub.status.busy":"2023-11-10T03:58:58.524044Z","iopub.execute_input":"2023-11-10T03:58:58.52444Z","iopub.status.idle":"2023-11-10T03:58:58.584038Z","shell.execute_reply.started":"2023-11-10T03:58:58.524397Z","shell.execute_reply":"2023-11-10T03:58:58.583022Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!kaggle competitions submit -c automatic-speech-recognition-asr -f submission.csv -m \"I made it!\"","metadata":{"id":"m1sZmEIs4yIz"},"execution_count":null,"outputs":[]}]}