{"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":"# Training Whisper for Bengali ASR 🎵👂➡️🇧🇩📝\n\nThe [Whisper models, developed by OpenAI](https://openai.com/research/whisper), are the best open source multilingual Automatic Speech Recognition models available.\n\n## Architecture 🏗️\n\n\nThe architecture is relatively straightforward: it's a sequence-to-sequence model containing an audio encoder and a text decoder. The feature extractor turns the 1d audio signal (amplitude over time) to a log-mel spectrogram. The encoder creates hidden states which are then passed to the decoder to generate text. It's basically BART with a few convolution layers at the input.\n\n![](https://huggingface.co/blog/assets/111_fine_tune_whisper/whisper_architecture.svg)\n\n\n## Training Data 📊\n\nWhile the architecture is nothing novel, the OpenAI team created an enormous dataset on nearly 700k labeled audio data. Of that data, 117k hours were on multilingual ASR. Sadly, it looks like there was barely any Bengali data (Less than 2 hours!) in that training set, but that's why we're here! This competition provides 1200 hours of data which can be used to fine-tune the existing Whisper models. Even though Whisper wasn't trained on much Bengali data, it will be able to learn very quickly.\n\nImage from [figure 11 on page 27 of whisper paper](https://arxiv.org/pdf/2212.04356.pdf)\n![](https://raw.githubusercontent.com/nbroad1881/kaggle-images/main/bengali-ai-asr/whisper-bn.png)\n\n\n\n\n## In this notebook 📓\n\n\n### Preprocessing 🎵➡️🔢\n\nI show how to preprocess the data and how to train. The preprocessing can be a bit slow, so it is recommended to use a CPU notebook, or better yet, a bulkier CPU VM in your favorite cloud. I used another instance to preprocess 10k training samples and 1k eval samples.\n\n### Training 🏋️\n\nI show how to train on 2x T4 GPUs using Hugging Face transformers in pytorch. These GPUs have tensor cores which makes them go fast while using mixed precision. The code doesn't do anything fancy, but it can serve as a starting point for understanding ASR training. [@mbmmurad](https://kaggle.com/mbmmurad) pointed out that `bangla-speech-processing/BanglaASR` is a whisper model that has already been trained on Bangla. This notebook will do more fine-tuning  on 10k out of 960k training samples. \n\n\n### Validation 🕵️\n\nI take a random split of files for train and validation, but since the domain is unknown, it is hard to get a good sense of how well the model will do on out-of-domain data. It would be nice if train.csv had domains as a column.\n\n\nNotebook Version | Model | WER \n- | - | -\n1 | openai/whisper-base | 0.69\n2 | bangla-speech-processing/BanglaASR | 0.529\n\n---\n\n## How to improve the model 💪\n\n1. Train on more data. \n  - Like I mentioned above, there over 900k files and I only used 10k of them. \n2. Use a larger model. \n  - I'm only using the small-sized model which has 244M params. The largest model [1550M params](https://huggingface.co/openai/whisper-large-v2), which is probably too big, but there is also a [769M](https://huggingface.co/openai/whisper-medium) model.\n3. Data augmentations.\n  - I use [spec augment](https://arxiv.org/abs/1904.08779)\n  - You could also consider [BPE dropout](https://arxiv.org/abs/1910.13267) \n4. Better hyperparameters.\n  - learning rate is usually the most important","metadata":{}},{"cell_type":"code","source":"# Necessary packages\n!pip install -U evaluate datasets transformers jiwer -q","metadata":{"execution":{"iopub.status.busy":"2023-07-19T00:13:12.879710Z","iopub.execute_input":"2023-07-19T00:13:12.880280Z","iopub.status.idle":"2023-07-19T00:13:41.258402Z","shell.execute_reply.started":"2023-07-19T00:13:12.880242Z","shell.execute_reply":"2023-07-19T00:13:41.257156Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Preprocessing\n\nThis should be done on CPU because this saves GPU time and the CPUs in CPU notebooks are faster than the CPUs in GPU notebooks. Better yet, use an even better CPU in a cloud VM.","metadata":{}},{"cell_type":"code","source":"%%writefile preprocess.py\n\nimport logging\nimport warnings\nfrom dataclasses import dataclass, field\nfrom pathlib import Path\nfrom typing import Any, Dict, List, Optional, Union\n\nimport datasets\nfrom datasets import DatasetDict, load_dataset\n\nfrom transformers import (\n    AutoConfig,\n    AutoFeatureExtractor,\n    AutoTokenizer,\n    HfArgumentParser,\n    Seq2SeqTrainingArguments,\n    set_seed,\n)\n\n\nwarnings.simplefilter(\"ignore\")\n\n\n@dataclass\nclass Config:\n    \"\"\"\n    Arguments pertaining to which model/config/tokenizer we are going to fine-tune from.\n    \"\"\"\n\n    model_name_or_path: str = field(\n        metadata={\n            \"help\": \"Path to pretrained model or model identifier from huggingface.co/models\"\n        }\n    )\n\n    apply_spec_augment: bool = field(\n        default=False,\n        metadata={\n            \"help\": \"Whether to apply *SpecAugment* data augmentation to the input features. This is currently only relevant for Wav2Vec2, HuBERT, WavLM and Whisper models.\"\n        },\n    )\n    overwrite_cache: bool = field(\n        default=False,\n        metadata={\"help\": \"Overwrite the cached training and evaluation sets\"},\n    )\n    preprocessing_num_workers: Optional[int] = field(\n        default=None,\n        metadata={\"help\": \"The number of processes to use for the preprocessing.\"},\n    )\n    forced_decoder_ids: List[List[int]] = field(\n        default=None,\n        metadata={\n            \"help\": (\n                \"A list of pairs of integers which indicates a mapping from generation indices to token indices \"\n                \"that will be forced before sampling. For example, [[0, 123]] means the first generated token \"\n                \"will always be a token of index 123.\"\n            )\n        },\n    )\n    suppress_tokens: List[int] = field(\n        default=None,\n        metadata={\"help\": \"A list of tokens that will be suppressed at generation.\"},\n    )\n    max_train_samples: Optional[int] = field(\n        default=None,\n        metadata={\n            \"help\": (\n                \"For debugging purposes or quicker training, truncate the number of training examples to this \"\n                \"value if set.\"\n            )\n        },\n    )\n    max_eval_samples: Optional[int] = field(\n        default=None,\n        metadata={\n            \"help\": (\n                \"For debugging purposes or quicker training, truncate the number of evaluation examples to this \"\n                \"value if set.\"\n            )\n        },\n    )\n    audio_column_name: str = field(\n        default=\"audio\",\n        metadata={\n            \"help\": \"The name of the dataset column containing the audio data. Defaults to 'audio'\"\n        },\n    )\n    text_column_name: str = field(\n        default=\"text\",\n        metadata={\n            \"help\": \"The name of the dataset column containing the text data. Defaults to 'text'\"\n        },\n    )\n    max_duration_in_seconds: float = field(\n        default=20.0,\n        metadata={\n            \"help\": (\n                \"Truncate audio files that are longer than `max_duration_in_seconds` seconds to\"\n                \" 'max_duration_in_seconds`\"\n            )\n        },\n    )\n    min_duration_in_seconds: float = field(\n        default=0.0,\n        metadata={\n            \"help\": \"Filter audio files that are shorter than `min_duration_in_seconds` seconds\"\n        },\n    )\n    preprocessing_only: bool = field(\n        default=False,\n        metadata={\n            \"help\": (\n                \"Whether to only do data preprocessing and skip training. This is especially useful when data\"\n                \" preprocessing errors out in distributed training due to timeout. In this case, one should run the\"\n                \" preprocessing in a non-distributed setup with `preprocessing_only=True` so that the cached datasets\"\n                \" can consequently be loaded in distributed training\"\n            )\n        },\n    )\n    language: str = field(\n        default=None,\n        metadata={\n            \"help\": (\n                \"Language for multilingual fine-tuning. This argument should be set for multilingual fine-tuning \"\n                \"only. For English speech recognition, it should be set to `None`.\"\n            )\n        },\n    )\n\n    data_dir: str = field(\n        default=\"/kaggle/input/bengaliai-speech\",\n        metadata={\n            \"help\": (\n                \"Language for multilingual fine-tuning. This argument should be set for multilingual fine-tuning \"\n                \"only. For English speech recognition, it should be set to `None`.\"\n            )\n        },\n    )\n\n\nlogger = logging.getLogger(__name__)\n\n\ndef main():\n    parser = HfArgumentParser((Config, Seq2SeqTrainingArguments))\n\n    cfg, training_args = parser.parse_args_into_dataclasses()\n\n    # Set seed before initializing model.\n    set_seed(training_args.seed)\n\n    config = AutoConfig.from_pretrained(cfg.model_name_or_path)\n\n    config.update(\n        {\n            \"forced_decoder_ids\": cfg.forced_decoder_ids,\n            \"suppress_tokens\": cfg.suppress_tokens,\n        }\n    )\n\n    # SpecAugment for whisper models\n    if getattr(config, \"model_type\", None) == \"whisper\":\n        config.update({\"apply_spec_augment\": cfg.apply_spec_augment})\n\n    feature_extractor = AutoFeatureExtractor.from_pretrained(cfg.model_name_or_path)\n    tokenizer = AutoTokenizer.from_pretrained(cfg.model_name_or_path)\n\n    raw_datasets = DatasetDict()\n\n    data_dir = Path(cfg.data_dir)\n\n    raw_ds = load_dataset(\"csv\", data_files=str(data_dir / \"train.csv\"), split=\"train\")\n\n    def add_mp3_path(examples):\n        return {\n            \"audio\": [str(data_dir / f\"train_mp3s/{id_}.mp3\") for id_ in examples[\"id\"]]\n        }\n\n    raw_ds = raw_ds.map(add_mp3_path, batched=True, num_proc=cfg.preprocessing_num_workers)\n    raw_ds = raw_ds.train_test_split(\n        test_size=0.2, seed=training_args.seed, shuffle=True\n    )\n\n    raw_datasets[\"train\"] = raw_ds[\"train\"]\n    raw_datasets[\"validation\"] = raw_ds[\"test\"]\n\n    if cfg.max_train_samples:\n        raw_datasets[\"train\"] = raw_datasets[\"train\"].select(\n            range(min(cfg.max_train_samples, len(raw_datasets[\"train\"])))\n        )\n\n    if cfg.max_eval_samples:\n        raw_datasets[\"validation\"] = raw_datasets[\"validation\"].select(\n            range(min(cfg.max_eval_samples, len(raw_datasets[\"validation\"])))\n        )\n\n    # cast to audio\n    raw_datasets = raw_datasets.cast_column(\n        cfg.audio_column_name,\n        datasets.features.Audio(sampling_rate=feature_extractor.sampling_rate),\n    )\n\n    if cfg.language is not None:\n        # We only need to set the task id when the language is specified (i.e. in a multilingual setting)\n        tokenizer.set_prefix_tokens(language=cfg.language, task=\"transcribe\")\n\n    # Preprocessing the datasets.\n    # We need to read the audio files as arrays and tokenize the targets.\n    max_input_length = cfg.max_duration_in_seconds * feature_extractor.sampling_rate\n    min_input_length = cfg.min_duration_in_seconds * feature_extractor.sampling_rate\n    audio_column_name = cfg.audio_column_name\n    num_workers = cfg.preprocessing_num_workers\n    text_column_name = cfg.text_column_name\n    model_input_name = feature_extractor.model_input_names[0]\n    # if SpecAugment is used for whisper models, return attention_mask to guide the mask along time axis\n    forward_attention_mask = (\n        getattr(config, \"model_type\", None) == \"whisper\"\n        and getattr(config, \"apply_spec_augment\", False)\n        and getattr(config, \"mask_time_prob\", 0) > 0\n    )\n\n    def prepare_dataset(batch):\n        # process audio\n        sample = batch[audio_column_name]\n        inputs = feature_extractor(\n            sample[\"array\"],\n            sampling_rate=sample[\"sampling_rate\"],\n            return_attention_mask=forward_attention_mask,\n        )\n        # process audio length\n        batch[model_input_name] = inputs.get(model_input_name)[0]\n        batch[\"input_length\"] = len(sample[\"array\"])\n        if forward_attention_mask:\n            batch[\"attention_mask\"] = inputs.get(\"attention_mask\")[0]\n\n        # process targets\n        input_str = batch[text_column_name]\n        batch[\"labels\"] = tokenizer(input_str).input_ids\n        return batch\n\n    with training_args.main_process_first(desc=\"dataset map pre-processing\"):\n        vectorized_datasets = raw_datasets.map(\n            prepare_dataset,\n            remove_columns=next(iter(raw_datasets.values())).column_names,\n            num_proc=cfg.preprocessing_num_workers,\n            desc=\"preprocess train dataset\",\n        )\n\n    # filter data that is shorter than min_input_length or longer than\n    # max_input_length\n    def is_audio_in_length_range(length):\n        return length > min_input_length and length < max_input_length\n\n    vectorized_datasets = vectorized_datasets.filter(\n        is_audio_in_length_range,\n        num_proc=num_workers,\n        input_columns=[\"input_length\"],\n    )\n\n    def save_chunks(ds, chunk_size, prefix):\n        for i in range(0, len(ds), chunk_size):\n            ii = min(i + chunk_size, len(ds))\n\n            ds.select(range(i, ii)).to_parquet(f\"{prefix}_{i}_to_{ii}.parquet\")\n\n    # for large datasets it is advised to run the preprocessing on a\n    # single machine first with `args.preprocessing_only` since there will mostly likely\n    # be a timeout when running the script in distributed mode.\n    # In a second step `args.preprocessing_only` can then be set to `False` to load the\n    # cached dataset\n    if cfg.preprocessing_only:\n        cache = {k: v.cache_files for k, v in vectorized_datasets.items()}\n        logger.info(f\"Data preprocessing finished. Files cached at {cache}.\")\n\n        save_chunks(\n            vectorized_datasets[\"train\"], 1000, f\"train_{training_args.output_dir}\"\n        )\n        save_chunks(\n            vectorized_datasets[\"validation\"], 1000, f\"eval_{training_args.output_dir}\"\n        )\n        return\n\n\nif __name__ == \"__main__\":\n    main()","metadata":{"execution":{"iopub.status.busy":"2023-07-19T00:13:41.261668Z","iopub.execute_input":"2023-07-19T00:13:41.262481Z","iopub.status.idle":"2023-07-19T00:13:41.278879Z","shell.execute_reply.started":"2023-07-19T00:13:41.262444Z","shell.execute_reply":"2023-07-19T00:13:41.277897Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# I've already uploaded a preprocessed dataset with 10k train samples and 1k eval samples so I won't run this","metadata":{}},{"cell_type":"code","source":"# !python preprocess.py \\\n#  --model_name_or_path \"openai/whisper-small\" \\\n#  --language \"Bengali\" \\\n#  --output_dir \"75k-samples\" \\\n#  --preprocessing_num_workers 90 \\\n#  --preprocessing_only \\\n#  --text_column_name \"sentence\" \\\n#  --data_dir \"data\" \\\n#  --min_duration_in_seconds 2 \\\n#  --max_duration_in_seconds 30 \\\n#  --max_train_samples 75000 \\\n#  --max_eval_samples 5000 \\\n#  --apply_spec_augment","metadata":{"execution":{"iopub.status.busy":"2023-07-19T00:13:41.280461Z","iopub.execute_input":"2023-07-19T00:13:41.281138Z","iopub.status.idle":"2023-07-19T00:13:41.291549Z","shell.execute_reply.started":"2023-07-19T00:13:41.281080Z","shell.execute_reply":"2023-07-19T00:13:41.290480Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training script\n\nAdapted from [here](https://github.com/huggingface/transformers/blob/main/examples/pytorch/speech-recognition/run_speech_recognition_seq2seq.py)","metadata":{}},{"cell_type":"code","source":"%%writefile train.py\n\n#!/usr/bin/env python\n# coding=utf-8\n# Copyright 2021 The HuggingFace Team. All rights reserved.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n#     http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nFine-tuning the library models for sequence to sequence speech recognition.\n\"\"\"\n# You can also adapt this script on your own sequence to sequence speech\n# recognition task. Pointers for this are left as comments.\n\nimport logging\nimport os\nimport sys\nfrom dataclasses import dataclass, field\nfrom pathlib import Path\nfrom typing import Any, Dict, List, Optional, Union\n\nimport datasets\nimport evaluate\nimport torch\nfrom datasets import DatasetDict, load_dataset\n\nimport transformers\nfrom transformers import (\n    AutoConfig,\n    AutoFeatureExtractor,\n    AutoModelForSpeechSeq2Seq,\n    AutoProcessor,\n    AutoTokenizer,\n    HfArgumentParser,\n    Seq2SeqTrainer,\n    Seq2SeqTrainingArguments,\n    set_seed,\n)\nfrom transformers.trainer_utils import get_last_checkpoint, is_main_process\n\nlogger = logging.getLogger(__name__)\n\n\n@dataclass\nclass Config:\n    \"\"\"\n    Arguments pertaining to which model/config/tokenizer we are going to fine-tune from.\n    \"\"\"\n\n    model_name_or_path: str = field(\n        metadata={\n            \"help\": \"Path to pretrained model or model identifier from huggingface.co/models\"\n        }\n    )\n    freeze_encoder: bool = field(\n        default=False,\n        metadata={\"help\": \"Whether to freeze the entire encoder of the seq2seq model.\"},\n    )\n    train_data_dir: str = field(default=None, metadata={\"help\": \"Path to train files\"})\n    validation_data_dir: str = field(\n        default=None, metadata={\"help\": \"Path to eval files\"}\n    )\n    forced_decoder_ids: List[List[int]] = field(\n        default=None,\n        metadata={\n            \"help\": (\n                \"A list of pairs of integers which indicates a mapping from generation indices to token indices \"\n                \"that will be forced before sampling. For example, [[0, 123]] means the first generated token \"\n                \"will always be a token of index 123.\"\n            )\n        },\n    )\n    suppress_tokens: List[int] = field(\n        default=None,\n        metadata={\"help\": \"A list of tokens that will be suppressed at generation.\"},\n    )\n    apply_spec_augment: bool = field(\n        default=False,\n        metadata={\n            \"help\": \"Whether to apply *SpecAugment* data augmentation to the input features. This is currently only relevant for Wav2Vec2, HuBERT, WavLM and Whisper models.\"\n        },\n    )\n    max_train_samples: Optional[int] = field(\n        default=None,\n        metadata={\n            \"help\": (\n                \"For debugging purposes or quicker training, truncate the number of training examples to this \"\n                \"value if set.\"\n            )\n        },\n    )\n    max_eval_samples: Optional[int] = field(\n        default=None,\n        metadata={\n            \"help\": (\n                \"For debugging purposes or quicker training, truncate the number of evaluation examples to this \"\n                \"value if set.\"\n            )\n        },\n    )\n    audio_column_name: str = field(\n        default=\"audio\",\n        metadata={\n            \"help\": \"The name of the dataset column containing the audio data. Defaults to 'audio'\"\n        },\n    )\n    text_column_name: str = field(\n        default=\"text\",\n        metadata={\n            \"help\": \"The name of the dataset column containing the text data. Defaults to 'text'\"\n        },\n    )\n\n    language: str = field(\n        default=None,\n        metadata={\n            \"help\": (\n                \"Language for multilingual fine-tuning. This argument should be set for multilingual fine-tuning \"\n                \"only. For English speech recognition, it should be set to `None`.\"\n            )\n        },\n    )\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    decoder_start_token_id: int\n    forward_attention_mask: bool\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        if self.forward_attention_mask:\n            batch[\"attention_mask\"] = torch.LongTensor(\n                [feature[\"attention_mask\"] for feature in features]\n            )\n\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.decoder_start_token_id).all().cpu().item():\n            labels = labels[:, 1:]\n\n        batch[\"labels\"] = labels\n\n        return batch\n\n\ndef main():\n    # 1. Parse input arguments\n    parser = HfArgumentParser((Config, Seq2SeqTrainingArguments))\n\n    cfg, training_args = parser.parse_args_into_dataclasses()\n\n    # 2. Detecting last checkpoint and eventually continue from last checkpoint\n    last_checkpoint = None\n    if (\n        os.path.isdir(training_args.output_dir)\n        and training_args.do_train\n        and not training_args.overwrite_output_dir\n    ):\n        last_checkpoint = get_last_checkpoint(training_args.output_dir)\n        if last_checkpoint is None and len(os.listdir(training_args.output_dir)) > 0:\n            raise ValueError(\n                f\"Output directory ({training_args.output_dir}) already exists and is not empty. \"\n                \"Use --overwrite_output_dir to overcome.\"\n            )\n        elif (\n            last_checkpoint is not None and training_args.resume_from_checkpoint is None\n        ):\n            logger.info(\n                f\"Checkpoint detected, resuming training at {last_checkpoint}. To avoid this behavior, change \"\n                \"the `--output_dir` or add `--overwrite_output_dir` to train from scratch.\"\n            )\n\n    # Set seed before initializing model.\n    set_seed(training_args.seed)\n\n    # 3. Load dataset\n    vectorized_datasets = DatasetDict()\n\n    if training_args.do_train:\n        train_files = list(map(str, Path(cfg.train_data_dir).glob(\"train*.parquet\")))\n        vectorized_datasets[\"train\"] = load_dataset(\n            \"parquet\", data_files=train_files, split=\"train\"\n        )\n\n    if training_args.do_eval:\n        eval_files = list(map(str, Path(cfg.validation_data_dir).glob(\"eval*.parquet\")))\n        vectorized_datasets[\"eval\"] = load_dataset(\n            \"parquet\", data_files=eval_files, split=\"train\"\n        )\n\n    # 4. Load pretrained model, tokenizer, and feature extractor\n    #\n    # Distributed training:\n    # The .from_pretrained methods guarantee that only one local process can concurrently\n    config = AutoConfig.from_pretrained(\n        cfg.model_name_or_path,\n    )\n\n    config.update(\n        {\n            \"forced_decoder_ids\": cfg.forced_decoder_ids,\n            \"suppress_tokens\": cfg.suppress_tokens,\n        }\n    )\n\n    # SpecAugment for whisper models\n    if getattr(config, \"model_type\", None) == \"whisper\":\n        config.update({\"apply_spec_augment\": cfg.apply_spec_augment})\n\n    feature_extractor = AutoFeatureExtractor.from_pretrained(\n        cfg.model_name_or_path,\n    )\n    tokenizer = AutoTokenizer.from_pretrained(\n        cfg.model_name_or_path,\n    )\n    model = AutoModelForSpeechSeq2Seq.from_pretrained(\n        cfg.model_name_or_path,\n        config=config,\n    )\n\n    if model.config.decoder_start_token_id is None:\n        raise ValueError(\n            \"Make sure that `config.decoder_start_token_id` is correctly defined\"\n        )\n\n    if cfg.freeze_encoder:\n        model.freeze_encoder()\n        model.model.encoder.gradient_checkpointing = False\n\n    if cfg.language is not None:\n        # We only need to set the task id when the language is specified (i.e. in a multilingual setting)\n        tokenizer.set_prefix_tokens(language=cfg.language, task=\"transcribe\")\n\n    if cfg.max_train_samples is not None:\n        vectorized_datasets[\"train\"] = vectorized_datasets[\"train\"].select(\n            range(cfg.max_train_samples)\n        )\n\n    if cfg.max_eval_samples is not None:\n        vectorized_datasets[\"eval\"] = vectorized_datasets[\"eval\"].select(\n            range(cfg.max_eval_samples)\n        )\n\n    # 5. Load Metric\n    metric = evaluate.load(\"wer\")\n\n    def compute_metrics(pred):\n        pred_ids = pred.predictions\n\n        pred.label_ids[pred.label_ids == -100] = tokenizer.pad_token_id\n\n        pred_str = tokenizer.batch_decode(pred_ids, skip_special_tokens=True)\n        # we do not want to group tokens when computing the metrics\n        label_str = tokenizer.batch_decode(pred.label_ids, skip_special_tokens=True)\n\n        wer = metric.compute(predictions=pred_str, references=label_str)\n\n        return {\"wer\": wer}\n\n    # 6. Create a single speech processor\n    # make sure all processes wait until data is saved\n    with training_args.main_process_first():\n        # only the main process saves them\n        if is_main_process(training_args.local_rank):\n            # save feature extractor, tokenizer and config\n            feature_extractor.save_pretrained(training_args.output_dir)\n            tokenizer.save_pretrained(training_args.output_dir)\n            config.save_pretrained(training_args.output_dir)\n\n    processor = AutoProcessor.from_pretrained(training_args.output_dir)\n\n    # if SpecAugment is used for whisper models, return attention_mask to guide the mask along time axis\n    forward_attention_mask = (\n        getattr(config, \"model_type\", None) == \"whisper\"\n        and getattr(config, \"apply_spec_augment\", False)\n        and getattr(config, \"mask_time_prob\", 0) > 0\n    )\n\n    # 7. Define data collator\n    data_collator = DataCollatorSpeechSeq2SeqWithPadding(\n        processor=processor,\n        decoder_start_token_id=model.config.decoder_start_token_id,\n        forward_attention_mask=forward_attention_mask,\n    )\n\n    # 8. Initialize Trainer\n    trainer = Seq2SeqTrainer(\n        model=model,\n        args=training_args,\n        train_dataset=vectorized_datasets[\"train\"] if training_args.do_train else None,\n        eval_dataset=vectorized_datasets[\"eval\"] if training_args.do_eval else None,\n        tokenizer=feature_extractor,\n        data_collator=data_collator,\n        compute_metrics=compute_metrics\n        if training_args.predict_with_generate\n        else None,\n    )\n\n    # 9. Training\n    if training_args.do_train:\n        checkpoint = None\n        if training_args.resume_from_checkpoint is not None:\n            checkpoint = training_args.resume_from_checkpoint\n        elif last_checkpoint is not None:\n            checkpoint = last_checkpoint\n        train_result = trainer.train(resume_from_checkpoint=checkpoint)\n        trainer.save_model()  # Saves the feature extractor too for easy upload\n\n        metrics = train_result.metrics\n        max_train_samples = (\n            cfg.max_train_samples\n            if cfg.max_train_samples is not None\n            else len(vectorized_datasets[\"train\"])\n        )\n        metrics[\"train_samples\"] = min(\n            max_train_samples, len(vectorized_datasets[\"train\"])\n        )\n        trainer.log_metrics(\"train\", metrics)\n        trainer.save_metrics(\"train\", metrics)\n        trainer.save_state()\n\n    # 10. Evaluation\n    results = {}\n    if training_args.do_eval:\n        logger.info(\"*** Evaluate ***\")\n        metrics = trainer.evaluate(\n            metric_key_prefix=\"eval\",\n            max_length=training_args.generation_max_length,\n            num_beams=training_args.generation_num_beams,\n        )\n        max_eval_samples = (\n            cfg.max_eval_samples\n            if cfg.max_eval_samples is not None\n            else len(vectorized_datasets[\"eval\"])\n        )\n        metrics[\"eval_samples\"] = min(\n            max_eval_samples, len(vectorized_datasets[\"eval\"])\n        )\n\n        trainer.log_metrics(\"eval\", metrics)\n        trainer.save_metrics(\"eval\", metrics)\n\n\nif __name__ == \"__main__\":\n    main()","metadata":{"execution":{"iopub.status.busy":"2023-07-19T00:13:41.294285Z","iopub.execute_input":"2023-07-19T00:13:41.294724Z","iopub.status.idle":"2023-07-19T00:13:41.316977Z","shell.execute_reply.started":"2023-07-19T00:13:41.294691Z","shell.execute_reply":"2023-07-19T00:13:41.316091Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_dir = \"/kaggle/input/bengali-ai-asr-10k\"\n\n!torchrun --nproc_per_node 2 train.py \\\n --model_name_or_path \"bangla-speech-processing/BanglaASR\" \\\n --train_data_dir $data_dir \\\n --validation_data_dir $data_dir \\\n --language \"Bengali\" \\\n --output_dir \"whisper-base-bn\" \\\n --do_train \\\n --do_eval \\\n --fp16 \\\n --group_by_length \\\n --predict_with_generate \\\n --dataloader_num_workers 1 \\\n --overwrite_output_dir \\\n --per_device_train_batch_size 4 \\\n --length_column_name \"input_length\" \\\n --report_to \"none\" \\\n --metric_for_best_model \"wer\" \\\n --greater_is_better False \\\n --evaluation_strategy \"epoch\" \\\n --save_strategy \"epoch\" \\\n --save_total_limit 1 \\\n --logging_steps 10 \\\n --gradient_checkpointing \\\n --warmup_steps 50 \\\n --apply_spec_augment True \\\n --num_train_epochs 3 \\\n --learning_rate \"1e-5\"","metadata":{"execution":{"iopub.status.busy":"2023-07-19T00:13:41.318244Z","iopub.execute_input":"2023-07-19T00:13:41.318690Z","iopub.status.idle":"2023-07-19T03:26:23.867155Z","shell.execute_reply.started":"2023-07-19T00:13:41.318657Z","shell.execute_reply":"2023-07-19T03:26:23.865966Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}