{"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":"**In this notebook we demonstarte:**\n* **trainig 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.316030","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-07-18T11:26:24.855413Z","iopub.execute_input":"2022-07-18T11:26:24.855740Z","iopub.status.idle":"2022-07-18T11:26:37.508916Z","shell.execute_reply.started":"2022-07-18T11:26:24.855651Z","shell.execute_reply":"2022-07-18T11:26:37.508009Z"},"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.670350","exception":false,"start_time":"2022-07-13T11:18:15.642726","status":"completed"},"tags":[]}},{"cell_type":"code","source":"from kaggle_datasets import KaggleDatasets\nGCS_PATH=KaggleDatasets().get_gcs_path(\"tfrecs-for-arijitx\")\nimport os \n#------------------------------\n# change able params\n#------------------------------\nTRAIN_GCS_PATTERNS      = [os.path.join(GCS_PATH,\"voted\",\"*/*.tfrecord\")]\n                           \nEVAL_GCS_PATTERNS       = [os.path.join(GCS_PATH,\"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   =[' ', '<empty>', '<empty>', '<empty>', '<empty>', '<empty>', '<empty>', '<empty>', '<empty>', \n          '<empty>', '<empty>', '<empty>', '<empty>', '<empty>', '<empty>', '<empty>', '<empty>', \n          '<empty>', '<empty>', '<empty>', '<empty>', '<empty>', '<empty>', '<empty>', '<empty>', \n          '<empty>', '<empty>', '<empty>', '<empty>', '<empty>', '।', 'ঁ', 'ং', 'ঃ', 'অ', 'আ', 'ই', \n          'ঈ', 'উ', 'ঊ', 'ঋ', 'এ', 'ঐ', 'ও', 'ঔ', 'ক', 'খ', 'গ', 'ঘ', 'ঙ', 'চ', 'ছ', 'জ', 'ঝ', 'ঞ', \n          'ট', 'ঠ', 'ড', 'ঢ', 'ণ', 'ত', 'থ', 'দ', 'ধ', 'ন', 'প', 'ফ', 'ব', 'ভ', 'ম', 'য', 'র', 'ল',\n          'শ', 'ষ', 'স', 'হ', '<empty>', 'া', 'ি', 'ী', 'ু', 'ূ', 'ৃ', 'ে', 'ৈ', 'ো', 'ৌ', '্', 'ৎ', \n          '<empty>', 'ড়', 'ঢ়', 'য়', '০', '১', '২', '৩', '৪', '৫', '৬', '৭', '৮', '৯', '<empty>', '<empty>', \n          '\\u200d', '<empty>', '<empty>', '', '<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-07-18T11:26:37.511760Z","iopub.execute_input":"2022-07-18T11:26:37.512172Z","iopub.status.idle":"2022-07-18T11:26:37.898541Z","shell.execute_reply.started":"2022-07-18T11:26:37.512121Z","shell.execute_reply":"2022-07-18T11:26:37.897689Z"},"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#-------------------------------\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-07-18T11:26:37.900365Z","iopub.execute_input":"2022-07-18T11:26:37.900682Z","iopub.status.idle":"2022-07-18T11:26:45.441074Z","shell.execute_reply.started":"2022-07-18T11:26:37.900641Z","shell.execute_reply":"2022-07-18T11:26:45.439972Z"},"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.\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() \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-07-18T11:26:45.442945Z","iopub.execute_input":"2022-07-18T11:26:45.443191Z","iopub.status.idle":"2022-07-18T11:26:51.934953Z","shell.execute_reply.started":"2022-07-18T11:26:45.443154Z","shell.execute_reply":"2022-07-18T11:26:51.934058Z"},"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-07-18T11:26:51.936076Z","iopub.execute_input":"2022-07-18T11:26:51.936317Z","iopub.status.idle":"2022-07-18T11:26:51.941880Z","shell.execute_reply.started":"2022-07-18T11:26:51.936275Z","shell.execute_reply":"2022-07-18T11:26:51.941027Z"},"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-07-18T11:26:51.942944Z","iopub.execute_input":"2022-07-18T11:26:51.943177Z","iopub.status.idle":"2022-07-18T11:26:51.957386Z","shell.execute_reply.started":"2022-07-18T11:26:51.943149Z","shell.execute_reply":"2022-07-18T11:26:51.956734Z"},"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-07-18T11:26:51.958212Z","iopub.execute_input":"2022-07-18T11:26:51.958451Z","iopub.status.idle":"2022-07-18T11:26:52.515069Z","shell.execute_reply.started":"2022-07-18T11:26:51.958425Z","shell.execute_reply":"2022-07-18T11:26:52.514200Z"},"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 eval_ds.take(1):\n    signal=x[0].numpy()\n    display(Audio(data=signal, rate=cfg.sample_rate))\n    label=y[0].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-07-18T11:26:52.516882Z","iopub.execute_input":"2022-07-18T11:26:52.517531Z","iopub.status.idle":"2022-07-18T11:26:58.031754Z","shell.execute_reply.started":"2022-07-18T11:26:52.517482Z","shell.execute_reply":"2022-07-18T11:26:58.031133Z"},"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.230340","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-07-18T11:26:58.032875Z","iopub.execute_input":"2022-07-18T11:26:58.033605Z","iopub.status.idle":"2022-07-18T11:26:58.051115Z","shell.execute_reply.started":"2022-07-18T11:26:58.033560Z","shell.execute_reply":"2022-07-18T11:26:58.050192Z"},"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-07-18T11:26:58.053206Z","iopub.execute_input":"2022-07-18T11:26:58.053851Z","iopub.status.idle":"2022-07-18T11:26:58.078652Z","shell.execute_reply.started":"2022-07-18T11:26:58.053821Z","shell.execute_reply":"2022-07-18T11:26:58.077908Z"},"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    model.load_weights(\"../input/arijitx-transfer-tf-weights/model.h5\")\nmodel.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-07-18T11:26:58.079997Z","iopub.execute_input":"2022-07-18T11:26:58.080397Z","iopub.status.idle":"2022-07-18T11:29:26.829700Z","shell.execute_reply.started":"2022-07-18T11:26:58.080370Z","shell.execute_reply":"2022-07-18T11:29:26.828054Z"},"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.0001,\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),\n                  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-07-18T11:29:26.831587Z","iopub.execute_input":"2022-07-18T11:29:26.833341Z","iopub.status.idle":"2022-07-18T11:29:26.892775Z","shell.execute_reply.started":"2022-07-18T11:29:26.833294Z","shell.execute_reply":"2022-07-18T11:29:26.891677Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history=model.fit(train_ds,\n                  epochs=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.669160","exception":false,"start_time":"2022-07-13T11:20:21.606228","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-18T11:29:26.894069Z","iopub.execute_input":"2022-07-18T11:29:26.894391Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"papermill":{"duration":4.677114,"end_time":"2022-07-13T17:54:15.646214","exception":false,"start_time":"2022-07-13T17:54:10.969100","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"curves={}\nfor key in history.history.keys():\n    curves[key]=history.history[key]\ncurves=pd.DataFrame(curves)\ncurves.to_csv(f\"history.csv\",index=False)","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":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"curves","metadata":{"executionInfo":{"elapsed":462,"status":"ok","timestamp":1657293740858,"user":{"displayName":"kugelblitz 1729","userId":"14824418471120817982"},"user_tz":-360},"id":"lwGtVNsnHCRF","outputId":"17138be5-b3e0-4e7d-9d0c-79fe1784a689","papermill":{"duration":4.663639,"end_time":"2022-07-13T17:54:34.322956","exception":false,"start_time":"2022-07-13T17:54:29.659317","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Caution: \n* the model is not trained on the whole dataset\n* spec-augmentation is not used due to eager execution problems\n* no masking is used (which is definitely a good idea to try)\n","metadata":{"papermill":{"duration":4.776325,"end_time":"2022-07-13T17:54:43.692808","exception":false,"start_time":"2022-07-13T17:54:38.916483","status":"completed"},"tags":[]}},{"cell_type":"code","source":"","metadata":{"id":"U89_QPqmaxBN","papermill":{"duration":4.717834,"end_time":"2022-07-13T17:54:53.042957","exception":false,"start_time":"2022-07-13T17:54:48.325123","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]}]}