{"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":"# Acknowledgement","metadata":{}},{"cell_type":"markdown","source":" \nThis notebook is the modified version of the [notebook](https://www.kaggle.com/code/nazmuddhohaansary/wave2vec2-tpu-transfer-learning-dl-sprint) provided by Nazmuddoha ansary, So thanks to him for providing this awesome notebook for training the model using TPU.\n    ","metadata":{}},{"cell_type":"markdown","source":"# Intro","metadata":{}},{"cell_type":"markdown","source":"**In this notebook we demonstarte:**\n* **training wave2vec2 with tensorflow - TPU transfer learning from:[arijitx/wav2vec2-xls-r-300m-bengali](https://huggingface.co/arijitx/wav2vec2-xls-r-300m-bengali)**\n\n\n**The [tfrecord creation script](https://www.kaggle.com/code/nazmuddhohaansary/tfrecords-for-transferlearning/notebook?scriptVersionId=101051948)** for this case\n\n![wav2vec2_structure](https://raw.githubusercontent.com/patrickvonplaten/scientific_images/master/xls_r.png)\n\n* for training with hugging-face torch follow this [notebook](https://www.kaggle.com/code/nazmuddhohaansary/wave2vec2-starter-for-dl-sprint-commonvoice)\n* training with this notebook is way faster due to tfrecords and TPU \n\n![TPU](https://diamond-thumbnails.s3.us-west-2.amazonaws.com/thinkbigcms/Product/logo/6e276569-90f0-4491-a0ed-45b07b8b05eb.png?hash=f54001f5e45cdaa749504dafc8d86bc4) \n\n\nTPU is an accelerator available on **colab** and **kaggle** which provides a way to train **tensorflow** models way faster than on gpus. To train any data on TPU the dataset has to be converted into **tfrecords** format. \n\n\n### Useful links to understand tfrecords and TPU's \n\n**TPU for (~20x)faster training** \n* [what is TPU and why do we need them](https://www.quora.com/What-is-TPU-and-GPU-Why-and-when-do-we-need-them)\n* [Kaggle TPU a-z](https://www.kaggle.com/docs/tpu)\n\n**TFRecords**\n* [Official Tensorflow Doc](https://www.tensorflow.org/tutorials/load_data/tfrecord)\n* [basics](https://www.kaggle.com/code/ryanholbrook/tfrecords-basics/notebook)","metadata":{"id":"8LlpVyXmLn8Z","papermill":{"duration":0.026054,"end_time":"2022-07-13T11:18:01.292419","exception":false,"start_time":"2022-07-13T11:18:01.266365","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"### **install dependencies**","metadata":{"id":"2jsHlsBpOY6T","papermill":{"duration":0.024088,"end_time":"2022-07-13T11:18:01.340118","exception":false,"start_time":"2022-07-13T11:18:01.31603","status":"completed"},"tags":[]}},{"cell_type":"code","source":"!pip install -q git+https://github.com/mnansary/gsoc-wav2vec2.git","metadata":{"executionInfo":{"elapsed":5454,"status":"ok","timestamp":1657278018808,"user":{"displayName":"kugelblitz 1729","userId":"14824418471120817982"},"user_tz":-360},"id":"2FjtlpKILn8e","outputId":"eb050008-442b-4b8e-f463-7f844e6cc907","papermill":{"duration":14.20055,"end_time":"2022-07-13T11:18:15.565426","exception":false,"start_time":"2022-07-13T11:18:01.364876","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-20T15:38:00.829613Z","iopub.execute_input":"2022-08-20T15:38:00.830448Z","iopub.status.idle":"2022-08-20T15:38:14.863636Z","shell.execute_reply.started":"2022-08-20T15:38:00.830338Z","shell.execute_reply":"2022-08-20T15:38:14.862729Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Access","metadata":{"id":"sr3PVY2aNmx9","papermill":{"duration":0.024399,"end_time":"2022-07-13T11:18:15.617793","exception":false,"start_time":"2022-07-13T11:18:15.593394","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"* We locate the tfrecods by file patterns.We use star as wild card entry\n* While training with tfrecords **we must not load the data locally**. We have to use **GCS Buckets** to load the data. \n    * kaggle_datasets api provides a way to access both public and private GCS(google cloud storage). Here we are using public data but private datasets can also be used.\n* **PER_REPLICA_BATCH_SIZE**  global batch size while training will be **8 times the PER_REPLICA_BATCH_SIZE** we provide \n\n* **REC_SIZE=256** simply means while creating the tfrecords , we stored 256 audio files with their labels in one tfrecord\n\n* for params\n```python\nPER_REPLICA_BATCH_SIZE  = 16      # this is a safe batch size \nEPOCHS                  = 50      # change this as needed .. keep the kaggle allowed TPU limit of 9 hours in mind    \n```\n* to use the full-dataset\n\n```python\nTRAIN_GCS_PATTERNS      = [os.path.join(GCS_PATH,\"voted\",\"*/*.tfrecord\")]\n\n```\n\n* VOCAB empty tags:\n    * see the tfrecord creation script above\n    * this is a way to purely transfer learn from the arijitx bengali model","metadata":{"id":"PEZHa0xJLn8g","papermill":{"duration":0.027624,"end_time":"2022-07-13T11:18:15.67035","exception":false,"start_time":"2022-07-13T11:18:15.642726","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import tensorflow as tf\ntry:\n    tpu = tf.distribute.cluster_resolver.TPUClusterResolver()  \n    print('Running on TPU ', tpu.master())\nexcept ValueError:\n    tpu = None\n\nif tpu:\n    tf.config.experimental_connect_to_cluster(tpu)\n    tf.tpu.experimental.initialize_tpu_system(tpu)\n    strategy = tf.distribute.experimental.TPUStrategy(tpu)\n    tf.config.optimizer.set_jit(True)\nelse:\n    strategy = tf.distribute.get_strategy()","metadata":{"execution":{"iopub.status.busy":"2022-08-20T15:38:14.866236Z","iopub.execute_input":"2022-08-20T15:38:14.866621Z","iopub.status.idle":"2022-08-20T15:38:27.775331Z","shell.execute_reply.started":"2022-08-20T15:38:14.866572Z","shell.execute_reply":"2022-08-20T15:38:27.774546Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nuser_credential = user_secrets.get_gcloud_credential()\nuser_secrets.set_tensorflow_credential(user_credential)","metadata":{"execution":{"iopub.status.busy":"2022-08-20T15:38:27.776569Z","iopub.execute_input":"2022-08-20T15:38:27.776881Z","iopub.status.idle":"2022-08-20T15:38:27.981659Z","shell.execute_reply.started":"2022-08-20T15:38:27.776847Z","shell.execute_reply":"2022-08-20T15:38:27.97974Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from kaggle_datasets import KaggleDatasets\n\nGCS_PATH_TRAIN1=KaggleDatasets().get_gcs_path(\"nrecords1\")#Tf records of first 70000 datas except downvote>upvote ones\nGCS_PATH_TRAIN2=KaggleDatasets().get_gcs_path(\"nrecord2\")#Tf records of 70000-140000 datas except downvote>upvote ones\nGCS_PATH_TRAIN3=KaggleDatasets().get_gcs_path(\"nrecords3\")#Tf records of remaining datas except downvote>upvote ones\n\n\nimport os \n#------------------------------\n# change able params\n#------------------------------\nTRAIN_GCS_PATTERNS      = [os.path.join(GCS_PATH_TRAIN1,\"train\",\"*/*.tfrecord\"),\n                          os.path.join(GCS_PATH_TRAIN2,\"train\",\"*/*.tfrecord\"),\n                          os.path.join(GCS_PATH_TRAIN3,\"train\",\"*/*.tfrecord\")\n                        ]\n                           \nEVAL_GCS_PATTERNS       = [os.path.join(GCS_PATH_TRAIN1,\"eval\",\"*/*.tfrecord\")]\n\nPER_REPLICA_BATCH_SIZE  = 16  # this is a safe batch size \nEPOCHS                  = 50      # change this as needed .. keep the kaggle allowed TPU limit of 9 hours in mind    \n\n#------------------------------\n# fixed params while creating the tfrecords\n#------------------------------\nREC_SIZE=256  \nVOCAB   =[\" \", \"_\", \"a\", \"b\", \"c\", \"d\", \"e\", \"f\", \"g\", \"h\", \"i\", \"j\", \"k\", \"l\", \n          \"m\", \"n\", \"o\", \"p\", \"r\", \"s\", \"t\", \"u\", \"v\", \"w\", \"x\", \"y\", \"z\", \"\",\n          \"\", \"œ\", \"।\", \"ঁ\", \"ং\", \"ঃ\", \"অ\", \"আ\", \"ই\", \"ঈ\", \"উ\", \"ঊ\", \"ঋ\", \n          \"এ\", \"ঐ\", \"ও\", \"ঔ\", \"ক\", \"খ\", \"গ\", \"ঘ\", \"ঙ\", \"চ\", \"ছ\", \"জ\", \"ঝ\", \n          \"ঞ\", \"ট\", \"ঠ\", \"ড\", \"ঢ\", \"ণ\", \"ত\", \"থ\", \"দ\", \"ধ\", \"ন\", \"প\", \"ফ\", \n          \"ব\", \"ভ\", \"ম\", \"য\", \"র\", \"ল\", \"শ\", \"ষ\", \"স\", \"হ\", \"়\", \"া\", \"ি\",\n          \"ী\", \"ু\", \"ূ\", \"ৃ\", \"ে\", \"ৈ\", \"ো\", \"ৌ\", \"্\", \"ৎ\", \"ৗ\", \"ড়\", \"ঢ়\",\n          \"য়\", \"০\", \"১\", \"২\", \"৩\", \"৪\", \"৫\", \"৬\", \"৭\", \"৮\", \"৯\", \"ৰ\", \"‌\", \n          \"‍\", \"‎\", \"⁇\",  \"\", \"<s>\", \"</s>\"]\n\nprint(\"Vocab Len:\",len(VOCAB))\nprint(\"Pad Id:\",VOCAB.index(\"\"))\n","metadata":{"id":"bB9uvms3Ln8h","papermill":{"duration":0.520515,"end_time":"2022-07-13T11:18:16.214944","exception":false,"start_time":"2022-07-13T11:18:15.694429","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-20T15:38:27.984174Z","iopub.execute_input":"2022-08-20T15:38:27.98446Z","iopub.status.idle":"2022-08-20T15:38:29.853397Z","shell.execute_reply.started":"2022-08-20T15:38:27.98443Z","shell.execute_reply":"2022-08-20T15:38:29.852451Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We import needed libraries here and collect the tfrecord paths that can be fed into [tf.data api](https://www.tensorflow.org/api_docs/python/tf/data/Dataset)which is the official way to use tfrecords ","metadata":{"id":"cWRDC6oALn8h","papermill":{"duration":0.022723,"end_time":"2022-07-13T11:18:16.261259","exception":false,"start_time":"2022-07-13T11:18:16.238536","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"### Imports and data ","metadata":{"id":"-h_SBtPBSXUT","papermill":{"duration":0.023917,"end_time":"2022-07-13T11:18:16.308703","exception":false,"start_time":"2022-07-13T11:18:16.284786","status":"completed"},"tags":[]}},{"cell_type":"code","source":"#-------------------------------\n# imports\n#-------------------------------\nimport os\nos.environ['TF_CPP_MIN_LOG_LEVEL'] = '3' \nimport random\nimport tensorflow as tf\nimport transformers\nimport matplotlib.pyplot as plt\nimport numpy as np \nimport pandas as pd\nfrom tqdm.auto import tqdm\nfrom IPython.display import display,Audio\nfrom wav2vec2 import RobustWav2Vec2Config,Wav2Vec2,CTCLoss\ntqdm.pandas()\n\n#--------------------------\n# GCS Paths and tfrecords\n#-------------------------\ntrain_recs=[]\neval_recs =[]\ndef get_tfrecs(gcs_pattern):\n    file_paths = tf.io.gfile.glob(gcs_pattern)\n    random.shuffle(file_paths)\n    print(\"found \",len(file_paths), \"tfrecords\")\n    return file_paths\n\nfor gcs in TRAIN_GCS_PATTERNS:\n    print(\"Looking into gcs path:\",gcs)\n    train_recs+=get_tfrecs(gcs)\nfor gcs in EVAL_GCS_PATTERNS:\n    print(gcs)\n    eval_recs+=get_tfrecs(gcs)\n\nprint(\"Total Eval-recs:\",len(eval_recs))\nprint(\"Total Train-recs:\",len(train_recs))\n#------------------------------------------------\n# change config\n#------------------------------------------------\nconfig = RobustWav2Vec2Config()\nconfig.pad_id=VOCAB.index(\"\")","metadata":{"executionInfo":{"elapsed":4450,"status":"ok","timestamp":1657278025337,"user":{"displayName":"kugelblitz 1729","userId":"14824418471120817982"},"user_tz":-360},"id":"Z7fTXNgsLn8i","outputId":"cd14876a-6b69-46ff-c901-859d70f457e5","papermill":{"duration":8.607519,"end_time":"2022-07-13T11:18:24.941917","exception":false,"start_time":"2022-07-13T11:18:16.334398","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-20T15:38:29.854944Z","iopub.execute_input":"2022-08-20T15:38:29.855647Z","iopub.status.idle":"2022-08-20T15:38:32.821986Z","shell.execute_reply.started":"2022-08-20T15:38:29.8556Z","shell.execute_reply":"2022-08-20T15:38:32.820926Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Initialize TPU\n* we initialize the tpu cluster for using\n* based on number of **replicas** or devices we fix:\n    * BATCH_SIZE\n    * STEPS_PER_EPOCH\n    * and evaluation steps within an epoch (EVAL_STEPS)","metadata":{"id":"gxC5GZBmLn8j","papermill":{"duration":0.023475,"end_time":"2022-07-13T11:18:24.989388","exception":false,"start_time":"2022-07-13T11:18:24.965913","status":"completed"},"tags":[]}},{"cell_type":"code","source":"#----------------------------------------------------------\n# Detect hardware, return appropriate distribution strategy\n#----------------------------------------------------------\n# TPU detection. No parameters necessary if TPU_NAME environment variable is set. On Kaggle this is always the case.\n \n    # default distribution strategy in Tensorflow. Works on CPU and single GPU.\n\nprint(\"REPLICAS: \", strategy.num_replicas_in_sync)\n\n#-------------------------------------\n# batching , strategy and steps\n#-------------------------------------\nif strategy.num_replicas_in_sync==1:\n    BATCH_SIZE = PER_REPLICA_BATCH_SIZE\nelse:\n    BATCH_SIZE = PER_REPLICA_BATCH_SIZE*strategy.num_replicas_in_sync\n\n# set    \nSTEPS_PER_EPOCH = (len(train_recs)*REC_SIZE)//(BATCH_SIZE)\nEVAL_STEPS      = (len(eval_recs)*REC_SIZE)//(2*BATCH_SIZE)\nprint(\"Batch Size:\",BATCH_SIZE)\nprint(\"Steps:\",STEPS_PER_EPOCH)\nprint(\"Eval Steps:\",EVAL_STEPS)","metadata":{"executionInfo":{"elapsed":11626,"status":"ok","timestamp":1657278036958,"user":{"displayName":"kugelblitz 1729","userId":"14824418471120817982"},"user_tz":-360},"id":"aoOTdXZpLn8j","outputId":"e365a2b1-b3bf-4e99-9bb5-7697a70b3e68","papermill":{"duration":6.493472,"end_time":"2022-07-13T11:18:31.506845","exception":false,"start_time":"2022-07-13T11:18:25.013373","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-20T15:38:32.823393Z","iopub.execute_input":"2022-08-20T15:38:32.823741Z","iopub.status.idle":"2022-08-20T15:38:32.833956Z","shell.execute_reply.started":"2022-08-20T15:38:32.8237Z","shell.execute_reply":"2022-08-20T15:38:32.833021Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Loader \n* cfg = our data config and some constant storing\n* config=actual wave2vec2 modeling config","metadata":{"id":"IdUqYVZ0Ln8k","papermill":{"duration":0.023923,"end_time":"2022-07-13T11:18:31.554898","exception":false,"start_time":"2022-07-13T11:18:31.530975","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class cfg:\n    audio_shape      =  (246000,)                   # this is actually fixed for the pretrained weights we are using -- highets audio length=15 secs\n    label_shape      =  (250,)                      # this is actually fixed for the pretrained weights we are using \n    sample_rate      =  16000\n    shuffle_buffer   =  1024\n    batch_size       =  BATCH_SIZE\n    vocab_len        =  len(VOCAB)                \n    embed_dim        =  1024\n    ","metadata":{"id":"ydv0NrOtSLvI","papermill":{"duration":0.036038,"end_time":"2022-07-13T11:18:31.615917","exception":false,"start_time":"2022-07-13T11:18:31.579879","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-20T15:38:32.835357Z","iopub.execute_input":"2022-08-20T15:38:32.835721Z","iopub.status.idle":"2022-08-20T15:38:32.853606Z","shell.execute_reply.started":"2022-08-20T15:38:32.835688Z","shell.execute_reply":"2022-08-20T15:38:32.85266Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#------------------------------\n# parsing tfrecords \n#------------------------------\ndef normalize(x):\n    # -> (1, seqlen)\n    mean = tf.reduce_mean(x, axis=-1, keepdims=True)\n    var = tf.math.reduce_variance(x, axis=-1, keepdims=True)\n    return tf.squeeze((x - mean) / tf.sqrt(var + 1e-5))\n\ndef read_raw_audio(audio):\n    wave,rate = tf.audio.decode_wav(audio, desired_channels=1, desired_samples=-1)\n    return tf.reshape(wave, shape=[-1]) \n    \ndef preprocess_example(audio,label):\n    with tf.device(\"/CPU:0\"):\n        signal = normalize(read_raw_audio(audio))\n        label = tf.strings.to_number(tf.strings.split(label), out_type=tf.int32)\n        return signal,label\n\ndef data_input_fn(recs): \n    '''\n      This Function generates data from gcs\n      * The parser function should look similiar now because of datasetEDA\n    '''\n    def _parser(example):   \n        feature ={  'audio' : tf.io.FixedLenFeature([],tf.string) ,\n                    'label' : tf.io.FixedLenFeature([],tf.string) \n        }    \n        example=tf.io.parse_single_example(example,feature)\n        audio,label=preprocess_example(**example)\n        return audio,label\n    # fixed code (for almost all tfrec training)\n    dataset = tf.data.TFRecordDataset(recs)\n    dataset = dataset.map(_parser)\n    dataset = dataset.shuffle(cfg.shuffle_buffer,reshuffle_each_iteration=True)\n    dataset = dataset.repeat()\n    dataset = dataset.padded_batch(cfg.batch_size, padded_shapes=(cfg.audio_shape[0],cfg.label_shape[0]), padding_values=(0.0,VOCAB.index(\"\")))\n    dataset = dataset.prefetch(tf.data.experimental.AUTOTUNE)\n    dataset = dataset.apply(tf.data.experimental.ignore_errors())\n    return dataset","metadata":{"id":"r6HBseySLn8k","papermill":{"duration":0.042851,"end_time":"2022-07-13T11:18:31.684622","exception":false,"start_time":"2022-07-13T11:18:31.641771","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-20T15:38:32.855099Z","iopub.execute_input":"2022-08-20T15:38:32.855367Z","iopub.status.idle":"2022-08-20T15:38:32.87088Z","shell.execute_reply.started":"2022-08-20T15:38:32.855337Z","shell.execute_reply":"2022-08-20T15:38:32.869908Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ds=data_input_fn(train_recs)\neval_ds =data_input_fn(eval_recs)","metadata":{"id":"Qt44lD2ZLn8m","papermill":{"duration":0.703473,"end_time":"2022-07-13T11:18:32.412591","exception":false,"start_time":"2022-07-13T11:18:31.709118","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-20T15:38:32.872417Z","iopub.execute_input":"2022-08-20T15:38:32.873805Z","iopub.status.idle":"2022-08-20T15:38:33.446212Z","shell.execute_reply.started":"2022-08-20T15:38:32.873755Z","shell.execute_reply":"2022-08-20T15:38:33.445165Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Visualize","metadata":{"id":"kr2LdngtSGfy","papermill":{"duration":0.023526,"end_time":"2022-07-13T11:18:32.460498","exception":false,"start_time":"2022-07-13T11:18:32.436972","status":"completed"},"tags":[]}},{"cell_type":"code","source":"#------------------------------\n# view data\n#------------------------------\nfor x,y in train_ds.take(1):\n    signal=x[1].numpy()\n    display(Audio(data=signal, rate=cfg.sample_rate))\n    label=y[1].numpy()\n    sen=\"\".join([VOCAB[int(i)] for i in label])\n    print(\"label:\",sen)\n    print(\"input shape:\",x.shape)\n    print(\"output shape:\",y.shape)","metadata":{"executionInfo":{"elapsed":15132,"status":"ok","timestamp":1657278064182,"user":{"displayName":"kugelblitz 1729","userId":"14824418471120817982"},"user_tz":-360},"id":"_5jRxL1OLn8m","outputId":"c2491bf9-4bbc-40dc-8a17-f6570b9e3fc3","papermill":{"duration":6.663317,"end_time":"2022-07-13T11:18:39.147769","exception":false,"start_time":"2022-07-13T11:18:32.484452","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-20T15:38:33.449568Z","iopub.execute_input":"2022-08-20T15:38:33.449949Z","iopub.status.idle":"2022-08-20T15:38:38.928816Z","shell.execute_reply.started":"2022-08-20T15:38:33.449905Z","shell.execute_reply":"2022-08-20T15:38:38.928025Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Modeling","metadata":{"id":"RapqzxDCLn8n","papermill":{"duration":0.040983,"end_time":"2022-07-13T11:18:39.23034","exception":false,"start_time":"2022-07-13T11:18:39.189357","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"## transfer-learning (torch to tensorflow conversion)\n**[SEE THIS NOTEBOOK FOR DETAILS](https://www.kaggle.com/code/nazmuddhohaansary/converting-pretrained-torch-models-to-tensorflow?scriptVersionId=100675925)**\n","metadata":{"papermill":{"duration":0.040984,"end_time":"2022-07-13T11:18:39.312113","exception":false,"start_time":"2022-07-13T11:18:39.271129","status":"completed"},"tags":[]}},{"cell_type":"code","source":"SUFFIX = \":0\"\nMAPPING = (\n    (\"layer_norm.weight\", \"layer_norm/gamma\"),\n    (\"layer_norm.bias\", \"layer_norm.beta\"),\n    (\"weight\", \"kernel\"),\n    (\".\", \"/\"),\n)\n\nSPECIAL_MAPPING_WITH_HEAD = {\n    \"wav2vec2.encoder.pos_conv_embed.conv.weight_g\": \"wav2vec2/encoder/pos_conv_embed/conv/weight_g:0\",\n    \"wav2vec2.encoder.pos_conv_embed.conv.weight_v\": \"wav2vec2/encoder/pos_conv_embed/conv/weight_v:0\",\n}\n\n\n\ndef replace(k):\n    \"\"\"\n    Converts PyTorch state_dict keys to TensorFlow varible name.\n    \"\"\"\n    for hf_v, tf_v in MAPPING:\n        k = k.replace(hf_v, tf_v)\n    return k + SUFFIX\n\n\ndef get_tf_pretrained_model(hf_model_id,tf_model):\n    \"\"\"\n    Converts HuggingFace PyTorch weights to TensorFlow compatible weights.\n    \"\"\"\n    hf_model = transformers.Wav2Vec2ForCTC.from_pretrained(hf_model_id)\n\n    hf_state_dict = hf_model.state_dict()\n\n    tf_variables = tf_model.variables\n    tf_variables_dict = {}\n    for v in tf_variables:\n        tf_variables_dict[v.name] = v\n\n    tf_weights = []\n    extra_keys = []\n    for k in tqdm(hf_state_dict):\n        if k in SPECIAL_MAPPING_WITH_HEAD:\n            new_k = (SPECIAL_MAPPING_WITH_HEAD[k])\n        else:\n            new_k = replace(k)\n        if new_k not in tf_variables_dict.keys():\n            extra_keys.append(k)\n            print(f\"SKIPPING {k}\")\n            continue\n\n        \n        array = hf_state_dict[k].numpy()\n\n        if k in SPECIAL_MAPPING_WITH_HEAD:\n            array = np.transpose(array, axes=(2, 1, 0))\n        elif \"kernel\" in new_k:\n            array = np.transpose(array)\n\n        tf_weights.append((tf_variables_dict[new_k], array))\n\n    print(\"EXTRA KEYS:\\n\", extra_keys)\n\n    tf.keras.backend.batch_set_value(tf_weights)\n    return tf_model","metadata":{"papermill":{"duration":0.061113,"end_time":"2022-07-13T11:18:39.416096","exception":false,"start_time":"2022-07-13T11:18:39.354983","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-20T15:38:38.929775Z","iopub.execute_input":"2022-08-20T15:38:38.930995Z","iopub.status.idle":"2022-08-20T15:38:38.946788Z","shell.execute_reply.started":"2022-08-20T15:38:38.930948Z","shell.execute_reply":"2022-08-20T15:38:38.945912Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Build the model","metadata":{"papermill":{"duration":0.040709,"end_time":"2022-07-13T11:18:39.498572","exception":false,"start_time":"2022-07-13T11:18:39.457863","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def create_model(cfg):\n    inputs = tf.keras.Input(shape=cfg.audio_shape)\n    # avoid using spec augmentation\n    config.apply_spec_augment=False\n    # freeze feature extractor\n    states = Wav2Vec2(config)(inputs)\n    logits= tf.keras.layers.Dense(cfg.vocab_len,name=\"lm_head\")(states)\n    model = tf.keras.Model(inputs=inputs, outputs=logits)\n    return model\n","metadata":{"papermill":{"duration":0.052328,"end_time":"2022-07-13T11:18:39.592254","exception":false,"start_time":"2022-07-13T11:18:39.539926","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-20T15:38:38.948132Z","iopub.execute_input":"2022-08-20T15:38:38.948544Z","iopub.status.idle":"2022-08-20T15:38:38.963049Z","shell.execute_reply.started":"2022-08-20T15:38:38.948501Z","shell.execute_reply":"2022-08-20T15:38:38.961989Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**model weights can be loaded from saved ones to continue training**\n```python\nmodel.load_weights(\"path to previously trained weights\")\n```","metadata":{"papermill":{"duration":0.043606,"end_time":"2022-07-13T11:18:39.678944","exception":false,"start_time":"2022-07-13T11:18:39.635338","status":"completed"},"tags":[]}},{"cell_type":"code","source":"with strategy.scope():\n    model=get_tf_pretrained_model(\"arijitx/wav2vec2-xls-r-300m-bengali\",create_model(cfg))\n    # freeze feature extractor\n    model.layers[1].freeze_feature_extractor()\n    \"\"\"Uncomment the following line if you want to resume training and \n       you already have trained weights. Replace your trained weights location with the location given. \n    \"\"\"\n    #model.load_weights(\"../input/tf-training-model-loc/model.h5\")\n    model.summary()","metadata":{"executionInfo":{"elapsed":5343,"status":"ok","timestamp":1657278099188,"user":{"displayName":"kugelblitz 1729","userId":"14824418471120817982"},"user_tz":-360},"id":"tauFmCJRLn8p","outputId":"e240cfe6-d25b-4637-dbfe-92685690822f","papermill":{"duration":101.600259,"end_time":"2022-07-13T11:20:21.322264","exception":false,"start_time":"2022-07-13T11:18:39.722005","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-20T15:38:38.964638Z","iopub.execute_input":"2022-08-20T15:38:38.965578Z","iopub.status.idle":"2022-08-20T15:40:48.771423Z","shell.execute_reply.started":"2022-08-20T15:38:38.965509Z","shell.execute_reply":"2022-08-20T15:40:48.770246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training\n* some ideas to extend: \n    * use different schedulers\n    * use callbacks to track some metrics\n    * reduce learning rate on plateau, early stopping setup might need some inspection ","metadata":{"id":"1rRC24zKOTVK","papermill":{"duration":0.043048,"end_time":"2022-07-13T11:20:21.408792","exception":false,"start_time":"2022-07-13T11:20:21.365744","status":"completed"},"tags":[]}},{"cell_type":"code","source":"    \n# early stopping\nearly_stopping = tf.keras.callbacks.EarlyStopping(patience=10, \n                                                  verbose=1, \n                                                  mode = 'auto') \nlr_reducer=tf.keras.callbacks.ReduceLROnPlateau( patience=3)\n\nmodel_save=tf.keras.callbacks.ModelCheckpoint(\"model.h5\",\n                                                save_best_only=True,\n                                                save_weights_only=True,\n                                                verbose=1)\ncallbacks = [model_save]\n\nwith strategy.scope():\n    lr_schedule = tf.keras.experimental.CosineDecay(initial_learning_rate=0.000065, #you can tune the learning rate.\n                                                             decay_steps=600000,\n                                                             alpha= 0.01)\n\n    loss_fn = CTCLoss(config, (PER_REPLICA_BATCH_SIZE,cfg.audio_shape[0]), division_factor=PER_REPLICA_BATCH_SIZE)\n    model.compile(optimizer=tf.keras.optimizers.Adam(lr_schedule),loss=loss_fn)","metadata":{"id":"ji6rjq9xR7ZA","papermill":{"duration":0.110017,"end_time":"2022-07-13T11:20:21.562484","exception":false,"start_time":"2022-07-13T11:20:21.452467","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-20T15:40:48.809314Z","iopub.execute_input":"2022-08-20T15:40:48.809706Z","iopub.status.idle":"2022-08-20T15:40:48.872482Z","shell.execute_reply.started":"2022-08-20T15:40:48.809672Z","shell.execute_reply":"2022-08-20T15:40:48.871359Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history=model.fit(train_ds,\n                  epochs=10,#EPOCHS,\n                  steps_per_epoch=STEPS_PER_EPOCH,\n                  verbose=1,\n                  validation_data=eval_ds,\n                  validation_steps=EVAL_STEPS,\n                  callbacks=callbacks)","metadata":{"executionInfo":{"elapsed":9777524,"status":"ok","timestamp":1657287889381,"user":{"displayName":"kugelblitz 1729","userId":"14824418471120817982"},"user_tz":-360},"id":"xsZRmDMyLn8q","outputId":"4449c29f-0b2d-4461-f0e5-9ea3f7f762fc","papermill":{"duration":23618.062932,"end_time":"2022-07-13T17:53:59.66916","exception":false,"start_time":"2022-07-13T11:20:21.606228","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-20T15:40:48.874045Z","iopub.execute_input":"2022-08-20T15:40:48.874619Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.save_weights('model_latest.h5') # saving the last model","metadata":{"execution":{"iopub.status.busy":"2022-08-20T14:28:36.993482Z","iopub.status.idle":"2022-08-20T14:28:36.993913Z","shell.execute_reply.started":"2022-08-20T14:28:36.993717Z","shell.execute_reply":"2022-08-20T14:28:36.993737Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ncurves={}\nfor key in history.history.keys():\n    curves[key]=history.history[key]\ncurves=pd.DataFrame(curves)\ncurves.to_csv(f\"history.csv\",index=False)\n","metadata":{"id":"zzgH25SYLn8q","papermill":{"duration":4.654906,"end_time":"2022-07-13T17:54:24.899969","exception":false,"start_time":"2022-07-13T17:54:20.245063","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-16T06:58:05.995371Z","iopub.status.idle":"2022-08-16T06:58:05.99589Z","shell.execute_reply.started":"2022-08-16T06:58:05.99571Z","shell.execute_reply":"2022-08-16T06:58:05.995728Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"curves","metadata":{"execution":{"iopub.status.busy":"2022-08-16T06:58:05.996936Z","iopub.status.idle":"2022-08-16T06:58:05.997694Z","shell.execute_reply.started":"2022-08-16T06:58:05.997467Z","shell.execute_reply":"2022-08-16T06:58:05.997492Z"},"trusted":true},"execution_count":null,"outputs":[]}]}