{"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":"import pandas as pd\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport librosa\n!pip install datasets jiwer path bnunicodenormalizer loralib\nfrom pathlib import Path\nfrom torch.utils.data import Dataset, DataLoader\n# import transformers\n# from transformers import Wav2Vec2CTCTokenizer,Wav2Vec2FeatureExtractor,Wav2Vec2ForCTC,Wav2Vec2Processor\nfrom functools import partial\nfrom torch.optim import AdamW\nfrom torch.nn import Module\nimport loralib as lora\nfrom torch import Tensor\nfrom typing import List, Optional, Tuple\n# from transformers import Trainer\nfrom datasets import load_metric\nimport numpy as np\nimport torch\nimport gc\nfrom bnunicodenormalizer import Normalizer \nimport torchaudio\nimport torchtext\nfrom torchtext.data import get_tokenizer\nfrom torch.nn.utils.rnn import pad_sequence\n\nimport os\nos.environ.pop('TPU_PROCESS_ADDRESSES')\nos.environ.pop('CLOUD_TPU_TASK_ID')\nos.environ['XLA_USE_BF16']='1'\n\n\nimport torch_xla.core.xla_model as xm\nimport torch_xla.distributed.parallel_loader as pl\nimport torch_xla.distributed.xla_multiprocessing as xmp\n\nTORCH_SEED=3\nimport random\nrandom.seed(TORCH_SEED)\nos.environ['PYTHONHASHSEED'] = str(TORCH_SEED)\nnp.random.seed(TORCH_SEED)\ntorch.manual_seed(TORCH_SEED)\n\ndf= pd.read_csv('/kaggle/input/create-wav-train-dataset/scores_train.csv')\ntrain_pth='/kaggle/input/bengaliai-speech/train_mp3s'\ntrain_df = pd.read_csv('/kaggle/input/bengaliai-speech/train.csv')\n\nbundle = torchaudio.pipelines.WAV2VEC2_XLSR_1B\n\nimport json\nfilename = \"/kaggle/input/bengali-sr-download-public-trained-models/wav2vec2-xls-r-300m-bengali/vocab.json\"\nwith open(filename, \"r\") as read_file:\n    vocab = json.load(read_file)\nvocab[' ']=110 #121 tokens in total\nfor i in range(10):\n    vocab[str(i)]= 111+i\n\nimport re\nbnorm= Normalizer()\nchars_to_ignore_regex = '[\\,\\?\\.\\!\\-\\;\\:\\\"\\—\\‘\\'\\‚\\“\\”\\…\\|]'\n\ndef prep(seq):\n    batch = re.sub(chars_to_ignore_regex, '', seq) + \" \"\n    _words = [bnorm(word)['normalized']  for word in batch.split()]\n    return \" \".join([word for word in _words if word is not None])\n\ntrain = train_df[train_df['split']=='train']\ntrain= train.sample(20000,weights= pd.merge(train,df,on='id').ykg_wer.apply(lambda x:np.exp(-8*x)))\nval = train_df[train_df['split']=='valid']\nval= val.sample(4000)\n\nclass SrDataset(Dataset):\n    def __init__(self,df,sr=16_000,vocab=vocab,is_train=True):\n        self.df= df\n        self.sr= sr\n        self.vocab= vocab\n        self.tokenize = get_tokenizer(lambda x: list(x),'be')\n        self.is_train=is_train\n    def __len__(self):\n        return len(self.df)\n    def __getitem__(self,i):\n        idx= self.df.iloc[i]\n        speech, sr = librosa.load((Path(train_pth) / idx.id).with_suffix('.mp3'), sr=self.sr)\n\n        label= self.tokenize(prep(idx.sentence))\n        label=list(map(lambda x:self.vocab.get(x,self.vocab['[UNK]']),label))\n#         assert(len(speech)>= len(label))\n        return (speech,label)\n\ndef collate_fn_pad(batch):\n    x= pad_sequence([torch.tensor(i[0],dtype=torch.float32) for i in batch],True,0)\n    y= pad_sequence([torch.tensor(i[1]) for i in batch],True,0)\n    return x,y\n\n\ndef loss_fn(y_pred,y_true):\n    trgt_length= torch.count_nonzero(y_true, axis=-1)\n    input_length= torch.full((y_true.shape[0],),y_pred.shape[1])\n    preds= torch.transpose(F.log_softmax(y_pred,dim=-1),0,1)\n    return F.ctc_loss(preds,y_true,input_length,trgt_length,blank= 109,zero_infinity=False)\n\nwer_metric = load_metric(\"wer\")\na_vocab= {j:i for i,j in vocab.items() if i not in ['[UNK]','[PAD]','|']}\n\ndef decode(seqs,pred=True):\n    x=[]\n    for seq in seqs:\n        if pred:\n            seq = torch.unique_consecutive(seq, dim=-1)\n        joined = \"\".join([a_vocab.get(int(i),'') for i in seq])\n        x.append(joined)\n    return x\n\ndef compute_metrics(y_pred,y_true):\n    y_pred = torch.argmax(y_pred, dim=-1)  # [num_seq,]\n    pred_str = decode(y_pred)\n    label_str = decode(y_true, False)\n\n    wer = wer_metric.compute(predictions=pred_str, references=label_str)\n\n    return wer","metadata":{"execution":{"iopub.status.busy":"2023-10-17T08:46:40.617549Z","iopub.execute_input":"2023-10-17T08:46:40.618039Z","iopub.status.idle":"2023-10-17T08:48:09.444836Z","shell.execute_reply.started":"2023-10-17T08:46:40.617995Z","shell.execute_reply":"2023-10-17T08:48:09.443843Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def model_load():\n    mdl = MyModel()\n    s=0\n    for i,j in mdl.named_parameters():\n        if '47' in i:\n            j.requires_grad=True\n            s+= j.numel()\n        else:\n            j.requires_grad=False\n    return mdl\n\nclass MyModel(nn.Module):\n    def __init__(self):\n        \n        super().__init__()\n        self.model = bundle.get_model()\n\n    def forward(self,x):\n        x= self.model(x)\n        return x[0]\nmodel_load()","metadata":{"execution":{"iopub.status.busy":"2023-10-17T08:48:49.238340Z","iopub.execute_input":"2023-10-17T08:48:49.238755Z","iopub.status.idle":"2023-10-17T08:48:49.246114Z","shell.execute_reply.started":"2023-10-17T08:48:49.238723Z","shell.execute_reply":"2023-10-17T08:48:49.245329Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SERIAL_EXEC = xmp.MpSerialExecutor()\ndef finetune(rank):\n  print('Starting', rank)\n\n  def get_dataset():\n    train_dataset= SrDataset(train.iloc[:64])\n    return train_dataset\n  \n  train_dataset = SERIAL_EXEC.run(get_dataset)\n\n  train_sampler = torch.utils.data.distributed.DistributedSampler(\n      train_dataset,\n      num_replicas=xm.xrt_world_size(),\n      rank=xm.get_ordinal(),\n      shuffle=True)\n\n  \n  rng = torch.Generator().manual_seed(TORCH_SEED)\n  train_loader = torch.utils.data.DataLoader(\n      train_dataset,\n      collate_fn=collate_fn_pad,\n      batch_size=4,\n      num_workers=1,\n      sampler=train_sampler,\n      drop_last=True,\n      generator=rng,\n  )\n    \n  device = xm.xla_device()\n  model = model_load().to(device)\n  optimizer = AdamW(filter(lambda p: p.requires_grad, model.parameters()), lr= 8e-5,weight_decay= 1e-2)\n\n  def train_loop_fn(loader, epoch):\n    loss_,met_,step=None,0,0\n#     loss_,step=None,0\n    model.train()\n    for s, (i,j) in enumerate(loader):\n      optimizer.zero_grad()\n      outputs = model(i)\n      loss = loss_fn(outputs,j)\n      loss.backward()\n      xm.optimizer_step(optimizer)\n      loss_log= loss.detach()\n      if loss_:\n        loss_= loss_+ loss_log\n      else: \n        loss_ = loss_log\n      step+=1\n      met_+= compute_metrics(outputs.detach().cpu(),j.detach().cpu())\n      del i,j,outputs,loss\n      gc.collect()\n      print('iter',s,'complete',xm.get_ordinal())\n    xm.master_print(f'loss: {loss_.item()/step:.3f} metric: {met_/step:.2f} ')\n    del loss_,step\n    gc.collect()\n  \n  train_device_loader =  pl.MpDeviceLoader(train_loader, device) \n\n  for epoch in range(1, 2+ 1):\n    xm.master_print(\"Started epoch {}\".format(epoch))\n    train_loop_fn(train_device_loader, epoch)","metadata":{"execution":{"iopub.status.busy":"2023-10-17T08:53:07.043242Z","iopub.execute_input":"2023-10-17T08:53:07.043704Z","iopub.status.idle":"2023-10-17T08:53:07.055447Z","shell.execute_reply.started":"2023-10-17T08:53:07.043658Z","shell.execute_reply":"2023-10-17T08:53:07.054724Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"xmp.spawn(finetune, args=(),start_method='fork')","metadata":{"execution":{"iopub.status.busy":"2023-10-17T08:53:09.242444Z","iopub.execute_input":"2023-10-17T08:53:09.242879Z","iopub.status.idle":"2023-10-17T09:01:11.781462Z","shell.execute_reply.started":"2023-10-17T08:53:09.242846Z","shell.execute_reply":"2023-10-17T09:01:11.780191Z"},"trusted":true},"execution_count":null,"outputs":[]}]}