{"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"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":52324,"databundleVersionId":6229904,"sourceType":"competition"},{"sourceId":6152784,"sourceType":"datasetVersion","datasetId":3528942}],"dockerImageVersionId":30528,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"%%capture\n! pip install transformers -q\n! pip install jiwer -q\n! pip install --upgrade wandb -q\n! pip install --upgrade librosa -q","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport json\nimport numpy as np\nimport pandas as pd\nfrom tqdm.auto import tqdm\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts\n\nfrom transformers import (\n    Wav2Vec2ForCTC,\n    Wav2Vec2Processor,\n    Wav2Vec2CTCTokenizer,\n    Wav2Vec2FeatureExtractor\n) \n\nimport wandb\nimport librosa\n\nimport warnings\nwarnings.simplefilter('ignore')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"Config = {\n    'audio_dir': '/kaggle/input/bengaliai-speech/train_mp3s',\n#     'model_name': 'facebook/wav2vec2-base',\n    'model_name': 'facebook/wav2vec2-large-xlsr-53',\n    'lr': 3e-4,\n    'wd': 1e-5,\n    'T_0': 10,\n    'T_mult': 2,\n    'eta_min': 1e-6,\n    'nb_epochs': 5,\n    'train_bs': 16,\n    'valid_bs': 16,\n    'sampling_rate': 16000,\n    '_wandb_kernel': 'bengali-ai-wav2vec2',\n}","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### About W&B:\n<center><img src=\"https://i.imgur.com/gb6B4ig.png\" width=\"400\" alt=\"Weights & Biases\"/></center><br>\n<p style=\"text-align:center\">WandB is a developer tool for companies turn deep learning research projects into deployed software by helping teams track their models, visualize model performance and easily automate training and improving models.\nWe will use their tools to log hyperparameters and output metrics from your runs, then visualize and compare results and quickly share findings with your colleagues.<br><br></p>","metadata":{}},{"cell_type":"code","source":"# W&B Logging\nfrom kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nwb_key = user_secrets.get_secret(\"be26cc81340e4b18cb24eefa49797a15333733fc\")\n\nwandb.login(key=wb_key)\n\nrun = wandb.init(\n    project='pytorch',\n    config=Config,\n    group='asr',\n    job_type='train',\n)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def read_audio(mp3_path, target_sr=16000):\n    \"\"\"\n    Loads an mp3 audio file and resamples it to 16kHz \n    Required for needed for Wav2Vec2 training\n    \"\"\"\n    audio, sr = librosa.load(mp3_path, sr=32000)\n    audio_array = librosa.resample(audio, orig_sr=sr, target_sr=target_sr)\n    return audio_array\n\ndef construct_vocab(texts):\n    \"\"\"\n    Get unique characters from all the text in a list\n    \"\"\"\n    all_text = \" \".join(texts)\n    vocab = list(set(all_text))\n    return vocab\n\ndef wandb_log(**kwargs):\n    for k, v in kwargs.items():\n        wandb.log({k: v})\n\ndef save_vocab(dataframe):\n    \"\"\"\n    Saves the processed vocab file as 'vocab.json', to be ingested by tokenizer\n    \"\"\"\n    vocab = construct_vocab(dataframe['sentence'].tolist())\n    vocab_dict = {v: k for k, v in enumerate(vocab)}\n    vocab_dict[\"__\"] = vocab_dict[\" \"]\n    _ = vocab_dict.pop(\" \")\n    vocab_dict[\"[UNK]\"] = len(vocab_dict)\n    vocab_dict[\"[PAD]\"] = len(vocab_dict)\n\n    with open('vocab.json', 'w') as fl:\n        json.dump(vocab_dict, fl)\n\n    print(\"Created Vocab file!\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ASRDataset(Dataset):\n    def __init__(self, df, config, is_test=False):\n        self.df = df\n        self.config = config\n        self.is_test = is_test\n    \n    def __getitem__(self, idx):\n        # First read and pre-process the audio file\n        audio = read_audio(self.df.loc[idx]['path'])\n        audio = processor(\n            audio, \n            sampling_rate=self.config['sampling_rate']\n        ).input_values[0]\n        \n        # Return -1 for label if in test-only mode\n        if self.is_test:\n            return {'audio': audio, 'label': -1}\n        else:\n            # If we are training/validating, also process the labels (actual sentences)\n            with processor.as_target_processor():\n                labels = processor(self.df.loc[idx]['sentence']).input_ids\n            return {'audio': audio, 'label': labels}\n        \n    def __len__(self):\n        return len(self.df)\n    \ndef ctc_data_collator(batch):\n    \"\"\"\n    Custom data collator function to dynamically pad the data\n    \"\"\"\n    input_features = [{\"input_values\": sample[\"audio\"]} for sample in batch]\n    label_features = [{\"input_ids\": sample[\"label\"]} for sample in batch]\n    batch = processor.pad(\n        input_features,\n        padding=True,\n        return_tensors=\"pt\",\n    )\n    with processor.as_target_processor():\n        labels_batch = processor.pad(\n            label_features,\n            padding=True,\n            return_tensors=\"pt\",\n        )\n        \n    labels = labels_batch[\"input_ids\"].masked_fill(labels_batch.attention_mask.ne(1), -100)\n    batch[\"labels\"] = labels\n    return batch","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_one_epoch(model, train_loader, optimizer, scheduler, device='cuda:0'):\n    model.train()\n    pbar = tqdm(train_loader, total=len(train_loader))\n    avg_loss = 0\n    for data in pbar:\n        data = {k: v.to(device) for k, v in data.items()}\n        loss = model(**data).loss\n        loss_itm = loss.item()\n        \n        avg_loss += loss_itm\n        pbar.set_description(f\"loss: {loss_itm:.4f}\")\n        wandb_log(train_step_loss=loss_itm)\n        \n        optimizer.zero_grad(set_to_none=True)\n        loss.backward()\n        optimizer.step()\n        scheduler.step()\n        \n    return avg_loss / len(train_loader)\n\n@torch.no_grad()\ndef valid_one_epoch(model, valid_loader, device='cuda:0'):\n    pbar = tqdm(valid_loader, total=len(valid_loader))\n    avg_loss = 0\n    for data in pbar:\n        data = {k: v.to(device) for k, v in data.items()}\n        loss = model(**data).loss\n        loss_itm = loss.item()\n        \n        avg_loss += loss_itm\n        pbar.set_description(f\"val_loss: {loss_itm:.4f}\")\n        wandb_log(valid_step_loss=loss_itm)\n\n    return avg_loss / len(valid_loader)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if __name__ == \"__main__\":\n    # Read in the dataframe and split by training and validation splits\n    df = pd.read_csv(\"/kaggle/input/bengaliai-speech/train.csv\")\n    \n    # Get a paths feature for reading in during dataloading\n    df['path'] = df['id'].apply(lambda x: os.path.join(Config['audio_dir'], x+'.mp3'))\n    train_df = df[df['split'] == 'train'].sample(frac=.05).reset_index(drop=True)\n    valid_df = df[df['split'] == 'valid'].sample(frac=.05).reset_index(drop=True)\n    print(f\"Training on samples: {len(train_df)}, Validation on samples: {len(valid_df)}\")\n\n    # Construct and save the vocab file\n    save_vocab(df)\n    \n    # Init the tokenizer, feature_extractor, processor and model\n    tokenizer = Wav2Vec2CTCTokenizer(\n        \"./vocab.json\", \n        unk_token=\"[UNK]\",\n        pad_token=\"[PAD]\",\n        word_delimiter_token=\"__\"\n    )\n    feature_extractor = Wav2Vec2FeatureExtractor(\n        feature_size=1, \n        sampling_rate=Config['sampling_rate'], \n        padding_value=0.0, \n        do_normalize=True, \n        return_attention_mask=False\n    )\n    processor = Wav2Vec2Processor(\n        feature_extractor=feature_extractor, \n        tokenizer=tokenizer\n    )\n\n    model = Wav2Vec2ForCTC.from_pretrained(\n        Config['model_name'],\n        ctc_loss_reduction=\"mean\", \n        pad_token_id=processor.tokenizer.pad_token_id,\n        vocab_size = len(tokenizer),\n    )\n    wandb.watch(model)\n    \n    # Freeze the feature encoder part since we won't be training it\n    model.to('cuda')\n    model.freeze_feature_encoder()\n    optimizer = torch.optim.AdamW(\n        model.parameters(), \n        lr=Config['lr'], \n        weight_decay=Config['wd']\n    )\n    scheduler = CosineAnnealingWarmRestarts(\n        optimizer,\n        T_0=Config['T_0'],\n        T_mult=Config['T_mult'],\n        eta_min=Config['eta_min']\n    )\n    \n    # Construct training and validation dataloaders\n    train_ds = ASRDataset(train_df, Config)\n    valid_ds = ASRDataset(valid_df, Config)\n    \n    train_loader = DataLoader(\n        train_ds, \n        batch_size=Config['train_bs'], \n        collate_fn=ctc_data_collator, \n    )\n    valid_loader = DataLoader(\n        valid_ds,\n        batch_size=Config['valid_bs'],\n        collate_fn=ctc_data_collator,\n    )\n    \n    # Train the model\n    best_loss = float('inf')\n    for epoch in range(Config['nb_epochs']):\n        print(f\"{'='*40} Epoch: {epoch+1} / {Config['nb_epochs']} {'='*40}\")\n        train_loss = train_one_epoch(model, train_loader, optimizer, scheduler)\n        valid_loss = valid_one_epoch(model, valid_loader)\n        wandb_log(train_loss=train_loss, val_loss=valid_loss)\n        print(f\"train_loss: {train_loss:.4f}, valid_loss: {valid_loss:.4f}\")\n        \n        if valid_loss < best_loss:\n            best_loss = valid_loss\n            torch.save(model.state_dict(), f\"wav2vec2_base_bengaliAI.pt\")\n            print(f\"Saved the best model so far with val_loss: {valid_loss:.4f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Once training is done, \nwandb.finish()","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}