{"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":6707460,"sourceType":"datasetVersion","datasetId":3865741},{"sourceId":8068743,"sourceType":"datasetVersion","datasetId":4760663},{"sourceId":8165253,"sourceType":"datasetVersion","datasetId":4815388},{"sourceId":8170150,"sourceType":"datasetVersion","datasetId":4812155},{"sourceId":8149165,"sourceType":"datasetVersion","datasetId":4819425}],"dockerImageVersionId":30684,"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 -q git+https://github.com/huggingface/accelerate.git@main\n# !pip install -q noisereduce","metadata":{"execution":{"iopub.status.busy":"2024-04-20T04:19:24.290027Z","iopub.execute_input":"2024-04-20T04:19:24.290472Z","iopub.status.idle":"2024-04-20T04:20:39.305018Z","shell.execute_reply.started":"2024-04-20T04:19:24.290431Z","shell.execute_reply":"2024-04-20T04:20:39.303830Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# imports\nimport 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":{"execution":{"iopub.status.busy":"2024-04-20T04:20:39.308245Z","iopub.execute_input":"2024-04-20T04:20:39.308688Z","iopub.status.idle":"2024-04-20T04:21:02.596110Z","shell.execute_reply.started":"2024-04-20T04:20:39.308646Z","shell.execute_reply":"2024-04-20T04:21:02.595133Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# BASE_DIR = '/kaggle/input/bengali-asr-dialect-dataset/dataset'\nBASE_DIR = '/kaggle/input/bengali-asr-dialect-augmented-dataset/dataset'\nval_data_dir = f'{BASE_DIR}/validation/validation_audio/'\nval_label = f'{BASE_DIR}/validation/validation.csv'\nval_df = pd.read_csv(val_label)\nval_df.shape","metadata":{"execution":{"iopub.status.busy":"2024-04-20T04:21:02.597276Z","iopub.execute_input":"2024-04-20T04:21:02.597824Z","iopub.status.idle":"2024-04-20T04:21:02.638606Z","shell.execute_reply.started":"2024-04-20T04:21:02.597797Z","shell.execute_reply":"2024-04-20T04:21:02.637579Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2024-04-20T04:21:02.641074Z","iopub.execute_input":"2024-04-20T04:21:02.641358Z","iopub.status.idle":"2024-04-20T04:21:02.647579Z","shell.execute_reply.started":"2024-04-20T04:21:02.641334Z","shell.execute_reply":"2024-04-20T04:21:02.646701Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_id='/kaggle/input/whisper-model/bengali-whisper-medium-pretrained-augmented-dataset-lora-added'\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)","metadata":{"execution":{"iopub.status.busy":"2024-04-20T04:21:02.648742Z","iopub.execute_input":"2024-04-20T04:21:02.649105Z","iopub.status.idle":"2024-04-20T04:21:42.400533Z","shell.execute_reply.started":"2024-04-20T04:21:02.649079Z","shell.execute_reply":"2024-04-20T04:21:42.399511Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_path = '/kaggle/input/bengali-whisper-medium-pretrained-lora-added/bengali-whisper-medium-pretrained-lora-added'   # add your new trained model \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])","metadata":{"execution":{"iopub.status.busy":"2024-04-20T04:21:42.401936Z","iopub.execute_input":"2024-04-20T04:21:42.402249Z","iopub.status.idle":"2024-04-20T04:22:19.201252Z","shell.execute_reply.started":"2024-04-20T04:21:42.402223Z","shell.execute_reply":"2024-04-20T04:22:19.200432Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# CHUNK_LENGTH_S = 20.1\n# ENABLE_BEAM = True\nPUNCT_WEIGHTS = [[1.0, 1.4, 1.0, 0.8]]\n# if ENABLE_BEAM: BATCH_SIZE = 4\n# else: BATCH_SIZE = 8\n\n# torch.cuda.empty_cache()\n# pipe = pipeline(\"automatic-speech-recognition\",model=model_path,tokenizer=model_path,chunk_length_s=CHUNK_LENGTH_S,device=\"cuda\")\n# pipe.model.config.forced_decoder_ids = pipe.tokenizer.get_decoder_prompt_ids(language=\"bn\", task=\"transcribe\")\n\n# print('Model Loaded')","metadata":{"execution":{"iopub.status.busy":"2024-04-20T04:22:19.202416Z","iopub.execute_input":"2024-04-20T04:22:19.202712Z","iopub.status.idle":"2024-04-20T04:22:19.207473Z","shell.execute_reply.started":"2024-04-20T04:22:19.202687Z","shell.execute_reply":"2024-04-20T04:22:19.206580Z"},"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\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# 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        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-04-20T04:22:19.209031Z","iopub.execute_input":"2024-04-20T04:22:19.209367Z","iopub.status.idle":"2024-04-20T04:22:19.223082Z","shell.execute_reply.started":"2024-04-20T04:22:19.209341Z","shell.execute_reply":"2024-04-20T04:22:19.222225Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"selected_files = []\nfor root, dirs, files in os.walk(val_data_dir):\n    files = sorted(files, key=pretty_sort)\n    selected_files = [os.path.join(root, f) for f in files]\n#     selected_files.extend([file for file in filews if os.path.basename(file) in sample_val_df['file_name'].values])\nprint('Selected files: ', len(selected_files))\n# selected_files=selected_files[:10]","metadata":{"execution":{"iopub.status.busy":"2024-04-20T04:22:19.224164Z","iopub.execute_input":"2024-04-20T04:22:19.224556Z","iopub.status.idle":"2024-04-20T04:22:19.356235Z","shell.execute_reply.started":"2024-04-20T04:22:19.224524Z","shell.execute_reply":"2024-04-20T04:22:19.355348Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import tqdm\n\n# Define a function to perform inference with a progress bar\ndef inference_with_progress(pipe, files):\n    texts = []\n    # Iterate over each file and perform inference\n    for file in tqdm.tqdm(files, desc=\"Inference Progress\"):\n        text = pipe(file)\n        texts.append(text)\n    return texts\n\n# Use the function to perform inference\ntexts = inference_with_progress(pipe, selected_files)\nprint('Validation Inference Done')","metadata":{"execution":{"iopub.status.busy":"2024-04-20T04:22:19.359538Z","iopub.execute_input":"2024-04-20T04:22:19.359868Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_pred = []\n\nwith open(\"compare.csv\", 'wt', encoding=\"utf8\") as csvfile:\n    writer = csv.writer(csvfile)\n    writer.writerow(['id', 'actual', 'predicted', 'region'])\n    for f, text in zip(selected_files, texts):\n        file_id = f.split('/')[-1]\n        region = extract_regions(file_id)\n        pred = text['text'].strip()\n        print(file_id)\n        actual = val_df[val_df['file_name'] == file_id]['transcripts'].to_list()[0]\n        print('\\nActual: ', actual)\n\n        pred = fix_repetition(pred, max_count=8)\n        pred = punctuate(pred)\n        try:\n            if pred[-1] not in ['।', '?', ',']:\n                pred = pred + '।'\n        except:\n            pred = pred +'<>'\n        \n        print('Predic: ', pred)\n        print()\n        # print(i, file_id, pred)\n        prediction = [file_id, actual, pred, region]\n        writer.writerow(prediction)\n        val_pred.append(prediction)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"comparison = pd.read_csv('/kaggle/working/compare.csv')\ncomparison['region'] = comparison['region'].str.capitalize()\nsolution = comparison[['id', 'actual', 'region']]\nsubmission = comparison[['id', 'predicted']]\nsolution = solution.rename(columns={'actual': 'sentence', 'region': 'domain'})\nsubmission = submission.rename(columns={'predicted': 'sentence'})","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"domain_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}\n\ndef mean_wer(solution, submission):\n   \n    joined = solution.merge(submission.rename(columns={'sentence': 'predicted'}))\n    domain_scores = joined.groupby('domain').apply(\n        lambda df: (\n#             print('Actual: ', df['sentence'].to_list()),\n#             print('Prediction: ',df['predicted'].to_list()),\n#             print(wer(df['sentence'].to_list(), df['predicted'].to_list())),\n            wer(df['sentence'].to_list(), df['predicted'].to_list())\n        )\n    )\n    for key, value in domain_weights.items():\n        domain_scores.loc[key] = domain_scores.loc[key].item()*value\n    print(domain_scores.sort_values(ascending=False))\n    return domain_scores.sum()\n\nprint(\"WER: \", mean_wer(solution, submission))","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}