{"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"}],"dockerImageVersionId":30527,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"%%capture\n!pip install evaluate\n!pip install jiwer","metadata":{"execution":{"iopub.status.busy":"2023-09-10T17:02:13.697654Z","iopub.execute_input":"2023-09-10T17:02:13.698070Z","iopub.status.idle":"2023-09-10T17:02:38.007084Z","shell.execute_reply.started":"2023-09-10T17:02:13.698034Z","shell.execute_reply":"2023-09-10T17:02:38.005652Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Import the necessary libraries","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport os\nfrom tqdm import tqdm\nfrom IPython.display import Audio\nimport librosa  \nfrom transformers import WhisperProcessor\nimport torch\nfrom dataclasses import dataclass\nfrom typing import Any, Dict, List, Union\nimport evaluate\nfrom transformers.models.whisper.english_normalizer import BasicTextNormalizer\nfrom transformers import WhisperForConditionalGeneration\nfrom functools import partial\nfrom transformers import Seq2SeqTrainingArguments\nfrom transformers import Seq2SeqTrainer\nfrom transformers.models.whisper.tokenization_whisper import TO_LANGUAGE_CODE\n\nTO_LANGUAGE_CODE['bengali']","metadata":{"execution":{"iopub.status.busy":"2023-09-10T16:59:44.969401Z","iopub.execute_input":"2023-09-10T16:59:44.970120Z","iopub.status.idle":"2023-09-10T16:59:59.223282Z","shell.execute_reply.started":"2023-09-10T16:59:44.970084Z","shell.execute_reply":"2023-09-10T16:59:59.222384Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load the dataset","metadata":{}},{"cell_type":"code","source":"path = '/kaggle/input/bengaliai-speech/'\nos.listdir(path)","metadata":{"execution":{"iopub.status.busy":"2023-09-10T16:59:59.224671Z","iopub.execute_input":"2023-09-10T16:59:59.226184Z","iopub.status.idle":"2023-09-10T16:59:59.237635Z","shell.execute_reply.started":"2023-09-10T16:59:59.226148Z","shell.execute_reply":"2023-09-10T16:59:59.236775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv(path +'/train.csv')\ntrain['path']= path + '/train_mp3s/' + train['id']+'.mp3'\ntrain.head()","metadata":{"execution":{"iopub.status.busy":"2023-09-10T16:59:59.240673Z","iopub.execute_input":"2023-09-10T16:59:59.241069Z","iopub.status.idle":"2023-09-10T17:00:03.907381Z","shell.execute_reply.started":"2023-09-10T16:59:59.241026Z","shell.execute_reply":"2023-09-10T17:00:03.906098Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_rate = 22500  # this is a common sample rate for audio\n\n# Load the audio file using librosa or any other audio processing library you prefer\naudio_path=path + '/train_mp3s/' + train['id'].iloc[7]+'.mp3'\naudio_data, _ = librosa.load(audio_path, sr=sample_rate)\nprint(train['sentence'].iloc[7])\n# Display the audio using IPython.display.Audio\ndisplay(Audio(data=audio_data, rate=sample_rate))","metadata":{"execution":{"iopub.status.busy":"2023-09-10T17:00:03.909185Z","iopub.execute_input":"2023-09-10T17:00:03.909660Z","iopub.status.idle":"2023-09-10T17:00:13.177605Z","shell.execute_reply.started":"2023-09-10T17:00:03.909626Z","shell.execute_reply":"2023-09-10T17:00:13.176777Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load the model and preprocessor","metadata":{}},{"cell_type":"code","source":"model_id = \"openai/whisper-small\"\n\nprocessor = WhisperProcessor.from_pretrained(\n    model_id, \n    language=\"bengali\", \n    task=\"transcribe\"\n)\n\nprocessor","metadata":{"execution":{"iopub.status.busy":"2023-09-10T17:00:13.178600Z","iopub.execute_input":"2023-09-10T17:00:13.179378Z","iopub.status.idle":"2023-09-10T17:00:17.427322Z","shell.execute_reply.started":"2023-09-10T17:00:13.179343Z","shell.execute_reply":"2023-09-10T17:00:17.426197Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sampling_rate = processor.feature_extractor.sampling_rate\nsampling_rate","metadata":{"execution":{"iopub.status.busy":"2023-09-10T17:00:17.428937Z","iopub.execute_input":"2023-09-10T17:00:17.429635Z","iopub.status.idle":"2023-09-10T17:00:17.436700Z","shell.execute_reply.started":"2023-09-10T17:00:17.429601Z","shell.execute_reply":"2023-09-10T17:00:17.435767Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_audio(audio_path):\n    audio_arrays, sampling_rate = librosa.load(audio_path)\n    return audio_arrays, sampling_rate\n\nload_audio(audio_path)","metadata":{"execution":{"iopub.status.busy":"2023-09-10T17:00:17.438120Z","iopub.execute_input":"2023-09-10T17:00:17.438704Z","iopub.status.idle":"2023-09-10T17:00:17.462674Z","shell.execute_reply.started":"2023-09-10T17:00:17.438672Z","shell.execute_reply":"2023-09-10T17:00:17.461655Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"- {'audio': Audio(sampling_rate=48000, mono=True, decode=True, id=None),\n-  'sentence': Value(dtype='string', id=None)}","metadata":{}},{"cell_type":"code","source":"processor.feature_extractor.model_input_names\naudio_arrays, _= librosa.load(train['path'].iloc[1], sr=sample_rate)\naudio_arrays.tolist()[0]","metadata":{"execution":{"iopub.status.busy":"2023-09-10T17:00:17.464843Z","iopub.execute_input":"2023-09-10T17:00:17.465846Z","iopub.status.idle":"2023-09-10T17:00:17.486411Z","shell.execute_reply.started":"2023-09-10T17:00:17.465813Z","shell.execute_reply":"2023-09-10T17:00:17.485155Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def prepare_dataset(example, sample_rate= sample_rate):\n    audio_arrays, sampling_rate = librosa.load(example['path'], sr=sample_rate)\n    \n    example = processor(\n        audio=audio_arrays,\n        sampling_rate=sampling_rate,\n        text=example[\"sentence\"],\n    )\n    # compute input length of audio sample in seconds\n    example[\"input_length\"] = len(audio_arrays) / sampling_rate\n\n    return example","metadata":{"execution":{"iopub.status.busy":"2023-09-10T17:00:17.488672Z","iopub.execute_input":"2023-09-10T17:00:17.489401Z","iopub.status.idle":"2023-09-10T17:00:17.496099Z","shell.execute_reply.started":"2023-09-10T17:00:17.489369Z","shell.execute_reply":"2023-09-10T17:00:17.495001Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_rate = processor.feature_extractor.sampling_rate\n\ninputs = prepare_dataset(train.iloc[0], sample_rate= sample_rate)\nprint('Audio files  :',inputs['input_features'])\nprint('Text vector  :',inputs['labels'])\nprint('Input length :',inputs['input_length'])","metadata":{"execution":{"iopub.status.busy":"2023-09-10T17:00:17.497647Z","iopub.execute_input":"2023-09-10T17:00:17.498019Z","iopub.status.idle":"2023-09-10T17:00:17.589592Z","shell.execute_reply.started":"2023-09-10T17:00:17.497977Z","shell.execute_reply":"2023-09-10T17:00:17.588270Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"decoded_text = processor.decode(inputs['labels'])\nprint(decoded_text)","metadata":{"execution":{"iopub.status.busy":"2023-09-10T17:00:17.595417Z","iopub.execute_input":"2023-09-10T17:00:17.596150Z","iopub.status.idle":"2023-09-10T17:00:17.604129Z","shell.execute_reply.started":"2023-09-10T17:00:17.596091Z","shell.execute_reply":"2023-09-10T17:00:17.603097Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from IPython.display import Audio\n\n# Convert the encoded audio back to waveform using librosa\ndecoded_audio = processor.feature_extractor.sampling_rate * librosa.feature.inverse.mel_to_audio(\n    inputs['input_features'][0], sr=processor.feature_extractor.sampling_rate\n)\n\n# Display the audio using IPython.display.Audio\ndisplay(Audio(decoded_audio, rate=sample_rate))","metadata":{"execution":{"iopub.status.busy":"2023-09-10T17:00:17.612501Z","iopub.execute_input":"2023-09-10T17:00:17.613422Z","iopub.status.idle":"2023-09-10T17:00:37.842353Z","shell.execute_reply.started":"2023-09-10T17:00:17.613389Z","shell.execute_reply":"2023-09-10T17:00:37.841031Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time \ntrain_data = []\nfor _, row in tqdm(train.iloc[:5000].iterrows()):\n    train_data.append(prepare_dataset(row, sample_rate=sample_rate))","metadata":{"execution":{"iopub.status.busy":"2023-09-10T17:00:37.844166Z","iopub.execute_input":"2023-09-10T17:00:37.845293Z","iopub.status.idle":"2023-09-10T17:00:44.503890Z","shell.execute_reply.started":"2023-09-10T17:00:37.845244Z","shell.execute_reply":"2023-09-10T17:00:44.502568Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time \nval_data = []\nfor _, row in tqdm(train.iloc[5000:7000].iterrows()):\n    val_data.append(prepare_dataset(row, sample_rate=sample_rate))","metadata":{"execution":{"iopub.status.busy":"2023-09-10T17:00:44.505841Z","iopub.execute_input":"2023-09-10T17:00:44.512541Z","iopub.status.idle":"2023-09-10T17:00:45.166498Z","shell.execute_reply.started":"2023-09-10T17:00:44.512493Z","shell.execute_reply":"2023-09-10T17:00:45.165280Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@dataclass\nclass DataCollatorSpeechSeq2SeqWithPadding:\n    processor: Any\n\n    def __call__(\n        self, features: List[Dict[str, Union[List[int], torch.Tensor]]]\n    ) -> Dict[str, torch.Tensor]:\n        # split inputs and labels since they have to be of different lengths and need different padding methods\n        # first treat the audio inputs by simply returning torch tensors\n        input_features = [\n            {\"input_features\": feature[\"input_features\"][0]} for feature in features\n        ]\n        batch = self.processor.feature_extractor.pad(input_features, return_tensors=\"pt\")\n\n        # get the tokenized label sequences\n        label_features = [{\"input_ids\": feature[\"labels\"]} for feature in features]\n        # pad the labels to max length\n        labels_batch = self.processor.tokenizer.pad(label_features, return_tensors=\"pt\")\n\n        # replace padding with -100 to ignore loss correctly\n        labels = labels_batch[\"input_ids\"].masked_fill(\n            labels_batch.attention_mask.ne(1), -100\n        )\n\n        # if bos token is appended in previous tokenization step,\n        # cut bos token here as it's append later anyways\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\n    \ndata_collator = DataCollatorSpeechSeq2SeqWithPadding(processor=processor)","metadata":{"execution":{"iopub.status.busy":"2023-09-10T17:00:45.168537Z","iopub.execute_input":"2023-09-10T17:00:45.174471Z","iopub.status.idle":"2023-09-10T17:00:45.193737Z","shell.execute_reply.started":"2023-09-10T17:00:45.174420Z","shell.execute_reply":"2023-09-10T17:00:45.192063Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"metric = evaluate.load(\"wer\")","metadata":{"execution":{"iopub.status.busy":"2023-09-10T17:02:38.009958Z","iopub.execute_input":"2023-09-10T17:02:38.010670Z","iopub.status.idle":"2023-09-10T17:02:39.197018Z","shell.execute_reply.started":"2023-09-10T17:02:38.010631Z","shell.execute_reply":"2023-09-10T17:02:39.196042Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"normalizer = BasicTextNormalizer()\n\n\ndef compute_metrics(pred):\n    pred_ids = pred.predictions\n    label_ids = pred.label_ids\n\n    # replace -100 with the pad_token_id\n    label_ids[label_ids == -100] = processor.tokenizer.pad_token_id\n\n    # we do not want to group tokens when computing the metrics\n    pred_str = processor.batch_decode(pred_ids, skip_special_tokens=True)\n    label_str = processor.batch_decode(label_ids, skip_special_tokens=True)\n\n    # compute orthographic wer\n    wer_ortho = 100 * metric.compute(predictions=pred_str, references=label_str)\n\n    # compute normalised WER\n    pred_str_norm = [normalizer(pred) for pred in pred_str]\n    label_str_norm = [normalizer(label) for label in label_str]\n    # filtering step to only evaluate the samples that correspond to non-zero references:\n    pred_str_norm = [\n        pred_str_norm[i] for i in range(len(pred_str_norm)) if len(label_str_norm[i]) > 0\n    ]\n    label_str_norm = [\n        label_str_norm[i]\n        for i in range(len(label_str_norm))\n        if len(label_str_norm[i]) > 0\n    ]\n\n    wer = 100 * metric.compute(predictions=pred_str_norm, references=label_str_norm)\n\n    return {\"wer_ortho\": wer_ortho, \"wer\": wer}","metadata":{"execution":{"iopub.status.busy":"2023-09-10T17:02:39.198346Z","iopub.execute_input":"2023-09-10T17:02:39.199301Z","iopub.status.idle":"2023-09-10T17:02:39.209644Z","shell.execute_reply.started":"2023-09-10T17:02:39.199267Z","shell.execute_reply":"2023-09-10T17:02:39.208262Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = WhisperForConditionalGeneration.from_pretrained(\"openai/whisper-small\")\n\n# disable cache during training since it's incompatible with gradient checkpointing\nmodel.config.use_cache = False\n\n# set language and task for generation and re-enable cache\nmodel.generate = partial(\n    model.generate, language=\"bengali\", task=\"transcribe\", use_cache=True\n)\n\nmodel","metadata":{"execution":{"iopub.status.busy":"2023-09-10T17:02:47.566764Z","iopub.execute_input":"2023-09-10T17:02:47.567185Z","iopub.status.idle":"2023-09-10T17:03:16.807338Z","shell.execute_reply.started":"2023-09-10T17:02:47.567152Z","shell.execute_reply":"2023-09-10T17:03:16.806217Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"training_args = Seq2SeqTrainingArguments(\n    output_dir=\"whisper-small-dv\", \n    per_device_train_batch_size=16,\n    gradient_accumulation_steps=1,  # increase by 2x for every 2x decrease in batch size\n    learning_rate=1e-5,\n    lr_scheduler_type=\"constant_with_warmup\",\n    warmup_steps=5,\n    max_steps=20,\n    gradient_checkpointing=True,\n    fp16=True,\n    fp16_full_eval=True,\n    evaluation_strategy=\"steps\",\n    per_device_eval_batch_size=16,\n    predict_with_generate=True,\n    generation_max_length=225,\n    save_steps=500,\n    eval_steps=500,\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=True,\n)","metadata":{"execution":{"iopub.status.busy":"2023-09-10T17:03:16.809589Z","iopub.execute_input":"2023-09-10T17:03:16.809983Z","iopub.status.idle":"2023-09-10T17:03:16.855441Z","shell.execute_reply.started":"2023-09-10T17:03:16.809931Z","shell.execute_reply":"2023-09-10T17:03:16.854535Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer = Seq2SeqTrainer(\n    args=training_args,\n    model=model,\n    train_dataset=train_data,\n    eval_dataset=val_data,\n    data_collator=data_collator,\n    compute_metrics=compute_metrics,\n    tokenizer=processor,\n)","metadata":{"execution":{"iopub.status.busy":"2023-09-10T17:03:16.856801Z","iopub.execute_input":"2023-09-10T17:03:16.857430Z","iopub.status.idle":"2023-09-10T17:03:16.890005Z","shell.execute_reply.started":"2023-09-10T17:03:16.857395Z","shell.execute_reply":"2023-09-10T17:03:16.889047Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\ntrainer.train()","metadata":{"execution":{"iopub.status.busy":"2023-09-10T17:03:16.892093Z","iopub.execute_input":"2023-09-10T17:03:16.892586Z","iopub.status.idle":"2023-09-10T17:04:56.783789Z","shell.execute_reply.started":"2023-09-10T17:03:16.892537Z","shell.execute_reply":"2023-09-10T17:04:56.782640Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define the directory where you want to save the model\nsave_directory = \"/kaggle/working/\"\n\n# Save the model and tokenizer\nmodel.save_pretrained(save_directory)\nprocessor.save_pretrained(save_directory)\n\nprint(\"Model and tokenizer saved to:\", save_directory)","metadata":{"execution":{"iopub.status.busy":"2023-09-10T17:05:11.879247Z","iopub.execute_input":"2023-09-10T17:05:11.879639Z","iopub.status.idle":"2023-09-10T17:05:13.733522Z","shell.execute_reply.started":"2023-09-10T17:05:11.879609Z","shell.execute_reply":"2023-09-10T17:05:13.732465Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}