{"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":"# This notebook will Train Data(with Audio Augmentation)","metadata":{}},{"cell_type":"code","source":"# split dataset to train and validation set based on split column \nfrom datasets import load_dataset,  Dataset, Audio\nimport pandas as pd \npd_input = pd.read_csv('/kaggle/input/bengaliai-speech/train.csv') \npd_input = pd_input.assign(audio ='/kaggle/input/bengaliai-speech/train_mp3s/'+pd_input['id']+'.mp3') \ngrouped = pd_input.groupby(['split'])\ntrainset = grouped.get_group(\"train\")\nvalidationset=grouped.get_group(\"valid\")\naudio_val_dataset = Dataset.from_dict({\"audio\": validationset['audio'], \"sentence\":validationset[\"sentence\"]})\naudio_dataset = Dataset.from_dict({\"audio\": trainset['audio'], \"sentence\":trainset[\"sentence\"]})\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#loading pretrained model \nfrom transformers import WhisperFeatureExtractor, WhisperTokenizer, WhisperProcessor, WhisperForConditionalGeneration\nmodel_nm = 'openai/whisper-small'\nfeature_extractor = WhisperFeatureExtractor.from_pretrained(model_nm,language=\"bengali\", task=\"transcribe\")\ntokenizer = WhisperTokenizer.from_pretrained(model_nm, language=\"bengali\", task=\"transcribe\")\nprocessor = WhisperProcessor.from_pretrained(model_nm, language=\"bengali\", task=\"transcribe\")\nmodel = WhisperForConditionalGeneration.from_pretrained(model_nm)\nmodel.config.forced_decoder_ids = None\nmodel.config.suppress_tokens = []","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# music dataset prepration\nfrom datasets import Dataset, Audio\nimport pandas as pd\nimport os\nmusic_files = os.listdir(\"/kaggle/input/musan-music\")        \npd_test_input = pd.DataFrame(\n                             {\n                                \"audio\": ['/kaggle/input/musan-music/'+filename for filename in music_files]\n                                \n                             }\n                            )\naudio_music_dataset = Dataset.from_dict({\"audio\": pd_test_input['audio']})\n\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# noise dataset prepration\nfrom datasets import Dataset, Audio\nimport pandas as pd\nimport os\nnoise_files = os.listdir(\"/kaggle/input/musan-noise\") \npd_test_input = pd.DataFrame(\n                             {\n                                \"audio\": ['/kaggle/input/musan-noise/'+filename for filename in noise_files]\n                                \n                             }\n                            )\naudio_noise_dataset = Dataset.from_dict({\"audio\": pd_test_input['audio']})\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#change speed function\nimport librosa\ndef change_speed(audio):\n    y_fast = librosa.effects.time_stretch(audio, rate = 2) \n    return y_fast","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#change pitch function\nimport librosa\ndef change_pitch(audio):\n    y_shifted = librosa.effects.pitch_shift(audio, sr=16000, n_steps=4) \n   \n    return y_shifted","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# add noise/music to audio\nimport numpy as np\nimport sys \n\ndef add_audio(original_audio, addition_audio):\n    \n    original_audio = original_audio\n    addition_audio , sr = librosa.load(addition_audio['audio'], sr=16000)\n    addition_audio = addition_audio\n    zeros = np.zeros_like(original_audio)\n    addition_audio = np.concatenate([addition_audio, zeros])\n    addition_audio = addition_audio[:len(original_audio)]\n    Es = np.sum(original_audio ** 2)\n    En = np.sum(addition_audio ** 2)\n    alpha = np.sqrt(Es/(25*En+sys.float_info.epsilon))\n    addition_audio *=  alpha\n    mixed_audio = original_audio + addition_audio\n    return mixed_audio","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#preparing dataset(adding audio augmentation) \nimport librosa\nimport random\nimport numpy as np\nfrom pydub import AudioSegment\nimport librosa\ndef prepare_dataset(batch):\n    audio = batch[\"audio\"]\n    audio , sr = librosa.load(audio, sr=16000)\n    random_audio_music = random.randint(0,len(audio_music_dataset))\n    random_audio_noise = random.randint(0,len(audio_noise_dataset))\n    random_number = random.randint(0,4)\n    if random_number==0:\n        mixed_audio = audio\n    elif random_number==1:\n        mixed_audio = change_speed(audio)\n    elif random_number==2:\n        mixed_audio = change_pitch(audio)\n    elif random_number==3:\n        music_audio = audio_music_dataset[random_audio_music]   \n        mixed_audio = add_audio(audio, music_audio)\n    elif random_number==4:\n        noise_audio = audio_noise_dataset[random_audio_noise]\n        mixed_audio = add_audio(audio, noise_audio )\n\n    normalized_audio = librosa.util.normalize(mixed_audio)\n    batch[\"input_features\"] = feature_extractor(normalized_audio , sampling_rate=16000).input_features[0]\n    batch[\"labels\"] = tokenizer(batch[\"sentence\"]).input_ids\n    return batch","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\n\nfrom dataclasses import dataclass\nfrom typing import Any, Dict, List, Union\n\n@dataclass\nclass DataCollatorSpeechSeq2SeqWithPadding:\n    processor: Any\n\n    def __call__(self, features: List[Dict[str, Union[List[int], torch.Tensor]]]) -> Dict[str, torch.Tensor]:\n      \n        features = list(map(prepare_dataset, features))\n        input_features = [{\"input_features\": feature[\"input_features\"]} for feature in features]\n        batch = self.processor.feature_extractor.pad(input_features, return_tensors=\"pt\")\n        label_features = [{\"input_ids\": feature[\"labels\"]} for feature in features]\n        labels_batch = self.processor.tokenizer.pad(label_features, return_tensors=\"pt\")\n        labels = labels_batch[\"input_ids\"].masked_fill(labels_batch.attention_mask.ne(1), -100)\n       \n        if (labels[:, 0] == self.processor.tokenizer.bos_token_id).all().cpu().item():\n            labels = labels[:, 1:]\n\n        batch[\"labels\"] = labels\n\n        return batch","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_collator = DataCollatorSpeechSeq2SeqWithPadding(processor=processor)\nimport evaluate\nmetric = evaluate.load(\"wer\")\ndef compute_metrics(pred):\n    pred_ids = pred.predictions\n    label_ids = pred.label_ids\n\n    label_ids[label_ids == -100] = tokenizer.pad_token_id\n\n    pred_str = tokenizer.batch_decode(pred_ids, skip_special_tokens=True)\n    label_str = tokenizer.batch_decode(label_ids, skip_special_tokens=True)\n\n    wer = metric.compute(predictions=pred_str, references=label_str)\n\n    return {\"wer\": wer}\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import Seq2SeqTrainingArguments, Seq2SeqTrainer\n\ntraining_args = Seq2SeqTrainingArguments(\n    output_dir=\"/kaggle/working/whisper-small-hi\", \n    per_device_train_batch_size=8,\n    gradient_accumulation_steps=1,  \n    learning_rate=1e-5,\n    warmup_steps=2,\n    max_steps=1000,\n    gradient_checkpointing=True,\n    # fp16=True,\n    evaluation_strategy=\"steps\",\n    per_device_eval_batch_size=8,\n    predict_with_generate=True,\n    generation_max_length=225,\n    save_steps=10,\n    eval_steps=10,\n    logging_steps=25,\n    report_to=[\"tensorboard\"],\n    load_best_model_at_end=True,\n    metric_for_best_model=\"wer\",\n    greater_is_better=False,\n    push_to_hub=False,\n    optim=\"adamw_torch\",\n    remove_unused_columns=False,\n    resume_from_checkpoint=True\n\n)\n\ntrainer = Seq2SeqTrainer(\n    args=training_args,\n    model=model,\n    train_dataset=audio_dataset,\n    eval_dataset=audio_val_dataset,\n    data_collator=data_collator,\n    compute_metrics=compute_metrics,\n    tokenizer=processor.feature_extractor,\n   \n\n)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.train()\ntrainer.save_model('/kaggle/working/trainer_Result')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plot loss curve \nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport numpy as np\ndf = pd.read_json(\"/kaggle/working/whisper-small-hi/checkpoint-1000/trainer_state.json\")#if the last checkpoint is 1000 otherwise edit to the last checkpoint's path'\ndf_plot = df['log_history']\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"eval_loss_history = np.array([(point.get('step', None), point.get('eval_loss', None)) for point in df_plot if point.get('eval_loss', None)])\ntrain_loss_history = np.array([(point.get('step', None), point.get('loss',None)) for point in df_plot if point.get('loss', None)])\neval_wer = np.array([(point.get('step', None), point.get('eval_wer',None)) for point in df_plot if point.get('eval_wer', None)])\n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plot eval WER\nplt.plot(eval_wer[:,0], eval_wer[:,1], 'g-', label=\"eval_wer\")\nplt.xlabel(\"step\")\nplt.ylabel(\"eval_wer\")\nplt.legend()\nplt.savefig('eval_wer.png')\n# plt.show()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plot eval loss and train loss\nplt.plot(eval_loss_history[:,0], eval_loss_history[:,1], 'r-', label=\"eval loss\")\nplt.xlabel(\"step\")\nplt.ylabel(\"eval_loss\")\nplt.plot(train_loss_history[:,0], train_loss_history[:,1], 'b-', label=\"train loss\")\nplt.xlabel(\"step\")\nplt.ylabel(\"loss\")\nplt.legend()\n# plt.show()\nplt.savefig('loss.png')","metadata":{},"execution_count":null,"outputs":[]}]}