{"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"}},"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-09-04T06:33:49.105018Z","iopub.execute_input":"2023-09-04T06:33:49.105447Z","iopub.status.idle":"2023-09-04T06:34:18.090123Z","shell.execute_reply.started":"2023-09-04T06:33:49.105413Z","shell.execute_reply":"2023-09-04T06:34:18.088504Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install bnunicodenormalizer","metadata":{"execution":{"iopub.status.busy":"2023-09-04T06:34:18.092382Z","iopub.execute_input":"2023-09-04T06:34:18.092820Z","iopub.status.idle":"2023-09-04T06:34:34.383935Z","shell.execute_reply.started":"2023-09-04T06:34:18.092786Z","shell.execute_reply":"2023-09-04T06:34:34.382569Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\ntrain_meta = pd.read_csv('/kaggle/input/bengaliai-speech/train.csv')\nmetadata = pd.read_csv('/kaggle/input/bengaliai-speech-train-metadata/train_metadata.csv')","metadata":{"execution":{"iopub.status.busy":"2023-09-04T06:35:59.436488Z","iopub.execute_input":"2023-09-04T06:35:59.436886Z","iopub.status.idle":"2023-09-04T06:36:29.514110Z","shell.execute_reply.started":"2023-09-04T06:35:59.436855Z","shell.execute_reply":"2023-09-04T06:36:29.513007Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_meta = train_meta.merge(metadata[['id', 'ykg_wer']], on=\"id\", how=\"left\")","metadata":{"execution":{"iopub.status.busy":"2023-09-04T06:36:55.725641Z","iopub.execute_input":"2023-09-04T06:36:55.726096Z","iopub.status.idle":"2023-09-04T06:36:57.272973Z","shell.execute_reply.started":"2023-09-04T06:36:55.726056Z","shell.execute_reply":"2023-09-04T06:36:57.271628Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_select = train_meta.loc[(train_meta.split=='valid')\n                             | ((train_meta.split=='train') & (train_meta.ykg_wer<=0.2))]","metadata":{"execution":{"iopub.status.busy":"2023-09-04T06:40:35.417933Z","iopub.execute_input":"2023-09-04T06:40:35.418418Z","iopub.status.idle":"2023-09-04T06:40:35.805454Z","shell.execute_reply.started":"2023-09-04T06:40:35.418384Z","shell.execute_reply":"2023-09-04T06:40:35.803955Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_select = train_select.drop(columns=['ykg_wer'])","metadata":{"execution":{"iopub.status.busy":"2023-09-04T06:44:28.295049Z","iopub.execute_input":"2023-09-04T06:44:28.295485Z","iopub.status.idle":"2023-09-04T06:44:28.314644Z","shell.execute_reply.started":"2023-09-04T06:44:28.295452Z","shell.execute_reply":"2023-09-04T06:44:28.313658Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_select","metadata":{"execution":{"iopub.status.busy":"2023-09-04T06:44:39.238702Z","iopub.execute_input":"2023-09-04T06:44:39.239234Z","iopub.status.idle":"2023-09-04T06:44:39.259656Z","shell.execute_reply.started":"2023-09-04T06:44:39.239192Z","shell.execute_reply":"2023-09-04T06:44:39.257913Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_select.to_csv(\"train_select.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2023-08-28T09:29:43.851384Z","iopub.execute_input":"2023-08-28T09:29:43.851814Z","iopub.status.idle":"2023-08-28T09:29:44.082966Z","shell.execute_reply.started":"2023-08-28T09:29:43.851779Z","shell.execute_reply":"2023-08-28T09:29:44.081855Z"},"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\nfrom bnunicodenormalizer import Normalizer\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='/kaggle/working/train_select.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.05, 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\" and \n        getattr(config, \"apply_spec_augment\", False) and \n        getattr(config, \"mask_time_prob\", 0) > 0\n    )\n    \n    bnorm = Normalizer()\n    def normalize(sentence):\n        word = [bnorm(word)['normalized']  for word in sentence.split()]\n        return \" \".join([w for w in word if w is not None])\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        input_str = normalize(input_str)\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\"], 10000, f\"train_{training_args.output_dir}\"\n        )\n        save_chunks(\n            vectorized_datasets[\"validation\"], 10000, 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 \"/kaggle/input/indicwhisper-bn/bengali_models/whisper-medium-bn_alldata_multigpu\" \\\n#  --language \"Bengali\" \\\n#  --output_dir \"select-samples\" \\\n#  --preprocessing_num_workers 4 \\\n#  --preprocessing_only \\\n#  --text_column_name \"sentence\" \\\n#  --data_dir \"/kaggle/input/bengaliai-speech\" \\\n#  --min_duration_in_seconds 0.5 \\\n#  --max_duration_in_seconds 30 \\\n#  --apply_spec_augment\n# #  --max_train_samples 75000 \\\n# #  --max_eval_samples 5000 \\","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.28108Z","shell.execute_reply":"2023-07-19T00:13:41.29048Z"},"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":"","metadata":{},"execution_count":null,"outputs":[]}],"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"}}