{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.8.17","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"tpu1vmV38","dataSources":[],"dockerImageVersionId":30527,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"from rich.progress import track\nimport sklearn.model_selection\nimport tensorflow as tf\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport pickle\nimport shutil\nimport math\nimport pandas as pd\nimport gc\nimport os\n\n!pip install ipywidgets","metadata":{"execution":{"iopub.status.busy":"2023-11-26T16:32:28.149588Z","iopub.execute_input":"2023-11-26T16:32:28.149903Z","iopub.status.idle":"2023-11-26T16:33:10.308594Z","shell.execute_reply.started":"2023-11-26T16:32:28.149877Z","shell.execute_reply":"2023-11-26T16:33:10.307224Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tf.__version__","metadata":{"execution":{"iopub.status.busy":"2023-11-26T16:33:15.855764Z","iopub.execute_input":"2023-11-26T16:33:15.856064Z","iopub.status.idle":"2023-11-26T16:33:15.864794Z","shell.execute_reply.started":"2023-11-26T16:33:15.856036Z","shell.execute_reply":"2023-11-26T16:33:15.864014Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_pth = '/kaggle/input/stanford-ribonanza-rna-folding/train_data.csv'\ndata = np.array(open(data_pth, 'r').read().split('\\n')[1:-1])\n\nmodel_save_dir = 'model_saves'\nprint('model saves dir:', model_save_dir)\nif not os.path.exists(model_save_dir):\n    os.makedirs(model_save_dir)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T16:33:15.866722Z","iopub.execute_input":"2023-11-26T16:33:15.867021Z","iopub.status.idle":"2023-11-26T16:33:47.364918Z","shell.execute_reply.started":"2023-11-26T16:33:15.866996Z","shell.execute_reply":"2023-11-26T16:33:47.364051Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def mk_data(fo, prefix):\n    words = {'A':1, 'C':2, 'G':3, 'U':4}\n    samples = {}\n    for i in track(range(len(fo))):\n        row = fo[i].split(',')\n        if prefix =='train':\n            if int(row[4])<100 or float(row[5]) <0.75:\n                continue\n        else:\n            if row[6] != '1':\n                continue\n        reactivity = np.array([float(s) if s!='' else np.nan for s in row[7:206+7]])\n        if 206 - np.isnan(reactivity).sum() <=0:\n            continue\n        reactivity = np.clip(reactivity, 0, 1)\n        if row[0] not in samples:\n            label = np.ones([206,2])*np.nan\n            seq  = np.array([words[s] for s in row[1]])\n            mask = np.ones([206])\n            mask[len(seq):] = 0\n            seq = np.pad(seq, (0, 206-seq.shape[0]), mode='constant', constant_values=0)\n            \n            if row[2] =='DMS_MaP':\n                label[...,0] = reactivity\n            else:\n                label[...,1] = reactivity\n            samples[row[0]] = {\n                \"seq\":seq,\n                'mask':mask,\n                'label':label,\n            }\n        else:\n\n            if row[2] =='DMS_MaP':\n                samples[row[0]]['label'][...,0] = reactivity\n            else:\n                samples[row[0]]['label'][...,1] = reactivity\n    return samples\n\n\ndef _parse_function(example_proto):\n    label,mask,seq = [],[],[]\n    for example in example_proto:\n        seq.append(example['seq'])\n        mask.append(example['mask'])\n        label.append(example['label'])\n    return (seq, mask), label\n\n\ndef _paste_mask(inputs, label):\n    seq, mask = inputs\n    mask = mask * tf.where(tf.random.uniform([206],0,1) < 0.15, tf.zeros_like(mask), tf.ones_like(mask))\n    label = tf.where(tf.tile(tf.expand_dims(mask, -1), [1,2])==0, tf.ones_like(label)*math.nan, label)\n    return (seq, mask), label\n# def _paste_mask(inputs, label):\n#     seq, mask = inputs\n#     random_ = tf.random.uniform([206],0,1) \n#     seq = mask * tf.cast(tf.where(random_ < 0.15, tf.ones_like(seq)*5, seq), mask.dtype)\n#     mask = mask * tf.where(random_ < 0.15, tf.zeros_like(mask), tf.ones_like(mask))\n#     return (seq, mask), label","metadata":{"execution":{"iopub.status.busy":"2023-11-26T16:33:47.366371Z","iopub.execute_input":"2023-11-26T16:33:47.366662Z","iopub.status.idle":"2023-11-26T16:33:47.380600Z","shell.execute_reply.started":"2023-11-26T16:33:47.366636Z","shell.execute_reply":"2023-11-26T16:33:47.379876Z"},"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-11-26T16:34:45.144416Z","iopub.execute_input":"2023-11-26T16:34:45.144896Z","iopub.status.idle":"2023-11-26T16:34:53.551011Z","shell.execute_reply.started":"2023-11-26T16:34:45.144861Z","shell.execute_reply":"2023-11-26T16:34:53.550164Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DEBUG = False\n\nPAD_x = 0.0\nPAD_y = np.nan\nX_max_len = 206\nbatch_size = 128\nval_batch_size = 5512\n\n\nif DEBUG:\n    batch_size = 2\n    val_batch_size = 2\n\nnum_vocab = 5\nhidden_dim = 192","metadata":{"execution":{"iopub.status.busy":"2023-11-26T16:34:53.552395Z","iopub.execute_input":"2023-11-26T16:34:53.552652Z","iopub.status.idle":"2023-11-26T16:34:53.557604Z","shell.execute_reply.started":"2023-11-26T16:34:53.552627Z","shell.execute_reply":"2023-11-26T16:34:53.556864Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class transformer_block(tf.keras.layers.Layer):\n    def __init__(self, dim, num_heads, feed_forward_dim, rate=0.1):\n        super().__init__()\n        self.att = tf.keras.layers.MultiHeadAttention(num_heads=num_heads, key_dim=dim//num_heads)\n        self.ffn = tf.keras.Sequential(\n            [\n                tf.keras.layers.Dense(feed_forward_dim, activation=\"relu\"),\n                tf.keras.layers.Dense(dim),\n            ]\n        )\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(rate)\n        self.dropout2 = tf.keras.layers.Dropout(rate)\n        self.supports_masking = True\n\n    def call(self, inputs, training, mask):\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        return self.layernorm2(out1 + ffn_output)\n\n\nclass positional_encoding_layer(tf.keras.layers.Layer):\n    def __init__(self, num_vocab=5, maxlen=500, 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.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","metadata":{"execution":{"iopub.status.busy":"2023-11-26T16:35:39.529090Z","iopub.execute_input":"2023-11-26T16:35:39.529540Z","iopub.status.idle":"2023-11-26T16:35:39.543877Z","shell.execute_reply.started":"2023-11-26T16:35:39.529507Z","shell.execute_reply":"2023-11-26T16:35:39.543070Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def loss_fn(labels, targets):\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","metadata":{"execution":{"iopub.status.busy":"2023-11-26T16:35:43.267777Z","iopub.execute_input":"2023-11-26T16:35:43.268214Z","iopub.status.idle":"2023-11-26T16:35:43.274264Z","shell.execute_reply.started":"2023-11-26T16:35:43.268181Z","shell.execute_reply":"2023-11-26T16:35:43.273460Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_model(hidden_dim = 384, max_len = 206):\n    with strategy.scope():\n        inp_x = tf.keras.Input([max_len])\n        inp_mask = tf.keras.Input([max_len])\n        x = inp_x\n\n        x = tf.keras.layers.Embedding(num_vocab, hidden_dim, mask_zero=True)(x)\n        x = positional_encoding_layer(num_vocab=num_vocab, maxlen=500, hidden_dim=hidden_dim)(x)\n\n        x = transformer_block(hidden_dim, 6, hidden_dim*4)(x, mask=inp_mask)\n        x = transformer_block(hidden_dim, 6, hidden_dim*4)(x, mask=inp_mask)\n        x = transformer_block(hidden_dim, 6, hidden_dim*4)(x, mask=inp_mask)\n        x = transformer_block(hidden_dim, 6, hidden_dim*4)(x, mask=inp_mask)\n\n        x = transformer_block(hidden_dim, 6, hidden_dim*4)(x, mask=inp_mask)\n        x = transformer_block(hidden_dim, 6, hidden_dim*4)(x, mask=inp_mask)\n        x = transformer_block(hidden_dim, 6, hidden_dim*4)(x, mask=inp_mask)\n        x = transformer_block(hidden_dim, 6, hidden_dim*4)(x, mask=inp_mask)\n\n        x = transformer_block(hidden_dim, 6, hidden_dim*4)(x, mask=inp_mask)\n        x = transformer_block(hidden_dim, 6, hidden_dim*4)(x, mask=inp_mask)\n        x = transformer_block(hidden_dim, 6, hidden_dim*4)(x, mask=inp_mask)\n        x = transformer_block(hidden_dim, 6, hidden_dim*4)(x, mask=inp_mask)\n\n        x = tf.keras.layers.Dropout(0.5)(x)\n        x = tf.keras.layers.Dense(2)(x)\n        #x = tf.keras.layers.Flatten()(x)\n        model = tf.keras.Model((inp_x, inp_mask), x)\n        loss = loss_fn\n        optimizer = tf.keras.optimizers.AdamW(learning_rate=0.0005)\n        model.compile(loss=loss, optimizer=optimizer, steps_per_execution = 100)\n        return model\n\ntf.keras.backend.clear_session()\n\nmodel = get_model(hidden_dim = 192,max_len = X_max_len)\n# model(batch[0])\nmodel.summary()","metadata":{"execution":{"iopub.status.busy":"2023-11-26T16:35:50.418902Z","iopub.execute_input":"2023-11-26T16:35:50.419328Z","iopub.status.idle":"2023-11-26T16:35:55.887167Z","shell.execute_reply.started":"2023-11-26T16:35:50.419296Z","shell.execute_reply":"2023-11-26T16:35:55.886304Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"N_EPOCHS = 400\nif DEBUG:\n    N_EPOCHS = 5\nN_WARMUP_EPOCHS = int(N_EPOCHS*0.2)\n# N_WARMUP_EPOCHS = 0\nLR_MAX = 5e-4\nWD_RATIO = 0.05\nWARMUP_METHOD = \"exp\"","metadata":{"execution":{"iopub.status.busy":"2023-11-26T16:35:55.888524Z","iopub.execute_input":"2023-11-26T16:35:55.888796Z","iopub.status.idle":"2023-11-26T16:35:55.893381Z","shell.execute_reply.started":"2023-11-26T16:35:55.888773Z","shell.execute_reply":"2023-11-26T16:35:55.892676Z"},"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    if current_step==0 and num_warmup_steps !=0:\n        return lr_max/(num_warmup_steps)\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        return lr_max/(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)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T16:35:56.836655Z","iopub.execute_input":"2023-11-26T16:35:56.837103Z","iopub.status.idle":"2023-11-26T16:35:59.168205Z","shell.execute_reply.started":"2023-11-26T16:35:56.837068Z","shell.execute_reply":"2023-11-26T16:35:59.167296Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class save_model_callback(tf.keras.callbacks.Callback):\n    def __init__(self, fold):\n        super().__init__()\n        self.fold = fold\n    def on_epoch_end(self, epoch: int, logs=None):\n        if epoch == 3 or (epoch+1)%25 == 0:\n            self.model.save_weights(os.path.join(model_save_dir, f'best_model_{self.fold}_{epoch}.h5'))","metadata":{"execution":{"iopub.status.busy":"2023-11-26T16:35:59.938810Z","iopub.execute_input":"2023-11-26T16:35:59.939218Z","iopub.status.idle":"2023-11-26T16:35:59.945543Z","shell.execute_reply.started":"2023-11-26T16:35:59.939188Z","shell.execute_reply":"2023-11-26T16:35:59.944675Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"kfold = sklearn.model_selection.KFold(n_splits=5,shuffle=True,random_state=0)\nfor fold,(tr_idx, va_idx) in enumerate(kfold.split(data)):\n    train_data = mk_data(data[tr_idx], 'train')\n    valid_data = mk_data(data[va_idx], 'valid')\n    train_dataset = tf.data.Dataset.from_tensor_slices(\n            _parse_function(train_data.values())).shuffle(len(tr_idx)).batch(128).prefetch(tf.data.AUTOTUNE)\n    valid_dataset = tf.data.Dataset.from_tensor_slices(\n            _parse_function(valid_data.values())).batch(128).prefetch(tf.data.AUTOTUNE)\n    \n    #steps_per_epoch = num_train//batch_size\n    #val_steps_per_epoch = num_val//val_batch_size\n    history = model.fit(\n    train_dataset,\n    validation_data=valid_dataset,\n    epochs=N_EPOCHS,\n    #steps_per_epoch = steps_per_epoch,\n    #validation_steps=val_steps_per_epoch,\n    callbacks=[\n        save_model_callback(fold),\n        lr_callback,\n    ]\n    )\n    break","metadata":{"execution":{"iopub.status.busy":"2023-11-26T16:37:33.465230Z","iopub.execute_input":"2023-11-26T16:37:33.465633Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}