{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":73047,"databundleVersionId":8149390,"sourceType":"competition"},{"sourceId":4143520,"sourceType":"datasetVersion","datasetId":2447262},{"sourceId":6707460,"sourceType":"datasetVersion","datasetId":3865741},{"sourceId":8068743,"sourceType":"datasetVersion","datasetId":4760663},{"sourceId":8072194,"sourceType":"datasetVersion","datasetId":4763163},{"sourceId":8165253,"sourceType":"datasetVersion","datasetId":4815388},{"sourceId":8209142,"sourceType":"datasetVersion","datasetId":4812155},{"sourceId":154204277,"sourceType":"kernelVersion"}],"dockerImageVersionId":30674,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install jiwer\n!pip install accelerate\n!pip install -q bitsandbytes\n!pip install -q git+https://github.com/huggingface/peft.git@main\n!pip install accelerate","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **0. Imports**","metadata":{}},{"cell_type":"code","source":"import os\nimport pandas as pd\nimport csv\nimport librosa\nimport torch\nfrom IPython.display import Audio\nfrom pathlib import Path\nfrom tqdm import tqdm \nfrom transformers import (\n    pipeline, \n    AutoModelForTokenClassification, \n    AutoTokenizer,BitsAndBytesConfig,\n    WhisperForConditionalGeneration,\n    WhisperTokenizer,\n    WhisperProcessor\n)\nfrom jiwer import wer\nfrom peft import (\n    prepare_model_for_kbit_training,\n    LoraConfig, \n    PeftModel, \n    LoraModel, \n    LoraConfig, \n    get_peft_model,\n    PeftConfig\n)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **1. Paths**","metadata":{}},{"cell_type":"code","source":"BASE_DIR = '/kaggle/input/bengali-asr-dialect-dataset/dataset'\ntest_data_dir = f\"{BASE_DIR}/test/test_audio/\"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **2. Setup Pipeline**","metadata":{}},{"cell_type":"code","source":"quant_config = BitsAndBytesConfig(\n    load_in_4bit=True,\n    bnb_4bit_quant_type=\"nf4\",\n    bnb_4bit_use_double_quant=True,\n    bnb_4bit_compute_dtype=torch.bfloat16,\n)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Change this with path of new models\nmodel_id='/kaggle/input/whisper-model/whisper-medium-tugisugi+lora+full-data+2-epoch'\nTASK='transcribe'\npeft_config = PeftConfig.from_pretrained(model_id)\nmodel = WhisperForConditionalGeneration.from_pretrained(peft_config.base_model_name_or_path, quantization_config=quant_config,attn_implementation=\"sdpa\", device_map=\"auto\")\nmodel = PeftModel.from_pretrained(model, model_id)\nmodel.config.use_cache = True\ntokenizer = WhisperTokenizer.from_pretrained(peft_config.base_model_name_or_path, language='bn', task=TASK)\nprocessor = WhisperProcessor.from_pretrained(peft_config.base_model_name_or_path, language='bn', task=TASK)\npipe = pipeline(\n    \"automatic-speech-recognition\",\n    model=model,\n    tokenizer=processor.tokenizer,\n    feature_extractor=processor.feature_extractor, \n    chunk_length_s=30, # chunk 15 secs\n    torch_dtype=torch.float16\n)\npipe.model.config.forced_decoder_ids = pipe.tokenizer.get_decoder_prompt_ids(language=\"bn\", task=TASK)\nprint('Model Loaded')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"peft_config.base_model_name_or_path","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **4. Pre-Process Input Data**","metadata":{}},{"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\n\ndef extract_regions(path):\n    return path.split('_')[1].split(' ')[0]\n\ndef pretty_sort(filename):\n    name, number_str = filename.split(\" (\")\n    number = int(number_str.split(\")\")[0])\n    return name, number \n\n# setup the punc model here\nPUNCT_WEIGHTS = [[1.0, 1.4, 1.0, 0.8]]\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]\nmodels = [AutoModelForTokenClassification.from_pretrained(f).eval().cuda() for f in PUNCT_MODELS]\ntokenizer = AutoTokenizer.from_pretrained(PUNCT_MODELS[0])\nprint('Punc Model Setup Done')\n\n# restore punctutaion on inferred text\ndef punctuate(text):\n    input_ids = tokenizer(text).input_ids\n    with torch.no_grad():\n        model = models[0]\n        \n        logits = torch.nn.functional.softmax(\n            model(input_ids=torch.LongTensor([input_ids]).cuda()).logits[0, 1:-1],\n            dim=1).cpu()\n\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            \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        \n        for index, token in enumerate(tokens):\n            token_str = tokenizer.decode(token)\n            if token_str == '>':\n                punct_text += token_str\n\n            elif '##' not in token_str:\n                punct_text += \" \" + token_str\n            else:\n                punct_text += token_str[2:]\n                \n            punct_text += ['', '।', ',', '?'][label_ids[index].item()]\n\n    punct_text = punct_text.strip()\n    return punct_text","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# clear CUDA Cache\ntorch.cuda.empty_cache()\nimport gc\ngc.collect()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **6. Inference**","metadata":{}},{"cell_type":"code","source":"# put swandip audios first\nfiles = []\nfor root, dirs, files in os.walk(test_data_dir):\n    files = sorted(files, key=pretty_sort)\n    shift = files[1070 : 1202]\n    files = shift + files[:1070] + files[1202:]\n    files = [test_data_dir+f for f in files]\n#     files = files[:10]\n\n# files[:10]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(files)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# inference\ntexts = pipe(files)  \nprint('Test set Inference Done')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(texts[:10])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# save inferred texts (without any post processing)\npredictions = []\nwith open(\"pure_inference_data.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        \n        file_id = f.split('/')[-1]\n        pred = text['text'].strip()\n        \n        prediction = [file_id, pred]\n        writer.writerow(prediction)\n        predictions.append(prediction)\n          \nprint(\"output saved without post processing\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# save inferred texts (with any post processing)\npredictions = []\n# texts = df['sentence'][:10]\nwith open(\"output.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 = f.split('/')[-1]\n        \n        pred = text['text'].strip()\n#         print('\\nNormal: ', pred)\n        \n        pred = fix_repetition(pred, max_count=8)\n#         print('Fix Repeatation: ', pred)\n        \n        pred = punctuate(pred)\n#         print('Fix Punctuation: ', pred)\n        \n        try:\n            if pred[-1] not in ['।', '?', ',']:\n                pred = pred + '।'\n        except:\n            pred = '<>'\n            \n        # print(i, file_id, pred)\n        prediction = [file_id, pred]\n        writer.writerow(prediction)\n        predictions.append(prediction)\n        \n        \nprint(\"output saved with post processing\", len(predictions))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **7. Create Submission CSV file**","metadata":{}},{"cell_type":"code","source":"out_df = pd.read_csv('/kaggle/working/output.csv')\nout_df.head()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **8. Save submission file**","metadata":{}},{"cell_type":"code","source":"out_df.to_csv(\"submission.csv\",index=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from IPython.display import FileLink\nFileLink(r'submission.csv')","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}