{"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":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Converting pretrained wave2vec2-xls-r300m torch model to tensorflow \n## WHY?\n* we want to train model in TPU to train them faster \n* we cant do that with **torch models** \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)\n\n**WE HAVE THE TFRECORDS PREPARED [HERE IN THIS DATASET](https://www.kaggle.com/datasets/ocrteamriad/dl-sprint-tfrecords)**","metadata":{}},{"cell_type":"markdown","source":"## TODO \n* we can implement wave2vec2 model from scratch \n    * wave2vec2 has 7 layers of convolutional blocks \n    * then there are transformer blocks for attention\n\n\n### FORTUNATELY FOR US THESE LAYERS ARE ALREADY IMPLEMENTED \n* **[HERE: transformer encoder](https://github.com/vasudevgupta7/gsoc-wav2vec2/blob/main/src/wav2vec2/encoder.py)**\n* and **[HERE: feature extractor ](https://github.com/vasudevgupta7/gsoc-wav2vec2/blob/main/src/wav2vec2/feature_extractor.py)**\n\nSo we will reuse this. ","metadata":{}},{"cell_type":"code","source":"!pip install git+https://github.com/mnansary/gsoc-wav2vec2.git@main","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"* **[there is code in this repo to convert the weights](https://github.com/vasudevgupta7/gsoc-wav2vec2/blob/main/src/convert_torch_to_tf.py)**\n\n**However this actually fail and wont serve our purpose**\n\n* the script only covers the following model conversion : \n```python\nACCEPTABLE_HF_IDS = [\"facebook/wav2vec2-base-960h\", \"facebook/wav2vec2-base\", \"facebook/wav2vec2-large-robust\", \"facebook/wav2vec2-large-xlsr-53\"]\n```\n\n**WE CAN HOWEVER REUSE THE FUNCTIONS WITH SOME CHANGES**\n\n","metadata":{}},{"cell_type":"markdown","source":"# Model Selection and conversion\n* The model we want to convert is **[arijitx/wav2vec2-xls-r-300m-bengali](https://huggingface.co/arijitx/wav2vec2-xls-r-300m-bengali)**\n* we inspect these two configs [tensorflow config](https://github.com/vasudevgupta7/gsoc-wav2vec2/blob/main/src/wav2vec2/config.py) and [hugging face config](https://huggingface.co/arijitx/wav2vec2-xls-r-300m-bengali/blob/main/config.json) and spot the differences\n\n| Tensorflow |Huggingface |\n|:---:|:---:|\n|num_heads: int = 12|\"num_attention_heads\": 16,|\n|num_layers: int = 12|\"num_hidden_layers\": 24,|\n|conv_bias: bool = False|\"conv_bias\": true,|\n|conv_bias: bool = False|\"conv_bias\": true,|\n|feature_extractor_norm_type: bool = \"group\"|\"feat_extract_norm\": \"layer\",|\n|hidden_size: int = 768|\"hidden_size\": 1024,|\n|intermediate_size: int = 3072|\"intermediate_size\": 4096,|\n\n**Note:we can safely ignore differences like dropout while conversion**\n\nThe changes can be executed by :\n\n```python\nconfig = Wav2Vec2Config()\nconfig.num_heads=16\nconfig.num_layers=24\nconfig.conv_bias=True\nconfig.feature_extractor_norm_type=\"layer\"\nconfig.hidden_size=1024\nconfig.intermediate_size=4096\n```\n\n**However to avoid complexity we can use the RobustModelConfig**","metadata":{}},{"cell_type":"markdown","source":"**Now we can modify [this script as needed](https://github.com/vasudevgupta7/gsoc-wav2vec2/blob/main/src/convert_torch_to_tf.py)**","metadata":{}},{"cell_type":"code","source":"from typing import Union\nimport tensorflow as tf\nimport transformers\nimport numpy as np\nfrom tqdm.auto import tqdm\nfrom wav2vec2 import Wav2Vec2Config, RobustWav2Vec2Config, Wav2Vec2ForCTC, Wav2Vec2Model\n\n\nSUFFIX = \":0\"\nMAPPING = (\n    (\"layer_norm.weight\", \"layer_norm/gamma\"),\n    (\"layer_norm.bias\", \"layer_norm.beta\"),\n    (\"weight\", \"kernel\"),\n    (\".\", \"/\"),\n)\n\n# fill-in PyTorch keys to ignore below\nKEYS_TO_IGNORE = []\n\nACCEPTABLE_HF_IDS = [\"facebook/wav2vec2-base-960h\", \n                     \"facebook/wav2vec2-base\", \n                     \"facebook/wav2vec2-large-robust\", \n                     \"facebook/wav2vec2-large-xlsr-53\",\n                     \"arijitx/wav2vec2-xls-r-300m-bengali\"]\n\nPREFIX_WITH_HEAD = \"wav2vec2-ctc/\"\nSPECIAL_MAPPING_WITH_HEAD = {\n    \"wav2vec2.encoder.pos_conv_embed.conv.weight_g\": f\"{PREFIX_WITH_HEAD}wav2vec2/encoder/pos_conv_embed/conv/weight_g:0\",\n    \"wav2vec2.encoder.pos_conv_embed.conv.weight_v\": f\"{PREFIX_WITH_HEAD}wav2vec2/encoder/pos_conv_embed/conv/weight_v:0\",\n}\n\nPREFIX_WITHOUT_HEAD = \"wav2vec2/\"\nSPECIAL_MAPPING_WITHOUT_HEAD = {\n    \"encoder.pos_conv_embed.conv.weight_g\": f\"{PREFIX_WITHOUT_HEAD}encoder/pos_conv_embed/conv/weight_g:0\",\n    \"encoder.pos_conv_embed.conv.weight_v\": f\"{PREFIX_WITHOUT_HEAD}encoder/pos_conv_embed/conv/weight_v:0\",\n}\n\n\ndef replace(k: str, prefix) -> str:\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 prefix + k + SUFFIX\n\n\ndef get_tf_pretrained_model(\n    config: Wav2Vec2Config,\n    hf_model_id: str,\n    verbose=False,\n    with_lm_head=True,\n) -> Union[Wav2Vec2ForCTC, Wav2Vec2Model]:\n    \"\"\"\n    Converts HuggingFace PyTorch weights to TensorFlow compatible weights.\n    Args:\n        config (:obj: `Wav2Vec2Config`):\n            Configuration of TF model.\n        hf_model_id (:obj: `str`):\n            model_id of HuggingFace PyTorch model.\n        with_lm_head (:obj: `bool`, default=True):\n            Whether to return Wav2Vec2ForCTC or Wav2Vec2Model\n    Returns:\n        Instance of `Wav2Vec2ForCTC` loaded with pre-trained weights.\n    \"\"\"\n    assert hf_model_id in ACCEPTABLE_HF_IDS, f\"{hf_model_id} is not acceptable\"\n\n    if with_lm_head:\n        tf_model = Wav2Vec2ForCTC(config)\n        prefix = PREFIX_WITH_HEAD\n        hf_model = transformers.Wav2Vec2ForCTC.from_pretrained(hf_model_id)\n    else:\n        tf_model = Wav2Vec2Model(config)\n        tf_model._init(input_shape=(1, 2048))\n        prefix = PREFIX_WITHOUT_HEAD\n        hf_model = transformers.Wav2Vec2Model.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, desc=\"hf -> tf\"):\n        if k in KEYS_TO_IGNORE:\n            continue\n\n        if k in SPECIAL_MAPPING_WITH_HEAD or k in SPECIAL_MAPPING_WITHOUT_HEAD:\n            new_k = (\n                SPECIAL_MAPPING_WITH_HEAD[k]\n                if with_lm_head\n                else SPECIAL_MAPPING_WITHOUT_HEAD[k]\n            )\n        else:\n            new_k = replace(k, prefix=prefix)\n\n        if new_k not in tf_variables_dict.keys():\n            extra_keys.append(k)\n            print(f\"SKIPPING {k}\")\n            continue\n\n        if verbose:\n            print(k, \"->\", new_k)\n\n        array = hf_state_dict[k].numpy()\n\n        # transpose the PyTorch weights for correct loading in TF-2\n        # Weights corresponding to `SPECIAL_MAPPING` are 3D array while other weights are 2D\n        # so we need to separate weights first & do special transpose on 3D weights\n        if k in SPECIAL_MAPPING_WITH_HEAD or k in SPECIAL_MAPPING_WITHOUT_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\n    return tf_model, hf_model\n\n\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"###########################\nis_robust= True \nwith_lm_head=True\nmodel_id =\"tf-wav2vec2-xls-r-300m-bengali\"\nhf_model_id=\"arijitx/wav2vec2-xls-r-300m-bengali\"\n###########################\nconfig = Wav2Vec2Config() if not is_robust else RobustWav2Vec2Config()\nconfig.vocab_size=112    \ntf_model, hf_model = get_tf_pretrained_model(config, hf_model_id, verbose=True, with_lm_head=with_lm_head)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Verify model ","metadata":{}},{"cell_type":"code","source":"tf_model.summary()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Test a random sample","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport librosa\n_path=\"../input/validation-fileswav-format/validation_files_wav/common_voice_bn_30620260.wav\"\ndef load_data(path):\n    \"\"\"loads a wav\"\"\"\n    wave,_= librosa.load(path, sr=16000, mono=True)\n    wave=np.trim_zeros(wave)\n    return wave\n\ndef _normalize(x):\n    \"\"\"You must call this before padding.\"\"\"\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\naudio=load_data(_path)\naudio=_normalize(audio)\nspeech = tf.constant(audio, dtype=tf.float32)\nspeech = tf.transpose(speech)\nspeech.shape","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tf_out = tf_model(tf.expand_dims(speech,axis=0), training=False)\ntf_out=tf.squeeze(tf.argmax(tf_out, axis=-1))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from IPython.display import display,Audio\ndisplay(Audio(data=speech, rate=16000))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### labels available here --> https://huggingface.co/arijitx/wav2vec2-xls-r-300m-bengali/blob/main/alphabet.json","metadata":{}},{"cell_type":"code","source":"LABELS=[\" \", \"_\", \"a\", \"b\", \"c\", \"d\", \"e\", \"f\", \"g\", \"h\", \"i\", \"j\", \"k\", \"l\", \"m\", \"n\", \"o\", \"p\", \"r\", \"s\", \"t\", \"u\", \"v\", \"w\", \n        \"x\", \"y\", \"z\", \"\", \"\", \"œ\", \"।\", \"ঁ\", \"ং\", \"ঃ\", \"অ\", \"আ\", \"ই\", \"ঈ\", \"উ\", \"ঊ\", \"ঋ\", \"এ\", \"ঐ\", \"ও\", \"ঔ\", \"ক\", \"খ\",\n        \"গ\", \"ঘ\", \"ঙ\", \"চ\", \"ছ\", \"জ\", \"ঝ\", \"ঞ\", \"ট\", \"ঠ\", \"ড\", \"ঢ\", \"ণ\", \"ত\", \"থ\", \"দ\", \"ধ\", \"ন\", \"প\", \"ফ\", \"ব\", \"ভ\", \"ম\", \"য\", \n        \"র\", \"ল\", \"শ\", \"ষ\", \"স\", \"হ\", \"়\", \"া\", \"ি\", \"ী\", \"ু\", \"ূ\", \"ৃ\", \"ে\", \"ৈ\", \"ো\", \"ৌ\", \"্\", \"ৎ\", \"ৗ\", \"ড়\", \"ঢ়\", \"য়\", \"০\",\n        \"১\", \"২\", \"৩\", \"৪\", \"৫\", \"৬\", \"৭\", \"৮\", \"৯\", \"ৰ\", \"‌\", \"‍\", \"‎\", \"⁇\",  \"\", \"<s>\", \"</s>\"]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"text=[LABELS[i] for i in tf_out]\ntext=\"\".join(text)\ntext","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"we can observe that our transfered model is the same as the original model (at least from inference)","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}