{"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":"Code changes:\n* Added config class for hyperparameters\n* Optimizer: add gradient clipping (`clipnorm=3.0`)\n* Loss: clip target reactivity values to between 0. and 1.\n* Model: add dropout to `MultiHeadAttention` and add option for `norm_first`\nV2:\n* Refactor, add checkpoint save and load\n* Add weight decay to optimizer","metadata":{}},{"cell_type":"code","source":"!nvidia-smi","metadata":{"execution":{"iopub.status.busy":"2023-10-25T03:08:03.474355Z","iopub.execute_input":"2023-10-25T03:08:03.474705Z","iopub.status.idle":"2023-10-25T03:08:04.507944Z","shell.execute_reply.started":"2023-10-25T03:08:03.474677Z","shell.execute_reply":"2023-10-25T03:08:04.506828Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"###---- Environment config ----###\n\n# MACHINE = \"COLAB\"\n# MACHINE = \"JAYOO_PC\"\nMACHINE = \"KAGGLE\"\n\n# device = \"TPU\"\ndevice = \"GPU\"\n# device = \"CPU\"\n\nDEBUG = True\n# DEBUG = False\nif DEBUG == True:\n    print(\"IN DEBUG MODE\")\n    \n# Set root directory\nif MACHINE == \"JAYOO_PC\":\n    ROOT = '/jayoo'  # local\nelif MACHINE == \"COLAB\":\n    ROOT = './drive/MyDrive/colab_env'\n    from google.colab import drive\n    drive.mount('/content/drive')\n    !pwd\nelse:\n    ROOT = ''  # Kaggle\n\nprint(f\"Machine: {MACHINE}, device: {device}, root: {ROOT}\")\n\nimport multiprocessing\nprint(multiprocessing.cpu_count())","metadata":{"execution":{"iopub.status.busy":"2023-10-25T03:08:04.510159Z","iopub.execute_input":"2023-10-25T03:08:04.510561Z","iopub.status.idle":"2023-10-25T03:08:04.524567Z","shell.execute_reply.started":"2023-10-25T03:08:04.510522Z","shell.execute_reply":"2023-10-25T03:08:04.523591Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"!pip install -q tensorflow-addons\n!pip install -q git+https://github.com/hoyso48/tf-utils@main","metadata":{"execution":{"iopub.status.busy":"2023-10-25T03:08:04.525875Z","iopub.execute_input":"2023-10-25T03:08:04.526297Z","iopub.status.idle":"2023-10-25T03:08:31.698485Z","shell.execute_reply.started":"2023-10-25T03:08:04.526265Z","shell.execute_reply":"2023-10-25T03:08:31.697431Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\n\nimport tensorflow as tf\nimport tensorflow_addons as tfa\n\nfrom tf_utils.schedules import OneCycleLR, ListedLR\nfrom tf_utils.callbacks import Snapshot, SWA\nfrom tf_utils.learners import FGM, AWP\n\nimport matplotlib.pyplot as plt\nfrom tqdm.autonotebook import tqdm\nimport sklearn\n\nimport os\nimport time\nimport pickle\nimport math\nimport random\nimport sys\nimport cv2\nimport gc\nimport glob\nimport datetime\nimport shutil\n\nimport ctypes\nlibc = ctypes.CDLL(\"libc.so.6\")","metadata":{"execution":{"iopub.status.busy":"2023-10-25T03:08:31.701481Z","iopub.execute_input":"2023-10-25T03:08:31.702165Z","iopub.status.idle":"2023-10-25T03:08:39.267294Z","shell.execute_reply.started":"2023-10-25T03:08:31.702121Z","shell.execute_reply":"2023-10-25T03:08:39.266521Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# TPU boilerplate code","metadata":{}},{"cell_type":"code","source":"\n# Configure Strategy. Assume TPU...if not set default for GPU\nif \"TPU\" in device:\n        print(\"connecting to TPU...\")\n        if device == 'TPU-VM':  # kaggle\n            tpu = 'local'\n            tpu = tf.distribute.cluster_resolver.TPUClusterResolver.connect(tpu=tpu)\n            strategy = tf.distribute.TPUStrategy(tpu)\n        if device == 'TPU':  # colab\n            tpu = tf.distribute.cluster_resolver.TPUClusterResolver()  # TPU detection\n            tf.config.experimental_connect_to_cluster(tpu)\n            tf.tpu.experimental.initialize_tpu_system(tpu)\n            strategy = tf.distribute.TPUStrategy(tpu)\n\n        IS_TPU = True\n\nif device == \"GPU\"  or device==\"CPU\":\n        IS_TPU = False\n        gpus = tf.config.experimental.list_physical_devices('GPU')\n        try:\n            for gpu in gpus:\n                tf.config.experimental.set_memory_growth(gpu, True)\n        except:\n            pass\n        ngpu = len(gpus)\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\nif device == \"GPU\":\n        print(\"Num GPUs Available: \", ngpu)\n\nAUTO = tf.data.experimental.AUTOTUNE\nREPLICAS = strategy.num_replicas_in_sync\nprint(f'REPLICAS: {REPLICAS}')","metadata":{"execution":{"iopub.status.busy":"2023-10-25T03:08:39.268350Z","iopub.execute_input":"2023-10-25T03:08:39.268842Z","iopub.status.idle":"2023-10-25T03:08:39.624326Z","shell.execute_reply.started":"2023-10-25T03:08:39.268816Z","shell.execute_reply":"2023-10-25T03:08:39.623269Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_name(CFG):\n    name = f\"drop{CFG.drop_rate}_d{CFG.model_dim}_L{CFG.n_layers}_h{CFG.n_heads}_lr{CFG.lr_max}_e{CFG.n_epochs}\"\n    return name\n\ndef clear_mem():\n    gc.collect()\n    libc.malloc_trim(0)\n    tf.keras.backend.clear_session()","metadata":{"execution":{"iopub.status.busy":"2023-10-25T03:08:39.625565Z","iopub.execute_input":"2023-10-25T03:08:39.626263Z","iopub.status.idle":"2023-10-25T03:08:39.633392Z","shell.execute_reply.started":"2023-10-25T03:08:39.626204Z","shell.execute_reply":"2023-10-25T03:08:39.632527Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Configs","metadata":{}},{"cell_type":"code","source":"PAD_x = 0.0\nPAD_y = np.nan\n# X_max_len = 256\n# batch_size = 1024\n# val_batch_size = 128\n\nNUM_VOCAB = 5\n# hidden_dim = 192","metadata":{"execution":{"iopub.status.busy":"2023-10-25T03:08:39.634441Z","iopub.execute_input":"2023-10-25T03:08:39.634740Z","iopub.status.idle":"2023-10-25T03:08:39.644842Z","shell.execute_reply.started":"2023-10-25T03:08:39.634715Z","shell.execute_reply":"2023-10-25T03:08:39.643943Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# data, model, and training hyperparameters\nclass CFG:\n    save_output = True\n    output_dir = './runs'\n    \n    # data\n    fold = 'all'\n    inp_max_len = 206\n    batch_size = 256\n    # grad_acc_steps = 2\n    val_batch_size = 512\n    \n    # model\n    model_dim = 256\n    n_layers = 16\n    n_heads = 4\n    drop_rate = 0.1\n    norm_first = False # If true, use pre layer norm transformer blocks\n    activation = 'relu' #'relu'  \n    precision = 'mixed_float16' #'mixed_bfloat16'\n    \n    # training\n    n_epochs = 100\n    warmup = 0.25\n    lr_max = 1e-3\n    lr_min = lr_max * 0.02\n    wd_ratio = 0.01\n    warmup_type = 'linear'\n    decay_type = 'cosine'\n    \n#     fgm = False\n#     awp = False #True\n#     awp_lr = 0.2\n#     awp_start_epoch = 0.1 * epoch\n    \n    resume = 0\n    resume_ckpt = ''\n    comment =  f\"d{model_dim}_L{n_layers}_h{n_heads}_lr{lr_max}_e{n_epochs}\"","metadata":{"execution":{"iopub.status.busy":"2023-10-25T03:08:39.645841Z","iopub.execute_input":"2023-10-25T03:08:39.646107Z","iopub.status.idle":"2023-10-25T03:08:39.654208Z","shell.execute_reply.started":"2023-10-25T03:08:39.646084Z","shell.execute_reply":"2023-10-25T03:08:39.653365Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(CFG.comment)","metadata":{"execution":{"iopub.status.busy":"2023-10-25T03:08:39.655478Z","iopub.execute_input":"2023-10-25T03:08:39.655710Z","iopub.status.idle":"2023-10-25T03:08:39.664353Z","shell.execute_reply.started":"2023-10-25T03:08:39.655689Z","shell.execute_reply":"2023-10-25T03:08:39.663594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if DEBUG:\n    CFG.n_epochs=2","metadata":{"execution":{"iopub.status.busy":"2023-10-25T03:08:39.668673Z","iopub.execute_input":"2023-10-25T03:08:39.668976Z","iopub.status.idle":"2023-10-25T03:08:39.673880Z","shell.execute_reply.started":"2023-10-25T03:08:39.668954Z","shell.execute_reply":"2023-10-25T03:08:39.673011Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data API pipeline\nThis section applies filtering, preprocessing, shuffling, paddings, and batchings. I already transformed all the data to TFRecords; you can find the TFRecords dataset [here](https://www.kaggle.com/datasets/shlomoron/srrf-tfrecords-ds). I shuffled the samples before creating the TFRecords.","metadata":{}},{"cell_type":"code","source":"tffiles_path = ROOT+'/kaggle/input/srrf-tfrecords-ds/tfds'\ntffiles = [f'{tffiles_path}/{x}.tfrecord' for x in range(164)]\n\nif DEBUG:\n    tffiles = tffiles[:4]\n\nprint(len(tffiles))","metadata":{"execution":{"iopub.status.busy":"2023-10-25T03:08:39.674766Z","iopub.execute_input":"2023-10-25T03:08:39.675011Z","iopub.status.idle":"2023-10-25T03:08:39.684154Z","shell.execute_reply.started":"2023-10-25T03:08:39.674989Z","shell.execute_reply":"2023-10-25T03:08:39.683362Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Decoding TFRecords","metadata":{}},{"cell_type":"code","source":"def decode_tfrec(record_bytes):\n    schema = {}\n    schema[\"id\"] = tf.io.VarLenFeature(dtype=tf.string)\n    schema[\"seq\"] = tf.io.VarLenFeature(dtype=tf.float32)\n    schema[\"dataset_name_2A3\"] = tf.io.VarLenFeature(dtype=tf.string)\n    schema[\"dataset_name_DMS\"] = tf.io.VarLenFeature(dtype=tf.string)\n    schema[\"reads_2A3\"] = tf.io.VarLenFeature(dtype=tf.float32)\n    schema[\"reads_DMS\"] = tf.io.VarLenFeature(dtype=tf.float32)\n    schema[\"signal_to_noise_2A3\"] = tf.io.VarLenFeature(dtype=tf.float32)\n    schema[\"signal_to_noise_DMS\"] = tf.io.VarLenFeature(dtype=tf.float32)\n    schema[\"SN_filter_2A3\"] = tf.io.VarLenFeature(dtype=tf.float32)\n    schema[\"SN_filter_DMS\"] = tf.io.VarLenFeature(dtype=tf.float32)\n    schema[\"reactivity_2A3\"] = tf.io.VarLenFeature(dtype=tf.float32)\n    schema[\"reactivity_DMS\"] = tf.io.VarLenFeature(dtype=tf.float32)\n    schema[\"reactivity_error_2A3\"] = tf.io.VarLenFeature(dtype=tf.float32)\n    schema[\"reactivity_error_DMS\"] = tf.io.VarLenFeature(dtype=tf.float32)\n    features = tf.io.parse_single_example(record_bytes, schema)\n\n    sample_id = tf.sparse.to_dense(features[\"id\"])\n    seq = tf.sparse.to_dense(features[\"seq\"])\n    dataset_name_2A3 = tf.sparse.to_dense(features[\"dataset_name_2A3\"])\n    dataset_name_DMS = tf.sparse.to_dense(features[\"dataset_name_DMS\"])\n    reads_2A3 = tf.sparse.to_dense(features[\"reads_2A3\"])\n    reads_DMS = tf.sparse.to_dense(features[\"reads_DMS\"])\n    signal_to_noise_2A3 = tf.sparse.to_dense(features[\"signal_to_noise_2A3\"])\n    signal_to_noise_DMS = tf.sparse.to_dense(features[\"signal_to_noise_DMS\"])\n    SN_filter_2A3 = tf.sparse.to_dense(features[\"SN_filter_2A3\"])\n    SN_filter_DMS = tf.sparse.to_dense(features[\"SN_filter_DMS\"])\n    reactivity_2A3 = tf.sparse.to_dense(features[\"reactivity_2A3\"])\n    reactivity_DMS = tf.sparse.to_dense(features[\"reactivity_DMS\"])\n    reactivity_error_2A3 = tf.sparse.to_dense(features[\"reactivity_error_2A3\"])\n    reactivity_error_DMS = tf.sparse.to_dense(features[\"reactivity_error_DMS\"])\n\n    out = {}\n    out['seq']  = seq\n    out['SN_filter_2A3']  = SN_filter_2A3\n    out['SN_filter_DMS']  = SN_filter_DMS\n    out['reads_2A3']  = reads_2A3\n    out['reads_DMS']  = reads_DMS\n    out['signal_to_noise_2A3']  = signal_to_noise_2A3\n    out['signal_to_noise_DMS']  = signal_to_noise_DMS\n    out['reactivity_2A3']  = reactivity_2A3\n    out['reactivity_DMS']  = reactivity_DMS\n    return out","metadata":{"execution":{"iopub.status.busy":"2023-10-25T03:08:39.685393Z","iopub.execute_input":"2023-10-25T03:08:39.685682Z","iopub.status.idle":"2023-10-25T03:08:39.700319Z","shell.execute_reply.started":"2023-10-25T03:08:39.685658Z","shell.execute_reply":"2023-10-25T03:08:39.699518Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Filtering","metadata":{}},{"cell_type":"code","source":"def f1(): return True\ndef f2(): return False\n\ndef filter_function_1(x):\n    SN_filter_2A3 = x['SN_filter_2A3']\n    SN_filter_DMS = x['SN_filter_DMS']\n    return tf.cond((SN_filter_2A3 == 1) and (SN_filter_DMS == 1) , true_fn=f1, false_fn=f2)\n\ndef filter_function_2(x):\n    reads_2A3 = x['reads_2A3']\n    reads_DMS = x['reads_DMS']\n    signal_to_noise_2A3 = x['signal_to_noise_2A3']\n    signal_to_noise_DMS = x['signal_to_noise_DMS']\n    cond = (reads_2A3>100 and signal_to_noise_2A3>0.75) or (reads_DMS>100 and signal_to_noise_DMS>0.75)\n    return tf.cond(cond, true_fn=f1, false_fn=f2)","metadata":{"execution":{"iopub.status.busy":"2023-10-25T03:08:39.701535Z","iopub.execute_input":"2023-10-25T03:08:39.701875Z","iopub.status.idle":"2023-10-25T03:08:39.711836Z","shell.execute_reply.started":"2023-10-25T03:08:39.701844Z","shell.execute_reply":"2023-10-25T03:08:39.710735Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Preprocessing","metadata":{}},{"cell_type":"code","source":"def nan_below_filter(x):\n    reads_2A3 = x['reads_2A3']\n    reads_DMS = x['reads_DMS']\n    signal_to_noise_2A3 = x['signal_to_noise_2A3']\n    signal_to_noise_DMS = x['signal_to_noise_DMS']\n    reactivity_2A3 = x['reactivity_2A3']\n    reactivity_DMS = x['reactivity_DMS']\n\n    if reads_2A3<100 or signal_to_noise_2A3<0.75:\n        reactivity_2A3 = np.nan+reactivity_2A3\n    if reads_DMS<100 or signal_to_noise_DMS<0.75:\n        reactivity_DMS = np.nan+reactivity_DMS\n\n    x['reactivity_2A3'] = reactivity_2A3\n    x['reactivity_DMS'] = reactivity_DMS\n    return x\n\ndef concat_target(x):\n    reactivity_2A3 = x['reactivity_2A3']\n    reactivity_DMS = x['reactivity_DMS']\n    target = tf.concat([reactivity_2A3[..., tf.newaxis], reactivity_DMS[..., tf.newaxis]], axis = 1)\n    target = tf.clip_by_value(target, 0, 1)\n    return x['seq'], target","metadata":{"execution":{"iopub.status.busy":"2023-10-25T03:08:39.712944Z","iopub.execute_input":"2023-10-25T03:08:39.713229Z","iopub.status.idle":"2023-10-25T03:08:39.723392Z","shell.execute_reply.started":"2023-10-25T03:08:39.713189Z","shell.execute_reply":"2023-10-25T03:08:39.722584Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## get_tfrec_dataset","metadata":{}},{"cell_type":"code","source":"def get_tfrec_dataset(tffiles, max_len, shuffle, batch_size, cache = False, to_filter = False,\n                      calculate_sample_num = True, repeat=False):\n    ds = tf.data.TFRecordDataset(\n        tffiles, num_parallel_reads=tf.data.AUTOTUNE, compression_type = 'GZIP').prefetch(tf.data.AUTOTUNE)\n\n    ds = ds.map(decode_tfrec, tf.data.AUTOTUNE)\n    if to_filter == 'filter_1':\n        ds = ds.filter(filter_function_1)\n    elif to_filter == 'filter_2':\n        ds = ds.filter(filter_function_2)\n    ds = ds.map(nan_below_filter, tf.data.AUTOTUNE)\n    ds = ds.map(concat_target, tf.data.AUTOTUNE)\n\n#     if DEBUG:\n#         ds = ds.take(8)\n\n    if cache:\n        ds = ds.cache()\n\n    samples_num = 0\n    if calculate_sample_num:\n        samples_num = ds.reduce(0, lambda x,_: x+1).numpy()\n        \n    if repeat:\n        ds = ds.repeat()\n\n    if shuffle:\n        if shuffle == -1:\n            ds = ds.shuffle(samples_num, reshuffle_each_iteration = True)\n        else:\n            ds = ds.shuffle(shuffle, reshuffle_each_iteration = True)\n\n    if batch_size:\n        ds = ds.padded_batch(\n            batch_size, padding_values=(PAD_x, PAD_y), padded_shapes=([max_len],[max_len, 2]), drop_remainder=True)\n    ds = ds.prefetch(tf.data.AUTOTUNE)\n    return ds, samples_num","metadata":{"execution":{"iopub.status.busy":"2023-10-25T03:08:39.724449Z","iopub.execute_input":"2023-10-25T03:08:39.726300Z","iopub.status.idle":"2023-10-25T03:08:39.736186Z","shell.execute_reply.started":"2023-10-25T03:08:39.726275Z","shell.execute_reply":"2023-10-25T03:08:39.735378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define datasets","metadata":{}},{"cell_type":"code","source":"val_len = 5\nif DEBUG:\n    val_len = 1\n\nval_files = tffiles[:val_len]\n\nif DEBUG:\n    train_files = tffiles[val_len:val_len+1]\nelse:\n    train_files = tffiles[val_len:]","metadata":{"execution":{"iopub.status.busy":"2023-10-25T03:08:39.737191Z","iopub.execute_input":"2023-10-25T03:08:39.737495Z","iopub.status.idle":"2023-10-25T03:08:39.747882Z","shell.execute_reply.started":"2023-10-25T03:08:39.737472Z","shell.execute_reply":"2023-10-25T03:08:39.747090Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Get datasets","metadata":{}},{"cell_type":"code","source":"train_dataset, num_train = get_tfrec_dataset(train_files, max_len=CFG.inp_max_len, shuffle = -1, batch_size = CFG.batch_size,\n                                                  cache = False, to_filter = 'filter_2', calculate_sample_num = True, repeat=True)\n\nval_dataset, num_val = get_tfrec_dataset(val_files, max_len=CFG.inp_max_len, shuffle = False, batch_size = CFG.val_batch_size,\n                                                  cache = True, to_filter = 'filter_1', calculate_sample_num = True, repeat=False)\nprint(num_train)\nprint(num_val)","metadata":{"execution":{"iopub.status.busy":"2023-10-25T03:08:39.748904Z","iopub.execute_input":"2023-10-25T03:08:39.749147Z","iopub.status.idle":"2023-10-25T03:08:44.572796Z","shell.execute_reply.started":"2023-10-25T03:08:39.749126Z","shell.execute_reply":"2023-10-25T03:08:44.571710Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"steps_per_epoch = num_train // CFG.batch_size\nvalidation_steps = num_val // CFG.val_batch_size\n\nprint(steps_per_epoch)\nprint(validation_steps)","metadata":{"execution":{"iopub.status.busy":"2023-10-25T03:08:44.573885Z","iopub.execute_input":"2023-10-25T03:08:44.574466Z","iopub.status.idle":"2023-10-25T03:08:44.580204Z","shell.execute_reply.started":"2023-10-25T03:08:44.574439Z","shell.execute_reply":"2023-10-25T03:08:44.579172Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch = next(iter(val_dataset))\nbatch[0].shape, batch[1].shape","metadata":{"execution":{"iopub.status.busy":"2023-10-25T03:08:44.581553Z","iopub.execute_input":"2023-10-25T03:08:44.581951Z","iopub.status.idle":"2023-10-25T03:08:44.640936Z","shell.execute_reply.started":"2023-10-25T03:08:44.581918Z","shell.execute_reply":"2023-10-25T03:08:44.639959Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"markdown","source":"## Model layers","metadata":{}},{"cell_type":"code","source":"\nclass transformer_block(tf.keras.layers.Layer):\n    def __init__(self, dim, num_heads, feed_forward_dim, drop_rate=0.1, activation='relu', norm_first=False):\n        super().__init__()\n        self.att = tf.keras.layers.MultiHeadAttention(num_heads=num_heads, key_dim=dim//num_heads, dropout=drop_rate)\n        self.ffn = tf.keras.Sequential(\n            [\n                tf.keras.layers.Dense(feed_forward_dim, activation=activation),\n                tf.keras.layers.Dense(dim),\n            ]\n        )\n        self.norm_first = norm_first\n        self.layernorm1 = tf.keras.layers.LayerNormalization(epsilon=1e-6)\n        self.layernorm2 = tf.keras.layers.LayerNormalization(epsilon=1e-6)\n        self.dropout1 = tf.keras.layers.Dropout(drop_rate)\n        self.dropout2 = tf.keras.layers.Dropout(drop_rate)\n        self.supports_masking = True\n\n    def call(self, inputs, training, mask):\n        if self.norm_first is False:\n            att_mask = tf.expand_dims(mask, axis=-1)\n            att_mask = tf.repeat(att_mask, repeats=tf.shape(att_mask)[1], axis=-1)\n\n            attn_output = self.att(inputs, inputs, attention_mask = att_mask)\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            output = self.layernorm2(out1 + ffn_output)\n        else:\n            att_mask = tf.expand_dims(mask, axis=-1)\n            att_mask = tf.repeat(att_mask, repeats=tf.shape(att_mask)[1], axis=-1)\n            \n            # apply layer norm before attn block and before ffn block\n            norm1 = self.layernorm1(inputs)\n            attn_output = self.att(norm1, norm1, attention_mask = att_mask)\n            attn_output = self.dropout1(attn_output, training=training)\n            out1 = inputs + attn_output\n            norm2 = self.layernorm2(out1)\n            ffn_output = self.ffn(norm2)\n            ffn_output = self.dropout2(ffn_output, training=training)\n            output = out1 + ffn_output\n            \n        return output\n\n\nclass positional_encoding_layer(tf.keras.layers.Layer):\n    def __init__(self, num_vocab=5, maxlen=512, hidden_dim=384):\n        super().__init__()\n        self.hidden_dim = hidden_dim\n        self.pos_emb = self.positional_encoding(maxlen-1, hidden_dim)\n        self.supports_masking = True\n\n    def call(self, x):\n        maxlen = tf.shape(x)[-2]\n        x = tf.cast(x, tf.float32)\n        x = tf.math.multiply(x, tf.math.sqrt(tf.cast(self.hidden_dim, tf.float32)))\n        return x + self.pos_emb[:maxlen, :]\n\n    def positional_encoding(self, maxlen, hidden_dim):\n        depth = hidden_dim/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\n    ","metadata":{"execution":{"iopub.status.busy":"2023-10-25T03:08:44.642457Z","iopub.execute_input":"2023-10-25T03:08:44.642934Z","iopub.status.idle":"2023-10-25T03:08:44.661428Z","shell.execute_reply.started":"2023-10-25T03:08:44.642849Z","shell.execute_reply":"2023-10-25T03:08:44.660343Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_model(hidden_dim = 384, max_len = 206, n_layers=12, n_heads=6, drop_rate=0.1, activation='relu', norm_first=False):\n    inp = tf.keras.Input([max_len])\n    x = inp\n\n    x = tf.keras.layers.Embedding(input_dim=NUM_VOCAB, output_dim=hidden_dim, mask_zero=True)(x)\n    x = positional_encoding_layer(num_vocab=NUM_VOCAB, maxlen=512, hidden_dim=hidden_dim)(x)\n\n    for i in range(n_layers):\n        x = transformer_block(hidden_dim, n_heads, hidden_dim*4, drop_rate=drop_rate, activation=activation, norm_first=norm_first)(x)\n\n    x = tf.keras.layers.Dropout(0.5)(x)\n    x = tf.keras.layers.Dense(2)(x)\n\n    # model = CustomTrainStep(n_gradients=CFG.grad_acc_steps, inputs=inp, outputs=x)\n    model = tf.keras.Model(inp, x)\n    return model","metadata":{"execution":{"iopub.status.busy":"2023-10-25T03:08:44.663016Z","iopub.execute_input":"2023-10-25T03:08:44.663551Z","iopub.status.idle":"2023-10-25T03:08:44.674144Z","shell.execute_reply.started":"2023-10-25T03:08:44.663516Z","shell.execute_reply":"2023-10-25T03:08:44.673238Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Loss function","metadata":{}},{"cell_type":"code","source":"def loss_fn(labels, targets):\n    targets = tf.clip_by_value(targets, 0, 1)\n    labels_mask = tf.math.is_nan(labels)\n    labels = tf.where(labels_mask, tf.zeros_like(labels), labels)\n    mask_count = tf.math.reduce_sum(tf.where(labels_mask, tf.zeros_like(labels), tf.ones_like(labels)))\n    loss = tf.math.abs(labels - targets)\n    loss = tf.where(labels_mask, tf.zeros_like(loss), loss)\n    loss = tf.math.reduce_sum(loss)/mask_count\n    return loss\n\n# metric fn\n# def MAE()","metadata":{"execution":{"iopub.status.busy":"2023-10-25T03:08:44.675414Z","iopub.execute_input":"2023-10-25T03:08:44.675735Z","iopub.status.idle":"2023-10-25T03:08:44.686787Z","shell.execute_reply.started":"2023-10-25T03:08:44.675706Z","shell.execute_reply":"2023-10-25T03:08:44.685922Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Callbacks","metadata":{}},{"cell_type":"code","source":"N_EPOCHS = CFG.n_epochs\nif DEBUG:\n    N_EPOCHS = 2\nN_WARMUP_EPOCHS = CFG.n_epochs * CFG.warmup\nLR_MAX = CFG.lr_max\nWD_RATIO = CFG.wd_ratio\nWARMUP_TYPE = CFG.warmup_type\nprint(f\"Epochs: {N_EPOCHS}\\nWarmup epochs: {N_WARMUP_EPOCHS}\\nLR: {LR_MAX}\\nWD_ratio: {WD_RATIO}\\nWarmup type: {WARMUP_TYPE}\")\n\ndef lrfn(current_step, num_warmup_steps, lr_max, num_cycles=0.50, num_training_steps=N_EPOCHS):\n    if current_step < num_warmup_steps:\n        if WARMUP_METHOD == 'log':\n            return lr_max * 0.10 ** (num_warmup_steps - current_step)\n        elif WARMUP_METHOD == 'exp':\n            return lr_max * 2 ** -(num_warmup_steps - current_step)\n        else: # linear\n            return lr_max * ((current_step + 0.01) / num_warmup_steps)\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\n\n# Plot Learning Rate Schedule\n# plot_lr_schedule(LR_SCHEDULE, epochs=N_EPOCHS)\n# Learning Rate Callback\n# lr_callback = tf.keras.callbacks.LearningRateScheduler(lambda step: LR_SCHEDULE[step], verbose=0)\n\n# save_folder = ROOT+'/kaggle/working/'+NAME\n# print(save_folder)\n# try:\n#     os.mkdir(f'{save_folder}')\n# except:\n#     pass\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 or (epoch+1)%25 == 0:\n            self.model.save_weights(f\"{save_folder}/model_epoch_{epoch}.h5\")\n            \nclass CSVLoggerV2(tf.keras.callbacks.CSVLogger):\n    def __init__(self, filename, separator=\",\", resume=0):\n        self.resume = resume\n        super().__init__(filename=filename, separator=separator, append=bool(resume))\n\n    def on_epoch_end(self, epoch, logs=None):\n        super(CSVLoggerV2,self).on_epoch_end(epoch=epoch+self.resume+1, logs=logs)","metadata":{"execution":{"iopub.status.busy":"2023-10-25T03:08:44.688165Z","iopub.execute_input":"2023-10-25T03:08:44.688504Z","iopub.status.idle":"2023-10-25T03:08:44.716697Z","shell.execute_reply.started":"2023-10-25T03:08:44.688474Z","shell.execute_reply":"2023-10-25T03:08:44.715397Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"code","source":"def train_model(CFG, train_ds, val_ds, train_steps, val_steps, strategy=strategy, summary=True):\n    clear_mem()\n    print(f'precision: {CFG.precision}')\n    if CFG.precision:\n        policy = tf.keras.mixed_precision.Policy(CFG.precision)\n        tf.keras.mixed_precision.set_global_policy(policy)\n    if device == \"GPU\":\n        tf.config.optimizer.set_jit(True)\n    \n    fold = CFG.fold\n    resume = CFG.resume\n    \n    # Create model\n    with strategy.scope():\n        model = get_model(hidden_dim=CFG.model_dim, max_len=CFG.inp_max_len, n_layers=CFG.n_layers, n_heads=CFG.n_heads, \n                          drop_rate=CFG.drop_rate, activation=CFG.activation, norm_first=CFG.norm_first)\n        \n        schedule = OneCycleLR(CFG.lr_max, CFG.n_epochs, warmup_epochs=CFG.n_epochs*CFG.warmup, steps_per_epoch=train_steps,\n                              resume_epoch=CFG.resume, decay_epochs=CFG.n_epochs, lr_min=CFG.lr_min, decay_type=CFG.decay_type, warmup_type=CFG.warmup_type)\n        decay_schedule = OneCycleLR(CFG.lr_max*CFG.wd_ratio, CFG.n_epochs, warmup_epochs=CFG.n_epochs*CFG.warmup, steps_per_epoch=train_steps,\n                                    resume_epoch=CFG.resume, decay_epochs=CFG.n_epochs, lr_min=CFG.lr_min*CFG.wd_ratio, decay_type=CFG.decay_type, warmup_type=CFG.warmup_type)\n        \n        # FGM, AWP ...\n        \n#         optimizer = tf.keras.optimizers.AdamW(learning_rate=schedule, weight_decay=decay_schedule, clipnorm=3.0)\n        optimizer = tfa.optimizers.AdamW(learning_rate=schedule, weight_decay=decay_schedule, clipnorm=3.0)\n        \n        model.compile(loss=loss_fn,\n                      optimizer=optimizer,\n                      steps_per_execution=train_steps,\n                      metrics=[loss_fn],\n                     )\n\n    if summary:\n        print()\n        print(CFG.comment)\n        print()\n        model.summary()\n        print()\n        print(train_ds, val_ds)\n        print()\n        schedule.plot()\n        print()\n    print(f'---------fold{fold}---------')\n    print(f'train batches:{train_steps} val batches:{val_steps}')\n    print()\n    \n    if CFG.resume != 0:\n        print(f'resume from epoch{CFG.resume}')\n        if CFG.resume_ckpt:\n            print(f'load weights from {CFG.resume_ckpt}')\n            model.load_weights(CFG.resume_ckpt)\n        else:\n            print(f'ckpt file not found')\n            return\n\n    # Callbacks\n    logger = CSVLoggerV2(f'{CFG.output_dir}/{CFG.comment}-fold{fold}-logs.csv', resume=resume)\n    mode = 'min'\n    if val_ds:\n        monitor = 'val_loss'\n    else:\n        monitor = 'loss'\n    if resume:\n        prev_best = pd.read_csv(f'{CFG.output_dir}/{CFG.comment}-fold{fold}-logs.csv')[monitor].agg(mode)\n    else:\n        prev_best = None\n    sv_loss = tf.keras.callbacks.ModelCheckpoint(f'{CFG.output_dir}/{CFG.comment}-fold{fold}-best.h5', monitor=monitor, verbose=0, save_best_only=True,\n                  save_weights_only=True, mode='min', save_freq='epoch', initial_value_threshold=prev_best)\n    snap = Snapshot(f'{CFG.output_dir}/{CFG.comment}-fold{fold}', snapshot_epochs=[])\n    # swa = SWA(f'{CFG.output_dir}/{CFG.comment}-fold{fold}', CFG.swa_epochs, strategy=strategy, train_ds=train_ds, valid_ds=valid_ds)\n    \n    callbacks = []\n    if CFG.save_output:\n        callbacks.append(logger)\n        callbacks.append(snap)\n        # callbacks.append(swa)\n        callbacks.append(sv_loss)\n\n    history = model.fit(\n        train_ds,\n        validation_data=val_ds,\n        epochs=CFG.n_epochs,\n        verbose = 2,\n        callbacks=callbacks,\n        steps_per_epoch=train_steps,\n        validation_steps=val_steps,\n    )\n    \n    return model, history","metadata":{"execution":{"iopub.status.busy":"2023-10-25T03:08:44.718394Z","iopub.execute_input":"2023-10-25T03:08:44.718754Z","iopub.status.idle":"2023-10-25T03:08:44.735661Z","shell.execute_reply.started":"2023-10-25T03:08:44.718677Z","shell.execute_reply":"2023-10-25T03:08:44.734623Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model, history = train_fold(CFG, train_dataset, val_dataset, steps_per_epoch, validation_steps, strategy)","metadata":{"execution":{"iopub.status.busy":"2023-10-25T03:08:45.701116Z","iopub.execute_input":"2023-10-25T03:08:45.701697Z","iopub.status.idle":"2023-10-25T03:10:45.290307Z","shell.execute_reply.started":"2023-10-25T03:08:45.701664Z","shell.execute_reply":"2023-10-25T03:10:45.289480Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Plotting loss","metadata":{}},{"cell_type":"code","source":"plt.plot(history.history['loss'])\nplt.plot(history.history['val_loss'])","metadata":{"execution":{"iopub.status.busy":"2023-10-25T03:10:45.291720Z","iopub.execute_input":"2023-10-25T03:10:45.292011Z","iopub.status.idle":"2023-10-25T03:10:45.547171Z","shell.execute_reply.started":"2023-10-25T03:10:45.291987Z","shell.execute_reply":"2023-10-25T03:10:45.546289Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference\nDoing the inference in a separate notebook was easier, so I split it. Sorry about that. Find the inference notebook [HERE](https://www.kaggle.com/code/shlomoron/srrf-transformer-tpu-inference).","metadata":{}}]}