{"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":"# Update (version 9):\n\nThis version (9) incorporates the latest updates to [ROHITH's notebbok](https://www.kaggle.com/code/irohith/aslfr-ctc-based-on-prev-comp-1st-place) (version 5)","metadata":{}},{"cell_type":"markdown","source":"# CTC on TPU\n\nThis modification of [ROHITH INGILELA's CTC notebook](https://www.kaggle.com/code/irohith/aslfr-ctc-based-on-prev-comp-1st-place) can run on TPU. I used [another implementation](https://github.com/alexeytochin/tf_seq2seq_losses) for CTC since TensorFlow implementation does not work on Kaggle's TPU. This implementation also does not run on Kaggle's TPU but was much easier to debug since 'This is a pure Python/TensorFlow implementation. We do not have to build or compile any C++/CUDA stuff.' as stated on the project page. This makes a HUGE difference, as the original TensorFlow implementation has some messy code that relies directly on TensorFlow's ops behind the curtains. When you dig deep enough, it does not work on Kaggle TPU due to a bug related to some TensorFlow ops. Then you have to start debugging some ugly C-related code (don't ask how I know 😉)...well, this is why I eventually searched for another route, haha. Anyway, As I said, this implementation also did not work at first on Kaggle's TPU, but after some debugging that resulted in two small changes to the source code, I could make it work. All this work was NOT easy by any means...took me about two days of trial and error. I would appreciate your votes🙏\n\n**RUNTIME:** With Kaggle TPU and some optimization to data loading (namely, load all the data to the RAM before converting it to a dataset instead of loading in batches from scratch each iteration, since TPU notebooks have more than enough rum), result in a MUCH faster runtime. ~47 min for the entire notebook (50 epochs) as opposed to the original notebook with ~6 hours for the whole notebook. More than seven times faster.\n\n**SUBMITTING:** To submit the TPU model, save it to h5 with model.save_weights, then build the model on a GPU notebook and load the weights with model.load_weights. If you submit directly from a TPU notebook, you will have to solve some problems that are not worth the time (if possible to solve at all).\n\n**P.S.** The original notebook accidentally trained on the validation score, which is, of course, a big NO-NO...and this is why it achieved a very low CTC validation score (~6). I fixed it and also added two callback functions, one for calculating the validation set Levenshtein distance, since this is our metric, and one for saving the models during the training (I saved the weights every five epochs, you can change that)","metadata":{}},{"cell_type":"code","source":"import gc\nimport json\nimport math\nimport pickle\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nimport tensorflow as tf\nimport tensorflow_addons as tfa\n\n!pip install cached-property\nfrom cached_property import cached_property\nfrom shutil import copyfile\n\n!pip install fastparquet\nimport fastparquet\n\n!pip install Levenshtein\nimport Levenshtein as lev","metadata":{"execution":{"iopub.status.busy":"2023-07-23T18:36:22.545146Z","iopub.execute_input":"2023-07-23T18:36:22.5455Z","iopub.status.idle":"2023-07-23T18:37:05.516441Z","shell.execute_reply.started":"2023-07-23T18:36:22.545474Z","shell.execute_reply":"2023-07-23T18:37:05.515164Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# copy our file into the working directory (make sure it has .py suffix)\ncopyfile(src = \"/kaggle/input/ctc-tpu/CTC_TPU.py\", dst = \"/kaggle/working//CTC_TPU.py\")\n\n# import all our functions\nfrom CTC_TPU import classic_ctc_loss","metadata":{"execution":{"iopub.status.busy":"2023-07-23T10:32:30.770264Z","iopub.execute_input":"2023-07-23T10:32:30.770754Z","iopub.status.idle":"2023-07-23T10:32:30.779262Z","shell.execute_reply.started":"2023-07-23T10:32:30.770708Z","shell.execute_reply":"2023-07-23T10:32:30.778221Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Configure Strategy. Assume TPU...if not set default for GPU\ntpu = None\ntry:\n    tpu = tf.distribute.cluster_resolver.TPUClusterResolver.connect(tpu=\"local\") # \"local\" for 1VM TPU\n    strategy = tf.distribute.TPUStrategy(tpu)\n    print(\"on TPU\")\n    print(\"REPLICAS: \", strategy.num_replicas_in_sync)\nexcept:\n    strategy = tf.distribute.get_strategy()","metadata":{"execution":{"iopub.status.busy":"2023-07-23T10:32:30.780774Z","iopub.execute_input":"2023-07-23T10:32:30.781795Z","iopub.status.idle":"2023-07-23T10:32:31.13024Z","shell.execute_reply.started":"2023-07-23T10:32:30.781755Z","shell.execute_reply":"2023-07-23T10:32:31.129197Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with open (\"/kaggle/input/asl-fingerspelling/character_to_prediction_index.json\", \"r\") as f:\n    char_to_num = json.load(f)\n\npad_token = '^'\npad_token_idx = 59\n\nchar_to_num[pad_token] = pad_token_idx\n\nnum_to_char = {j:i for i,j in char_to_num.items()}\ndf = pd.read_csv('/kaggle/input/asl-fingerspelling/train.csv')\n\nLIP = [\n    61, 185, 40, 39, 37, 0, 267, 269, 270, 409,\n    291, 146, 91, 181, 84, 17, 314, 405, 321, 375,\n    78, 191, 80, 81, 82, 13, 312, 311, 310, 415,\n    95, 88, 178, 87, 14, 317, 402, 318, 324, 308,\n]\nLPOSE = [13, 15, 17, 19, 21]\nRPOSE = [14, 16, 18, 20, 22]\nPOSE = LPOSE + RPOSE\n\nX = [f'x_right_hand_{i}' for i in range(21)] + [f'x_left_hand_{i}' for i in range(21)] + [f'x_pose_{i}' for i in POSE] + [f'x_face_{i}' for i in LIP]\nY = [f'y_right_hand_{i}' for i in range(21)] + [f'y_left_hand_{i}' for i in range(21)] + [f'y_pose_{i}' for i in POSE] + [f'y_face_{i}' for i in LIP]\nZ = [f'z_right_hand_{i}' for i in range(21)] + [f'z_left_hand_{i}' for i in range(21)] + [f'z_pose_{i}' for i in POSE] + [f'z_face_{i}' for i in LIP]\n\nSEL_COLS = X + Y + Z\nFRAME_LEN = 128\nMAX_PHRASE_LENGTH = 64\n\nLIP_IDX_X   = [i for i, col in enumerate(SEL_COLS)  if  \"face\" in col and \"x\" in col]\nRHAND_IDX_X = [i for i, col in enumerate(SEL_COLS)  if \"right\" in col and \"x\" in col]\nLHAND_IDX_X = [i for i, col in enumerate(SEL_COLS)  if  \"left\" in col and \"x\" in col]\nRPOSE_IDX_X = [i for i, col in enumerate(SEL_COLS)  if  \"pose\" in col and int(col[-2:]) in RPOSE and \"x\" in col]\nLPOSE_IDX_X = [i for i, col in enumerate(SEL_COLS)  if  \"pose\" in col and int(col[-2:]) in LPOSE and \"x\" in col]\n\nLIP_IDX_Y   = [i for i, col in enumerate(SEL_COLS)  if  \"face\" in col and \"y\" in col]\nRHAND_IDX_Y = [i for i, col in enumerate(SEL_COLS)  if \"right\" in col and \"y\" in col]\nLHAND_IDX_Y = [i for i, col in enumerate(SEL_COLS)  if  \"left\" in col and \"y\" in col]\nRPOSE_IDX_Y = [i for i, col in enumerate(SEL_COLS)  if  \"pose\" in col and int(col[-2:]) in RPOSE and \"y\" in col]\nLPOSE_IDX_Y = [i for i, col in enumerate(SEL_COLS)  if  \"pose\" in col and int(col[-2:]) in LPOSE and \"y\" in col]\n\nLIP_IDX_Z   = [i for i, col in enumerate(SEL_COLS)  if  \"face\" in col and \"z\" in col]\nRHAND_IDX_Z = [i for i, col in enumerate(SEL_COLS)  if \"right\" in col and \"z\" in col]\nLHAND_IDX_Z = [i for i, col in enumerate(SEL_COLS)  if  \"left\" in col and \"z\" in col]\nRPOSE_IDX_Z = [i for i, col in enumerate(SEL_COLS)  if  \"pose\" in col and int(col[-2:]) in RPOSE and \"z\" in col]\nLPOSE_IDX_Z = [i for i, col in enumerate(SEL_COLS)  if  \"pose\" in col and int(col[-2:]) in LPOSE and \"z\" in col]\n\nRHM = np.load(\"/kaggle/input/aslfr-dataset-tfrecords/mean_std/rh_mean.npy\")\nLHM = np.load(\"/kaggle/input/aslfr-dataset-tfrecords/mean_std/lh_mean.npy\")\nRPM = np.load(\"/kaggle/input/aslfr-dataset-tfrecords/mean_std/rp_mean.npy\")\nLPM = np.load(\"/kaggle/input/aslfr-dataset-tfrecords/mean_std/lp_mean.npy\")\nLIPM = np.load(\"/kaggle/input/aslfr-dataset-tfrecords/mean_std/lip_mean.npy\")\n\nRHS = np.load(\"/kaggle/input/aslfr-dataset-tfrecords/mean_std/rh_std.npy\")\nLHS = np.load(\"/kaggle/input/aslfr-dataset-tfrecords/mean_std/lh_std.npy\")\nRPS = np.load(\"/kaggle/input/aslfr-dataset-tfrecords/mean_std/rp_std.npy\")\nLPS = np.load(\"/kaggle/input/aslfr-dataset-tfrecords/mean_std/lp_std.npy\")\nLIPS = np.load(\"/kaggle/input/aslfr-dataset-tfrecords/mean_std/lip_std.npy\")","metadata":{"execution":{"iopub.status.busy":"2023-07-23T10:32:31.134658Z","iopub.execute_input":"2023-07-23T10:32:31.135047Z","iopub.status.idle":"2023-07-23T10:32:31.273384Z","shell.execute_reply.started":"2023-07-23T10:32:31.135018Z","shell.execute_reply":"2023-07-23T10:32:31.272327Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_relevant_data_subset(pq_path):\n    return pd.read_parquet(pq_path, columns=SEL_COLS)\n\nfile_id = df.file_id.iloc[0]\ninpdir = \"/kaggle/input/asl-fingerspelling/train_landmarks\"\npqfile = f\"{inpdir}/{file_id}.parquet\"\nseq_refs = df.loc[df.file_id == file_id]\nseqs = load_relevant_data_subset(pqfile)\n\nseq_id = seq_refs.sequence_id.iloc[0]\nframes = seqs.iloc[seqs.index == seq_id]\nphrase = str(df.loc[df.sequence_id == seq_id].phrase.iloc[0])","metadata":{"execution":{"iopub.status.busy":"2023-07-23T10:32:31.275214Z","iopub.execute_input":"2023-07-23T10:32:31.275638Z","iopub.status.idle":"2023-07-23T10:32:31.778061Z","shell.execute_reply.started":"2023-07-23T10:32:31.275599Z","shell.execute_reply":"2023-07-23T10:32:31.776749Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@tf.function()\ndef resize_pad(x):\n    if tf.shape(x)[0] < FRAME_LEN:\n        x = tf.pad(x, ([[0, FRAME_LEN-tf.shape(x)[0]], [0, 0], [0, 0]]), constant_values=float(\"NaN\"))\n    else:\n        x = tf.image.resize(x, (FRAME_LEN, tf.shape(x)[1]))\n    return x\n\n@tf.function(jit_compile=True)\ndef pre_process0(x):\n    lip_x = tf.gather(x, LIP_IDX_X, axis=1)\n    lip_y = tf.gather(x, LIP_IDX_Y, axis=1)\n    lip_z = tf.gather(x, LIP_IDX_Z, axis=1)\n\n    rhand_x = tf.gather(x, RHAND_IDX_X, axis=1)\n    rhand_y = tf.gather(x, RHAND_IDX_Y, axis=1)\n    rhand_z = tf.gather(x, RHAND_IDX_Z, axis=1)\n    \n    lhand_x = tf.gather(x, LHAND_IDX_X, axis=1)\n    lhand_y = tf.gather(x, LHAND_IDX_Y, axis=1)\n    lhand_z = tf.gather(x, LHAND_IDX_Z, axis=1)\n\n    rpose_x = tf.gather(x, RPOSE_IDX_X, axis=1)\n    rpose_y = tf.gather(x, RPOSE_IDX_Y, axis=1)\n    rpose_z = tf.gather(x, RPOSE_IDX_Z, axis=1)\n    \n    lpose_x = tf.gather(x, LPOSE_IDX_X, axis=1)\n    lpose_y = tf.gather(x, LPOSE_IDX_Y, axis=1)\n    lpose_z = tf.gather(x, LPOSE_IDX_Z, axis=1)\n    \n    lip   = tf.concat([lip_x[..., tf.newaxis], lip_y[..., tf.newaxis], lip_z[..., tf.newaxis]], axis=-1)\n    rhand = tf.concat([rhand_x[..., tf.newaxis], rhand_y[..., tf.newaxis], rhand_z[..., tf.newaxis]], axis=-1)\n    lhand = tf.concat([lhand_x[..., tf.newaxis], lhand_y[..., tf.newaxis], lhand_z[..., tf.newaxis]], axis=-1)\n    rpose = tf.concat([rpose_x[..., tf.newaxis], rpose_y[..., tf.newaxis], rpose_z[..., tf.newaxis]], axis=-1)\n    lpose = tf.concat([lpose_x[..., tf.newaxis], lpose_y[..., tf.newaxis], lpose_z[..., tf.newaxis]], axis=-1)\n    \n    hand = tf.concat([rhand, lhand], axis=1)\n    hand = tf.where(tf.math.is_nan(hand), 0.0, hand)\n    mask = tf.math.not_equal(tf.reduce_sum(hand, axis=[1, 2]), 0.0)\n\n    lip = lip[mask]\n    rhand = rhand[mask]\n    lhand = lhand[mask]\n    rpose = rpose[mask]\n    lpose = lpose[mask]\n\n    return lip, rhand, lhand, rpose, lpose\n\n@tf.function()\ndef pre_process1(lip, rhand, lhand, rpose, lpose):\n    lip   = (resize_pad(lip) - LIPM) / LIPS\n    rhand = (resize_pad(rhand) - RHM) / RHS\n    lhand = (resize_pad(lhand) - LHM) / LHS\n    rpose = (resize_pad(rpose) - RPM) / RPS\n    lpose = (resize_pad(lpose) - LPM) / LPS\n\n    x = tf.concat([lip, rhand, lhand, rpose, lpose], axis=1)\n    s = tf.shape(x)\n    x = tf.reshape(x, (s[0], s[1]*s[2]))\n    x = tf.where(tf.math.is_nan(x), 0.0, x)\n    return x\n\n#This fail on TPU\n'''\npre0 = pre_process0(frames)\npre1 = pre_process1(*pre0)\nINPUT_SHAPE = list(pre1.shape)\nprint(INPUT_SHAPE)\npre1\n'''\nINPUT_SHAPE = [128, 276]","metadata":{"execution":{"iopub.status.busy":"2023-07-23T10:32:31.779951Z","iopub.execute_input":"2023-07-23T10:32:31.780846Z","iopub.status.idle":"2023-07-23T10:32:31.800276Z","shell.execute_reply.started":"2023-07-23T10:32:31.780804Z","shell.execute_reply":"2023-07-23T10:32:31.799352Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def decode_fn(record_bytes):\n    schema = {\n        \"lip\": tf.io.VarLenFeature(tf.float32),\n        \"rhand\": tf.io.VarLenFeature(tf.float32),\n        \"lhand\": tf.io.VarLenFeature(tf.float32),\n        \"rpose\": tf.io.VarLenFeature(tf.float32),\n        \"lpose\": tf.io.VarLenFeature(tf.float32),\n        \"phrase\": tf.io.VarLenFeature(tf.int64)\n    }\n    x = tf.io.parse_single_example(record_bytes, schema)\n\n    lip = tf.reshape(tf.sparse.to_dense(x[\"lip\"]), (-1, 40, 3))\n    rhand = tf.reshape(tf.sparse.to_dense(x[\"rhand\"]), (-1, 21, 3))\n    lhand = tf.reshape(tf.sparse.to_dense(x[\"lhand\"]), (-1, 21, 3))\n    rpose = tf.reshape(tf.sparse.to_dense(x[\"rpose\"]), (-1, 5, 3))\n    lpose = tf.reshape(tf.sparse.to_dense(x[\"lpose\"]), (-1, 5, 3))\n    phrase = tf.sparse.to_dense(x[\"phrase\"])\n\n    return lip, rhand, lhand, rpose, lpose, phrase\n\ndef pre_process_fn(lip, rhand, lhand, rpose, lpose, phrase):\n    phrase = tf.pad(phrase, [[0, MAX_PHRASE_LENGTH-tf.shape(phrase)[0]]], constant_values=pad_token_idx)\n    return pre_process1(lip, rhand, lhand, rpose, lpose), phrase\n    \ntffiles = [f\"/kaggle/input/aslfr-dataset-tfrecords/tfds/{file_id}.tfrecord\" for file_id in df.file_id.unique()]\nval_len = 1#int(0.05 * len(tffiles))\nprint('val_len: ' + str(val_len))\ntrain_batch_size = 32\nval_batch_size = 32\n'''\ntrain_dataset =  tf.data.TFRecordDataset(tffiles).prefetch(tf.data.AUTOTUNE).shuffle(5000).map(decode_fn, num_parallel_calls=tf.data.AUTOTUNE).map(pre_process_fn, num_parallel_calls=tf.data.AUTOTUNE).batch(train_batch_size).prefetch(tf.data.AUTOTUNE)\nval_dataset =  tf.data.TFRecordDataset(tffiles[:val_len]).prefetch(tf.data.AUTOTUNE).map(decode_fn, num_parallel_calls=tf.data.AUTOTUNE).map(pre_process_fn, num_parallel_calls=tf.data.AUTOTUNE).batch(train_batch_size).prefetch(tf.data.AUTOTUNE)\n'''\ntrain_dataset_pre =  tf.data.TFRecordDataset(tffiles[val_len:]).prefetch(tf.data.AUTOTUNE).map(decode_fn, num_parallel_calls=tf.data.AUTOTUNE).map(pre_process_fn, num_parallel_calls=tf.data.AUTOTUNE)\nval_dataset_pre =  tf.data.TFRecordDataset(tffiles[:val_len]).prefetch(tf.data.AUTOTUNE).map(decode_fn, num_parallel_calls=tf.data.AUTOTUNE).map(pre_process_fn, num_parallel_calls=tf.data.AUTOTUNE)\n\n'''\nbatch = next(iter(val_dataset))\nbatch[0].shape, batch[1].shape\n'''","metadata":{"execution":{"iopub.status.busy":"2023-07-23T10:32:31.801589Z","iopub.execute_input":"2023-07-23T10:32:31.80267Z","iopub.status.idle":"2023-07-23T10:32:32.3639Z","shell.execute_reply.started":"2023-07-23T10:32:31.802629Z","shell.execute_reply":"2023-07-23T10:32:32.363098Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_items = [x for x in val_dataset_pre]\nval_items_X = [x[0] for x in val_items]\nval_items_y = [tf.cast(x[1], dtype = tf.int32) for x in val_items]\nval_dataset = tf.data.Dataset.from_tensor_slices((val_items_X, val_items_y)).prefetch(tf.data.AUTOTUNE).batch(\n    val_batch_size, drop_remainder=True).prefetch(tf.data.AUTOTUNE)\n\ntrain_items = [x for x in train_dataset_pre]\ntrain_items_X = [x[0] for x in train_items]\ntrain_items_y = [tf.cast(x[1], dtype = tf.int32) for x in train_items]\ntrain_dataset = tf.data.Dataset.from_tensor_slices((train_items_X, train_items_y))\ntrain_dataset = train_dataset.prefetch(tf.data.AUTOTUNE).shuffle(\n    buffer_size=5000, reshuffle_each_iteration = True).batch(train_batch_size, drop_remainder=True).prefetch(\n    tf.data.AUTOTUNE)\n\nbatch = next(iter(val_dataset))\nbatch[0].shape, batch[1].shape","metadata":{"execution":{"iopub.status.busy":"2023-07-23T10:32:32.407313Z","iopub.status.idle":"2023-07-23T10:32:32.40797Z","shell.execute_reply.started":"2023-07-23T10:32:32.407713Z","shell.execute_reply":"2023-07-23T10:32:32.407734Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Copied from previous comp 1st place model: https://www.kaggle.com/code/hoyso48/1st-place-solution-training\nclass ECA(tf.keras.layers.Layer):\n    def __init__(self, kernel_size=5, **kwargs):\n        super().__init__(**kwargs)\n        self.supports_masking = True\n        self.kernel_size = kernel_size\n        self.conv = tf.keras.layers.Conv1D(1, kernel_size=kernel_size, strides=1, padding=\"same\", use_bias=False)\n\n    def call(self, inputs, mask=None):\n        nn = tf.keras.layers.GlobalAveragePooling1D()(inputs, mask=mask)\n        nn = tf.expand_dims(nn, -1)\n        nn = self.conv(nn)\n        nn = tf.squeeze(nn, -1)\n        nn = tf.nn.sigmoid(nn)\n        nn = nn[:,None,:]\n        return inputs * nn\n\nclass CausalDWConv1D(tf.keras.layers.Layer):\n    def __init__(self, \n        kernel_size=17,\n        dilation_rate=1,\n        use_bias=False,\n        depthwise_initializer='glorot_uniform',\n        name='', **kwargs):\n        super().__init__(name=name,**kwargs)\n        self.causal_pad = tf.keras.layers.ZeroPadding1D((dilation_rate*(kernel_size-1),0),name=name + '_pad')\n        self.dw_conv = tf.keras.layers.DepthwiseConv1D(\n                            kernel_size,\n                            strides=1,\n                            dilation_rate=dilation_rate,\n                            padding='valid',\n                            use_bias=use_bias,\n                            depthwise_initializer=depthwise_initializer,\n                            name=name + '_dwconv')\n        self.supports_masking = True\n        \n    def call(self, inputs):\n        x = self.causal_pad(inputs)\n        x = self.dw_conv(x)\n        return x\n\ndef Conv1DBlock(channel_size,\n          kernel_size,\n          dilation_rate=1,\n          drop_rate=0.0,\n          expand_ratio=2,\n          se_ratio=0.25,\n          activation='swish',\n          name=None):\n    '''\n    efficient conv1d block, @hoyso48\n    '''\n    if name is None:\n        name = str(tf.keras.backend.get_uid(\"mbblock\"))\n    # Expansion phase\n    def apply(inputs):\n        channels_in = tf.keras.backend.int_shape(inputs)[-1]\n        channels_expand = channels_in * expand_ratio\n\n        skip = inputs\n\n        x = tf.keras.layers.Dense(\n            channels_expand,\n            use_bias=True,\n            activation=activation,\n            name=name + '_expand_conv')(inputs)\n\n        # Depthwise Convolution\n        x = CausalDWConv1D(kernel_size,\n            dilation_rate=dilation_rate,\n            use_bias=False,\n            name=name + '_dwconv')(x)\n\n        x = tf.keras.layers.BatchNormalization(momentum=0.95, name=name + '_bn')(x)\n\n        x  = ECA()(x)\n\n        x = tf.keras.layers.Dense(\n            channel_size,\n            use_bias=True,\n            name=name + '_project_conv')(x)\n\n        if drop_rate > 0:\n            x = tf.keras.layers.Dropout(drop_rate, noise_shape=(None,1,1), name=name + '_drop')(x)\n\n        if (channels_in == channel_size):\n            x = tf.keras.layers.add([x, skip], name=name + '_add')\n        return x\n\n    return apply\n\nclass MultiHeadSelfAttention(tf.keras.layers.Layer):\n    def __init__(self, dim=256, num_heads=4, dropout=0, **kwargs):\n        super().__init__(**kwargs)\n        self.dim = dim\n        self.scale = self.dim ** -0.5\n        self.num_heads = num_heads\n        self.qkv = tf.keras.layers.Dense(3 * dim, use_bias=False)\n        self.drop1 = tf.keras.layers.Dropout(dropout)\n        self.proj = tf.keras.layers.Dense(dim, use_bias=False)\n        self.supports_masking = True\n\n    def call(self, inputs, mask=None):\n        qkv = self.qkv(inputs)\n        qkv = tf.keras.layers.Permute((2, 1, 3))(tf.keras.layers.Reshape((-1, self.num_heads, self.dim * 3 // self.num_heads))(qkv))\n        q, k, v = tf.split(qkv, [self.dim // self.num_heads] * 3, axis=-1)\n\n        attn = tf.matmul(q, k, transpose_b=True) * self.scale\n\n        if mask is not None:\n            mask = mask[:, None, None, :]\n\n        attn = tf.keras.layers.Softmax(axis=-1)(attn, mask=mask)\n        attn = self.drop1(attn)\n\n        x = attn @ v\n        x = tf.keras.layers.Reshape((-1, self.dim))(tf.keras.layers.Permute((2, 1, 3))(x))\n        x = self.proj(x)\n        return x\n\n\ndef TransformerBlock(dim=256, num_heads=6, expand=4, attn_dropout=0.2, drop_rate=0.2, activation='swish'):\n    def apply(inputs):\n        x = inputs\n        x = tf.keras.layers.LayerNormalization(epsilon=1e-6)(x)\n        x = MultiHeadSelfAttention(dim=dim,num_heads=num_heads,dropout=attn_dropout)(x)\n        x = tf.keras.layers.Dropout(drop_rate, noise_shape=(None,1,1))(x)\n        x = tf.keras.layers.Add()([inputs, x])\n        attn_out = x\n\n        x = tf.keras.layers.LayerNormalization(epsilon=1e-6)(x)\n        x = tf.keras.layers.Dense(dim*expand, use_bias=False, activation=activation)(x)\n        x = tf.keras.layers.Dense(dim, use_bias=False)(x)\n        x = tf.keras.layers.Dropout(drop_rate, noise_shape=(None,1,1))(x)\n        x = tf.keras.layers.Add()([attn_out, x])\n        return x\n    return apply\n\ndef positional_encoding(maxlen, num_hid):\n        depth = num_hid/2\n        positions = tf.range(maxlen, dtype = tf.float32)[..., tf.newaxis]\n        depths = tf.range(depth, dtype = tf.float32)[np.newaxis, :]/depth\n        angle_rates = tf.math.divide(1, tf.math.pow(tf.cast(10000, tf.float32), depths))\n        angle_rads = tf.linalg.matmul(positions, angle_rates)\n        pos_encoding = tf.concat(\n          [tf.math.sin(angle_rads), tf.math.cos(angle_rads)],\n          axis=-1)\n        return pos_encoding","metadata":{"execution":{"iopub.status.busy":"2023-07-23T10:32:32.409088Z","iopub.status.idle":"2023-07-23T10:32:32.409623Z","shell.execute_reply.started":"2023-07-23T10:32:32.409443Z","shell.execute_reply":"2023-07-23T10:32:32.409463Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def CTCLoss(labels, logits):\n    label_length = tf.reduce_sum(tf.cast(labels != pad_token_idx, tf.int32), axis=-1)\n    logit_length = tf.ones(tf.shape(logits)[0], dtype=tf.int32) * tf.shape(logits)[1]\n    \n    loss = classic_ctc_loss(\n            labels=labels,\n            logits=logits,\n            label_length=label_length,\n            logit_length=logit_length,\n            blank_index=pad_token_idx,\n        )\n    '''\n    loss = tf.nn.ctc_loss(\n            labels=labels,\n            logits=logits,\n            label_length=label_length,\n            logit_length=logit_length,\n            blank_index=pad_token_idx,\n            logits_time_major=False\n        )\n    '''\n    loss = tf.reduce_mean(loss)\n    return loss","metadata":{"execution":{"iopub.status.busy":"2023-07-23T10:32:32.410633Z","iopub.status.idle":"2023-07-23T10:32:32.411263Z","shell.execute_reply.started":"2023-07-23T10:32:32.411059Z","shell.execute_reply":"2023-07-23T10:32:32.411082Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_model(dim = 384, dropout_step=0):\n    with strategy.scope():\n        inp = tf.keras.Input(INPUT_SHAPE)\n        x = inp\n        \n        x = tf.keras.layers.Masking(mask_value=0.0)(x)\n        #x = tf.keras.layers.Dense(dim, use_bias=False,name='stem_conv')(x)\n        x = tf.keras.layers.Dense(dim, use_bias=False,name='stem_conv')(x) + positional_encoding(INPUT_SHAPE[0], dim)\n        x = tf.keras.layers.BatchNormalization(momentum=0.95,name='stem_bn')(x)\n\n        x = Conv1DBlock(dim,11,drop_rate=0.2)(x)\n        x = Conv1DBlock(dim,5,drop_rate=0.2)(x)\n        x = Conv1DBlock(dim,3,drop_rate=0.2)(x)\n        x = TransformerBlock(dim,expand=2)(x)\n        \n        x = Conv1DBlock(dim,11,drop_rate=0.2)(x)\n        x = Conv1DBlock(dim,5,drop_rate=0.2)(x)\n        x = Conv1DBlock(dim,3,drop_rate=0.2)(x)\n        x = TransformerBlock(dim,expand=2)(x)\n        \n        x = Conv1DBlock(dim,11,drop_rate=0.2)(x)\n        x = Conv1DBlock(dim,5,drop_rate=0.2)(x)\n        x = Conv1DBlock(dim,3,drop_rate=0.2)(x)\n        x = TransformerBlock(dim,expand=2)(x)\n        \n        x = tf.keras.layers.Dense(dim*2,activation='relu',name='top_conv')(x)\n        x = tf.keras.layers.Dropout(0.4)(x)\n        x = tf.keras.layers.Dense(len(char_to_num))(x)\n\n        model = tf.keras.Model(inp, x)\n\n        loss = CTCLoss\n\n        # Adam Optimizer\n        optimizer = tfa.optimizers.RectifiedAdam(sma_threshold=4)\n        optimizer = tfa.optimizers.Lookahead(optimizer, sync_period=5)\n\n        model.compile(loss=loss, optimizer=optimizer)\n\n        return model\n\ntf.keras.backend.clear_session()\n\nmodel = get_model()\nmodel(batch[0])\nmodel.summary()","metadata":{"execution":{"iopub.status.busy":"2023-07-23T10:32:32.412496Z","iopub.status.idle":"2023-07-23T10:32:32.412905Z","shell.execute_reply.started":"2023-07-23T10:32:32.412703Z","shell.execute_reply":"2023-07-23T10:32:32.412721Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def num_to_char_fn(y):\n    return [num_to_char.get(x, \"\") for x in y]\n\n@tf.function()\ndef decode_phrase(pred):\n    x = tf.argmax(pred, axis=1)\n    diff = tf.not_equal(x[:-1], x[1:])\n    adjacent_indices = tf.where(diff)[:, 0]\n    x = tf.gather(x, adjacent_indices)\n    mask = x != pad_token_idx\n    x = tf.boolean_mask(x, mask, axis=0)\n    return x\n\n# A utility function to decode the output of the network\ndef decode_batch_predictions(pred):\n    output_text = []\n    for result in pred:\n        result = \"\".join(num_to_char_fn(decode_phrase(result).numpy()))\n        output_text.append(result)\n    return output_text","metadata":{"execution":{"iopub.status.busy":"2023-07-23T10:32:32.414181Z","iopub.status.idle":"2023-07-23T10:32:32.414581Z","shell.execute_reply.started":"2023-07-23T10:32:32.41439Z","shell.execute_reply":"2023-07-23T10:32:32.414408Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n# A callback class to output a few transcriptions during training\nclass CallbackEval(tf.keras.callbacks.Callback):\n    \"\"\"Displays a batch of outputs after every epoch.\"\"\"\n\n    def __init__(self, dataset):\n        super().__init__()\n        self.dataset = dataset\n\n    def on_epoch_end(self, epoch: int, logs=None):\n        model.save_weights(\"model.h5\")\n        predictions = []\n        targets = []\n        for batch in self.dataset:\n            X, y = batch\n            batch_predictions = model(X)\n            batch_predictions = decode_batch_predictions(batch_predictions)\n            predictions.extend(batch_predictions)\n            for label in y:\n                label = \"\".join(num_to_char_fn(label.numpy()))\n                targets.append(label)\n        print(\"-\" * 100)\n        # for i in np.random.randint(0, len(predictions), 2):\n        for i in range(32):\n            print(f\"Target    : {targets[i]}\")\n            print(f\"Prediction: {predictions[i]}, len: {len(predictions[i])}\")\n            print(\"-\" * 100)\n\n# Callback function to check transcription on the val set.\nvalidation_callback = CallbackEval(val_dataset.take(1))","metadata":{"execution":{"iopub.status.busy":"2023-07-23T10:32:32.415917Z","iopub.status.idle":"2023-07-23T10:32:32.416306Z","shell.execute_reply.started":"2023-07-23T10:32:32.416122Z","shell.execute_reply":"2023-07-23T10:32:32.41614Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"N_EPOCHS = 50\nN_WARMUP_EPOCHS = 10\nLR_MAX = 1e-3\nWD_RATIO = 0.05\nWARMUP_METHOD = \"exp\"","metadata":{"execution":{"iopub.status.busy":"2023-07-23T10:32:32.418211Z","iopub.status.idle":"2023-07-23T10:32:32.419035Z","shell.execute_reply.started":"2023-07-23T10:32:32.418788Z","shell.execute_reply":"2023-07-23T10:32:32.418812Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### This is for calculating the validation set' Levenshtein distance during training","metadata":{}},{"cell_type":"code","source":"val_set = [x for x in val_dataset]\n\nwith open (\"/kaggle/input/asl-fingerspelling/character_to_prediction_index.json\", \"r\") as f:\n    character_map = json.load(f)\nrev_character_map = {j:i for i,j in character_map.items()}\n\nclass val_lev_callback(tf.keras.callbacks.Callback):\n    def __init__(self):\n        super().__init__()\n    def on_epoch_end(self, epoch: int, logs=None):\n        calculate_val_lev()\n        \ndef calculate_val_lev():\n    preds = []\n    targets = []\n    scores = []\n    for batch_idx in range(len(val_set)):\n        preds_batch = model.predict(val_set[batch_idx][0], verbose = 0)\n        targets_batch = val_set[batch_idx][1]\n        for pred_idx in range(len(preds_batch)):\n            preds.append(\"\".join([rev_character_map.get(s, \"\") for s in decode_phrase(preds_batch[pred_idx]).numpy()]))\n            targets.append(\"\".join([rev_character_map.get(s, \"\") for s in targets_batch[pred_idx].numpy()]))\n\n    N = [len(phrase) for phrase in targets]\n    lev_dist = [lev.distance(preds[i], targets[i]) for i in range(len(targets))]\n    print('Lev distance: '+str((np.sum(N) - np.sum(lev_dist))/np.sum(N)))","metadata":{"execution":{"iopub.status.busy":"2023-07-23T10:32:32.420393Z","iopub.status.idle":"2023-07-23T10:32:32.42079Z","shell.execute_reply.started":"2023-07-23T10:32:32.420609Z","shell.execute_reply":"2023-07-23T10:32:32.420628Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def lrfn(current_step, num_warmup_steps, lr_max, num_cycles=0.50, num_training_steps=N_EPOCHS):\n    \n    if current_step < num_warmup_steps:\n        if WARMUP_METHOD == 'log':\n            return lr_max * 0.10 ** (num_warmup_steps - current_step)\n        else:\n            return lr_max * 2 ** -(num_warmup_steps - current_step)\n    else:\n        progress = float(current_step - num_warmup_steps) / float(max(1, num_training_steps - num_warmup_steps))\n\n        return max(0.0, 0.5 * (1.0 + math.cos(math.pi * float(num_cycles) * 2.0 * progress))) * lr_max\n    \ndef plot_lr_schedule(lr_schedule, epochs):\n    fig = plt.figure(figsize=(20, 10))\n    plt.plot([None] + lr_schedule + [None])\n    # X Labels\n    x = np.arange(1, epochs + 1)\n    x_axis_labels = [i if epochs <= 40 or i % 5 == 0 or i == 1 else None for i in range(1, epochs + 1)]\n    plt.xlim([1, epochs])\n    plt.xticks(x, x_axis_labels) # set tick step to 1 and let x axis start at 1\n    \n    # Increase y-limit for better readability\n    plt.ylim([0, max(lr_schedule) * 1.1])\n    \n    # Title\n    schedule_info = f'start: {lr_schedule[0]:.1E}, max: {max(lr_schedule):.1E}, final: {lr_schedule[-1]:.1E}'\n    plt.title(f'Step Learning Rate Schedule, {schedule_info}', size=18, pad=12)\n    \n    # Plot Learning Rates\n    for x, val in enumerate(lr_schedule):\n        if epochs <= 40 or x % 5 == 0 or x is epochs - 1:\n            if x < len(lr_schedule) - 1:\n                if lr_schedule[x - 1] < val:\n                    ha = 'right'\n                else:\n                    ha = 'left'\n            elif x == 0:\n                ha = 'right'\n            else:\n                ha = 'left'\n            plt.plot(x + 1, val, 'o', color='black');\n            offset_y = (max(lr_schedule) - min(lr_schedule)) * 0.02\n            plt.annotate(f'{val:.1E}', xy=(x + 1, val + offset_y), size=12, ha=ha)\n    \n    plt.xlabel('Epoch', size=16, labelpad=5)\n    plt.ylabel('Learning Rate', size=16, labelpad=5)\n    plt.grid()\n    plt.show()\n\n# Learning rate for encoder\nLR_SCHEDULE = [lrfn(step, num_warmup_steps=N_WARMUP_EPOCHS, lr_max=LR_MAX, num_cycles=0.50) for step in range(N_EPOCHS)]\n# Plot Learning Rate Schedule\nplot_lr_schedule(LR_SCHEDULE, epochs=N_EPOCHS)\n# Learning Rate Callback\nlr_callback = tf.keras.callbacks.LearningRateScheduler(lambda step: LR_SCHEDULE[step], verbose=0)\n\n# Custom callback to update weight decay with learning rate\nclass WeightDecayCallback(tf.keras.callbacks.Callback):\n    def __init__(self, wd_ratio=WD_RATIO):\n        self.step_counter = 0\n        self.wd_ratio = wd_ratio\n    \n    def on_epoch_begin(self, epoch, logs=None):\n        model.optimizer.weight_decay = model.optimizer.learning_rate * self.wd_ratio\n        print(f'learning rate: {model.optimizer.learning_rate.numpy():.2e}, weight decay: {model.optimizer.weight_decay.numpy():.2e}')\n\nclass save_model_callback(tf.keras.callbacks.Callback):\n    def __init__(self):\n        super().__init__()\n    def on_epoch_end(self, epoch: int, logs=None):\n        if (epoch+1)%5 == 0:\n            self.model.save_weights(f\"model_epoch_{epoch}.h5\")","metadata":{"execution":{"iopub.status.busy":"2023-07-23T10:32:32.422292Z","iopub.status.idle":"2023-07-23T10:32:32.422671Z","shell.execute_reply.started":"2023-07-23T10:32:32.422491Z","shell.execute_reply":"2023-07-23T10:32:32.42251Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history = model.fit(\n    train_dataset,\n    validation_data=val_dataset,\n    epochs=50,\n    verbose = 2,\n    callbacks=[\n        save_model_callback(),\n        lr_callback,\n        WeightDecayCallback(),\n        val_lev_callback(),\n    ]\n)","metadata":{"execution":{"iopub.status.busy":"2023-07-23T10:32:32.425666Z","iopub.status.idle":"2023-07-23T10:32:32.426081Z","shell.execute_reply.started":"2023-07-23T10:32:32.425886Z","shell.execute_reply":"2023-07-23T10:32:32.425907Z"},"trusted":true},"execution_count":null,"outputs":[]}]}