{"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":6172319,"sourceType":"datasetVersion","datasetId":3541690},{"sourceId":6212334,"sourceType":"datasetVersion","datasetId":3567413},{"sourceId":137898680,"sourceType":"kernelVersion"},{"sourceId":137917217,"sourceType":"kernelVersion"}],"dockerImageVersionId":30559,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"- This is a training demo, you can run this code locally, using better GPUs.\n- The inference part is here: [Bengali SR wav2vec_v1_bengali [Inference]](https://www.kaggle.com/takanashihumbert/bengali-sr-wav2vec-v1-bengali-inference), it scores **0.445** on the leaderboard.\n- Feel free to upvote, thanks!","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!cp -r ../input/python-packages2 ./\n\n# Remove existing python-Levenshtein\n!pip uninstall -y python-Levenshtein\n\n# Install python-Levenshtein from the unofficial wheels repository\n!pip install python-Levenshtein==0.12.2\n\n# Install jiwer\n!tar xvfz ./python-packages2/jiwer.tgz\n!pip install ./jiwer/jiwer-2.3.0-py3-none-any.whl -f ./ --no-index\n!pip install -U jiwer\n\n# Install bnunicodenormalizer\n!tar xvfz ./python-packages2/normalizer.tgz\n!pip install ./normalizer/bnunicodenormalizer-0.0.24.tar.gz -f ./ --no-index\n\n# Install pyctcdecode and dependencies\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# Install pypi-kenlm\n!tar xvfz ./python-packages2/pypikenlm.tgz\n!pip install ./pypikenlm/pypi-kenlm-0.1.20220713.tar.gz -f ./ --no-index --no-deps\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-12-24T03:26:33.729706Z","iopub.execute_input":"2023-12-24T03:26:33.730070Z","iopub.status.idle":"2023-12-24T03:28:23.167821Z","shell.execute_reply.started":"2023-12-24T03:26:33.730040Z","shell.execute_reply":"2023-12-24T03:28:23.166714Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch \nimport torch.nn as nn\nimport torchaudio\nimport torchaudio.transforms as tat\nfrom datasets import load_dataset, load_metric, Audio\nimport os\n\nimport typing as tp\nfrom pathlib import Path\nfrom functools import partial\nfrom dataclasses import dataclass, field\nfrom typing import Any, Dict, List, Optional, Union\n\nimport pandas as pd\nimport pyctcdecode\nimport numpy as np\nfrom tqdm.notebook import tqdm\n\nimport librosa\nimport gc\nimport jiwer\nimport pyctcdecode\nimport kenlm\nimport torch\nfrom transformers import Wav2Vec2Processor, Wav2Vec2ProcessorWithLM, Wav2Vec2ForCTC\nfrom transformers import TrainingArguments, Trainer, EarlyStoppingCallback\nfrom bnunicodenormalizer import Normalizer\nimport warnings\nwarnings.filterwarnings('ignore')\ntorchaudio.set_audio_backend(\"soundfile\")","metadata":{"execution":{"iopub.status.busy":"2023-12-24T03:28:23.169750Z","iopub.execute_input":"2023-12-24T03:28:23.170069Z","iopub.status.idle":"2023-12-24T03:28:36.601118Z","shell.execute_reply.started":"2023-12-24T03:28:23.170038Z","shell.execute_reply":"2023-12-24T03:28:36.600327Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"### hyper-parameters\nSR = 16000\ntorch.backends.cudnn.benchmark = True\noutput_dir = \"./\"\nMODEL_PATH = \"/kaggle/input/ai4bharat-indicwav2vec-v1-bengali/indicwav2vec_v1_bengali\"\nLM_PATH = \"/kaggle/input/arijitx-full-model/wav2vec2-xls-r-300m-bengali/language_model\"","metadata":{"execution":{"iopub.status.busy":"2023-12-24T03:28:36.602354Z","iopub.execute_input":"2023-12-24T03:28:36.602691Z","iopub.status.idle":"2023-12-24T03:28:36.608690Z","shell.execute_reply.started":"2023-12-24T03:28:36.602659Z","shell.execute_reply":"2023-12-24T03:28:36.607689Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"processor = Wav2Vec2Processor.from_pretrained(MODEL_PATH)\nvocab_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    str(LM_PATH+\"/unigrams.txt\"),\n)\nprocessor_with_lm = Wav2Vec2ProcessorWithLM(\n    feature_extractor=processor.feature_extractor,\n    tokenizer=processor.tokenizer,\n    decoder=decoder\n)","metadata":{"execution":{"iopub.status.busy":"2023-12-24T03:28:36.611354Z","iopub.execute_input":"2023-12-24T03:28:36.611749Z","iopub.status.idle":"2023-12-24T03:29:07.066001Z","shell.execute_reply.started":"2023-12-24T03:28:36.611708Z","shell.execute_reply":"2023-12-24T03:29:07.065235Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"- From @mbmmurad's [Dataset overlaps with CommonVoice 11 bn](https://www.kaggle.com/code/mbmmurad/dataset-overlaps-with-commonvoice-11-bn), The competition dataset might contain the audios of the mozilla-foundation/common_voice_11_0 dataset. Here I just simply exclude them from the validation set.\n- Also, I use @UmongSain's normalized data [here](https://www.kaggle.com/code/umongsain/macro-normalization/notebook). Thanks to him!","metadata":{}},{"cell_type":"code","source":"sentences = pd.read_csv(\"/kaggle/input/macro-normalization/normalized.csv\")\nindexes = set(pd.read_csv(\"/kaggle/input/dataset-overlaps-with-commonvoice-11-bn/indexes.csv\")['id'])\nprint(len(sentences))\nsentences = sentences[~((sentences.index.isin(indexes))&(sentences['split']=='train'))].reset_index(drop=True)\nprint(len(sentences))","metadata":{"execution":{"iopub.status.busy":"2023-12-24T03:29:07.067296Z","iopub.execute_input":"2023-12-24T03:29:07.067634Z","iopub.status.idle":"2023-12-24T03:29:15.163577Z","shell.execute_reply.started":"2023-12-24T03:29:07.067598Z","shell.execute_reply":"2023-12-24T03:29:15.162666Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"* sample 10% data from \"valid\" part into validation set, 90% into training set.\n* sample 5% data from \"train\" part, and additionally sample 8% from it into validation set, 92% into training set.\n* There will be **57776** train data, **5667** valid data.","metadata":{}},{"cell_type":"code","source":"data_0 = sentences.loc[sentences['split']=='valid'].reset_index(drop=True)\nvalid_0 = data_0.sample(frac=0.1, random_state=42)\ntrain_0 = data_0[~data_0.index.isin(valid_0.index)]\n\ndata_1 = sentences.loc[sentences['split']=='train'].reset_index(drop=True).sample(frac=0.05, random_state=42)\nvalid_1 = data_1.sample(frac=0.08, random_state=42)\ntrain_1 = data_1[~data_1.index.isin(valid_1.index)]\n\ntrain = pd.concat([train_0, train_1], axis=0).sample(frac=1, random_state=42).reset_index(drop=True)\nvalid = pd.concat([valid_0, valid_1], axis=0).sample(frac=1, random_state=42).reset_index(drop=True)\n\ndel data_0, data_1, valid_0, valid_1, train_0, train_1\nall_ids = sentences['id'].to_list()\ntrain_ids = train['id'].to_list()\nvalid_ids = valid['id'].to_list()\n\n# in kaggle notebook, validating is very time-consuming, so here I use a very small validation set, rather than 5667.\nvalid = valid.sample(n=500, random_state=42)\n\nprint(len(all_ids))\nprint(len(train_ids))\nprint(len(valid_ids))","metadata":{"execution":{"iopub.status.busy":"2023-12-24T03:29:15.167790Z","iopub.execute_input":"2023-12-24T03:29:15.168095Z","iopub.status.idle":"2023-12-24T03:29:15.644860Z","shell.execute_reply.started":"2023-12-24T03:29:15.168068Z","shell.execute_reply":"2023-12-24T03:29:15.644016Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class W2v2Dataset(torch.utils.data.Dataset):\n    def __init__(self, df):\n        self.df = df\n        self.pathes = df['id'].values\n        self.sentences = df['normalized'].values\n        self.resampler = tat.Resample(32000, SR)\n\n    def __getitem__(self, idx):\n        apath = f'/kaggle/input/bengaliai-speech/train_mp3s/{self.pathes[idx]}.mp3'\n        waveform, sample_rate = torchaudio.load(apath, format=\"mp3\")\n        waveform = self.resampler(waveform)\n        batch = dict()\n        y = processor(waveform.reshape(-1), sampling_rate=SR).input_values[0] \n        batch[\"input_values\"] = y\n        with processor.as_target_processor():\n            batch[\"labels\"] = processor(self.sentences[idx]).input_ids       \n        \n        return batch\n\n    def __len__(self):\n        return len(self.df)\n\ntrain_dataset = W2v2Dataset(train)\nvalid_dataset = W2v2Dataset(valid)","metadata":{"execution":{"iopub.status.busy":"2023-12-24T03:29:15.646037Z","iopub.execute_input":"2023-12-24T03:29:15.646386Z","iopub.status.idle":"2023-12-24T03:29:15.739569Z","shell.execute_reply.started":"2023-12-24T03:29:15.646351Z","shell.execute_reply":"2023-12-24T03:29:15.738867Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@dataclass\nclass DataCollatorCTCWithPadding:\n    \"\"\"\n    Data collator that will dynamically pad the inputs received.\n    Args:\n        processor (:class:`~transformers.Wav2Vec2Processor`)\n            The processor used for proccessing the data.\n        padding (:obj:`bool`, :obj:`str` or :class:`~transformers.tokenization_utils_base.PaddingStrategy`, `optional`, defaults to :obj:`True`):\n            Select a strategy to pad the returned sequences (according to the model's padding side and padding index)\n            among:\n            * :obj:`True` or :obj:`'longest'`: Pad to the longest sequence in the batch (or no padding if only a single\n              sequence if provided).\n            * :obj:`'max_length'`: Pad to a maximum length specified with the argument :obj:`max_length` or to the\n              maximum acceptable input length for the model if that argument is not provided.\n            * :obj:`False` or :obj:`'do_not_pad'` (default): No padding (i.e., can output a batch with sequences of\n              different lengths).\n        max_length (:obj:`int`, `optional`):\n            Maximum length of the ``input_values`` of the returned list and optionally padding length (see above).\n        max_length_labels (:obj:`int`, `optional`):\n            Maximum length of the ``labels`` returned list and optionally padding length (see above).\n        pad_to_multiple_of (:obj:`int`, `optional`):\n            If set will pad the sequence to a multiple of the provided value.\n            This is especially useful to enable the use of Tensor Cores on NVIDIA hardware with compute capability >=\n            7.5 (Volta).\n    \"\"\"\n\n    processor: Wav2Vec2Processor\n    padding: Union[bool, str] = True\n    max_length: Optional[int] = None\n    max_length_labels: Optional[int] = None\n    pad_to_multiple_of: Optional[int] = None\n    pad_to_multiple_of_labels: Optional[int] = None\n\n    def __call__(self, features: List[Dict[str, Union[List[int], torch.Tensor]]]) -> Dict[str, torch.Tensor]:\n        # split inputs and labels since they have to be of different lenghts and need\n        # different padding methods\n        input_features = [{\"input_values\": feature[\"input_values\"]} for feature in features]\n        label_features = [{\"input_ids\": feature[\"labels\"]} for feature in features]\n\n        batch = self.processor.pad(\n            input_features,\n            padding=self.padding,\n            max_length=self.max_length,\n            pad_to_multiple_of=self.pad_to_multiple_of,\n            return_tensors=\"pt\",\n        )\n        with self.processor.as_target_processor():\n            labels_batch = self.processor.pad(\n                label_features,\n                padding=self.padding,\n                max_length=self.max_length_labels,\n                pad_to_multiple_of=self.pad_to_multiple_of_labels,\n                return_tensors=\"pt\",\n            )\n\n        # replace padding with -100 to ignore loss correctly\n        labels = labels_batch[\"input_ids\"].masked_fill(labels_batch.attention_mask.ne(1), -100)\n\n        batch[\"labels\"] = labels\n\n        return batch","metadata":{"execution":{"iopub.status.busy":"2023-12-24T03:29:15.740924Z","iopub.execute_input":"2023-12-24T03:29:15.741513Z","iopub.status.idle":"2023-12-24T03:29:15.766835Z","shell.execute_reply.started":"2023-12-24T03:29:15.741478Z","shell.execute_reply":"2023-12-24T03:29:15.765868Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_collator = DataCollatorCTCWithPadding(processor=processor, padding=True)","metadata":{"execution":{"iopub.status.busy":"2023-12-24T03:29:15.768181Z","iopub.execute_input":"2023-12-24T03:29:15.768432Z","iopub.status.idle":"2023-12-24T03:29:15.781994Z","shell.execute_reply.started":"2023-12-24T03:29:15.768409Z","shell.execute_reply":"2023-12-24T03:29:15.781114Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"- In kaggle notebook, there is an error: **cannot import name 'compute_measures' from 'jiwer' (unknown location)**. But in my local notebook, there is no such error.","metadata":{}},{"cell_type":"code","source":"##from jiwer import compute_measures\n\n\n#from jiwer import wer\nwer_metric = load_metric(\"wer\")\n\ndef compute_metrics(pred):\n    pred_logits = pred.predictions\n    pred_ids = np.argmax(pred_logits, axis=-1)\n\n    pred.label_ids[pred.label_ids == -100] = processor.tokenizer.pad_token_id\n\n    pred_str = processor.batch_decode(pred_ids)\n    # we do not want to group tokens when computing the metrics\n    label_str = processor.batch_decode(pred.label_ids, group_tokens=False)\n\n    wer = wer_metric.compute(predictions=pred_str, references=label_str)\n\n    return {\"wer\": wer}\n    ","metadata":{"execution":{"iopub.status.busy":"2023-12-24T03:29:15.785362Z","iopub.execute_input":"2023-12-24T03:29:15.785610Z","iopub.status.idle":"2023-12-24T03:29:16.470331Z","shell.execute_reply.started":"2023-12-24T03:29:15.785588Z","shell.execute_reply":"2023-12-24T03:29:16.469523Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = Wav2Vec2ForCTC.from_pretrained(\n    MODEL_PATH,\n    attention_dropout=0.1,\n    hidden_dropout=0.1,\n    feat_proj_dropout=0.0,\n    mask_time_prob=0.05,\n    layerdrop=0.1,\n    #gradient_checkpointing=True, \n    ctc_loss_reduction=\"mean\", \n    pad_token_id=processor.tokenizer.pad_token_id,\n    vocab_size=len(processor.tokenizer),\n    ctc_zero_infinity=True,\n    diversity_loss_weight=100 \n)","metadata":{"execution":{"iopub.status.busy":"2023-12-24T03:29:16.471424Z","iopub.execute_input":"2023-12-24T03:29:16.471698Z","iopub.status.idle":"2023-12-24T03:29:27.037559Z","shell.execute_reply.started":"2023-12-24T03:29:16.471673Z","shell.execute_reply":"2023-12-24T03:29:27.036620Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# you can freeze some params\nmodel.freeze_feature_extractor()\n#model.freeze_feature_encoder()","metadata":{"execution":{"iopub.status.busy":"2023-12-24T03:29:27.038657Z","iopub.execute_input":"2023-12-24T03:29:27.038914Z","iopub.status.idle":"2023-12-24T03:29:27.043916Z","shell.execute_reply.started":"2023-12-24T03:29:27.038891Z","shell.execute_reply":"2023-12-24T03:29:27.042983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"- As a demo, \"**num_train_epochs**\", \"**eval_steps**\" and \"**early_stopping_patience**\" are set to very small values, you can make them larger.\n- If there is no error about jiwer, you can set **metric_for_best_model**=\"wer\", and remember to set **greater_is_better**=False and use **compute_metrics**.","metadata":{}},{"cell_type":"code","source":"training_args = TrainingArguments(\n    output_dir=output_dir,\n    overwrite_output_dir=True,\n    group_by_length=False,\n    lr_scheduler_type='cosine',\n    weight_decay=0.01,\n    per_device_train_batch_size=4,\n    per_device_eval_batch_size=16,\n    gradient_accumulation_steps=1,\n    evaluation_strategy=\"steps\",\n    save_strategy=\"steps\",\n    max_steps=100, # you can change to \"num_train_epochs\"\n    fp16=True,\n    save_steps=20,\n    eval_steps=20,\n    logging_steps=20,\n    learning_rate=2e-5,\n    warmup_steps=600,\n    save_total_limit=1,\n    load_best_model_at_end=True,\n    #metric_for_best_model=\"wer\",\n    #greater_is_better=False,\n    prediction_loss_only=False,\n    auto_find_batch_size=True,\n    report_to=\"none\"\n)","metadata":{"execution":{"iopub.status.busy":"2023-12-24T03:29:27.046437Z","iopub.execute_input":"2023-12-24T03:29:27.046734Z","iopub.status.idle":"2023-12-24T03:29:27.081243Z","shell.execute_reply.started":"2023-12-24T03:29:27.046708Z","shell.execute_reply":"2023-12-24T03:29:27.080454Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer = Trainer(\n    model=model,\n    data_collator=data_collator,\n    args=training_args,\n    #compute_metrics=compute_metrics,\n    train_dataset=train_dataset,\n    eval_dataset=valid_dataset,\n    tokenizer=processor.feature_extractor,\n    callbacks=[EarlyStoppingCallback(early_stopping_patience=1)],\n)","metadata":{"execution":{"iopub.status.busy":"2023-12-24T03:29:27.082247Z","iopub.execute_input":"2023-12-24T03:29:27.082517Z","iopub.status.idle":"2023-12-24T03:29:31.723374Z","shell.execute_reply.started":"2023-12-24T03:29:27.082493Z","shell.execute_reply":"2023-12-24T03:29:31.722573Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.train()\ntrainer.save_model(output_dir)","metadata":{"execution":{"iopub.status.busy":"2023-12-24T03:29:31.724656Z","iopub.execute_input":"2023-12-24T03:29:31.725033Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"- To improve scores you can: \n    * use different pretrained models\n    * alter the parameters\n    * choose more data\n    * filter data in another way.","metadata":{}},{"cell_type":"code","source":"# import jiwer\n\n# # ... (your existing code)\n\n# # Assuming you have your solution and predictions DataFrames defined\n# # solution and predictions should have columns 'id', 'domain', and 'sentence'\n\n# def compute_wer(solution, predictions):\n#     joined = solution.merge(predictions.rename(columns={'sentence': 'predicted'}))\n#     domain_scores = joined.groupby('domain').apply(\n#         lambda df: jiwer.wer(df['sentence'].to_list(), df['predicted'].to_list())\n#     )\n    \n#     mean_wer = domain_scores.mean()\n    \n#     return mean_wer\n\n# # Assuming your model has generated predictions in a DataFrame called 'predictions'\n# # It should have columns 'id' and 'sentence'\n# # Adjust the column names accordingly if needed\n\n# # Make sure you have a DataFrame named 'solution' with columns 'id', 'domain', and 'sentence'\n# # Adjust the column names accordingly if needed\n# wer_results = compute_wer(solution, predictions)\n\n# # Now you can print the WER Results directly\n# print(\"WER Results\")\n# print(f\"The model achieved competitive results in Bengali speech recognition across out-of-distribution domains.\")\n# print(f\"The mean WER, considering the diverse set of domains, reflects the model's effectiveness in handling the linguistic variations present in the test set.\")\n# print(f\"Mean WER: {wer_results}\")\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}