{"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":"import numpy as np\nimport pandas as pd\nimport os\nimport time\nimport numpy as np\nimport tensorflow as tf\nfrom sklearn.model_selection import KFold\nimport random\nimport joblib\nimport logging","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-11-13T11:54:06.176611Z","iopub.execute_input":"2023-11-13T11:54:06.177170Z","iopub.status.idle":"2023-11-13T11:54:41.667435Z","shell.execute_reply.started":"2023-11-13T11:54:06.177142Z","shell.execute_reply":"2023-11-13T11:54:41.666633Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_file = \"train_file.csv\"\n#model = 'model_DMS.ml'\nmodel = 'model.ml'","metadata":{"execution":{"iopub.status.busy":"2023-11-13T11:54:41.669134Z","iopub.execute_input":"2023-11-13T11:54:41.669572Z","iopub.status.idle":"2023-11-13T11:54:41.673087Z","shell.execute_reply.started":"2023-11-13T11:54:41.669547Z","shell.execute_reply":"2023-11-13T11:54:41.672442Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# read and create data","metadata":{}},{"cell_type":"markdown","source":"###    Configs","metadata":{}},{"cell_type":"code","source":"PAD_x = 0.0\nPAD_y = np.nan\nX_max_len = 206\nbatch_size = 128\nval_batch_size = 5512","metadata":{"execution":{"iopub.status.busy":"2023-11-13T11:54:41.673885Z","iopub.execute_input":"2023-11-13T11:54:41.674096Z","iopub.status.idle":"2023-11-13T11:54:41.686531Z","shell.execute_reply.started":"2023-11-13T11:54:41.674076Z","shell.execute_reply":"2023-11-13T11:54:41.685917Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tffiles_path = '/kaggle/input/srrf-tfrecords-ds/tfds'\ntffiles = [f'{tffiles_path}/{x}.tfrecord' for x in range(164)]","metadata":{"execution":{"iopub.status.busy":"2023-11-13T11:54:41.687420Z","iopub.execute_input":"2023-11-13T11:54:41.687709Z","iopub.status.idle":"2023-11-13T11:54:41.698073Z","shell.execute_reply.started":"2023-11-13T11:54:41.687684Z","shell.execute_reply":"2023-11-13T11:54:41.697307Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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-11-13T11:54:41.700298Z","iopub.execute_input":"2023-11-13T11:54:41.700532Z","iopub.status.idle":"2023-11-13T11:54:41.712720Z","shell.execute_reply.started":"2023-11-13T11:54:41.700511Z","shell.execute_reply":"2023-11-13T11:54:41.712051Z"},"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-11-13T11:54:41.713452Z","iopub.execute_input":"2023-11-13T11:54:41.713694Z","iopub.status.idle":"2023-11-13T11:54:41.727485Z","shell.execute_reply.started":"2023-11-13T11:54:41.713673Z","shell.execute_reply":"2023-11-13T11:54:41.726819Z"},"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-11-13T11:54:41.728464Z","iopub.execute_input":"2023-11-13T11:54:41.728743Z","iopub.status.idle":"2023-11-13T11:54:41.739413Z","shell.execute_reply.started":"2023-11-13T11:54:41.728720Z","shell.execute_reply":"2023-11-13T11:54:41.738684Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### get_tfrec_dataset","metadata":{}},{"cell_type":"code","source":"def get_tfrec_dataset(tffiles, shuffle, batch_size, cache = False, to_filter = False,\n                      calculate_sample_num = True, to_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 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 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 to_repeat:\n        ds = ds.repeat()\n         \n    if batch_size:\n        ds = ds.padded_batch(\n            batch_size, padding_values=(PAD_x, PAD_y), padded_shapes=([X_max_len],[X_max_len, 2]), drop_remainder=True)\n    ds = ds.prefetch(tf.data.AUTOTUNE)\n    return ds, samples_num","metadata":{"execution":{"iopub.status.busy":"2023-11-13T11:54:41.740286Z","iopub.execute_input":"2023-11-13T11:54:41.740514Z","iopub.status.idle":"2023-11-13T11:54:41.753963Z","shell.execute_reply.started":"2023-11-13T11:54:41.740494Z","shell.execute_reply":"2023-11-13T11:54:41.753319Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_len = 5\nval_files = tffiles[:val_len]\ntrain_files = tffiles[val_len:]","metadata":{"execution":{"iopub.status.busy":"2023-11-13T11:54:41.754770Z","iopub.execute_input":"2023-11-13T11:54:41.754989Z","iopub.status.idle":"2023-11-13T11:54:41.768663Z","shell.execute_reply.started":"2023-11-13T11:54:41.754969Z","shell.execute_reply":"2023-11-13T11:54:41.767909Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset, num_train = get_tfrec_dataset(train_files, shuffle = -1, batch_size = batch_size,\n                                                  cache = True, to_filter = 'filter_2', calculate_sample_num = True,\n                                            to_repeat = True)\n\nval_dataset, num_val = get_tfrec_dataset(val_files, shuffle = False, batch_size = val_batch_size,\n                                                  cache = True, to_filter = 'filter_1', calculate_sample_num = True)\nprint(num_train)\nprint(num_val)","metadata":{"execution":{"iopub.status.busy":"2023-11-13T11:54:41.769483Z","iopub.execute_input":"2023-11-13T11:54:41.769743Z","iopub.status.idle":"2023-11-13T11:55:45.658412Z","shell.execute_reply.started":"2023-11-13T11:54:41.769722Z","shell.execute_reply":"2023-11-13T11:55:45.657163Z"},"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-11-13T11:55:45.659854Z","iopub.execute_input":"2023-11-13T11:55:45.660180Z","iopub.status.idle":"2023-11-13T11:55:45.729583Z","shell.execute_reply.started":"2023-11-13T11:55:45.660152Z","shell.execute_reply":"2023-11-13T11:55:45.728526Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# model","metadata":{}},{"cell_type":"code","source":"def positional_encoding(length, depth):\n  depth = depth/2\n\n  positions = np.arange(length)[:, np.newaxis]     # (seq, 1)\n  depths = np.arange(depth)[np.newaxis, :]/depth   # (1, depth)\n\n  angle_rates = 1 / (10000**depths)         # (1, depth)\n  angle_rads = positions * angle_rates      # (pos, depth)\n\n  pos_encoding = np.concatenate(\n      [np.sin(angle_rads), np.cos(angle_rads)],\n      axis=-1) \n\n  return tf.cast(pos_encoding, dtype=tf.float32)\n                 \nclass PositionalEmbedding(tf.keras.layers.Layer):\n  def __init__(self, vocab_size, d_model):\n    super().__init__()\n    self.d_model = d_model\n    self.embedding = tf.keras.layers.Embedding(vocab_size, d_model, mask_zero=True) \n    self.pos_encoding = positional_encoding(length=2048, depth=d_model)\n\n  def compute_mask(self, *args, **kwargs):\n    return self.embedding.compute_mask(*args, **kwargs)\n\n  def call(self, x):\n    length = tf.shape(x)[1]\n    x = self.embedding(x)\n    # This factor sets the relative scale of the embedding and positonal_encoding.\n    x *= tf.math.sqrt(tf.cast(self.d_model, tf.float32))\n    x = x + self.pos_encoding[tf.newaxis, :length, :]\n    return x","metadata":{"execution":{"iopub.status.busy":"2023-11-13T11:55:45.730769Z","iopub.execute_input":"2023-11-13T11:55:45.731060Z","iopub.status.idle":"2023-11-13T11:55:45.740657Z","shell.execute_reply.started":"2023-11-13T11:55:45.731034Z","shell.execute_reply":"2023-11-13T11:55:45.739678Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class BaseAttention(tf.keras.layers.Layer):\n  def __init__(self, **kwargs):\n    super().__init__()\n    self.mha = tf.keras.layers.MultiHeadAttention(**kwargs)\n    self.layernorm = tf.keras.layers.LayerNormalization()\n    self.add = tf.keras.layers.Add()","metadata":{"execution":{"iopub.status.busy":"2023-11-13T11:55:45.741797Z","iopub.execute_input":"2023-11-13T11:55:45.742073Z","iopub.status.idle":"2023-11-13T11:55:45.755524Z","shell.execute_reply.started":"2023-11-13T11:55:45.742048Z","shell.execute_reply":"2023-11-13T11:55:45.754685Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class GlobalSelfAttention(BaseAttention):\n  def call(self, x):\n    attn_output = self.mha(\n        query=x,\n        value=x,\n        key=x)\n    x = self.add([x, attn_output])\n    x = self.layernorm(x)\n    return x","metadata":{"execution":{"iopub.status.busy":"2023-11-13T11:55:45.759080Z","iopub.execute_input":"2023-11-13T11:55:45.759333Z","iopub.status.idle":"2023-11-13T11:55:45.765107Z","shell.execute_reply.started":"2023-11-13T11:55:45.759313Z","shell.execute_reply":"2023-11-13T11:55:45.764225Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class FeedForward(tf.keras.layers.Layer):\n  def __init__(self, d_model, dff, dropout_rate=0.1):\n    super().__init__()\n    self.seq = tf.keras.Sequential([\n      tf.keras.layers.Dense(dff, activation='relu'),\n      tf.keras.layers.Dense(d_model),\n      tf.keras.layers.Dropout(dropout_rate)\n    ])\n    self.add = tf.keras.layers.Add()\n    self.layer_norm = tf.keras.layers.LayerNormalization()\n\n  def call(self, x):\n    x = self.add([x, self.seq(x)])\n    x = self.layer_norm(x) \n    return x","metadata":{"execution":{"iopub.status.busy":"2023-11-13T11:55:45.765993Z","iopub.execute_input":"2023-11-13T11:55:45.766211Z","iopub.status.idle":"2023-11-13T11:55:45.776610Z","shell.execute_reply.started":"2023-11-13T11:55:45.766191Z","shell.execute_reply":"2023-11-13T11:55:45.775610Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class EncoderLayer(tf.keras.layers.Layer):\n  def __init__(self,*, d_model, num_heads, dff, dropout_rate=0.1):\n    super().__init__()\n\n    self.self_attention = GlobalSelfAttention(\n        num_heads=num_heads,\n        key_dim=d_model,\n        dropout=dropout_rate)\n\n    self.ffn = FeedForward(d_model, dff)\n\n  def call(self, x):\n    x = self.self_attention(x)\n    x = self.ffn(x)\n    return x","metadata":{"execution":{"iopub.status.busy":"2023-11-13T11:55:45.777550Z","iopub.execute_input":"2023-11-13T11:55:45.777887Z","iopub.status.idle":"2023-11-13T11:55:45.791061Z","shell.execute_reply.started":"2023-11-13T11:55:45.777864Z","shell.execute_reply":"2023-11-13T11:55:45.790209Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RnnModel(tf.keras.Model):\n  def __init__(self, *, num_layers, d_model, num_heads,\n               dff, vocab_size, dropout_rate=0.1):\n    super().__init__()\n\n    self.d_model = d_model\n    self.num_layers = num_layers\n\n    self.pos_embedding = PositionalEmbedding(\n        vocab_size=vocab_size, d_model=d_model)\n\n    self.enc_layers = [\n        EncoderLayer(d_model=d_model,\n                     num_heads=num_heads,\n                     dff=dff,\n                     dropout_rate=dropout_rate)\n        for _ in range(num_layers)]\n    \n    self.dropout = tf.keras.layers.Dropout(dropout_rate)\n    self.layer1_100 = tf.keras.layers.Dense(400, activation='relu')\n    self.layer1_10 = tf.keras.layers.Dense(40,activation='relu')\n    self.out1_layer = tf.keras.layers.Dense(2)\n    self.gausi1 = tf.keras.layers.GaussianNoise(0.01)\n    self.gausi2 = tf.keras.layers.GaussianNoise(0.01)\n  def call(self, x):\n    #print(x.shape)\n    # `x` is token-IDs shape: (batch, seq_len)\n    x = self.pos_embedding(x)  # Shape `(batch_size, seq_len, d_model)`.\n    # Add dropout.\n    x = self.dropout(x)\n    for i in range(self.num_layers):\n        x = self.enc_layers[i](x)\n    x = self.gausi1(x)\n    x1 = self.layer1_100(x)\n    x1 = self.gausi2(x1)\n    x1 = self.layer1_10(x1)\n    o1 = self.out1_layer(x1) \n    return o1 ","metadata":{"execution":{"iopub.status.busy":"2023-11-13T11:55:45.791964Z","iopub.execute_input":"2023-11-13T11:55:45.792183Z","iopub.status.idle":"2023-11-13T11:55:45.803079Z","shell.execute_reply.started":"2023-11-13T11:55:45.792164Z","shell.execute_reply":"2023-11-13T11:55:45.802102Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def scheduler(epoch, lr):\n    if epoch < 1:\n        return lr\n    return lr * tf.math.exp(-0.1)\ndef 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-13T11:55:45.804461Z","iopub.execute_input":"2023-11-13T11:55:45.804743Z","iopub.status.idle":"2023-11-13T11:55:45.816766Z","shell.execute_reply.started":"2023-11-13T11:55:45.804719Z","shell.execute_reply":"2023-11-13T11:55:45.815967Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# train","metadata":{}},{"cell_type":"code","source":"steps_per_epoch = num_train//batch_size\nval_steps_per_epoch = num_val//val_batch_size\nN_EPOCHS = 300\nprint(steps_per_epoch)\nprint(val_steps_per_epoch)","metadata":{"execution":{"iopub.status.busy":"2023-11-13T11:55:45.817669Z","iopub.execute_input":"2023-11-13T11:55:45.817885Z","iopub.status.idle":"2023-11-13T11:55:45.828375Z","shell.execute_reply.started":"2023-11-13T11:55:45.817866Z","shell.execute_reply":"2023-11-13T11:55:45.827556Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import math\nN_WARMUP_EPOCHS = 0\nLR_MAX = 1.8978e-04\nWD_RATIO = 0.05\nWARMUP_METHOD = \"exp\"\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        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# 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# Learning Rate Callback\nlr_callback = tf.keras.callbacks.LearningRateScheduler(lambda step: LR_SCHEDULE[step], verbose=0)","metadata":{"execution":{"iopub.status.busy":"2023-11-13T11:55:45.829379Z","iopub.execute_input":"2023-11-13T11:55:45.829600Z","iopub.status.idle":"2023-11-13T11:55:45.839443Z","shell.execute_reply.started":"2023-11-13T11:55:45.829580Z","shell.execute_reply":"2023-11-13T11:55:45.838602Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"try:\n    tpu = tf.distribute.cluster_resolver.TPUClusterResolver.connect(tpu=\"local\") # \"local\" for 1VM TPU\n    print('Running on TPU ')#, tpu.cluster_spec().as_dict()['worker'])\nexcept ValueError:\n    tpu = None\nif tpu:\n    strategy = tf.distribute.TPUStrategy(tpu)\n    print(\"on TPU\")\n    print(\"REPLICAS: \", strategy.num_replicas_in_sync)\nelse:\n    strategy = tf.distribute.get_strategy()\n'''\nstrategy = tf.distribute.MirroredStrategy()\n'''\nmodel_list = []\nwith strategy.scope():\n    learning_rate = 0.0003#1e-4\n    epsilon = 1e-10\n    loss = loss_fn\n    optimizer = tf.keras.optimizers.Adam(learning_rate=learning_rate, epsilon=epsilon)\n\n    #model_T = RnnModel(num_layers=6,d_model=198,num_heads=8,dff=500,vocab_size=5)\n    model_T = joblib.load(\"/kaggle/input/stanford-rrf-tensorflow-tpu/models.pkl2\")[0]\n    model_T.compile(optimizer=optimizer, loss=loss)\n\n    callback1 = tf.keras.callbacks.EarlyStopping(monitor='loss', patience=5)\n    #callback2 = tf.keras.callbacks.LearningRateScheduler(scheduler)\n    model_T.fit(train_dataset,\n                validation_data=val_dataset,\n                epochs=200,#N_EPOCHS,\n                steps_per_epoch = steps_per_epoch,\n                validation_steps=val_steps_per_epoch,\n                verbose = 2,\n                callbacks=[callback1,lr_callback])\n    model_list.append(model_T)\njoblib.dump(model_list,\"models.pkl2\", compress=3)","metadata":{"execution":{"iopub.status.busy":"2023-11-13T11:55:45.840404Z","iopub.execute_input":"2023-11-13T11:55:45.840633Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}