{"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\n!pip install bnunicodenormalizer\n\n!pip install aksharamukha\n!pip install -q torchaudio omegaconf\n\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-03-12T08:15:55.908159Z","iopub.execute_input":"2023-03-12T08:15:55.908534Z","iopub.status.idle":"2023-03-12T08:17:43.780267Z","shell.execute_reply.started":"2023-03-12T08:15:55.908503Z","shell.execute_reply":"2023-03-12T08:17:43.778837Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"![](https://www.respeecher.com/hubfs/What-is-Text-to-Speech-TTS%29-Initial-Speech-Synthesis-Explained-Respeecher-voice-cloning-software.jpeg)\n\n![](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":"pip install aksharamukha","metadata":{"execution":{"iopub.status.busy":"2023-03-12T09:08:35.467041Z","iopub.execute_input":"2023-03-12T09:08:35.467825Z","iopub.status.idle":"2023-03-12T09:08:49.536895Z","shell.execute_reply.started":"2023-03-12T09:08:35.467788Z","shell.execute_reply":"2023-03-12T09:08:49.535444Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nprint(sys.version)","metadata":{"execution":{"iopub.status.busy":"2023-03-12T09:28:11.765322Z","iopub.execute_input":"2023-03-12T09:28:11.765727Z","iopub.status.idle":"2023-03-12T09:28:11.770891Z","shell.execute_reply.started":"2023-03-12T09:28:11.765694Z","shell.execute_reply":"2023-03-12T09:28:11.769931Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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\nimport torchaudio\nfrom IPython.display import Audio, display\nfrom aksharamukha import transliterate\nimport random\n\nfrom bnunicodenormalizer import Normalizer \n\ntqdm.pandas()\nimport warnings\nwarnings.filterwarnings(\"ignore\")\nfrom pandarallel import pandarallel\n\npandarallel.initialize(progress_bar=True,nb_workers=8)\n\n\nprint(torch.__version__)\nprint(torchaudio.__version__)\n\nbnorm=Normalizer()","metadata":{"execution":{"iopub.status.busy":"2023-03-12T09:09:51.040469Z","iopub.execute_input":"2023-03-12T09:09:51.041386Z","iopub.status.idle":"2023-03-12T09:09:51.090160Z","shell.execute_reply.started":"2023-03-12T09:09:51.041343Z","shell.execute_reply":"2023-03-12T09:09:51.088693Z"},"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    post_asr_corrector = False\n    \n\n","metadata":{"execution":{"iopub.status.busy":"2023-03-12T08:17:43.928394Z","iopub.status.idle":"2023-03-12T08:17:43.928724Z","shell.execute_reply.started":"2023-03-12T08:17:43.928565Z","shell.execute_reply":"2023-03-12T08:17:43.928581Z"},"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":"2023-03-12T08:17:43.930375Z","iopub.status.idle":"2023-03-12T08:17:43.931164Z","shell.execute_reply.started":"2023-03-12T08:17:43.930892Z","shell.execute_reply":"2023-03-12T08:17:43.930919Z"},"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":"2023-03-12T08:17:43.933263Z","iopub.status.idle":"2023-03-12T08:17:43.933818Z","shell.execute_reply.started":"2023-03-12T08:17:43.933565Z","shell.execute_reply":"2023-03-12T08:17:43.933592Z"},"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":"2023-03-12T08:17:43.935821Z","iopub.status.idle":"2023-03-12T08:17:43.936299Z","shell.execute_reply.started":"2023-03-12T08:17:43.936059Z","shell.execute_reply":"2023-03-12T08:17:43.936083Z"},"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":"2023-03-12T08:17:43.938120Z","iopub.status.idle":"2023-03-12T08:17:43.938603Z","shell.execute_reply.started":"2023-03-12T08:17:43.938348Z","shell.execute_reply":"2023-03-12T08:17:43.938380Z"},"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)\n\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":"2023-03-12T08:17:43.940761Z","iopub.status.idle":"2023-03-12T08:17:43.941744Z","shell.execute_reply.started":"2023-03-12T08:17:43.941489Z","shell.execute_reply":"2023-03-12T08:17:43.941513Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.empty_cache() \ngc.collect()\n!nvidia-smi","metadata":{"execution":{"iopub.status.busy":"2023-03-12T08:17:43.942925Z","iopub.status.idle":"2023-03-12T08:17:43.943777Z","shell.execute_reply.started":"2023-03-12T08:17:43.943527Z","shell.execute_reply":"2023-03-12T08:17:43.943551Z"},"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\n","metadata":{"execution":{"iopub.status.busy":"2023-03-12T08:17:43.945118Z","iopub.status.idle":"2023-03-12T08:17:43.945928Z","shell.execute_reply.started":"2023-03-12T08:17:43.945670Z","shell.execute_reply":"2023-03-12T08:17:43.945694Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Unicode Normalizer\n","metadata":{}},{"cell_type":"markdown","source":"from [webinar suplimentary notebook :: DL SPRINT](https://www.kaggle.com/code/nazmuddhohaansary/webinar-suplimentary-notebook-dl-sprint)","metadata":{}},{"cell_type":"code","source":"\ndef normalize(sen):\n    _words = [bnorm(word)['normalized']  for word in sen.split()]\n    return \" \".join([word for word in _words if word is not None]) \n\ndf.predictions= df.predictions.parallel_apply(lambda x:normalize(x))\ndf.references= df.references.parallel_apply(lambda x:normalize(x))\ndf.to_csv('./results.csv',index = False) #use it for error analysis and other stuffs\ndf.head(10)","metadata":{"execution":{"iopub.status.busy":"2023-03-12T08:17:43.947256Z","iopub.status.idle":"2023-03-12T08:17:43.948085Z","shell.execute_reply.started":"2023-03-12T08:17:43.947810Z","shell.execute_reply":"2023-03-12T08:17:43.947835Z"},"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=df.predictions, references=df.references)\nprint(\"validation cer_score -> \",cer_score)\nwer_score = wer.compute(predictions=df.predictions, references=df.references)\nprint(\"validation wer_score -> \",wer_score)","metadata":{"execution":{"iopub.status.busy":"2023-03-12T08:17:43.949455Z","iopub.status.idle":"2023-03-12T08:17:43.950271Z","shell.execute_reply.started":"2023-03-12T08:17:43.950033Z","shell.execute_reply":"2023-03-12T08:17:43.950058Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"* without bnunicodenormalizer our cv was (as discussed [here](https://www.kaggle.com/competitions/dlsprint/discussion/334951)) :\n\n**validation cer_score -> 0.09787704766628824**\n\n**validation wer_score -> 0.30921300101701055**\n\n* with bnunicodenormalizer our cv is :\n\n**validation cer_score -> 0.09668792125091967**\n\n**validation wer_score -> 0.3049220524108723**","metadata":{}},{"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(df.predictions)):\n    if(df.predictions[i][-1] == '।'):\n        continue\n    else:\n        df.predictions[i] = df.predictions[i]+'।'","metadata":{"execution":{"iopub.status.busy":"2023-03-12T08:17:43.951587Z","iopub.status.idle":"2023-03-12T08:17:43.953387Z","shell.execute_reply.started":"2023-03-12T08:17:43.953113Z","shell.execute_reply":"2023-03-12T08:17:43.953138Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cer_score = cer.compute(predictions=df.predictions, references=df.references)\nprint(\"Final validation cer_score -> \",cer_score)\nwer_score = wer.compute(predictions=df.predictions, references=df.references)\nprint(\"Final validation wer_score -> \",wer_score)","metadata":{"execution":{"iopub.status.busy":"2023-03-12T08:17:43.954777Z","iopub.status.idle":"2023-03-12T08:17:43.955568Z","shell.execute_reply.started":"2023-03-12T08:17:43.955304Z","shell.execute_reply":"2023-03-12T08:17:43.955330Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"* without bnunicodenormalizer our post processed model's cv was (as discussed [here](https://www.kaggle.com/competitions/dlsprint/discussion/334951#1856670) ) :\n\n**validation cer_score -> 0.09301847750965592**\n\n**validation wer_score -> 0.28501372267654884**\n\n* with bnunicodenormalizer our final cv (with post processing) is :\n\n**validation cer_score -> 0.09173650517129109**\n\n**validation wer_score -> 0.28054166260326835**\n\n#  with bnunicodenormalizer we've got 0.004472060073280493 WER improvement","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":"2023-03-12T08:17:43.957132Z","iopub.status.idle":"2023-03-12T08:17:43.957910Z","shell.execute_reply.started":"2023-03-12T08:17:43.957656Z","shell.execute_reply":"2023-03-12T08:17:43.957681Z"},"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        \ndf.sentence=df.sentence.parallel_apply(lambda x:normalize(x)) #unicode normalizer\n","metadata":{"execution":{"iopub.status.busy":"2023-03-12T08:17:43.959337Z","iopub.status.idle":"2023-03-12T08:17:43.960134Z","shell.execute_reply.started":"2023-03-12T08:17:43.959862Z","shell.execute_reply":"2023-03-12T08:17:43.959887Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.head(3)","metadata":{"execution":{"iopub.status.busy":"2023-03-12T08:17:43.961611Z","iopub.status.idle":"2023-03-12T08:17:43.962381Z","shell.execute_reply.started":"2023-03-12T08:17:43.962117Z","shell.execute_reply":"2023-03-12T08:17:43.962141Z"},"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":"2023-03-12T08:17:43.963773Z","iopub.status.idle":"2023-03-12T08:17:43.964551Z","shell.execute_reply.started":"2023-03-12T08:17:43.964287Z","shell.execute_reply":"2023-03-12T08:17:43.964311Z"},"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":"2023-03-12T08:17:43.965923Z","iopub.status.idle":"2023-03-12T08:17:43.966699Z","shell.execute_reply.started":"2023-03-12T08:17:43.966443Z","shell.execute_reply":"2023-03-12T08:17:43.966467Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.sentence[0]","metadata":{"execution":{"iopub.status.busy":"2023-03-12T08:17:43.968202Z","iopub.status.idle":"2023-03-12T08:17:43.968970Z","shell.execute_reply.started":"2023-03-12T08:17:43.968712Z","shell.execute_reply":"2023-03-12T08:17:43.968736Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.sentence[80]","metadata":{"execution":{"iopub.status.busy":"2023-03-12T08:17:43.970388Z","iopub.status.idle":"2023-03-12T08:17:43.971155Z","shell.execute_reply.started":"2023-03-12T08:17:43.970887Z","shell.execute_reply":"2023-03-12T08:17:43.970910Z"},"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":"2023-03-12T08:17:43.972560Z","iopub.status.idle":"2023-03-12T08:17:43.973376Z","shell.execute_reply.started":"2023-03-12T08:17:43.973103Z","shell.execute_reply":"2023-03-12T08:17:43.973130Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install .","metadata":{"execution":{"iopub.status.busy":"2023-03-12T08:17:43.974769Z","iopub.status.idle":"2023-03-12T08:17:43.975543Z","shell.execute_reply.started":"2023-03-12T08:17:43.975277Z","shell.execute_reply":"2023-03-12T08:17:43.975301Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.chdir('..')\n!ls","metadata":{"execution":{"iopub.status.busy":"2023-03-12T08:17:43.976996Z","iopub.status.idle":"2023-03-12T08:17:43.977814Z","shell.execute_reply.started":"2023-03-12T08:17:43.977570Z","shell.execute_reply":"2023-03-12T08:17:43.977595Z"},"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":"2023-03-12T08:17:43.979198Z","iopub.status.idle":"2023-03-12T08:17:43.979959Z","shell.execute_reply.started":"2023-03-12T08:17:43.979708Z","shell.execute_reply":"2023-03-12T08:17:43.979732Z"},"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":"2023-03-12T08:17:43.981377Z","iopub.status.idle":"2023-03-12T08:17:43.982141Z","shell.execute_reply.started":"2023-03-12T08:17:43.981871Z","shell.execute_reply":"2023-03-12T08:17:43.981895Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.empty_cache()\ngc.collect()\n!nvidia-smi","metadata":{"execution":{"iopub.status.busy":"2023-03-12T08:17:43.983528Z","iopub.status.idle":"2023-03-12T08:17:43.984294Z","shell.execute_reply.started":"2023-03-12T08:17:43.984046Z","shell.execute_reply":"2023-03-12T08:17:43.984072Z"},"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":"2023-03-12T08:17:43.985689Z","iopub.status.idle":"2023-03-12T08:17:43.986457Z","shell.execute_reply.started":"2023-03-12T08:17:43.986193Z","shell.execute_reply":"2023-03-12T08:17:43.986217Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"\n![](https://i.ibb.co/741XzvD/post-stt-corrector.png)","metadata":{}},{"cell_type":"markdown","source":"tried to train in versionn 5 already of this notebook [commonvoice_bn xls-r metric calculation](https://www.kaggle.com/code/mobassir/commonvoice-bn-xls-r-metric-calculation)","metadata":{}},{"cell_type":"code","source":"%%time\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\n\nif (CFG.post_asr_corrector):\n    print(\"training POST ASR Corrector...\\n\")\n    epochs = 6000\n    batch_size = 1024\n    PATH = './post_ASR_corrector.pt'\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    print(device)\n    X= X.to(device)\n    Y= 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\n\n    model = post_asr_corrector()\n    model.to(device)\n    model.train()\n    model.fit(X, Y, epochs = epochs, progress_bar = 1,batch_size = batch_size)\n    model.eval()\n    torch.save(model.state_dict(), PATH)","metadata":{"execution":{"iopub.status.busy":"2023-03-12T08:17:43.987904Z","iopub.status.idle":"2023-03-12T08:17:43.988683Z","shell.execute_reply.started":"2023-03-12T08:17:43.988431Z","shell.execute_reply":"2023-03-12T08:17:43.988455Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# load and infer","metadata":{}},{"cell_type":"code","source":"if (CFG.post_asr_corrector):\n    del model\n    torch.cuda.empty_cache()\n    gc.collect()\n\n\n    model = post_asr_corrector()\n\n    model.load_state_dict(torch.load(PATH))\n    model.to(device)\n    model.eval()\n    model","metadata":{"execution":{"iopub.status.busy":"2023-03-12T08:17:43.990107Z","iopub.status.idle":"2023-03-12T08:17:43.990882Z","shell.execute_reply.started":"2023-03-12T08:17:43.990626Z","shell.execute_reply":"2023-03-12T08:17:43.990651Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if (CFG.post_asr_corrector):\n    # test data\n    test = 'শীতকালীন গেমসে এখনো কোন পদক জিততে পারেনি।'\n    reference = 'শীতকালীন গেমস এ কোন কোন পদক জিততে পারেনি।'\n    new_source = [list(test)]\n    X_new = source_index.text2tensor(new_source).to(device)\n\n    # plain beam search\n    predictions, log_probabilities = seq2seq.beam_search(\n        model, \n        X_new,\n        progress_bar = 0\n    )\n    just_beam = target_index.tensor2text(predictions[:, 0, :])[0]\n    just_beam = re.sub(r\"<START>|<PAD>|<UNK>|<END>.*\", \"\", just_beam)\n    print(log_probabilities)\n    print(\"\\nresults\")\n    print(\"  test data                      \", test)\n    print(\"  plain beam search              \", just_beam)\n    #print(\"  disjoint windows, beam search  \", disjoint_beam)\n","metadata":{"execution":{"iopub.status.busy":"2023-03-12T08:17:43.992288Z","iopub.status.idle":"2023-03-12T08:17:43.993073Z","shell.execute_reply.started":"2023-03-12T08:17:43.992802Z","shell.execute_reply":"2023-03-12T08:17:43.992825Z"},"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":"2023-03-12T08:17:43.994517Z","iopub.status.idle":"2023-03-12T08:17:43.995323Z","shell.execute_reply.started":"2023-03-12T08:17:43.995062Z","shell.execute_reply":"2023-03-12T08:17:43.995088Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# silero-tts demo for bangla","metadata":{}},{"cell_type":"markdown","source":"**------------------------>>>>>>>>>>>>>>>>>>>>>>>>>>>  high level overview**\n\n![](https://miro.medium.com/max/1400/1*MwgQEqWrRMeQXdPBmwCh4g.png)\n![](https://www.researchgate.net/profile/Suhas-Mache/publication/304601298/figure/fig3/AS:667867041247263@1536243319820/Block-diagram-of-Text-to-Speech-System-Techniques-of-speech-synthesis-5-a.png)","metadata":{}},{"cell_type":"markdown","source":"even though this competition is all about STT (speech to text) recognition,however the dataset isn't limited to STT domain,you can work on TTS (text to speech) system using the dataset of this competition.we know that **Deep learning is data-hungry**. one good idea for possible improvement of your ASR system could be to try retraining your STT model with augmented data included in your pipeline like this [Transcoding & Augmenting Audio On-The-Fly](https://www.kaggle.com/code/shahruk10/transcoding-augmenting-audio-on-the-fly) but i was curious and also thinking that **can we also use the prediction result of a tts system for training our stt model? this is also similar to data augmentation technique but without having background noise,no?** \nsorry if i am wrong,i don't know if it will help or not,i am not expert in ASR domain,just a beginner who is sharing his naive thoughts. let's see how silero tts of torch.hub works with bangla.\n\n**KEEP IN MIND -> because bengali is a low resource language, that's why except silero tts we don't have any better free tts system for bangla. beside bangla STT,we also need a powerful bangla TTS system as well.**\n","metadata":{}},{"cell_type":"code","source":"\nprint(torch.__version__)\nprint(torchaudio.__version__)\n\n\nsentences = df.sentence.tolist()\n\nsample_rate = 48000\n# Loading model\nmodel, example_text = torch.hub.load(repo_or_dir='snakers4/silero-models',\n                                     model='silero_tts',\n                                     language='indic',\n                                     speaker='v3_indic')\nfor i in range(10):\n    idx = random.randint(i, 7740)\n    orig_text = sentences[idx]\n    print(f\"\\n\\n{idx} orig_text -> \",orig_text)\n    roman_text = transliterate.process('Bengali', 'ISO', orig_text)\n    print(\"\\n\\nroman_text -> \",roman_text)\n    if (i % 2) == 1:\n        audio = model.apply_tts(roman_text,\n                            speaker='bengali_male')\n    else:\n        audio = model.apply_tts(roman_text,\n                            speaker='bengali_female')\n    \n    torchaudio.save(f'idx_{idx}.wav', audio.unsqueeze(0), sample_rate)\n    display(Audio(audio, rate=sample_rate))","metadata":{"execution":{"iopub.status.busy":"2023-03-12T08:17:43.996741Z","iopub.status.idle":"2023-03-12T08:17:43.997525Z","shell.execute_reply.started":"2023-03-12T08:17:43.997258Z","shell.execute_reply":"2023-03-12T08:17:43.997282Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# improvement ideas\n\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":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}