{"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":"code","source":"!pip install -q transformers datasets librosa evaluate jiwer gradio bitsandbytes #accelerate bitsandbytes==0.37\n!pip install -q git+https://github.com/huggingface/peft.git@main\n!pip install -q git+https://github.com/huggingface/accelerate.git@main\n!pip install -q git+https://github.com/huggingface/datasets.git@main","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-10-07T19:52:21.630703Z","iopub.execute_input":"2023-10-07T19:52:21.631209Z","iopub.status.idle":"2023-10-07T19:53:50.823116Z","shell.execute_reply.started":"2023-10-07T19:52:21.631169Z","shell.execute_reply":"2023-10-07T19:53:50.821691Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_name_or_path = \"openai/whisper-large-v2\"\ntask = \"transcribe\"","metadata":{"execution":{"iopub.status.busy":"2023-10-07T19:53:50.825491Z","iopub.execute_input":"2023-10-07T19:53:50.826518Z","iopub.status.idle":"2023-10-07T19:53:50.831797Z","shell.execute_reply.started":"2023-10-07T19:53:50.826480Z","shell.execute_reply":"2023-10-07T19:53:50.830799Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import WhisperFeatureExtractor\n\nfeature_extractor = WhisperFeatureExtractor.from_pretrained(model_name_or_path)","metadata":{"execution":{"iopub.status.busy":"2023-10-07T19:53:50.834211Z","iopub.execute_input":"2023-10-07T19:53:50.834760Z","iopub.status.idle":"2023-10-07T19:53:53.088537Z","shell.execute_reply.started":"2023-10-07T19:53:50.834728Z","shell.execute_reply":"2023-10-07T19:53:53.087650Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import WhisperTokenizer\n\ntokenizer = WhisperTokenizer.from_pretrained(model_name_or_path, language='bn', task=task)","metadata":{"execution":{"iopub.status.busy":"2023-10-07T19:53:53.091045Z","iopub.execute_input":"2023-10-07T19:53:53.092115Z","iopub.status.idle":"2023-10-07T19:53:54.277723Z","shell.execute_reply.started":"2023-10-07T19:53:53.092077Z","shell.execute_reply":"2023-10-07T19:53:54.276738Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import WhisperProcessor\n\nprocessor = WhisperProcessor.from_pretrained(model_name_or_path, language='bn', task=task)","metadata":{"execution":{"iopub.status.busy":"2023-10-07T19:53:54.278917Z","iopub.execute_input":"2023-10-07T19:53:54.279267Z","iopub.status.idle":"2023-10-07T19:53:54.501793Z","shell.execute_reply.started":"2023-10-07T19:53:54.279235Z","shell.execute_reply":"2023-10-07T19:53:54.500834Z"},"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    \n@dataclass\nclass DataCollatorSpeechSeq2SeqWithPadding:\n    \"\"\"\n    Data collator that will dynamically pad the inputs received.\n    Args:\n        processor ([`WhisperProcessor`])\n            The processor used for processing the data.\n        decoder_start_token_id (`int`)\n            The begin-of-sentence of the decoder.\n        forward_attention_mask (`bool`)\n            Whether to return attention_mask.\n    \"\"\"\n\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\n        # different padding methods\n        model_input_name = self.processor.model_input_names[0]\n        input_features = [\n            {model_input_name: feature[model_input_name]} for feature in features\n        ]\n        label_features = [{\"input_ids\": feature[\"labels\"]} for feature in features]\n\n        batch = self.processor.feature_extractor.pad(\n            input_features, return_tensors=\"pt\"\n        )\n\n        labels_batch = self.processor.tokenizer.pad(label_features, return_tensors=\"pt\")\n        \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\n        batch[\"labels\"] = labels\n\n        return batch\n\ndata_collator = DataCollatorSpeechSeq2SeqWithPadding(\n    processor=processor,\n\n)","metadata":{"execution":{"iopub.status.busy":"2023-10-07T19:53:54.503350Z","iopub.execute_input":"2023-10-07T19:53:54.503911Z","iopub.status.idle":"2023-10-07T19:53:56.154121Z","shell.execute_reply.started":"2023-10-07T19:53:54.503877Z","shell.execute_reply":"2023-10-07T19:53:56.153149Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import evaluate\n\nmetric = evaluate.load(\"wer\")","metadata":{"execution":{"iopub.status.busy":"2023-10-07T19:53:56.155358Z","iopub.execute_input":"2023-10-07T19:53:56.156647Z","iopub.status.idle":"2023-10-07T19:54:07.318848Z","shell.execute_reply.started":"2023-10-07T19:53:56.156612Z","shell.execute_reply":"2023-10-07T19:54:07.317995Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import WhisperForConditionalGeneration\nfrom transformers import BitsAndBytesConfig\n\nnf4_config = BitsAndBytesConfig(\n   load_in_4bit=True,\n   bnb_4bit_quant_type=\"nf4\",\n   bnb_4bit_use_double_quant=True,\n   bnb_4bit_compute_dtype=torch.bfloat16\n)\n\n\nmodel = WhisperForConditionalGeneration.from_pretrained(model_name_or_path,quantization_config=nf4_config,\n                                                        device_map=\"auto\")","metadata":{"execution":{"iopub.status.busy":"2023-10-07T19:54:07.320462Z","iopub.execute_input":"2023-10-07T19:54:07.321228Z","iopub.status.idle":"2023-10-07T20:03:01.778646Z","shell.execute_reply.started":"2023-10-07T19:54:07.321192Z","shell.execute_reply":"2023-10-07T20:03:01.777474Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from peft import prepare_model_for_kbit_training\n\nmodel = prepare_model_for_kbit_training(model)","metadata":{"execution":{"iopub.status.busy":"2023-10-07T20:03:01.780426Z","iopub.execute_input":"2023-10-07T20:03:01.780797Z","iopub.status.idle":"2023-10-07T20:03:01.853698Z","shell.execute_reply.started":"2023-10-07T20:03:01.780762Z","shell.execute_reply":"2023-10-07T20:03:01.852651Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def make_inputs_require_grad(module, input, output):\n    output.requires_grad_(True)\n\nmodel.model.encoder.conv1.register_forward_hook(make_inputs_require_grad)","metadata":{"execution":{"iopub.status.busy":"2023-10-07T20:03:01.856980Z","iopub.execute_input":"2023-10-07T20:03:01.857604Z","iopub.status.idle":"2023-10-07T20:03:01.865354Z","shell.execute_reply.started":"2023-10-07T20:03:01.857567Z","shell.execute_reply":"2023-10-07T20:03:01.864304Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from peft import LoraConfig, PeftModel, LoraModel, LoraConfig, get_peft_model\n\nlora_config = LoraConfig(r=16, lora_alpha=32, target_modules=[\"q_proj\", \"v_proj\"], lora_dropout=0.05, bias=\"none\")\n\nmodel = get_peft_model(model, lora_config)\nmodel.print_trainable_parameters()","metadata":{"execution":{"iopub.status.busy":"2023-10-07T20:03:01.867162Z","iopub.execute_input":"2023-10-07T20:03:01.867828Z","iopub.status.idle":"2023-10-07T20:03:05.396759Z","shell.execute_reply.started":"2023-10-07T20:03:01.867794Z","shell.execute_reply":"2023-10-07T20:03:05.395704Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"### load data\nimport datasets\nfrom datasets import DatasetDict, load_dataset\nfrom pathlib import Path\n\nvectorized_datasets = DatasetDict()\ntrain_data_dir = '/kaggle/input/bengali-ai-asr-10k'\nvalidation_data_dir = '/kaggle/input/bengali-ai-asr-10k'\n\ntrain_files = list(map(str, Path(train_data_dir).glob(\"train*.parquet\")))\nvectorized_datasets[\"train\"] = load_dataset(\"parquet\", data_files=train_files[:1], split=\"train\")\n\neval_files = list(map(str, Path(validation_data_dir).glob(\"eval*.parquet\")))\nvectorized_datasets[\"eval\"] = load_dataset(\n    \"parquet\", data_files=eval_files, split=\"train\"\n)","metadata":{"execution":{"iopub.status.busy":"2023-10-07T20:03:05.398255Z","iopub.execute_input":"2023-10-07T20:03:05.399211Z","iopub.status.idle":"2023-10-07T20:03:26.772365Z","shell.execute_reply.started":"2023-10-07T20:03:05.399178Z","shell.execute_reply":"2023-10-07T20:03:26.771230Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import Seq2SeqTrainingArguments\n\ntraining_args = Seq2SeqTrainingArguments(\n    output_dir=\"lora/test\",  # change to a repo name of your choice\n    report_to=\"none\", ### comment this out to login to wandb\n    per_device_train_batch_size=8,\n    gradient_accumulation_steps=1,  # increase by 2x for every 2x decrease in batch size\n    learning_rate=1e-5,\n    warmup_steps=50,\n    num_train_epochs=1,\n    evaluation_strategy=\"steps\",\n    fp16=True,\n    per_device_eval_batch_size=8,\n    optim = \"paged_adamw_8bit\" , \n    logging_steps=100,\n    remove_unused_columns=False,  # required as the PeftModel forward doesn't have the signature of the wrapped model's forward\n    label_names=[\"labels\"],  # same reason as above\n)","metadata":{"execution":{"iopub.status.busy":"2023-10-07T20:03:26.773895Z","iopub.execute_input":"2023-10-07T20:03:26.774815Z","iopub.status.idle":"2023-10-07T20:03:26.832089Z","shell.execute_reply.started":"2023-10-07T20:03:26.774776Z","shell.execute_reply":"2023-10-07T20:03:26.830999Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import Seq2SeqTrainer, TrainerCallback, TrainingArguments, TrainerState, TrainerControl\nfrom transformers.trainer_utils import PREFIX_CHECKPOINT_DIR\n\n# This callback helps to save only the adapter weights and remove the base model weights.\nclass SavePeftModelCallback(TrainerCallback):\n    def on_save(\n        self,\n        args: TrainingArguments,\n        state: TrainerState,\n        control: TrainerControl,\n        **kwargs,\n    ):\n        checkpoint_folder = os.path.join(args.output_dir, f\"{PREFIX_CHECKPOINT_DIR}-{state.global_step}\")\n\n        peft_model_path = os.path.join(checkpoint_folder, \"adapter_model\")\n        kwargs[\"model\"].save_pretrained(peft_model_path)\n\n        pytorch_model_path = os.path.join(checkpoint_folder, \"pytorch_model.bin\")\n        if os.path.exists(pytorch_model_path):\n            os.remove(pytorch_model_path)\n        return control\n\n\ntrainer = Seq2SeqTrainer(\n    args=training_args,\n    model=model,\n    train_dataset=vectorized_datasets[\"train\"],\n    eval_dataset=vectorized_datasets[\"eval\"],\n    data_collator=data_collator,\n    tokenizer=processor.feature_extractor,\n    callbacks=[SavePeftModelCallback],\n\n)\nmodel.config.use_cache = False  # silence the warnings. Please re-enable for inference!","metadata":{"execution":{"iopub.status.busy":"2023-10-07T20:03:26.833751Z","iopub.execute_input":"2023-10-07T20:03:26.834173Z","iopub.status.idle":"2023-10-07T20:03:28.048962Z","shell.execute_reply.started":"2023-10-07T20:03:26.834135Z","shell.execute_reply":"2023-10-07T20:03:28.047682Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.train()\ntrainer.save_model()","metadata":{"execution":{"iopub.status.busy":"2023-10-07T20:03:28.050572Z","iopub.execute_input":"2023-10-07T20:03:28.051199Z","iopub.status.idle":"2023-10-07T20:04:11.824652Z","shell.execute_reply.started":"2023-10-07T20:03:28.051165Z","shell.execute_reply":"2023-10-07T20:04:11.823228Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}