{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install -q tensorflow-addons==0.20.0\n!pip install -q git+https://github.com/hoyso48/tf-utils@main","metadata":{"execution":{"iopub.status.busy":"2023-07-07T07:58:51.907641Z","iopub.execute_input":"2023-07-07T07:58:51.908658Z","iopub.status.idle":"2023-07-07T07:59:17.037234Z","shell.execute_reply.started":"2023-07-07T07:58:51.908621Z","shell.execute_reply":"2023-07-07T07:59:17.036012Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport matplotlib as mpl\nimport seaborn as sn\nimport pyarrow.parquet as pq\nfrom tqdm.notebook import tqdm\n\nimport tensorflow as tf\nfrom tensorflow import keras\nfrom tensorflow.keras import layers\nimport tensorflow_addons as tfa\nfrom tf_utils.schedules import OneCycleLR\n\nfrom leven import levenshtein\nfrom skimage.transform import resize\n\nimport glob\nimport sys\nimport os\nimport math\nimport gc\nimport sys\nimport sklearn\nimport time\nimport json\nimport random\n\ntqdm.pandas()\n\nprint(f'Tensorflow Version {tf.__version__}')\nprint(f'Python Version: {sys.version}')","metadata":{"execution":{"iopub.status.busy":"2023-07-07T08:01:40.248273Z","iopub.execute_input":"2023-07-07T08:01:40.248738Z","iopub.status.idle":"2023-07-07T08:01:40.264434Z","shell.execute_reply.started":"2023-07-07T08:01:40.248701Z","shell.execute_reply":"2023-07-07T08:01:40.263382Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Configs","metadata":{}},{"cell_type":"code","source":"CONFIG = {\"exp_name\": 'test',\n          \"lr\": 0.001,\n          \"lr_min\": 1e-6,\n          \"weight_decay\": 0.1,\n          \"warmup\": 0,\n          \"device\": \"GPU\", # CPU, GPU, TPU, TPU-VM\n          \"max_len\" : 128,\n          \"epoch\": 30,\n          \"use_valid\": True,\n          \"valid_size\": 0.2,\n          \"batch_size\": 64,\n          \"seeds\": [42],\n          \"fold\": \"part\", # \"all\", \"part\"\n          \"float16\": True}\n\nroot = '/kaggle/input/asl-fingerspelling'\nroot_data = '/kaggle/input/asl-tfr'\nroot_save = '/kaggle/working'\nprint(CONFIG[\"exp_name\"])","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-07-07T08:01:40.283031Z","iopub.execute_input":"2023-07-07T08:01:40.283731Z","iopub.status.idle":"2023-07-07T08:01:40.295299Z","shell.execute_reply.started":"2023-07-07T08:01:40.283697Z","shell.execute_reply":"2023-07-07T08:01:40.294435Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(seed=42):\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    random.seed(seed)\n    np.random.seed(seed)\n    tf.random.set_seed(seed)\n    \ndef get_strategy(device='TPU-VM'):\n    IS_TPU = False\n    if \"TPU\" in device:\n        tpu = 'local' if device=='TPU-VM' else None\n        print(\"connecting to TPU...\")\n        tpu = tf.distribute.cluster_resolver.TPUClusterResolver.connect(tpu=tpu)\n        strategy = tf.distribute.TPUStrategy(tpu)\n        IS_TPU = True\n\n    if device == \"GPU\"  or device==\"CPU\":\n        ngpu = len(tf.config.experimental.list_physical_devices('GPU'))\n        if ngpu>1:\n            print(\"Using multi GPU\")\n            strategy = tf.distribute.MirroredStrategy()\n        elif ngpu==1:\n            print(\"Using single GPU\")\n            strategy = tf.distribute.get_strategy()\n        else:\n            print(\"Using CPU\")\n            strategy = tf.distribute.get_strategy()\n\n    if device == \"GPU\":\n        print(\"Num GPUs Available: \", ngpu)\n\n    AUTO     = tf.data.experimental.AUTOTUNE\n    REPLICAS = strategy.num_replicas_in_sync\n    print(f'REPLICAS: {REPLICAS}')\n    \n    return strategy, REPLICAS, IS_TPU\n\nSTRATEGY, N_REPLICAS, IS_TPU = get_strategy(CONFIG[\"device\"])","metadata":{"execution":{"iopub.status.busy":"2023-07-07T08:01:40.324679Z","iopub.execute_input":"2023-07-07T08:01:40.325395Z","iopub.status.idle":"2023-07-07T08:01:40.343891Z","shell.execute_reply.started":"2023-07-07T08:01:40.325358Z","shell.execute_reply":"2023-07-07T08:01:40.343005Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Process Data","metadata":{}},{"cell_type":"code","source":"train = pd.read_csv(f'{root}/train.csv')\ntrain_supplemental = pd.read_csv(f'{root}/supplemental_metadata.csv')\nNUM_DATA = len(train)","metadata":{"execution":{"iopub.status.busy":"2023-07-07T08:01:40.403795Z","iopub.execute_input":"2023-07-07T08:01:40.405937Z","iopub.status.idle":"2023-07-07T08:01:40.63143Z","shell.execute_reply.started":"2023-07-07T08:01:40.405904Z","shell.execute_reply":"2023-07-07T08:01:40.630363Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define Features\nLIP = [\n    61, 185, 40, 39, 37, 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]\n\nLPOSE = [13, 15, 17, 19, 21]\nRPOSE = [14, 16, 18, 20, 22]\nPOSE = LPOSE + RPOSE\n\nFACE = [f'x_face_{i}' for i in LIP] + [f'y_face_{i}' for i in LIP] + [f'z_face_{i}' for i in LIP]\n\nLHAND = [f'x_left_hand_{i}' for i in range(21)] + [f'y_left_hand_{i}' for i in range(21)] + [f'z_left_hand_{i}' for i in range(21)]\nRHAND = [f'x_right_hand_{i}' for i in range(21)] + [f'y_right_hand_{i}' for i in range(21)] + [f'z_right_hand_{i}' for i in range(21)]\nPOSE = [f'x_pose_{i}' for i in POSE] + [f'y_pose_{i}' for i in POSE] + [f'z_pose_{i}' for i in POSE]\n\nFEATURE_COL = FACE + LHAND + RHAND + POSE\nFRAME_LEN = CONFIG['max_len']\nprint(len(FEATURE_COL))","metadata":{"execution":{"iopub.status.busy":"2023-07-07T08:01:40.636389Z","iopub.execute_input":"2023-07-07T08:01:40.638522Z","iopub.status.idle":"2023-07-07T08:01:40.655681Z","shell.execute_reply.started":"2023-07-07T08:01:40.638487Z","shell.execute_reply":"2023-07-07T08:01:40.654631Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_IDX = [i for i, col in enumerate(FEATURE_COL) if \"x_\" in col]\nY_IDX = [i for i, col in enumerate(FEATURE_COL) if \"y_\" in col]\nZ_IDX = [i for i, col in enumerate(FEATURE_COL) if \"z_\" in col]\n\nRHAND_IDX = [i for i, col in enumerate(FEATURE_COL) if \"right\" in col]\nLHAND_IDX = [i for i, col in enumerate(FEATURE_COL) if \"left\" in col]\nRPOSE_IDX = [i for i, col in enumerate(FEATURE_COL) if \"pose\" in col and int(col[-2:]) in RPOSE]\nLPOSE_IDX = [i for i, col in enumerate(FEATURE_COL) if \"pose\" in col and int(col[-2:]) in LPOSE]","metadata":{"execution":{"iopub.status.busy":"2023-07-07T08:01:40.660864Z","iopub.execute_input":"2023-07-07T08:01:40.663483Z","iopub.status.idle":"2023-07-07T08:01:40.674342Z","shell.execute_reply.started":"2023-07-07T08:01:40.663449Z","shell.execute_reply":"2023-07-07T08:01:40.673447Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with open(f'{root}/character_to_prediction_index.json') as json_file:\n    CHAR2ORD = json.load(json_file)\nORD2CHAR = {j:i for i,j in CHAR2ORD.items()}\n# Character to Ordinal Encoding Mapping   \n# display(pd.Series(CHAR2ORD).to_frame('Ordinal Encoding'))\n\nPAD = -100\nPAD_TOKEN = len(CHAR2ORD)\nSOS_TOKEN = len(CHAR2ORD) + 1 # Start Of Sentence\nEOS_TOKEN = len(CHAR2ORD) + 2 # End Of Sentence","metadata":{"execution":{"iopub.status.busy":"2023-07-07T08:01:40.679955Z","iopub.execute_input":"2023-07-07T08:01:40.682613Z","iopub.status.idle":"2023-07-07T08:01:40.690879Z","shell.execute_reply.started":"2023-07-07T08:01:40.682578Z","shell.execute_reply":"2023-07-07T08:01:40.69Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# TODO\n# still the same as starter\ndef resize_pad(x, max_len):\n    if tf.shape(x)[0] < max_len:\n        x = tf.pad(x, ([[0, max_len-tf.shape(x)[0]], [0, 0], [0, 0]]))\n    else:\n        x = tf.image.resize(x, (max_len, tf.shape(x)[1]))\n    return x\n\n# Detect the dominant hand from the number of NaN values.\n# Dominant hand will have less NaN values since it is in frame moving.\ndef preprocess(landmarks, phrase, augment=True, max_len=64):\n    x = landmarks\n    rhand = tf.gather(x, RHAND_IDX, axis=1)\n    lhand = tf.gather(x, LHAND_IDX, axis=1)\n    rpose = tf.gather(x, RPOSE_IDX, axis=1)\n    lpose = tf.gather(x, LPOSE_IDX, axis=1)\n    \n    rnan_idx = tf.reduce_any(tf.math.is_nan(rhand), axis=1)\n    lnan_idx = tf.reduce_any(tf.math.is_nan(lhand), axis=1)\n    \n    rnans = tf.math.count_nonzero(rnan_idx)\n    lnans = tf.math.count_nonzero(lnan_idx)\n    \n    # For dominant hand\n    if rnans > lnans:\n        hand = lhand\n        pose = lpose\n        \n        hand_x = hand[:, 0*(len(LHAND_IDX)//3) : 1*(len(LHAND_IDX)//3)]\n        hand_y = hand[:, 1*(len(LHAND_IDX)//3) : 2*(len(LHAND_IDX)//3)]\n        hand_z = hand[:, 2*(len(LHAND_IDX)//3) : 3*(len(LHAND_IDX)//3)]\n        hand = tf.concat([1-hand_x, hand_y, hand_z], axis=1)\n        \n        pose_x = pose[:, 0*(len(LPOSE_IDX)//3) : 1*(len(LPOSE_IDX)//3)]\n        pose_y = pose[:, 1*(len(LPOSE_IDX)//3) : 2*(len(LPOSE_IDX)//3)]\n        pose_z = pose[:, 2*(len(LPOSE_IDX)//3) : 3*(len(LPOSE_IDX)//3)]\n        pose = tf.concat([1-pose_x, pose_y, pose_z], axis=1)\n    else:\n        hand = rhand\n        pose = rpose\n    \n    hand_x = hand[:, 0*(len(LHAND_IDX)//3) : 1*(len(LHAND_IDX)//3)]\n    hand_y = hand[:, 1*(len(LHAND_IDX)//3) : 2*(len(LHAND_IDX)//3)]\n    hand_z = hand[:, 2*(len(LHAND_IDX)//3) : 3*(len(LHAND_IDX)//3)]\n    hand = tf.concat([hand_x[..., tf.newaxis], hand_y[..., tf.newaxis], hand_z[..., tf.newaxis]], axis=-1)\n    \n    mean = tf.math.reduce_mean(hand, axis=1)[:, tf.newaxis, :]\n    std = tf.math.reduce_std(hand, axis=1)[:, tf.newaxis, :]\n    hand = (hand - mean) / std\n\n    pose_x = pose[:, 0*(len(LPOSE_IDX)//3) : 1*(len(LPOSE_IDX)//3)]\n    pose_y = pose[:, 1*(len(LPOSE_IDX)//3) : 2*(len(LPOSE_IDX)//3)]\n    pose_z = pose[:, 2*(len(LPOSE_IDX)//3) : 3*(len(LPOSE_IDX)//3)]\n    pose = tf.concat([pose_x[..., tf.newaxis], pose_y[..., tf.newaxis], pose_z[..., tf.newaxis]], axis=-1)\n    \n    x = tf.concat([hand, pose], axis=1)\n    x = resize_pad(x, max_len)\n    \n    x = tf.where(tf.math.is_nan(x), tf.zeros_like(x), x)\n    x = tf.reshape(x, (FRAME_LEN, len(LHAND_IDX) + len(LPOSE_IDX)))\n    return x, phrase","metadata":{"execution":{"iopub.status.busy":"2023-07-07T08:01:40.696463Z","iopub.execute_input":"2023-07-07T08:01:40.698558Z","iopub.status.idle":"2023-07-07T08:01:40.726541Z","shell.execute_reply.started":"2023-07-07T08:01:40.698527Z","shell.execute_reply":"2023-07-07T08:01:40.725678Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"table = tf.lookup.StaticHashTable(\n    initializer=tf.lookup.KeyValueTensorInitializer(\n        keys=list(CHAR2ORD.keys()),\n        values=list(CHAR2ORD.values()),\n    ),\n    default_value=tf.constant(-1),\n    name=\"class_weight\"\n)\n\ndef decode_fn(record_bytes):\n    schema = {COL: tf.io.VarLenFeature(dtype=tf.float32) for COL in FEATURE_COL}\n    schema[\"phrase\"] = tf.io.FixedLenFeature([], dtype=tf.string)\n    features = tf.io.parse_single_example(record_bytes, schema)\n    phrase = features[\"phrase\"]\n    landmarks = ([tf.sparse.to_dense(features[COL]) for COL in FEATURE_COL])\n    # Transpose to maintain the original shape of landmarks data.\n    landmarks = tf.transpose(landmarks)\n    \n    return landmarks, phrase\n\ndef convert_fn(landmarks, phrase):\n    # Add start and end pointers to phrase.\n    phrase = tf.strings.bytes_split(phrase)\n    phrase = table.lookup(phrase)\n    # Vectorize and add padding.\n    phrase = tf.pad(phrase, paddings=[[0, 62 - tf.shape(phrase)[0]]], mode = 'CONSTANT',\n                    constant_values = PAD_TOKEN)\n    t1 = [SOS_TOKEN]\n    t2 = [EOS_TOKEN]\n    phrase = tf.concat([t1, phrase, t2], axis=0)\n\n    return landmarks, phrase","metadata":{"execution":{"iopub.status.busy":"2023-07-07T08:01:40.7326Z","iopub.execute_input":"2023-07-07T08:01:40.73522Z","iopub.status.idle":"2023-07-07T08:01:40.754299Z","shell.execute_reply.started":"2023-07-07T08:01:40.735189Z","shell.execute_reply":"2023-07-07T08:01:40.753468Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# see tfr data\nfor file_id in train.file_id:\n    tffile = f\"/kaggle/input/asl-tfr/supplemental/1032110484.tfrecord\"\n    for batch in tf.data.TFRecordDataset([tffile]).map(decode_fn).take(5):\n        print(list(batch)[0].shape, list(batch)[1])\n    break","metadata":{"execution":{"iopub.status.busy":"2023-07-07T08:01:40.759589Z","iopub.execute_input":"2023-07-07T08:01:40.762021Z","iopub.status.idle":"2023-07-07T08:01:41.682493Z","shell.execute_reply.started":"2023-07-07T08:01:40.761985Z","shell.execute_reply":"2023-07-07T08:01:41.681378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Get Dataset","metadata":{}},{"cell_type":"code","source":"def get_tfrec_dataset(tfrecords, batch_size=64, max_len=64, drop_remainder=False, augment=False, shuffle=False, repeat=False):\n    print(\"## Getting Dataset\")\n    ds = tf.data.TFRecordDataset(tfrecords, num_parallel_reads=tf.data.AUTOTUNE)\n    ds = ds.map(decode_fn, tf.data.AUTOTUNE)\n    ds = ds.map(convert_fn, tf.data.AUTOTUNE)\n    ds = ds.map(lambda landmarks, phrase: preprocess(landmarks, phrase, augment=augment, max_len=max_len), tf.data.AUTOTUNE)\n\n    if repeat: \n        ds = ds.repeat()\n        \n    if shuffle:\n        ds = ds.shuffle(shuffle)\n        options = tf.data.Options()\n        options.experimental_deterministic = (False)\n        ds = ds.with_options(options)\n    \n    if batch_size:\n        # ds = ds.padded_batch(batch_size, padding_values=PAD, padded_shapes=([CONFIG[\"max_len\"], len(FEATURE_COL)],[64]), drop_remainder=drop_remainder)\n        ds = ds.batch(batch_size).prefetch(tf.data.AUTOTUNE).cache()\n\n    ds = ds.prefetch(tf.data.AUTOTUNE)\n    print(\"## Done\")\n    print()\n        \n    return ds","metadata":{"execution":{"iopub.status.busy":"2023-07-07T08:01:41.683805Z","iopub.execute_input":"2023-07-07T08:01:41.68416Z","iopub.status.idle":"2023-07-07T08:01:41.693774Z","shell.execute_reply.started":"2023-07-07T08:01:41.684134Z","shell.execute_reply":"2023-07-07T08:01:41.69284Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TRAIN_FILENAMES = train.file_id.map(lambda x: f'{root_data}/train/{x}.tfrecord').unique()\nTRAIN_SUPPLEMENTAL_FILENAMES = train_supplemental.file_id.map(lambda x: f'{root_data}/supplemental/{x}.tfrecord').unique()\nprint(f\"Train: {len(TRAIN_FILENAMES)} TFRecord files.\")\nprint(f\"Train Supplemental: {len(TRAIN_SUPPLEMENTAL_FILENAMES)} TFRecord files.\")","metadata":{"execution":{"iopub.status.busy":"2023-07-07T08:01:41.697602Z","iopub.execute_input":"2023-07-07T08:01:41.698651Z","iopub.status.idle":"2023-07-07T08:01:41.791802Z","shell.execute_reply.started":"2023-07-07T08:01:41.698597Z","shell.execute_reply":"2023-07-07T08:01:41.790759Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = get_tfrec_dataset(TRAIN_FILENAMES, \n                            batch_size=CONFIG[\"batch_size\"]*N_REPLICAS, \n                            max_len=CONFIG[\"max_len\"], \n                            drop_remainder=True, augment=True, \n                            repeat=True, shuffle=32768)\ndataset_supplemental = get_tfrec_dataset(TRAIN_SUPPLEMENTAL_FILENAMES, \n                                         batch_size=CONFIG[\"batch_size\"]*N_REPLICAS, \n                                         max_len=CONFIG[\"max_len\"], \n                                         drop_remainder=True, augment=True, \n                                         repeat=True, shuffle=32768)\n\nfor batch in dataset.take(1):\n    print(list(batch)[0][0].shape, list(batch)[1][0])","metadata":{"execution":{"iopub.status.busy":"2023-07-07T08:01:41.79306Z","iopub.execute_input":"2023-07-07T08:01:41.793401Z","iopub.status.idle":"2023-07-07T08:02:54.317896Z","shell.execute_reply.started":"2023-07-07T08:01:41.793369Z","shell.execute_reply":"2023-07-07T08:02:54.316841Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CONFIG[\"use_valid\"]:\n    len_train = int(len(TRAIN_FILENAMES) * (1 - CONFIG[\"valid_size\"]))\n    train_dataset = get_tfrec_dataset(TRAIN_FILENAMES[:len_train], \n                                      batch_size=CONFIG[\"batch_size\"]*N_REPLICAS, \n                                      max_len=CONFIG[\"max_len\"], \n                                      drop_remainder=True, augment=True, \n                                      repeat=True, shuffle=32768)\n    valid_dataset = get_tfrec_dataset(TRAIN_FILENAMES[len_train:], \n                                      batch_size=CONFIG[\"batch_size\"]*N_REPLICAS, \n                                      max_len=CONFIG[\"max_len\"], \n                                      drop_remainder=False, \n                                      repeat=False, \n                                      shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2023-07-07T08:02:54.319651Z","iopub.execute_input":"2023-07-07T08:02:54.320039Z","iopub.status.idle":"2023-07-07T08:02:55.60404Z","shell.execute_reply.started":"2023-07-07T08:02:54.320002Z","shell.execute_reply":"2023-07-07T08:02:55.603062Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Utils","metadata":{}},{"cell_type":"code","source":"def outputs2phrase(outputs):\n    if outputs.ndim == 2:\n        outputs = np.argmax(outputs, axis=1)\n    \n    return ''.join([ORD2CHAR.get(s, '') for s in outputs])\n\ndef levenhstein_distance(labels, outputs):\n    edit_dist = tf.edit_distance(tf.sparse.from_dense(labels), \n                                 tf.sparse.from_dense(tf.cast(tf.argmax(outputs, axis=1), tf.int32)))\n    edit_dist = tf.reduce_mean(edit_dist)\n    return edit_dist","metadata":{"execution":{"iopub.status.busy":"2023-07-07T08:02:55.606395Z","iopub.execute_input":"2023-07-07T08:02:55.607046Z","iopub.status.idle":"2023-07-07T08:02:55.613693Z","shell.execute_reply.started":"2023-07-07T08:02:55.607011Z","shell.execute_reply":"2023-07-07T08:02:55.612686Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Get Model","metadata":{}},{"cell_type":"markdown","source":"#### Embedding","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 = tf.keras.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\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)","metadata":{"execution":{"iopub.status.busy":"2023-07-07T08:02:55.615048Z","iopub.execute_input":"2023-07-07T08:02:55.615732Z","iopub.status.idle":"2023-07-07T08:02:55.628006Z","shell.execute_reply.started":"2023-07-07T08:02:55.6157Z","shell.execute_reply":"2023-07-07T08:02:55.626978Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Encoder","metadata":{}},{"cell_type":"code","source":"class 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            [\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_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)","metadata":{"execution":{"iopub.status.busy":"2023-07-07T08:02:55.629267Z","iopub.execute_input":"2023-07-07T08:02:55.62988Z","iopub.status.idle":"2023-07-07T08:02:55.641675Z","shell.execute_reply.started":"2023-07-07T08:02:55.629848Z","shell.execute_reply":"2023-07-07T08:02:55.64063Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Decoder","metadata":{}},{"cell_type":"code","source":"# Customized to add `training` variable\n# Reference: https://www.kaggle.com/code/shlomoron/aslfr-a-simple-transformer/notebook\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            [\n                layers.Dense(feed_forward_dim, activation=\"relu\"),\n                layers.Dense(embed_dim),\n            ]\n        )\n\n    def causal_attention_mask(self, batch_size, n_dest, n_src, dtype):\n        \"\"\"Masks the upper half of the dot product matrix in self attention.\n\n        This prevents flow of information from future tokens to current token.\n        1's in the lower triangle, counting from the lower right corner.\n        \"\"\"\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        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        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        ffn_out = self.ffn(enc_out_norm)\n        ffn_out_norm = self.layernorm3(enc_out_norm + self.ffn_dropout(ffn_out, training = training))\n        return ffn_out_norm","metadata":{"execution":{"iopub.status.busy":"2023-07-07T08:02:55.64279Z","iopub.execute_input":"2023-07-07T08:02:55.643122Z","iopub.status.idle":"2023-07-07T08:02:55.65815Z","shell.execute_reply.started":"2023-07-07T08:02:55.643092Z","shell.execute_reply":"2023-07-07T08:02:55.657071Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Transformer","metadata":{}},{"cell_type":"code","source":"class 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, pad_token_idx))\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, pad_token_idx))\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","metadata":{"execution":{"iopub.status.busy":"2023-07-07T08:02:55.66021Z","iopub.execute_input":"2023-07-07T08:02:55.660722Z","iopub.status.idle":"2023-07-07T08:02:55.68655Z","shell.execute_reply.started":"2023-07-07T08:02:55.660684Z","shell.execute_reply":"2023-07-07T08:02:55.685902Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training","metadata":{}},{"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    return 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)","metadata":{"execution":{"iopub.status.busy":"2023-07-07T08:02:55.687935Z","iopub.execute_input":"2023-07-07T08:02:55.688641Z","iopub.status.idle":"2023-07-07T08:02:55.699185Z","shell.execute_reply.started":"2023-07-07T08:02:55.688605Z","shell.execute_reply":"2023-07-07T08:02:55.698101Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_fold(seed, train_dataset, valid_dataset, strategy):\n    # set seed\n    print(\"## Seeding\")\n    seed_everything(seed=seed)\n    tf.keras.backend.clear_session()\n    gc.collect()\n    tf.config.optimizer.set_jit(True)\n\n    if CONFIG['float16']:\n        try:\n            policy = tf.keras.mixed_precision.Policy('mixed_bfloat16')\n            tf.keras.mixed_precision.set_global_policy(policy)\n        except:\n            policy = tf.keras.mixed_precision.Policy('mixed_float16')\n            tf.keras.mixed_precision.set_global_policy(policy)\n    else:\n        policy = tf.keras.mixed_precision.Policy('float32')\n        tf.keras.mixed_precision.set_global_policy(policy)\n    \n    num_train = (NUM_DATA * (1 - CONFIG[\"valid_size\"]))\n    num_valid = (NUM_DATA * (CONFIG[\"valid_size\"]))\n    steps_per_epoch = num_train//CONFIG[\"batch_size\"]\n    \n    with strategy.scope():\n        print(\"## Getting Model\")\n        model = Transformer(num_hid=200,\n                            num_head=4,\n                            num_feed_forward=400,\n                            source_maxlen = FRAME_LEN,\n                            target_maxlen=64,\n                            num_layers_enc=2,\n                            num_layers_dec=1,\n                            num_classes=62)\n        \n        schedule = OneCycleLR(CONFIG[\"lr\"], CONFIG[\"epoch\"], \n                              warmup_epochs=CONFIG[\"epoch\"]*CONFIG[\"warmup\"], \n                              steps_per_epoch=steps_per_epoch, decay_epochs=CONFIG[\"epoch\"], \n                              lr_min=CONFIG[\"lr_min\"], \n                              decay_type='cosine', \n                              warmup_type='linear')\n        decay_schedule = OneCycleLR(CONFIG[\"lr\"]*CONFIG[\"weight_decay\"], CONFIG[\"epoch\"], \n                                    warmup_epochs=CONFIG[\"epoch\"]*CONFIG[\"warmup\"], \n                                    steps_per_epoch=steps_per_epoch, decay_epochs=CONFIG[\"epoch\"], \n                                    lr_min=CONFIG[\"lr_min\"]*CONFIG[\"weight_decay\"], \n                                    decay_type='cosine', \n                                    warmup_type='linear')\n        \n        optimizer = tfa.optimizers.RectifiedAdam(learning_rate=schedule, weight_decay=decay_schedule, sma_threshold=4)\n        optimizer = tfa.optimizers.Lookahead(optimizer, sync_period=5)\n        \n        model.compile(optimizer=optimizer,\n                      loss=[CTCLoss],\n                      metrics=['levehnstein dist', levenhstein_distance],\n                      steps_per_execution=steps_per_epoch)\n    \n    model.summary()\n    \n    print(f'## start training fold {seed}')\n    print(f'## train:{num_train} valid:{num_valid}')\n    \n    logger = tf.keras.callbacks.CSVLogger(f'{root_save}/fold{seed}-logs.csv')\n    sv_loss = tf.keras.callbacks.ModelCheckpoint(f'{root_save}/fold{seed}-best.h5', monitor='val_loss', verbose=0, save_best_only=True,\n                                                 save_weights_only=True, mode='min', save_freq='epoch')\n    callbacks = []\n    if valid_dataset is not None:\n        callbacks.append(sv_loss)\n    \n    history = model.fit(train_dataset,\n                        epochs=CONFIG[\"epoch\"],\n                        steps_per_epoch=steps_per_epoch,\n                        callbacks=callbacks,\n                        validation_data=valid_dataset,\n                        verbose=2,\n                        validation_steps=-(num_valid//-CONFIG[\"batch_size\"]))\n    try:\n        model.load_weights(f'/kaggle/working/fold{seed}-best.h5')\n    except:\n        pass\n    \n    if val_dataset is not None:\n        val = model.evaluate(valid_dataset, verbose=2, validation_steps=-(num_valid//-CONFIG[\"batch_size\"]))\n    else:\n        val = None\n    \n    return model, cv, history\n    \ndef run_training(seeds=[42], fold=\"all\", strategy=STRATEGY):\n    for seed in seeds:\n        if fold == 'all':\n            train_fold(seed=seed, train_dataset=dataset, valid_dataset=None, strategy=strategy)\n        elif fold == \"part\":\n            train_fold(seed=seed, train_dataset=train_dataset, valid_dataset=valid_dataset, strategy=strategy)\n    return","metadata":{"execution":{"iopub.status.busy":"2023-07-07T08:08:29.985413Z","iopub.execute_input":"2023-07-07T08:08:29.985788Z","iopub.status.idle":"2023-07-07T08:08:30.006618Z","shell.execute_reply.started":"2023-07-07T08:08:29.985759Z","shell.execute_reply":"2023-07-07T08:08:30.005588Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"run_training(seeds=CONFIG['seeds'], fold=CONFIG['fold'], strategy=STRATEGY)","metadata":{"execution":{"iopub.status.busy":"2023-07-07T08:08:30.008752Z","iopub.execute_input":"2023-07-07T08:08:30.009225Z","iopub.status.idle":"2023-07-07T08:08:30.482487Z","shell.execute_reply.started":"2023-07-07T08:08:30.009156Z","shell.execute_reply":"2023-07-07T08:08:30.480911Z"},"trusted":true},"execution_count":null,"outputs":[]}]}