{"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":"# Hugging Face🤗 Wav2Vec2.0 Training Notebook\n### This Notebook is forked my private notebook. Due to the limitations of Kaggle notebook, the following points differ from the original.\n* smaller datasets\n* shorter epochs\n\n## [This code was very helpful in creating this notebook.](https://colab.research.google.com/github/patrickvonplaten/notebooks/blob/master/Fine_Tune_XLSR_Wav2Vec2_on_Turkish_ASR_with_%F0%9F%A4%97_Transformers.ipynb)","metadata":{}},{"cell_type":"markdown","source":"## Import Libralies","metadata":{}},{"cell_type":"code","source":"import torch \nimport torch.nn as nn\nimport torchaudio\nimport torchaudio.transforms as tat\nimport numpy as np\nfrom datasets import load_dataset, load_metric, Audio\nimport os","metadata":{"execution":{"iopub.status.busy":"2023-08-10T05:16:30.213550Z","iopub.execute_input":"2023-08-10T05:16:30.214723Z","iopub.status.idle":"2023-08-10T05:16:34.913563Z","shell.execute_reply.started":"2023-08-10T05:16:30.214678Z","shell.execute_reply":"2023-08-10T05:16:34.912375Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import warnings\nwarnings.filterwarnings('ignore')","metadata":{"execution":{"iopub.status.busy":"2023-08-10T05:16:34.915555Z","iopub.execute_input":"2023-08-10T05:16:34.916246Z","iopub.status.idle":"2023-08-10T05:16:34.922309Z","shell.execute_reply.started":"2023-08-10T05:16:34.916204Z","shell.execute_reply":"2023-08-10T05:16:34.920297Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torchaudio.set_audio_backend(\"soundfile\")\n# Default is set \"sox\", but it cannnot load mp3 (on experiment)","metadata":{"execution":{"iopub.status.busy":"2023-08-10T05:16:34.923876Z","iopub.execute_input":"2023-08-10T05:16:34.924600Z","iopub.status.idle":"2023-08-10T05:16:34.934244Z","shell.execute_reply.started":"2023-08-10T05:16:34.924565Z","shell.execute_reply":"2023-08-10T05:16:34.933255Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Set hyper-parameter","metadata":{}},{"cell_type":"code","source":"### hyper-parameters\nSR = 16000 # Wav2Vec2.0 requires samplerate 16000.\nTRAIN_ALL_WEIGHTS = True\nTRAIN_SIZE = 20000\nNUM_TRAIN_EPOCHS = 13\nPER_DEVICE_TRAIN_BATCH_SIZE = 1 # Long audio samples are present; little batchsize to prevent OOM\nTARGET_BATCHSIZE = 16 # Used for gradient accumulation and lr warmup\nLR = 2e-5 / 16 * TARGET_BATCHSIZE # Using GradAcc, so lr can be made a little larger.\ntorch.backends.cudnn.benchmark = True ","metadata":{"execution":{"iopub.status.busy":"2023-08-10T05:16:34.938643Z","iopub.execute_input":"2023-08-10T05:16:34.938944Z","iopub.status.idle":"2023-08-10T05:16:34.946105Z","shell.execute_reply.started":"2023-08-10T05:16:34.938918Z","shell.execute_reply":"2023-08-10T05:16:34.945134Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Make dataset","metadata":{}},{"cell_type":"code","source":"import pandas as pd\ntrain = pd.read_csv(f'/kaggle/input/hf-wav2vec2-0-preprocess-baseline/train_bengali.csv', index_col=0)\nval = pd.read_csv(f'/kaggle/input/hf-wav2vec2-0-preprocess-baseline/val_bengali.csv', index_col=0)\n\n# make datasets smaller\n\n# train = train[:50000]\nval = val[:300]\ntrain.head()","metadata":{"execution":{"iopub.status.busy":"2023-08-10T05:16:34.947661Z","iopub.execute_input":"2023-08-10T05:16:34.947974Z","iopub.status.idle":"2023-08-10T05:16:39.962236Z","shell.execute_reply.started":"2023-08-10T05:16:34.947939Z","shell.execute_reply":"2023-08-10T05:16:39.961086Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val = val.reset_index()\ntrain = train.reset_index()\nlen(train), len(val)","metadata":{"execution":{"iopub.status.busy":"2023-08-10T05:16:39.964065Z","iopub.execute_input":"2023-08-10T05:16:39.964421Z","iopub.status.idle":"2023-08-10T05:16:40.029878Z","shell.execute_reply.started":"2023-08-10T05:16:39.964389Z","shell.execute_reply":"2023-08-10T05:16:40.028623Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Prepare to train in transformers' Wav2Vec2.0","metadata":{}},{"cell_type":"code","source":"from transformers import Wav2Vec2CTCTokenizer\n\ntokenizer = Wav2Vec2CTCTokenizer(f\"/kaggle/input/hf-wav2vec2-0-preprocess-baseline/vocab_bengali.json\", unk_token=\"[UNK]\", pad_token=\"[PAD]\", word_delimiter_token=\"|\")\ntokenizer","metadata":{"execution":{"iopub.status.busy":"2023-08-10T05:16:40.031917Z","iopub.execute_input":"2023-08-10T05:16:40.032319Z","iopub.status.idle":"2023-08-10T05:16:41.983597Z","shell.execute_reply.started":"2023-08-10T05:16:40.032284Z","shell.execute_reply":"2023-08-10T05:16:41.982652Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import Wav2Vec2FeatureExtractor\n\nfeature_extractor = Wav2Vec2FeatureExtractor(feature_size=1, sampling_rate=SR, padding_value=0.0, do_normalize=True, return_attention_mask=True)\nfeature_extractor","metadata":{"execution":{"iopub.status.busy":"2023-08-10T05:16:41.985182Z","iopub.execute_input":"2023-08-10T05:16:41.985893Z","iopub.status.idle":"2023-08-10T05:16:41.998760Z","shell.execute_reply.started":"2023-08-10T05:16:41.985857Z","shell.execute_reply":"2023-08-10T05:16:41.997500Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import Wav2Vec2Processor\n\nprocessor = Wav2Vec2Processor(feature_extractor=feature_extractor, tokenizer=tokenizer)\nprocessor","metadata":{"execution":{"iopub.status.busy":"2023-08-10T05:16:42.000667Z","iopub.execute_input":"2023-08-10T05:16:42.001201Z","iopub.status.idle":"2023-08-10T05:16:42.016686Z","shell.execute_reply.started":"2023-08-10T05:16:42.001164Z","shell.execute_reply":"2023-08-10T05:16:42.015610Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class W2v2Dataset(torch.utils.data.Dataset):\n    def __init__(self, df, is_train=False):\n        self.df = df\n        self.pathes = df['id'].values\n        self.sentences = df['sentence'].values\n        self.resampler = tat.Resample(32000, SR)\n        self.is_train = is_train\n\n    def __getitem__(self, idx):\n        if self.is_train:\n            idx = torch.randint(0, len(self.df), (1,))[0].numpy()\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        batch[\"input_values\"] = processor(waveform.reshape(-1), sampling_rate=SR).input_values[0]  \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 TRAIN_SIZE if self.is_train else len(self.df)\n\ntrain_dataset = W2v2Dataset(train, is_train=True)\nval_dataset = W2v2Dataset(val)","metadata":{"execution":{"iopub.status.busy":"2023-08-10T05:16:42.021584Z","iopub.execute_input":"2023-08-10T05:16:42.021872Z","iopub.status.idle":"2023-08-10T05:16:42.130836Z","shell.execute_reply.started":"2023-08-10T05:16:42.021846Z","shell.execute_reply":"2023-08-10T05:16:42.129597Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from dataclasses import dataclass, field\nfrom typing import Any, Dict, List, Optional, Union\n\n@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-08-10T05:16:42.132762Z","iopub.execute_input":"2023-08-10T05:16:42.133208Z","iopub.status.idle":"2023-08-10T05:16:42.148495Z","shell.execute_reply.started":"2023-08-10T05:16:42.133167Z","shell.execute_reply":"2023-08-10T05:16:42.147406Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_collator = DataCollatorCTCWithPadding(processor=processor, padding=True)","metadata":{"execution":{"iopub.status.busy":"2023-08-10T05:16:42.150288Z","iopub.execute_input":"2023-08-10T05:16:42.151007Z","iopub.status.idle":"2023-08-10T05:16:42.165377Z","shell.execute_reply.started":"2023-08-10T05:16:42.150968Z","shell.execute_reply":"2023-08-10T05:16:42.164050Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define metric","metadata":{}},{"cell_type":"code","source":"# Using Levenshtein distance as metric\n# (explained with Japanese language)\n#\n# レーベンシュタイン距離を用いて，\n# 認識結果の誤り数を算出します．\n#\n\nimport numpy as np\nimport copy\n\ndef calculate_error(hypothesis, reference):\n    ''' レーベンシュタイン距離を計算し，\n        置換誤り，削除誤り，挿入誤りを出力する\n    hypothesis:       認識結果(トークン毎に区切ったリスト形式)\n    reference:        正解(同上)\n    total_error:      総誤り数\n    substitute_error: 置換誤り数\n    delete_error:     削除誤り数\n    insert_error:     挿入誤り数\n    len_ref:          正解文のトークン数\n    '''\n    # 認識結果および正解系列の長さを取得\n    len_hyp = len(hypothesis)\n    len_ref = len(reference)\n\n    # 累積コスト行列を作成する\n    # 行列の各要素には，トータルコスト，\n    # 置換コスト，削除コスト，挿入コストの\n    # 累積値が辞書形式で定義される．\n    cost_matrix = [[{\"total\":0, \n                     \"substitute\":0,\n                     \"delete\":0,\n                     \"insert\":0} \\\n                     for j in range(len_ref+1)] \\\n                         for i in range(len_hyp+1)]\n\n    # 0列目と0行目の入力\n    for i in range(1, len_hyp+1):\n        # 縦方向への遷移は，削除処理を意味する\n        cost_matrix[i][0][\"delete\"] = i\n        cost_matrix[i][0][\"total\"] = i\n    for j in range(1, len_ref+1):\n        # 横方向への遷移は，挿入処理を意味する\n        cost_matrix[0][j][\"insert\"] = j\n        cost_matrix[0][j][\"total\"] = j\n\n    # 1列目と1行目以降の累積コストを計算していく\n    for i in range(1, len_hyp+1):\n        for j in range(1, len_ref+1):\n            #\n            # 各処理のコストを計算する\n            #\n            # 斜め方向の遷移時，文字が一致しない場合は，\n            # 置換処理により累積コストが1増加\n            substitute_cost = \\\n                cost_matrix[i-1][j-1][\"total\"] \\\n                + (0 if hypothesis[i-1] == reference[j-1] else 1)\n            # 縦方向の遷移時は，削除処理により累積コストが1増加\n            delete_cost = cost_matrix[i-1][j][\"total\"] + 1\n            # 横方向の遷移時は，挿入処理により累積コストが1増加\n            insert_cost = cost_matrix[i][j-1][\"total\"] + 1\n\n            # 置換処理，削除処理，挿入処理のうち，\n            # どの処理を行えば累積コストが最も小さくなるかを計算\n            cost = [substitute_cost, delete_cost, insert_cost]\n            min_index = np.argmin(cost)\n\n            if min_index == 0:\n                # 置換処理が累積コスト最小となる場合\n\n                # 遷移元の累積コスト情報をコピー\n                cost_matrix[i][j] = \\\n                    copy.copy(cost_matrix[i-1][j-1])\n                # 文字が一致しない場合は，\n                # 累積置換コストを1増加させる\n                cost_matrix[i][j][\"substitute\"] \\\n                    += (0 if hypothesis[i-1] \\\n                        == reference[j-1] else 1)\n            elif min_index == 1:\n                # 削除処理が累積コスト最小となる場合\n                \n                # 遷移元の累積コスト情報をコピー\n                cost_matrix[i][j] = copy.copy(cost_matrix[i-1][j])\n                # 累積削除コストを1増加させる\n                cost_matrix[i][j][\"delete\"] += 1\n            else:\n                # 置換処理が累積コスト最小となる場合\n                \n                # 遷移元の累積コスト情報をコピー\n                cost_matrix[i][j] = copy.copy(cost_matrix[i][j-1])\n                # 累積挿入コストを1増加させる\n                cost_matrix[i][j][\"insert\"] += 1\n\n            # 累積トータルコスト(置換+削除+挿入コスト)を更新\n            cost_matrix[i][j][\"total\"] = cost[min_index]\n\n    #\n    # エラーの数を出力する\n    # このとき，削除コストは挿入誤り，\n    # 挿入コストは削除誤りになる点に注意．\n    # (削除コストが1である\n    #    = 1文字削除しないと正解文にならない \n    #    = 認識結果は1文字分余計に挿入されている\n    #    = 挿入誤りが1である)\n    #\n\n    # 累積コスト行列の右下の要素が最終的なコストとなる．\n    total_error = cost_matrix[len_hyp][len_ref][\"total\"]\n    substitute_error = cost_matrix[len_hyp][len_ref][\"substitute\"]\n    # 削除誤り = 挿入コスト\n    delete_error = cost_matrix[len_hyp][len_ref][\"insert\"]\n    # 挿入誤り = 削除コスト\n    insert_error = cost_matrix[len_hyp][len_ref][\"delete\"]\n    \n    # 各誤り数と，正解文の文字数\n    # (誤り率を算出する際に分母として用いる)を出力\n    return (total_error, \n            substitute_error,\n            delete_error,\n            insert_error,\n            len_ref)","metadata":{"execution":{"iopub.status.busy":"2023-08-10T05:16:42.169506Z","iopub.execute_input":"2023-08-10T05:16:42.169845Z","iopub.status.idle":"2023-08-10T05:16:42.191563Z","shell.execute_reply.started":"2023-08-10T05:16:42.169814Z","shell.execute_reply":"2023-08-10T05:16:42.190517Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def mtr_leven(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    # print example\n    for i in range(5):\n        print(pred_str[i])\n        print(label_str[i])\n        \n    print('#' * 50)\n        \n    # 各誤りの総数(エラー率算出時の分子)\n    total_err = 0\n    total_sub = 0\n    total_del = 0\n    total_ins = 0\n    # 正解文の総文字数(エラー率算出時の分母)\n    total_length = 0\n    for i in range(len(pred_str)):\n        (error, substitute, delete, insert, ref_length) \\\n                = calculate_error(pred_str[i], label_str[i])\n\n        # 総誤り数を累積する\n        total_err += error\n        total_sub += substitute\n        total_del += delete\n        total_ins += insert\n        total_length += ref_length\n        \n    err_rate = 100.0 * total_err / total_length\n    sub_rate = 100.0 * total_sub / total_length\n    del_rate = 100.0 * total_del / total_length\n    ins_rate = 100.0 * total_ins / total_length\n\n    return {\"err\": err_rate, 'subr':sub_rate, 'delr':del_rate, 'insr':ins_rate}","metadata":{"execution":{"iopub.status.busy":"2023-08-10T05:16:42.193500Z","iopub.execute_input":"2023-08-10T05:16:42.194136Z","iopub.status.idle":"2023-08-10T05:16:42.207325Z","shell.execute_reply.started":"2023-08-10T05:16:42.194099Z","shell.execute_reply":"2023-08-10T05:16:42.206108Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define model","metadata":{}},{"cell_type":"code","source":"def get_wav2vec2_model():\n    from transformers import Wav2Vec2ForCTC\n\n    model = Wav2Vec2ForCTC.from_pretrained(\n        'facebook/wav2vec2-large-xlsr-53',\n        attention_dropout=0.2,\n        hidden_dropout=0.2,\n        feat_proj_dropout=0.2,\n        mask_time_prob=0.1,\n        layerdrop=0.2,\n        ctc_loss_reduction=\"mean\", \n        ctc_zero_infinity=True, # dark magic to avoid nan \n        pad_token_id=processor.tokenizer.pad_token_id,\n        diversity_loss_weight=100 # dark magic to avoid nan \n    )\n    model.lm_head = nn.Linear(1024, 112)\n    model.config.vocab_size = 112\n    return model","metadata":{"execution":{"iopub.status.busy":"2023-08-10T05:16:42.209136Z","iopub.execute_input":"2023-08-10T05:16:42.209525Z","iopub.status.idle":"2023-08-10T05:16:42.222984Z","shell.execute_reply.started":"2023-08-10T05:16:42.209472Z","shell.execute_reply":"2023-08-10T05:16:42.222046Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = get_wav2vec2_model()\nmodel_file = f'/kaggle/input/yellowking-dlsprint-model/YellowKing_model/pytorch_model.bin'\nmodel.load_state_dict(torch.load(model_file))\nmodel.lm_head = nn.Linear(1024, len(processor.tokenizer))\nmodel.config.vocab_size = len(processor.tokenizer)","metadata":{"execution":{"iopub.status.busy":"2023-08-10T05:16:42.225670Z","iopub.execute_input":"2023-08-10T05:16:42.226022Z","iopub.status.idle":"2023-08-10T05:17:26.520067Z","shell.execute_reply.started":"2023-08-10T05:16:42.225989Z","shell.execute_reply":"2023-08-10T05:17:26.518868Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Print model details","metadata":{}},{"cell_type":"code","source":"model.config","metadata":{"execution":{"iopub.status.busy":"2023-08-10T05:17:26.521868Z","iopub.execute_input":"2023-08-10T05:17:26.522248Z","iopub.status.idle":"2023-08-10T05:17:26.532063Z","shell.execute_reply.started":"2023-08-10T05:17:26.522213Z","shell.execute_reply":"2023-08-10T05:17:26.530867Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model","metadata":{"execution":{"iopub.status.busy":"2023-08-10T05:17:26.533843Z","iopub.execute_input":"2023-08-10T05:17:26.534235Z","iopub.status.idle":"2023-08-10T05:17:26.550197Z","shell.execute_reply.started":"2023-08-10T05:17:26.534199Z","shell.execute_reply":"2023-08-10T05:17:26.549011Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"total_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\ntotal_params","metadata":{"execution":{"iopub.status.busy":"2023-08-10T05:17:26.551918Z","iopub.execute_input":"2023-08-10T05:17:26.552278Z","iopub.status.idle":"2023-08-10T05:17:26.568199Z","shell.execute_reply.started":"2023-08-10T05:17:26.552245Z","shell.execute_reply":"2023-08-10T05:17:26.567235Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if TRAIN_ALL_WEIGHTS:\n    for param in model.parameters():\n        param.requires_grad = True\nelse:\n    model.freeze_feature_extractor()","metadata":{"execution":{"iopub.status.busy":"2023-08-10T05:17:26.570958Z","iopub.execute_input":"2023-08-10T05:17:26.572141Z","iopub.status.idle":"2023-08-10T05:17:26.580486Z","shell.execute_reply.started":"2023-08-10T05:17:26.572108Z","shell.execute_reply":"2023-08-10T05:17:26.579371Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train settings","metadata":{}},{"cell_type":"code","source":"warmup_steps = (NUM_TRAIN_EPOCHS / 20) * TRAIN_SIZE // TARGET_BATCHSIZE\nnum_total_steps = NUM_TRAIN_EPOCHS * TRAIN_SIZE // TARGET_BATCHSIZE","metadata":{"execution":{"iopub.status.busy":"2023-08-10T05:17:26.581783Z","iopub.execute_input":"2023-08-10T05:17:26.582289Z","iopub.status.idle":"2023-08-10T05:17:26.592592Z","shell.execute_reply.started":"2023-08-10T05:17:26.582244Z","shell.execute_reply":"2023-08-10T05:17:26.591561Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import TrainingArguments\ntraining_args = TrainingArguments(\n  output_dir=f\"./ASR_Bengali\",\n  group_by_length=False,\n  per_device_train_batch_size=PER_DEVICE_TRAIN_BATCH_SIZE,\n  gradient_accumulation_steps=TARGET_BATCHSIZE // PER_DEVICE_TRAIN_BATCH_SIZE,\n  per_device_eval_batch_size=PER_DEVICE_TRAIN_BATCH_SIZE,\n  evaluation_strategy=\"epoch\",\n  num_train_epochs=NUM_TRAIN_EPOCHS,\n  save_strategy='epoch',\n  fp16_full_eval=False,\n  fp16=False,\n  fp16_backend=False,\n  half_precision_backend=False,\n  logging_steps=10,\n  learning_rate=LR,\n  warmup_steps=warmup_steps,\n  save_total_limit=3,\n  weight_decay=1e-5,\n  dataloader_num_workers=os.cpu_count()*20,\n  prediction_loss_only=False,\n  lr_scheduler_type='linear',\n  report_to='none',\n  auto_find_batch_size=True,\n)\ntraining_args","metadata":{"execution":{"iopub.status.busy":"2023-08-10T05:17:26.594296Z","iopub.execute_input":"2023-08-10T05:17:26.594720Z","iopub.status.idle":"2023-08-10T05:17:26.664181Z","shell.execute_reply.started":"2023-08-10T05:17:26.594689Z","shell.execute_reply":"2023-08-10T05:17:26.663013Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import Trainer\n\ntrainer = Trainer(\n    model=model,\n    data_collator=data_collator,\n    args=training_args,\n    compute_metrics=mtr_leven,\n    train_dataset=train_dataset,\n    eval_dataset=val_dataset,\n    tokenizer=processor.feature_extractor\n)\ntrainer","metadata":{"execution":{"iopub.status.busy":"2023-08-10T05:17:26.667191Z","iopub.execute_input":"2023-08-10T05:17:26.667894Z","iopub.status.idle":"2023-08-10T05:17:33.260068Z","shell.execute_reply.started":"2023-08-10T05:17:26.667852Z","shell.execute_reply":"2023-08-10T05:17:33.258980Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Let's train!","metadata":{}},{"cell_type":"code","source":"trainer.train()","metadata":{"execution":{"iopub.status.busy":"2023-08-10T05:17:33.261607Z","iopub.execute_input":"2023-08-10T05:17:33.261957Z"},"trusted":true},"execution_count":null,"outputs":[]}]}