{"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":"!pip -q install https://github.com/kpu/kenlm/archive/master.zip pyctcdecode","metadata":{"execution":{"iopub.status.busy":"2022-07-05T00:01:44.475152Z","iopub.execute_input":"2022-07-05T00:01:44.475548Z","iopub.status.idle":"2022-07-05T00:01:55.865225Z","shell.execute_reply.started":"2022-07-05T00:01:44.475508Z","shell.execute_reply":"2022-07-05T00:01:55.863861Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import re\nimport pandas as pd","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"chars_to_ignore_regex = '[\\,\\?\\.\\!\\-\\;\\:\\\"\\“\\%\\‘\\”\\�\\।\\’]'\n\ntrain_df = pd.read_csv('../input/dlsprint/train.csv')\nvalid_df = pd.read_csv('../input/dlsprint/validation.csv')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with open('text.txt', 'w') as f:\n    for sentence in train_df['sentence']:\n        f.write(re.sub(chars_to_ignore_regex, '', sentence))\n        f.write(' ')\n        \n    for sentence in valid_df['sentence']:\n        f.write(re.sub(chars_to_ignore_regex, '', sentence))\n        f.write(' ')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! sudo apt -y install build-essential cmake libboost-system-dev libboost-thread-dev libboost-program-options-dev libboost-test-dev libeigen3-dev zlib1g-dev libbz2-dev liblzma-dev\n! wget -O - https://kheafield.com/code/kenlm.tar.gz | tar xz\n! mkdir kenlm/build && cd kenlm/build && cmake .. && make -j2\n! ls kenlm/build/bin\n! kenlm/build/bin/lmplz -o 2 < \"text.txt\" > \"2gram.arpa\"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with open(\"2gram.arpa\", \"r\") as read_file, open(\"2gram_correct.arpa\", \"w\") as write_file:\n  has_added_eos = False\n  for line in read_file:\n    if not has_added_eos and \"ngram 1=\" in line:\n      count=line.strip().split(\"=\")[-1]\n      write_file.write(line.replace(f\"{count}\", f\"{int(count)+1}\"))\n    elif not has_added_eos and \"<s>\" in line:\n      write_file.write(line)\n      write_file.write(line.replace(\"<s>\", \"</s>\"))\n      has_added_eos = True\n    else:\n      write_file.write(line)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import Wav2Vec2Processor\nfrom transformers import Wav2Vec2ForCTC\nfrom transformers import Wav2Vec2ProcessorWithLM\nfrom transformers import Wav2Vec2FeatureExtractor","metadata":{"execution":{"iopub.status.busy":"2022-07-05T00:01:57.834128Z","iopub.execute_input":"2022-07-05T00:01:57.834831Z","iopub.status.idle":"2022-07-05T00:02:05.659025Z","shell.execute_reply.started":"2022-07-05T00:01:57.834794Z","shell.execute_reply":"2022-07-05T00:02:05.658078Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = Wav2Vec2ForCTC.from_pretrained('arijitx/wav2vec2-xls-r-300m-bengali')\nprocessor = Wav2Vec2Processor.from_pretrained('arijitx/wav2vec2-xls-r-300m-bengali')\n\nvocab_dict = processor.tokenizer.get_vocab()\nsorted_vocab_dict = {k.lower(): v for k, v in sorted(vocab_dict.items(), key=lambda item: item[1])}","metadata":{"execution":{"iopub.status.busy":"2022-07-05T00:04:08.263578Z","iopub.execute_input":"2022-07-05T00:04:08.264938Z","iopub.status.idle":"2022-07-05T00:07:34.535868Z","shell.execute_reply.started":"2022-07-05T00:04:08.2649Z","shell.execute_reply":"2022-07-05T00:07:34.534808Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from pyctcdecode import build_ctcdecoder\n\ndecoder = build_ctcdecoder(\n    labels=list(sorted_vocab_dict.keys()),\n    kenlm_model_path=\"2gram_correct.arpa\",\n)\n\nfrom transformers import Wav2Vec2ProcessorWithLM\n\nprocessor = Wav2Vec2ProcessorWithLM(\n    feature_extractor=processor.feature_extractor,\n    tokenizer=processor.tokenizer,\n    decoder=decoder\n)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\n\nimport torch\nimport torchaudio\nimport numpy as np","metadata":{"execution":{"iopub.status.busy":"2022-07-05T00:08:32.242543Z","iopub.execute_input":"2022-07-05T00:08:32.243152Z","iopub.status.idle":"2022-07-05T00:08:32.526023Z","shell.execute_reply.started":"2022-07-05T00:08:32.243115Z","shell.execute_reply":"2022-07-05T00:08:32.524596Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df['path'] = train_df['path'].map(lambda path: os.path.join('../input/dlsprint/train_files', path))\n\nvalid_df['path'] = valid_df['path'].map(lambda path: os.path.join('../input/dlsprint/validation_files', path))\n\nsubmit_df = pd.read_csv('../input/dlsprint/sample_submission.csv')\nsubmit_df['path'] = submit_df['path'].map(lambda path: os.path.join('../input/dlsprint/test_files', path))","metadata":{"execution":{"iopub.status.busy":"2022-07-05T00:08:56.422803Z","iopub.execute_input":"2022-07-05T00:08:56.423189Z","iopub.status.idle":"2022-07-05T00:08:58.730598Z","shell.execute_reply.started":"2022-07-05T00:08:56.423141Z","shell.execute_reply":"2022-07-05T00:08:58.729593Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TrainDS(torch.utils.data.Dataset):\n    \n    def __init__(self, df):\n        self.df = df\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        paths = self.df['path'][idx]\n        arrays = []\n        for path in paths:\n            array, sampling_rate = torchaudio.load(path)\n            resampler = torchaudio.transforms.Resample(sampling_rate, 16_000)\n            array = resampler(array)\n            input_values = processor(array, sampling_rate=16_000).input_values[0]\n            arrays.append({\"input_values\":input_values[0]})\n        \n        batch = processor.pad(\n            arrays,\n            return_tensors=\"pt\",\n        )\n        \n        return batch['input_values'].to('cuda'), batch['attention_mask'].to('cuda')","metadata":{"execution":{"iopub.status.busy":"2022-07-05T00:09:25.711131Z","iopub.execute_input":"2022-07-05T00:09:25.712118Z","iopub.status.idle":"2022-07-05T00:09:25.72082Z","shell.execute_reply.started":"2022-07-05T00:09:25.712082Z","shell.execute_reply":"2022-07-05T00:09:25.719252Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ds = TrainDS(train_df)\nvalid_ds = TrainDS(valid_df)\nsubmit_ds = TrainDS(submit_df)","metadata":{"execution":{"iopub.status.busy":"2022-07-05T00:09:27.447116Z","iopub.execute_input":"2022-07-05T00:09:27.448048Z","iopub.status.idle":"2022-07-05T00:09:27.456617Z","shell.execute_reply.started":"2022-07-05T00:09:27.448004Z","shell.execute_reply":"2022-07-05T00:09:27.455498Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.to('cuda');","metadata":{"execution":{"iopub.status.busy":"2022-07-05T00:10:36.790955Z","iopub.execute_input":"2022-07-05T00:10:36.791342Z","iopub.status.idle":"2022-07-05T00:10:36.80381Z","shell.execute_reply.started":"2022-07-05T00:10:36.791309Z","shell.execute_reply":"2022-07-05T00:10:36.802603Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm import tqdm\n\nwith torch.no_grad():\n    for i, j in zip(tqdm(range(0, 7725, 25)), range(20, 7750, 25)):\n        input_values, attention_mask = submit_ds[i:j]\n        logits = model(input_values=input_values, attention_mask=attention_mask).logits\n        submit_df['sentence'][i:j] = processor.batch_decode(logits=logits.cpu().numpy()).text","metadata":{"execution":{"iopub.status.busy":"2022-07-05T00:30:03.470626Z","iopub.execute_input":"2022-07-05T00:30:03.471011Z","iopub.status.idle":"2022-07-05T00:52:56.715747Z","shell.execute_reply.started":"2022-07-05T00:30:03.470976Z","shell.execute_reply":"2022-07-05T00:52:56.714097Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with torch.no_grad():\n    i = 7725\n    input_values, attention_mask = submit_ds[i:]\n    logits = model(input_values=input_values, attention_mask=attention_mask).logits\n    submit_df['sentence'][i:] = processor.batch_decode(logits.cpu().numpy()).text\n    \nsubmit_df.tail()","metadata":{"execution":{"iopub.status.busy":"2022-07-05T00:55:11.240489Z","iopub.execute_input":"2022-07-05T00:55:11.241278Z","iopub.status.idle":"2022-07-05T00:55:14.780054Z","shell.execute_reply.started":"2022-07-05T00:55:11.241239Z","shell.execute_reply":"2022-07-05T00:55:14.778607Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submit_df['path'] = pd.read_csv('../input/dlsprint/sample_submission.csv')['path']\nsubmit_df.to_csv('submit.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2022-07-05T00:55:35.038079Z","iopub.execute_input":"2022-07-05T00:55:35.038822Z","iopub.status.idle":"2022-07-05T00:55:35.118376Z","shell.execute_reply.started":"2022-07-05T00:55:35.038779Z","shell.execute_reply":"2022-07-05T00:55:35.117378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submit_df.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-05T00:59:18.899617Z","iopub.execute_input":"2022-07-05T00:59:18.90009Z","iopub.status.idle":"2022-07-05T00:59:18.911293Z","shell.execute_reply.started":"2022-07-05T00:59:18.900057Z","shell.execute_reply":"2022-07-05T00:59:18.909232Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submit_df.tail()","metadata":{"execution":{"iopub.status.busy":"2022-07-05T00:59:33.238984Z","iopub.execute_input":"2022-07-05T00:59:33.239362Z","iopub.status.idle":"2022-07-05T00:59:33.250306Z","shell.execute_reply.started":"2022-07-05T00:59:33.239329Z","shell.execute_reply":"2022-07-05T00:59:33.249244Z"},"trusted":true},"execution_count":null,"outputs":[]}]}