{ "cells": [ { "cell_type": "markdown", "metadata": { "id": "Bp8t2AI8i7uP" }, "source": [ "##### Copyright 2022 The TensorFlow Authors." ] }, { "cell_type": "code", "execution_count": 1, "metadata": { "cellView": "form", "execution": { "iopub.execute_input": "2023-12-14T12:09:17.840503Z", "iopub.status.busy": "2023-12-14T12:09:17.839835Z", "iopub.status.idle": "2023-12-14T12:09:17.843684Z", "shell.execute_reply": "2023-12-14T12:09:17.843133Z" }, "id": "rxPj2Lsni9O4" }, "outputs": [], "source": [ "#@title Licensed under the Apache License, Version 2.0 (the \"License\");\n", "# you may not use this file except in compliance with the License.\n", "# You may obtain a copy of the License at\n", "#\n", "# https://www.apache.org/licenses/LICENSE-2.0\n", "#\n", "# Unless required by applicable law or agreed to in writing, software\n", "# distributed under the License is distributed on an \"AS IS\" BASIS,\n", "# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n", "# See the License for the specific language governing permissions and\n", "# limitations under the License." ] }, { "cell_type": "markdown", "metadata": { "id": "6xS-9i5DrRvO" }, "source": [ "# Customizing a Transformer Encoder" ] }, { "cell_type": "markdown", "metadata": { "id": "Mwb9uw1cDXsa" }, "source": [ "\n", " \n", " \n", " \n", " \n", "
\n", " View on TensorFlow.org\n", " \n", " Run in Google Colab\n", " \n", " View source on GitHub\n", " \n", " Download notebook\n", "
" ] }, { "cell_type": "markdown", "metadata": { "id": "iLrcV4IyrcGX" }, "source": [ "## Learning objectives\n", "\n", "The [TensorFlow Models NLP library](https://github.com/tensorflow/models/tree/master/official/nlp/modeling) is a collection of tools for building and training modern high performance natural language models.\n", "\n", "The `tfm.nlp.networks.EncoderScaffold` is the core of this library, and lots of new network architectures are proposed to improve the encoder. In this Colab notebook, we will learn how to customize the encoder to employ new network architectures." ] }, { "cell_type": "markdown", "metadata": { "id": "YYxdyoWgsl8t" }, "source": [ "## Install and import" ] }, { "cell_type": "markdown", "metadata": { "id": "fEJSFutUsn_h" }, "source": [ "### Install the TensorFlow Model Garden pip package\n", "\n", "* `tf-models-official` is the stable Model Garden package. Note that it may not include the latest changes in the `tensorflow_models` github repo. To include latest changes, you may install `tf-models-nightly`,\n", "which is the nightly Model Garden package created daily automatically.\n", "* `pip` will install all models and dependencies automatically." ] }, { "cell_type": "code", "execution_count": 2, "metadata": { "execution": { "iopub.execute_input": "2023-12-14T12:09:17.847626Z", "iopub.status.busy": "2023-12-14T12:09:17.847188Z", "iopub.status.idle": "2023-12-14T12:09:22.397594Z", "shell.execute_reply": "2023-12-14T12:09:22.396446Z" }, "id": "mfHI5JyuJ1y9" }, "outputs": [], "source": [ "!pip install -q opencv-python" ] }, { "cell_type": "code", "execution_count": 3, "metadata": { "execution": { "iopub.execute_input": "2023-12-14T12:09:22.402397Z", "iopub.status.busy": "2023-12-14T12:09:22.401542Z", "iopub.status.idle": "2023-12-14T12:09:32.488117Z", "shell.execute_reply": "2023-12-14T12:09:32.486964Z" }, "id": "thsKZDjhswhR" }, "outputs": [], "source": [ "!pip install -q tf-models-official" ] }, { "cell_type": "markdown", "metadata": { "id": "hpf7JPCVsqtv" }, "source": [ "### Import Tensorflow and other libraries" ] }, { "cell_type": "code", "execution_count": 4, "metadata": { "execution": { "iopub.execute_input": "2023-12-14T12:09:32.492064Z", "iopub.status.busy": "2023-12-14T12:09:32.491769Z", "iopub.status.idle": "2023-12-14T12:09:37.185821Z", "shell.execute_reply": "2023-12-14T12:09:37.184726Z" }, "id": "my4dp-RMssQe" }, "outputs": [ { "name": "stderr", "output_type": "stream", "text": [ "2023-12-14 12:09:32.926415: E external/local_xla/xla/stream_executor/cuda/cuda_dnn.cc:9261] Unable to register cuDNN factory: Attempting to register factory for plugin cuDNN when one has already been registered\n", "2023-12-14 12:09:32.926462: E external/local_xla/xla/stream_executor/cuda/cuda_fft.cc:607] Unable to register cuFFT factory: Attempting to register factory for plugin cuFFT when one has already been registered\n", "2023-12-14 12:09:32.927992: E external/local_xla/xla/stream_executor/cuda/cuda_blas.cc:1515] Unable to register cuBLAS factory: Attempting to register factory for plugin cuBLAS when one has already been registered\n" ] } ], "source": [ "import numpy as np\n", "import tensorflow as tf\n", "\n", "import tensorflow_models as tfm\n", "nlp = tfm.nlp" ] }, { "cell_type": "markdown", "metadata": { "id": "vjDmVsFfs85n" }, "source": [ "## Canonical BERT encoder\n", "\n", "Before learning how to customize the encoder, let's firstly create a canonical BERT enoder and use it to instantiate a `bert_classifier.BertClassifier` for classification task." ] }, { "cell_type": "code", "execution_count": 5, "metadata": { "execution": { "iopub.execute_input": "2023-12-14T12:09:37.190371Z", "iopub.status.busy": "2023-12-14T12:09:37.189615Z", "iopub.status.idle": "2023-12-14T12:09:40.961070Z", "shell.execute_reply": "2023-12-14T12:09:40.960348Z" }, "id": "Oav8sbgstWc-" }, "outputs": [], "source": [ "cfg = {\n", " \"vocab_size\": 100,\n", " \"hidden_size\": 32,\n", " \"num_layers\": 3,\n", " \"num_attention_heads\": 4,\n", " \"intermediate_size\": 64,\n", " \"activation\": tfm.utils.activations.gelu,\n", " \"dropout_rate\": 0.1,\n", " \"attention_dropout_rate\": 0.1,\n", " \"max_sequence_length\": 16,\n", " \"type_vocab_size\": 2,\n", " \"initializer\": tf.keras.initializers.TruncatedNormal(stddev=0.02),\n", "}\n", "bert_encoder = nlp.networks.BertEncoder(**cfg)\n", "\n", "def build_classifier(bert_encoder):\n", " return nlp.models.BertClassifier(bert_encoder, num_classes=2)\n", "\n", "canonical_classifier_model = build_classifier(bert_encoder)" ] }, { "cell_type": "markdown", "metadata": { "id": "Qe2UWI6_tsHo" }, "source": [ "`canonical_classifier_model` can be trained using the training data. For details about how to train the model, please see the [Fine tuning bert](https://www.tensorflow.org/text/tutorials/fine_tune_bert) notebook. We skip the code that trains the model here.\n", "\n", "After training, we can apply the model to do prediction.\n" ] }, { "cell_type": "code", "execution_count": 6, "metadata": { "execution": { "iopub.execute_input": "2023-12-14T12:09:40.965393Z", "iopub.status.busy": "2023-12-14T12:09:40.964878Z", "iopub.status.idle": "2023-12-14T12:09:42.062738Z", "shell.execute_reply": "2023-12-14T12:09:42.062071Z" }, "id": "csED2d-Yt5h6" }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "tf.Tensor(\n", "[[ 0.03545166 0.30729884]\n", " [ 0.00677404 0.17251147]\n", " [-0.07276718 0.17345032]], shape=(3, 2), dtype=float32)\n" ] } ], "source": [ "def predict(model):\n", " batch_size = 3\n", " np.random.seed(0)\n", " word_ids = np.random.randint(\n", " cfg[\"vocab_size\"], size=(batch_size, cfg[\"max_sequence_length\"]))\n", " mask = np.random.randint(2, size=(batch_size, cfg[\"max_sequence_length\"]))\n", " type_ids = np.random.randint(\n", " cfg[\"type_vocab_size\"], size=(batch_size, cfg[\"max_sequence_length\"]))\n", " print(model([word_ids, mask, type_ids], training=False))\n", "\n", "predict(canonical_classifier_model)" ] }, { "cell_type": "markdown", "metadata": { "id": "PzKStEK9t_Pb" }, "source": [ "## Customize BERT encoder\n", "\n", "One BERT encoder consists of an embedding network and multiple transformer blocks, and each transformer block contains an attention layer and a feedforward layer." ] }, { "cell_type": "markdown", "metadata": { "id": "rmwQfhj6fmKz" }, "source": [ "We provide easy ways to customize each of those components via (1)\n", "[EncoderScaffold](https://github.com/tensorflow/models/blob/master/official/nlp/modeling/networks/encoder_scaffold.py) and (2) [TransformerScaffold](https://github.com/tensorflow/models/blob/master/official/nlp/modeling/layers/transformer_scaffold.py)." ] }, { "cell_type": "markdown", "metadata": { "id": "xsMgEVHAui11" }, "source": [ "### Use EncoderScaffold\n", "\n", "`networks.EncoderScaffold` allows users to provide a custom embedding subnetwork\n", " (which will replace the standard embedding logic) and/or a custom hidden layer class (which will replace the `Transformer` instantiation in the encoder)." ] }, { "cell_type": "markdown", "metadata": { "id": "-JBabpa2AOz8" }, "source": [ "#### Without Customization\n", "\n", "Without any customization, `networks.EncoderScaffold` behaves the same the canonical `networks.BertEncoder`.\n", "\n", "As shown in the following example, `networks.EncoderScaffold` can load `networks.BertEncoder`'s weights and output the same values:" ] }, { "cell_type": "code", "execution_count": 7, "metadata": { "execution": { "iopub.execute_input": "2023-12-14T12:09:42.066614Z", "iopub.status.busy": "2023-12-14T12:09:42.066361Z", "iopub.status.idle": "2023-12-14T12:09:42.784432Z", "shell.execute_reply": "2023-12-14T12:09:42.783667Z" }, "id": "ktNzKuVByZQf" }, "outputs": [ { "name": "stderr", "output_type": "stream", "text": [ "WARNING:absl:The `Transformer` layer is deprecated. Please directly use `TransformerEncoderBlock`.\n" ] }, { "name": "stderr", "output_type": "stream", "text": [ "WARNING:absl:The `Transformer` layer is deprecated. Please directly use `TransformerEncoderBlock`.\n" ] }, { "name": "stderr", "output_type": "stream", "text": [ "WARNING:absl:The `Transformer` layer is deprecated. Please directly use `TransformerEncoderBlock`.\n" ] }, { "name": "stdout", "output_type": "stream", "text": [ "tf.Tensor(\n", "[[ 0.03545166 0.30729884]\n", " [ 0.00677404 0.17251147]\n", " [-0.07276718 0.17345032]], shape=(3, 2), dtype=float32)\n" ] } ], "source": [ "default_hidden_cfg = dict(\n", " num_attention_heads=cfg[\"num_attention_heads\"],\n", " intermediate_size=cfg[\"intermediate_size\"],\n", " intermediate_activation=cfg[\"activation\"],\n", " dropout_rate=cfg[\"dropout_rate\"],\n", " attention_dropout_rate=cfg[\"attention_dropout_rate\"],\n", " kernel_initializer=cfg[\"initializer\"],\n", ")\n", "default_embedding_cfg = dict(\n", " vocab_size=cfg[\"vocab_size\"],\n", " type_vocab_size=cfg[\"type_vocab_size\"],\n", " hidden_size=cfg[\"hidden_size\"],\n", " initializer=cfg[\"initializer\"],\n", " dropout_rate=cfg[\"dropout_rate\"],\n", " max_seq_length=cfg[\"max_sequence_length\"]\n", ")\n", "default_kwargs = dict(\n", " hidden_cfg=default_hidden_cfg,\n", " embedding_cfg=default_embedding_cfg,\n", " num_hidden_instances=cfg[\"num_layers\"],\n", " pooled_output_dim=cfg[\"hidden_size\"],\n", " return_all_layer_outputs=True,\n", " pooler_layer_initializer=cfg[\"initializer\"],\n", ")\n", "\n", "encoder_scaffold = nlp.networks.EncoderScaffold(**default_kwargs)\n", "classifier_model_from_encoder_scaffold = build_classifier(encoder_scaffold)\n", "classifier_model_from_encoder_scaffold.set_weights(\n", " canonical_classifier_model.get_weights())\n", "predict(classifier_model_from_encoder_scaffold)" ] }, { "cell_type": "markdown", "metadata": { "id": "sMaUmLyIuwcs" }, "source": [ "#### Customize Embedding\n", "\n", "Next, we show how to use a customized embedding network.\n", "\n", "We first build an embedding network that would replace the default network. This one will have 2 inputs (`mask` and `word_ids`) instead of 3, and won't use positional embeddings." ] }, { "cell_type": "code", "execution_count": 8, "metadata": { "execution": { "iopub.execute_input": "2023-12-14T12:09:42.788065Z", "iopub.status.busy": "2023-12-14T12:09:42.787462Z", "iopub.status.idle": "2023-12-14T12:09:42.817176Z", "shell.execute_reply": "2023-12-14T12:09:42.816557Z" }, "id": "LTinnaG6vcsw" }, "outputs": [ { "name": "stderr", "output_type": "stream", "text": [ "/tmpfs/src/tf_docs_env/lib/python3.9/site-packages/keras/src/initializers/initializers.py:120: UserWarning: The initializer TruncatedNormal is unseeded and being called multiple times, which will return identical values each time (even if the initializer is unseeded). Please update your code to provide a seed to the initializer, or avoid using the same initializer instance more than once.\n", " warnings.warn(\n" ] } ], "source": [ "word_ids = tf.keras.layers.Input(\n", " shape=(cfg['max_sequence_length'],), dtype=tf.int32, name=\"input_word_ids\")\n", "mask = tf.keras.layers.Input(\n", " shape=(cfg['max_sequence_length'],), dtype=tf.int32, name=\"input_mask\")\n", "embedding_layer = nlp.layers.OnDeviceEmbedding(\n", " vocab_size=cfg['vocab_size'],\n", " embedding_width=cfg['hidden_size'],\n", " initializer=cfg[\"initializer\"],\n", " name=\"word_embeddings\")\n", "word_embeddings = embedding_layer(word_ids)\n", "attention_mask = nlp.layers.SelfAttentionMask()([word_embeddings, mask])\n", "new_embedding_network = tf.keras.Model([word_ids, mask],\n", " [word_embeddings, attention_mask])" ] }, { "cell_type": "markdown", "metadata": { "id": "HN7_yu-6O3qI" }, "source": [ "Inspecting `new_embedding_network`, we can see it takes two inputs:\n", "`input_word_ids` and `input_mask`." ] }, { "cell_type": "code", "execution_count": 9, "metadata": { "execution": { "iopub.execute_input": "2023-12-14T12:09:42.820481Z", "iopub.status.busy": "2023-12-14T12:09:42.819979Z", "iopub.status.idle": "2023-12-14T12:09:42.937604Z", "shell.execute_reply": "2023-12-14T12:09:42.936850Z" }, "id": "fO9zKFE4OpHp" }, "outputs": [ { "data": { "image/png": "iVBORw0KGgoAAAANSUhEUgAAAbAAAACTCAIAAABgY9tRAAAABmJLR0QA/wD/AP+gvaeTAAAgAElEQVR4nO2de1hTR/rH51AFFgJRsoJ4A1EEAxYWKVZLV1tgwZUWlIuCFrXwINYbFesFrMBTLy1g0dpqeeqltq5VwIJVKq1s6VZW4NGuBBBXpIoFgQjhYhS5JfP7Y9jzS5MQck8I7+evcyZzeWfmzZs558z5hsIYIwAAAAAhI10bAAAAoC+MUb2KnJwc1SsBAML8+fOnTJmiayuAUYoaVoiZmZmqVzIS0ZOOl5aWlpaW6toK9VBWVmYwfQFGImpYIU6ZMiUsLEz1ekYcOTk5+tNx/bEEAEYuo+georOzc0BAgHpz6idlZWUURbm7u5PT1tbWgwcP6tYkhNDt27fPnDlDjvPz8ymK8vX11a1JACCGRgJid3d3eHi4JmpWhZKSEqnpktYOlVN+dD4C0dHRFRUV5Dg+Pn7t2rXx8fEURaWkpJBEV1dXiqL27t2rrhaLi4u9vLzy8/PplNLSUjc3N1tb27y8PISQi4tLXV1deXk5Qig4OLi5uVldTQOAutBIQDQzM8vOzla6eE9Pz8svv6xGe2SjorWaqFONI8Dlctva2qysrA4dOhQVFXXgwIGrV68ihKqrq3fs2LF79261tIIQoijKy8uLPu3r64uIiMjIyKisrDx+/DhJDAgIOHXqlLpaBAC1o5GAGBoaOnHiRPogLi7O0tLy9OnTCCFfX19nZ+ewsDAmk5mamooQ8vb2Jl9+d3f3wMBAhFBgYGB5eTlFUaLLDYTQzJkzKYoqKirKyspKT09vbGykKGrjxo0CgSAmJobJZLq7u1dVVZF2bWxsIiMjTUxMvv3222XLlllYWBw5ckS2tX19fWI57927x2azjYyMSAadj4AS1NXVWVtbk2NbW9v9+/evXLmyqamJziB19MRsFggE0dHRlpaWCxYsaGxslNrQokWLRE+vX7/u4uLi5+c3YcKEgoIC2oC7d++q2CMA0BwaCYi5ubkkHOTm5pqamsbGxpaXl2dlZSGEioqKWltb9+zZU1NTc+rUqd9++62oqIiUunz5Mn0wb948jHFwcLBotZcvXw4NDfX19f3pp58uXbo0ZcqU1NTUI0eOZGdnc7ncxsbG3bt3v/POO6RdMzOz5OTk3t7ep0+fCoXC5uZme3t72daeO3dOLGdBQcGrr77a3t7e0tKiqxEQCoWenp59fX0KGUCDMaYoij5NSEhYtGjR8uXLBwYGSIrU0ROzOScnp6Oj4+HDhwkJCSSID8ujR4+MjY1nz57NYrEOHz4s1RgA0Dc0/lCFyWR6eHjMnj27u7ubpEyaNGnOnDmTJ0/29PSsq6ujc9LvzAz1nXF2dm5oaHj8+LGDg0Nra2tLS4uxsTFFUXfu3AkICLCwsAgODq6pqaHbdXJyQgjdu3fPz8+PwWAsXrxYtqmSOSMjI9va2qZNm7Zr1y5djYCRkdHNmzeNjY2Va93R0ZHL5YqmnDx5ksfjJSYmktOhRk/U5pqamry8PCsrq9DQ0Bs3bsjTLoPBuH//fklJya1bt9LS0jo6OhBCXC7X0dFRuY4AgBbQeECUjG7Nzc0cDqepqenmzZuOjo4mJiaPHz/u6uoqLi4mGYyNjfl8fnl5ueQt/6CgoM2bNwcFBYWEhMTHx5NnwWw2u7CwkM/n5+fns9lssSKzZs26evXq06dPf/zxR9mmSua0tra+cOECh8PJzc199OiRPoyAotjY2LBYrPb2djqFwWBcuHAhKyuLz+ejIUZPzGY2mx0TE9Pd3Y0xpp/VyMbLy8vIyEgoFBoZGVEURSq8cuXKmjVrVOwRAGgQrDJhYWFiKSEhIQihkJCQ1atXI4Sio6ODgoIQQgkJCRhjNze3iIgIS0vLlJQUkn/jxo3m5uZJSUkIoQMHDmCMIyIiWCwWh8MRq/n+/fvTp0/HGFdVVTk5OZHEgYGB6OhoCwsLNze3yspKjDFpl2To7e1dunSpubk5WeUlJycPZa1kzuTkZISQqalpVFSUUCgctuOaGAGBQODp6dnb2ztUW9nZ2dnZ2aIpZG+zm5sbOeVyuRkZGVu2bEEILVy4kCTm5OTExcXJGD1Rm0keBoMxceLEEydOCAQCV1dXMTP8/f2JR6Wnp5OUY8eOWVtb29jYHDt2DGNcXV391VdfkY/Ic2cfH59h+wIA2oTCKos7hIeHK/RE1d3dXc5Vhp6jaMdp1DsC5NVJbW7MPnPmDIPBELvDqxa03xcAEEXbG7N9fX05HE5MTIyc+SkR6D10KqKJOuVH0RHQQ1atWqWJaAgAOkcNr+4pBP1EVU5UX8Bqp075UXQEAADQGqPo1T0AAADZQEAEAAAYRA0PVdzd3WfNmqUWa0YWFRUVtICCDiGvjhiGhmBjY+O7774LD1UAXaGGe4izZs1S+7vAIwKlnzKrF0N6Mgtiw4BugUtmAwTkvwBAOXQQEFNSUsiWl87OTv2pysAA+S8AUALdBMSQkJDS0tJx48YpUVxUGkvFqnSCQtJequuAgfwXAMiPLi+ZNSSNtW/fPgsLi7CwsP7+/p6eHrKErK+vDw4OpiiKx+OJKVmJaoXRejPKISmlJdt+sf5qQgcM5L8AQH50GRDVKw5Gk5SUxOfzFy5ceOXKFVNT04KCgvXr19vb2ycmJubl5V29elVMyUpUK4yEIaWRlNKSbb9Yf8kgDJVZOR0wDPJfACA3On6ookZxMEJNTc2cOXPGjBmzadOm1tZWhNDixYuvX7/e1dV18eLFN954Q6qSFa0VpiJSpbRk2y+1v1IzK6cDBvJfACA/Og6I6pXG2rZtW3Z2tp+fX2dn5/bt20lYoSgqKipq37591tbWL7zwghJKVvIjKaU1rP2i/XV1dVW7DhjIfwGAAqgumCNDBUsqRFMLIUTUpVSRxqKrIrzyyiscDsfBwWHq1KmRkZEIoYaGBoxxR0fHhAkTeDwe/p/aFa1khf+oFaZixyWltGTbL9lfJXTAQP4LANSFDuS/ZKMJcbCenp6PPvpILHqqjlo6rnp/Qf4LANSFfm3M1oQ0VmBgoIODQ0REhBrrVBcjVAoM5L8AQ0Xb8l+y0YQ0loo7aTQKSIEBgF6hXytEAAAAHQIBEQAAYBA1PFSZOnXq/Pnz1WLNyKK2tlZ+3bOuri4mk6kJM3g8HkKIxWJponItA/JfgG5RQ0AE5GH58uXnz5/XtRUAAMgCLpkBAAAGgYAIAAAwCAREAACAQSAgAgAADAIBEQAAYBAIiAAAAINAQAQAABgEAiIAAMAgEBABAAAGgYAIAAAwCAREAACAQSAgAgAADAIBEQAAYBAIiAAAAINAQAQAABgEAiIAAMAgIBCrWXp7e19//fXe3t729nYrKyuBQJCZmblo0SJd2wUAgBT061/3DA8TExOBQPDrr78ihB48eDB16tQXX3xR10YBACAduGTWOOvWrTM1NSXHdnZ2VlZWurUHAIChgICoccLDw62trRFCpqam69at07U5AAAMCQREjWNubj5jxgyEkI2NzbJly3RtDgAAQwIBURts2rTJzMzMycnJzMxM17YAADAkEBC1wZIlSyiK2rhxo64NAQBAFoPbbsLDw3VtieEg9Z/jq6qqXFxcjIy0+gtUW1s7a9YsbbaoQ0pLSxsaGmTnMQA/l+pdOsHAvCs7OxuJbrsh54Dq5OTkIITCwsJEE3t7e01MTLRsSXh4+OiZVjmD3UgfEKnepRMMybto59HqgsXZ2TkgIECVGnx9fZ2dnZVuTnUDlEb70VAJWltbDx48qGsr0O3bt8+cOaNrK3SGQk6uV5SVlVEU5e7ujvTGl5CIO+Xn51MU5evrKyOzVgNiSUmJijUUFRXRe/qUaI5O6e7u1snVk67alZP4+Pi1a9fGx8dTFJWSkkISXV1dKYrau3evulopLi728vLKz8+nU0pLS93c3GxtbfPy8hBCLi4udXV15eXl6mpRBno4Iwo5uSj60Jfo6OiKigr0P18iB3riTsHBwc3NzbJrG6UPVczMzHSy2lex3Z6enpdfflmN9ojC5XLb2tqsrKwOHToUFRV14MCBq1evIoSqq6t37Nixe/dudTVEUZSXlxd92tfXFxERkZGRUVlZefz4cZIYEBBw6tQpdbUoA32eEUXRn77QvoQQGlnupFhAnDlzJkVRRUVFWVlZ6enpjY2N5OGpQCCIiYlhMpnu7u5VVVUIodDQUBsbm8jISBMTk2+//XbZsmUWFhZHjhyRWq1AIIiOjra0tFywYEFjY2NoaOjEiRPj4uKsrKy+//57b2/vSZMmVVdX0/kjIyOZTGZqaqpkWTIiYs1JppAm6IO4uDhLS8vTp0/TmRkMRmJiIvllu3fvHpvNNjIyIkVUQUa75EIpLCyM7pq3tzdxUHd398DAQIRQYGBgeXk5RVGiv4fqoq6ujmwgRwjZ2tru379/5cqVTU1NdAapsyzWC8npkETsVe7r16+7uLj4+flNmDChoKCANuDu3btq76MkmpsR2W68b98+CwuLsLCw/v7+oRyMfN1+/vlnnfdFUUR9CcnnTsr5ElK7O2GMMcZhYWFYDu7cuRMaGooxDg8Pf/XVVzHGqampQqHw7NmzgYGBT548ycnJ8fb2Jpnt7e3/+9//YoxPnz4dFBTE5/NPnjzp7+8vWe0333yzdOnS9vb23NzcmJgYjLGdnV1FRcW1a9dmzpzZ1NT02Wefvf/++yQzi8X6z3/+09jYaGdnV1dXJ1lWsjmpBri5uZEDOzu7X3/9taamZv78+aKZT58+TTJnZmbGxsZ2dHTIM0QY4+zs7Ozs7KE+HapdjLGVlVVlZSXdtefPn8+bNw9j3NDQsGTJEowxnYIxFggEc+fO7e3tlWGJnNNKuHbt2ltvvUWOd+zYQYp7e3v39/eTU6mzLNYLyemQyoYNG/Ly8sjxmTNngoODnZ2dyeKUJD548OC1116T33g5Oys1j7pmRBIZbkw4cuTIxYsXJR3Mzc2tsbExLS2tra1NrE599q7S0tLo6Gj8R1/CcruTcr6EFXGn5uZmHx8fGX1RbIXo7Ozc0NDw+PFjBweH1tbWlpYWY2NjiqLu3LkTEBBgYWERHBxcU1NDMjOZTCcnJ4TQvXv3/Pz8GAzG4sWLpVZbU1OTl5dnZWUVGhp648YNUtbNzc3Z2XnatGm2trZsNvvJkyd0sP/LX/4yefJkT0/Puro6ybKSzck2gMlkenh4zJ49u7u7WzSzv78/yRAZGdnW1jZt2rRdu3YpNFyyEWsXITRp0qQ5c+bQXaNz4v8pElEURScaGRndvHnT2NhYXfY4OjpyuVzRlJMnT/J4vMTERHI61CyL9kJyOoaFwWDcv3+/pKTk1q1baWlpHR0dCCEul+vo6KiursmJijMitUKpblxTUzNnzpwxY8Zs2rSptbVV0sF4PF5ycnJvb6/S22t0612SvoTkcCfVfQmp7E4K30MMCgravHlzUFBQSEhIfHw8eWjLZrMLCwv5fH5+fj6bzRYrMmvWrKtXrz59+vTHH3+UWiebzY6Jienu7sYYkzuy9NxIOlxTU1NVVVVTU9PNmzcdHR0ly0o2J9sAsSbozIWFhSTF2tr6woULHA4nNzf30aNHioyWLCS71tzczOFw6K6ZmJg8fvy4q6uruLiYZDA2Nubz+eXl5Wq8J01jY2PDYrHa29vpFAaDceHChaysLD6fj4aYZbFeSE7HsHh5eRkZGQmFQiMjI4qiSIVXrlxZs2aNuromJ2qfkaHcODs728/Pr7Ozc/v27RhjSQdjsVjHjx9//PgxfRdM531RCElfQnK4k+q+hFR3p6FWv0Nx//796dOnY4yrqqqcnJxI4sDAQHR0tIWFhZubW2VlJcZ49erVCCGSobe3d+nSpebm5uQHMDk5WaxOUpzBYEycOPHEiROkbHR09Lx58xBCmZmZ5ubmpFc+Pj5OTk6hoaEWFhYpKSmSZaU2J5kSEhKCEAoJCaHbCgoKQgglJCTQmRMTE8mlRHJyMkLI1NQ0KipKKBQOO0QyLmpktIsxdnNzi4iIsLS0JF3DGG/cuNHc3DwpKQkhdODAAYxxREQEi8XicDgCgcDT01ONl8wYYy6Xm5GRsWXLFoTQwoULSWJOTk5cXBweepZFeyE2HQKBwNXVVawVeumdnp5OUo4dO2ZtbW1jY3Ps2DGMcXV19VdffaWQ5XJ2VjKPGmdErGYZbszhcBwcHKZOnRoZGUnyiDoYMWDDhg2kYFJSkmi1+uxdpaWlCCFy2U58CWMspzsN60sYYxXdiTx0ln3JrHBAHCUMDAx8/vnnW7duVaKs7Ls8MqBvAKkLnU/r119/Td/c0TRK30OUgdpnRHVGs3dpzp2UvIeoLigR6A1KegLGODEx8c9//vPXX3+9bds2rbXr6+vL4XBiYmK01qIWWLVqVXBwsK6tUBJFZ0SfvdowvEsL7qQbxWysx/9bQFHU/v379+/fr+V2i4qKtNwiIBtFZ0SfvRq8S05G6cZsAAAASSAgAgAADDIo/+Xj46MPgkKGAdlVP2XKFF0bgioqKsib9qOB2traYTdnGICfg3dpAh6P989//hPR9xBZLJbBKPnoHBBo0gny6BoYgJ+Dd2kC3ch/AQAA6DMGGBDll5PTK8FEraE/QnVKMMqlEvUcA9dDvHv3rr+/P5PJdHZ2/vLLL4fKlpKSQvZejRs3buHChf/+97/lsVJ+7Ta6fsKwCkXyy8npoWCiFtAT0UNJMMZpaWnjx4+fMWMGeZ/s2bNnq1evNjMz8/DwqKysRJqRSqQdrLOzU99qG3EYrB7is2fPAgMDV69e/ejRo9zc3MzMzEuXLknNmZKSEhISUlpaWl9fHxsb++abb8qjtCO/dhtdP9lK7u3tLU8pFdG+YKKianTKqdfpleihGA8fPrx7925TU9OHH35IvjCXLl3y8PBob29fsWIFEa1CGpBKpB1s3LhxytUgOheq16YJtONdNAaoh1hQUPDiiy9GRkYyGAxXV9e9e/cePXoUSZNaoxk3btzKlSvXr19PXmIV1TLr6ekhP5v19fXBwcFk4Uq024gs2rhx4/72t79xudxhRdAyMjLQcGJzSKZmotoFE4cdZUklQdlqdJpTr9Mr0UMx7O3tT5w4MXbsWD6fT256rFixYsuWLaampgsWLKCfq2pUKlFDSoK0+iGfzxf7InR2dkrqgdJaopcvXx7WZv3xLpqRq4c4ZECsr6+fPn06fTp9+vTff/8dIZSbm2tqahobG1teXp6VlSVZkM1mP3z4MCcnp6Oj4+HDhwkJCampqaampgUFBevXr7e3t09MTMzLyysqKiIRJzs7m8vlPnz4cP369adOnRIrSFc7f/58iqJoBU1ixvr167/77rstW7bk5OTs3r2bXtM1Nja+9957NTU1p06d+u2338TqPHfunFAobG5utre3J/klU3Jzc0lbkv0lmVtaWpydnf39/eUJiKSPjY2Nu3fvfuedd5DImwO0x1++fJmo0QUHBxcVFbW2tu7Zs4fuguz8QqHQ09Ozr69vWEswxqKaIgkJCYsWLVq+fPnAwMBQpkqOwFBzJINHjx4ZGxvPnj2bxWIdPnxYRk5nZ+e9e/fGx8fTKQMDA+fOnduzZ4/ULqgXyc4qOhdSq01KSuLz+QsXLiwuLhb7IowbN05sPHNzc83MzIgoCQlPstEf76KRnKNhPU0tboaG8DT5fWbIgGhnZ1dfX0+fPnjwgI6PklJroty+fdvOzk5Sy2zx4sVlZWVdXV3ffffdm2++SecnsmhMJnPp0qU7d+4cSgSNXDK3tLTQKUprJqpdMHFYpCoJEvAQynoaUq/TK9FDqdTW1p4/f54OBBjjnTt3bt26ld4/qGmpRLWrIoqpHy5evPj69etdXV0XL15844030BB6oERLVB70x7toDFAPccmSJRUVFefPn3/69Ont27f37NmzYcMG8tFQ09/Z2Xn27NnPP/88OjpaqsThmjVr9u3bZ2NjI/r3xEQW7dmzZ/SpbBG0nTt3ipmhqGai2gUTh0VS+m1YNToNqdfpleihGPn5+enp6f39/WPHjuVyuWQpkZaWFhMT4+DgQN910rRUonqVBLdt2yamfkhRVFRU1L59+6ytrV944QWkrPAfjf54F41h6iHeuXPHx8eHwWA4OjqeOnWKJErKlhG5QIQQk8n861//WlJSgqXJFGKMnzx5Ymtr297ejkW02+icU6ZM+fHHHyUL0vUTiMaiKpqJahdMFENSoElSSRAPp6ynFvU6qdOqJ6KHkqW6u7tjY2OZTOaUKVO++OILjLHoczwy7zKkEpWW/6IdrKOjQ3UlQTF3feWVV8TUDxsaGjo6OiZMmMDj8UhtUvVAabFRMfTZu0APcfQiQzBRacU6UdSiXqedaVVOpU7t2naa0EMkqF1J8Pnz53QkUpTR5l2iGKwe4ogGa14wcWSp1ymnUjdSpBLVPheBgYEODg4RERHqqlBRRpZ3iWKweogjGi0IJoJ6nf6g9rmQZyeNRgHvkgGsEAEAAAb5/xUiUdEAVIfcWtYaz58/NzU1lfrctrGxcfRMq4y9u6KMrAF5/vz5n/70J9EULXuXDAzSuwb1EMvKyhoaGnRtDKAMt27dKiws9PHxeemllzS3Y3lEMKwo1gjycx6P99133z19+nTjxo2jfFq1wNSpU8m7OoMBERjR9PX1ffnll3l5eREREStXriS724ARyu+//56Zmfno0aOdO3d6eHjo2pzRBQREw6Gvr+/cuXO5ubmhoaEQFkci9fX1hw8fbmtrS0xMnD17tq7NGY1AQDQ0SFj8xz/+ERIS8vbbb48ZAxsJRgAPHjxIS0vr6upKTk6W/6U9QO1AQDRM+vv7v/nmGwiL+k9NTc3HH38sFAp37dql0Xe0AXmAgGjIkLB49uzZZcuWrV27duzYsbq2CPh/qqurDx06JBAIkpKSZs6cqWtzAIQgII4GhELhhQsXTpw4sXjx4ri4OBMTE11bNNqprKw8cuTImDFjkpKS9OH/8wAaCIijBdGwuG7dOjn/ZQFQLxUVFYcPH2Yyme+9997kyZN1bQ4gDgTE0YVQKCwoKDh+/Pjrr78eGxsrtukX0BzXr18/cuSIjY3Njh07bG1tdW0OIB0IiKMRjPHly5ePHTu2cOHCzZs3Q1jUKCUlJUePHrW2tt61a5eNjY2uzQFkAQFx9ELC4ueff/7Xv/5106ZNZmZmurbI0CgpKfn4449nzZq1Y8eO8ePH69ocYHggIAKoqKgoMzPzpZde2rp1q6Wlpa7NMQRKSkoOHz7s4uISHx+vV3+/B8gGAiIwSFFR0aFDhzw9Pd99910mk6lrc0YqMIwjGgiIwB+ApY1ywP0HwwACIiCFkpKSgwcPOjk5wc2vYSEP7uEJlWEAAREYEngmIBvY2ml4QEAEhoHsGpk2bdr27dutrKx0bY5eAC//GCoQEAG5oPcV79y5c+LEibo2R2eIvh4OqhmGBwREQAFKS0s/+eST0fm6hajcZGRkJIRCgwQCIqAwHA7n0KFDTCZz+/btkyZNQgjV1tYmJSWdP3/eyMhA/rbsgw8+mDVr1vLlyxEIko8mICACSlJZWZmenm5ubv7+++/Hxsb+9NNPf//733Nzcw3gD0A++uij9PR0Gxub8vLyEydOFBYWRkRErFq1ymDCPTAUEBABlaioqEhNTS0rK2tpaTE3N/f39x/pMfHQoUN79+7l8XiWlpbe3t7r169fsmTJiO4RID/wiweohLu7u1AobGlpQQg9e/assLBw5cqVI/dX9pNPPtm3bx+Px0MIPXnypL6+HqLhqAICIqASDx48+Ne//mVpaUmiRnd39+XLl99++21d26UMn3766QcffNDW1kZOzczMmpqafvjhB91aBWiTUXTJHB4ermsTDBCMcU9PT1NT0wsvvNDV1cXn83t6ep4/f+7g4ODi4qJlY3g8HovFUq5sQ0NDRUWFqampsbExg8GwtLS0sLAwNzc3MzODpyiaIDs7W9cmSGF0bR3QzzkwAMLDwyXHFmOs5YtNqWbIifatHc3o7eoELpkBTTGy4svIshbQEBAQh8TX19fZ2ZkcHz161Nzc3N7eXqcWaRDRzhoSZWVlFEW5u7uT09bW1oMHD+rWJKW5ffv2mTNnZGTQk96J2pmfn09RlK+vr25Nkh8IiENSVFREv67/xRdfNDY21tfXa6Hd7u5u7V9QiHZWkmfPnq1evdrMzMzDw6OyslKhmnXSHVGio6MrKirIcXx8/Nq1a+Pj4ymKSklJIYmurq4URe3du1ddLRYXF3t5eeXn59MppaWlbm5utra2eXl5UotgjNPS0saPHz9jxozi4mIkbcxdXFzq6urKy8uHaldPeidqZ3BwcHNzs7qa1gIQEOUCYyyP3EtPT8/LL78seaxQQTMzM32713np0iUPD4/29vYVK1akpqYqVFbF7ig0jLLhcrltbW1WVlaHDh2Kioo6cODA1atXEULV1dU7duzYvXu3WlpBCFEU5eXlRZ/29fVFRERkZGRUVlYeP35capGHDx/evXu3qanpww8/JMFL6pgHBAScOnVK/3snw049Z1QHxHv37rHZbCMjI6JWIBAIoqOjLS0tFyxY0NjYSGcLDQ3lcDgURZWVlYnVsG/fPgsLi7CwsP7+foRQYGBgeXk5RVH5+fmix2I1h4aGTpw4MS4uztLS8vTp02IFyafEnpiYGCaT6e7uXlVVRSwRKygK/amVldX333/v7e09adKk6upqqaaK9Z1m5syZFEX9/PPPookrVqzYsmWLqanpggULFP0fYbo7ksaT6/SwsDAmk0m+897e3iT8ubu7BwYGio2MQu1KUldXZ21tTY5tbW3379+/cuXKpqYmOoM8Az6Uk4iyaNEi0dPr16+7uLj4+flNmDChoKBAahF7e/sTJ06MHTuWz+eTexdSx9zW1vbu3bv63zsZduo7eNQQFhYmlpKZmRkbG9vR0UFOv/nmm6VLl7a3t+fm5sbExGCM3dzcyEf0gVSOHDly8eJFjGPNCbcAAAhCSURBVPHz58/nzZtHEkWPJWu2s7P79ddfa2pq5s+fL5aZbu7s2bOBgYFPnjzJycnx9vYmH4kVFMPOzq6iouLatWszZ85samr67LPP3n//fammivWdNNrY2JiWltbW1ia1m/39/Rs2bJD6qeTYikKPnqTxVlZWlZWVjY2NdnZ2dXV19Dg0NDQsWbJEbGQEAsHcuXN7e3uHakjSjNLS0ujoaHJ87dq1t956ixzv2LGD5Pf29u7v7yen8gy45FRKZcOGDXl5eeT4zJkzwcHBzs7OZPkmY6BmzJgxffr02tpaOkVszB88ePDaa69JLatXvRO1s7m52cfHR6wG2Q6jQ0b1CjEyMrKtrW3atGm7du1CCNXU1OTl5VlZWYWGht64cWPY4jU1NXPmzBkzZsymTZtaW1vRH59Uih5L1sxkMj08PGbPnt3d3Y2GeMR5586dgIAACwuL4ODgmpoakihWUAwmk+nm5ubs7Dxt2jRbW1s2m/3kyROppor1HSHE4/GSk5N7e3ulbuXDGO/cuXPr1q1Kb/STavykSZPmzJkzefJkT0/Puro60ebIgejIGBkZ3bx509jYWLnWHR0duVyuaMrJkyd5PF5iYiI5lWfAFXUShBCDwbh//35JScmtW7fS0tI6OjqGyllbW3v+/HmyNEbSxpzL5To6Oup/72TYqeeM6oBobW194cIFDoeTm5v76NEjNpsdExPT3d2NMaZvw8sgOzvbz8+vs7Nz+/bt5AtsbGzM5/PLy8v37t0reixZs1gEFM1MJ7LZ7MLCQj6fn5+fz2azSaLs3SH0p2LZJE0V6ztCiMViHT9+/PHjx1Lvc6WlpcXExDg4OKhyN0rS+ObmZg6H09TUdPPmTUdHRxMTk8ePH3d1dZEHC2iIkVEOGxsbFovV3t5OpzAYjAsXLmRlZfH5fCTfgCvqJAghLy8vIyMjoVBoZGREUZTUGczPz09PT+/v7x87diyXyx0YGEDSxvzKlStr1qzR/97JsFPf0dnaVOtIrtKTk5MRQqamplFRUUKhcGBgIDo6msFgTJw48cSJEz4+Pgih6OjokJAQhBCTyRQrzuFwHBwcpk6dGhkZiRBqaGjAGEdERLBYLA6HI3osVvPq1atJzUFBQQihhIQE0cykuZCQEFLKwsLCzc2tsrISYyy1IA396bx58xBCmZmZ5ubmCKErV65ImirWd1Lhhg0bSNmkpCTRmi9dukQ7jJOTkzxjS0N3R6rxbm5uERERlpaWKSkpJP/GjRvNzc2TkpIQQgcOHBAdGYFA4OnpqeglM0KIvmbncrkZGRlbtmxBCC1cuJAk5uTkxMXFYYzlGXCxqRQIBK6urmKN+vv7k7FKT08nKceOHbO2traxsTl27BjGWLJUd3d3bGwsk8mcMmXKF198IXXMq6urv/rqK5Jfart60jtRO8lz5xF0yTyqAyKgLpQeW9k3Z7VmhtJ8/fXX9N00TZdSbw26bUVvv4wQEBVDdHGdnJyseoX6aY+i1So3tvQaXImyajQD0D56O1Oj611m1cF6poWhIXu0082ioiIttAIA8jOqH6oAAACIMopWiDweLycnR9dWGCZNTU36MLZ6YgYwLESCVw+BFSIAAMAgo2iFyGKxwsLCdG2FYZKTk6MPY6snZgDDorcLeVghKolBioMZqgiY5tATxS2lGVZSbLQBAVE627Zts7KymjRp0qeffioUCiUzyBAHO3/+PNnerC4UUtDatm0bRVHr1q2jU3744QeKory9vYctK1sEDJCEKG6RA32TFJOn1LCSYqMNCIhSqKys/OWXX+7evVtbW1tVVUW/+DkU+I/iYJcuXaqvryeCItoXBMvIyFizZs3Zs2dptZKjR49Onjy5pKREzhq0g6K6XmrUAVMXtOIWQkgfRLeUKzVypbo0AQREKVhaWvb29nZ1dTEYjKysLFdXVyS3OFh9fb2FhUVQUBBxODkFwZCEFpMqgmAsFuvtt99OS0tDCBUXF7/22mu0taqIgCmEmKmydb20rAOmLkQVt5B8olvKKW4hpSTF5Cw1gqW6NIFu94VrE4U2x+fm5np6erLZ7IyMjIGBASy3ONju3bsLCwuvXLliZWVF/n9OTkEw/EctJlUEwRISEtra2mxtbVtaWtatW/f8+fPJkyeL9k4VETCpSI6tmKmydb2wmnTAtPz+g6jiFpZbdEs5xS2srKSYQlJdWkNv31SBFaJ0QkJCbty4UVhY+Msvv5C7zvIoIwmFwvz8/Ndff93Hx0coFH777bfyC4KhP2oxqSgIxmKx4uPjV65c6e/vP2bMGLpRFUXA5EeqqWgIXS+kXR0wdSGpuIXkEN1SXXELKSIpNmypkSvVpQkgIEqhuro6Li6Oz+czmczx48c/ffoUyaeMVFhYWF1dbWxsbGxs3NnZefz4cfkFwdAfv/OqC4Jt3rx50aJFS5cupVNUFwGTHzFTh9X10qYOmLqQVNxCcohuqa64heSTFJOz1AiW6tIEulyeahf5V+kCgeDjjz+2t7dnMBhvvvnmkydP8P/Uk2SIg5EvwNy5c0klc+fORQitW7dOHkEwLE2LSTlBMKLrRS4zCU5OTgghf39/FUXAFBpbSVNl6HphNemAaf9CjChuYYzlFN0aVnELDyHtpYSkmDylRKW6tIneXjJTWM/UCjSHKv9iDshG9bF1d3eXf32kOTP0gTNnzjAYjODgYC2U0hV6O1Oj6E0VQG/x9fXlcDgxMTEqXqcbBqtWrdJaKUAMCIiA7gEdMEBPgIcqAAAAg0BABAAAGGR0PVTRtQkGC4/HU3HfoiGZAciDfj5UGUUBEQAAQDZwyQwAADDI/wFLahrsvfVezgAAAABJRU5ErkJggg==", "text/plain": [ "" ] }, "execution_count": 9, "metadata": {}, "output_type": "execute_result" } ], "source": [ "tf.keras.utils.plot_model(new_embedding_network, show_shapes=True, dpi=48)" ] }, { "cell_type": "markdown", "metadata": { "id": "9cOaGQHLv12W" }, "source": [ "We can then build a new encoder using the above `new_embedding_network`." ] }, { "cell_type": "code", "execution_count": 10, "metadata": { "execution": { "iopub.execute_input": "2023-12-14T12:09:42.940943Z", "iopub.status.busy": "2023-12-14T12:09:42.940693Z", "iopub.status.idle": "2023-12-14T12:09:43.501467Z", "shell.execute_reply": "2023-12-14T12:09:43.500740Z" }, "id": "mtFDMNf2vIl9" }, "outputs": [ { "name": "stderr", "output_type": "stream", "text": [ "WARNING:absl:The `Transformer` layer is deprecated. Please directly use `TransformerEncoderBlock`.\n" ] }, { "name": "stderr", "output_type": "stream", "text": [ "WARNING:absl:The `Transformer` layer is deprecated. Please directly use `TransformerEncoderBlock`.\n" ] }, { "name": "stderr", "output_type": "stream", "text": [ "WARNING:absl:The `Transformer` layer is deprecated. Please directly use `TransformerEncoderBlock`.\n" ] }, { "name": "stderr", "output_type": "stream", "text": [ "/tmpfs/src/tf_docs_env/lib/python3.9/site-packages/keras/src/initializers/initializers.py:120: UserWarning: The initializer TruncatedNormal is unseeded and being called multiple times, which will return identical values each time (even if the initializer is unseeded). Please update your code to provide a seed to the initializer, or avoid using the same initializer instance more than once.\n", " warnings.warn(\n" ] }, { "name": "stdout", "output_type": "stream", "text": [ "[, ]\n" ] } ], "source": [ "kwargs = dict(default_kwargs)\n", "\n", "# Use new embedding network.\n", "kwargs['embedding_cls'] = new_embedding_network\n", "kwargs['embedding_data'] = embedding_layer.embeddings\n", "\n", "encoder_with_customized_embedding = nlp.networks.EncoderScaffold(**kwargs)\n", "classifier_model = build_classifier(encoder_with_customized_embedding)\n", "# ... Train the model ...\n", "print(classifier_model.inputs)\n", "\n", "# Assert that there are only two inputs.\n", "assert len(classifier_model.inputs) == 2" ] }, { "cell_type": "markdown", "metadata": { "id": "Z73ZQDtmwg9K" }, "source": [ "#### Customized Transformer\n", "\n", "Users can also override the `hidden_cls` argument in `networks.EncoderScaffold`'s constructor employ a customized Transformer layer.\n", "\n", "See [the source of `nlp.layers.ReZeroTransformer`](https://github.com/tensorflow/models/blob/master/official/nlp/modeling/layers/rezero_transformer.py) for how to implement a customized Transformer layer.\n", "\n", "The following is an example of using `nlp.layers.ReZeroTransformer`:\n" ] }, { "cell_type": "code", "execution_count": 11, "metadata": { "execution": { "iopub.execute_input": "2023-12-14T12:09:43.505084Z", "iopub.status.busy": "2023-12-14T12:09:43.504552Z", "iopub.status.idle": "2023-12-14T12:09:44.240060Z", "shell.execute_reply": "2023-12-14T12:09:44.239337Z" }, "id": "uAIarLZgw6pA" }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "tf.Tensor(\n", "[[-0.08663296 0.09281035]\n", " [-0.07291833 0.36477187]\n", " [-0.08730186 0.1503254 ]], shape=(3, 2), dtype=float32)\n" ] } ], "source": [ "kwargs = dict(default_kwargs)\n", "\n", "# Use ReZeroTransformer.\n", "kwargs['hidden_cls'] = nlp.layers.ReZeroTransformer\n", "\n", "encoder_with_rezero_transformer = nlp.networks.EncoderScaffold(**kwargs)\n", "classifier_model = build_classifier(encoder_with_rezero_transformer)\n", "# ... Train the model ...\n", "predict(classifier_model)\n", "\n", "# Assert that the variable `rezero_alpha` from ReZeroTransformer exists.\n", "assert 'rezero_alpha' in ''.join([x.name for x in classifier_model.trainable_weights])" ] }, { "cell_type": "markdown", "metadata": { "id": "6PMHFdvnxvR0" }, "source": [ "### Use `nlp.layers.TransformerScaffold`\n", "\n", "The above method of customizing the model requires rewriting the whole `nlp.layers.Transformer` layer, while sometimes you may only want to customize either attention layer or feedforward block. In this case, `nlp.layers.TransformerScaffold` can be used.\n" ] }, { "cell_type": "markdown", "metadata": { "id": "D6FejlgwyAy_" }, "source": [ "#### Customize Attention Layer\n", "\n", "User can also override the `attention_cls` argument in `layers.TransformerScaffold`'s constructor to employ a customized Attention layer.\n", "\n", "See [the source of `nlp.layers.TalkingHeadsAttention`](https://github.com/tensorflow/models/blob/master/official/nlp/modeling/layers/talking_heads_attention.py) for how to implement a customized `Attention` layer.\n", "\n", "Following is an example of using `nlp.layers.TalkingHeadsAttention`:" ] }, { "cell_type": "code", "execution_count": 12, "metadata": { "execution": { "iopub.execute_input": "2023-12-14T12:09:44.243540Z", "iopub.status.busy": "2023-12-14T12:09:44.243284Z", "iopub.status.idle": "2023-12-14T12:09:45.344434Z", "shell.execute_reply": "2023-12-14T12:09:45.343602Z" }, "id": "nFrSMrZuyNeQ" }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "tf.Tensor(\n", "[[-0.20591784 0.09203205]\n", " [-0.0056177 -0.10278902]\n", " [-0.21681327 -0.12282 ]], shape=(3, 2), dtype=float32)\n" ] } ], "source": [ "# Use TalkingHeadsAttention\n", "hidden_cfg = dict(default_hidden_cfg)\n", "hidden_cfg['attention_cls'] = nlp.layers.TalkingHeadsAttention\n", "\n", "kwargs = dict(default_kwargs)\n", "kwargs['hidden_cls'] = nlp.layers.TransformerScaffold\n", "kwargs['hidden_cfg'] = hidden_cfg\n", "\n", "encoder = nlp.networks.EncoderScaffold(**kwargs)\n", "classifier_model = build_classifier(encoder)\n", "# ... Train the model ...\n", "predict(classifier_model)\n", "\n", "# Assert that the variable `pre_softmax_weight` from TalkingHeadsAttention exists.\n", "assert 'pre_softmax_weight' in ''.join([x.name for x in classifier_model.trainable_weights])" ] }, { "cell_type": "code", "execution_count": 13, "metadata": { "execution": { "iopub.execute_input": "2023-12-14T12:09:45.347489Z", "iopub.status.busy": "2023-12-14T12:09:45.347221Z", "iopub.status.idle": "2023-12-14T12:09:45.497660Z", "shell.execute_reply": "2023-12-14T12:09:45.496874Z" }, "id": "tKkZ8spzYmpc" }, "outputs": [ { "data": { "image/png": "", "text/plain": [ "" ] }, "execution_count": 13, "metadata": {}, "output_type": "execute_result" } ], "source": [ "tf.keras.utils.plot_model(encoder_with_rezero_transformer, show_shapes=True, dpi=48)" ] }, { "cell_type": "markdown", "metadata": { "id": "kuEJcTyByVvI" }, "source": [ "#### Customize Feedforward Layer\n", "\n", "Similiarly, one could also customize the feedforward layer.\n", "\n", "See [the source of `nlp.layers.GatedFeedforward`](https://github.com/tensorflow/models/blob/master/official/nlp/modeling/layers/gated_feedforward.py) for how to implement a customized feedforward layer.\n", "\n", "Following is an example of using `nlp.layers.GatedFeedforward`:" ] }, { "cell_type": "code", "execution_count": 14, "metadata": { "execution": { "iopub.execute_input": "2023-12-14T12:09:45.501208Z", "iopub.status.busy": "2023-12-14T12:09:45.500920Z", "iopub.status.idle": "2023-12-14T12:09:46.364316Z", "shell.execute_reply": "2023-12-14T12:09:46.363561Z" }, "id": "XAbKy_l4y_-i" }, "outputs": [ { "name": "stderr", "output_type": "stream", "text": [ "/tmpfs/src/tf_docs_env/lib/python3.9/site-packages/keras/src/initializers/initializers.py:120: UserWarning: The initializer TruncatedNormal is unseeded and being called multiple times, which will return identical values each time (even if the initializer is unseeded). Please update your code to provide a seed to the initializer, or avoid using the same initializer instance more than once.\n", " warnings.warn(\n" ] }, { "name": "stdout", "output_type": "stream", "text": [ "tf.Tensor(\n", "[[-0.10270456 -0.10999684]\n", " [-0.03512481 0.15430304]\n", " [-0.23601504 -0.18162844]], shape=(3, 2), dtype=float32)\n" ] } ], "source": [ "# Use GatedFeedforward\n", "hidden_cfg = dict(default_hidden_cfg)\n", "hidden_cfg['feedforward_cls'] = nlp.layers.GatedFeedforward\n", "\n", "kwargs = dict(default_kwargs)\n", "kwargs['hidden_cls'] = nlp.layers.TransformerScaffold\n", "kwargs['hidden_cfg'] = hidden_cfg\n", "\n", "encoder_with_gated_feedforward = nlp.networks.EncoderScaffold(**kwargs)\n", "classifier_model = build_classifier(encoder_with_gated_feedforward)\n", "# ... Train the model ...\n", "predict(classifier_model)\n", "\n", "# Assert that the variable `gate` from GatedFeedforward exists.\n", "assert 'gate' in ''.join([x.name for x in classifier_model.trainable_weights])" ] }, { "cell_type": "markdown", "metadata": { "id": "a_8NWUhkzeAq" }, "source": [ "### Build a new Encoder\n", "\n", "Finally, you could also build a new encoder using building blocks in the modeling library.\n", "\n", "See [the source for `nlp.networks.AlbertEncoder`](https://github.com/tensorflow/models/blob/master/official/nlp/modeling/networks/albert_encoder.py) as an example of how to do this. \n", "\n", "Here is an example using `nlp.networks.AlbertEncoder`:\n" ] }, { "cell_type": "code", "execution_count": 15, "metadata": { "execution": { "iopub.execute_input": "2023-12-14T12:09:46.367627Z", "iopub.status.busy": "2023-12-14T12:09:46.367375Z", "iopub.status.idle": "2023-12-14T12:09:46.924698Z", "shell.execute_reply": "2023-12-14T12:09:46.923962Z" }, "id": "xsiA3RzUzmUM" }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "tf.Tensor(\n", "[[-0.00369881 -0.2540995 ]\n", " [ 0.1235221 -0.2959229 ]\n", " [-0.08698564 -0.17653546]], shape=(3, 2), dtype=float32)\n" ] } ], "source": [ "albert_encoder = nlp.networks.AlbertEncoder(**cfg)\n", "classifier_model = build_classifier(albert_encoder)\n", "# ... Train the model ...\n", "predict(classifier_model)" ] }, { "cell_type": "markdown", "metadata": { "id": "MeidDfhlHKSO" }, "source": [ "Inspecting the `albert_encoder`, we see it stacks the same `Transformer` layer multiple times (note the loop-back on the \"Transformer\" block below.." ] }, { "cell_type": "code", "execution_count": 16, "metadata": { "execution": { "iopub.execute_input": "2023-12-14T12:09:46.927962Z", "iopub.status.busy": "2023-12-14T12:09:46.927677Z", "iopub.status.idle": "2023-12-14T12:09:47.076921Z", "shell.execute_reply": "2023-12-14T12:09:47.076078Z" }, "id": "Uv_juT22HERW" }, "outputs": [ { "data": { "image/png": "", "text/plain": [ "" ] }, "execution_count": 16, "metadata": {}, "output_type": "execute_result" } ], "source": [ "tf.keras.utils.plot_model(albert_encoder, show_shapes=True, dpi=48)" ] } ], "metadata": { "colab": { "collapsed_sections": [], "name": "customize_encoder.ipynb", "provenance": [], "toc_visible": true }, "kernelspec": { "display_name": "Python 3", "name": "python3" }, "language_info": { "codemirror_mode": { "name": "ipython", "version": 3 }, "file_extension": ".py", "mimetype": "text/x-python", "name": "python", "nbconvert_exporter": "python", "pygments_lexer": "ipython3", "version": "3.9.18" } }, "nbformat": 4, "nbformat_minor": 0 }