{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-09-08T00:19:37.740895Z","iopub.execute_input":"2023-09-08T00:19:37.741281Z","iopub.status.idle":"2023-09-08T00:19:37.793825Z","shell.execute_reply.started":"2023-09-08T00:19:37.741253Z","shell.execute_reply":"2023-09-08T00:19:37.792866Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 1. Define constants\n\n## a. Imports","metadata":{}},{"cell_type":"code","source":"import os\nimport json\n\nimport keras\nimport keras.layers as layers\n\nimport datetime\nimport math\n\nimport numpy as np\nimport pandas as pd\nimport tensorflow as tf\n\nimport matplotlib.pyplot as plt\n\nimport pyarrow\nfrom pyarrow import parquet as pq\n\nfrom tqdm import tqdm","metadata":{"execution":{"iopub.status.busy":"2023-09-08T00:18:43.32385Z","iopub.execute_input":"2023-09-08T00:18:43.324467Z","iopub.status.idle":"2023-09-08T00:18:52.573208Z","shell.execute_reply.started":"2023-09-08T00:18:43.324437Z","shell.execute_reply":"2023-09-08T00:18:52.571978Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## b. File paths","metadata":{}},{"cell_type":"code","source":"TRAINING_CSV = '/kaggle/input/asl-fingerspelling/train.csv'\nTRAIN_FOLDER = '/kaggle/input/asl-fingerspelling/train_landmarks'\nSUPPLEMENTAL_CSV = '/kaggle/input/asl-fingerspelling/supplemental_metadata.csv'\nSUPPLEMENTAL_FOLDER = '/kaggle/input/asl-fingerspelling/supplemental_landmarks'\n\nOUTPUT_MODEL = '/kaggle/working/model.tflite'\nMAPPING_FILE = '/kaggle/input/asl-fingerspelling/character_to_prediction_index.json'\n\nPREPROCESSED_DIR = '/kaggle/working/preprocessed'","metadata":{"execution":{"iopub.status.busy":"2023-09-08T00:22:36.896526Z","iopub.execute_input":"2023-09-08T00:22:36.89699Z","iopub.status.idle":"2023-09-08T00:22:36.903368Z","shell.execute_reply.started":"2023-09-08T00:22:36.896955Z","shell.execute_reply":"2023-09-08T00:22:36.902191Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## c. Model constants","metadata":{}},{"cell_type":"code","source":"FRAME_COUNT = 512\n\nPAD = 'P'\nSTART = '<'\nEND = '>'\n\nEPOCHS=50\nBATCH_SIZE=192","metadata":{"execution":{"iopub.status.busy":"2023-09-08T00:19:06.841848Z","iopub.execute_input":"2023-09-08T00:19:06.842467Z","iopub.status.idle":"2023-09-08T00:19:06.847948Z","shell.execute_reply.started":"2023-09-08T00:19:06.842425Z","shell.execute_reply":"2023-09-08T00:19:06.846893Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## d. Competition constants\n(i.e., parquet features, character mappings)","metadata":{}},{"cell_type":"code","source":"PARQUET_FEATURE_LIST = [\n    *[\n        f'{coord}_{hand}_{i}'\n        for hand in ['left_hand', 'right_hand']\n        for coord in ['x', 'y', 'z']\n        for i in range(21)\n    ],\n    'frame',\n    'sequence_id'\n]\n\nPARQUET_LH_FEATURES = [i for i in range(0, 63)]\nPARQUET_RH_FEATURES = [i for i in range(63, 126)]\n\nPARQUET_X = slice(0, 21)\nPARQUET_Y = slice(21, 42)\nPARQUET_Z = slice(42, 63)\n\ndef load_charmap():\n    with open(MAPPING_FILE, 'r') as mapping_file:\n        forward = json.load(mapping_file)\n\n    # See `https://www.kaggle.com/code/gusthema/asl-fingerspelling-recognition-w-tensorflow\n    # #Load-character_to_prediction-json-file`\n    forward[PAD] = 59\n    forward[START] = 60\n    forward[END] = 61\n\n    backward = {\n        index: char\n        for char, index in forward.items()\n    }\n\n    return forward, backward\n\nCHAR_TO_IDX, IDX_TO_CHAR = load_charmap()\n\nMAPPING_LOOKUP_TABLE = tf.lookup.StaticHashTable(\n    initializer=tf.lookup.KeyValueTensorInitializer(\n        keys=list(CHAR_TO_IDX.keys()),\n        values=list(CHAR_TO_IDX.values())\n    ),\n    default_value=tf.constant(-1),\n    name=\"class_weight\"\n)","metadata":{"execution":{"iopub.status.busy":"2023-09-08T00:20:33.218743Z","iopub.execute_input":"2023-09-08T00:20:33.219219Z","iopub.status.idle":"2023-09-08T00:20:33.339025Z","shell.execute_reply.started":"2023-09-08T00:20:33.21918Z","shell.execute_reply":"2023-09-08T00:20:33.338081Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## e. MediaPipe features\nThe values shown here are copied from the diagram [shown in the MediaPipe documentation](https://developers.google.com/mediapipe/solutions/vision/hand_landmarker#models). \n\nFor reference\n\n![Hand joints](https://upload.wikimedia.org/wikipedia/commons/b/b6/814_Radiograph_of_Hand.jpg)\n\n* `CMC`: carpometacarpal joint\n* `MCP`: metacarpophalangeal joint\n* `IP`: interphalangeal joint\n* `PIP/DIP`: proximal / distal interphalangeal joint\n* `TIP`: fingertip :)","metadata":{}},{"cell_type":"code","source":"WRIST = 0\n\nTHUMB_CMC, THUMB_MCP, THUMB_IP, THUMB_TIP = 1, 2, 3, 4\nTHUMB_IN = [THUMB_CMC, THUMB_MCP, THUMB_IP, THUMB_TIP]\n\nINDEX_FINGER_MCP, INDEX_FINGER_PIP, INDEX_FINGER_DIP, INDEX_FINGER_TIP = 5, 6, 7, 8\nINDEX_IN = [INDEX_FINGER_MCP, INDEX_FINGER_PIP, INDEX_FINGER_DIP, INDEX_FINGER_TIP]\n\nMIDDLE_FINGER_MCP, MIDDLE_FINGER_PIP, MIDDLE_FINGER_DIP, MIDDLE_FINGER_TIP = 9, 10, 11, 12\nMIDDLE_IN = [MIDDLE_FINGER_MCP, MIDDLE_FINGER_PIP, MIDDLE_FINGER_DIP, MIDDLE_FINGER_TIP]\n\nRING_FINGER_MCP, RING_FINGER_PIP, RING_FINGER_DIP, RING_FINGER_TIP = 13, 14, 15, 16\nRING_IN = [RING_FINGER_MCP, RING_FINGER_PIP, RING_FINGER_DIP, RING_FINGER_TIP]\n\nPINKY_MCP, PINKY_PIP, PINKY_DIP, PINKY_TIP = 17, 18, 19, 20\nPINKY_IN = [PINKY_MCP, PINKY_PIP, PINKY_DIP, PINKY_TIP]\n\nINBOUND_FEATURE_GROUPS = [THUMB_IN, INDEX_IN, MIDDLE_IN, RING_IN, PINKY_IN]\nINBOUND_FEATURES = [*THUMB_IN, *INDEX_IN, *MIDDLE_IN, *RING_IN, *PINKY_IN]\nINBOUND_FEATURE_COUNT = 21","metadata":{"execution":{"iopub.status.busy":"2023-09-08T00:20:39.775005Z","iopub.execute_input":"2023-09-08T00:20:39.775389Z","iopub.status.idle":"2023-09-08T00:20:39.785156Z","shell.execute_reply.started":"2023-09-08T00:20:39.77536Z","shell.execute_reply":"2023-09-08T00:20:39.78384Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## f. Internal feature labels\n\nThe constants used here are to make the rest of the source code more readable.\n\n* The `WRIST_POS_[X/Y/Z]`, `INDEX_MCP_[X/Y/Z]`, and `PINKY_MCP_[X/Y/Z]` features are provided to give the position of the hand in 3D space. Initially, only the wrist position was provided, but providing the base of the index finger and pinky allows the model to see how the hand moves, which may be useful for certain signing behaviors (such as signing double letters).\n* The `WRIST_POS_[dX/dY/dZ]` features (and corresponding features for the index and pinky fingers) are provided so that the model can see the motion of these three points in space between frames.","metadata":{}},{"cell_type":"code","source":"WRIST_POS_X, WRIST_POS_Y, WRIST_POS_Z = 0, 1, 2\nWRIST_POS_dX, WRIST_POS_dY, WRIST_POS_dZ = 3, 4, 5\n\nINDEX_MCP_X, INDEX_MCP_Y, INDEX_MCP_Z = 6, 7, 8\nINDEX_MCP_dX, INDEX_MCP_dY, INDEX_MCP_dZ = 9, 10, 11\n\nPINKY_MCP_X, PINKY_MCP_Y, PINKY_MCP_Z = 12, 13, 14\nPINKY_MCP_dX, PINKY_MCP_dY, PINKY_MCP_dZ = 15, 16, 17\n\nHAND_POS = [\n    WRIST_POS_X, WRIST_POS_Y, WRIST_POS_Z,\n    INDEX_MCP_X, INDEX_MCP_Y, INDEX_MCP_Z,\n    PINKY_MCP_X, PINKY_MCP_Y, PINKY_MCP_Z\n]\n\nHAND_POS_DELTA = [\n    WRIST_POS_dX, WRIST_POS_dY, WRIST_POS_dZ,\n    INDEX_MCP_dX, INDEX_MCP_dY, INDEX_MCP_dZ,\n    PINKY_MCP_dX, PINKY_MCP_dY, PINKY_MCP_dZ\n]","metadata":{"execution":{"iopub.status.busy":"2023-09-08T00:20:43.141422Z","iopub.execute_input":"2023-09-08T00:20:43.142328Z","iopub.status.idle":"2023-09-08T00:20:43.15122Z","shell.execute_reply.started":"2023-09-08T00:20:43.142272Z","shell.execute_reply":"2023-09-08T00:20:43.149953Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Note**: Rather than using the Cartesian coordinates for the position of each of the fingers, we use angles here. To understand the `_PITCH` and `_YAW` features below, first place your hand palm-up on a desk. The pitch is the vertical angle of each finger, and the yaw is the horizontal angle. Only the first joint of each finger (closest to the palm) can move in two axes; the remaining joints only move along one axis, so their articulation is represented using the `_TILT` features.\n\n![Hand joints](https://upload.wikimedia.org/wikipedia/commons/a/ab/Scheme_human_hand_bones-en.svg)\n\n* `PP`: proximal phalanx\n* `IP`: intermediate phalanx\n* `DP`: distal phalanx","metadata":{}},{"cell_type":"code","source":"THUMB_PITCH, THUMB_YAW, THUMB_PP_TILT, THUMB_DP_TILT = 18, 19, 20, 21\nTHUMB_OUT = [THUMB_PITCH, THUMB_YAW, THUMB_PP_TILT, THUMB_DP_TILT]\n\nINDEX_PITCH, INDEX_YAW, INDEX_IP_TILT, INDEX_DP_TILT = 22, 23, 24, 25\nINDEX_OUT = [INDEX_PITCH, INDEX_YAW, INDEX_IP_TILT, INDEX_DP_TILT]\n\nMIDDLE_PITCH, MIDDLE_YAW, MIDDLE_IP_TILT, MIDDLE_DP_TILT = 26, 27, 28, 29\nMIDDLE_OUT = [MIDDLE_PITCH, MIDDLE_YAW, MIDDLE_IP_TILT, MIDDLE_DP_TILT]\n\nRING_PITCH, RING_YAW, RING_IP_TILT, RING_DP_TILT = 30, 31, 32, 33\nRING_OUT = [RING_PITCH, RING_YAW, RING_IP_TILT, RING_DP_TILT]\n\nPINKY_PITCH, PINKY_YAW, PINKY_IP_TILT, PINKY_DP_TILT = 34, 35, 36, 37\nPINKY_OUT = [PINKY_PITCH, PINKY_YAW, PINKY_IP_TILT, PINKY_DP_TILT]\n\nOUTBOUND_POLAR_FEATURES = [*THUMB_OUT, *INDEX_OUT, *MIDDLE_OUT, *RING_OUT, *PINKY_OUT]\nOUTBOUND_START_INDEX = THUMB_PITCH\nPOLAR_COUNT = len(OUTBOUND_POLAR_FEATURES)\nOUTBOUND_FEATURE_GROUPS = [THUMB_OUT, INDEX_OUT, MIDDLE_OUT, RING_OUT, PINKY_OUT]\nOUTBOUND_FEATURE_COUNT = 38","metadata":{"execution":{"iopub.status.busy":"2023-09-08T00:20:55.48267Z","iopub.execute_input":"2023-09-08T00:20:55.483061Z","iopub.status.idle":"2023-09-08T00:20:55.493005Z","shell.execute_reply.started":"2023-09-08T00:20:55.483033Z","shell.execute_reply":"2023-09-08T00:20:55.491868Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 2. Implement Transformer model\n\nThe Transformer model here is copied from the [ASL Fingerspelling Recognition with TensorFlow](https://www.kaggle.com/code/gusthema/asl-fingerspelling-recognition-w-tensorflow) guide. There's not much commentary for (or many comments in) the code below, as this was the portion of the code that I touched the least.","metadata":{}},{"cell_type":"code","source":"class TokenEmbedding(layers.Layer):\n    def __init__(self, num_vocab=1000, maxlen=100, num_hid=64):\n        super().__init__()\n        self.emb = layers.Embedding(num_vocab, num_hid)\n        self.pos_emb = layers.Embedding(input_dim=maxlen, output_dim=num_hid)\n\n    def call(self, x):\n        maxlen = tf.shape(x)[-1]\n        x = self.emb(x)\n        positions = tf.range(start=0, limit=maxlen, delta=1)\n        positions = self.pos_emb(positions)\n        return x + positions\n\nclass LandmarkEmbedding(layers.Layer):\n    def __init__(self, num_hid=64, maxlen=100):\n        super().__init__()\n        self.conv1 = tf.keras.layers.Conv1D(\n            num_hid, 11, strides=2, padding='same', activation='relu'\n        )\n        self.conv2 = tf.keras.layers.Conv1D(\n            num_hid, 11, strides=2, padding='same', activation='relu'\n        )\n        self.conv3 = tf.keras.layers.Conv1D(\n            num_hid, 11, strides=2, padding='same', activation='relu'\n        )\n        self.pos_emb = layers.Embedding(input_dim=maxlen, output_dim=num_hid)\n\n    def call(self, x):\n        x = self.conv1(x)\n        x = self.conv2(x)\n        return self.conv3(x)\n\nclass TransformerEncoder(layers.Layer):\n    def __init__(self, embed_dim, num_heads, feed_forward_dim, rate=0.1):\n        super().__init__()\n        self.att = layers.MultiHeadAttention(num_heads=num_heads, key_dim=embed_dim)\n        self.ffn = keras.Sequential([\n            layers.Dense(feed_forward_dim, activation='relu'),\n            layers.Dense(embed_dim),\n        ])\n\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_out = self.att(inputs, inputs)\n        attn_out = self.dropout1(attn_out, training=training)\n        out1 = self.layernorm1(inputs + attn_out)\n\n        ffn_out = self.ffn(out1)\n        ffn_out = self.dropout2(ffn_out, training=training)\n        return self.layernorm2(out1 + ffn_out)\n\nclass TransformerDecoder(layers.Layer):\n    def __init__(self, embed_dim, num_heads, feed_forward_dim, dropout_rate=0.1):\n        super().__init__()\n        self.layernorm1 = layers.LayerNormalization(epsilon=1e-6)\n        self.layernorm2 = layers.LayerNormalization(epsilon=1e-6)\n        self.layernorm3 = layers.LayerNormalization(epsilon=1e-6)\n        self.self_att = layers.MultiHeadAttention(\n            num_heads=num_heads, key_dim=embed_dim\n        )\n        self.enc_att = layers.MultiHeadAttention(num_heads=num_heads, key_dim=embed_dim)\n        self.self_dropout = layers.Dropout(0.5)\n        self.enc_dropout = layers.Dropout(0.1)\n        self.ffn_dropout = layers.Dropout(0.1)\n        self.ffn = keras.Sequential([\n            layers.Dense(feed_forward_dim, activation=\"relu\"),\n            layers.Dense(embed_dim),\n        ])\n\n    def causal_attention_mask(self, batch_size, n_dest, n_src, dtype):\n        i = tf.range(n_dest)[:, None]\n        j = tf.range(n_src)\n        m = i >= j - n_src + n_dest\n        mask = tf.cast(m, dtype)\n        mask = tf.reshape(mask, [1, n_dest, n_src])\n        mult = tf.concat(\n            [batch_size[..., tf.newaxis], tf.constant([1, 1], dtype=tf.int32)], 0\n        )\n        return tf.tile(mask, mult)\n\n    def call(self, enc_out, target, training):\n        input_shape = tf.shape(target)\n        batch_size = input_shape[0]\n        seq_len = input_shape[1]\n        causal_mask = self.causal_attention_mask(batch_size, seq_len, seq_len, tf.bool)\n\n        target_att = self.self_att(target, target, attention_mask=causal_mask)\n        target_norm = self.layernorm1(target + self.self_dropout(target_att, training=training))\n\n        enc_out = self.enc_att(target_norm, enc_out)\n        enc_out_norm = self.layernorm2(self.enc_dropout(enc_out, training=training) + target_norm)\n\n        ffn_out = self.ffn(enc_out_norm)\n        ffn_out_norm = self.layernorm3(enc_out_norm + self.ffn_dropout(ffn_out, training=training))\n\n        return ffn_out_norm\n\nclass Transformer(keras.Model):\n    def __init__(\n            self,\n            num_hid=64,\n            num_head=2,\n            num_feed_forward=128,\n            source_maxlen=100,\n            target_maxlen=100,\n            num_layers_enc=4,\n            num_layers_dec=1,\n            num_classes=60,\n    ):\n        super().__init__()\n        self.loss_metric = keras.metrics.Mean(name=\"loss\")\n        self.acc_metric = keras.metrics.Mean(name=\"edit_dist\")\n        self.num_layers_enc = num_layers_enc\n        self.num_layers_dec = num_layers_dec\n        self.target_maxlen = target_maxlen\n        self.num_classes = num_classes\n\n        self.enc_input = LandmarkEmbedding(num_hid=num_hid, maxlen=source_maxlen)\n        self.dec_input = TokenEmbedding(\n            num_vocab=num_classes, maxlen=target_maxlen, num_hid=num_hid\n        )\n\n        self.encoder = keras.Sequential(\n            [self.enc_input]\n            + [\n                TransformerEncoder(num_hid, num_head, num_feed_forward)\n                for _ in range(num_layers_enc)\n            ]\n        )\n\n        for i in range(num_layers_dec):\n            setattr(\n                self,\n                f\"dec_layer_{i}\",\n                TransformerDecoder(num_hid, num_head, num_feed_forward),\n            )\n\n        self.classifier = layers.Dense(num_classes)\n\n    def decode(self, enc_out, target, training):\n        y = self.dec_input(target)\n        for i in range(self.num_layers_dec):\n            y = getattr(self, f\"dec_layer_{i}\")(enc_out, y, training)\n        return y\n\n    def call(self, inputs, training):\n        source = inputs[0]\n        target = inputs[1]\n        x = self.encoder(source, training)\n        y = self.decode(x, target, training)\n        return self.classifier(y)\n\n    @property\n    def metrics(self):\n        return [self.loss_metric]\n\n    def train_step(self, batch):\n        \"\"\"Processes one batch inside model.fit().\"\"\"\n        source = batch[0]\n        target = batch[1]\n\n        input_shape = tf.shape(target)\n        batch_size = input_shape[0]\n\n        dec_input = target[:, :-1]\n        dec_target = target[:, 1:]\n        with tf.GradientTape() as tape:\n            preds = self([source, dec_input])\n            one_hot = tf.one_hot(dec_target, depth=self.num_classes)\n            mask = tf.math.logical_not(tf.math.equal(dec_target, CHAR_TO_IDX[PAD]))\n            loss = self.compiled_loss(one_hot, preds, sample_weight=mask)\n        trainable_vars = self.trainable_variables\n        gradients = tape.gradient(loss, trainable_vars)\n        self.optimizer.apply_gradients(zip(gradients, trainable_vars))\n        # Computes the Levenshtein distance between sequences since the evaluation\n        # metric for this contest is the normalized total levenshtein distance.\n        edit_dist = tf.edit_distance(tf.sparse.from_dense(target),\n                                     tf.sparse.from_dense(tf.cast(tf.argmax(preds, axis=1), tf.int32)))\n        edit_dist = tf.reduce_mean(edit_dist)\n        self.acc_metric.update_state(edit_dist)\n        self.loss_metric.update_state(loss)\n        return {\"loss\": self.loss_metric.result(), \"edit_dist\": self.acc_metric.result()}\n\n    def test_step(self, batch):\n        source = batch[0]\n        target = batch[1]\n\n        input_shape = tf.shape(target)\n        batch_size = input_shape[0]\n\n        dec_input = target[:, :-1]\n        dec_target = target[:, 1:]\n        preds = self([source, dec_input])\n        one_hot = tf.one_hot(dec_target, depth=self.num_classes)\n        mask = tf.math.logical_not(tf.math.equal(dec_target, CHAR_TO_IDX[PAD]))\n        loss = self.compiled_loss(one_hot, preds, sample_weight=mask)\n        # Computes the Levenshtein distance between sequences since the evaluation\n        # metric for this contest is the normalized total levenshtein distance.\n        edit_dist = tf.edit_distance(tf.sparse.from_dense(target),\n                                     tf.sparse.from_dense(tf.cast(tf.argmax(preds, axis=1), tf.int32)))\n        edit_dist = tf.reduce_mean(edit_dist)\n        self.acc_metric.update_state(edit_dist)\n        self.loss_metric.update_state(loss)\n        return {\"loss\": self.loss_metric.result(), \"edit_dist\": self.acc_metric.result()}\n\n    def generate(self, source, target_start_token_idx):\n        \"\"\"Performs inference over one batch of inputs using greedy decoding.\"\"\"\n        bs = tf.shape(source)[0]\n        enc = self.encoder(source, training=False)\n        dec_input = tf.ones((bs, 1), dtype=tf.int32) * target_start_token_idx\n        dec_logits = []\n        for i in range(self.target_maxlen - 1):\n            dec_out = self.decode(enc, dec_input, training=False)\n            logits = self.classifier(dec_out)\n            logits = tf.argmax(logits, axis=-1, output_type=tf.int32)\n            last_logit = logits[:, -1][..., tf.newaxis]\n            dec_logits.append(last_logit)\n            dec_input = tf.concat([dec_input, last_logit], axis=-1)\n        return dec_input\n\nclass DisplayOutputs(keras.callbacks.Callback):\n    def __init__(\n        self, batch, idx_to_token, target_start_token_idx=60, target_end_token_idx=61\n    ):\n        \"\"\"Displays a batch of outputs after every 4 epoch\n\n        Args:\n            batch: A test batch\n            idx_to_token: A List containing the vocabulary tokens corresponding to their indices\n            target_start_token_idx: A start token index in the target vocabulary\n            target_end_token_idx: An end token index in the target vocabulary\n        \"\"\"\n        self.batch = batch\n        self.target_start_token_idx = target_start_token_idx\n        self.target_end_token_idx = target_end_token_idx\n        self.idx_to_char = idx_to_token\n\n    def on_epoch_end(self, epoch, logs=None):\n        if epoch % 4 != 0:\n            return\n        source = self.batch[0]\n        target = self.batch[1].numpy()\n        bs = tf.shape(source)[0]\n        preds = self.model.generate(source, self.target_start_token_idx)\n        preds = preds.numpy()\n        for i in range(bs):\n            target_text = \"\".join([self.idx_to_char[_] for _ in target[i, :]])\n            prediction = \"\"\n            for idx in preds[i, :]:\n                prediction += self.idx_to_char[idx]\n                if idx == self.target_end_token_idx:\n                    break\n            print(f\"target:     {target_text.replace('-','')}\")\n            print(f\"prediction: {prediction}\\n\")","metadata":{"execution":{"iopub.status.busy":"2023-09-08T00:21:00.927214Z","iopub.execute_input":"2023-09-08T00:21:00.927613Z","iopub.status.idle":"2023-09-08T00:21:00.986951Z","shell.execute_reply.started":"2023-09-08T00:21:00.927581Z","shell.execute_reply":"2023-09-08T00:21:00.985637Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 3. Load / prepare dataset","metadata":{}},{"cell_type":"code","source":"def load_parquet(filename) -> pd.DataFrame:\n    parquet_array = pq.read_table(filename, columns=PARQUET_FEATURE_LIST, memory_map=True).to_pandas()\n    return parquet_array\n\ndef preload(source, source_folder):\n    labels = pd.read_csv(source)\n\n    for file in tqdm(labels.file_id.unique()):\n        pq_data = load_parquet(os.path.join(source_folder, f'{file}.parquet'))\n        entries = labels.loc[labels['file_id'] == file]\n\n        with tf.io.TFRecordWriter(os.path.join(PREPROCESSED_DIR, f'{file}.tfrecord')) as tfrec:\n            for index, row in entries.iterrows():\n                seq_id = row['sequence_id']\n                frames = pq_data[pq_data.index == seq_id]\n                frames = frames.to_numpy()\n\n                # Code introduced from `https://www.kaggle.com/code/gusthema/asl-fingerspelling\n                # -recognition-w-tensorflow?scriptVersionId=135798259&cellId=34`; logical\n                # explanation (in comments below) derived independently\n\n                # Compute how many frames have zero missing features from the left hand\n                lh_full_frames = np.sum(\n                    # Compute how many features are missing in each frame\n                    np.sum(\n                        # Compute, for each feature in each frame, whether that feature is NaN (missing)\n                        np.isnan(frames[:, PARQUET_LH_FEATURES]),\n                        axis=1\n                    ) == 0\n                )\n\n                # Do the same for the right hand, then pick the hand with more frames\n                rh_full_frames = np.sum(np.sum(np.isnan(frames[:, PARQUET_RH_FEATURES]), axis=1) == 0)\n                most_full_frames = max(lh_full_frames, rh_full_frames)\n\n                phrase = row['phrase']\n\n                # If there are at least twice as many frames as there are characters, copy the data\n                # into a TFRecord\n                if most_full_frames > 2 * len(phrase):\n                    features = {\n                        label: tf.train.Feature(float_list=tf.train.FloatList(value=frames[:, i]))\n                        for i, label in enumerate(PARQUET_FEATURE_LIST[:-1])\n                    }\n\n                    features['phrase'] = tf.train.Feature(\n                        bytes_list=tf.train.BytesList(value=[bytes(phrase, 'utf-8')])\n                    )\n\n                    record_as_bytes = tf.train.Example(\n                        features=tf.train.Features(feature=features)\n                    ).SerializeToString()\n                    tfrec.write(record_as_bytes)\n\ndef prepare_data():\n    preload(TRAINING_CSV, TRAIN_FOLDER)\n    preload(SUPPLEMENTAL_CSV, SUPPLEMENTAL_FOLDER)","metadata":{"execution":{"iopub.status.busy":"2023-09-08T00:21:07.061327Z","iopub.execute_input":"2023-09-08T00:21:07.0617Z","iopub.status.idle":"2023-09-08T00:21:07.076449Z","shell.execute_reply.started":"2023-09-08T00:21:07.061672Z","shell.execute_reply":"2023-09-08T00:21:07.075117Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Optionally, you can disable `run_preparation` below for subsequent runs of the model.","metadata":{}},{"cell_type":"code","source":"!mkdir -p '/kaggle/working/preprocessed'","metadata":{"execution":{"iopub.status.busy":"2023-09-08T00:23:41.370781Z","iopub.execute_input":"2023-09-08T00:23:41.371231Z","iopub.status.idle":"2023-09-08T00:23:42.501894Z","shell.execute_reply.started":"2023-09-08T00:23:41.371196Z","shell.execute_reply":"2023-09-08T00:23:42.500194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"run_preparation = True\nif run_preparation:\n    prepare_data()","metadata":{"execution":{"iopub.status.busy":"2023-09-08T00:23:48.110563Z","iopub.execute_input":"2023-09-08T00:23:48.111031Z","iopub.status.idle":"2023-09-08T00:42:13.247917Z","shell.execute_reply.started":"2023-09-08T00:23:48.110994Z","shell.execute_reply":"2023-09-08T00:42:13.245995Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 4. Convert hand coordinates\n\n`hand_rel_coords_tf()` converts the hand coordinates provided in 3D space to a relative coordinate system centered on the palm of the hand. This allows us to consider the shape that the fingers are forming entirely independently of the hand movement. Note that we still factor the position of the hand into the model, as doing so allows us to handle certain gestures (like sliding the hand sideways) which commonly indicate a repeated letter.\n\nThe coordinate system in use here is pretty straightforward. First, be sure to check the diagram at the link below:\nhttps://developers.google.com/mediapipe/solutions/vision/hand_landmarker#models\n\n![Diagram from above link](https://developers.google.com/static/mediapipe/images/solutions/hand-landmarks.png)\n\nWe can form a triangle in 3D space from points `0` (A), `5` (B), and `17` (C) in the diagram (respectively, the wrist, the index finger's knuckle, and the pinky's knuckle). We can then multiply vectors `AB x AC` to get a normal vector `N` pointing inwards (i.e. straight up if you lay your hand palm-side-up on the desk). They call it the right-hand rule for a reason :) (Left hands get mirrored at a separate stage in the data processing pipeline.)\n\nWe scale / rotate all of the coordinates so that the centroid of the triangle is at `(0, 0, 0)`, `N = (0, 0, 1)`, and `A = (0, -1, 0)`. From there, for each finger, we perform three separate steps:\n\n1. For the segment of the finger closest to the palm, we calculate articulation in two dimensions. Lay your hand flat on your desk. The PITCH of a finger is its vertical rotation (which increases as you curl your fingers upward). The YAW of a finger is its lateral rotation (side-to-side articulation).\n2. For the middle segment of the finger, we calculate the articulation in only one dimension: the angle of the middle segment relative to the inner segment (the one closest to the palm). For example, if the finger is completely straight, this segment's articulation will be 0 radians; if the middle part of the finger is bent at a 90 degree angle, this segment's articulation will be pi/4 radians.\n3. We calculate the same thing as step 2 for the angle of the fingertip relative to the middle segment of the finger.\n\nOnce all of these features are calculated, we have a standardized hand shape representation that can be fed to the machine learning model!\n\n**Note on functions**:\n\n* The `arcsin(y)`, `arccos(x)`, and `cross(a, b)` functions are implemented below because TFLite does not natively support these functions.\n* The `finger_angles_tf(finger)` function handles the angle data for an individual finger.\n* The `hand_rel_coords_tf(hand, prior)` function handles the entire hand's data.","metadata":{}},{"cell_type":"code","source":"@tf.function\ndef arcsin(y):\n    return tf.math.atan2(y, tf.math.sqrt(1 - (y * y)))\n\n@tf.function\ndef arccos(x):\n    return tf.math.atan2(tf.math.sqrt(1 - (x * x)), x)\n\n@tf.function(input_signature=(tf.TensorSpec(shape=[3], dtype=tf.float32),\n             tf.TensorSpec(shape=[3], dtype=tf.float32)))\ndef cross(a: tf.Tensor, b: tf.Tensor):\n    return tf.stack([\n        a[1] * b[2] - a[2] * b[1],\n        a[2] * b[0] - a[0] * b[2],\n        a[0] * b[1] - a[1] * b[0]\n    ])\n\n@tf.function(input_signature=(tf.TensorSpec(shape=[None, None], dtype=tf.float32),), experimental_relax_shapes=True)\ndef finger_angles_tf(finger: tf.Tensor):\n    print('tracing finger_angles_tf')\n    IN_KNUCKLE, IN_PIP, IN_DIP, IN_TIP = 0, 1, 2, 3\n    pip_rel = finger[IN_PIP] - finger[IN_KNUCKLE]\n    pip_rel /= tf.linalg.norm(pip_rel)\n\n    dip_rel = finger[IN_DIP] - finger[IN_PIP]\n    dip_rel /= tf.linalg.norm(dip_rel)\n\n    tip_rel = finger[IN_TIP] - finger[IN_DIP]\n    tip_rel /= tf.linalg.norm(tip_rel)\n\n    return tf.stack([\n        arcsin(pip_rel[2]),                                                              # Pitch of proximal phalanx\n        tf.math.atan2(pip_rel[1], pip_rel[0]),                                           # Yaw of proximal phalanx\n        arccos(tf.clip_by_value(tf.reduce_sum(tf.multiply(pip_rel, dip_rel)), -1, 1)),   # Angle of medial phalanx\n        arccos(tf.clip_by_value(tf.reduce_sum(tf.multiply(dip_rel, tip_rel)), -1, 1)),   # Angle of distal phalanx\n    ])\n\n@tf.function(input_signature=(tf.TensorSpec(shape=[None, None], dtype=tf.float32), tf.TensorSpec(shape=[None, None], dtype=tf.float32)),\n             experimental_relax_shapes=True)\ndef hand_rel_coords_tf(hand: tf.Tensor, prior: tf.Tensor):\n    print('tracing hand_rel_coords_tf')\n    if tf.reduce_sum(tf.cast(tf.math.is_finite(hand), dtype=tf.int32)) < INBOUND_FEATURE_COUNT:\n        return tf.zeros([OUTBOUND_FEATURE_COUNT])\n\n    if prior is None or tf.reduce_sum(tf.cast(tf.math.is_finite(prior), dtype=tf.int32)) < INBOUND_FEATURE_COUNT:\n        prior = hand\n\n    palm_triangle = tf.concat([\n        hand[WRIST],\n        hand[INDEX_FINGER_MCP],\n        hand[PINKY_MCP]\n    ], axis=0)\n\n    palm_centroid = tf.reduce_mean(palm_triangle, axis=0)\n    centered = hand - palm_centroid\n\n    a = centered[WRIST]\n    b = centered[INDEX_FINGER_MCP]\n    c = centered[PINKY_MCP]\n\n    j = b - a\n    j /= tf.linalg.norm(j)\n\n    k = cross(b - a, c - a)\n    k /= tf.linalg.norm(k)\n\n    i = cross(j, k)\n    i /= tf.linalg.norm(i)\n\n    basis = tf.concat([\n        tf.reshape(i, [1, 3]),\n        tf.reshape(j, [1, 3]),\n        tf.reshape(k, [1, 3]),\n    ], axis=0)\n\n    # If it's not invertible\n    if tf.linalg.det(basis) == 0:\n        return tf.zeros([OUTBOUND_FEATURE_COUNT])\n\n    basis_inv = tf.linalg.inv(basis)\n\n    # Matrix-multiply all of the centered data points to get their positions in the new coordinate space\n    transformed = tf.transpose(tf.matmul(basis_inv, tf.transpose(centered)))\n\n    # 5 fingers, four data points per finger, three dimensions per data point (60 features total, excluding wrist)\n    hand_pos_data = tf.reshape(transformed[1:], [5, 4, 3])\n    finger_angles = tf.map_fn(finger_angles_tf, hand_pos_data)\n\n    output = tf.concat([\n        tf.reshape(centered[WRIST], [-1]),                                       # Wrist\n        tf.reshape(centered[WRIST] - prior[WRIST], [-1]),                        # Wrist movement\n        tf.reshape(centered[INDEX_FINGER_MCP], [-1]),                            # Base of index finger\n        tf.reshape(centered[INDEX_FINGER_MCP] - prior[INDEX_FINGER_MCP], [-1]),  # Base of index finger's movement\n        tf.reshape(centered[PINKY_MCP], [-1]),                                   # Base of pinky\n        tf.reshape(centered[PINKY_MCP] - prior[PINKY_MCP], [-1]),                # Base of pinky's movement\n        tf.reshape(finger_angles, [-1]),\n    ], axis=0)\n\n    return tf.reshape(output, [OUTBOUND_FEATURE_COUNT])                         # OUTBOUND_FEATURE_COUNT = 38","metadata":{"execution":{"iopub.status.busy":"2023-09-08T01:32:08.119838Z","iopub.execute_input":"2023-09-08T01:32:08.120487Z","iopub.status.idle":"2023-09-08T01:32:08.18791Z","shell.execute_reply.started":"2023-09-08T01:32:08.120453Z","shell.execute_reply":"2023-09-08T01:32:08.186372Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The function below converts an entire training example to the polar-coordinate system.","metadata":{}},{"cell_type":"code","source":"@tf.function(input_signature=(tf.TensorSpec(shape=[None, None], dtype=tf.float32),), experimental_relax_shapes=True)\ndef preprocess_data(data):\n    print('tracing preprocess_data')\n    rh_data = tf.gather(data, PARQUET_RH_FEATURES, axis=1)\n    lh_data = tf.gather(data, PARQUET_LH_FEATURES, axis=1)\n\n    rh_nans = tf.math.count_nonzero(tf.reduce_any(tf.math.is_nan(rh_data), axis=1))\n    lh_nans = tf.math.count_nonzero(tf.reduce_any(tf.math.is_nan(lh_data), axis=1))\n\n    # If there is more data for the left hand, choose it and flip the coordinates to be right-handed.\n    if lh_nans < rh_nans:\n        hand = lh_data\n        x = hand[:, PARQUET_X]\n        y = hand[:, PARQUET_Y]\n        z = hand[:, PARQUET_Z]\n        hand = tf.concat([1 - x, y, z], axis=1)\n\n    else:\n        hand = rh_data\n\n    x = hand[:, PARQUET_X]\n    y = hand[:, PARQUET_Y]\n    z = hand[:, PARQUET_Z]\n    hand = tf.concat([\n        x[..., tf.newaxis],\n        y[..., tf.newaxis],\n        z[..., tf.newaxis],\n    ], axis=-1)\n\n    mean = tf.math.reduce_mean(hand, axis=1)[:, tf.newaxis, :]\n    stdev = tf.math.reduce_std(hand, axis=1)[:, tf.newaxis, :]\n    hand = (hand - mean) / stdev\n\n    priors = tf.concat([\n        [hand[0]],\n        hand\n    ], axis=0)\n\n    mapped = tf.map_fn(\n        fn=lambda arg: hand_rel_coords_tf(arg[0], arg[1]),\n        elems=(hand, priors),\n        fn_output_signature=tf.TensorSpec(shape=[38], dtype=tf.float32)\n    )[:FRAME_COUNT]\n\n    feature_sum = tf.reduce_sum(tf.abs(mapped), axis=1)\n    invalid_removed = tf.boolean_mask(mapped, tf.cast(feature_sum, dtype=tf.bool))\n    valid_frames = tf.shape(invalid_removed)[0]\n\n    pad = tf.cond(valid_frames < FRAME_COUNT, true_fn=lambda: FRAME_COUNT - valid_frames, false_fn=lambda: 0)\n    return tf.pad(invalid_removed, [[0, pad], [0, 0]], 'CONSTANT')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 5. Data augmentation\n\nWe have two ways of augmenting the data:\n\n1. **Temporal resampling.** Suppose that there are 10 frames of data in a training example. Those frames (1-indexed) are `[1, 2, 3, 4, 5, 6, 7, 8, 9, 10]`. We can resample an example to be played back faster or slower by sampling in between real frames (e.g., slower playback looks like `[1, 1.4, 1.8, 2.2, 2.6, 3, 3.4, 3.8, 4.2, 4.6, 5, 5.4, ...]` and faster playback looks like `[1, 2.5, 4, 5.5, 7, 8.5, 10]`). We then simply use linear interpolation (lerp) of all of the features to get a reasonable estimate of the hand position at that time step. \n2. **Spatial perturbation.** One of the winning participants in the [ASL Sign Recognition competition](https://www.kaggle.com/competitions/asl-signs) augmented the data by slightly varying the position of the hand landmarks, a sensible technique since signers won't be perfectly uniform.\n\n**Note**: both augmentation methods *should* work better with polar coordinates than with Cartesian coordinates, due to the angles being simpler to interpolate correctly compared to Cartesian coordinates. For example, if the fingertip starts at `(0, 1, 0)` and curls to `(0, 0, 1)`, the linear interpolation at `t=0.5` is `(0, 0.5, 0.5)`, which is not accurate because the finger curls in a circular fashion. On the other hand, the same fingertip starts at `0` radians and ends at `pi/2` radians, and the linear interpolation of the angle at `t=0.5` is `pi/4` radians, which is correct.\n\n**Function below**:\n* `lerp_and_perturb(data, scale, args)` does all of the heavy mathematical lifting (linear interpolation, calculating angles, etc.) to augment a single frame","metadata":{}},{"cell_type":"code","source":"# TensorFlow doesn't permit tf.Variables to be re-declared, so these need to be declared outside the function.\nlerp_result = tf.Variable(tf.zeros([OUTBOUND_FEATURE_COUNT, 1]))\nperturbs = tf.Variable(tf.zeros([3, 3]))\n\nframe_data_building = tf.Variable(tf.zeros(shape=[FRAME_COUNT, OUTBOUND_FEATURE_COUNT]))\n\n@tf.function\ndef lerp_and_perturb(data, scale, args) -> tf.Variable:\n    i, position, angle_perturbs = args[0], args[1], args[2:]\n    pos_int = tf.cast(position, dtype=tf.int32)\n    before = data[pos_int - 1]\n\n    t = position - tf.cast(pos_int, dtype=tf.float32)\n    after = data[pos_int]\n\n    delta = after - before\n    result = before + (t * delta)\n    lerp_result.assign(tf.reshape(result, [OUTBOUND_FEATURE_COUNT, 1]))\n\n    # ($): Create three random thetas, phis, and rs (each, for 9 total random numbers)\n    # so that each of the three hand coordinates (wrist, base of index, and base of pinky)\n    # can have be perturbed in a random direction \n    thetas = tf.random.uniform([3], minval=0, maxval=2*math.pi)\n    phis = tf.random.uniform([3], minval=0, maxval=math.pi)\n    rs = scale * tf.clip_by_value(\n        tf.random.normal([3], mean=0.5, stddev=1.0/6),\n        clip_value_min=0,\n        clip_value_max=1\n    )\n\n    # (@): Notice below the code for converting the above variables into their Cartesian\n    # equivalent:\n    # \n    # tf.stack([\n    #     r * tf.math.cos(theta) * tf.math.sin(phi),\n    #     r * tf.math.sin(theta) * tf.math.sin(phi),\n    #     r * tf.math.cos(phi)\n    # ])\n    #\n    # We see that each value is multiplied by the radius, so initialize a 3x3 tensor\n    # containing the 3 radii COLUMN-WISE. This means that each COLUMN represents one\n    # vector, and each ROW represents a dimension. This is convenient for doing element-wise\n    # multiplication below, but must be transposed to get the \"natural\" orientation\n    # (where each row is a vector).\n    perturbs.assign(tf.stack([rs, rs, rs]))\n\n    # x = r*cos(theta)*sin(phi)\n    perturbs[0].assign(\n        tf.math.multiply(perturbs[0], \n            tf.math.multiply(tf.math.cos(thetas), tf.math.sin(phis))\n        )\n    )\n\n    # y = r*sin(theta)*sin(phi)\n    perturbs[1].assign(\n        tf.math.multiply(perturbs[1],\n            tf.math.multiply(tf.math.sin(thetas), tf.math.sin(phis))\n        )\n    )\n\n    # z = r*cos(phi)\n    perturbs[2].assign(tf.math.multiply(perturbs[2], tf.math.cos(phis)))\n\n    # 1. Convert Variable to Tensor\n    # 2. Transpose. As per the comment marked (@) above, the tensor is flipped from\n    #    its natural orientation, so we need to transpose it to correct that.\n    # 3. Reshape to be compatible with scatter_nd_add\n    p = tf.reshape(tf.transpose(tf.convert_to_tensor(perturbs)), [9, 1])\n    angle_perturbs = tf.reshape(angle_perturbs, [POLAR_COUNT, 1])\n\n    lerp_result.scatter_nd_add(tf.reshape(HAND_POS, [9, 1]), p)\n    lerp_result.scatter_nd_add(tf.reshape(OUTBOUND_POLAR_FEATURES, [POLAR_COUNT, 1]), angle_perturbs)\n\n    frame_data_building[int(i)].assign(lerp_result)\n    return lerp_result","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We have some hyperparameters for the temporal and spatial augmentations.\n\n#### Temporal resampling parameters:\n* `slow_step=0.67`: the slowest permitted playback speed, relative to 1x.\n* `fast_step=1.5`: the fastest permitted playback speed.\n* `max_zdelta=0.2`: the rate of playback is stabilized so that we don't sporadically speed up and slow down throughout the original data. The amount of time skipped (in frames) is determined by sampling from a Normal distribution N(0, 1/3), then scaling the sampled value between `slow_step` and `fast_step`. The standard deviation is set to 1/3 so that 99.7% of values fall into the range `[-1, 1]`. For each frame, we store the \"z-score\" (not really, since `sigma != 1`, but just go with it) and constrain the next sampling from the normal distribution into the range `[last_z - max_zdelta, last_z + max_zdelta]` (and, of course, the range `[-1, 1]`). Thus, the smaller `max_zdelta` is, the less the playback speed can vary from frame to frame.\n\n#### Spatial perturbation parameters\n* `pos_perturb=0.05`: the maximum amount that the coordinates of the wrist, index finger's base, and pinky finger's base can be adjusted. The value here represents the maximum possible length of each of the three vectors which will be added to each of the three existing landmark coordinates. The chosen length of the vector will be selected from a normal distribution N(0.5, 1/6), and the direction of the vector will be determined from a uniform random distribution.\n* `angle_perturb=0.1` the maximum amount, in radians, that the angular coordinates (for each of the joints in each finger) can be adjusted. The actual values will be selected from a normal distribution N(0, 1/3), multiplied by `angle_perturb`.\n* `max_tdelta=0.05`: like with `max_zdelta`, this value limits the amount by which the perturbation for a given feature can differ from the previous frame's perturbation. Note that this applies only to the features using angular coordinates. (The `t` is for `theta`.)\n\n**Functions below**:\n* `augment_example(data, time_warp, perturb)` handles the augmentation for a single example\n* `augment_callable(data, label)`: Keras handle for augmentation","metadata":{}},{"cell_type":"code","source":"# Note: this variable is transposed to make scatter_nd_update work more easily\nframe_data_var = tf.Variable(tf.zeros(shape=[OUTBOUND_FEATURE_COUNT, FRAME_COUNT]))\n\nframe_times = tf.Variable(tf.zeros([FRAME_COUNT]))\nlast_time = tf.Variable(1.0, dtype=tf.float32)\nlast_z = tf.Variable(0.0, dtype=tf.float32)\ni = tf.Variable(0, dtype=tf.int32)\nframes_sampled = tf.Variable(0, dtype=tf.int32)\n\nlerp_args = tf.Variable(tf.zeros([frames_sampled, POLAR_COUNT + 2]))\n\n@tf.function\ndef augment_example(data, time_warp=(0.67, 1.5, 0.2), perturb=(0.05, 0.1, 0.05)):\n    global frame_times\n    start = datetime.datetime.now()\n    slow_step, fast_step, max_zdelta = tf.constant(time_warp[0]), tf.constant(time_warp[1]), tf.constant(time_warp[2])\n    pos_perturb, angle_perturb, max_tdelta = tf.constant(perturb[0]), tf.constant(perturb[1]), tf.constant(perturb[2])\n\n    slowdown_factor = 1.0 / slow_step\n\n    # if slow_step > 1 or slow_step <= 0:\n    #     raise ValueError(f'augment_example(): invalid time_warp=(slow_step, fast_step, max_zdelta) '\n    #                      f'(requires 0 < slow_step <= 1, given {slow_step=})')\n    # if fast_step < 1:\n    #     raise ValueError(f'augment_example(): invalid time_warp=(slow_step, fast_step, max_zdelta) '\n    #                      f'(requires fast_step >= 1, given {fast_step=})')\n    # if slow_step > fast_step:\n    #     raise ValueError(f'augment_example(): invalid time_warp=(slow_step, fast_step, max_zdelta) '\n    #                      f'(requires fast_step > slow_step, given {slow_step=}, {fast_step=})')\n\n    base_frame_count = tf.cast(tf.shape(data)[0], dtype=tf.float32)\n\n    # Note that we are NOT averaging between the slow_step and fast_step values to determine the mean of the\n    # normal distribution we sample from; for example, the simple average of 0.67 and 1.5 would be 1.085, but\n    # we want the average playback speed to be 1x, not 1.085x. So we sample from a normal distribution that,\n    # according to the Empirical Rule, will genuinely generate a value `z` between -1 and 1 in 99.7% of cases\n    # (and the remaining 0.3% of samples will be forcibly clipped to that same range). Next, we scale the\n    # playback rate according to the sampled value. At z = -1, we skip `slow_step` frames; at z = 0, we skip\n    # exactly 1 frame; at z = 1, we skip `fast_step` frames.\n\n    @tf.function\n    def next_frametime(last_time, last_z, i):\n        global frame_times\n        z = tf.clip_by_value(\n            tf.random.normal([1], mean=0, stddev=1.0/3),\n            clip_value_min=tf.math.maximum(-1.0, last_z - max_zdelta),\n            clip_value_max=tf.math.minimum(1.0, last_z + max_zdelta)\n        )[0]\n\n        if z == 0:\n            time_skip = tf.constant(1.0)\n        elif z < 0:\n            slow_rate = 1.0 - z * (slowdown_factor - 1)\n            time_skip = 1.0 / slow_rate\n        else:\n            time_skip = 1.0 + (z * (fast_step - 1))\n\n        next_time = last_time + time_skip\n        if next_time <= base_frame_count:\n            frame_times[i].assign(next_time)\n            tf.add(frames_sampled, 1)\n        \n        return [next_time, z, tf.add(i, 1)]\n\n    tf.while_loop(\n        cond=lambda lt, lz, i: tf.math.logical_and(lt < base_frame_count, i < FRAME_COUNT),\n        body=next_frametime,\n        loop_vars=[last_time, last_z, i],\n        parallel_iterations=1\n    )\n\n    frame_times = frame_times[:frames_sampled]\n\n    # STEP 2. PERTURBATION\n    prior_perturbs = tf.zeros([POLAR_COUNT])\n\n    @tf.function\n    def next_perturb(i, prior):\n        new_perturbs = tf.clip_by_value(\n            angle_perturb * tf.math.abs(tf.random.normal([POLAR_COUNT], mean=0, stddev=1/3)),\n            clip_value_min=tf.math.maximum(tf.zeros([POLAR_COUNT]), prior - max_tdelta),\n            clip_value_max=tf.math.minimum(tf.zeros([POLAR_COUNT]) + angle_perturb, prior + max_tdelta)\n        )\n\n        lerp_args[i].assign(tf.concat([[i, frame_times[i]], new_perturbs], 0))\n        return [tf.add(i, 1), new_perturbs]\n\n    i.assign(0)\n    tf.while_loop(\n        cond=lambda i, p: i < frames_sampled,\n        body=next_perturb,\n        loop_vars=[i, prior_perturbs],\n        parallel_iterations=1\n    )\n\n    frame_data = tf.zeros([OUTBOUND_FEATURE_COUNT, FRAME_COUNT])\n    tf.map_fn(fn=lambda args: lerp_and_perturb(data, pos_perturb, args), elems=lerp_args)\n\n    new_frame_data = tf.convert_to_tensor(frame_data_building)\n\n    offset_frames = tf.transpose(\n        tf.concat([tf.reshape(new_frame_data[0], [1, OUTBOUND_FEATURE_COUNT]) , new_frame_data[:-1]], 0)\n    )\n\n    # Transpose because the scatter_nd_function allows us to easily select rows, but not columns\n    nfd_transpose = tf.transpose(new_frame_data)\n    deltas = nfd_transpose - offset_frames\n    \n    frame_data_var.assign(nfd_transpose)\n    \n    # Calculate dX, dY, dZ features for each of the three positions\n    frame_data_var.scatter_nd_update(\n        [[WRIST_POS_dX], [WRIST_POS_dY], [WRIST_POS_dZ]], \n        [deltas[WRIST_POS_X], deltas[WRIST_POS_Y], deltas[WRIST_POS_Z]]\n    )\n\n    frame_data_var.scatter_nd_update(\n        [[INDEX_MCP_dX], [INDEX_MCP_dY], [INDEX_MCP_dZ]], \n        [deltas[INDEX_MCP_X], deltas[INDEX_MCP_Y], deltas[INDEX_MCP_Z]]\n    )\n\n    frame_data_var.scatter_nd_update(\n        [[PINKY_MCP_dX], [PINKY_MCP_dY], [PINKY_MCP_dZ]], \n        [deltas[PINKY_MCP_X], deltas[PINKY_MCP_Y], deltas[PINKY_MCP_Z]]\n    )\n\n    return tf.transpose(tf.convert_to_tensor(frame_data_var))[:frames_sampled]\n\n@tf.function\ndef augment_callable(data, label):\n    return augment_example(data), label","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 6. Load data and train\n\nThe `tfrec_load` and `tf_prepare` functions are largely copied from [the same guide](https://www.kaggle.com/code/gusthema/asl-fingerspelling-recognition-w-tensorflow) as the model architecture.","metadata":{}},{"cell_type":"code","source":"@tf.function\ndef tfrec_load(record_bytes):\n    schema = {\n        column: tf.io.VarLenFeature(dtype=tf.float32)\n        for column in PARQUET_FEATURE_LIST[:-2]\n    }\n\n    schema['phrase'] = tf.io.FixedLenFeature([], dtype=tf.string)\n    features = tf.io.parse_single_example(record_bytes, schema)\n    phrase = features['phrase']\n\n    landmarks = ([tf.sparse.to_dense(features[column]) for column in PARQUET_FEATURE_LIST[:-2]])\n    landmarks = tf.transpose(landmarks)\n    return landmarks, phrase\n\n\ndef tf_prepare(landmarks, phrase):\n    phrase = START + phrase + END\n    phrase = tf.strings.bytes_split(phrase)\n    phrase = MAPPING_LOOKUP_TABLE.lookup(phrase)\n    phrase = tf.pad(phrase, paddings=[[0, 64 - tf.shape(phrase)[0]]], mode='CONSTANT',\n                    constant_values=CHAR_TO_IDX[PAD])\n    return preprocess_data(landmarks), phrase\n\n\ndef train():\n    DEBUG = False\n    tf_records = [\n        os.path.join(PREPROCESSED_DIR, file) for file in os.listdir(PREPROCESSED_DIR)\n        if file.endswith('.tfrecord')\n    ]\n\n    batch_size = BATCH_SIZE\n    train_size = int(0.8 * len(tf_records))\n    tensorboard_callback = keras.callbacks.TensorBoard(log_dir='log')\n\n    train_ds = tf.data.TFRecordDataset(tf_records[:train_size]) \\\n        .map(tfrec_load) \\\n        .map(tf_prepare) \\\n        .map(augment_callable, num_parallel_calls=tf.data.AUTOTUNE) \\\n        .batch(batch_size) \\\n        .prefetch(buffer_size=tf.data.AUTOTUNE) \\\n        .cache()\n\n    valid_ds = tf.data.TFRecordDataset(tf_records[train_size:]) \\\n        .map(tfrec_load) \\\n        .map(tf_prepare) \\\n        .batch(batch_size) \\\n        .prefetch(buffer_size=tf.data.AUTOTUNE) \\\n        .cache()\n\n    print(valid_ds)\n    batch = next(iter(valid_ds))\n\n    # The vocabulary to convert predicted indices into characters\n    display_cb = DisplayOutputs(\n        batch,\n        IDX_TO_CHAR,\n        target_start_token_idx=CHAR_TO_IDX[START],\n        target_end_token_idx=CHAR_TO_IDX[END]\n    )  # set the arguments as per vocabulary index for '<' and '>'\n\n    model = Transformer(\n        num_hid=200,\n        num_head=4,\n        num_feed_forward=400,\n        source_maxlen=FRAME_COUNT,\n        target_maxlen=64,\n        num_layers_enc=2,\n        num_layers_dec=1,\n        num_classes=62\n    )\n\n    loss_fn = tf.keras.losses.CategoricalCrossentropy(\n        from_logits=True, label_smoothing=0.1,\n    )\n\n    optimizer = keras.optimizers.Adam(0.0001)\n    model.compile(optimizer=optimizer, loss=loss_fn)\n\n    history = model.fit(train_ds, validation_data=valid_ds, callbacks=[display_cb], epochs=EPOCHS)\n\n    plt.plot(history.history['loss'])\n    plt.plot(history.history['val_loss'])\n    plt.legend(['training loss', 'val_loss'])\n\n    model.save('model.keras')\n    return model","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"time = datetime.datetime.now()\nmodel = train()\n\nprint(model.summary())\nprint(f'Started {time}, finished {datetime.datetime.now()}')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 7. Convert to TFLite","metadata":{}},{"cell_type":"code","source":"class TFLiteModel(tf.Module):\n    def __init__(self, model):\n        super(TFLiteModel, self).__init__()\n        self.target_start_token_idx = CHAR_TO_IDX[START]\n        self.target_end_token_idx = CHAR_TO_IDX[END]\n        self.model = model\n\n    @tf.function(input_signature=[tf.TensorSpec(shape=[None, INBOUND_FEATURE_COUNT], dtype=tf.float32, name='inputs')])\n    def __call__(self, inputs, training=False):\n        # Preprocess Data\n        x = tf.cast(inputs, tf.float32)\n        x = x[None]\n        x = tf.cond(tf.shape(x)[1] == 0, lambda: tf.zeros((1, 1, INBOUND_FEATURE_COUNT)), lambda: tf.identity(x))\n        x = x[0]\n        x = preprocess_data(x)\n        x = x[None]\n        x = self.model.generate(x, self.target_start_token_idx)\n        x = x[0]\n        idx = tf.argmax(tf.cast(tf.equal(x, self.target_end_token_idx), tf.int32))\n        idx = tf.where(tf.math.less(idx, 1), tf.constant(2, dtype=tf.int64), idx)\n        x = x[1:idx]\n        x = tf.one_hot(x, 59)\n        return {'outputs': x}","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tfmodel = TFLiteModel(model)\nkeras_model_converter = tf.lite.TFLiteConverter.from_keras_model(tfmodel)\nkeras_model_converter.target_spec.supported_ops = [\n    tf.lite.OpsSet.TFLITE_BUILTINS,\n    tf.lite.OpsSet.SELECT_TF_OPS\n]\n\nkeras_model_converter.allow_custom_ops = True\n\ntflmodel = keras_model_converter.convert()\n\nwith open(OUTPUT_MODEL, 'wb') as file:\n    file.write(tflmodel)\n\ntfargs = {'selected_columns': PARQUET_FEATURE_LIST}\nwith open('inference_args.json', 'w') as json_file:\n    json.dump(tfargs, json_file)","metadata":{},"execution_count":null,"outputs":[]}],"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"}}