{"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: trainig wave2vec2 with tensorflow - TPU**\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"}},{"cell_type":"markdown","source":"### **install dependencies**","metadata":{"id":"2jsHlsBpOY6T"}},{"cell_type":"code","source":"!pip install -q git+https://github.com/vasudevgupta7/gsoc-wav2vec2@main","metadata":{"id":"2FjtlpKILn8e","executionInfo":{"status":"ok","timestamp":1657278018808,"user_tz":-360,"elapsed":5454,"user":{"displayName":"kugelblitz 1729","userId":"14824418471120817982"}},"outputId":"eb050008-442b-4b8e-f463-7f844e6cc907","execution":{"iopub.status.busy":"2022-07-09T04:27:00.836233Z","iopub.execute_input":"2022-07-09T04:27:00.836533Z","iopub.status.idle":"2022-07-09T04:27:10.881450Z","shell.execute_reply.started":"2022-07-09T04:27:00.836460Z","shell.execute_reply":"2022-07-09T04:27:10.880013Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Access","metadata":{"id":"sr3PVY2aNmx9"}},{"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  = 32      # 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                           os.path.join(GCS_PATH,\"unverified\",\"*/*.tfrecord\"),]\n\n```","metadata":{"id":"PEZHa0xJLn8g"}},{"cell_type":"code","source":"from kaggle_datasets import KaggleDatasets\nGCS_PATH=KaggleDatasets().get_gcs_path(\"dl-sprint-tfrecords\")\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  = 32      # 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   =[ 'pad','start','end','\\u200d',\n        ' ','!',\"'\",',','-','.',':',';','=','?','।',\n        'ঁ','ং','ঃ',\n        'অ','আ','ই','ঈ','উ','ঊ','ঋ','এ','ঐ','ও','ঔ',\n        'ক','খ','গ','ঘ','ঙ',\n        'চ','ছ','জ','ঝ','ঞ',\n        'ট','ঠ','ড','ঢ','ণ',\n        'ত','থ','দ','ধ','ন',\n        'প','ফ','ব','ভ','ম',\n        'য','র','ল',\n        'শ','ষ','স','হ',\n        'া','ি','ী','ু','ূ','ৃ','ে','ৈ','ো','ৌ','্',\n        'ৎ','ড়','ঢ়','য়',\n        '০','১','২','৩','৪','৫','৬','৭','৮','৯']\n","metadata":{"id":"bB9uvms3Ln8h","execution":{"iopub.status.busy":"2022-07-09T04:27:10.883574Z","iopub.execute_input":"2022-07-09T04:27:10.883901Z","iopub.status.idle":"2022-07-09T04:27:11.333171Z","shell.execute_reply.started":"2022-07-09T04:27:10.883860Z","shell.execute_reply":"2022-07-09T04:27:11.332199Z"},"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"}},{"cell_type":"markdown","source":"### Imports and data ","metadata":{"id":"-h_SBtPBSXUT"}},{"cell_type":"code","source":"#-------------------------------\n# imports\n#-------------------------------\nos.environ['TF_CPP_MIN_LOG_LEVEL'] = '3' \nimport random\nimport tensorflow as tf\nimport tensorflow_hub as hub\nimport matplotlib.pyplot as plt\nimport numpy as np \nfrom tqdm.auto import tqdm\nfrom IPython.display import display,Audio\nfrom wav2vec2 import Wav2Vec2Config,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 = Wav2Vec2Config()\nconfig.vocab_size=len(VOCAB)+1\nconfigUseful","metadata":{"id":"Z7fTXNgsLn8i","executionInfo":{"status":"ok","timestamp":1657278025337,"user_tz":-360,"elapsed":4450,"user":{"displayName":"kugelblitz 1729","userId":"14824418471120817982"}},"outputId":"cd14876a-6b69-46ff-c901-859d70f457e5","execution":{"iopub.status.busy":"2022-07-09T04:27:11.335018Z","iopub.execute_input":"2022-07-09T04:27:11.335625Z","iopub.status.idle":"2022-07-09T04:27:13.902732Z","shell.execute_reply.started":"2022-07-09T04:27:11.335580Z","shell.execute_reply":"2022-07-09T04:27:13.901990Z"},"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"}},{"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":{"id":"aoOTdXZpLn8j","executionInfo":{"status":"ok","timestamp":1657278036958,"user_tz":-360,"elapsed":11626,"user":{"displayName":"kugelblitz 1729","userId":"14824418471120817982"}},"outputId":"e365a2b1-b3bf-4e99-9bb5-7697a70b3e68","execution":{"iopub.status.busy":"2022-07-09T04:27:13.904957Z","iopub.execute_input":"2022-07-09T04:27:13.905347Z","iopub.status.idle":"2022-07-09T04:27:22.588273Z","shell.execute_reply.started":"2022-07-09T04:27:13.905315Z","shell.execute_reply":"2022-07-09T04:27:22.587129Z"},"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"}},{"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)+1                # the additional vocab can account for <UNK>\n    ","metadata":{"id":"ydv0NrOtSLvI","execution":{"iopub.status.busy":"2022-07-09T04:27:22.589770Z","iopub.execute_input":"2022-07-09T04:27:22.590060Z","iopub.status.idle":"2022-07-09T04:27:22.595607Z","shell.execute_reply.started":"2022-07-09T04:27:22.590030Z","shell.execute_reply":"2022-07-09T04:27:22.594554Z"},"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, 0))\n    dataset = dataset.prefetch(tf.data.experimental.AUTOTUNE)\n    dataset = dataset.apply(tf.data.experimental.ignore_errors())\n    return dataset","metadata":{"id":"r6HBseySLn8k","execution":{"iopub.status.busy":"2022-07-09T04:27:22.597473Z","iopub.execute_input":"2022-07-09T04:27:22.597733Z","iopub.status.idle":"2022-07-09T04:27:22.613670Z","shell.execute_reply.started":"2022-07-09T04:27:22.597695Z","shell.execute_reply":"2022-07-09T04:27:22.612712Z"},"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","execution":{"iopub.status.busy":"2022-07-09T04:27:22.614945Z","iopub.execute_input":"2022-07-09T04:27:22.615299Z","iopub.status.idle":"2022-07-09T04:27:23.194630Z","shell.execute_reply.started":"2022-07-09T04:27:22.615258Z","shell.execute_reply":"2022-07-09T04:27:23.193506Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Visualize","metadata":{"id":"kr2LdngtSGfy"}},{"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 if i > VOCAB.index(\"end\")])\n    print(\"label:\",sen)\n    print(\"input shape:\",x.shape)\n    print(\"output shape:\",y.shape)","metadata":{"id":"_5jRxL1OLn8m","executionInfo":{"status":"ok","timestamp":1657278064182,"user_tz":-360,"elapsed":15132,"user":{"displayName":"kugelblitz 1729","userId":"14824418471120817982"}},"outputId":"c2491bf9-4bbc-40dc-8a17-f6570b9e3fc3","execution":{"iopub.status.busy":"2022-07-09T04:27:23.196299Z","iopub.execute_input":"2022-07-09T04:27:23.196656Z","iopub.status.idle":"2022-07-09T04:27:26.320783Z","shell.execute_reply.started":"2022-07-09T04:27:23.196612Z","shell.execute_reply":"2022-07-09T04:27:26.319777Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Modeling","metadata":{"id":"RapqzxDCLn8n"}},{"cell_type":"code","source":"def create_model(cfg):\n    load_locally = tf.saved_model.LoadOptions(experimental_io_device='/job:localhost')\n    pretrained_layer = hub.KerasLayer(\"https://tfhub.dev/vasudevgupta7/wav2vec2/1\",load_options=load_locally,trainable=True)\n    inputs = tf.keras.Input(shape=cfg.audio_shape)\n    states = pretrained_layer(inputs)\n    logits= tf.keras.layers.Dense(cfg.vocab_len)(states)\n    model = tf.keras.Model(inputs=inputs, outputs=logits)\n    return model","metadata":{"execution":{"iopub.status.busy":"2022-07-09T04:27:26.322494Z","iopub.execute_input":"2022-07-09T04:27:26.322788Z","iopub.status.idle":"2022-07-09T04:27:26.329314Z","shell.execute_reply.started":"2022-07-09T04:27:26.322757Z","shell.execute_reply":"2022-07-09T04:27:26.328502Z"},"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":{}},{"cell_type":"code","source":"with strategy.scope():\n    model=create_model(cfg)\n    # model.load_weights(\"model.h5\")\nmodel.summary()","metadata":{"id":"tauFmCJRLn8p","executionInfo":{"status":"ok","timestamp":1657278099188,"user_tz":-360,"elapsed":5343,"user":{"displayName":"kugelblitz 1729","userId":"14824418471120817982"}},"outputId":"e240cfe6-d25b-4637-dbfe-92685690822f","execution":{"iopub.status.busy":"2022-07-09T04:27:26.330556Z","iopub.execute_input":"2022-07-09T04:27:26.330801Z","iopub.status.idle":"2022-07-09T04:27:43.047963Z","shell.execute_reply.started":"2022-07-09T04:27:26.330773Z","shell.execute_reply":"2022-07-09T04:27:43.046950Z"},"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"}},{"cell_type":"code","source":"    \n# early stopping\nearly_stopping = tf.keras.callbacks.EarlyStopping(patience=5, \n                                                  verbose=1, \n                                                  mode = 'auto') \ncallbacks = [tf.keras.callbacks.ModelCheckpoint(\"model.h5\",\n                                                save_best_only=True,\n                                                save_weights_only=True,\n                                                verbose=1),\n             early_stopping]\n\nwith strategy.scope():\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(5e-5),\n                  loss=loss_fn)","metadata":{"id":"ji6rjq9xR7ZA","execution":{"iopub.status.busy":"2022-07-09T04:31:43.927553Z","iopub.execute_input":"2022-07-09T04:31:43.927916Z","iopub.status.idle":"2022-07-09T04:31:44.566798Z","shell.execute_reply.started":"2022-07-09T04:31:43.927885Z","shell.execute_reply":"2022-07-09T04:31:44.565931Z"},"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":{"id":"xsZRmDMyLn8q","executionInfo":{"status":"ok","timestamp":1657287889381,"user_tz":-360,"elapsed":9777524,"user":{"displayName":"kugelblitz 1729","userId":"14824418471120817982"}},"outputId":"4449c29f-0b2d-4461-f0e5-9ea3f7f762fc","execution":{"iopub.status.busy":"2022-07-09T04:31:45.723816Z","iopub.execute_input":"2022-07-09T04:31:45.724401Z"},"trusted":true},"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","execution":{"iopub.status.busy":"2022-07-09T04:28:15.504423Z","iopub.status.idle":"2022-07-09T04:28:15.505040Z","shell.execute_reply.started":"2022-07-09T04:28:15.504810Z","shell.execute_reply":"2022-07-09T04:28:15.504849Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"curves","metadata":{"id":"lwGtVNsnHCRF","executionInfo":{"status":"ok","timestamp":1657293740858,"user_tz":-360,"elapsed":462,"user":{"displayName":"kugelblitz 1729","userId":"14824418471120817982"}},"outputId":"17138be5-b3e0-4e7d-9d0c-79fe1784a689","execution":{"iopub.status.busy":"2022-07-09T04:28:15.506429Z","iopub.status.idle":"2022-07-09T04:28:15.507134Z","shell.execute_reply.started":"2022-07-09T04:28:15.506925Z","shell.execute_reply":"2022-07-09T04:28:15.506946Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"id":"U89_QPqmaxBN"},"execution_count":null,"outputs":[]}]}