{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":52324,"databundleVersionId":6229904,"sourceType":"competition"},{"sourceId":4143520,"sourceType":"datasetVersion","datasetId":2447262},{"sourceId":6707460,"sourceType":"datasetVersion","datasetId":3865741}],"dockerImageVersionId":30528,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"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":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-07-02T18:13:34.229534Z","iopub.execute_input":"2024-07-02T18:13:34.229837Z","iopub.status.idle":"2024-07-02T18:14:05.956411Z","shell.execute_reply.started":"2024-07-02T18:13:34.229809Z","shell.execute_reply":"2024-07-02T18:14:05.955103Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport csv\nimport time\nimport glob\n\nMODEL = '/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]\nCHUNK_LENGTH_S = 20.1\nENABLE_BEAM = True\n# none alone 0.9 0.037964275588318684\n# none alone 0.7 0.03592288063510065\n# none alone 0.4 0.03416501275871846\nPUNCT_WEIGHTS = [[1.0, 1.4, 1.0, 0.8]]\n\nif ENABLE_BEAM:\n    BATCH_SIZE = 4\nelse:\n    BATCH_SIZE = 8\n\nif len(glob.glob(\"/kaggle/input/bengaliai-speech/test_mp3s/*.mp3\")) > 10:\n    EVAL = False\n    DATASET_PATH = '/kaggle/input/bengaliai-speech/test_mp3s/'\nelse:\n    EVAL = True\n    DATASET_PATH = '/kaggle/input/bengaliai-speech/test_mp3s/'\n    \nimport csv\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\nfiles = list(glob.glob(DATASET_PATH + '/' + '*.wav'))\nfiles += list(glob.glob(DATASET_PATH + '/' + '*.mp3'))\nfiles.sort()\n\npipe = pipeline(task=\"automatic-speech-recognition\",\n                model=MODEL,\n                tokenizer=MODEL,\n                chunk_length_s=CHUNK_LENGTH_S, device=0, batch_size=BATCH_SIZE)\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-07-02T18:14:05.958663Z","iopub.execute_input":"2024-07-02T18:14:05.958958Z","iopub.status.idle":"2024-07-02T18:15:11.274563Z","shell.execute_reply.started":"2024-07-02T18:14:05.958932Z","shell.execute_reply":"2024-07-02T18:15:11.273503Z"},"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-07-02T18:15:11.276682Z","iopub.execute_input":"2024-07-02T18:15:11.276980Z","iopub.status.idle":"2024-07-02T18:15:11.283385Z","shell.execute_reply.started":"2024-07-02T18:15:11.276954Z","shell.execute_reply":"2024-07-02T18:15:11.282319Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if ENABLE_BEAM:\n    texts = pipe(files, generate_kwargs={\"max_length\": 260, \"num_beams\": 4})\nelse:\n    texts = pipe(files)","metadata":{"execution":{"iopub.status.busy":"2024-07-02T18:15:11.286180Z","iopub.execute_input":"2024-07-02T18:15:11.286674Z","iopub.status.idle":"2024-07-02T18:15:16.494255Z","shell.execute_reply.started":"2024-07-02T18:15:11.286639Z","shell.execute_reply":"2024-07-02T18:15:16.493325Z"},"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-07-02T18:15:16.495407Z","iopub.execute_input":"2024-07-02T18:15:16.495690Z","iopub.status.idle":"2024-07-02T18:16:14.159598Z","shell.execute_reply.started":"2024-07-02T18:15:16.495667Z","shell.execute_reply":"2024-07-02T18:16:14.158689Z"},"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(['id', '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        pred = punctuate(pred)\n        if pred[-1] not in ['।', '?', ',']:\n            pred = pred + '।'\n        # print(i, file_id, pred)\n        prediction = [file_id, pred]\n        writer.writerow(prediction)\n        predictions.append(prediction)\nprint(\"inference finished!\")","metadata":{"execution":{"iopub.status.busy":"2024-07-02T18:16:14.160803Z","iopub.execute_input":"2024-07-02T18:16:14.161083Z","iopub.status.idle":"2024-07-02T18:16:14.291548Z","shell.execute_reply.started":"2024-07-02T18:16:14.161060Z","shell.execute_reply":"2024-07-02T18:16:14.290550Z"},"trusted":true},"execution_count":null,"outputs":[]}]}