{"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":" # <h1 style='background:#F7B2B0; border:0; color:black'><center>THE GATEDTABTRANSFORMER - AN ENHANCED DEEP LEARNING ARCHITECTURE FOR TABULAR MODELING.</center></h1> ","metadata":{}},{"cell_type":"markdown","source":"This is a two part tutorial series containing the implementation of [GatedTabTransformer](https://arxiv.org/pdf/2201.00199.pdf) paper in both FLAX and TensorFlow (TPU)\n\n[Part 1 : GatedTabTransformer in FLAX](https://www.kaggle.com/code/usharengaraju/gatedtabtransformer-flax)\n\n[Part 2 : GatedTabTransformer in TensorFlow + TPU](https://www.kaggle.com/code/usharengaraju/tensorflow-tpu-gatedtabtransformer)","metadata":{}},{"cell_type":"markdown","source":"### Abstract\n\nOver the last few years the research towards using deep learning for tabular data has been on the rise .The state of the art TabTransformer incorporates an attention mechanism to better track relationships between categorical features and then makes use of a standard MLP to output its final logits.GatedTabTransformer implements linear projections are implemented in the MLP block and the paper also experiments with several activation functions .\n\n### Introduction\n\nTabular data is the most commonly used data type in real world applications . Tree based ensemble methods like LightGBM , XGBoost are the current state of the art approaches for tabular data . Over the last few years , there is increasing interest in the usage of deep learning techniques for tabular data primarily because of the eliminating the need for manual embedding and feature engineering . Some of the neural networks architectures which have comparable performance with Tree based ensemble methods are TabNet , DNF-Net etc.\n\nThere is an increasing usage of attention-based architectures like Transformers which was originally used to  handle NLP tasks to solve tabular data problems . TabTransformer is one such architecture which focuses on using Multi-Head Self Attention blocks to model relationships between the categorical features in tabular data, transforming them into robust contextual embeddings.The transformed categorical features are concatenated with continuous values and then fed through a standard multilayer perceptron which makes TabTransformer significantly outperform other deep learning counterparts like TabNet and MLP . GatedTabTransformer further enhances TabTransformer by replacing the final MLP block with a gated multi-layer perceptron (gMLP) , a simple MLP-based network with spatial gating projections, which aims to be on par with Transformers in terms of performance on sequential data\n\n\n### TabTransformer    \n\nThe TabTransformer model, outperforms the other state-of-the-art deep learning methods for tabular data by at least 1.0% on mean AUROC. It consists of a column embedding layer, a stack of N Transformer layers, and a multilayer perceptron . The inputted tabular features are split in two parts for the categorical and continuous values. Column embedding is performed for each categorical feature .It generates parametric embeddings which are inputted to a stack of Transformer layers. Each Transformer layer consists of a multi-head self-attention layer followed by a position-wise feed-forward layer. After the processing of categorical values , they are concatenated along with the continuous values to form a final feature vector which is inputted to a standard multilayer perceptron\n\n![](https://i.imgur.com/mnv2bLy.png)\n\n\n### gMLP model\n\nThe gMLP model consists of a stack of multiple identically structured blocks ,activation function and linear projections along the channel dimension and the spatial gating unit which captures spatial cross-token interactions. The weights are initialized as near-zero values and the biases as ones at the beginning of training. This structure does not require positional embeddings because relevant information will be captured in the gating units. gMLP has been proposed as an alternative to Transformers for NLP and vision tasks having up to 66% less trainable parameters. GatedTabTransformers replaces the pure MLP block in the TabTransformer with gMLP \n\n![](https://i.imgur.com/SRyVmYY.png)\n\n\n\n### GatedTabTransformer\n\nThe column embeddings are generated from categorical data features and continuous values are passed through a normalization layer. The categorical embeddings are then processed by a Transformer block. \n\nTransformer block represents the encoder part of a Transformer . It has two sub-layers - a multi-head self-attention mechanism, and a simple, position wise fully connected feed-forward network. In the final layer MLP is replaced by gMLP and the architecture is adapted to output classification logits and works best for optimization of cross entropy or binary cross entropy loss \n\n![](https://i.imgur.com/ROdKcQy.png)\n","metadata":{}},{"cell_type":"markdown","source":"<img src=\"https://camo.githubusercontent.com/dd842f7b0be57140e68b2ab9cb007992acd131c48284eaf6b1aca758bfea358b/68747470733a2f2f692e696d6775722e636f6d2f52557469567a482e706e67\">\n\n> I will be integrating W&B for visualizations and logging artifacts!\n> \n> [GatedTabTransformer in FLAX project on W&B Dashboard](https://wandb.ai/usharengaraju/GatedTabTransformer_FLAX)\n> \n> - To get the API key, create an account in the [website](https://wandb.ai/site) .\n> - Use secrets to use API Keys more securely","metadata":{}},{"cell_type":"markdown","source":"# **<span style=\"color:#F7B2B0;\">W & B Artifacts</span>**\n\nAn artifact as a versioned folder of data.Entire datasets can be directly stored as artifacts .\n\nW&B Artifacts are used for dataset versioning, model versioning . They are also used for tracking dependencies and results across machine learning pipelines.Artifact references can be used to point to data in other systems like S3, GCP, or your own system.\n\nYou can learn more about W&B artifacts [here](https://docs.wandb.ai/guides/artifacts)\n\n![](https://drive.google.com/uc?id=1JYSaIMXuEVBheP15xxuaex-32yzxgglV)","metadata":{}},{"cell_type":"code","source":"import wandb\nfrom kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nwandb_key = user_secrets.get_secret(\"api_key\")\nwandb.login(key = wandb_key)","metadata":{"execution":{"iopub.status.busy":"2022-11-25T16:37:52.746096Z","iopub.execute_input":"2022-11-25T16:37:52.746849Z","iopub.status.idle":"2022-11-25T16:37:53.139552Z","shell.execute_reply.started":"2022-11-25T16:37:52.746804Z","shell.execute_reply":"2022-11-25T16:37:53.138889Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Save training data to W&B Artifacts\nrun = wandb.init(project='GatedTabTransformer_FLAX', name='processed_data') \nartifact = wandb.Artifact(name='processed_data',type='dataset')\nartifact.add_file(\"/kaggle/input/amex-tfrecords/minidata (1).csv\")\nwandb.log_artifact(artifact)\nwandb.finish()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"![](https://i.imgur.com/tRDVISy.png)","metadata":{}},{"cell_type":"code","source":"!pip install --quiet tensorflow-addons","metadata":{"_kg_hide-output":true,"execution":{"iopub.execute_input":"2022-08-06T18:23:32.254967Z","iopub.status.busy":"2022-08-06T18:23:32.253793Z","iopub.status.idle":"2022-08-06T18:23:41.206329Z","shell.execute_reply":"2022-08-06T18:23:41.205588Z"},"papermill":{"duration":8.975822,"end_time":"2022-08-06T18:23:41.206496","exception":false,"start_time":"2022-08-06T18:23:32.230674","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import tensorflow as tf\nfrom kaggle_datasets import KaggleDatasets\ndef configure_device():\n    try:\n        tpu = tf.distribute.cluster_resolver.TPUClusterResolver.connect()  # connect to tpu cluster\n        strategy = tf.distribute.TPUStrategy(tpu) # get strategy for tpu\n        print('Num of TPUs: ', strategy.num_replicas_in_sync)\n        device='TPU'\n    except: # otherwise detect GPUs\n        tpu = None\n        gpus = tf.config.list_logical_devices('GPU') # get logical gpus\n        ngpu = len(gpus)\n        if ngpu: # if number of GPUs are 0 then CPU\n            strategy = tf.distribute.MirroredStrategy(gpus) # single-GPU or multi-GPU\n            print(\"> Running on GPU\", end=' | ')\n            print(\"Num of GPUs: \", ngpu)\n            device='GPU'\n        else:\n            print(\"> Running on CPU\")\n            strategy = tf.distribute.get_strategy() # connect to single gpu or cpu\n            device='CPU'\n    return strategy, device, tpu","metadata":{"_kg_hide-output":true,"execution":{"iopub.execute_input":"2022-08-06T18:23:41.253296Z","iopub.status.busy":"2022-08-06T18:23:41.252642Z","iopub.status.idle":"2022-08-06T18:23:46.663110Z","shell.execute_reply":"2022-08-06T18:23:46.662495Z"},"papermill":{"duration":5.437163,"end_time":"2022-08-06T18:23:46.663257","exception":false,"start_time":"2022-08-06T18:23:41.226094","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"strategy, device, tpu = configure_device()\nAUTO = tf.data.experimental.AUTOTUNE","metadata":{"_kg_hide-output":true,"execution":{"iopub.execute_input":"2022-08-06T18:23:46.715460Z","iopub.status.busy":"2022-08-06T18:23:46.714799Z","iopub.status.idle":"2022-08-06T18:23:52.775172Z","shell.execute_reply":"2022-08-06T18:23:52.775687Z"},"papermill":{"duration":6.091953,"end_time":"2022-08-06T18:23:52.775876","exception":false,"start_time":"2022-08-06T18:23:46.683923","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\n\nimport pandas as pd\ndf  = pd.read_csv('../input/amex-tfrecords/minidata (1).csv')","metadata":{"execution":{"iopub.execute_input":"2022-08-06T18:23:52.822362Z","iopub.status.busy":"2022-08-06T18:23:52.821669Z","iopub.status.idle":"2022-08-06T18:23:54.237940Z","shell.execute_reply":"2022-08-06T18:23:54.237012Z","shell.execute_reply.started":"2022-08-06T17:54:01.039442Z"},"papermill":{"duration":1.442091,"end_time":"2022-08-06T18:23:54.238098","exception":false,"start_time":"2022-08-06T18:23:52.796007","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.select_dtypes('int').columns","metadata":{"execution":{"iopub.execute_input":"2022-08-06T18:23:54.283563Z","iopub.status.busy":"2022-08-06T18:23:54.282616Z","iopub.status.idle":"2022-08-06T18:23:54.301469Z","shell.execute_reply":"2022-08-06T18:23:54.300933Z","shell.execute_reply.started":"2022-08-06T17:54:07.662409Z"},"papermill":{"duration":0.043651,"end_time":"2022-08-06T18:23:54.301653","exception":false,"start_time":"2022-08-06T18:23:54.258002","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cat_features  = ['B_30', 'B_38', 'D_63', 'D_64', 'D_66', 'D_68', 'D_114', 'D_116', 'D_117', 'D_120', 'D_126']\ncont_features = [x for x in list(df.columns) if x not in cat_features]","metadata":{"execution":{"iopub.execute_input":"2022-08-06T18:23:54.350683Z","iopub.status.busy":"2022-08-06T18:23:54.349969Z","iopub.status.idle":"2022-08-06T18:23:54.353752Z","shell.execute_reply":"2022-08-06T18:23:54.353190Z","shell.execute_reply.started":"2022-08-06T17:54:11.198022Z"},"papermill":{"duration":0.030467,"end_time":"2022-08-06T18:23:54.353897","exception":false,"start_time":"2022-08-06T18:23:54.323430","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nff = list(df.select_dtypes('int').columns)\ntt  = [x for x in list(df.columns) if x not in ff]\n\ndef parse_tfr_element(element):\n  data = {}\n  for col in ff:\n    data[col] = tf.io.FixedLenFeature([], tf.int64)\n\n  for col in tt:\n    data[col] = tf.io.FixedLenFeature([], tf.float32)\n  \n    \n  content = tf.io.parse_single_example(element, data)\n  \n  my_arr = []\n  for col in list(df.columns):\n    my_arr.append(content[col])\n  return my_arr\n\ndef get_dataset_small(filename):\n  dataset = tf.data.TFRecordDataset(filename)\n  dataset = dataset.map(\n      parse_tfr_element\n  )\n  return dataset\ngspath = KaggleDatasets().get_gcs_path('amex-tfrecords')\nfilename = tf.io.gfile.glob(gspath + '/*.tfrecord')\ndataset_small = get_dataset_small(filename[0])","metadata":{"_kg_hide-output":true,"execution":{"iopub.execute_input":"2022-08-06T18:23:54.406953Z","iopub.status.busy":"2022-08-06T18:23:54.406240Z","iopub.status.idle":"2022-08-06T18:23:54.831881Z","shell.execute_reply":"2022-08-06T18:23:54.832438Z"},"papermill":{"duration":0.458315,"end_time":"2022-08-06T18:23:54.832650","exception":false,"start_time":"2022-08-06T18:23:54.374335","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm import tqdm\narr = []\nfor i in tqdm(dataset_small.take(1000)):\n  temp = [j.numpy() for j in i]\n  arr.append(temp)","metadata":{"_kg_hide-output":true,"execution":{"iopub.execute_input":"2022-08-06T18:23:54.883807Z","iopub.status.busy":"2022-08-06T18:23:54.882895Z","iopub.status.idle":"2022-08-06T18:28:38.678826Z","shell.execute_reply":"2022-08-06T18:28:38.676892Z"},"papermill":{"duration":283.82517,"end_time":"2022-08-06T18:28:38.678996","exception":false,"start_time":"2022-08-06T18:23:54.853826","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"arr = np.array(arr)\narr.shape","metadata":{"execution":{"iopub.execute_input":"2022-08-06T18:28:39.309892Z","iopub.status.busy":"2022-08-06T18:28:39.308903Z","iopub.status.idle":"2022-08-06T18:28:39.314882Z","shell.execute_reply":"2022-08-06T18:28:39.314208Z","shell.execute_reply.started":"2022-08-06T18:00:03.751237Z"},"papermill":{"duration":0.337025,"end_time":"2022-08-06T18:28:39.315028","exception":false,"start_time":"2022-08-06T18:28:38.978003","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df1 = pd.DataFrame(arr,columns=list(df.columns))\ndf1","metadata":{"execution":{"iopub.execute_input":"2022-08-06T18:28:39.918373Z","iopub.status.busy":"2022-08-06T18:28:39.917386Z","iopub.status.idle":"2022-08-06T18:28:39.948595Z","shell.execute_reply":"2022-08-06T18:28:39.948017Z","shell.execute_reply.started":"2022-08-06T18:00:08.127382Z"},"papermill":{"duration":0.334848,"end_time":"2022-08-06T18:28:39.948743","exception":false,"start_time":"2022-08-06T18:28:39.613895","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_X_from_groups(feature_set, groups):\n    result = []\n    for group in groups:\n        result.append(feature_set[group])\n    return result\n\ndef get_X_from_features(feature_set, cont_features, cat_features):\n    groups = [cont_features]\n    groups.extend(cat_features)\n    return get_X_from_groups(feature_set, groups)","metadata":{"execution":{"iopub.execute_input":"2022-08-06T18:28:40.547823Z","iopub.status.busy":"2022-08-06T18:28:40.547137Z","iopub.status.idle":"2022-08-06T18:28:40.551595Z","shell.execute_reply":"2022-08-06T18:28:40.552075Z","shell.execute_reply.started":"2022-08-06T18:00:11.077954Z"},"papermill":{"duration":0.30389,"end_time":"2022-08-06T18:28:40.552264","exception":false,"start_time":"2022-08-06T18:28:40.248374","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cont_features.remove('target')\nlen(cont_features)","metadata":{"execution":{"iopub.execute_input":"2022-08-06T18:28:41.146466Z","iopub.status.busy":"2022-08-06T18:28:41.145799Z","iopub.status.idle":"2022-08-06T18:28:41.150255Z","shell.execute_reply":"2022-08-06T18:28:41.150783Z","shell.execute_reply.started":"2022-08-06T18:00:14.310641Z"},"papermill":{"duration":0.301853,"end_time":"2022-08-06T18:28:41.150971","exception":false,"start_time":"2022-08-06T18:28:40.849118","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cats = []\nfor col in cat_features:\n  cats.append(df1[col].unique().shape[0])","metadata":{"execution":{"iopub.execute_input":"2022-08-06T18:28:41.744513Z","iopub.status.busy":"2022-08-06T18:28:41.743887Z","iopub.status.idle":"2022-08-06T18:28:41.747372Z","shell.execute_reply":"2022-08-06T18:28:41.747959Z","shell.execute_reply.started":"2022-08-06T18:00:18.133863Z"},"papermill":{"duration":0.304984,"end_time":"2022-08-06T18:28:41.748122","exception":false,"start_time":"2022-08-06T18:28:41.443138","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X = get_X_from_features(df1.drop(columns=['target'],axis=1),cont_features,cat_features)\ny = df1['target']","metadata":{"execution":{"iopub.execute_input":"2022-08-06T18:28:42.432898Z","iopub.status.busy":"2022-08-06T18:28:42.432128Z","iopub.status.idle":"2022-08-06T18:28:42.434292Z","shell.execute_reply":"2022-08-06T18:28:42.434793Z","shell.execute_reply.started":"2022-08-06T18:00:21.902970Z"},"papermill":{"duration":0.30692,"end_time":"2022-08-06T18:28:42.434975","exception":false,"start_time":"2022-08-06T18:28:42.128055","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import tensorflow as tf \nfrom tensorflow import keras\nfrom tensorflow.keras import layers\nimport tensorflow_addons as tfa\nclass gMLPLayer(layers.Layer):\n    def __init__(self, num_patches, embedding_dim, dropout_rate, *args, **kwargs):\n        super(gMLPLayer, self).__init__(*args, **kwargs)\n\n        self.channel_projection1 = keras.Sequential(\n            [\n                layers.Dense(units=embedding_dim * 2),\n                tfa.layers.GELU(),\n                layers.Dropout(rate=dropout_rate),\n            ]\n        )\n\n        self.channel_projection2 = layers.Dense(units=embedding_dim)\n\n        self.spatial_projection = layers.Dense(\n            units=num_patches, bias_initializer=\"Ones\"\n        )\n\n        self.normalize1 = layers.LayerNormalization(epsilon=1e-6)\n        self.normalize2 = layers.LayerNormalization(epsilon=1e-6)\n\n    def spatial_gating_unit(self, x):\n        # Split x along the channel dimensions.\n        # Tensors u and v will in th shape of [batch_size, num_patchs, embedding_dim].\n        u, v = tf.split(x, num_or_size_splits=2, axis=2)\n        # Apply layer normalization.\n        v = self.normalize2(v)\n        # Apply spatial projection.\n        v_channels = tf.linalg.matrix_transpose(v)\n        v_projected = self.spatial_projection(v_channels)\n        v_projected = tf.linalg.matrix_transpose(v_projected)\n        # Apply element-wise multiplication.\n        return u * v_projected\n\n    def call(self, inputs):\n        # Apply layer normalization.\n        x = self.normalize1(inputs)\n        # Apply the first channel projection. x_projected shape: [batch_size, num_patches, embedding_dim * 2].\n        x_projected = self.channel_projection1(x)\n        # Apply the spatial gating unit. x_spatial shape: [batch_size, num_patches, embedding_dim].\n        x_spatial = self.spatial_gating_unit(x_projected)\n        # Apply the second channel projection. x_projected shape: [batch_size, num_patches, embedding_dim].\n        x_projected = self.channel_projection2(x_spatial)\n        # Add skip connection.\n        return x + x_projected","metadata":{"execution":{"iopub.execute_input":"2022-08-06T18:28:43.035504Z","iopub.status.busy":"2022-08-06T18:28:43.034865Z","iopub.status.idle":"2022-08-06T18:28:43.162816Z","shell.execute_reply":"2022-08-06T18:28:43.162192Z","shell.execute_reply.started":"2022-08-06T18:03:57.793729Z"},"papermill":{"duration":0.42824,"end_time":"2022-08-06T18:28:43.162984","exception":false,"start_time":"2022-08-06T18:28:42.734744","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TransformerBlock(layers.Layer):\n    def __init__(self, embed_dim, num_heads, ff_dim, rate=0.1):\n        super(TransformerBlock, self).__init__()\n        # parametreleri\n        self.att = layers.MultiHeadAttention(num_heads=num_heads, key_dim=embed_dim)\n        self.ffn = keras.Sequential(\n            [layers.Dense(ff_dim, activation=\"relu\"), layers.Dense(embed_dim),]\n        )\n        # batch-layer\n        self.layernorm1 = layers.LayerNormalization(epsilon=1e-6)\n        self.layernorm2 = layers.LayerNormalization(epsilon=1e-6)\n        self.dropout1 = layers.Dropout(rate)\n        self.dropout2 = layers.Dropout(rate)\n\n    def call(self, inputs, training):\n        attn_output = self.att(inputs, inputs)\n        attn_output = self.dropout1(attn_output, training=training)\n        out1 = self.layernorm1(inputs + attn_output)\n        ffn_output = self.ffn(out1)\n        ffn_output = self.dropout2(ffn_output, training=training)\n        return self.layernorm2(out1 + ffn_output)\n\nclass GatedTabTransformer(keras.Model):\n\n    def __init__(self, \n            categories,\n            num_continuous,\n            dim,\n            dim_out,\n            depth,\n            embedding_dim,\n            heads,\n            attn_dropout,\n            ff_dropout,\n            gmlp_blocks,\n            normalize_continuous = True):\n\n        super(GatedTabTransformer, self).__init__()\n\n        # --> continuous inputs\n        self.embedding_dim = embedding_dim\n        self.normalize_continuous = normalize_continuous\n        if normalize_continuous:\n            self.continuous_normalization = layers.LayerNormalization()\n\n        # --> categorical inputs\n\n        # embedding\n        self.embedding_layers = []\n        for number_of_classes in categories:\n            self.embedding_layers.append(layers.Embedding(input_dim = number_of_classes, output_dim = dim))\n\n        # concatenation\n        self.embedded_concatenation = layers.Concatenate(axis=1)\n\n        # adding transformers\n        self.transformers = []\n        for _ in range(depth):\n            self.transformers.append(TransformerBlock(dim, heads, dim))\n        self.flatten_transformer_output = layers.Flatten()\n\n        # --> MLP\n        self.pre_mlp_concatenation = layers.Concatenate()\n\n        # mlp layers\n        self.gmlp_layers = []\n        for _ in range(gmlp_blocks):\n            self.gmlp_layers.append(gMLPLayer(1,self.embedding_dim,0.2))\n        self.embedder2 = layers.Dense(self.embedding_dim)\n        self.output_layer = layers.Dense(dim_out,activation='sigmoid')\n\n    def call(self, inputs):\n        continuous_inputs  = inputs[0]\n        categorical_inputs = inputs[1:]\n        \n        # --> continuous\n        if self.normalize_continuous:\n            continuous_inputs = self.continuous_normalization(continuous_inputs)\n\n        # --> categorical\n        embedding_outputs = []\n        for categorical_input, embedding_layer in zip(categorical_inputs, self.embedding_layers):\n            embedding_outputs.append(embedding_layer(categorical_input))\n        categorical_inputs = self.embedded_concatenation(embedding_outputs)\n        # print(embedding_outputs[0].shape)\n        \n        for transformer in self.transformers:\n            categorical_inputs = transformer(categorical_inputs)\n        contextual_embedding = self.flatten_transformer_output(categorical_inputs)\n        # print(categorical_inputs.shape)\n        # --> MLP\n        mlp_input = self.pre_mlp_concatenation([continuous_inputs, contextual_embedding])\n        gmlp_input = tf.expand_dims(self.embedder2(mlp_input),axis=1)\n        for gmlp_layer in self.gmlp_layers:\n            gmlp_input = gmlp_layer(gmlp_input)\n        gmlp_input = tf.math.reduce_mean(gmlp_input,axis=1)\n        return self.output_layer(gmlp_input)\n  ","metadata":{"execution":{"iopub.execute_input":"2022-08-06T18:28:43.762789Z","iopub.status.busy":"2022-08-06T18:28:43.762075Z","iopub.status.idle":"2022-08-06T18:28:43.779542Z","shell.execute_reply":"2022-08-06T18:28:43.780093Z","shell.execute_reply.started":"2022-08-06T18:04:01.848923Z"},"papermill":{"duration":0.316197,"end_time":"2022-08-06T18:28:43.780271","exception":false,"start_time":"2022-08-06T18:28:43.464074","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with strategy.scope():\n  model = GatedTabTransformer(\n      categories = cats, # number of unique elements in each categorical feature\n      num_continuous = 177,      # number of numerical features\n      dim = 16,                # embedding/transformer dimension\n      dim_out = 1,             # dimension of the model output\n      depth = 6,  \n      embedding_dim=256,             # number of transformer layers in the stack\n      heads = 8,               # number of attention heads\n      attn_dropout = 0.1,      # attention layer dropout in transformers\n      ff_dropout = 0.1,        # feed-forward layer dropout in transformers\n      gmlp_blocks = 6 # mlp layer dimensions and activations\n  )\n\n  model.compile(loss='binary_crossentropy', optimizer='adam', metrics=['accuracy'])\n  history = model.fit(X,y,epochs=20,validation_split=0.2,batch_size=32,verbose=1)","metadata":{"execution":{"iopub.execute_input":"2022-08-06T18:28:44.395404Z","iopub.status.busy":"2022-08-06T18:28:44.394770Z","iopub.status.idle":"2022-08-06T18:29:46.381610Z","shell.execute_reply":"2022-08-06T18:29:46.382153Z","shell.execute_reply.started":"2022-08-06T18:04:06.505894Z"},"papermill":{"duration":62.299353,"end_time":"2022-08-06T18:29:46.382352","exception":false,"start_time":"2022-08-06T18:28:44.082999","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df2 = df1.drop(columns=['target'],axis=1)\ny = y.to_numpy()","metadata":{"execution":{"iopub.execute_input":"2022-08-06T18:29:47.100002Z","iopub.status.busy":"2022-08-06T18:29:47.098820Z","iopub.status.idle":"2022-08-06T18:29:47.101277Z","shell.execute_reply":"2022-08-06T18:29:47.101793Z","shell.execute_reply.started":"2022-08-06T18:05:13.802427Z"},"papermill":{"duration":0.3668,"end_time":"2022-08-06T18:29:47.101994","exception":false,"start_time":"2022-08-06T18:29:46.735194","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import KFold\nkfold = KFold(n_splits=5, shuffle=True)\nacc_per_fold=[]\nloss_per_fold = []\n# K-fold Cross Validation model evaluation\nfold_no = 1\nwith strategy.scope():\n  for train, test in kfold.split(df2, y):\n    train = list(train)\n    test = list(test)\n    train_df = df2.iloc[train]\n    test_df = df2.iloc[test]\n    train_y = y[train]\n    test_y = y[test]\n    train_X = get_X_from_features(train_df,cont_features,cat_features)\n    test_X = get_X_from_features(test_df,cont_features,cat_features)\n    \n    model = GatedTabTransformer(\n        categories = cats, # number of unique elements in each categorical feature\n        num_continuous = 177,      # number of numerical features\n        dim = 16,                # embedding/transformer dimension\n        dim_out = 1,             # dimension of the model output\n        depth = 6,  \n        embedding_dim=256,             # number of transformer layers in the stack\n        heads = 8,               # number of attention heads\n        attn_dropout = 0.1,      # attention layer dropout in transformers\n        ff_dropout = 0.1,        # feed-forward layer dropout in transformers\n        gmlp_blocks = 6 # mlp layer dimensions and activations\n    )\n\n    model.compile(loss='binary_crossentropy', optimizer='adam', metrics=['accuracy'])\n    \n    print('------------------------------------------------------------------------')\n    print(f'Training for fold {fold_no} ...')\n    history = model.fit(train_X,train_y,epochs=20,validation_split=0.2,batch_size=32,verbose=1)\n    \n\n    # Generate generalization metrics\n    scores = model.evaluate(test_X, test_y, verbose=0)\n    print(f'Score for fold {fold_no}: {model.metrics_names[0]} of {scores[0]}; {model.metrics_names[1]} of {scores[1]*100}%')\n    acc_per_fold.append(scores[1] * 100)\n    loss_per_fold.append(scores[0])\n\n    # Increase fold number\n    fold_no = fold_no + 1","metadata":{"_kg_hide-output":true,"execution":{"iopub.execute_input":"2022-08-06T18:29:47.829147Z","iopub.status.busy":"2022-08-06T18:29:47.828317Z","iopub.status.idle":"2022-08-06T18:35:02.759050Z","shell.execute_reply":"2022-08-06T18:35:02.759594Z"},"papermill":{"duration":315.300185,"end_time":"2022-08-06T18:35:02.759799","exception":false,"start_time":"2022-08-06T18:29:47.459614","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.save_weights('gatedtab.h5', overwrite=True)","metadata":{"execution":{"iopub.execute_input":"2022-08-06T18:35:03.957391Z","iopub.status.busy":"2022-08-06T18:35:03.956700Z","iopub.status.idle":"2022-08-06T18:35:04.442958Z","shell.execute_reply":"2022-08-06T18:35:04.442353Z","shell.execute_reply.started":"2022-08-06T18:10:33.140705Z"},"papermill":{"duration":1.084345,"end_time":"2022-08-06T18:35:04.443112","exception":false,"start_time":"2022-08-06T18:35:03.358767","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nprint(history.history.keys())\n# summarize history for accuracy\nplt.plot(history.history['accuracy'])\nplt.plot(history.history['val_accuracy'])\nplt.title('model accuracy')\nplt.ylabel('accuracy')\nplt.xlabel('epoch')\nplt.legend(['train', 'test'], loc='upper left')\nplt.show()\n# summarize history for loss\nplt.plot(history.history['loss'])\nplt.plot(history.history['val_loss'])\nplt.title('model loss')\nplt.ylabel('loss')\nplt.xlabel('epoch')\nplt.legend(['train', 'test'], loc='upper left')\nplt.show()","metadata":{"execution":{"iopub.execute_input":"2022-08-06T18:35:05.646029Z","iopub.status.busy":"2022-08-06T18:35:05.645351Z","iopub.status.idle":"2022-08-06T18:35:06.130172Z","shell.execute_reply":"2022-08-06T18:35:06.130663Z","shell.execute_reply.started":"2022-08-06T18:10:39.426475Z"},"papermill":{"duration":1.094335,"end_time":"2022-08-06T18:35:06.130833","exception":false,"start_time":"2022-08-06T18:35:05.036498","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pip install --quiet shap","metadata":{"_kg_hide-output":true,"execution":{"iopub.execute_input":"2022-08-06T18:35:07.411201Z","iopub.status.busy":"2022-08-06T18:35:07.409864Z","iopub.status.idle":"2022-08-06T18:35:15.489731Z","shell.execute_reply":"2022-08-06T18:35:15.489143Z"},"papermill":{"duration":8.77077,"end_time":"2022-08-06T18:35:15.489885","exception":false,"start_time":"2022-08-06T18:35:06.719115","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import shap\n\n# print the JS visualization code to the notebook\nshap.initjs()","metadata":{"_kg_hide-output":true,"execution":{"iopub.execute_input":"2022-08-06T18:35:16.677246Z","iopub.status.busy":"2022-08-06T18:35:16.676255Z","iopub.status.idle":"2022-08-06T18:35:21.784822Z","shell.execute_reply":"2022-08-06T18:35:21.785358Z"},"papermill":{"duration":5.705664,"end_time":"2022-08-06T18:35:21.785564","exception":false,"start_time":"2022-08-06T18:35:16.079900","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\ndef f(inp):\n    inp = pd.DataFrame(inp,columns=list(df.columns))\n    X = get_X_from_features(inp.drop(columns=['target'],axis=1),cont_features,cat_features)\n    return model.predict(X).flatten()\n\nexplainer = shap.KernelExplainer(f, df1.iloc[:50,:])\nshap_values = explainer.shap_values(df1.iloc[299,:], nsamples=500)\nshap.force_plot(explainer.expected_value, shap_values, df1.iloc[299,:])","metadata":{"execution":{"iopub.execute_input":"2022-08-06T18:35:23.012772Z","iopub.status.busy":"2022-08-06T18:35:23.011732Z","iopub.status.idle":"2022-08-06T18:35:42.351275Z","shell.execute_reply":"2022-08-06T18:35:42.350120Z","shell.execute_reply.started":"2022-08-06T18:11:03.738658Z"},"papermill":{"duration":19.959649,"end_time":"2022-08-06T18:35:42.351592","exception":false,"start_time":"2022-08-06T18:35:22.391943","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"shap_values50 = explainer.shap_values(df1.iloc[280:330,:], nsamples=500)\nfig = shap.summary_plot(shap_values50, df1.iloc[280:330,:],show=False)\nplt.savefig('shap.png', dpi=600, bbox_inches='tight')\nplt.show()","metadata":{"execution":{"iopub.execute_input":"2022-08-06T18:35:43.627659Z","iopub.status.busy":"2022-08-06T18:35:43.626840Z","iopub.status.idle":"2022-08-06T18:44:05.487308Z","shell.execute_reply":"2022-08-06T18:44:05.486672Z","shell.execute_reply.started":"2022-08-06T18:11:37.828699Z"},"papermill":{"duration":502.47455,"end_time":"2022-08-06T18:44:05.487463","exception":false,"start_time":"2022-08-06T18:35:43.012913","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"TabTransformers are highly robust against missing and noisy data and provide better interpretability . By replacing the final MLP block with gated Multilayer perceptron , GatedTabTranformers are able to achieve high accuracy in binary classification tasks.\n\n### References\n\nhttps://arxiv.org/pdf/2201.00199.pdf\n\nhttps://arxiv.org/pdf/2012.06678.pdf\n\nhttps://www.tensorflow.org/\n\nhttps://flax.readthedocs.io/\n\nPytorch Implementation : https://github.com/radi-cho/GatedTabTransformer","metadata":{}}]}