{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":4143520,"sourceType":"datasetVersion","datasetId":2447262},{"sourceId":6707460,"sourceType":"datasetVersion","datasetId":3865741},{"sourceId":8310846,"sourceType":"datasetVersion","datasetId":4879926}],"dockerImageVersionId":30674,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Introduction","metadata":{}},{"cell_type":"markdown","source":"Continuation of the pseudolabeling pipeline described in https://www.kaggle.com/code/reasat/pseudolabeling-step-1-download-speech-audio\n\nModel weight and inference notebook copied from: https://www.kaggle.com/competitions/bengaliai-speech/discussion/447970\n\n## STT Model:\n\n* OpenAI whisper-medium\n* Huggingface trainer\n* Trained on 8x 48GB RTX A6000\n* bs=8 and lr=1e-5\n* Train steps 50k\n* Spectrogram dithering\n* Spectrogram time and frequency masking\n* Resampling 16khz->8khz->16khz as augmentation\n* Inference with max_length=260, num_beams=4 and chunk_length_s=20.1s\n* Libsonic based speed/pitch augmentation\n* Datasets: OpenSLR 37, OpenSLR 53, MadASR, Shrutilipi, Macro, Kathbath, GoogleTTS generated audios and pseudo labeled YouTube videos\n\n## Punctuation Model:\n\n* AutoModelForTokenClassification google/muril-base-cased\n* Huggingface trainer\n* Labels: period, comma and question mark\n* bs=64, lr=2e-4 and max_seq_length=512\n* Ensemble of 4 models (using 6, 8, 11 and 12 layers of google/muril-base-cased)\n* Normalized IndicCorp v2 Bangla dataset\n\n","metadata":{}},{"cell_type":"code","source":"# !cp /kaggle/input/bengali-eval-data/predict.py .\n\n!cp -r ../input/python-packages2 ./\n!tar xvfz ./python-packages2/jiwer.tgz\n!pip install ./jiwer/python-Levenshtein-0.12.2.tar.gz -f ./ --no-index\n!pip install ./jiwer/jiwer-2.3.0-py3-none-any.whl -f ./ --no-index","metadata":{"execution":{"iopub.status.busy":"2024-06-08T19:13:05.303162Z","iopub.execute_input":"2024-06-08T19:13:05.304052Z","iopub.status.idle":"2024-06-08T19:13:37.677252Z","shell.execute_reply.started":"2024-06-08T19:13:05.303990Z","shell.execute_reply":"2024-06-08T19:13:37.676035Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport csv\nimport time\nimport glob\n\n# MODEL = '/kaggle/input/bengali-ai-asr-submission/bengali-whisper-medium/'\nPUNCT_MODELS = [\n    '/kaggle/input/bengali-ai-asr-submission/punct-model-6layers/',\n    '/kaggle/input/bengali-ai-asr-submission/punct-model-8layers/',\n    '/kaggle/input/bengali-ai-asr-submission/punct-model-11layers/',\n    '/kaggle/input/bengali-ai-asr-submission/punct-model-12layers/'\n]\nPUNCT_WEIGHTS = [[1.0, 1.4, 1.0, 0.8]]\n\nCHUNK_LENGTH_S = 20.1\nENABLE_BEAM = True\n\n\n\nif ENABLE_BEAM:\n    BATCH_SIZE = 4\nelse:\n    BATCH_SIZE = 8\n\nDATASET_PATH = '/kaggle/input/iut-comp-dataset/16_kHz_test_audio'\nMODEL = 'Reasat/tugstugi_bengaliai-asr_whisper-medium'    ","metadata":{"execution":{"iopub.status.busy":"2024-06-08T19:24:05.704593Z","iopub.execute_input":"2024-06-08T19:24:05.705464Z","iopub.status.idle":"2024-06-08T19:24:05.711249Z","shell.execute_reply.started":"2024-06-08T19:24:05.705433Z","shell.execute_reply":"2024-06-08T19:24:05.710292Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import csv\nimport pandas as pd\nimport glob\nimport shutil\nimport librosa\nimport argparse\nimport warnings\nfrom pathlib import Path\nimport transformers\nprint(transformers.__version__)\nfrom transformers import pipeline, AutoModelForTokenClassification, AutoTokenizer\n\nimport warnings\n\nwarnings.filterwarnings(\"ignore\")\n\n# files = []\ndf_sub = pd.read_csv('/kaggle/input/iut-comp-dataset/test.csv')\nprint(df_sub.head())\ndf_sub['paths'] = df_sub['file_name'].apply(lambda x: os.path.join(DATASET_PATH, x))\nfiles = df_sub['paths'].to_list()\n# files += list(glob.glob(DATASET_PATH + '/' + '*.mp3'))\nprint('files', len(files))\nprint(files[:3])\n# NOTE: running on a few samples for demonstration\n# files = files[:10]\n\n# files.sort()","metadata":{"execution":{"iopub.status.busy":"2024-06-08T19:26:08.803939Z","iopub.execute_input":"2024-06-08T19:26:08.804839Z","iopub.status.idle":"2024-06-08T19:26:08.843449Z","shell.execute_reply.started":"2024-06-08T19:26:08.804807Z","shell.execute_reply":"2024-06-08T19:26:08.842452Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pipe = pipeline(task=\"automatic-speech-recognition\",\n                model=MODEL,\n                tokenizer=MODEL,\n                chunk_length_s=CHUNK_LENGTH_S, device=0, \n#                 batch_size=BATCH_SIZE\n               )\npipe.model.config.forced_decoder_ids = pipe.tokenizer.get_decoder_prompt_ids(language=\"bn\", task=\"transcribe\")\n\nprint(\"model loaded!\")","metadata":{"execution":{"iopub.status.busy":"2024-06-08T19:27:10.967488Z","iopub.execute_input":"2024-06-08T19:27:10.968130Z","iopub.status.idle":"2024-06-08T19:27:29.157111Z","shell.execute_reply.started":"2024-06-08T19:27:10.968098Z","shell.execute_reply":"2024-06-08T19:27:29.156123Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def fix_repetition(text, max_count):\n    uniq_word_counter = {}\n    words = text.split()\n    for word in text.split():\n        if word not in uniq_word_counter:\n            uniq_word_counter[word] = 1\n        else:\n            uniq_word_counter[word] += 1\n\n    for word, count in uniq_word_counter.items():\n        if count > max_count:\n            words = [w for w in words if w != word]\n    text = \" \".join(words)\n    return text","metadata":{"execution":{"iopub.status.busy":"2024-06-08T19:27:33.445838Z","iopub.execute_input":"2024-06-08T19:27:33.446215Z","iopub.status.idle":"2024-06-08T19:27:33.453148Z","shell.execute_reply.started":"2024-06-08T19:27:33.446187Z","shell.execute_reply":"2024-06-08T19:27:33.452145Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm.auto import tqdm\nimport torch\nimport time\ndef batchify(inputs, batch_size):\n    for i in range(0, len(inputs), batch_size):\n        yield inputs[i:i + batch_size]\n    \n# ENABLE_BEAM = 0\ngenerate_kwargs = {\"max_length\": 260, \"num_beams\": 4} if ENABLE_BEAM else None\nstart = time.time()\ntexts = []\nfor batch in batchify(files, BATCH_SIZE):\n    texts+=pipe(batch, generate_kwargs = generate_kwargs)\n    elapsed = time.time()-start\n    print('completed: {}, average time: {:.2f}'.format(len(texts), elapsed/len(texts)))\nprint('total time: {:.2f}'.format(time.time()-start))","metadata":{"execution":{"iopub.status.busy":"2024-06-08T19:27:34.834900Z","iopub.execute_input":"2024-06-08T19:27:34.836184Z","iopub.status.idle":"2024-06-08T20:17:00.067983Z","shell.execute_reply.started":"2024-06-08T19:27:34.836146Z","shell.execute_reply":"2024-06-08T20:17:00.066299Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# del pipe\nimport torch\nmodels = [\n    AutoModelForTokenClassification.from_pretrained(f).eval().cuda() for f in PUNCT_MODELS\n]\ntokenizer = AutoTokenizer.from_pretrained(PUNCT_MODELS[0])\ndef punctuate(text):\n    input_ids = tokenizer(text).input_ids\n    with torch.no_grad():\n        model = models[0]\n        logits = torch.nn.functional.softmax(\n            model(input_ids=torch.LongTensor([input_ids]).cuda()).logits[0, 1:-1],\n            dim=1).cpu()\n        for model in models[1:]:\n            logits += torch.nn.functional.softmax(\n                model(input_ids=torch.LongTensor([input_ids]).cuda()).logits[0, 1:-1],\n                dim=1).cpu()\n        logits = logits / len(models)\n        logits *= torch.FloatTensor(PUNCT_WEIGHTS)\n        label_ids = torch.argmax(logits, dim=-1)\n\n        tokens = tokenizer(text, add_special_tokens=False).input_ids\n        punct_text = \"\"\n        for index, token in enumerate(tokens):\n            token_str = tokenizer.decode(token)\n            if '##' not in token_str:\n                punct_text += \" \" + token_str\n            else:\n                punct_text += token_str[2:]\n            punct_text += ['', '।', ',', '?'][label_ids[index].item()]\n\n    punct_text = punct_text.strip()\n    return punct_text","metadata":{"execution":{"iopub.status.busy":"2024-06-08T08:08:14.849947Z","iopub.status.idle":"2024-06-08T08:08:14.850305Z","shell.execute_reply.started":"2024-06-08T08:08:14.850132Z","shell.execute_reply":"2024-06-08T08:08:14.850146Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# texts = [{'text': 'text'} for _ in range(len(files))]","metadata":{"execution":{"iopub.status.busy":"2024-06-08T20:18:33.487957Z","iopub.execute_input":"2024-06-08T20:18:33.488401Z","iopub.status.idle":"2024-06-08T20:18:33.494383Z","shell.execute_reply.started":"2024-06-08T20:18:33.488372Z","shell.execute_reply":"2024-06-08T20:18:33.493280Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions = []\nwith open(\"submission.csv\", 'wt', encoding=\"utf8\") as csvfile:\n    writer = csv.writer(csvfile)\n    writer.writerow(['file_name', 'sentence'])\n    for f, text in zip(files, texts):\n        file_id = Path(f).stem\n        pred = text['text'].strip()\n        pred = fix_repetition(pred, max_count=8)\n        if len(pred) == 0:\n            pred = ' '\n#         pred = punctuate(pred)\n#         if pred[-1] not in ['।', '?', ',']:\n#             pred = pred + '।'\n        # print(i, file_id, pred)\n        prediction = [file_id, pred]\n        if len(pred)==0:\n            pred = ' '\n        writer.writerow(prediction)\n        predictions.append(prediction)\nprint(\"inference finished!\")","metadata":{"execution":{"iopub.status.busy":"2024-06-08T20:18:37.950176Z","iopub.execute_input":"2024-06-08T20:18:37.950571Z","iopub.status.idle":"2024-06-08T20:18:37.985165Z","shell.execute_reply.started":"2024-06-08T20:18:37.950534Z","shell.execute_reply":"2024-06-08T20:18:37.984273Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nsubmission = pd.read_csv(\"/kaggle/working/submission.csv\")\nsubmission","metadata":{"execution":{"iopub.status.busy":"2024-06-08T20:28:51.180195Z","iopub.execute_input":"2024-06-08T20:28:51.180882Z","iopub.status.idle":"2024-06-08T20:28:51.199333Z","shell.execute_reply.started":"2024-06-08T20:28:51.180852Z","shell.execute_reply":"2024-06-08T20:28:51.198235Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(submission))","metadata":{"execution":{"iopub.status.busy":"2024-06-08T20:23:12.628871Z","iopub.execute_input":"2024-06-08T20:23:12.629753Z","iopub.status.idle":"2024-06-08T20:23:12.634588Z","shell.execute_reply.started":"2024-06-08T20:23:12.629722Z","shell.execute_reply":"2024-06-08T20:23:12.633651Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"solution = pd.read_csv('/kaggle/input/iut-comp-dataset/test.csv')\nsolution = solution.rename(columns = {'transcripts': 'sentence', 'district': 'domain'})\nsolution['file_name'] = solution['file_name'].apply(lambda x: x.replace('.wav', ''))\nsolution.head()","metadata":{"execution":{"iopub.status.busy":"2024-06-08T20:31:11.478899Z","iopub.execute_input":"2024-06-08T20:31:11.479639Z","iopub.status.idle":"2024-06-08T20:31:11.513196Z","shell.execute_reply.started":"2024-06-08T20:31:11.479608Z","shell.execute_reply":"2024-06-08T20:31:11.512275Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"solution.domain.unique()","metadata":{"execution":{"iopub.status.busy":"2024-06-08T20:31:13.187194Z","iopub.execute_input":"2024-06-08T20:31:13.187887Z","iopub.status.idle":"2024-06-08T20:31:13.194431Z","shell.execute_reply.started":"2024-06-08T20:31:13.187857Z","shell.execute_reply":"2024-06-08T20:31:13.193518Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import jiwer  # you may need to install this library\ndomain_weights = {\n        'Barishal': 0.125,\n        'Chittagong': 0.083,\n        'Habiganj': 0.125,\n        'Kishoreganj': 0.083,\n        'Narail': 0.083, \n        'Narsingdi': 0.083,\n        'Rangpur': 0.083,\n        'Sylhet': 0.125,\n        'Sandwip': 0.125,\n        'Tangail': 0.083,\n    }\ndomain_weights = { key.lower(): value for key, value in domain_weights.items()}\n# unseen: Habiganj, Barishal, Sylhet, Sandwip\ndef mean_wer(solution, submission):\n    joined = solution.merge(submission.rename(columns={'sentence': 'predicted'}))\n#     print(joined)\n    domain_scores = joined.groupby('domain').apply(\n        # note that jiwer.wer computes a weighted average wer by default when given lists of strings\n        lambda df: jiwer.wer(df['sentence'].to_list(), df['predicted'].to_list()),\n    )\n    print(domain_scores)\n    for key, value in domain_weights.items():\n        domain_scores.loc[key] = domain_scores.loc[key].item()*value\n    print(domain_scores)\n    return domain_scores.sum()","metadata":{"execution":{"iopub.status.busy":"2024-06-08T20:31:18.893843Z","iopub.execute_input":"2024-06-08T20:31:18.894172Z","iopub.status.idle":"2024-06-08T20:31:18.902472Z","shell.execute_reply.started":"2024-06-08T20:31:18.894148Z","shell.execute_reply":"2024-06-08T20:31:18.901394Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mean_wer(solution, submission)","metadata":{"execution":{"iopub.status.busy":"2024-06-08T20:31:19.567771Z","iopub.execute_input":"2024-06-08T20:31:19.568140Z","iopub.status.idle":"2024-06-08T20:31:19.688991Z","shell.execute_reply.started":"2024-06-08T20:31:19.568110Z","shell.execute_reply":"2024-06-08T20:31:19.688048Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}