{"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":"code","source":"%%capture\n!pip install pyctcdecode\n!python -m pip install pypi-kenlm\n!pip install jiwer","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-08-20T06:38:24.640473Z","iopub.execute_input":"2022-08-20T06:38:24.641409Z","iopub.status.idle":"2022-08-20T06:39:45.195851Z","shell.execute_reply.started":"2022-08-20T06:38:24.641310Z","shell.execute_reply":"2022-08-20T06:39:45.194576Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"![](https://developer-blogs.nvidia.com/wp-content/uploads/2019/12/automatic-speech-recognition_updated.png)\n\n![](https://www.researchgate.net/profile/Diana-Militaru/publication/299594444/figure/fig1/AS:346834426974208@1459703179403/The-block-diagram-of-an-automatic-speech-recognition-and-understanding-system.png)","metadata":{}},{"cell_type":"markdown","source":"in this notebook we will try to demonstrate how to calculate CER,WER metric on validation dataset using xls-r wav2vec2 model,we will be using public best available pretrained model from huggingface to demonstrate the metric calculation process. for understanding how to train wav2vec2 on this dataset please check our past work [wav2vec2 starter](https://www.kaggle.com/code/nazmuddhohaansary/wave2vec2-starter-for-dl-sprint-commonvoice)","metadata":{}},{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"import os\nimport numpy as np\nfrom tqdm.auto import tqdm\nfrom glob import glob\nfrom transformers import AutoFeatureExtractor, pipeline\nimport pandas as pd\nimport librosa\nimport IPython\nfrom datasets import load_metric\nfrom tqdm.auto import tqdm\nfrom torch.utils.data import Dataset, DataLoader\nimport torch\nimport re\nimport gc\nimport wave\nfrom scipy.io import wavfile\nimport scipy.signal as sps\n\ntqdm.pandas()\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\n","metadata":{"execution":{"iopub.status.busy":"2022-08-20T06:39:45.198760Z","iopub.execute_input":"2022-08-20T06:39:45.199227Z","iopub.status.idle":"2022-08-20T06:39:55.501187Z","shell.execute_reply.started":"2022-08-20T06:39:45.199187Z","shell.execute_reply":"2022-08-20T06:39:55.500095Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Configs","metadata":{}},{"cell_type":"code","source":"#according to our experiment this is the best model -> arijitx/wav2vec2-xls-r-300m-bengali\nclass CFG:\n    model_name = 'arijitx/wav2vec2-xls-r-300m-bengali' #arijitx/wav2vec2-large-xlsr-bengali,arijitx/wav2vec2-xls-r-300m-bengali, Tahsin-Mayeesha/wav2vec2-bn-300m\n    valid_df_path = '../input/dlsprint/validation.csv'\n    sample_sub_df_path = '../input/dlsprint/sample_submission.csv'\n    valid = \"../input/dlsprint/validation_files/\"\n    test = \"../input/dlsprint/test_files/\"\n    valid_wav = '../input/validation-fileswav-format/validation_files_wav/'\n    test_wav = '../input/test-wav-files-dl-sprint/test_files_wav/'\n    batch_size = 48#not using this param now\n    single_SPEECH_FILE = \"../input/dlsprint/validation_files/common_voice_bn_30620258.mp3\"\n    \n\n","metadata":{"execution":{"iopub.status.busy":"2022-08-20T06:39:55.503089Z","iopub.execute_input":"2022-08-20T06:39:55.504170Z","iopub.status.idle":"2022-08-20T06:39:55.514826Z","shell.execute_reply.started":"2022-08-20T06:39:55.504130Z","shell.execute_reply":"2022-08-20T06:39:55.512836Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# single sample inference demo","metadata":{}},{"cell_type":"code","source":"asr = pipeline(\"automatic-speech-recognition\", model=CFG.model_name, device=0)\nfeature_extractor = AutoFeatureExtractor.from_pretrained(\n        CFG.model_name, cache_dir=None, use_auth_token=False\n    )\nspeech, sr = librosa.load(CFG.single_SPEECH_FILE, sr=feature_extractor.sampling_rate)\nprediction = asr(\n            speech, chunk_length_s=112, stride_length_s=None\n        )\n\npred = prediction[\"text\"]\npred\n","metadata":{"execution":{"iopub.status.busy":"2022-08-20T06:39:55.521053Z","iopub.execute_input":"2022-08-20T06:39:55.521361Z","iopub.status.idle":"2022-08-20T06:43:14.184951Z","shell.execute_reply.started":"2022-08-20T06:39:55.521334Z","shell.execute_reply":"2022-08-20T06:43:14.183787Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# check the original audio","metadata":{}},{"cell_type":"code","source":"IPython.display.Audio(CFG.single_SPEECH_FILE)","metadata":{"execution":{"iopub.status.busy":"2022-08-20T06:43:14.186571Z","iopub.execute_input":"2022-08-20T06:43:14.188113Z","iopub.status.idle":"2022-08-20T06:43:14.198012Z","shell.execute_reply.started":"2022-08-20T06:43:14.188073Z","shell.execute_reply":"2022-08-20T06:43:14.197039Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Fix paths","metadata":{}},{"cell_type":"code","source":"df = pd.read_csv('../input/dlsprint/validation.csv')\ndirectory =\"../input/dlsprint/validation_files/\"\ndf[\"path\"]=df[\"path\"].progress_apply(lambda x:os.path.join(directory,str(x)))\ndf.head(3)","metadata":{"execution":{"iopub.status.busy":"2022-08-20T06:43:14.200150Z","iopub.execute_input":"2022-08-20T06:43:14.200939Z","iopub.status.idle":"2022-08-20T06:43:14.381626Z","shell.execute_reply.started":"2022-08-20T06:43:14.200865Z","shell.execute_reply":"2022-08-20T06:43:14.380632Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Custom dataset class","metadata":{}},{"cell_type":"markdown","source":"librosa with mp3 is super slow,so we will be using wav files for faster inference","metadata":{}},{"cell_type":"code","source":"class bn_asr_Dataset(Dataset):\n    '''\n    args:\n        df      : path of the dataframe\n        dir     : directory of sound files\n    '''\n    def __init__(self,df,dir):\n        self.df = pd.read_csv(df)\n        self.dir = dir\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, i):\n   \n        #speech, _ = librosa.load(self.dir+self.df.path[i], sr=feature_extractor.sampling_rate) \n        path = self.dir+self.df.path[i]\n        path = os.path.splitext(path)[0]+'.wav'\n        # Read file\n        sampling_rate, data = wavfile.read(path)\n        # Resample data\n        number_of_samples = round(len(data) * float(feature_extractor.sampling_rate) / sampling_rate)\n        speech = sps.resample(data, number_of_samples)\n        return speech\n  \n","metadata":{"execution":{"iopub.status.busy":"2022-08-20T06:43:14.383214Z","iopub.execute_input":"2022-08-20T06:43:14.383593Z","iopub.status.idle":"2022-08-20T06:43:14.391138Z","shell.execute_reply.started":"2022-08-20T06:43:14.383558Z","shell.execute_reply":"2022-08-20T06:43:14.389988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# making prediction on whole validation set","metadata":{}},{"cell_type":"code","source":"%%time\n#single image inference\n''' \n#super slow inference...\n\npredictions = []\nreferences = []\nfor i in range(len(df.path)):\n    speech, sr = librosa.load(df.path[i], sr=feature_extractor.sampling_rate)\n    prediction = asr(speech, chunk_length_s=112, stride_length_s=None)\n    pred = prediction[\"text\"]\n    predictions.append(pred)\n    references.append(df.sentence[i])\n    \nprint(len(predictions),len(references))\n'''\n\ndf = pd.read_csv(CFG.valid_df_path)\nvalid_dataset = bn_asr_Dataset(CFG.valid_df_path,CFG.valid_wav)#CFG.valid\npredictions = []\nreferences = []\n# for i,pred_sentence in enumerate(tqdm(asr(valid_dataset, chunk_length_s=112, stride_length_s=None,batch_size=CFG.batch_size), total=len(valid_dataset))):\n#     references.append(df.sentence[i])\n#     predictions.append(pred_sentence['text'])\n    \nfor i in range(len(valid_dataset)):\n    pred = asr(valid_dataset.__getitem__(i), chunk_length_s=112, stride_length_s=None)\n    references.append(df.sentence[i])\n    predictions.append(pred['text'])\n  ","metadata":{"execution":{"iopub.status.busy":"2022-08-20T06:43:14.392776Z","iopub.execute_input":"2022-08-20T06:43:14.393151Z","iopub.status.idle":"2022-08-20T07:03:44.367081Z","shell.execute_reply.started":"2022-08-20T06:43:14.393116Z","shell.execute_reply":"2022-08-20T07:03:44.366000Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.empty_cache() \ngc.collect()\n!nvidia-smi","metadata":{"execution":{"iopub.status.busy":"2022-08-20T07:03:44.368541Z","iopub.execute_input":"2022-08-20T07:03:44.369117Z","iopub.status.idle":"2022-08-20T07:03:55.353093Z","shell.execute_reply.started":"2022-08-20T07:03:44.369077Z","shell.execute_reply":"2022-08-20T07:03:55.351901Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# WER (word error rate) calculation process\n\n","metadata":{}},{"cell_type":"markdown","source":"![](https://miro.medium.com/max/700/1*MUGLdWm3zMYK7dLmyo3pqA.png)\n\n**WER = 100 (insertions(INS) + substitutions(SUB) + deletions(DEL))**\n\n![](http://www.italk2learn.eu/wp-content/uploads/2015/02/speech-bubble-image.png)","metadata":{}},{"cell_type":"markdown","source":"# CER (character error rate) calculation process\n\n","metadata":{}},{"cell_type":"markdown","source":"character error rate (cer) is a common metric of the performance of an automatic speech recognition system. This value indicates the percentage of characters that were incorrectly predicted. The lower the value, the better the performance of the ASR system with a CER of 0 being a perfect score.\n\nCER calculation is based on the concept of [Levenshtein distance](https://towardsdatascience.com/evaluating-ocr-output-quality-with-character-error-rate-cer-and-word-error-rate-wer-853175297510#9bd1), where we count the minimum number of character-level operations required to transform the ground truth text (aka reference text) into the OCR output.\n\nCharacter Error Rate (CER) formula :\n\n![](https://miro.medium.com/max/700/1*KsWFDKnLI7mudmhbzGjc4w.png)\n\nwhere:\n\n* S = Number of Substitutions\n* D = Number of Deletions\n* I = Number of Insertions\n* N = Number of characters in reference text (aka ground truth)\n\nLet’s look at an example:\n\n**Ground Truth Reference Text**: 809475127\n\n**ASR Transcribed Output Text**: 80g475Z7\n\nSeveral errors require edits to transform ASR output into the ground truth:\n\n1. g instead of 9 (at reference text character 3)\n2. Missing 1 (at reference text character 7)\n3. Z instead of 2 (at reference text character 8)\n\nWith that, here are the values to input into the equation:\n\n* Number of Substitutions (S) = 2\n* Number of Deletions (D) = 1\n* Number of Insertions (I) = 0\n* Number of characters in reference text (N) = 9\n\nBased on the above, we get (2 + 1 + 0) / 9 = 0.3333. When converted to a percentage value, the CER becomes 33.33%. This implies that every 3rd character in the sequence was incorrectly transcribed.\n\nWe repeat this calculation for all the pairs of transcribed output and corresponding ground truth, and take the mean of these values to obtain an overall CER percentage.\n\n**Reference :** [Evaluate OCR Output Quality with Character Error Rate (CER) and Word Error Rate (WER)](https://towardsdatascience.com/evaluating-ocr-output-quality-with-character-error-rate-cer-and-word-error-rate-wer-853175297510#5aec)","metadata":{}},{"cell_type":"markdown","source":"# calculating metric on whole validation set","metadata":{}},{"cell_type":"code","source":"\ndf = pd.DataFrame(columns=['predictions', 'references'])\ndf.predictions = predictions\ndf.references = references\ndf.to_csv('./results.csv',index = False) #use it for error analysis and other stuffs\ndf.head(10)","metadata":{"execution":{"iopub.status.busy":"2022-08-20T07:03:55.358763Z","iopub.execute_input":"2022-08-20T07:03:55.360918Z","iopub.status.idle":"2022-08-20T07:03:55.427873Z","shell.execute_reply.started":"2022-08-20T07:03:55.360880Z","shell.execute_reply":"2022-08-20T07:03:55.426777Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Without Post Processing","metadata":{}},{"cell_type":"code","source":"cer = load_metric(\"cer\")\nwer = load_metric(\"wer\")\n\ncer_score = cer.compute(predictions=predictions, references=references)\nprint(\"validation cer_score -> \",cer_score)\nwer_score = wer.compute(predictions=predictions, references=references)\nprint(\"validation wer_score -> \",wer_score)","metadata":{"execution":{"iopub.status.busy":"2022-08-20T07:03:55.429329Z","iopub.execute_input":"2022-08-20T07:03:55.429695Z","iopub.status.idle":"2022-08-20T07:04:00.229557Z","shell.execute_reply.started":"2022-08-20T07:03:55.429659Z","shell.execute_reply":"2022-08-20T07:04:00.228109Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# With  post processing","metadata":{}},{"cell_type":"markdown","source":"during error analysis using the results.csv file we've seen that the model is frequently missing to predict punctuations, almost all the sentences in ground truth ends with '।' but while predicting using the public best trained model we can see that the model is missing to predict '।' most of the times, so in the simple post processing code below we will check if the predicted sentence ends with '।' or not,if no then we forcefully add '।' at the end of the predicted sentence.","metadata":{}},{"cell_type":"code","source":"for i in range(len(predictions)):\n    if(predictions[i][-1] == '।'):\n        continue\n    else:\n        predictions[i] = predictions[i]+'।'","metadata":{"execution":{"iopub.status.busy":"2022-08-20T07:04:00.231335Z","iopub.execute_input":"2022-08-20T07:04:00.231716Z","iopub.status.idle":"2022-08-20T07:04:00.247124Z","shell.execute_reply.started":"2022-08-20T07:04:00.231678Z","shell.execute_reply":"2022-08-20T07:04:00.242209Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cer_score = cer.compute(predictions=predictions, references=references)\nprint(\"Final validation cer_score -> \",cer_score)\nwer_score = wer.compute(predictions=predictions, references=references)\nprint(\"Final validation wer_score -> \",wer_score)","metadata":{"execution":{"iopub.status.busy":"2022-08-20T07:04:00.249035Z","iopub.execute_input":"2022-08-20T07:04:00.249886Z","iopub.status.idle":"2022-08-20T07:04:01.573933Z","shell.execute_reply.started":"2022-08-20T07:04:00.249832Z","shell.execute_reply":"2022-08-20T07:04:01.572756Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**WOW great!!!\nwith above post processing word error rate improved from 0.30921300101701055 to 0.28501372267654884 and that's 0.024199278340461705 improvement,not bad no?**\n","metadata":{}},{"cell_type":"markdown","source":"# Submission with post processing","metadata":{}},{"cell_type":"code","source":"df = pd.read_csv('../input/dlsprint/sample_submission.csv')\nlen(df.path)","metadata":{"execution":{"iopub.status.busy":"2022-08-20T07:04:01.575332Z","iopub.execute_input":"2022-08-20T07:04:01.575905Z","iopub.status.idle":"2022-08-20T07:04:01.629936Z","shell.execute_reply.started":"2022-08-20T07:04:01.575848Z","shell.execute_reply":"2022-08-20T07:04:01.628820Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n\ntest_dataset = bn_asr_Dataset(CFG.sample_sub_df_path,CFG.test_wav)\n\n# for i,prediction in enumerate(tqdm(asr(test_dataset, chunk_length_s=112, stride_length_s=None,batch_size=CFG.batch_size), total=len(test_dataset))):\n#     df.sentence[i] = prediction[\"text\"]\n    \nfor i in range(len(test_dataset)):\n    pred = asr(test_dataset.__getitem__(i), chunk_length_s=112, stride_length_s=None)\n    \n    #applying simple post processing with error handler\n    try:\n        if(pred[\"text\"][-1] == '।'):\n            df.sentence[i] = pred[\"text\"]\n        else:\n            df.sentence[i] = pred[\"text\"]+'।'\n    except:\n        print(\"predicted text at idx \",i,\" is -> \",pred[\"text\"])\n        df.sentence[i] = pred[\"text\"]+'।'\n","metadata":{"execution":{"iopub.status.busy":"2022-08-20T07:04:01.631610Z","iopub.execute_input":"2022-08-20T07:04:01.632050Z","iopub.status.idle":"2022-08-20T07:25:14.200137Z","shell.execute_reply.started":"2022-08-20T07:04:01.632011Z","shell.execute_reply":"2022-08-20T07:25:14.198901Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.head(3)","metadata":{"execution":{"iopub.status.busy":"2022-08-20T07:25:14.201894Z","iopub.execute_input":"2022-08-20T07:25:14.202305Z","iopub.status.idle":"2022-08-20T07:25:14.213053Z","shell.execute_reply.started":"2022-08-20T07:25:14.202267Z","shell.execute_reply":"2022-08-20T07:25:14.211930Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.to_csv('./submission.csv',index = False)\ndf.sentence[1]","metadata":{"execution":{"iopub.status.busy":"2022-08-20T07:25:14.214603Z","iopub.execute_input":"2022-08-20T07:25:14.215229Z","iopub.status.idle":"2022-08-20T07:25:14.254877Z","shell.execute_reply.started":"2022-08-20T07:25:14.215193Z","shell.execute_reply":"2022-08-20T07:25:14.253983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"IPython.display.Audio('../input/dlsprint/test_files/common_voice_bn_31675220.mp3')","metadata":{"execution":{"iopub.status.busy":"2022-08-20T07:25:14.256248Z","iopub.execute_input":"2022-08-20T07:25:14.256608Z","iopub.status.idle":"2022-08-20T07:25:14.270580Z","shell.execute_reply.started":"2022-08-20T07:25:14.256582Z","shell.execute_reply":"2022-08-20T07:25:14.269535Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.sentence[0]","metadata":{"execution":{"iopub.status.busy":"2022-08-20T07:25:14.271946Z","iopub.execute_input":"2022-08-20T07:25:14.272399Z","iopub.status.idle":"2022-08-20T07:25:14.279822Z","shell.execute_reply.started":"2022-08-20T07:25:14.272361Z","shell.execute_reply":"2022-08-20T07:25:14.278470Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.sentence[80]","metadata":{"execution":{"iopub.status.busy":"2022-08-20T07:25:14.281377Z","iopub.execute_input":"2022-08-20T07:25:14.282121Z","iopub.status.idle":"2022-08-20T07:25:14.289555Z","shell.execute_reply.started":"2022-08-20T07:25:14.282083Z","shell.execute_reply":"2022-08-20T07:25:14.288450Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# optional (post ASR correction attempt)","metadata":{}},{"cell_type":"markdown","source":"in this section we will try to implement the recent best research on POST OCR (optical character recognition) CORRECTION titled[ Post-OCR Document Correction with large Ensembles of Character\nSequence-to-Sequence Models](https://arxiv.org/pdf/2109.06264.pdf) this research work was done in ocr domain and not in ASR domain so i was thinking what will happen if we try this approach in ASR domain? **well if you never try you'll never know**.\nThe core of this system is a standard sequence-to-sequence model that can correct sequences of characters. In the below implementation, we used a Transformer as the sequence model, which takes as input a segment of characters from the document to correct, and the output is the corrected segment. To train this sequence model, it is necessary to align the raw documents with their corresponding correct transcriptions, which is not always straightforward.Since the output is not necessarily of the same length as the input (because of possible insertions or deletions of characters), a decoding method like Greedy Search or Beam Search\nis needed to produce the most likely corrected sequence according to the model.\nfor the below experiment we will be using results.csv where references column contains actual clean annotation and predictions column contains output of STT model including errors\n","metadata":{}},{"cell_type":"code","source":"!git clone https://github.com/jarobyte91/post_ocr_correction.git\nos.chdir('./post_ocr_correction')\n!pwd\n!ls","metadata":{"execution":{"iopub.status.busy":"2022-08-20T07:25:14.291229Z","iopub.execute_input":"2022-08-20T07:25:14.291894Z","iopub.status.idle":"2022-08-20T07:25:19.568305Z","shell.execute_reply.started":"2022-08-20T07:25:14.291834Z","shell.execute_reply":"2022-08-20T07:25:19.567047Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install .","metadata":{"execution":{"iopub.status.busy":"2022-08-20T07:25:19.571469Z","iopub.execute_input":"2022-08-20T07:25:19.571912Z","iopub.status.idle":"2022-08-20T07:25:58.552677Z","shell.execute_reply.started":"2022-08-20T07:25:19.571849Z","shell.execute_reply":"2022-08-20T07:25:58.551502Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.chdir('..')\n!ls","metadata":{"execution":{"iopub.status.busy":"2022-08-20T07:25:58.554428Z","iopub.execute_input":"2022-08-20T07:25:58.554823Z","iopub.status.idle":"2022-08-20T07:25:59.659156Z","shell.execute_reply.started":"2022-08-20T07:25:58.554783Z","shell.execute_reply":"2022-08-20T07:25:59.657977Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"results = pd.read_csv('../input/commonvoice-bn-xls-r-metric-calculation/results.csv')\nprint(len(results))\npreds = results.predictions.tolist()\nrefs = results.references.tolist()\nresults.head()","metadata":{"execution":{"iopub.status.busy":"2022-08-20T07:25:59.661120Z","iopub.execute_input":"2022-08-20T07:25:59.661894Z","iopub.status.idle":"2022-08-20T07:25:59.755465Z","shell.execute_reply.started":"2022-08-20T07:25:59.661832Z","shell.execute_reply":"2022-08-20T07:25:59.754349Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Be careful,\nif you predict on train set using the best publicly available ASR bangla model from huggingface you will see the model making NaN prediction for many audio samples in train set,to get the index of those NaN output files i used the code below ","metadata":{}},{"cell_type":"code","source":"#no nan output in results.csv but they exist in train.csv (give it a try)\nidx = [i for i, x in zip(range(len(preds)), preds) if not isinstance(x,str)]\nfor ele in sorted(idx, reverse = True):\n    del preds[ele]\n    del refs[ele]\nlen(preds),len(refs)\n","metadata":{"execution":{"iopub.status.busy":"2022-08-20T07:25:59.757021Z","iopub.execute_input":"2022-08-20T07:25:59.757359Z","iopub.status.idle":"2022-08-20T07:25:59.769229Z","shell.execute_reply.started":"2022-08-20T07:25:59.757323Z","shell.execute_reply":"2022-08-20T07:25:59.768018Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.empty_cache()\ngc.collect()\n!nvidia-smi","metadata":{"execution":{"iopub.status.busy":"2022-08-20T07:25:59.771012Z","iopub.execute_input":"2022-08-20T07:25:59.771765Z","iopub.status.idle":"2022-08-20T07:26:11.475128Z","shell.execute_reply.started":"2022-08-20T07:25:59.771728Z","shell.execute_reply":"2022-08-20T07:26:11.473825Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from pytorch_beam_search import seq2seq\nfrom post_ocr_correction import correction\n\nfor i in range(len(refs)):\n    refs[i] = list(refs[i])\n    preds[i] = list(preds[i])\n\n\n# train data and model\nsource = preds\ntarget = refs\nsource_index = seq2seq.Index(source)\ntarget_index = seq2seq.Index(target)\nX = source_index.text2tensor(source)\nY = target_index.text2tensor(target)\nprint(source_index)\nprint(\".....\")\nprint(target_index)\nprint(\".....\")\nprint(X)\nprint(\".....\")\nprint(Y)","metadata":{"execution":{"iopub.status.busy":"2022-08-20T07:26:11.480506Z","iopub.execute_input":"2022-08-20T07:26:11.483165Z","iopub.status.idle":"2022-08-20T07:26:13.907230Z","shell.execute_reply.started":"2022-08-20T07:26:11.483119Z","shell.execute_reply":"2022-08-20T07:26:13.906032Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"\n![](https://i.ibb.co/741XzvD/post-stt-corrector.png)","metadata":{}},{"cell_type":"code","source":"%%time\n\nepochs = 20\nbatch_size = 1024\nPATH = './post_ASR_corrector.pt'\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(device)\nX= X.to(device)\nY= Y.to(device)\n\n# def post_asr_corrector():\n#     model = seq2seq.Transformer(source_index, target_index,max_sequence_length = len(results)+4,\n#                     embedding_dimension = 512,\n#                     feedforward_dimension = 1024,\n#                     attention_heads = 2,\n#                     encoder_layers = 2,\n#                     decoder_layers = 2)\n#     return model\n\ndef post_asr_corrector():\n    model = seq2seq.Transformer(source_index, target_index,max_sequence_length = 256,dropout = 0.0,embedding_dimension = 192)\n    return model\nmodel = post_asr_corrector()\nmodel.to(device)\nmodel.train()\nmodel.fit(X, Y, epochs = epochs, progress_bar = 1,batch_size = batch_size)\nmodel.eval()\ntorch.save(model.state_dict(), PATH)","metadata":{"execution":{"iopub.status.busy":"2022-08-20T07:26:13.913396Z","iopub.execute_input":"2022-08-20T07:26:13.913713Z","iopub.status.idle":"2022-08-20T07:27:37.477692Z","shell.execute_reply.started":"2022-08-20T07:26:13.913685Z","shell.execute_reply":"2022-08-20T07:27:37.476547Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# load and infer","metadata":{}},{"cell_type":"code","source":"del model\ntorch.cuda.empty_cache()\ngc.collect()\n\n\nmodel = post_asr_corrector()\n\nmodel.load_state_dict(torch.load(PATH))\nmodel.to(device)\nmodel.eval()\nmodel","metadata":{"execution":{"iopub.status.busy":"2022-08-20T07:27:37.479288Z","iopub.execute_input":"2022-08-20T07:27:37.479828Z","iopub.status.idle":"2022-08-20T07:27:48.614988Z","shell.execute_reply.started":"2022-08-20T07:27:37.479785Z","shell.execute_reply":"2022-08-20T07:27:48.613986Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# test data\ntest = 'শীতকালীন গেমসে এখনো কোন পদক জিততে পারেনি।'\nreference = 'শীতকালীন গেমস এ কোন কোন পদক জিততে পারেনি।'\nnew_source = [list(test)]\nX_new = source_index.text2tensor(new_source).to(device)\n\n# plain beam search\npredictions, log_probabilities = seq2seq.beam_search(\n    model, \n    X_new,\n    progress_bar = 16\n)\njust_beam = target_index.tensor2text(predictions[:, 0, :])[0]\njust_beam = re.sub(r\"<START>|<PAD>|<UNK>|<END>.*\", \"\", just_beam)\njust_beam","metadata":{"execution":{"iopub.status.busy":"2022-08-20T07:27:48.616415Z","iopub.execute_input":"2022-08-20T07:27:48.617273Z","iopub.status.idle":"2022-08-20T07:27:49.497259Z","shell.execute_reply.started":"2022-08-20T07:27:48.617235Z","shell.execute_reply":"2022-08-20T07:27:49.496255Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"log_probabilities","metadata":{"execution":{"iopub.status.busy":"2022-08-20T07:27:49.498778Z","iopub.execute_input":"2022-08-20T07:27:49.499947Z","iopub.status.idle":"2022-08-20T07:27:49.519037Z","shell.execute_reply.started":"2022-08-20T07:27:49.499892Z","shell.execute_reply":"2022-08-20T07:27:49.517984Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# post ASR correction\n# disjoint_beam = correction.disjoint(\n#     test,\n#     model,\n#     source_index,\n#     target_index,\n#     50,\n#     \"beam_search\",\n# )","metadata":{"execution":{"iopub.status.busy":"2022-08-20T07:27:49.520765Z","iopub.execute_input":"2022-08-20T07:27:49.521159Z","iopub.status.idle":"2022-08-20T07:27:49.525625Z","shell.execute_reply.started":"2022-08-20T07:27:49.521123Z","shell.execute_reply":"2022-08-20T07:27:49.524520Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"\\nresults\")\nprint(\"  test data                      \", test)\nprint(\"  plain beam search              \", just_beam)\n#print(\"  disjoint windows, beam search  \", disjoint_beam)\n","metadata":{"execution":{"iopub.status.busy":"2022-08-20T07:27:49.527706Z","iopub.execute_input":"2022-08-20T07:27:49.528145Z","iopub.status.idle":"2022-08-20T07:27:49.536146Z","shell.execute_reply.started":"2022-08-20T07:27:49.528110Z","shell.execute_reply":"2022-08-20T07:27:49.534987Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# improvement ideas\n\nfor better post ASR correction,i would like to recommend going deep in [ROBART](https://arxiv.org/pdf/2202.01157.pdf)\n![](https://i.ibb.co/YWVGRVF/post-asr-corrector.png)\n\n![](https://i.ibb.co/9sZbvD7/post-asr.png)\none example implementation of levenshtein transformer can be found [here](https://github.com/nmfisher/levenshtein_transformer/blob/master/Untitled.ipynb)\n\nmore about post ASR correction was discussed [here](https://www.kaggle.com/competitions/dlsprint/discussion/335411)","metadata":{}},{"cell_type":"markdown","source":"![](https://images.unsplash.com/photo-1499744937866-d7e566a20a61?ixlib=rb-1.2.1&ixid=MnwxMjA3fDB8MHxwaG90by1wYWdlfHx8fGVufDB8fHx8&auto=format&fit=crop&w=870&q=80)","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}