{"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":"markdown","source":"#### Reference\n\nFirstly, Please upvote/refer to [@snnclsr](https://www.kaggle.com/snnclsr) discussions and inference [notebook]https://www.kaggle.com/code/snnclsr/0-444-optimize-decoding-parameters-with-optuna).\n\n\nThird, please upvote this one :)","metadata":{}},{"cell_type":"markdown","source":"## What this notebook features??\n\n- I wanted to showcase the impact of finetuning the models on competition dataset.\n- Current version comprises of finetuned model only with 10% of competition training data.\n- I will publish the training code in upcoming days. You can refer to this [dataset]()\n\nPublic models from hugging faces:\n* `https://huggingface.co/ai4bharat/indicwav2vec_v1_bengali` for Wav2vec2CTC Model only\n* `https://huggingface.co/arijitx/wav2vec2-xls-r-300m-bengali` for Language Model\n\nI didn't trained these models using the competitaion data at all. I just want to know public models score as baseline.  \n\nSo we may get higher and higher score by fine-tuning on competition data.\n\n**Note: I only finetuned the indicwav2vec_v1_bengali which is a CTC model. I am still using the public LM model mentioned above.**\n\n\n### Everything above PLUS\n\n- How to find the best decoding params using optuna on valid dataset.\n\nIt's being suggested by the authors of the [pyctcdecode](https://github.com/kensho-technologies/pyctcdecode/tree/main) developers that we should perform a parameter search because it can improve our results on a specific tasks other than English such as ours.\n\n> (Note: pyctcdecode contains several free hyperparameters that can strongly influence error rate and wall time. Default values for these parameters were (merely) chosen in order to yield good performance for one particular use case. For best results, especially when working with languages other than English, users are encouraged to perform a hyperparameter optimization study on their own data.)\n\nSo we will give it a try to find the best parameters in the validation split of the train dataset (because of the time constraints we will only use 5k). Here are the list of decoding params for easy access:\n\n```python\n# from: https://github.com/kensho-technologies/pyctcdecode/blob/main/pyctcdecode/constants.py\n# default parameters for decoding (can be modified)\nDEFAULT_ALPHA = 0.515\nDEFAULT_BETA = 1.665\nDEFAULT_UNK_LOGP_OFFSET = -10.0\nDEFAULT_BEAM_WIDTH = 100\nDEFAULT_HOTWORD_WEIGHT = 10.0\nDEFAULT_PRUNE_LOGP = -10.0\nDEFAULT_PRUNE_BEAMS = False\nDEFAULT_MIN_TOKEN_LOGP = -5.0\nDEFAULT_SCORE_LM_BOUNDARY = True\n\n# other constants for decoding\nAVG_TOKEN_LEN = 6  # average number of characters expected per token (used for UNK scoring)\nMIN_TOKEN_CLIP_P = 1e-15  # clipping to avoid underflow in case of malformed logit input\nLOG_BASE_CHANGE_FACTOR = 1.0 / math.log10(math.e)  # kenlm returns base10 but we like natural\n```","metadata":{}},{"cell_type":"markdown","source":"## Import","metadata":{}},{"cell_type":"code","source":"!cp -r ../input/python-packages2 ./\n\n!tar xvfz ./python-packages2/jiwer.tgz\n!pip install ./jiwer/jiwer-2.3.0-py3-none-any.whl -f ./ --no-index\n!tar xvfz ./python-packages2/normalizer.tgz\n!pip install ./normalizer/bnunicodenormalizer-0.0.24.tar.gz -f ./ --no-index\n!tar xvfz ./python-packages2/pyctcdecode.tgz\n!pip install ./pyctcdecode/attrs-22.1.0-py2.py3-none-any.whl -f ./ --no-index --no-deps\n!pip install ./pyctcdecode/exceptiongroup-1.0.0rc9-py3-none-any.whl -f ./ --no-index --no-deps\n!pip install ./pyctcdecode/hypothesis-6.54.4-py3-none-any.whl -f ./ --no-index --no-deps\n!pip install ./pyctcdecode/numpy-1.21.6-cp37-cp37m-manylinux_2_12_x86_64.manylinux2010_x86_64.whl -f ./ --no-index --no-deps\n!pip install ./pyctcdecode/pygtrie-2.5.0.tar.gz -f ./ --no-index --no-deps\n!pip install ./pyctcdecode/sortedcontainers-2.4.0-py2.py3-none-any.whl -f ./ --no-index --no-deps\n!pip install ./pyctcdecode/pyctcdecode-0.4.0-py2.py3-none-any.whl -f ./ --no-index --no-deps\n\n!tar xvfz ./python-packages2/pypikenlm.tgz\n!pip install ./pypikenlm/pypi-kenlm-0.1.20220713.tar.gz -f ./ --no-index --no-deps","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-09-12T20:07:46.032807Z","iopub.execute_input":"2023-09-12T20:07:46.033128Z","iopub.status.idle":"2023-09-12T20:08:58.826621Z","shell.execute_reply.started":"2023-09-12T20:07:46.033101Z","shell.execute_reply":"2023-09-12T20:08:58.825449Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install ../input/jiwer-3-0-3/jiwer-3.0.3-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2023-09-12T20:08:58.829417Z","iopub.execute_input":"2023-09-12T20:08:58.829819Z","iopub.status.idle":"2023-09-12T20:09:30.21527Z","shell.execute_reply.started":"2023-09-12T20:08:58.829783Z","shell.execute_reply":"2023-09-12T20:09:30.214065Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rm -r python-packages2 jiwer normalizer pyctcdecode pypikenlm","metadata":{"execution":{"iopub.status.busy":"2023-09-12T20:09:30.217357Z","iopub.execute_input":"2023-09-12T20:09:30.217751Z","iopub.status.idle":"2023-09-12T20:09:31.166558Z","shell.execute_reply.started":"2023-09-12T20:09:30.217717Z","shell.execute_reply":"2023-09-12T20:09:31.165242Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import typing as tp\nfrom pathlib import Path\nfrom functools import partial\nfrom dataclasses import dataclass, field\n\nimport pandas as pd\nimport pyctcdecode\nimport numpy as np\nfrom tqdm.notebook import tqdm\n\nimport librosa\n\nimport pyctcdecode\nimport kenlm\nimport torch\nfrom transformers import Wav2Vec2Processor, Wav2Vec2ProcessorWithLM, Wav2Vec2ForCTC\nfrom bnunicodenormalizer import Normalizer\n\nimport cloudpickle as cpkl","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-09-12T20:09:31.170241Z","iopub.execute_input":"2023-09-12T20:09:31.170639Z","iopub.status.idle":"2023-09-12T20:09:43.52614Z","shell.execute_reply.started":"2023-09-12T20:09:31.170608Z","shell.execute_reply":"2023-09-12T20:09:43.525173Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"FIND_PARAMS = False\n\nROOT = Path.cwd().parent\nINPUT = ROOT / \"input\"\nDATA = INPUT / \"bengaliai-speech\"\nTRAIN = DATA / \"train_mp3s\"\nTEST = DATA / \"test_mp3s\"\n\nSAMPLING_RATE = 16_000\nMODEL_PATH = INPUT / \"bengali-wav2vec2-finetuned/\"\nLM_PATH = INPUT / \"bengali-sr-download-public-trained-models/wav2vec2-xls-r-300m-bengali/language_model/\"","metadata":{"execution":{"iopub.status.busy":"2023-09-12T20:09:43.527633Z","iopub.execute_input":"2023-09-12T20:09:43.527971Z","iopub.status.idle":"2023-09-12T20:09:43.536593Z","shell.execute_reply.started":"2023-09-12T20:09:43.527937Z","shell.execute_reply":"2023-09-12T20:09:43.535536Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### load model, processor, decoder","metadata":{}},{"cell_type":"code","source":"model = Wav2Vec2ForCTC.from_pretrained(MODEL_PATH)\nprocessor = Wav2Vec2Processor.from_pretrained(MODEL_PATH)","metadata":{"execution":{"iopub.status.busy":"2023-09-12T20:09:43.537872Z","iopub.execute_input":"2023-09-12T20:09:43.538399Z","iopub.status.idle":"2023-09-12T20:09:57.944506Z","shell.execute_reply.started":"2023-09-12T20:09:43.538364Z","shell.execute_reply":"2023-09-12T20:09:57.943547Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"vocab_dict = processor.tokenizer.get_vocab()\nsorted_vocab_dict = {k: v for k, v in sorted(vocab_dict.items(), key=lambda item: item[1])}\n\ndecoder = pyctcdecode.build_ctcdecoder(\n    list(sorted_vocab_dict.keys()),\n    str(LM_PATH / \"5gram.bin\"),\n)","metadata":{"execution":{"iopub.status.busy":"2023-09-12T20:09:57.94599Z","iopub.execute_input":"2023-09-12T20:09:57.946416Z","iopub.status.idle":"2023-09-12T20:10:39.791321Z","shell.execute_reply.started":"2023-09-12T20:09:57.946383Z","shell.execute_reply":"2023-09-12T20:10:39.790285Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"processor_with_lm = Wav2Vec2ProcessorWithLM(\n    feature_extractor=processor.feature_extractor,\n    tokenizer=processor.tokenizer,\n    decoder=decoder\n)","metadata":{"execution":{"iopub.status.busy":"2023-09-12T20:10:39.793114Z","iopub.execute_input":"2023-09-12T20:10:39.793838Z","iopub.status.idle":"2023-09-12T20:10:39.801395Z","shell.execute_reply.started":"2023-09-12T20:10:39.793777Z","shell.execute_reply":"2023-09-12T20:10:39.800375Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## prepare dataloader","metadata":{}},{"cell_type":"code","source":"class BengaliSRTestDataset(torch.utils.data.Dataset):\n    \n    def __init__(\n        self,\n        audio_paths: list[str],\n        sampling_rate: int\n    ):\n        self.audio_paths = audio_paths\n        self.sampling_rate = sampling_rate\n        \n    def __len__(self,):\n        return len(self.audio_paths)\n    \n    def __getitem__(self, index: int):\n        audio_path = self.audio_paths[index]\n        sr = self.sampling_rate\n        w = librosa.load(audio_path, sr=sr, mono=False)[0]\n        \n        return w","metadata":{"execution":{"iopub.status.busy":"2023-09-12T20:10:39.80289Z","iopub.execute_input":"2023-09-12T20:10:39.803329Z","iopub.status.idle":"2023-09-12T20:10:39.813786Z","shell.execute_reply.started":"2023-09-12T20:10:39.803298Z","shell.execute_reply":"2023-09-12T20:10:39.812693Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not torch.cuda.is_available():\n    device = torch.device(\"cpu\")\nelse:\n    device = torch.device(\"cuda\")\n\nmodel = model.to(device)\nmodel = model.eval()\nmodel = model.half()","metadata":{"execution":{"iopub.status.busy":"2023-09-12T20:10:39.817925Z","iopub.execute_input":"2023-09-12T20:10:39.818438Z","iopub.status.idle":"2023-09-12T20:10:45.383677Z","shell.execute_reply.started":"2023-09-12T20:10:39.818406Z","shell.execute_reply":"2023-09-12T20:10:45.382535Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Finding the best decoding params","metadata":{}},{"cell_type":"code","source":"import jiwer\n\nbnorm = Normalizer()\n\ndef postprocess(sentence):\n    period_set = set([\".\", \"?\", \"!\", \"।\"])\n    _words = [bnorm(word)['normalized']  for word in sentence.split()]\n    sentence = \" \".join([word for word in _words if word is not None])\n    try:\n        if sentence[-1] not in period_set:\n            sentence+=\"।\"\n    except:\n        sentence = \"।\"\n    return sentence\n\n\ndef score(gts, preds):\n    return jiwer.wer(gts, preds)\n\n\ndef inference(m, data_loader):\n    logits = []\n    with torch.no_grad():\n        for batch in tqdm(data_loader):\n            x = batch[\"input_values\"]\n            x = x.to(device, non_blocking=True)\n            with torch.cuda.amp.autocast(True):\n                y = model(x).logits\n            y = y.detach().cpu().numpy()\n            logits.extend(y)\n    return logits\n\n\ndef decode(logits, params={\"beam_width\": 512}, pp=True):    \n    pred_sentence_list = [processor_with_lm.decode(sentence, **params).text for sentence in tqdm(logits)]\n    if pp:\n        pred_sentence_list = [postprocess(s) for s in pred_sentence_list]\n    return pred_sentence_list","metadata":{"execution":{"iopub.status.busy":"2023-09-12T20:10:45.385511Z","iopub.execute_input":"2023-09-12T20:10:45.385975Z","iopub.status.idle":"2023-09-12T20:10:45.474626Z","shell.execute_reply.started":"2023-09-12T20:10:45.385934Z","shell.execute_reply":"2023-09-12T20:10:45.473623Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"constants = \"\"\"\n# from: https://github.com/kensho-technologies/pyctcdecode/blob/main/pyctcdecode/constants.py\n# default parameters for decoding (can be modified)\nDEFAULT_ALPHA = 0.495\nDEFAULT_BETA = 1.275\nDEFAULT_UNK_LOGP_OFFSET = -10.0\nDEFAULT_BEAM_WIDTH = 100\nDEFAULT_HOTWORD_WEIGHT = 10.0\nDEFAULT_PRUNE_LOGP = -10.0\nDEFAULT_PRUNE_BEAMS = False\nDEFAULT_MIN_TOKEN_LOGP = -5.0\nDEFAULT_SCORE_LM_BOUNDARY = True\n\n# other constants for decoding\nAVG_TOKEN_LEN = 6  # average number of characters expected per token (used for UNK scoring)\nMIN_TOKEN_CLIP_P = 1e-15  # clipping to avoid underflow in case of malformed logit input\nLOG_BASE_CHANGE_FACTOR = 1.0 / math.log10(math.e)  # kenlm returns base10 but we like natural\n\"\"\"","metadata":{"execution":{"iopub.status.busy":"2023-09-12T20:10:45.476029Z","iopub.execute_input":"2023-09-12T20:10:45.476468Z","iopub.status.idle":"2023-09-12T20:10:45.482016Z","shell.execute_reply.started":"2023-09-12T20:10:45.476426Z","shell.execute_reply":"2023-09-12T20:10:45.481104Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def objective(trial):\n    \"\"\"\n    alpha: weight for language model during shallow fusion\n    beta: weight for length score adjustment of during scoring\n    unk_score_offset: amount of log score offset for unknown tokens\n    lm_score_boundary: whether to have kenlm respect boundaries when scoring\n    \"\"\"\n    alpha = trial.suggest_float(\"alpha\", 0.0, 2.15)\n    beta = trial.suggest_float(\"beta\", 0.0, 2.05)\n    beam_width = trial.suggest_categorical(\"beam_width\", [256, 512, 768])\n    gts = valid[\"sentence\"].values.tolist()\n    decode_params = {\n        \"alpha\": alpha,\n        \"beta\": beta,\n        \"beam_width\": beam_width\n    }\n    preds = decode(logits, params=decode_params, pp=True)\n    wer_score = score(gts, preds)\n    return wer_score","metadata":{"execution":{"iopub.status.busy":"2023-09-12T20:10:45.483358Z","iopub.execute_input":"2023-09-12T20:10:45.484277Z","iopub.status.idle":"2023-09-12T20:10:45.495962Z","shell.execute_reply.started":"2023-09-12T20:10:45.484238Z","shell.execute_reply":"2023-09-12T20:10:45.495057Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Default decoding configuration in the public notebook.\nbest_params = {\"beam_width\": 512}\n\nif FIND_PARAMS:\n    import optuna\n    from optuna.trial import TrialState\n    \n    valid = pd.read_csv(DATA / \"train.csv\") # dtype={\"id\": str}\n    valid = valid.query('split==\"valid\"').sample(n=5000, random_state=42).reset_index(drop=True)\n    valid_audio_paths = [str(TRAIN / f\"{aid}.mp3\") for aid in valid[\"id\"].values]\n\n    valid_dataset = BengaliSRTestDataset(\n        valid_audio_paths, SAMPLING_RATE\n    )\n\n    collate_func = partial(\n        processor_with_lm.feature_extractor,\n        return_tensors=\"pt\", sampling_rate=SAMPLING_RATE,\n        padding=True,\n    )\n\n    valid_loader = torch.utils.data.DataLoader(\n        valid_dataset, batch_size=8, shuffle=False,\n        num_workers=2, collate_fn=collate_func, drop_last=False,\n        pin_memory=True,\n    )\n    # Calculating the base score\n    print(constants)\n    logits = inference(model, valid_loader)\n    base_preds = decode(logits)\n    gts = valid[\"sentence\"].values.tolist()\n    base_wer_score = score(gts, base_preds)\n    print(f\"Base wer score: {base_wer_score}\")\n\n    study = optuna.create_study(direction=\"minimize\")\n    study.optimize(objective, n_trials=25)\n\n    pruned_trials = study.get_trials(deepcopy=False, states=[TrialState.PRUNED])\n    complete_trials = study.get_trials(deepcopy=False, states=[TrialState.COMPLETE])\n\n    print(\"Study statistics: \")\n    print(\"  Number of finished trials: \", len(study.trials))\n    print(\"  Number of pruned trials: \", len(pruned_trials))\n    print(\"  Number of complete trials: \", len(complete_trials))\n\n    print(\"Best trial:\")\n    trial = study.best_trial\n\n    print(\"  Value: \", trial.value)\n\n    print(\"  Params: \")\n    for key, value in trial.params.items():\n        print(\"    {}: {}\".format(key, value))\n    \n    if study.best_value < base_wer_score:\n        print(f\"Base score improved to {study.best_value} from {base_wer_score}. Assigning {study.best_params} to best_params\")\n        best_params = study.best_params","metadata":{"execution":{"iopub.status.busy":"2023-09-12T20:10:45.497293Z","iopub.execute_input":"2023-09-12T20:10:45.49817Z","iopub.status.idle":"2023-09-12T20:10:45.512523Z","shell.execute_reply.started":"2023-09-12T20:10:45.498136Z","shell.execute_reply":"2023-09-12T20:10:45.511355Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference with the best params","metadata":{}},{"cell_type":"code","source":"# Please see the Version 3. of this notebook to see the results.\nbest_params = {'alpha': 0.345, 'beta': 0.06, 'beam_width': 768}","metadata":{"execution":{"iopub.status.busy":"2023-09-12T20:10:45.513898Z","iopub.execute_input":"2023-09-12T20:10:45.514433Z","iopub.status.idle":"2023-09-12T20:10:45.527589Z","shell.execute_reply.started":"2023-09-12T20:10:45.514403Z","shell.execute_reply":"2023-09-12T20:10:45.526613Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f\"Running the inference with params: {best_params}\")","metadata":{"execution":{"iopub.status.busy":"2023-09-12T20:10:45.528851Z","iopub.execute_input":"2023-09-12T20:10:45.52929Z","iopub.status.idle":"2023-09-12T20:10:45.538701Z","shell.execute_reply.started":"2023-09-12T20:10:45.529259Z","shell.execute_reply":"2023-09-12T20:10:45.537553Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test = pd.read_csv(DATA / \"sample_submission.csv\", dtype={\"id\": str})\ntest_audio_paths = [str(TEST / f\"{aid}.mp3\") for aid in test[\"id\"].values]\n\ntest_dataset = BengaliSRTestDataset(\n    test_audio_paths, SAMPLING_RATE\n)\ncollate_func = partial(\n    processor_with_lm.feature_extractor,\n    return_tensors=\"pt\", sampling_rate=SAMPLING_RATE,\n    padding=True,\n)\ntest_loader = torch.utils.data.DataLoader(\n    test_dataset, batch_size=8, shuffle=False,\n    num_workers=2, collate_fn=collate_func, drop_last=False,\n    pin_memory=True,\n)\n\npred_sentence_list = []\n\nwith torch.no_grad():\n    for batch in tqdm(test_loader):\n        x = batch[\"input_values\"]\n        x = x.to(device, non_blocking=True)\n        with torch.cuda.amp.autocast(True):\n            y = model(x).logits\n        y = y.detach().cpu().numpy()\n        \n        for l in y:  \n            sentence = processor_with_lm.decode(l, **best_params).text\n            pred_sentence_list.append(sentence)\n\n\npp_pred_sentence_list = [postprocess(s) for s in tqdm(pred_sentence_list)]","metadata":{"execution":{"iopub.status.busy":"2023-09-12T20:10:45.540221Z","iopub.execute_input":"2023-09-12T20:10:45.540679Z","iopub.status.idle":"2023-09-12T20:11:01.110759Z","shell.execute_reply.started":"2023-09-12T20:10:45.540647Z","shell.execute_reply":"2023-09-12T20:11:01.109676Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Make Submission","metadata":{}},{"cell_type":"code","source":"test[\"sentence\"] = pp_pred_sentence_list\ntest.to_csv(\"submission.csv\", index=False)\nprint(test.head())","metadata":{"execution":{"iopub.status.busy":"2023-09-12T20:11:01.112593Z","iopub.execute_input":"2023-09-12T20:11:01.11324Z","iopub.status.idle":"2023-09-12T20:11:01.130296Z","shell.execute_reply.started":"2023-09-12T20:11:01.113202Z","shell.execute_reply":"2023-09-12T20:11:01.129137Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## EOF","metadata":{}}]}